File size: 2,297 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
"""Validate native text-prefix caching against the exact uncached transformer."""
import argparse
import json
import time
from pathlib import Path
import mlx.core as mx
from mlx_vlm.models.qwen_image.weights import load_transformer
from mlx_vlm.models.qwen_image.kv_cache import QwenImageKVCache
from scripts.common import write_json

def main():
    ap=argparse.ArgumentParser()
    ap.add_argument('--model',type=Path,required=True)
    ap.add_argument('--output',type=Path,required=True)
    args=ap.parse_args()
    model=load_transformer(args.model); mx.eval(model.parameters())
    mx.random.seed(20260926)
    emb=mx.random.normal((1,48,4096)).astype(mx.bfloat16)
    latent=mx.random.normal((1,1024,64)).astype(mx.bfloat16)
    cache=QwenImageKVCache(len(model.transformer_blocks))
    kw=dict(encoder_hidden_states=emb,img_shape=(1,32,32))
    init=model(hidden_states=latent,timestep=mx.array([0.9],dtype=mx.bfloat16),kv_cache=cache,kv_cache_mode='extract',**kw)
    mx.eval(init,cache.arrays())
    rows=[]
    for t in (0.7,0.3):
        latent=mx.random.normal((1,1024,64)).astype(mx.bfloat16)
        timestep=mx.array([t],dtype=mx.bfloat16)
        start=time.perf_counter(); expected=model(hidden_states=latent,timestep=timestep,**kw); mx.eval(expected)
        uncached=time.perf_counter()-start
        start=time.perf_counter(); actual=model(hidden_states=latent,timestep=timestep,kv_cache=cache,kv_cache_mode='cached',**kw); mx.eval(actual)
        cached=time.perf_counter()-start
        delta=actual.astype(mx.float32)-expected.astype(mx.float32)
        rel=float(mx.sqrt(mx.mean(delta**2)/mx.mean(expected.astype(mx.float32)**2)).item())
        row=dict(timestep=t,max_abs_error=float(mx.max(mx.abs(delta)).item()),relative_rmse=rel,
                 uncached_seconds=uncached,cached_seconds=cached)
        rows.append(row); print(row,flush=True)
        # Kernel splitting can change BF16 rounding; guard against material drift.
        if rel>0.01: raise ValueError(f'KV cache changes predictions materially: {rel}')
    write_json(args.output,dict(model=str(args.model),passed=True,rows=rows,
                                note='Real transformer with synthetic conditioning/latents; BF16 numerical tolerance, not pixel equality.'))

if __name__=='__main__': main()