Normalize PUCT Q + add mirror-descent policy target
Codex follow-up diagnostics identified two MCTS+training-loop issues that
together cap finetune-from-BC at the heuristic ceiling:
1. PUCT Q is in raw score units (~±100 for value_scale=100), but the
exploration bonus c_puct * prior * sqrt(N) / (1+n) is on order of 1-10
for our parameter ranges. Result: a single bad backup pushes q_eff
well below the bonus floor and that action is effectively never
visited again. With only 50 sims/move this is catastrophic for the
policy-improvement operator. Fix: divide q_eff by config.q_scale
(default 100, configurable) inside _select_action. Backups and value
targets remain in raw score units; only the selection signal is
normalized. AlphaZero canonical convention.
2. The current kl_anchor_beta path adds KL(current || ref) directly to
the loss. That preserves BC but prevents improvement (gradient
actively pulls policy back to reference). The standard regularized
policy improvement operator is to mix the target instead:
pi_target = softmax(alpha * log(pi_mcts) + (1-alpha) * log(pi_ref))
Anneal alpha from low (rely on BC) to high (rely on MCTS) over
training. Network learns to follow the regularized target, which
stays near BC early but lets MCTS-discovered improvements through
later.
Config additions:
- mcts.q_scale (default 100.0): PUCT Q divisor
- training.md_target_ref_ckpt: reference policy path (alternative to kl_anchor)
- training.md_target_alpha_start / _end / _iters: linear alpha schedule
Both Python mcts.py and Cython mcts.pyx updated; parity test passes.
Tests: 19/19.
Hypothesis: with normalized PUCT the network can actually explore and
exploit prior knowledge competently at 50 sims, and the mirror-descent
target lets self-play improvement happen while BC anchors the trajectory.
This is the operator-side fix that c9 (no anchor, collapsed) and c10/c11
(loss-side KL anchor, preserved-but-stuck) both missed.
This commit is contained in:
@@ -32,6 +32,12 @@ class MctsConfig(StrictModel):
|
||||
eval_n_simulations: int = 0
|
||||
root_dirichlet_alpha: float = 0.0
|
||||
root_dirichlet_epsilon: float = 0.0
|
||||
# Divisor applied to Q values inside PUCT to bring them onto roughly the
|
||||
# same scale as the exploration bonus. With value_scale=100 score units,
|
||||
# raw Q can swing ±100 while c_puct * prior * sqrt(N) is ~1-10, so a single
|
||||
# bad backup permanently kills an action. Setting q_scale=100 normalizes Q
|
||||
# to ~[-1, 1] (consistent with AlphaZero's convention).
|
||||
q_scale: float = 100.0
|
||||
|
||||
@field_validator("n_simulations", "max_depth", "parallel_simulations")
|
||||
@classmethod
|
||||
@@ -74,6 +80,16 @@ class TrainingConfig(StrictModel):
|
||||
# catastrophic forgetting / drift to weak self-play equilibria.
|
||||
kl_anchor_ckpt: str | None = None
|
||||
kl_anchor_beta: float = 0.0
|
||||
# Mirror-descent target mixing for policy loss. Alternative to kl_anchor;
|
||||
# blends MCTS visit distribution with the reference (BC) policy in log
|
||||
# space, then trains the network to match. pi_target = softmax(
|
||||
# alpha * log(pi_mcts) + (1 - alpha) * log(pi_ref)
|
||||
# ). Anneal alpha from low (rely on BC) to high (rely on MCTS) over
|
||||
# training. Requires kl_anchor_ckpt to be set as the reference source.
|
||||
md_target_ref_ckpt: str | None = None
|
||||
md_target_alpha_start: float = 0.3
|
||||
md_target_alpha_end: float = 0.8
|
||||
md_target_alpha_iters: int = 500
|
||||
|
||||
@field_validator(
|
||||
"games_per_iter",
|
||||
|
||||
@@ -263,6 +263,7 @@ class IsMctsSearcher:
|
||||
sqrt_total = math.sqrt(max(1, total_visits))
|
||||
best_score = -float("inf")
|
||||
best_action = legal_actions[0]
|
||||
q_scale = float(getattr(self.config, "q_scale", 100.0)) or 1.0
|
||||
for action in legal_actions:
|
||||
n = node.visits.get(action, 0)
|
||||
virtual = node.virtual_visits.get(action, 0)
|
||||
@@ -274,7 +275,8 @@ class IsMctsSearcher:
|
||||
q_eff = (
|
||||
node.value_sum.get(action, 0.0) - virtual * self.config.virtual_loss_value
|
||||
) / n_eff
|
||||
score = q_eff + self.config.c_puct * prior * sqrt_total / (1 + n_eff)
|
||||
# Normalize Q to match exploration-bonus scale; see mcts.pyx for details.
|
||||
score = q_eff / q_scale + self.config.c_puct * prior * sqrt_total / (1 + n_eff)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_action = action
|
||||
|
||||
@@ -553,13 +553,17 @@ cdef class IsMctsSearcher:
|
||||
cdef double sqrt_total
|
||||
cdef double prior
|
||||
cdef double q_eff
|
||||
cdef double q_normalized
|
||||
cdef double score
|
||||
cdef double best_score = -float("inf")
|
||||
cdef int best_action = int(legal_actions[0])
|
||||
cdef double q_scale = float(getattr(self.config, "q_scale", 100.0))
|
||||
cdef _ArrayMap visits = <_ArrayMap>node.visits
|
||||
cdef _ArrayMap virtual_visits = <_ArrayMap>node.virtual_visits
|
||||
cdef _ArrayMap priors = <_ArrayMap>node.priors
|
||||
cdef _ArrayMap value_sum = <_ArrayMap>node.value_sum
|
||||
if q_scale <= 0.0:
|
||||
q_scale = 1.0
|
||||
for action in legal_actions:
|
||||
total_visits += visits.get_int(action, 0) + virtual_visits.get_int(action, 0)
|
||||
sqrt_total = math.sqrt(max(1, total_visits))
|
||||
@@ -575,7 +579,12 @@ cdef class IsMctsSearcher:
|
||||
value_sum.get_float(action, 0.0)
|
||||
- virtual * float(self.config.virtual_loss_value)
|
||||
) / n_eff
|
||||
score = q_eff + float(self.config.c_puct) * prior * sqrt_total / (1 + n_eff)
|
||||
# Normalize Q to roughly [-1, 1] so the exploration bonus
|
||||
# (c_puct * prior * sqrt(N) / (1+n)) competes on the right scale.
|
||||
# Without this, raw score-units Q (±100) dominates and a single
|
||||
# noisy backup kills exploration of low-prior actions.
|
||||
q_normalized = q_eff / q_scale
|
||||
score = q_normalized + float(self.config.c_puct) * prior * sqrt_total / (1 + n_eff)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_action = action
|
||||
|
||||
@@ -104,6 +104,29 @@ class IsMctsTrainer:
|
||||
f"[trainer] KL anchor: ref={ref_path} beta={self.kl_anchor_beta}",
|
||||
flush=True,
|
||||
)
|
||||
# Optional mirror-descent reference policy: blend MCTS visit dist
|
||||
# with this reference in log space before computing CE loss.
|
||||
self.md_target_ref: AlphaZeroNet | None = None
|
||||
if config.training.md_target_ref_ckpt:
|
||||
md_ref_path = Path(config.training.md_target_ref_ckpt)
|
||||
if not md_ref_path.exists():
|
||||
raise FileNotFoundError(f"md_target_ref_ckpt not found: {md_ref_path}")
|
||||
md_payload = torch.load(md_ref_path, map_location=self.device, weights_only=False)
|
||||
self.md_target_ref = AlphaZeroNet.from_config(
|
||||
self.input_dim, self.action_size, config
|
||||
).to(self.device)
|
||||
self.md_target_ref.load_state_dict(md_payload["network"])
|
||||
self.md_target_ref.eval()
|
||||
for p in self.md_target_ref.parameters():
|
||||
p.requires_grad = False
|
||||
print(
|
||||
f"[trainer] mirror-descent ref={md_ref_path} "
|
||||
f"alpha {config.training.md_target_alpha_start} -> "
|
||||
f"{config.training.md_target_alpha_end} over "
|
||||
f"{config.training.md_target_alpha_iters} iters",
|
||||
flush=True,
|
||||
)
|
||||
self._current_md_alpha = float(config.training.md_target_alpha_start)
|
||||
|
||||
def _resolve_device(self, device: torch.device | str) -> torch.device:
|
||||
token = str(device)
|
||||
@@ -138,6 +161,15 @@ class IsMctsTrainer:
|
||||
return metrics
|
||||
|
||||
def run_iteration(self, iteration: int) -> IterationMetrics:
|
||||
# Update mirror-descent alpha schedule (linear from start to end over alpha_iters).
|
||||
if self.md_target_ref is not None:
|
||||
cfg_t = self.config.training
|
||||
n = max(1, int(cfg_t.md_target_alpha_iters))
|
||||
frac = min(1.0, float(iteration) / float(n))
|
||||
self._current_md_alpha = float(
|
||||
cfg_t.md_target_alpha_start
|
||||
+ (cfg_t.md_target_alpha_end - cfg_t.md_target_alpha_start) * frac
|
||||
)
|
||||
print(
|
||||
f"[iter {iteration}] self-play start (workers={self.config.training.num_workers})",
|
||||
flush=True,
|
||||
@@ -280,7 +312,25 @@ class IsMctsTrainer:
|
||||
)
|
||||
logits, value_pred = self.network(info, legal)
|
||||
log_probs = torch.log_softmax(logits, dim=-1)
|
||||
policy_loss = -(pi * log_probs).sum(dim=-1).mean()
|
||||
# Optional mirror-descent target: blend MCTS visit distribution with
|
||||
# the frozen reference (BC) policy in log space, then train CE to that
|
||||
# blended target. This is the standard regularized policy improvement
|
||||
# operator: pi_target = softmax(alpha * log(pi_mcts) + (1-alpha) * log(pi_ref)).
|
||||
if self.md_target_ref is not None:
|
||||
with torch.no_grad():
|
||||
ref_logits, _ref_value = self.md_target_ref(info, legal)
|
||||
ref_log_probs = torch.log_softmax(ref_logits, dim=-1)
|
||||
alpha = float(self._current_md_alpha)
|
||||
# Clamp pi to avoid log(0); MCTS visit dist already has only legal
|
||||
# actions positive, so this affects illegal actions which the mask
|
||||
# in the network forward already zeroed out via -inf logits.
|
||||
log_pi = torch.log(pi.clamp_min(1.0e-12))
|
||||
mixed = alpha * log_pi + (1.0 - alpha) * ref_log_probs
|
||||
mixed = mixed.masked_fill(~legal, torch.finfo(mixed.dtype).min)
|
||||
pi_target = torch.softmax(mixed, dim=-1)
|
||||
policy_loss = -(pi_target * log_probs).sum(dim=-1).mean()
|
||||
else:
|
||||
policy_loss = -(pi * log_probs).sum(dim=-1).mean()
|
||||
v_scale = float(self.network.value_scale)
|
||||
value_loss = nn.functional.mse_loss(value_pred / v_scale, value_target / v_scale)
|
||||
value_weight = float(self.config.training.value_loss_weight)
|
||||
|
||||
Reference in New Issue
Block a user