File size: 2,888 Bytes
4f03424
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
"""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()