"""Mechanism analysis of the loop operator — runs on the REAL model or the tiny one. Measures (and saves to JSON) four things that explain *why* the loop behaves as it does, on Laguna's actual architecture: A. Output drift vs K — RMS logit drift from baseline as K grows, for damped-RK vs naive. "naive diverges, damping stays bounded." B. MoE routing churn — fraction of tokens whose top-k expert set changes across loop iterations, block- vs layer-mode (the routing-thrash claim). C. Anchor beta sweep — drift vs beta (beta=1 -> 0, i.e. baseline). D. Per-layer drift — loop ONLY each window layer (K=3) and measure the logit drift it causes. Tests whether the GLOBAL-attention layer is "load-bearing" (perturbs most) — the mechanistic basis for "looping global layers hurts". Real model: uv run python scripts/tf_analyze_dynamics.py --device cuda --output results_dynamics.json Tiny check: uv run python scripts/tf_analyze_dynamics.py --tiny """ from __future__ import annotations import argparse import json import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import torch from looped_laguna import LoopConfig, load_model_and_tokenizer, patch, unpatch DEFAULT_TOKENIZER = str(Path(__file__).resolve().parent.parent / "laguna_src") PROMPT = ("Question: A block slides down a frictionless ramp and then across a rough floor. " "Which quantity is conserved during the slide down the ramp, and why does the block " "eventually stop on the floor? Explain step by step.\nAnswer:") def rms(a: torch.Tensor, b: torch.Tensor) -> float: return (a - b).pow(2).mean().sqrt().item() @torch.no_grad() def main() -> None: p = argparse.ArgumentParser(description="Loop-operator mechanism analysis (real or tiny model).") p.add_argument("--model", default="poolside/Laguna-XS.2") p.add_argument("--tokenizer", default=DEFAULT_TOKENIZER) p.add_argument("--dtype", default="bfloat16") p.add_argument("--device", default="cuda") p.add_argument("--tiny", action="store_true") p.add_argument("--tiny-layers", type=int, default=8) p.add_argument("--width", type=int, default=4) p.add_argument("--center", type=float, default=0.5) p.add_argument("--output", default=None, help="save results JSON for offline plotting") args = p.parse_args() model, tok = load_model_and_tokenizer( args.model, args.tokenizer, dtype=args.dtype, device=args.device, tiny=args.tiny, tiny_layers=args.tiny_layers, ) device = next(model.parameters()).device n = model.config.num_hidden_layers window = LoopConfig.from_depth_fraction(n, center_frac=args.center, width=args.width).window lt = model.config.layer_types ids = tok(PROMPT, return_tensors="pt").input_ids.to(device) print(f"model: {'tiny' if args.tiny else args.model} | layers={n} | window={window} | " f"types={[lt[i] for i in range(window[0], window[1]+1)]} | tokens={ids.shape[1]}") def logits(cfg: LoopConfig | None): unpatch(model) if cfg is None else patch(model, cfg) out = model(input_ids=ids, use_cache=False).logits unpatch(model) return out base = logits(None) out: dict = {"meta": {"model": "tiny" if args.tiny else args.model, "n_layers": n, "window": list(window), "window_types": [lt[i] for i in range(window[0], window[1] + 1)]}} # ---- A. output drift vs K ---- print("\n[A] RMS logit drift vs K:") print(f" {'K':>2} {'layer-damped':>13} {'block-damped':>13} {'layer-naive':>12} {'block-naive':>12}") out["drift_vs_K"] = {} for K in range(1, 7): row = {} for mode, naive, key in [("layer", False, "layer_damped"), ("block", False, "block_damped"), ("layer", True, "layer_naive"), ("block", True, "block_naive")]: row[key] = rms(logits(LoopConfig(window=window, K=K, mode=mode, naive=naive)), base) out["drift_vs_K"][K] = row print(f" {K:>2} {row['layer_damped']:>13.3f} {row['block_damped']:>13.3f} " f"{row['layer_naive']:>12.3f} {row['block_naive']:>12.3f}") # ---- B. routing churn (block vs layer) ---- captures: dict[int, list[torch.Tensor]] = {} handles = [] for i in range(window[0], window[1] + 1): gate = getattr(model.model.layers[i].mlp, "gate", None) if gate is not None: handles.append(gate.register_forward_hook( lambda _m, _i, o, i=i: captures.setdefault(i, []).append(o[2].detach().clone()))) def churn(cfg: LoopConfig) -> dict[int, float]: captures.clear() logits(cfg) pl = {} for i, caps in captures.items(): first = [set(r.tolist()) for r in caps[0]] diffs = [sum(set(r.tolist()) != f for r, f in zip(caps[k], first)) / len(first) for k in range(1, len(caps))] pl[i] = sum(diffs) / len(diffs) if diffs else 0.0 return pl print("\n[B] MoE routing churn at K=3 (frac. tokens whose top-k expert set changes vs 1st iter):") out["routing_churn"] = {} for mode in ("layer", "block"): pl = churn(LoopConfig(window=window, K=3, mode=mode)) out["routing_churn"][mode] = {str(i): pl.get(i, 0.0) for i in range(window[0], window[1] + 1)} vals = list(out["routing_churn"][mode].values()) print(f" {mode + '-mode':>11}: " + " ".join(f"L{i}={v:.3f}" for i, v in zip(range(window[0], window[1] + 1), vals)) + f" mean={sum(vals)/len(vals):.3f}") for h in handles: h.remove() # ---- C. beta sweep ---- print("\n[C] RMS drift vs anchor beta (layer, K=3); beta=1 -> 0:") out["beta_sweep"] = {} for beta in (0.0, 0.25, 0.5, 0.75, 1.0): d = rms(logits(LoopConfig(window=window, K=3, mode="layer", beta=beta)), base) out["beta_sweep"][beta] = d print(f" beta={beta:.2f} drift={d:.4f}") # ---- D. per-layer drift (loop ONLY layer i) ---- print("\n[D] Per-layer drift — loop ONLY each window layer (K=3), RMS logit drift:") out["per_layer_drift"] = {} for i in range(window[0], window[1] + 1): d = rms(logits(LoopConfig(layers=(i,), K=3, mode="layer")), base) out["per_layer_drift"][str(i)] = {"drift": d, "type": lt[i]} print(f" layer {i} ({lt[i]:18s}): drift={d:.4f}") if args.output: Path(args.output).write_text(json.dumps(out, indent=2)) print(f"\nwrote {args.output}") if __name__ == "__main__": main()