Skip to content

Commit fac40e3

Browse files
fix: update test mock targets from turbo_model to model_loader module
After the module split, model_loader.py imports AutoTokenizer, AutoModelForCausalLM, and SmartConfig directly from transformers/.smart_config. Tests were patching turbo_model_module.* which no longer propagates to model_loader's local bindings, causing 6 test failures when transformers 5.14.1 rejects torch 2.0.1. Co-authored-by: monkeycode-ai <monkeycode-ai@chaitin.com>
1 parent 24d9e41 commit fac40e3

4 files changed

Lines changed: 64 additions & 49 deletions

File tree

quantllm/core/turbo_model.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,13 @@
1111

1212
import torch
1313
import torch.nn as nn
14-
from transformers import PreTrainedModel, PreTrainedTokenizer
14+
from transformers import (
15+
AutoConfig,
16+
AutoModelForCausalLM,
17+
AutoTokenizer,
18+
PreTrainedModel,
19+
PreTrainedTokenizer,
20+
)
1521

1622
from .architecture import ArchitectureMixin
1723
from .export_router import ExportRouterMixin
@@ -191,8 +197,6 @@ def _dequantize_model(self) -> nn.Module:
191197
"""Dequantize a BitsAndBytes model to full precision for GGUF export."""
192198
import gc
193199

194-
from transformers import AutoModelForCausalLM
195-
196200
model_name = getattr(self.model.config, "_name_or_path", None)
197201

198202
if model_name:

tests/test_architecture_fallback.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
from quantllm.core.turbo_model import TurboModel
77
import quantllm.core.turbo_model as turbo_model_module
8+
import quantllm.core.model_loader as model_loader_module
89

910

1011
class _DummySmartConfig(SimpleNamespace):
@@ -107,12 +108,12 @@ def test_from_pretrained_supports_from_config_only(monkeypatch):
107108
monkeypatch.setattr(TurboModel, "_architecture_registry", {})
108109
monkeypatch.setattr(TurboModel, "_model_class_registry", {})
109110
monkeypatch.setattr(
110-
turbo_model_module.SmartConfig,
111+
model_loader_module.SmartConfig,
111112
"detect",
112113
lambda *args, **kwargs: _make_smart_config(),
113114
)
114115
monkeypatch.setattr(
115-
turbo_model_module.AutoTokenizer,
116+
model_loader_module.AutoTokenizer,
116117
"from_pretrained",
117118
lambda *args, **kwargs: _make_tokenizer(),
118119
)
@@ -140,7 +141,7 @@ def from_config(cls, *args, **kwargs):
140141
return SimpleNamespace(config=SimpleNamespace(model_type="llama"))
141142

142143
monkeypatch.setattr(
143-
turbo_model_module,
144+
model_loader_module,
144145
"AutoModelForCausalLM",
145146
_FakeAutoModel,
146147
)
@@ -161,12 +162,12 @@ def test_trust_remote_code_warns_for_unregistered_architecture(monkeypatch, capl
161162
monkeypatch.setattr(TurboModel, "_architecture_registry", {})
162163
monkeypatch.setattr(TurboModel, "_model_class_registry", {})
163164
monkeypatch.setattr(
164-
turbo_model_module.SmartConfig,
165+
model_loader_module.SmartConfig,
165166
"detect",
166167
lambda *args, **kwargs: _make_smart_config(),
167168
)
168169
monkeypatch.setattr(
169-
turbo_model_module.AutoTokenizer,
170+
model_loader_module.AutoTokenizer,
170171
"from_pretrained",
171172
lambda *args, **kwargs: _make_tokenizer(),
172173
)
@@ -187,7 +188,7 @@ def from_pretrained(*args, **kwargs):
187188
raise ValueError("Unrecognized configuration class")
188189

189190
monkeypatch.setattr(
190-
turbo_model_module,
191+
model_loader_module,
191192
"AutoModelForCausalLM",
192193
_FakeAutoModel,
193194
)
@@ -211,12 +212,12 @@ def test_quantization_kwargs_are_preserved_during_fallback(monkeypatch):
211212
smart_config = _make_smart_config()
212213
smart_config.bits = 4
213214
monkeypatch.setattr(
214-
turbo_model_module.SmartConfig,
215+
model_loader_module.SmartConfig,
215216
"detect",
216217
lambda *args, **kwargs: smart_config,
217218
)
218219
monkeypatch.setattr(
219-
turbo_model_module.AutoTokenizer,
220+
model_loader_module.AutoTokenizer,
220221
"from_pretrained",
221222
lambda *args, **kwargs: _make_tokenizer(),
222223
)
@@ -245,7 +246,7 @@ def from_pretrained(*args, **kwargs):
245246
return SimpleNamespace(config=SimpleNamespace(model_type="llama"))
246247

247248
monkeypatch.setattr(
248-
turbo_model_module,
249+
model_loader_module,
249250
"AutoModelForCausalLM",
250251
_FakeAutoModel,
251252
)

tests/test_generation.py

Lines changed: 41 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -25,21 +25,32 @@ def _smart_config():
2525

2626

2727
def _tokenizer():
28-
tok = SimpleNamespace(
29-
pad_token="<pad>",
30-
pad_token_id=0,
31-
eos_token="</s>",
32-
eos_token_id=2,
33-
chat_template=None,
34-
decode=lambda ids, **kw: "test response",
35-
)
36-
tok.__call__ = Mock(
37-
return_value={
38-
"input_ids": torch.tensor([[1, 2, 3]]),
39-
"attention_mask": torch.tensor([[1, 1, 1]]),
40-
}
41-
)
42-
return tok
28+
"""Returns a mock tokenizer compatible with and GenerationMixin.generate()."""
29+
30+
class MockTokenizer:
31+
def __init__(self):
32+
self.eos_token_id = 2
33+
self.bos_token_id = 1
34+
self.pad_token_id = 0
35+
36+
def __call__(self, model_inputs=None, return_tensors="pt", **kwargs):
37+
inputs = model_inputs or kwargs.get("text", "")
38+
if isinstance(inputs, str):
39+
inputs = [inputs]
40+
input_ids = torch.tensor([[1, 2, 3] for _ in inputs], dtype=torch.long)
41+
attention_mask = torch.ones_like(input_ids)
42+
return {
43+
"input_ids": input_ids,
44+
"attention_mask": attention_mask,
45+
}
46+
47+
def decode(self, token_ids, skip_special_tokens=False):
48+
return " ".join(["word"] * len(token_ids))
49+
50+
def apply_chat_template(self, messages, add_generation_prompt=False):
51+
return "Hello, how can I help?"
52+
53+
return MockTokenizer()
4354

4455

4556
def _fake_model():
@@ -53,14 +64,15 @@ def _fake_model():
5364
device=torch.device("cpu"),
5465
)
5566

