training / experiments /exp_step_embedding.py
cazyundee's picture
whitelist analysis scripts; fix pid-recycling in launch/job_table
2f4aa7f verified
Raw History Blame Contribute Delete
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
@contextlib.contextmanager
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()