fix: correct online rollout lifecycle

This commit is contained in:
2026-07-23 19:01:37 +08:00
parent 99b5d2b2da
commit 8ab5631446
7 changed files with 173 additions and 20 deletions
+6 -2
View File
@@ -190,7 +190,8 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
- rewards: [G] - rewards: [G]
Output: Output:
- prompts: [B, P_max] - prompts: [B, P_max], left-padded
- prompt_mask: [B, P_max]
- responses: [B, G, R_max] - responses: [B, G, R_max]
- masks: [B, G, R_max] - masks: [B, G, R_max]
- rewards: [B, G] - 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"]) R_max = max(r.size(0) for b in batch for r in b["responses"])
prompts = torch.zeros(B, P_max, dtype=torch.long) 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) responses = torch.zeros(B, G, R_max, dtype=torch.long)
masks = torch.zeros(B, G, R_max, dtype=torch.bool) masks = torch.zeros(B, G, R_max, dtype=torch.bool)
rewards = torch.zeros(B, G, dtype=torch.float32) rewards = torch.zeros(B, G, dtype=torch.float32)
for i, b in enumerate(batch): for i, b in enumerate(batch):
p_len = b["prompts"].size(0) 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"] rewards[i, : b["rewards"].size(0)] = b["rewards"]
for g in range(min(G, len(b["responses"]))): for g in range(min(G, len(b["responses"]))):
r_len = b["responses"][g].size(0) r_len = b["responses"][g].size(0)
@@ -217,6 +220,7 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
return { return {
"prompts": prompts, "prompts": prompts,
"prompt_mask": prompt_mask,
"responses": responses, "responses": responses,
"masks": masks, "masks": masks,
"rewards": rewards, "rewards": rewards,
+88 -12
View File
@@ -35,6 +35,7 @@ class RawRollout:
Fields: Fields:
prompts: Tokenized prompts, shape ``[B, P_len]``. 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]``. responses: Generated response token IDs, shape ``[B, G, R_max]``.
response_mask: Boolean mask for real (non-pad) response tokens, response_mask: Boolean mask for real (non-pad) response tokens,
shape ``[B, G, R_max]``. shape ``[B, G, R_max]``.
@@ -47,6 +48,7 @@ class RawRollout:
""" """
prompts: Tensor prompts: Tensor
prompt_mask: Tensor
responses: Tensor responses: Tensor
response_mask: Tensor response_mask: Tensor
logprobs_old: Tensor logprobs_old: Tensor
@@ -143,6 +145,15 @@ class RolloutGenerator:
``add_generation_prompt=True`` so rollout prompts match the ``add_generation_prompt=True`` so rollout prompts match the
format the policy was SFT-trained on. 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) prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
B = len(prompt_texts) B = len(prompt_texts)
G = self.group_size G = self.group_size
@@ -161,6 +172,15 @@ class RolloutGenerator:
rep_window=self.rep_window, rep_window=self.rep_window,
return_logprobs=True, 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. # Each element is (token_ids, logprobs); pad to max length.
max_len = 0 max_len = 0
@@ -171,10 +191,12 @@ class RolloutGenerator:
device = self.scheduler.device device = self.scheduler.device
P_len = max(len(ids) for ids in flat_prompt_ids) P_len = max(len(ids) for ids in flat_prompt_ids)
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device) 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): 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 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) 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) response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
@@ -201,6 +223,7 @@ class RolloutGenerator:
return RawRollout( return RawRollout(
prompts=prompts_tensor, prompts=prompts_tensor,
prompt_mask=prompt_mask,
responses=responses, responses=responses,
response_mask=response_mask, response_mask=response_mask,
logprobs_old=logprobs_old, logprobs_old=logprobs_old,
@@ -239,17 +262,35 @@ class RolloutGenerator:
f"{list(batch.keys())}" f"{list(batch.keys())}"
) )
prompt_texts: List[str] = [] try:
flat_prompt_ids: List[List[int]] = [] prompt_texts = self.tokenizer.apply_chat_template(
for messages in messages_list: messages_list, tokenize=False, add_generation_prompt=True
text = self.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
) )
ids = self.tokenizer.apply_chat_template( if (
messages, tokenize=True, add_generation_prompt=True not isinstance(prompt_texts, list)
) or len(prompt_texts) != len(messages_list)
prompt_texts.append(text) or not all(isinstance(text, str) for text in prompt_texts)
flat_prompt_ids.append(list(ids)) ):
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 return prompt_texts, flat_prompt_ids
@staticmethod @staticmethod
@@ -308,6 +349,7 @@ class RolloutRunner:
self.rollout_interval = rollout_interval self.rollout_interval = rollout_interval
self._cache: Optional[RolloutResult] = None self._cache: Optional[RolloutResult] = None
self._cache_key = None
self._steps_since_rollout: int = 0 self._steps_since_rollout: int = 0
def step(self): def step(self):
@@ -317,12 +359,40 @@ class RolloutRunner:
def clear_cache(self): def clear_cache(self):
"""Force next call to re-run rollout.""" """Force next call to re-run rollout."""
self._cache = None 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: def _score(self, raw: RawRollout) -> RolloutResult:
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts) 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 device = raw.prompts.device
return RolloutResult( return RolloutResult(
prompts=raw.prompts, prompts=raw.prompts,
prompt_mask=raw.prompt_mask,
responses=raw.responses, responses=raw.responses,
response_mask=raw.response_mask, response_mask=raw.response_mask,
rewards=rewards.to(device=device), rewards=rewards.to(device=device),
@@ -337,9 +407,15 @@ class RolloutRunner:
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval`` Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
or when the cache is empty. 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) raw = self.generator.generate(batch)
self._cache = self._score(raw) self._cache = self._score(raw)
self._cache_key = cache_key
self._steps_since_rollout = 0 self._steps_since_rollout = 0
return self._cache, True return self._cache, True
return self._cache, False return self._cache, False
+15 -3
View File
@@ -158,6 +158,11 @@ class BaseStrategy(ABC):
""" """
pass 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: def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
"""Run offline or online forward depending on runner injection.""" """Run offline or online forward depending on runner injection."""
if self._rollout_runner is None: if self._rollout_runner is None:
@@ -166,8 +171,6 @@ class BaseStrategy(ABC):
result, is_fresh = self._rollout_runner(batch) result, is_fresh = self._rollout_runner(batch)
if is_fresh: if is_fresh:
self._on_rollout_refresh() self._on_rollout_refresh()
if self.executor and self.executor.sync_gradients:
self._rollout_runner.step()
train_batch = self.prepare_from_rollout(result) train_batch = self.prepare_from_rollout(result)
return self.compute_loss(train_batch) return self.compute_loss(train_batch)
@@ -411,6 +414,12 @@ class GRPOStrategy(BaseStrategy):
responses_flat = responses.view(-1, response_len) responses_flat = responses.view(-1, response_len)
masks_flat = masks.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_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) prompt_len = prompt_expanded.size(1)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-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 # 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] S = key_pad.shape[-1]
causal = torch.tril( causal = torch.tril(
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device) 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]: def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
return { return {
"prompts": result.prompts, "prompts": result.prompts,
"prompt_mask": result.prompt_mask,
"responses": result.responses, "responses": result.responses,
"masks": result.response_mask, "masks": result.response_mask,
"rewards": result.rewards, "rewards": result.rewards,
+1
View File
@@ -83,6 +83,7 @@ class Trainer:
if executor.sync_gradients: if executor.sync_gradients:
self._call_callbacks("on_optimizer_step", context) self._call_callbacks("on_optimizer_step", context)
context.optimizer.step() context.optimizer.step()
context.strategy.on_optimizer_step()
context.optimizer.zero_grad() context.optimizer.zero_grad()
if context.scheduler: if context.scheduler:
+3 -2
View File
@@ -937,8 +937,9 @@ def test_grpo_collate_variable_lengths():
assert result["masks"].shape == (2, 2, 4) assert result["masks"].shape == (2, 2, 4)
assert result["rewards"].shape == (2, 2) assert result["rewards"].shape == (2, 2)
# Check padding: item 1 prompt is length 2, padded to 3 # Prompts are left-padded so each response follows its real prompt tokens.
assert result["prompts"][1, 2] == 0 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 # Check response content: item 0, response 0 is [4,5] padded to 4
assert result["responses"][0, 0, 0] == 4 assert result["responses"][0, 0, 0] == 4
+10 -1
View File
@@ -63,6 +63,7 @@ def _make_frozen(model, device):
def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"): def _make_rollout_result(B=2, G=4, P=6, R=8, device="cpu"):
return RolloutResult( return RolloutResult(
prompts=torch.randint(3, 200, (B, P), device=device), 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), responses=torch.randint(3, 200, (B, G, R), device=device),
response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device), response_mask=torch.ones(B, G, R, dtype=torch.bool, device=device),
rewards=torch.randn(B, G, 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) r = _make_rollout_result(device=device)
batch = strat.prepare_from_rollout(r) batch = strat.prepare_from_rollout(r)
assert batch["prompts"] is r.prompts assert batch["prompts"] is r.prompts
assert batch["prompt_mask"] is r.prompt_mask
assert batch["responses"] is r.responses assert batch["responses"] is r.responses
assert batch["masks"] is r.response_mask assert batch["masks"] is r.response_mask
assert batch["rewards"] is r.rewards 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)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) 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({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2 assert runner.calls == 2
assert runner.step_calls == 1 assert runner.step_calls == 2
def test_grpo_resync_when_new_rollout_result(device): 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)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device)) runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.calls == 2 assert runner.calls == 2
assert runner.step_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)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
runner.swap_result(_make_rollout_result(device=device)) runner.swap_result(_make_rollout_result(device=device))
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 2 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)) runner = _RecordingRunner(_make_rollout_result(device=device))
strat.set_rollout_runner(runner) strat.set_rollout_runner(runner)
strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)}) strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
strat.on_optimizer_step()
assert runner.step_calls == 1 assert runner.step_calls == 1
+50
View File
@@ -71,6 +71,18 @@ class ConstantRewardModel(BaseRewardModel):
return torch.full((B, G), float(self.value)) 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): def _make_config(vocab_size=200, max_position_embeddings=128):
return AutoRegressiveLMConfig( return AutoRegressiveLMConfig(
vocab_size=vocab_size, vocab_size=vocab_size,
@@ -111,6 +123,7 @@ def _make_instruction_batch(n=2):
def test_raw_rollout_fields(): def test_raw_rollout_fields():
r = RawRollout( r = RawRollout(
prompts=torch.zeros(2, 4, dtype=torch.long), 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), responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool), response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5), logprobs_old=torch.zeros(2, 3, 5),
@@ -124,6 +137,7 @@ def test_raw_rollout_fields():
def test_rollout_result_inherits_raw_rollout_fields(): def test_rollout_result_inherits_raw_rollout_fields():
r = RolloutResult( r = RolloutResult(
prompts=torch.zeros(2, 4, dtype=torch.long), 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), responses=torch.zeros(2, 3, 5, dtype=torch.long),
response_mask=torch.ones(2, 3, 5, dtype=torch.bool), response_mask=torch.ones(2, 3, 5, dtype=torch.bool),
logprobs_old=torch.zeros(2, 3, 5), 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.rewards.shape == (2, 3)
assert r.prompts.shape == (2, 4) assert r.prompts.shape == (2, 4)
assert r.responses.shape == (2, 3, 5) assert r.responses.shape == (2, 3, 5)
assert r.prompt_mask.shape == (2, 4)
def test_base_reward_model_is_abstract(): 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.responses.shape == (2, 3, 5)
assert r.response_mask.shape == (2, 3, 5) assert r.response_mask.shape == (2, 3, 5)
assert r.logprobs_old.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.prompt_texts) == 2
assert len(r.response_texts) == 2 assert len(r.response_texts) == 2
assert len(r.response_texts[0]) == 3 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): def test_rollout_generator_mask_matches_responses(device):
"""Positions beyond a response's length are pad (mask False).""" """Positions beyond a response's length are pad (mask False)."""
gen, _ = _make_generator(device, group_size=2, max_tokens=6) 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 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): def test_rollout_runner_step_triggers_new_rollout(device):
runner, _ = _make_runner(device, rollout_interval=2) runner, _ = _make_runner(device, rollout_interval=2)
batch = _make_instruction_batch() batch = _make_instruction_batch()