refactor: separate old policy and ref model in GRPO strategy

- Split single ref_model into old_model (importance sampling ratio) and ref_model (frozen KL regularizer)
- Move ref_model/old_model creation from strategy __init__ to TrainContextBuilder, pass as explicit parameters
- Remove periodic sync_ref_model + sync_interval; add sync_old_model for external rollout loop to call
- DPOStrategy also receives ref_model from builder
- Fix std to use unbiased=False (population std per GRPO paper)
- Remove redundant tests (test_grpo_kl_zero_at_init, test_grpo_no_sync_interval_param)
- Remove --grpo_sync_interval CLI arg
This commit is contained in:
2026-07-14 20:03:45 +08:00
parent 3e0007fc91
commit 2c7a71a9c0
4 changed files with 80 additions and 58 deletions
+29 -24
View File
@@ -223,14 +223,13 @@ class DPOStrategy(BaseStrategy):
self,
model: nn.Module,
device: str,
ref_model: nn.Module,
beta: float = 0.1,
reduction: str = "mean",
**kwargs,
):
super().__init__(model, device, **kwargs)
self.ref_model = create_ref_model(
self.model_fn, self.executor.unwrap_model(model)
).to(device=self.device)
self.ref_model = ref_model
self.beta = beta
self.reduction = reduction
@@ -272,40 +271,40 @@ class GRPOStrategy(BaseStrategy):
broadcast across all response tokens. The loss is computed **only on
response tokens** — prompt tokens are masked out.
The strategy expects offline-collected batches (``responses`` / ``rewards``
pre-generated by the current or a recent policy). Call ``sync_ref_model()``
after each data-generation round so ``ref_model`` tracks the sampling policy.
Three model roles are distinguished:
* **Policy** ``self.model`` — the model being trained.
* **Old policy** ``self.old_model`` — the behaviour policy that generated
the responses. Used for the importance sampling ratio
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
* **Reference model** ``self.ref_model`` — a frozen copy of the initial
policy (typically the SFT checkpoint) used **only** for the KL
regularisation term. It is never updated during training.
"""
def __init__(
self,
model: nn.Module,
device: str,
old_model: nn.Module,
ref_model: nn.Module,
clip_eps: float = 0.2,
kl_coef: float = 0.01,
group_size: int = 4,
sync_interval: int = 200,
**kwargs,
):
super().__init__(model, device, **kwargs)
self.ref_model = create_ref_model(
self.model_fn, self.executor.unwrap_model(model)
).to(device=self.device)
self.old_model = old_model
self.ref_model = ref_model
self.clip_eps = clip_eps
self.kl_coef = kl_coef
self.group_size = group_size
self.sync_interval = sync_interval
self._step = 0
def sync_ref_model(self):
"""Copy current model weights to ref model."""
self.ref_model.load_state_dict(self.executor.unwrap_model(self.model))
def sync_old_model(self):
"""Copy current policy weights to old model."""
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
self._step += 1
if self._step % self.sync_interval == 0:
self.sync_ref_model()
batch = move_to_device(batch, self.device)
prompts = batch["prompts"]
responses = batch["responses"]
@@ -332,25 +331,29 @@ class GRPOStrategy(BaseStrategy):
self.model, full_sequences, full_masks, "none"
)[:, prompt_len - 1 :]
with torch.no_grad():
token_log_probs_old = get_logprobs(
self.old_model, full_sequences, full_masks, "none"
)[:, prompt_len - 1 :]
token_log_probs_ref = get_logprobs(
self.ref_model, full_sequences, full_masks, "none"
)[:, prompt_len - 1 :]
# Reshape to [B, G, response_len]
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
token_masks = masks_flat.view(batch_size, group_size, -1).float()
# Group-normalized advantages from scalar per-response rewards.
eps = 1e-8
mean = rewards.mean(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True, unbiased=False)
advantages = (rewards - mean) / (std + eps)
# Broadcast scalar advantage to every response token: [B, G, 1]
advantages = advantages.unsqueeze(-1)
# Token-level ratio and PPO clipping.
log_ratio = token_log_probs_policy - token_log_probs_ref
# Token-level ratio (π_θ / π_old) and PPO clipping.
log_ratio = token_log_probs_policy - token_log_probs_old
ratio = torch.exp(log_ratio)
surr1 = ratio * advantages
@@ -359,8 +362,10 @@ class GRPOStrategy(BaseStrategy):
token_count = token_masks.sum().clamp(min=1.0)
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
# KL penalty with k1 estimator (non-negative): r - log(r) - 1, r=π_ref/π_θ.
r = torch.exp(-log_ratio)
# KL penalty to frozen reference model with k1 estimator (non-negative):
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
r = torch.exp(log_ref_ratio)
kl_per_token = r - torch.log(r + eps) - 1.0
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count