From d850070ed469409729c9b7ed16152a325bebb825 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Mon, 11 May 2026 15:24:14 +0900 Subject: [PATCH] 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. --- .../games/classic/ismcts/config.py | 7 ++++ .../games/classic/ismcts/trainer.py | 41 +++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/src/coolrl_lost_cities/games/classic/ismcts/config.py b/src/coolrl_lost_cities/games/classic/ismcts/config.py index e08588d..d5e5c53 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/config.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/config.py @@ -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", diff --git a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py index e509bff..18990f8 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py @@ -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]: