Code / experiments/micro/expF_error_accumulation/benchmark.py
experiments/micro/expF_error_accumulation/benchmark.py
181 lines
#!/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()