diff --git a/deepspeed/runtime/rollout/hybrid_engine_rollout.py b/deepspeed/runtime/rollout/hybrid_engine_rollout.py index 1249e0f68d43..0db4965f5077 100644 --- a/deepspeed/runtime/rollout/hybrid_engine_rollout.py +++ b/deepspeed/runtime/rollout/hybrid_engine_rollout.py @@ -57,6 +57,8 @@ def generate(self, request: RolloutRequest, sampling: SamplingConfig) -> Rollout pad_token_id = self.tokenizer.pad_token_id if pad_token_id is None: pad_token_id = self.tokenizer.eos_token_id + if pad_token_id is None: + raise ValueError("The tokenizer must define pad_token_id or eos_token_id") module = self.engine.module @@ -96,6 +98,9 @@ def generate(self, request: RolloutRequest, sampling: SamplingConfig) -> Rollout prompt_ids, attention_mask=prompt_attn, max_new_tokens=max_new_tokens, + # ZeRO-3 gathers parameters during each decode forward, so every + # data-parallel rank must execute the same number of iterations. + eos_token_id=None, do_sample=do_sample, temperature=temperature if do_sample else 1.0, top_p=sampling.top_p if do_sample else 1.0, @@ -109,10 +114,22 @@ def generate(self, request: RolloutRequest, sampling: SamplingConfig) -> Rollout accelerator.synchronize() generation_end = time.perf_counter() + # Generation deliberately ignores EOS above so ZeRO-3 ranks execute + # the same number of parameter-gather collectives. Restore the usual + # generation semantics before returning: retain the first EOS in each + # response and replace every later token with padding. + output_ids, response_attn = self._pad_after_eos( + output_ids, + response_start=prompt_len, + eos_token_id=self.tokenizer.eos_token_id, + pad_token_id=pad_token_id, + ) + # Build attention mask: pad positions (both left padding from prompt # and right padding from EOS / shorter sequences) are 0. response_start = prompt_len attention_mask = (output_ids != pad_token_id).long() + attention_mask[:, response_start:] = response_attn for i in range(total): prompt_valid = request.prompt_attention_mask[i // n if B > 1 else 0] attention_mask[i, :prompt_len] = prompt_valid @@ -191,6 +208,29 @@ def expand_prompt_output(_module, _args, _kwargs, output): post_handle = module.register_forward_hook(expand_prompt_output, with_kwargs=True) return pre_handle, post_handle + @staticmethod + def _pad_after_eos(output_ids, response_start, eos_token_id, pad_token_id): + """Retain the first response EOS and pad every subsequent position.""" + response_ids = output_ids[:, response_start:] + response_attn = (response_ids != pad_token_id) + + if eos_token_id is None or response_ids.shape[1] == 0: + return output_ids, response_attn.long() + + eos_ids = torch.as_tensor(eos_token_id, device=response_ids.device, dtype=response_ids.dtype).flatten() + is_eos = (response_ids.unsqueeze(-1) == eos_ids).any(dim=-1) + has_eos = is_eos.any(dim=-1) + first_eos_idx = is_eos.long().argmax(dim=-1) + positions = torch.arange(response_ids.shape[1], device=response_ids.device).unsqueeze(0) + after_first_eos = has_eos.unsqueeze(1) & (positions > first_eos_idx.unsqueeze(1)) + first_eos = has_eos.unsqueeze(1) & (positions == first_eos_idx.unsqueeze(1)) + + output_ids = output_ids.clone() + output_ids[:, response_start:].masked_fill_(after_first_eos, pad_token_id) + # EOS is a valid generated token even when pad_token_id == eos_token_id. + response_attn = ((response_ids != pad_token_id) | first_eos) & ~after_first_eos + return output_ids, response_attn.long() + # ------------------------------------------------------------------ # Graph capture decode loop (greedy only) # ------------------------------------------------------------------ diff --git a/tests/unit/runtime/rollout/test_hybrid_engine_rollout.py b/tests/unit/runtime/rollout/test_hybrid_engine_rollout.py index 9f34b3e134e4..475c73f5854d 100644 --- a/tests/unit/runtime/rollout/test_hybrid_engine_rollout.py +++ b/tests/unit/runtime/rollout/test_hybrid_engine_rollout.py @@ -414,3 +414,80 @@ def test_generate_calls_graph_capture_when_enabled(): rollout.generate(req, sampling) rollout._generate_graph.assert_called_once() + + +def test_generate_keeps_ranks_in_lockstep_and_pads_after_eos(): + engine = _make_engine() + tok = _make_tokenizer() + rollout = HybridEngineRollout(engine, tok) + engine.module.generate.return_value = torch.tensor([[10, 11, 5, 2, 7, 8]]) + + req = MagicMock() + req.prompt_ids = torch.tensor([[10, 11]]) + req.prompt_attention_mask = torch.ones(1, 2, dtype=torch.long) + sampling = MagicMock() + sampling.temperature = 0 + sampling.n_samples_per_prompt = 1 + sampling.max_new_tokens = 4 + sampling.top_p = 1.0 + + result = rollout.generate(req, sampling) + + assert engine.module.generate.call_args.kwargs['eos_token_id'] is None + assert result.input_ids.tolist() == [[10, 11, 5, 2, 0, 0]] + assert result.attention_mask.tolist() == [[1, 1, 1, 1, 0, 0]] + + +def test_pad_after_eos_handles_different_lengths_and_missing_eos(): + output_ids = torch.tensor([ + [10, 11, 2, 7, 8, 9], + [10, 11, 5, 6, 2, 9], + [10, 11, 5, 6, 7, 8], + ]) + + padded, response_attn = HybridEngineRollout._pad_after_eos(output_ids, 2, eos_token_id=2, pad_token_id=0) + + assert padded.tolist() == [ + [10, 11, 2, 0, 0, 0], + [10, 11, 5, 6, 2, 0], + [10, 11, 5, 6, 7, 8], + ] + assert response_attn.tolist() == [ + [1, 0, 0, 0], + [1, 1, 1, 0], + [1, 1, 1, 1], + ] + + +def test_pad_after_eos_keeps_eos_attended_when_eos_is_pad(): + output_ids = torch.tensor([[10, 11, 5, 2, 7, 8]]) + + padded, response_attn = HybridEngineRollout._pad_after_eos(output_ids, 2, eos_token_id=2, pad_token_id=2) + + assert padded.tolist() == [[10, 11, 5, 2, 2, 2]] + assert response_attn.tolist() == [[1, 1, 0, 0]] + + +def test_pad_after_eos_supports_multiple_eos_ids(): + output_ids = torch.tensor([[10, 11, 5, 3, 7, 2]]) + + padded, response_attn = HybridEngineRollout._pad_after_eos(output_ids, 2, eos_token_id=[2, 3], pad_token_id=0) + + assert padded.tolist() == [[10, 11, 5, 3, 0, 0]] + assert response_attn.tolist() == [[1, 1, 0, 0]] + + +def test_generate_accepts_zero_pad_token_id(): + engine = _make_engine() + tok = _make_tokenizer() + rollout = HybridEngineRollout(engine, tok) + engine.module.generate.return_value = torch.tensor([[10, 11, 5, 6]]) + + req = MagicMock() + req.prompt_ids = torch.tensor([[10, 11]]) + req.prompt_attention_mask = torch.ones(1, 2, dtype=torch.long) + sampling = MagicMock(temperature=0, n_samples_per_prompt=1, max_new_tokens=2, top_p=1.0) + + rollout.generate(req, sampling) + + assert engine.module.generate.call_args.kwargs['pad_token_id'] == 0