fix: correct online rollout lifecycle
This commit is contained in:
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
- rewards: [G]
|
||||
|
||||
Output:
|
||||
- prompts: [B, P_max]
|
||||
- prompts: [B, P_max], left-padded
|
||||
- prompt_mask: [B, P_max]
|
||||
- responses: [B, G, R_max]
|
||||
- masks: [B, G, R_max]
|
||||
- rewards: [B, G]
|
||||
@@ -201,13 +202,15 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||
|
||||
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
|
||||
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
p_len = b["prompts"].size(0)
|
||||
prompts[i, :p_len] = b["prompts"]
|
||||
prompts[i, -p_len:] = b["prompts"]
|
||||
prompt_mask[i, -p_len:] = True
|
||||
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||
for g in range(min(G, len(b["responses"]))):
|
||||
r_len = b["responses"][g].size(0)
|
||||
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"prompt_mask": prompt_mask,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
|
||||
+88
-12
@@ -35,6 +35,7 @@ class RawRollout:
|
||||
|
||||
Fields:
|
||||
prompts: Tokenized prompts, shape ``[B, P_len]``.
|
||||
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
|
||||
responses: Generated response token IDs, shape ``[B, G, R_max]``.
|
||||
response_mask: Boolean mask for real (non-pad) response tokens,
|
||||
shape ``[B, G, R_max]``.
|
||||
@@ -47,6 +48,7 @@ class RawRollout:
|
||||
"""
|
||||
|
||||
prompts: Tensor
|
||||
prompt_mask: Tensor
|
||||
responses: Tensor
|
||||
response_mask: Tensor
|
||||
logprobs_old: Tensor
|
||||
@@ -143,6 +145,15 @@ class RolloutGenerator:
|
||||
``add_generation_prompt=True`` so rollout prompts match the
|
||||
format the policy was SFT-trained on.
|
||||
"""
|
||||
model = self.scheduler._executor.model
|
||||
was_training = model.training
|
||||
model.eval()
|
||||
try:
|
||||
return self._generate_eval(batch)
|
||||
finally:
|
||||
model.train(was_training)
|
||||
|
||||
def _generate_eval(self, batch: Dict) -> RawRollout:
|
||||
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
|
||||
B = len(prompt_texts)
|
||||
G = self.group_size
|
||||
@@ -161,6 +172,15 @@ class RolloutGenerator:
|
||||
rep_window=self.rep_window,
|
||||
return_logprobs=True,
|
||||
)
|
||||
if len(results) != B * G:
|
||||
raise RuntimeError(
|
||||
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
|
||||
)
|
||||
for token_ids, logprobs in results:
|
||||
if len(token_ids) != len(logprobs):
|
||||
raise RuntimeError(
|
||||
"Rollout scheduler returned misaligned token IDs and logprobs"
|
||||
)
|
||||
|
||||
# Each element is (token_ids, logprobs); pad to max length.
|
||||
max_len = 0
|
||||
@@ -171,10 +191,12 @@ class RolloutGenerator:
|
||||
device = self.scheduler.device
|
||||
P_len = max(len(ids) for ids in flat_prompt_ids)
|
||||
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
|
||||
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
|
||||
for i, ids in enumerate(flat_prompt_ids):
|
||||
prompts_tensor[i, : len(ids)] = torch.tensor(
|
||||
prompts_tensor[i, -len(ids) :] = torch.tensor(
|
||||
ids, dtype=torch.long, device=device
|
||||
)
|
||||
prompt_mask[i, -len(ids) :] = True
|
||||
|
||||
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
||||
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
||||
@@ -201,6 +223,7 @@ class RolloutGenerator:
|
||||
|
||||
return RawRollout(
|
||||
prompts=prompts_tensor,
|
||||
prompt_mask=prompt_mask,
|
||||
responses=responses,
|
||||
response_mask=response_mask,
|
||||
logprobs_old=logprobs_old,
|
||||
@@ -239,17 +262,35 @@ class RolloutGenerator:
|
||||
f"{list(batch.keys())}"
|
||||
)
|
||||
|
||||
prompt_texts: List[str] = []
|
||||
flat_prompt_ids: List[List[int]] = []
|
||||
for messages in messages_list:
|
||||
text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
try:
|
||||
prompt_texts = self.tokenizer.apply_chat_template(
|
||||
messages_list, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
ids = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
prompt_texts.append(text)
|
||||
flat_prompt_ids.append(list(ids))
|
||||
if (
|
||||
not isinstance(prompt_texts, list)
|
||||
or len(prompt_texts) != len(messages_list)
|
||||
or not all(isinstance(text, str) for text in prompt_texts)
|
||||
):
|
||||
raise TypeError("Tokenizer does not support batched chat templates")
|
||||
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
|
||||
if len(flat_prompt_ids) != len(messages_list) or not all(
|
||||
isinstance(ids, list) for ids in flat_prompt_ids
|
||||
):
|
||||
raise TypeError("Tokenizer does not support batched encoding")
|
||||
except (TypeError, IndexError, KeyError):
|
||||
# Keep compatibility with lightweight tokenizer adapters that only
|
||||
# implement the single-conversation template API.
|
||||
prompt_texts = []
|
||||
flat_prompt_ids = []
|
||||
for messages in messages_list:
|
||||
text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
ids = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
prompt_texts.append(text)
|
||||
flat_prompt_ids.append(list(ids))
|
||||
return prompt_texts, flat_prompt_ids
|
||||
|
||||
@staticmethod
|
||||
@@ -308,6 +349,7 @@ class RolloutRunner:
|
||||
self.rollout_interval = rollout_interval
|
||||
|
||||
self._cache: Optional[RolloutResult] = None
|
||||
self._cache_key = None
|
||||
self._steps_since_rollout: int = 0
|
||||
|
||||
def step(self):
|
||||
@@ -317,12 +359,40 @@ class RolloutRunner:
|
||||
def clear_cache(self):
|
||||
"""Force next call to re-run rollout."""
|
||||
self._cache = None
|
||||
self._cache_key = None
|
||||
|
||||
@staticmethod
|
||||
def _batch_key(batch: Dict):
|
||||
"""Build a stable key for the prompt fields accepted by the generator."""
|
||||
|
||||
def freeze(value):
|
||||
if isinstance(value, dict):
|
||||
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
|
||||
if isinstance(value, (list, tuple)):
|
||||
return tuple(freeze(item) for item in value)
|
||||
return value
|
||||
|
||||
fields = ("messages", "instruction", "input", "output")
|
||||
return tuple(
|
||||
(field, freeze(batch[field])) for field in fields if field in batch
|
||||
)
|
||||
|
||||
def _score(self, raw: RawRollout) -> RolloutResult:
|
||||
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
|
||||
if not isinstance(rewards, Tensor):
|
||||
rewards = torch.as_tensor(rewards, dtype=torch.float32)
|
||||
expected_shape = raw.responses.shape[:2]
|
||||
if rewards.shape != expected_shape:
|
||||
raise ValueError(
|
||||
f"Reward model returned shape {tuple(rewards.shape)}, "
|
||||
f"expected {tuple(expected_shape)}"
|
||||
)
|
||||
if not torch.isfinite(rewards).all():
|
||||
raise ValueError("Reward model returned non-finite values")
|
||||
device = raw.prompts.device
|
||||
return RolloutResult(
|
||||
prompts=raw.prompts,
|
||||
prompt_mask=raw.prompt_mask,
|
||||
responses=raw.responses,
|
||||
response_mask=raw.response_mask,
|
||||
rewards=rewards.to(device=device),
|
||||
@@ -337,9 +407,15 @@ class RolloutRunner:
|
||||
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
|
||||
or when the cache is empty.
|
||||
"""
|
||||
if self._cache is None or self._steps_since_rollout >= self.rollout_interval:
|
||||
cache_key = self._batch_key(batch)
|
||||
if (
|
||||
self._cache is None
|
||||
or cache_key != self._cache_key
|
||||
or self._steps_since_rollout >= self.rollout_interval
|
||||
):
|
||||
raw = self.generator.generate(batch)
|
||||
self._cache = self._score(raw)
|
||||
self._cache_key = cache_key
|
||||
self._steps_since_rollout = 0
|
||||
return self._cache, True
|
||||
return self._cache, False
|
||||
|
||||
@@ -158,6 +158,11 @@ class BaseStrategy(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def on_optimizer_step(self):
|
||||
"""Advance online rollout state after a successful optimizer step."""
|
||||
if self._rollout_runner is not None:
|
||||
self._rollout_runner.step()
|
||||
|
||||
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
"""Run offline or online forward depending on runner injection."""
|
||||
if self._rollout_runner is None:
|
||||
@@ -166,8 +171,6 @@ class BaseStrategy(ABC):
|
||||
result, is_fresh = self._rollout_runner(batch)
|
||||
if is_fresh:
|
||||
self._on_rollout_refresh()
|
||||
if self.executor and self.executor.sync_gradients:
|
||||
self._rollout_runner.step()
|
||||
|
||||
train_batch = self.prepare_from_rollout(result)
|
||||
return self.compute_loss(train_batch)
|
||||
@@ -411,6 +414,12 @@ class GRPOStrategy(BaseStrategy):
|
||||
responses_flat = responses.view(-1, response_len)
|
||||
masks_flat = masks.view(-1, response_len)
|
||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||
prompt_mask = batch.get("prompt_mask")
|
||||
if prompt_mask is None:
|
||||
prompt_mask = prompts.ne(0)
|
||||
prompt_mask_expanded = (
|
||||
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
|
||||
)
|
||||
prompt_len = prompt_expanded.size(1)
|
||||
|
||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||
@@ -423,7 +432,9 @@ class GRPOStrategy(BaseStrategy):
|
||||
)
|
||||
|
||||
# Build full attention mask: key-padding + causal
|
||||
key_pad = full_sequences.bool()[:, None, None, :]
|
||||
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
|
||||
:, None, None, :
|
||||
]
|
||||
S = key_pad.shape[-1]
|
||||
causal = torch.tril(
|
||||
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
|
||||
@@ -485,6 +496,7 @@ class GRPOStrategy(BaseStrategy):
|
||||
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||
return {
|
||||
"prompts": result.prompts,
|
||||
"prompt_mask": result.prompt_mask,
|
||||
"responses": result.responses,
|
||||
"masks": result.response_mask,
|
||||
"rewards": result.rewards,
|
||||
|
||||
@@ -83,6 +83,7 @@ class Trainer:
|
||||
if executor.sync_gradients:
|
||||
self._call_callbacks("on_optimizer_step", context)
|
||||
context.optimizer.step()
|
||||
context.strategy.on_optimizer_step()
|
||||
context.optimizer.zero_grad()
|
||||
|
||||
if context.scheduler:
|
||||
|
||||
@@ -937,8 +937,9 @@ def test_grpo_collate_variable_lengths():
|
||||
assert result["masks"].shape == (2, 2, 4)
|
||||
assert result["rewards"].shape == (2, 2)
|
||||
|
||||
# Check padding: item 1 prompt is length 2, padded to 3
|
||||
assert result["prompts"][1, 2] == 0
|
||||
# Prompts are left-padded so each response follows its real prompt tokens.
|
||||
assert torch.equal(result["prompts"][1], torch.tensor([0, 10, 11]))
|
||||
assert torch.equal(result["prompt_mask"][1], torch.tensor([False, True, True]))
|
||||
|
||||
# Check response content: item 0, response 0 is [4,5] padded to 4
|
||||
assert result["responses"][0, 0, 0] == 4
|
||||
|
||||
@@ -63,6 +63,7 @@ def _make_frozen(model, device):
|
||||
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
|
||||
return RolloutResult(
|
||||
prompts=torch.randint(3, 200, (B, P), device=device),
|
||||
prompt_mask=torch.ones(B, P, dtype=torch.bool, device=device),
|
||||
responses=torch.randint(3, 200, (B, G, R), device=device),
|
||||
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
|
||||
rewards=torch.randn(B, G, device=device),
|
||||
@@ -177,6 +178,7 @@ def test_grpo_prepare_from_rollout_mapping(device):
|
||||
r = _make_rollout_result(device=device)
|
||||
batch = strat.prepare_from_rollout(r)
|
||||
assert batch["prompts"] is r.prompts
|
||||
assert batch["prompt_mask"] is r.prompt_mask
|
||||
assert batch["responses"] is r.responses
|
||||
assert batch["masks"] is r.response_mask
|
||||
assert batch["rewards"] is r.rewards
|
||||
@@ -255,9 +257,11 @@ def test_grpo_no_resync_when_same_cached_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 1
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
def test_grpo_resync_when_new_rollout_result(device):
|
||||
@@ -265,8 +269,10 @@ def test_grpo_resync_when_new_rollout_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.calls == 2
|
||||
assert runner.step_calls == 2
|
||||
|
||||
@@ -281,8 +287,10 @@ def test_dpo_no_sync_hook_when_new_rollout_result(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
runner.swap_result(_make_rollout_result(device=device))
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 2
|
||||
|
||||
|
||||
@@ -301,6 +309,7 @@ def test_step_called_when_sync_gradients_true(device):
|
||||
runner = _RecordingRunner(_make_rollout_result(device=device))
|
||||
strat.set_rollout_runner(runner)
|
||||
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
||||
strat.on_optimizer_step()
|
||||
assert runner.step_calls == 1
|
||||
|
||||
|
||||
|
||||
@@ -71,6 +71,18 @@ class ConstantRewardModel(BaseRewardModel):
|
||||
return torch.full((B, G), float(self.value))
|
||||
|
||||
|
||||
class BadShapeRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
return torch.zeros(len(prompts))
|
||||
|
||||
|
||||
class NonFiniteRewardModel(BaseRewardModel):
|
||||
def score(self, prompts, responses):
|
||||
B = len(prompts)
|
||||
G = len(responses[0]) if B else 0
|
||||
return torch.full((B, G), float("nan"))
|
||||
|
||||
|
||||
def _make_config(vocab_size=200, max_position_embeddings=128):
|
||||
return AutoRegressiveLMConfig(
|
||||
vocab_size=vocab_size,
|
||||
@@ -111,6 +123,7 @@ def _make_instruction_batch(n=2):
|
||||
def test_raw_rollout_fields():
|
||||
r = RawRollout(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
@@ -124,6 +137,7 @@ def test_raw_rollout_fields():
|
||||
def test_rollout_result_inherits_raw_rollout_fields():
|
||||
r = RolloutResult(
|
||||
prompts=torch.zeros(2, 4, dtype=torch.long),
|
||||
prompt_mask=torch.ones(2, 4, dtype=torch.bool),
|
||||
responses=torch.zeros(2, 3, 5, dtype=torch.long),
|
||||
response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
|
||||
logprobs_old=torch.zeros(2, 3, 5),
|
||||
@@ -132,6 +146,7 @@ def test_rollout_result_inherits_raw_rollout_fields():
|
||||
assert r.rewards.shape == (2, 3)
|
||||
assert r.prompts.shape == (2, 4)
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == (2, 4)
|
||||
|
||||
|
||||
def test_base_reward_model_is_abstract():
|
||||
@@ -179,11 +194,28 @@ def test_rollout_generator_shapes(device):
|
||||
assert r.responses.shape == (2, 3, 5)
|
||||
assert r.response_mask.shape == (2, 3, 5)
|
||||
assert r.logprobs_old.shape == (2, 3, 5)
|
||||
assert r.prompt_mask.shape == r.prompts.shape
|
||||
assert len(r.prompt_texts) == 2
|
||||
assert len(r.response_texts) == 2
|
||||
assert len(r.response_texts[0]) == 3
|
||||
|
||||
|
||||
def test_rollout_generator_uses_eval_and_restores_mode(device):
|
||||
gen, model = _make_generator(device, group_size=1, max_tokens=2)
|
||||
model.train()
|
||||
seen_training = []
|
||||
original = gen.scheduler.run_batch
|
||||
|
||||
def recording_run_batch(*args, **kwargs):
|
||||
seen_training.append(model.training)
|
||||
return original(*args, **kwargs)
|
||||
|
||||
gen.scheduler.run_batch = recording_run_batch
|
||||
gen.generate(_make_instruction_batch(n=1))
|
||||
assert seen_training == [False]
|
||||
assert model.training is True
|
||||
|
||||
|
||||
def test_rollout_generator_mask_matches_responses(device):
|
||||
"""Positions beyond a response's length are pad (mask False)."""
|
||||
gen, _ = _make_generator(device, group_size=2, max_tokens=6)
|
||||
@@ -291,6 +323,24 @@ def test_rollout_runner_cache_returns_stale_flag(device):
|
||||
assert fresh2 is False
|
||||
|
||||
|
||||
def test_rollout_runner_refreshes_for_different_batch(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=100)
|
||||
r1, fresh1 = runner(_make_instruction_batch(n=1))
|
||||
batch2 = {"instruction": ["Different prompt"], "input": [""]}
|
||||
r2, fresh2 = runner(batch2)
|
||||
assert fresh1 is True
|
||||
assert fresh2 is True
|
||||
assert r2 is not r1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("reward_model", [BadShapeRewardModel, NonFiniteRewardModel])
|
||||
def test_rollout_runner_rejects_invalid_rewards(device, reward_model):
|
||||
generator, _ = _make_generator(device, group_size=2, max_tokens=2)
|
||||
runner = RolloutRunner(generator, reward_model(), rollout_interval=1)
|
||||
with pytest.raises(ValueError):
|
||||
runner(_make_instruction_batch(n=1))
|
||||
|
||||
|
||||
def test_rollout_runner_step_triggers_new_rollout(device):
|
||||
runner, _ = _make_runner(device, rollout_interval=2)
|
||||
batch = _make_instruction_batch()
|
||||
|
||||
Reference in New Issue
Block a user