DeepSeek-V4.1-Flash-NVFP4 / verify_transcode.py
jon1012's picture
Sync verify_transcode.py (streamed Engram + GPU option, verifier, synthetic engram test)
1abd948 verified
Raw
History Blame Contribute Delete
2.91 kB
#!/usr/bin/env python3
"""Prove the MXFP4->NVFP4 expert transcode is bit-exact, and measure the Engram FP4 loss.
Reconstructs values from BOTH representations and compares:
source : FP4_TABLE[nibble] * scale_e8m0
output : FP4_TABLE[nibble] * scale_e4m3 * global_f32
A single non-identical element fails the run.
"""
import json, os, sys
import torch
from safetensors import safe_open
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from dsv41_fp4_stream import (FP4_TABLE, SRC_BLOCK, NVFP4_BLOCK, EXPERT_RE, ENGRAM_RE,
unpack_e2m1)
src_dir, out_dir, shard = sys.argv[1], sys.argv[2], sys.argv[3]
sf = safe_open(os.path.join(src_dir, shard), framework="pt")
of = safe_open(os.path.join(out_dir, shard), framework="pt")
src_keys, out_keys = set(sf.keys()), set(of.keys())
n_exp = n_bad = 0
max_abs = 0.0
checked_elems = 0
for name in sorted(src_keys):
if not EXPERT_RE.match(name):
continue
base = name[:-len(".weight")]
a = unpack_e2m1(sf.get_tensor(name))
sa = sf.get_tensor(base + ".scale").float()
ref = a.view(a.shape[0], -1, SRC_BLOCK) * sa.unsqueeze(-1)
ref = ref.reshape(a.shape)
b = unpack_e2m1(of.get_tensor(base + ".weight"))
sb = of.get_tensor(base + ".weight_scale").float()
g = of.get_tensor(base + ".weight_scale_2").float()
got = b.view(b.shape[0], -1, NVFP4_BLOCK) * (sb * g).unsqueeze(-1)
got = got.reshape(b.shape)
d = (ref - got).abs().max().item()
max_abs = max(max_abs, d)
if d != 0.0:
n_bad += 1
if n_bad <= 3:
print(f" MISMATCH {base} max|d|={d:.6g}")
n_exp += 1
checked_elems += ref.numel()
if n_exp >= 200: # 200 experts is plenty and keeps the check quick
break
print(f"[experts] {n_exp} weights checked ({checked_elems/1e6:.1f}M elements), "
f"mismatches={n_bad}, max|delta|={max_abs:.6g}")
n_eng = 0
for name in sorted(src_keys):
if not ENGRAM_RE.match(name):
continue
base = name[:-len(".weight")]
rows = 65536
w = sf.get_slice(name)[:rows]
s = sf.get_slice(base + ".scale")[:rows].float()
ref = w.float().view(rows, -1, SRC_BLOCK) * s.unsqueeze(-1)
b = unpack_e2m1(of.get_slice(base + ".weight")[:rows])
sb = of.get_slice(base + ".scale")[:rows].float()
got = b.view(rows, -1, SRC_BLOCK) * sb.unsqueeze(-1)
a2, b2 = ref.flatten(1), got.flatten(1)
num = (a2 * b2).sum(1); den = a2.norm(dim=1) * b2.norm(dim=1)
ok = den > 0
cos = (num[ok] / den[ok])
rel = ((a2 - b2).norm(dim=1) / a2.norm(dim=1).clamp(min=1e-30))[ok]
print(f"[engram] {base} rows={rows} cos mean={cos.mean():.6f} min={cos.min():.6f} "
f"rel-err mean={rel.mean():.4f}")
n_eng += 1
if n_eng == 0:
print("[engram] none in this shard")
print("RESULT:", "LOSSLESS" if n_bad == 0 and n_exp > 0 else ("FAIL" if n_bad else "no experts here"))