#!/usr/bin/env python3 # ============================================================================= # Project : localvm-research # File : experiments/micro/expF_error_accumulation/benchmark.py # Purpose : Layer-sensitivity map — degrade-one / repair-one / repair-top-k # (does layer-restricted escalation cut bytes-per-escalation?) # Author : Simon-Pierre Boucher # Contact : contact@spboucher.ai # Created : 2026-08-12 # Modified : 2026-08-12 # Platform : macOS / Apple Silicon (arm64) — MLX / Metal # License : All rights reserved (research code) # ============================================================================= """Experiment F — error accumulation / layer sensitivity (charter §9.F). Per depth-group 4-bit degradation and repair on Qwen3-1.7B, teacher-forced over reference greedy trajectories. Usage: .venv/bin/python benchmark.py [--per-domain 8] [--gen-tokens 128] [--groups 7] """ from __future__ import annotations import argparse import json import sys import time from datetime import datetime, timezone from pathlib import Path import mlx.core as mx import mlx.nn as nn import numpy as np from mlx_lm import load REPO_ROOT = Path(__file__).resolve().parents[3] sys.path.insert(0, str(REPO_ROOT / "benchmarks")) sys.path.insert(0, str(REPO_ROOT / "src")) from hardware_manifest import collect_manifest # noqa: E402 from localvm.quality.decision_stats import greedy_generate, kl_ref_vs, teacher_forced_stats # noqa: E402 GROUP_SIZE = 64 BITS = 4 def layer_index(path: str) -> int | None: parts = path.split(".") for i, p in enumerate(parts): if p == "layers" and i + 1 < len(parts) and parts[i + 1].isdigit(): return int(parts[i + 1]) return None def quantized_weights_by_layer(model) -> dict[str, tuple[int, mx.array]]: """{param_path: (layer_idx, 4-bit-dequantized bf16 weight)} for all divisible Linear layers inside transformer blocks.""" out = {} for path, module in model.named_modules(): li = layer_index(path) if li is None or not isinstance(module, nn.Linear): continue if module.weight.shape[-1] % GROUP_SIZE != 0: continue w = module.weight.astype(mx.float32) qw, sc, bi = mx.quantize(w, group_size=GROUP_SIZE, bits=BITS) deq = mx.dequantize(qw, sc, bi, group_size=GROUP_SIZE, bits=BITS).astype(mx.bfloat16) mx.eval(deq) out[path] = (li, deq) return out def apply_config(model, qweights: dict, originals: dict, degrade_layers: set[int]) -> None: """Set each eligible Linear to 4-bit dequant if its layer ∈ degrade_layers, else restore the original bf16 weight.""" for path, module in model.named_modules(): if path in qweights: li, deq = qweights[path] module.weight = deq if li in degrade_layers else originals[path] def evaluate(model, trajectories, ref_stats) -> dict: agrees, kls = [], [] for t, ref in zip(trajectories, ref_stats): qs = teacher_forced_stats(model, t["full_ids"], t["start"]) ref_next = np.array(t["full_ids"][t["start"]:]) agrees.append((qs["argmax"] == ref_next).astype(np.int8)) kls.append(kl_ref_vs(qs["logprobs"], ref["logprobs"])) return { "agreement_rate": float(np.concatenate(agrees).mean()), "mean_kl": float(np.mean(np.concatenate(kls))), } def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--model", default="mlx-community/Qwen3-1.7B-bf16") ap.add_argument("--gen-tokens", type=int, default=128) ap.add_argument("--per-domain", type=int, default=8) ap.add_argument("--groups", type=int, default=7) args = ap.parse_args() domains = json.loads((REPO_ROOT / "benchmarks/datasets/eval_prompts.json").read_text())["domains"] print(f"loading {args.model} …", flush=True) model, tokenizer = load(args.model) n_layers = len(model.model.layers) bounds = np.linspace(0, n_layers, args.groups + 1).astype(int) groups = [set(range(bounds[i], bounds[i + 1])) for i in range(args.groups)] trajectories = [] t0 = time.time() for domain, plist in domains.items(): for prompt in plist[: args.per_domain]: ids = tokenizer.apply_chat_template( [{"role": "user", "content": prompt}], add_generation_prompt=True) gen = greedy_generate(model, tokenizer, ids, args.gen_tokens) if len(gen) >= 8: trajectories.append({"domain": domain, "full_ids": list(ids) + gen, "start": len(ids)}) print(f"{len(trajectories)} trajectories in {time.time()-t0:.0f}s", flush=True) ref_stats = [teacher_forced_stats(model, t["full_ids"], t["start"]) for t in trajectories] print("precomputing 4-bit weights …", flush=True) qweights = quantized_weights_by_layer(model) originals = {p: m.weight for p, m in model.named_modules() if p in qweights} all_layers = set(range(n_layers)) runs: dict[str, dict] = {} def run(tag: str, degrade: set[int]) -> dict: apply_config(model, qweights, originals, degrade) r = evaluate(model, trajectories, ref_stats) r["degraded_layers"] = sorted(degrade) runs[tag] = r print(f" {tag:>24}: agree={r['agreement_rate']:.4f} KL={r['mean_kl']:.4f}", flush=True) return r print("all-4-bit floor:", flush=True) floor = run("all_4bit", all_layers) print("degrade-one (rest bf16):", flush=True) for gi, g in enumerate(groups): run(f"degrade_g{gi}_L{min(g)}-{max(g)}", g) print("repair-one (rest 4-bit):", flush=True) for gi, g in enumerate(groups): run(f"repair_g{gi}_L{min(g)}-{max(g)}", all_layers - g) # repair-top-k by measured repair value lost = 1.0 - floor["agreement_rate"] repair_value = { gi: runs[f"repair_g{gi}_L{min(g)}-{max(g)}"]["agreement_rate"] - floor["agreement_rate"] for gi, g in enumerate(groups) } order = sorted(repair_value, key=repair_value.get, reverse=True) print("repair-top-k (best groups bf16):", flush=True) for k in (2, 3): keep = set().union(*(groups[gi] for gi in order[:k])) run(f"repair_top{k}_groups_{sorted(order[:k])}", all_layers - keep) apply_config(model, qweights, originals, set()) # restore ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") out_dir = REPO_ROOT / "results" / "expF_error_accumulation" / ts out_dir.mkdir(parents=True) (out_dir / "results.json").write_text(json.dumps({ "experiment": "expF_error_accumulation", "author": "Simon-Pierre Boucher", "contact": "contact@spboucher.ai", "manifest": collect_manifest(), "config": vars(args), "bits": BITS, "group_size": GROUP_SIZE, "n_layers": n_layers, "layer_groups": [sorted(g) for g in groups], "agreement_lost_all4bit": lost, "repair_value_by_group": {str(k): v for k, v in repair_value.items()}, "runs": runs, }, indent=2)) print(f"\nwrote {out_dir / 'results.json'}") if __name__ == "__main__": main()