#!/usr/bin/env python3 """Distil the engine's search into a 30 KB network, trained with MLX. The teacher is the engine's own alpha-beta search. Every training position carries the score that search returned at a fixed node count, plus the result of the self-play game it came from. The student is a static evaluation that never searches -- the idea behind DeepMind's searchless grandmaster-level chess, at a size that fits in L1 cache rather than a TPU pod. 934 -> 32 (per perspective, shared weights) -> clipped ReLU -> 1 of 8 buckets Two things here are deliberate and load-bearing. Features come from the engine, never from this file. `featdump` writes the active indices for each position and we read them back. Re-deriving them in Python would mean two implementations of one feature map, and when those drift you get a network that loads, runs, and is quietly wrong -- the worst kind of bug to chase. Quantisation is part of the objective rather than a post-processing step. Weights are projected back into the int8 box after every optimiser step, so the exported network computes the function the trainer actually converged to. """ import os import struct import subprocess import sys import time import numpy as np import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim IN = 934 # feature rows; must match net.rs MAX_F = 96 # feature slots per perspective PAD = IN # index of the permanent zero row H = int(os.environ.get("NET_H", "32")) BUCKETS = int(os.environ.get("NET_B", "8")) QA, QB = 127, 64 # int8 scales for the two layers SCALE = 400 # network units -> centipawns EVAL_WEIGHT = float(os.environ.get("EVAL_W", "0.9")) WEIGHTED = os.environ.get("WEIGHTED", "1") != "0" WEIGHT_CLAMP = float(os.environ.get("WEIGHT_CLAMP", "3")) SHARD_DECAY = float(os.environ.get("SHARD_DECAY", "1.0")) SIGMOID_K = float(os.environ.get("SIG_K", "400")) ENGINE = os.environ.get("ENGINE", "./target/release/sable") FEAT_MAGIC = b"SBF2" # featdump stream header # --------------------------------------------------------------------------- # Data # --------------------------------------------------------------------------- def read_labels(paths, limit): """Pull FEN text and labels out of the augmented self-play shards. Positions are deduplicated across shards, later shard wins. The engine deduplicates within one generation run, but two runs months apart still rediscover the same openings, and in a mixed set the duplicate carries the *older*, weaker teacher's label. Keeping the last occurrence means a mixed set is the union of what each teacher saw with the better label on the overlap, rather than a set where the overlap is graded twice. Each shard also carries a provenance weight: with SHARD_DECAY below 1, older shards count for less, so a mixed set can lean on the newer teacher without throwing the older positions away. """ seen = {} fens, sc, wdl, src = [], [], [], [] for si, path in enumerate(paths): with open(path, "rb") as fh: for line in fh: parts = line.split(b"|") if len(parts) < 3: continue try: score = int(parts[1]) result = int(parts[2]) except ValueError: continue # shard truncated mid-line fen = parts[0].strip() if len(fen.split()) < 2: continue white = fen.split()[1] == b"w" # Everything is stored from the mover's point of view. if not white: score = -score result = 2 - result prev = seen.get(fen) if prev is None: seen[fen] = len(fens) fens.append(fen) sc.append(score) wdl.append(result * 0.5) src.append(si) else: # Later shard wins: newer teacher, better label. sc[prev] = score wdl[prev] = result * 0.5 src[prev] = si if limit and len(fens) >= limit: break print(f" {path}: {len(fens)} unique positions", flush=True) if limit and len(fens) >= limit: break n_shards = max(src) + 1 if src else 1 age = np.array([SHARD_DECAY ** (n_shards - 1 - i) for i in src], np.float32) return fens, np.array(sc, np.float32), np.array(wdl, np.float32), age def parse_features(path, n_expected): """Unpack a featdump stream into padded index arrays. The stream is self-describing: a header, then one variable-length record per position. Records are packed rather than padded to MAX_F, so the array width here is the widest position actually seen -- typically well under the ninety-six slots the engine allows for, which is the difference between the feature cache fitting in memory and not. """ with open(path, "rb") as fh: head = fh.read(8) if len(head) < 8 or head[:4] != FEAT_MAGIC: return None n_in, max_f = struct.unpack(" 1 else None epochs = int(sys.argv[2]) if len(sys.argv) > 2 else 15 seed = int(os.environ.get("SEED", "42")) np.random.seed(seed) mx.random.seed(seed) fens, sc, wdl, age = read_labels(shards, limit) us, them, buckets, width = dump_features(fens, f"data/feat_{limit or 'all'}.bin") n = len(sc) print( f"{n} positions, {epochs} epochs, {(us.nbytes + them.nbytes)/1e6:.0f} MB of " f"features at width {width} of {MAX_F}, seed {seed}" ) # Held-out positions never touched by the optimiser. The exported network # is the epoch that did best here, not whatever the last epoch happened # to land on. val_n = min(200_000, n // 20) val_idx = np.random.permutation(n)[:val_n] train_mask = np.ones(n, bool) train_mask[val_idx] = False train_idx = np.flatnonzero(train_mask) # Blend the teacher's score with the game result. The score is precise but # only as good as the search; the result is noisy but grounded in truth. target = ( EVAL_WEIGHT / (1.0 + np.exp(-sc / SIGMOID_K)) + (1 - EVAL_WEIGHT) * wdl ).astype(np.float32) # Per-position weights. Self-play spends most of its plies in the middlegame, # so the material buckets are far from evenly filled and an unweighted mean # quietly trains the crowded buckets at the expense of the sparse ones. The # correction is inverse bucket frequency, clamped: bucket balance is worth # nudging, not worth letting a handful of endgame positions dominate. counts = np.bincount(buckets, minlength=BUCKETS).astype(np.float64) share = np.where(counts > 0, counts.sum() / (BUCKETS * np.maximum(counts, 1)), 1.0) share = np.clip(share, 1.0 / WEIGHT_CLAMP, WEIGHT_CLAMP) weight = (share[buckets] * age).astype(np.float32) if not WEIGHTED: weight = age.astype(np.float32) print( " bucket counts " + " ".join(f"{int(c)}" for c in counts) + "\n" " weights " + " ".join(f"{w:.2f}" for w in share) ) model = Net() mx.eval(model.parameters()) batch = 16384 steps = len(train_idx) // batch base_lr = float(os.environ.get("LR", "1e-2")) opt = optim.AdamW(learning_rate=base_lr, weight_decay=0.0) def loss_fn(model, u, t, bk, y, w): # SCALE / SIGMOID_K converts network units into the sigmoid's argument. err = (mx.sigmoid(model(u, t, bk) * (SCALE / SIGMOID_K)) - y) ** 2 return mx.sum(err * w) / mx.sum(w) grad_fn = nn.value_and_grad(model, loss_fn) def val_loss(): total, m = 0.0, 0 for i in range(0, val_n, batch): idx = val_idx[i : i + batch] total += float( loss_fn( model, mx.array(us[idx].astype(np.int32)), mx.array(them[idx].astype(np.int32)), mx.array(buckets[idx]), mx.array(target[idx]), mx.array(weight[idx]), ) ) * len(idx) m += len(idx) return total / m best = (float("inf"), None) warmup = 1 for ep in range(epochs): # One linear warmup epoch, then cosine to ~0: early steps explore, # late steps settle into the quantisation grid. if ep < warmup: opt.learning_rate = base_lr * (ep + 1) / warmup else: import math t = (ep - warmup) / max(1, epochs - warmup - 1) opt.learning_rate = base_lr * 0.5 * (1 + math.cos(math.pi * t)) perm = train_idx[np.random.permutation(len(train_idx))] total, t0 = 0.0, time.time() for i in range(steps): idx = perm[i * batch : (i + 1) * batch] loss, grads = grad_fn( model, mx.array(us[idx].astype(np.int32)), mx.array(them[idx].astype(np.int32)), mx.array(buckets[idx]), mx.array(target[idx]), mx.array(weight[idx]), ) opt.update(model, grads) mx.eval(model.parameters(), opt.state) clip_weights(model) total += float(loss) vl = val_loss() star = "" if vl < best[0]: best = (vl, [mx.array(p) for p in (model.ft, model.ft_b, model.out, model.out_b)]) star = " *" print( f"epoch {ep+1:2d}/{epochs} loss {total/steps:.5f} val {vl:.5f}{star} " f"lr {opt.learning_rate.item():.5f} {time.time()-t0:.0f}s", flush=True, ) if best[1] is not None: model.ft, model.ft_b, model.out, model.out_b = best[1] print(f"exporting best-val checkpoint (val {best[0]:.5f})") ftq, fbq, oq, obq = export(model) # How well does the quantised network track the teacher it was distilled # from? Reported on the quantised weights, since those are what ship. # Reported on held-out positions: this is generalisation, not memorisation. k = min(6000, val_n) pick = np.random.choice(val_idx, k, replace=False) pred = np.array( [quantised_eval(ftq, fbq, oq, obq, us[i], them[i], buckets[i]) for i in pick] ) truth = sc[pick] r = np.corrcoef(pred, truth)[0, 1] print( f" quantised net vs teacher: r={r:.4f} " f"mae={np.mean(np.abs(pred-truth)):5.1f}cp " f"rmse={np.sqrt(np.mean((pred-truth)**2)):5.1f}cp" ) if __name__ == "__main__": main()