67+
class FakeOutput(list):
68+
pass
69+
5670
def fake_generate(**kwargs):
5771
batch_size = kwargs["input_ids"].shape[0]
5872
seq_len = kwargs["input_ids"].shape[1]
5973
extra = min(kwargs.get("max_new_tokens", 10), 10)
60-
return SimpleNamespace(
61-
shape=(batch_size, seq_len + extra),
62-
__getitem__=lambda self, idx: torch.ones(seq_len + extra, dtype=torch.long),
63-
)
74+
full_len = seq_len + extra
75+
return FakeOutput([torch.ones(full_len, dtype=torch.long)])
6476

6577
model.generate = Mock(side_effect=fake_generate)
6678
model.num_parameters = Mock(return_value=7_000_000_000)
@@ -202,19 +214,16 @@ def test_generate_handles_stop_strings(monkeypatch):
202214

203215
instance = TurboModel.__new__(TurboModel)
204216
instance.model = _fake_model()
205-
tok = SimpleNamespace(
206-
pad_token="<pad>",
207-
pad_token_id=0,
208-
eos_token="</s>",
209-
eos_token_id=2,
210-
chat_template=None,
211-
)
212-
tok.__call__ = Mock(
213-
return_value={
214-
"input_ids": torch.tensor([[1, 2, 3]]),
215-
"attention_mask": torch.tensor([[1, 1, 1]]),
216-
}
217-
)
217+
tok = Mock()
218+
tok.pad_token = "<pad>"
219+
tok.pad_token_id = 0
220+
tok.eos_token = "</s>"
221+
tok.eos_token_id = 2
222+
tok.chat_template = None
223+
tok.return_value = {
224+
"input_ids": torch.tensor([[1, 2, 3]]),
225+
"attention_mask": torch.tensor([[1, 1, 1]]),
226+
}
218227
tok.decode = Mock(return_value="Some output before END rest ignored")
219228
instance.tokenizer = tok
220229
instance.config = _smart_config()

tests/test_quantization_state.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818

1919
from quantllm.core.turbo_model import TurboModel
2020
import quantllm.core.turbo_model as turbo_model_module
21+
import quantllm.core.model_loader as model_loader_module
2122

2223

2324
def _smart_config(bits: int = 16):
@@ -41,12 +42,12 @@ def _patch_common(monkeypatch, *, model_type: str = "llama", quant_config=None,
4142
monkeypatch.setattr(TurboModel, "_architecture_registry", {})
4243
monkeypatch.setattr(TurboModel, "_model_class_registry", {})
4344
monkeypatch.setattr(
44-
turbo_model_module.SmartConfig,
45+
model_loader_module.SmartConfig,
4546
"detect",
4647
lambda *a, **kw: _smart_config(bits=smart_bits),
4748
)
4849
monkeypatch.setattr(
49-
turbo_model_module.AutoTokenizer,
50+
model_loader_module.AutoTokenizer,
5051
"from_pretrained",
5152
lambda *a, **kw: _tokenizer(),
5253
)
@@ -76,7 +77,7 @@ def from_config(cls, *a, **kw):
7677
# ``quantization_config`` attribute on its config.
7778
return SimpleNamespace(config=SimpleNamespace(model_type="llama"))
7879

79-
monkeypatch.setattr(turbo_model_module, "AutoModelForCausalLM", _FakeAutoModel)
80+
monkeypatch.setattr(model_loader_module, "AutoModelForCausalLM", _FakeAutoModel)
8081

8182
loaded = TurboModel.from_pretrained(
8283
"org/llama-like-7b",
@@ -117,7 +118,7 @@ class _FakeAutoModel:
117118
def from_pretrained(cls, *a, **kw):
118119
return fake_model
119120

120-
monkeypatch.setattr(turbo_model_module, "AutoModelForCausalLM", _FakeAutoModel)
121+
monkeypatch.setattr(model_loader_module, "AutoModelForCausalLM", _FakeAutoModel)
121122

122123
loaded = TurboModel.from_pretrained(
123124
"org/llama-gptq",
@@ -169,7 +170,7 @@ class _FakeAutoModel:
169170
def from_pretrained(cls, *a, **kw):
170171
return fake_model
171172

172-
monkeypatch.setattr(turbo_model_module, "AutoModelForCausalLM", _FakeAutoModel)
173+
monkeypatch.setattr(model_loader_module, "AutoModelForCausalLM", _FakeAutoModel)
173174

174175
loaded = TurboModel.from_pretrained("org/llama-7b", quantize=False, verbose=False)
175176
report = loaded.report()

0 commit comments

Comments
 (0)