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:
+29
-24
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user