Image21-MLX-8bit / scripts /verify_cache.py
ixim's picture
Add files using upload-large-folder tool
4f03424 verified
Raw History Blame Contribute Delete
2.3 kB
"""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()