Download experiments/exp_step_embedding.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 10.1 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/experiments/exp_step_embedding.py
- Command line
-
hf download hf://spaces/cazyundee/training/experiments/exp_step_embedding.py
-
curl -L -o exp_step_embedding.py https://huggingface.co/spaces/cazyundee/training/resolve/main/experiments/exp_step_embedding.py
10.1 kB
| #!/usr/bin/env python3 | |
| """ | |
| Does inference-depth degradation come from step-embedding extrapolation (B) | |
| or from genuine recurrent saturation (A)? | |
| HYPOTHESIS | |
| The core applies `step_emb[min(step, max_steps)]`. Beyond `max_steps` every | |
| additional step receives the SAME embedding as the last trained step, so the | |
| core gets an out-of-distribution "which step am I on" signal. If that is what | |
| breaks deep inference, then removing the step embedding entirely (making the | |
| core genuinely step-invariant) should change the shape of the depth curve. | |
| DESIGN (weights are FIXED and IDENTICAL in every condition) | |
| condition `clamped` : stock behaviour, step_emb[min(step, max_steps)] | |
| condition `zeroed` : step_emb forced to 0 at every step -> step-invariant | |
| condition `frozen_k` : every step uses step_emb[k] for a fixed in-distribution | |
| k, isolating "OOD signal" from "no signal at all" | |
| Each condition is evaluated at depths 1/2/4/8/16/32 against a FIXED opponent | |
| pool (random-legal and one-ply material), with identical games, temperature, | |
| seed and adjudication. Only the depth and the embedding policy vary. | |
| DISCRIMINATION | |
| A (saturation) -> all conditions degrade at the same depth | |
| B (OOD step embeddings) -> `zeroed`/`frozen_k` keep improving where `clamped` | |
| degrades | |
| Anything else -> report as ambiguous | |
| Also records a weight-free signal: policy agreement with the depth-1 policy and | |
| mean entropy per depth, to show whether deep inference is drifting or collapsing. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import contextlib | |
| import json | |
| import os | |
| import sys | |
| import time | |
| import torch | |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) | |
| from tinychess.arena import (MaterialAgent, RandomAgent, elo_from_score, | |
| head_to_head, wilson_interval) | |
| from tinychess.arena import policy_diagnostics | |
| from tinychess.batching import encode_batch | |
| from tinychess.checkpoint import load_model_only | |
| def step_embedding_mode(model, mode: str, k: int = 0): | |
| """Temporarily replace the core's step embedding policy. Weights unchanged.""" | |
| core = model.core | |
| original = core.step_emb.data.clone() | |
| try: | |
| if mode == "clamped": | |
| pass # stock behaviour | |
| elif mode == "zeroed": | |
| core.step_emb.data.zero_() # step-invariant core | |
| elif mode == "frozen": | |
| row = original[min(k, original.shape[0] - 1)].clone() | |
| core.step_emb.data[:] = row # same in-distribution signal | |
| else: | |
| raise ValueError(mode) | |
| yield | |
| finally: | |
| core.step_emb.data.copy_(original) | |
| def make_boards(n=48, seed=0): | |
| """Fixed, reproducible mid-game position set (same recipe as exp_compute_scaling).""" | |
| import random | |
| import chess | |
| rng = random.Random(seed) | |
| boards = [] | |
| while len(boards) < n: | |
| b = chess.Board() | |
| for _ in range(rng.randint(0, 40)): | |
| ms = list(b.legal_moves) | |
| if not ms: | |
| break | |
| b.push(rng.choice(ms)) | |
| if not b.is_game_over(): | |
| boards.append(b) | |
| return boards | |
| def policy_stats(model, boards, steps): | |
| """Top-1 move ids, entropy, top-1 prob and illegal mass at a given depth.""" | |
| pb = encode_batch(boards) | |
| with torch.no_grad(): | |
| out = model(**pb.model_args(), steps=steps) | |
| logits = out["logits"] | |
| d = policy_diagnostics(model, boards, steps=steps) | |
| return logits.argmax(-1), d["mean_entropy"], d["mean_top1_prob"], d["illegal_prob_mass"] | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--ckpt", required=True) | |
| ap.add_argument("--depths", default="1,2,4,8,16,32") | |
| ap.add_argument("--modes", default="clamped,zeroed,frozen") | |
| ap.add_argument("--frozen-k", type=int, default=-1, | |
| help="step index for 'frozen' mode; -1 = last trained step") | |
| ap.add_argument("--games", type=int, default=40) | |
| ap.add_argument("--max-plies", type=int, default=140) | |
| ap.add_argument("--temperature", type=float, default=0.3) | |
| ap.add_argument("--adjudicate", default="material", | |
| choices=["material", "none"]) | |
| ap.add_argument("--seed", type=int, default=0) | |
| ap.add_argument("--out", default="runs/step_embedding.json") | |
| a = ap.parse_args() | |
| depths = [int(x) for x in a.depths.split(",")] | |
| modes = [m.strip() for m in a.modes.split(",")] | |
| model, cfg = load_model_only(a.ckpt) | |
| rep = model.param_report() | |
| # `cfg.max_steps` is the TRAINED depth. Capture it before any mutation: | |
| # model.forward clamps n to cfg.max_steps, so without lifting the cap every | |
| # depth > max_steps silently runs only max_steps iterations. | |
| trained_max_steps = int(cfg.max_steps) | |
| max_steps = trained_max_steps | |
| if max(depths) > trained_max_steps: | |
| import torch.nn as nn | |
| old = model.core.step_emb.data | |
| new = torch.zeros(max(depths) + 1, old.shape[1]) | |
| new[:old.shape[0]] = old | |
| new[old.shape[0]:] = old[-1] # stock behaviour: repeat last row | |
| model.core.step_emb = nn.Parameter(new) | |
| model.cfg.max_steps = max(depths) | |
| model.core.cfg.max_steps = max(depths) | |
| print(f"[lifted] step cap {trained_max_steps} -> {max(depths)} " | |
| f"(step_emb rows {old.shape[0]} -> {new.shape[0]})") | |
| k = a.frozen_k if a.frozen_k >= 0 else trained_max_steps | |
| print(f"ckpt={a.ckpt}") | |
| print(f"params={rep['total']:,} trained max_steps={trained_max_steps} frozen_k={k}") | |
| print(f"depths={depths} modes={modes} games={a.games} adj={a.adjudicate}\n") | |
| boards = make_boards(48, seed=a.seed) | |
| opponents = {"random_legal": RandomAgent(), "material_greedy": MaterialAgent()} | |
| res = {"checkpoint": a.ckpt, "params": rep["total"], | |
| "trained_max_steps": trained_max_steps, | |
| "frozen_k": k, "args": vars(a), "conditions": {}} | |
| ref_top1 = {} | |
| for mode in modes: | |
| res["conditions"][mode] = {} | |
| with step_embedding_mode(model, mode, k): | |
| for d in depths: | |
| entry = {} | |
| t0 = time.time() | |
| top1, ent, top1p, illegal = policy_stats(model, boards, d) | |
| entry["mean_entropy"] = round(ent, 4) | |
| entry["mean_top1_prob"] = round(top1p, 4) | |
| entry["illegal_prob_mass"] = float(illegal) | |
| entry["ms_per_position"] = round( | |
| 1000 * (time.time() - t0) / len(boards), 3) | |
| if d == depths[0]: | |
| ref_top1[mode] = top1 | |
| entry["top1_agreement_vs_depth1"] = round( | |
| float((top1 == ref_top1[mode]).float().mean()), 4) | |
| for oname, opp in opponents.items(): | |
| h = head_to_head(model, opp, n_games=a.games, steps=d, | |
| max_plies=a.max_plies, | |
| temperature=a.temperature, | |
| adjudicate=a.adjudicate, seed=100 + a.seed) | |
| ci = wilson_interval(h["wins"], h["draws"], h["games"]) | |
| entry[oname] = { | |
| "score": round(h["score"], 4), | |
| "ci95": [round(ci[0], 3), round(ci[1], 3)], | |
| "elo": round(elo_from_score(h["score"]), 1), | |
| "W": h["wins"], "D": h["draws"], "L": h["losses"], | |
| "mean_plies": round(h["mean_plies"], 1), | |
| } | |
| res["conditions"][mode][str(d)] = entry | |
| r = entry["random_legal"]; mt = entry["material_greedy"] | |
| print(f" [{mode:8s}] d={d:3d} rand={r['score']:.3f}" | |
| f"{str(r['ci95']):16s} mat={mt['score']:.3f}" | |
| f"{str(mt['ci95']):16s} ent={ent:.2f} " | |
| f"agree_d1={entry['top1_agreement_vs_depth1']:.2f} " | |
| f"{entry['ms_per_position']:.1f}ms") | |
| # ---- verdict ---- | |
| def curve(mode, opp): | |
| return [(d, res["conditions"][mode][str(d)][opp]["score"]) for d in depths] | |
| verdict = {} | |
| for opp in opponents: | |
| v = {} | |
| for mode in modes: | |
| c = curve(mode, opp) | |
| best_d, best_s = max(c, key=lambda t: t[1]) | |
| deep = [s for d, s in c if d > trained_max_steps] | |
| shallow = [s for d, s in c if d <= trained_max_steps] | |
| ms = (sum(shallow) / len(shallow)) if shallow else None | |
| md = (sum(deep) / len(deep)) if deep else None | |
| v[mode] = {"curve": c, "best_depth": best_d, "best_score": best_s, | |
| "n_shallow": len(shallow), "n_deep": len(deep), | |
| "mean_shallow": round(ms, 4) if ms is not None else None, | |
| "mean_deep": round(md, 4) if md is not None else None, | |
| "deep_minus_shallow": (round(md - ms, 4) | |
| if (ms is not None and md is not None) | |
| else None)} | |
| verdict[opp] = v | |
| res["verdict"] = verdict | |
| print("\n=== deep-vs-shallow (mean score beyond trained max_steps minus within) ===") | |
| for opp, v in verdict.items(): | |
| print(f" {opp}:") | |
| for mode, d in v.items(): | |
| if d["deep_minus_shallow"] is None: | |
| print(f" {mode:8s} best_depth={d['best_depth']:3d} " | |
| f"shallow={d['mean_shallow']} deep=n/a " | |
| f"(no depths beyond trained max_steps)") | |
| else: | |
| print(f" {mode:8s} best_depth={d['best_depth']:3d} " | |
| f"shallow={d['mean_shallow']:.3f} deep={d['mean_deep']:.3f} " | |
| f"delta={d['deep_minus_shallow']:+.3f}") | |
| os.makedirs(os.path.dirname(os.path.abspath(a.out)), exist_ok=True) | |
| json.dump(res, open(a.out, "w"), indent=2) | |
| print(f"\n[saved] {a.out}") | |
| if __name__ == "__main__": | |
| main() | |