import argparse
import json
import math
import os
import struct
import sys
import numpy as np
CODEBOOK_4BIT = np.array([
-2.7325896, -2.0690172, -1.6180464, -1.2562312,
-0.9423405, -0.6567591, -0.3880483, -0.1283950,
0.1283950, 0.3880483, 0.6567591, 0.9423405,
1.2562312, 1.6180464, 2.0690172, 2.7325896,
], dtype=np.float32)
def fwht_inplace(x: np.ndarray) -> None:
n = x.shape[-1]
h = 1
while h < n:
x_view = x.reshape(x.shape[:-1] + (n // (h * 2), h * 2))
a = x_view[..., :h].copy()
b = x_view[..., h:].copy()
x_view[..., :h] = a + b
x_view[..., h:] = a - b
h *= 2
x *= 1.0 / math.sqrt(n)
def nrmse(reference: np.ndarray, actual: np.ndarray) -> float:
diff = actual - reference
sum_sq_diff = float(np.sum(diff.astype(np.float64) ** 2))
sum_sq_ref = float(np.sum(reference.astype(np.float64) ** 2))
return math.sqrt(sum_sq_diff / max(sum_sq_ref, 1e-30))
def dequant_packed(packed: np.ndarray, norms: np.ndarray, hd: int) -> np.ndarray:
nkv, P, _ = packed.shape
assert packed.shape[2] == hd // 2, f"packed hd/2 mismatch: {packed.shape[2]} != {hd//2}"
assert norms.shape == (nkv, P), f"norms shape mismatch: {norms.shape} != ({nkv},{P})"
p_u8 = packed.astype(np.uint8) low = (p_u8 & 0x0F).astype(np.int32) high = ((p_u8 >> 4) & 0x0F).astype(np.int32)
indices = np.empty((nkv, P, hd), dtype=np.int32)
indices[..., 0::2] = low indices[..., 1::2] = high
rotated = CODEBOOK_4BIT[indices]
inv_sqrt_hd = 1.0 / math.sqrt(hd)
scale = norms * inv_sqrt_hd rotated *= scale[:, :, np.newaxis]
fwht_inplace(rotated)
return rotated
def forward_quantize(x: np.ndarray, hd: int):
nkv, P, _ = x.shape
val = x.copy()
fwht_inplace(val)
norms = np.sqrt(np.sum(val.astype(np.float64) ** 2, axis=-1)).astype(np.float32)
safe_norms = np.where(norms > 1e-10, norms, 1.0)
scaled = val / safe_norms[:, :, np.newaxis] * math.sqrt(hd)
boundaries = 0.5 * (CODEBOOK_4BIT[:-1] + CODEBOOK_4BIT[1:]) cmp = (scaled[:, :, :, np.newaxis] > boundaries[np.newaxis, np.newaxis, np.newaxis, :])
indices = cmp.sum(axis=-1).astype(np.uint8)
even = indices[..., 0::2] & 0xF odd = indices[..., 1::2] & 0xF packed = (even | (odd << 4)).astype(np.uint8)
return packed, norms
def run_self_test(hd: int = 256, n_vecs: int = 1000, seed: int = 42) -> float:
rng = np.random.default_rng(seed)
nkv, P = 4, n_vecs // 4
x = rng.standard_normal((nkv, P, hd)).astype(np.float32)
packed, enc_norms = forward_quantize(x, hd)
x_hat = dequant_packed(packed, enc_norms, hd)
err = nrmse(x, x_hat)
return err
def load_f32_bin(path: str, shape) -> np.ndarray:
data = np.fromfile(path, dtype=np.float32)
return data.reshape(shape)
def load_u8_bin(path: str, shape) -> np.ndarray:
data = np.fromfile(path, dtype=np.uint8)
return data.reshape(shape)
def run_analysis(
tq_root: str,
dense_root: str,
csv_path: str,
summary_path: str,
layers: list,
):
rows = []
for layer_cfg in layers:
ll = layer_cfg["layer"]
nkv = layer_cfg["nkv"]
hd = layer_cfg["hd"]
kv_seq_len = layer_cfg["kv_seq_len"]
P = kv_seq_len tag = f"L{ll:02d}"
k_packed = load_u8_bin(
os.path.join(tq_root, f"hf2q_k_packed_layer{ll:02d}_pos22.u8.bin"),
(nkv, P, hd // 2),
)
k_norms = load_f32_bin(
os.path.join(tq_root, f"hf2q_k_norms_layer{ll:02d}_pos22.f32.bin"),
(nkv, P),
)
v_packed = load_u8_bin(
os.path.join(tq_root, f"hf2q_v_packed_layer{ll:02d}_pos22.u8.bin"),
(nkv, P, hd // 2),
)
v_norms = load_f32_bin(
os.path.join(tq_root, f"hf2q_v_norms_layer{ll:02d}_pos22.f32.bin"),
(nkv, P),
)
k_dense_full = load_f32_bin(
os.path.join(dense_root, f"hf2q_cache_k_layer{ll:02d}_pos22.bin"),
(nkv, 23, hd),
)
k_dense = k_dense_full[:, :P, :]
v_dense_full = load_f32_bin(
os.path.join(dense_root, f"hf2q_cache_v_layer{ll:02d}_pos22.bin"),
(nkv, 23, hd),
)
v_dense = v_dense_full[:, :P, :]
k_tq = dequant_packed(k_packed, k_norms, hd) v_tq = dequant_packed(v_packed, v_norms, hd)
for (op_label, tq_data, dense_data) in [
("k", k_tq, k_dense),
("v", v_tq, v_dense),
]:
for h in range(nkv):
for p in range(P):
ref = dense_data[h, p, :] act = tq_data[h, p, :] diff = act - ref
mad = float(np.max(np.abs(diff)))
err = nrmse(ref, act)
rows.append({
"layer": ll,
"op": op_label,
"head": h,
"pos": p,
"max_abs_diff": mad,
"nrmse": err,
})
os.makedirs(os.path.dirname(csv_path), exist_ok=True)
with open(csv_path, "w") as f:
f.write("layer,op,head,pos,max_abs_diff,nrmse\n")
for r in rows:
f.write(f"{r['layer']},{r['op']},{r['head']},{r['pos']},"
f"{r['max_abs_diff']:.6f},{r['nrmse']:.6f}\n")
print(f"CSV written: {csv_path} ({len(rows)} rows)")
from collections import defaultdict
summary = defaultdict(lambda: {"max_nrmse": 0.0, "worst_head": -1, "worst_pos": -1,
"max_abs": 0.0, "violations": 0})
for r in rows:
key = (r["layer"], r["op"])
s = summary[key]
if r["nrmse"] > s["max_nrmse"]:
s["max_nrmse"] = r["nrmse"]
s["worst_head"] = r["head"]
s["worst_pos"] = r["pos"]
if r["max_abs_diff"] > s["max_abs"]:
s["max_abs"] = r["max_abs_diff"]
if r["nrmse"] > 0.15:
s["violations"] += 1
print("\n--- Summary (per layer x op) ---")
print(f"{'layer':>5} {'op':>2} {'max_nrmse':>10} {'worst_head':>11} {'worst_pos':>10} "
f"{'max_abs_diff':>13} {'violations>0.15':>16}")
for (ll, op) in sorted(summary.keys()):
s = summary[(ll, op)]
print(f"{ll:>5} {op:>2} {s['max_nrmse']:>10.6f} {s['worst_head']:>11} {s['worst_pos']:>10} "
f"{s['max_abs']:>13.6f} {s['violations']:>16}")
all_nrmse_ok = all(s["max_nrmse"] <= 0.15 for s in summary.values())
all_abs_ok = all(s["max_abs"] <= 1.0 for s in summary.values())
any_violation = any(s["violations"] > 0 for s in summary.values())
if all_nrmse_ok and all_abs_ok:
verdict = "E1"
verdict_reason = "All nrmse <= 0.15 and max_abs_diff <= 1.0: packed cache dequantizes within kernel bound. Bug NOT in H3 (encode/cache). Downstream: H1 kernel / H2 FWHT / H4 dispatch."
elif any_violation:
verdict = "E2"
worst_r = max(rows, key=lambda r: r["nrmse"])
verdict_reason = (f"E2: nrmse violation at layer={worst_r['layer']} op={worst_r['op']} "
f"head={worst_r['head']} pos={worst_r['pos']} nrmse={worst_r['nrmse']:.6f}. "
f"Encode/cache is broken.")
else:
verdict = "E3"
verdict_reason = "E3: mixed result — partial encode errors, check per-layer breakdown."
print(f"\nVERDICT: {verdict}")
print(f" {verdict_reason}")
worst_overall = max(rows, key=lambda r: r["nrmse"])
worst_abs_overall = max(rows, key=lambda r: r["max_abs_diff"])
total_violations = sum(1 for r in rows if r["nrmse"] > 0.15)
with open(summary_path, "w") as f:
f.write("# TQ C0b Localize — Dequant Diff Summary\n\n")
f.write(f"Date: 2026-04-21 | CFA session: cfa-20260421-C0b-localize | Worker 2\n\n")
f.write(f"## Verdict: {verdict}\n\n")
f.write(f"{verdict_reason}\n\n")
f.write("## Per-layer x op worst-case\n\n")
f.write("| layer | op | max_nrmse | worst_head | worst_pos | max_abs_diff | nrmse_violations |\n")
f.write("|------:|---:|----------:|-----------:|----------:|-------------:|-----------------:|\n")
for (ll, op) in sorted(summary.keys()):
s = summary[(ll, op)]
f.write(f"| {ll} | {op} | {s['max_nrmse']:.6f} | {s['worst_head']} | {s['worst_pos']} "
f"| {s['max_abs']:.6f} | {s['violations']} |\n")
f.write("\n## Worst-case cell (nrmse)\n\n")
f.write(f"layer={worst_overall['layer']} op={worst_overall['op']} "
f"head={worst_overall['head']} pos={worst_overall['pos']} "
f"nrmse={worst_overall['nrmse']:.6f} max_abs_diff={worst_overall['max_abs_diff']:.6f}\n\n")
f.write("## Worst-case cell (max_abs_diff)\n\n")
f.write(f"layer={worst_abs_overall['layer']} op={worst_abs_overall['op']} "
f"head={worst_abs_overall['head']} pos={worst_abs_overall['pos']} "
f"nrmse={worst_abs_overall['nrmse']:.6f} max_abs_diff={worst_abs_overall['max_abs_diff']:.6f}\n\n")
f.write(f"## Total cells with nrmse > 0.15: {total_violations}\n\n")
f.write("## Kernel bound status\n\n")
f.write(f"nrmse bound (< 0.15) holds: {'YES' if all_nrmse_ok else 'NO'}\n")
f.write(f"max_abs_diff bound (< 1.0) holds: {'YES' if all_abs_ok else 'NO'}\n")
print(f"Summary written: {summary_path}")
return verdict, summary, rows
def main():
parser = argparse.ArgumentParser(description="TQ dequantizer + diff — CFA C0b Worker 2")
parser.add_argument("--self-test", action="store_true", help="Run round-trip self-test only")
parser.add_argument("--hd", type=int, default=256, help="head_dim for self-test")
parser.add_argument(
"--tq-root",
default="/tmp/cfa-20260421-C0b-localize/dumps/tq",
help="Root dir for TQ dumps",
)
parser.add_argument(
"--dense-root",
default="/tmp/cfa-20260421-C0b-localize/dumps/dense",
help="Root dir for dense F32 dumps",
)
parser.add_argument(
"--csv",
default="/opt/hf2q/docs/tq-c0b-localize-2026-04-21-raw.csv",
help="Output CSV path",
)
parser.add_argument(
"--summary",
default="/opt/hf2q/docs/tq-c0b-localize-2026-04-21-summary.md",
help="Output summary markdown path",
)
args = parser.parse_args()
if args.self_test:
print("Running round-trip self-test (1000 vectors, hd=256)...")
err = run_self_test(hd=256, n_vecs=1000, seed=42)
print(f" hd=256: nrmse={err:.6f}")
err512 = run_self_test(hd=512, n_vecs=1000, seed=99)
print(f" hd=512: nrmse={err512:.6f}")
worst = max(err, err512)
if worst < 0.12:
print(f"PASS nrmse={worst:.6f} (< 0.12 gate; expected ~0.097 for correct 4-bit)")
sys.exit(0)
else:
print(f"FAIL nrmse={worst:.6f} (>= 0.12 gate) -- dequantizer is wrong")
sys.exit(1)
layers = [
{"layer": 0, "nkv": 8, "hd": 256, "kv_seq_len": 22},
{"layer": 5, "nkv": 2, "hd": 512, "kv_seq_len": 22},
]
verdict, summary, rows = run_analysis(
tq_root=args.tq_root,
dense_root=args.dense_root,
csv_path=args.csv,
summary_path=args.summary,
layers=layers,
)
worst_r = max(rows, key=lambda r: r["nrmse"])
worst_abs = max(rows, key=lambda r: r["max_abs_diff"])
print(f"\nFINAL: verdict={verdict} worst_nrmse={worst_r['nrmse']:.6f} "
f"(L{worst_r['layer']:02d} {worst_r['op']} h{worst_r['head']} p{worst_r['pos']}) "
f"worst_abs={worst_abs['max_abs_diff']:.6f}")
if __name__ == "__main__":
main()