@@ -208,21 +208,78 @@ def test_generate_calls_graph_capture_when_enabled():
208208 rollout ._generate_graph .assert_called_once ()
209209
210210
211- def test_generate_disables_eos_early_exit ():
211+ def test_generate_keeps_ranks_in_lockstep_and_pads_after_eos ():
212212 engine = _make_engine ()
213213 tok = _make_tokenizer ()
214214 rollout = HybridEngineRollout (engine , tok )
215- engine .module .generate .return_value = torch .tensor ([[1 , 2 , 3 , 4 ]])
215+ engine .module .generate .return_value = torch .tensor ([[10 , 11 , 5 , 2 , 7 , 8 ]])
216216
217217 req = MagicMock ()
218- req .prompt_ids = torch .tensor ([[1 , 2 ]])
218+ req .prompt_ids = torch .tensor ([[10 , 11 ]])
219219 req .prompt_attention_mask = torch .ones (1 , 2 , dtype = torch .long )
220220 sampling = MagicMock ()
221221 sampling .temperature = 0
222222 sampling .n_samples_per_prompt = 1
223- sampling .max_new_tokens = 2
223+ sampling .max_new_tokens = 4
224224 sampling .top_p = 1.0
225225
226- rollout .generate (req , sampling )
226+ result = rollout .generate (req , sampling )
227227
228228 assert engine .module .generate .call_args .kwargs ['eos_token_id' ] is None
229+ assert result .input_ids .tolist () == [[10 , 11 , 5 , 2 , 0 , 0 ]]
230+ assert result .attention_mask .tolist () == [[1 , 1 , 1 , 1 , 0 , 0 ]]
231+
232+
233+ def test_pad_after_eos_handles_different_lengths_and_missing_eos ():
234+ output_ids = torch .tensor ([
235+ [10 , 11 , 2 , 7 , 8 , 9 ],
236+ [10 , 11 , 5 , 6 , 2 , 9 ],
237+ [10 , 11 , 5 , 6 , 7 , 8 ],
238+ ])
239+
240+ padded , response_attn = HybridEngineRollout ._pad_after_eos (output_ids , 2 , eos_token_id = 2 , pad_token_id = 0 )
241+
242+ assert padded .tolist () == [
243+ [10 , 11 , 2 , 0 , 0 , 0 ],
244+ [10 , 11 , 5 , 6 , 2 , 0 ],
245+ [10 , 11 , 5 , 6 , 7 , 8 ],
246+ ]
247+ assert response_attn .tolist () == [
248+ [1 , 0 , 0 , 0 ],
249+ [1 , 1 , 1 , 0 ],
250+ [1 , 1 , 1 , 1 ],
251+ ]
252+
253+
254+ def test_pad_after_eos_keeps_eos_attended_when_eos_is_pad ():
255+ output_ids = torch .tensor ([[10 , 11 , 5 , 2 , 7 , 8 ]])
256+
257+ padded , response_attn = HybridEngineRollout ._pad_after_eos (output_ids , 2 , eos_token_id = 2 , pad_token_id = 2 )
258+
259+ assert padded .tolist () == [[10 , 11 , 5 , 2 , 2 , 2 ]]
260+ assert response_attn .tolist () == [[1 , 1 , 0 , 0 ]]
261+
262+
263+ def test_pad_after_eos_supports_multiple_eos_ids ():
264+ output_ids = torch .tensor ([[10 , 11 , 5 , 3 , 7 , 2 ]])
265+
266+ padded , response_attn = HybridEngineRollout ._pad_after_eos (output_ids , 2 , eos_token_id = [2 , 3 ], pad_token_id = 0 )
267+
268+ assert padded .tolist () == [[10 , 11 , 5 , 3 , 0 , 0 ]]
269+ assert response_attn .tolist () == [[1 , 1 , 0 , 0 ]]
270+
271+
272+ def test_generate_accepts_zero_pad_token_id ():
273+ engine = _make_engine ()
274+ tok = _make_tokenizer ()
275+ rollout = HybridEngineRollout (engine , tok )
276+ engine .module .generate .return_value = torch .tensor ([[10 , 11 , 5 , 6 ]])
277+
278+ req = MagicMock ()
279+ req .prompt_ids = torch .tensor ([[10 , 11 ]])
280+ req .prompt_attention_mask = torch .ones (1 , 2 , dtype = torch .long )
281+ sampling = MagicMock (temperature = 0 , n_samples_per_prompt = 1 , max_new_tokens = 2 , top_p = 1.0 )
282+
283+ rollout .generate (req , sampling )
284+
285+ assert engine .module .generate .call_args .kwargs ['pad_token_id' ] == 0
0 commit comments