ixim's picture
Add files using upload-large-folder tool
4f03424 verified
Raw History Blame Contribute Delete
2.89 kB
"""Stratified weight reconstruction and GPU matmul probe, not an image-quality score."""
import argparse
import gc
import json
import time
from pathlib import Path
import mlx.core as mx
from scripts.common import SOURCE, eligible, tensors, write_json
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--source', type=Path, default=SOURCE)
ap.add_argument('--output', default='artifacts/precision-probe.json')
args = ap.parse_args()
results = []
for component, layers in [('transformer', ('0', '15', '31')), ('text_encoder', ('0','17','35'))]:
for path, name, info in tensors(args.source / component):
if not eligible(component, name, info['shape']):
continue
if not any(f'.{i}.' in name for i in layers):
continue
if not any(x in name for x in ('to_q.weight', 'img_mlp.out.weight', 'q_proj.weight', 'down_proj.weight')):
continue
all_weights = mx.load(str(path))
# 256 evenly spaced rows, all columns/groups: bounded memory and reproducible.
full = all_weights[name]
weight = full[mx.linspace(0, full.shape[0]-1, 256).astype(mx.int32)]
mx.eval(weight)
del full, all_weights
wf = weight.astype(mx.float32)
mx.random.seed(2026)
x = mx.random.normal((1,256,weight.shape[-1])).astype(mx.bfloat16)
expected = x @ weight.T
mx.eval(expected)
for bits in (4,6,8):
q,s,b = mx.quantize(weight, bits=bits, group_size=64)
rec = mx.dequantize(q,s,b,bits=bits,group_size=64).astype(mx.float32)
err = float(mx.sqrt(mx.mean((rec-wf)**2) / mx.mean(wf**2)).item())
fn = lambda: mx.quantized_matmul(x,q,s,b,bits=bits,group_size=64,transpose=True)
for _ in range(3): mx.eval(fn())
start=time.perf_counter()
for _ in range(10): mx.eval(fn())
ms=(time.perf_counter()-start)*100
y=fn().astype(mx.float32); e=expected.astype(mx.float32)
out_err=float(mx.sqrt(mx.mean((y-e)**2)/mx.mean(e**2)).item())
row=dict(component=component,tensor=name,source_shape=info['shape'],sample_rows=256,
bits=bits,weight_relative_rmse=err,synthetic_matmul_relative_rmse=out_err,matmul_ms=ms)
results.append(row)
print(json.dumps(row),flush=True)
del weight,wf,x,expected,q,s,b,rec,y,e
gc.collect(); mx.clear_cache()
write_json(args.output,dict(group_size=64,mode='affine',device=mx.metal.device_info(),results=results,
limitations='Sampled rows and synthetic inputs; not calibrated activations, image quality or end-to-end speed.'))
if __name__ == '__main__': main()