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