@@ -25,21 +25,32 @@ def _smart_config():
2525
2626
2727def _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
4556def _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 ()
0 commit comments