Files
AstrAI/tests/trainer/test_grpo_strategy.py
T
ViperEkura 5ba21f4eb3 refactor: eliminate test duplication via shared helpers
- Add tests/helpers.py with shared config, dataset, tokenizer, executor, and assertion helpers
- Replace 15 copies of device one-liner with session-scoped fixture
- Collapse 5 near-identical Dataset subclasses into RandomTokenDataset
- Remove duplicate _make_config/_make_model/_make_frozen and FakeTokenizer/FakeExecutor definitions
- Make test_callbacks and test_early_stopping use existing train_config_factory
- Replace 6 duplicate meta.json read blocks with load_shard_meta
- Fix mkdtemp leaks in test_lora.py with TemporaryDirectory
2026-07-27 22:34:53 +08:00

178 lines
5.7 KiB
Python

import pytest
import torch
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.strategy import GRPOStrategy
from tests.helpers import FakeExecutor, make_frozen, make_model, make_rollout_config
def _make_batch(
batch_size=2, group_size=4, prompt_len=8, response_len=12, device="cpu"
):
"""Construct a GRPO batch with deterministic shapes.
Returns dict with prompts [B, P], responses [B, G, R], masks [B, G, R],
rewards [B, G].
"""
prompts = torch.randint(0, 200, (batch_size, prompt_len), device=device)
responses = torch.randint(
0, 200, (batch_size, group_size, response_len), device=device
)
masks = torch.ones(batch_size, group_size, response_len, device=device)
rewards = torch.randn(batch_size, group_size, device=device)
return {
"prompts": prompts,
"responses": responses,
"masks": masks,
"rewards": rewards,
}
@pytest.fixture
def grpo_strategy(device):
"""Build a GRPOStrategy with a small real model and fake executor."""
model, config = make_model(device)
old_model = make_frozen(model, device)
ref_model = make_frozen(model, device)
strategy = GRPOStrategy(
model=model,
device=device,
old_model=old_model,
ref_model=ref_model,
clip_eps=0.2,
kl_coef=0.01,
group_size=4,
model_fn=lambda c=config: AutoRegressiveLM(c).to(device=device),
executor=FakeExecutor(),
)
return strategy, device
def test_grpo_loss_is_finite(grpo_strategy):
"""compute_loss returns a finite scalar."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
assert loss.dim() == 0
assert torch.isfinite(loss).item()
def test_grpo_loss_backward(grpo_strategy):
"""Loss is differentiable w.r.t. policy model parameters."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
has_grad = any(
p.grad is not None and p.grad.abs().sum().item() > 0
for p in strategy.model.parameters()
)
assert has_grad
def test_grpo_ref_model_not_updated(grpo_strategy):
"""Backward should not populate gradients on ref_model."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
for p in strategy.ref_model.parameters():
assert p.grad is None
def test_grpo_old_model_not_updated(grpo_strategy):
"""Backward should not populate gradients on old_model."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
loss.backward()
for p in strategy.old_model.parameters():
assert p.grad is None
def test_grpo_prompt_tokens_masked(grpo_strategy):
"""When only prompt-equivalent tokens are unmasked (response mask all 0),
the policy loss should be zero (no valid tokens contribute)."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
batch["masks"] = torch.zeros_like(batch["masks"])
loss = strategy.compute_loss(batch)
assert loss.item() == pytest.approx(0.0, abs=1e-6)
def test_grpo_identical_rewards_zero_advantage(grpo_strategy):
"""When all group rewards are identical, advantage is 0 -> policy_loss is 0.
Only the KL term remains (which is 0 when policy == ref at init)."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
batch["rewards"] = torch.ones(batch["rewards"].shape, device=device)
loss = strategy.compute_loss(batch)
assert loss.item() == pytest.approx(0.0, abs=1e-5)
def test_grpo_sync_old_model(grpo_strategy):
"""sync_old_model copies current policy weights into old_model."""
strategy, device = grpo_strategy
with torch.no_grad():
for p in strategy.model.parameters():
p.add_(0.05)
policy_sd = strategy.model.state_dict()
old_sd = strategy.old_model.state_dict()
differs_before = any(
not torch.allclose(policy_sd[k], old_sd[k]) for k in policy_sd if k in old_sd
)
assert differs_before
strategy.sync_old_model()
old_sd_after = strategy.old_model.state_dict()
matches = all(
torch.allclose(policy_sd[k], old_sd_after[k])
for k in policy_sd
if k in old_sd_after
)
assert matches
def test_grpo_partial_mask(grpo_strategy):
"""Only the first half of response tokens are valid."""
strategy, device = grpo_strategy
batch = _make_batch(device=device)
B, G, R = batch["masks"].shape
half = R // 2
batch["masks"][:, :, half:] = 0.0
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()
def test_grpo_clipping_effect(grpo_strategy):
"""After diverging policy from ref, ratio should be clipped to [1-eps, 1+eps]
on the surrogate. Verify loss is finite and non-zero for distinct rewards."""
strategy, device = grpo_strategy
with torch.no_grad():
for p in strategy.model.parameters():
p.add_(0.3)
batch = _make_batch(device=device)
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()
assert loss.abs().item() > 1e-4
def test_grpo_no_reduction_param():
"""GRPOStrategy.__init__ must not accept ``reduction`` (removed)."""
import inspect
sig = inspect.signature(GRPOStrategy.__init__)
assert "reduction" not in sig.parameters
def test_grpo_shapes_3d_batch(grpo_strategy):
"""Verify compute_loss handles non-square prompt/response lengths."""
strategy, device = grpo_strategy
batch = _make_batch(
batch_size=3, group_size=4, prompt_len=10, response_len=8, device=device
)
loss = strategy.compute_loss(batch)
assert torch.isfinite(loss).item()