Skip to content

Commit 9637a65

Browse files
committed
Restore EOS padding in synchronized OPSD rollout
Signed-off-by: LiRunGuo <li19107254665@gmail.com>
1 parent f7695ac commit 9637a65

2 files changed

Lines changed: 99 additions & 5 deletions

File tree

deepspeed/runtime/rollout/hybrid_engine_rollout.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,8 @@ def generate(self, request: RolloutRequest, sampling: SamplingConfig) -> Rollout
5555
pad_token_id = self.tokenizer.pad_token_id
5656
if pad_token_id is None:
5757
pad_token_id = self.tokenizer.eos_token_id
58+
if pad_token_id is None:
59+
raise ValueError("The tokenizer must define pad_token_id or eos_token_id")
5860

5961
module = self.engine.module
6062

@@ -99,10 +101,22 @@ def generate(self, request: RolloutRequest, sampling: SamplingConfig) -> Rollout
99101
accelerator.synchronize()
100102
generation_end = time.perf_counter()
101103

104+
# Generation deliberately ignores EOS above so ZeRO-3 ranks execute
105+
# the same number of parameter-gather collectives. Restore the usual
106+
# generation semantics before returning: retain the first EOS in each
107+
# response and replace every later token with padding.
108+
output_ids, response_attn = self._pad_after_eos(
109+
output_ids,
110+
response_start=prompt_len,
111+
eos_token_id=self.tokenizer.eos_token_id,
112+
pad_token_id=pad_token_id,
113+
)
114+
102115
# Build attention mask: pad positions (both left padding from prompt
103116
# and right padding from EOS / shorter sequences) are 0.
104117
response_start = prompt_len
105118
attention_mask = (output_ids != pad_token_id).long()
119+
attention_mask[:, response_start:] = response_attn
106120
for i in range(total):
107121
prompt_valid = request.prompt_attention_mask[i // n if B > 1 else 0]
108122
attention_mask[i, :prompt_len] = prompt_valid
@@ -144,6 +158,29 @@ def get_last_profile(self):
144158
"""Return the most recent profiling snapshot for this rollout instance."""
145159
return self._last_profile
146160

161+
@staticmethod
162+
def _pad_after_eos(output_ids, response_start, eos_token_id, pad_token_id):
163+
"""Retain the first response EOS and pad every subsequent position."""
164+
response_ids = output_ids[:, response_start:]
165+
response_attn = (response_ids != pad_token_id)
166+
167+
if eos_token_id is None or response_ids.shape[1] == 0:
168+
return output_ids, response_attn.long()
169+
170+
eos_ids = torch.as_tensor(eos_token_id, device=response_ids.device, dtype=response_ids.dtype).flatten()
171+
is_eos = (response_ids.unsqueeze(-1) == eos_ids).any(dim=-1)
172+
has_eos = is_eos.any(dim=-1)
173+
first_eos_idx = is_eos.long().argmax(dim=-1)
174+
positions = torch.arange(response_ids.shape[1], device=response_ids.device).unsqueeze(0)
175+
after_first_eos = has_eos.unsqueeze(1) & (positions > first_eos_idx.unsqueeze(1))
176+
first_eos = has_eos.unsqueeze(1) & (positions == first_eos_idx.unsqueeze(1))
177+
178+
output_ids = output_ids.clone()
179+
output_ids[:, response_start:].masked_fill_(after_first_eos, pad_token_id)
180+
# EOS is a valid generated token even when pad_token_id == eos_token_id.
181+
response_attn = ((response_ids != pad_token_id) | first_eos) & ~after_first_eos
182+
return output_ids, response_attn.long()
183+
147184
# ------------------------------------------------------------------
148185
# Graph capture decode loop (greedy only)
149186
# ------------------------------------------------------------------

tests/unit/runtime/rollout/test_hybrid_engine_rollout.py

Lines changed: 62 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)