From 8ab5631446ee99da706729925259ff5727dd22d7 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Thu, 23 Jul 2026 19:01:37 +0800 Subject: [PATCH] fix: correct online rollout lifecycle --- astrai/dataset/dataset.py | 8 ++- astrai/trainer/rollout.py | 100 ++++++++++++++++++++++---- astrai/trainer/strategy.py | 18 ++++- astrai/trainer/trainer.py | 1 + tests/data/test_dataset.py | 5 +- tests/trainer/test_online_strategy.py | 11 ++- tests/trainer/test_rollout.py | 50 +++++++++++++ 7 files changed, 173 insertions(+), 20 deletions(-) diff --git a/astrai/dataset/dataset.py b/astrai/dataset/dataset.py index 03ebec7..2dcb1d4 100644 --- a/astrai/dataset/dataset.py +++ b/astrai/dataset/dataset.py @@ -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, diff --git a/astrai/trainer/rollout.py b/astrai/trainer/rollout.py index 713e22d..e1c6621 100644 --- a/astrai/trainer/rollout.py +++ b/astrai/trainer/rollout.py @@ -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 diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index 0c1ad9b..fdb8347 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -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, diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index 2e1dcea..3f908cd 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -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: diff --git a/tests/data/test_dataset.py b/tests/data/test_dataset.py index bacab09..93899b6 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -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 diff --git a/tests/trainer/test_online_strategy.py b/tests/trainer/test_online_strategy.py index f299a56..34c88d1 100644 --- a/tests/trainer/test_online_strategy.py +++ b/tests/trainer/test_online_strategy.py @@ -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 diff --git a/tests/trainer/test_rollout.py b/tests/trainer/test_rollout.py index 3c73df4..4adf0b9 100644 --- a/tests/trainer/test_rollout.py +++ b/tests/trainer/test_rollout.py @@ -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()