Add KL anchor to BC reference policy in trainer

Self-play drift fix: regularize loss with KL(current || BC_reference).
Config: training.kl_anchor_ckpt + training.kl_anchor_beta. Loaded once
at trainer init, frozen. KL computed over legal actions only.
Hypothesis: appropriate beta keeps pretrained competence during self-play
finetune, escaping the c9 catastrophic forgetting.
This commit is contained in:
2026-05-11 15:24:14 +09:00
parent 9fdfa88b23
commit d850070ed4
2 changed files with 48 additions and 0 deletions
@@ -67,6 +67,13 @@ class TrainingConfig(StrictModel):
# head learn faster relative to policy loss. Useful when value_prediction_error
# is large but loss/value is tiny because of the normalization.
value_loss_weight: float = 1.0
# Optional KL anchor to a reference (e.g. behavior-cloned) policy. The
# reference network is loaded once at trainer start and frozen; on every
# gradient step we add `kl_anchor_beta * KL(current || reference)` to the
# loss. Anchors self-play training to the pretrained policy and prevents
# catastrophic forgetting / drift to weak self-play equilibria.
kl_anchor_ckpt: str | None = None
kl_anchor_beta: float = 0.0
@field_validator(
"games_per_iter",
@@ -83,6 +83,28 @@ class IsMctsTrainer:
self.metrics_path = self.run_dir / "metrics.jsonl"
self.rng = random.Random(config.run.seed)
# Optional KL anchor: load a frozen reference network whose policy we
# use to regularize updates (KL(current || reference) added to loss).
# Prevents drift from a BC pretrained baseline during self-play.
self.kl_anchor_ref: AlphaZeroNet | None = None
self.kl_anchor_beta = float(config.training.kl_anchor_beta)
if config.training.kl_anchor_ckpt and self.kl_anchor_beta > 0.0:
ref_path = Path(config.training.kl_anchor_ckpt)
if not ref_path.exists():
raise FileNotFoundError(f"kl_anchor_ckpt not found: {ref_path}")
ref_payload = torch.load(ref_path, map_location=self.device, weights_only=False)
self.kl_anchor_ref = AlphaZeroNet.from_config(
self.input_dim, self.action_size, config
).to(self.device)
self.kl_anchor_ref.load_state_dict(ref_payload["network"])
self.kl_anchor_ref.eval()
for p in self.kl_anchor_ref.parameters():
p.requires_grad = False
print(
f"[trainer] KL anchor: ref={ref_path} beta={self.kl_anchor_beta}",
flush=True,
)
def _resolve_device(self, device: torch.device | str) -> torch.device:
token = str(device)
if token == "auto":
@@ -263,6 +285,22 @@ class IsMctsTrainer:
value_loss = nn.functional.mse_loss(value_pred / v_scale, value_target / v_scale)
value_weight = float(self.config.training.value_loss_weight)
loss = policy_loss + value_weight * value_loss
# KL anchor: KL(current || reference) over legal actions only.
# Encourages current policy to stay close to the reference (BC) policy.
kl_anchor_loss = 0.0
if self.kl_anchor_ref is not None and self.kl_anchor_beta > 0.0:
with torch.no_grad():
ref_logits, _ref_value = self.kl_anchor_ref(info, legal)
ref_log_probs = torch.log_softmax(ref_logits, dim=-1)
# KL(current || ref) = sum_a p_cur(a) * (log p_cur(a) - log p_ref(a))
cur_probs = log_probs.exp()
legal_f = legal.float()
kl_per_action = cur_probs * (log_probs - ref_log_probs) * legal_f
kl = kl_per_action.sum(dim=-1).mean()
loss = loss + self.kl_anchor_beta * kl
kl_anchor_loss = float(kl.item())
self.optimizer.zero_grad(set_to_none=True)
loss.backward()
if self.config.optimization.grad_clip > 0:
@@ -271,6 +309,9 @@ class IsMctsTrainer:
self.config.optimization.grad_clip,
)
self.optimizer.step()
# Stash auxiliary loss for the IterationMetrics path; we return only
# the three primary scalars for backward compatibility.
self._last_kl_anchor_loss = kl_anchor_loss
return float(policy_loss.item()), float(value_loss.item()), float(loss.item())
def _evaluate(self, iteration: int) -> dict[str, float | int]: