"""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()