"""Split prefill vs upstream prefill on identical step-0 inputs: output and prefix-KV agreement.""" import sys, types, torch from PIL import Image from diffusers.models.transformers.transformer_qwenimage21 import QwenImage21KVCache from common import ROOT, BASE from prompts import EVALUATION precision = sys.argv[1] if precision == 'fp8': from fp8_runtime import load_pipeline else: from nvfp4_runtime import load_pipeline from acceleration import accelerate_pipeline pipe = accelerate_pipeline(load_pipeline(str(BASE), str(ROOT / precision / 'release' / 'transformer'))) tf = pipe.transformer; split = tf.forward; original = tf._image21_original_forward e = lambda a, b: ((a.float() - b.float()).square().mean() / b.float().square().mean()).sqrt().item() def check(*args, **kw): if kw.get('kv_cache_mode') == 'extract': scratch = QwenImage21KVCache(len(tf.transformer_blocks)) ref = original(*args, **{**kw, 'kv_cache': scratch})[0] out = split(*args, **kw)[0] n = out.shape[1] kv = max(e(a.k, b.k) for a, b in zip(kw['kv_cache'].layer_caches, scratch.layer_caches)) print({'target_output_nrmse': e(out, ref[:, -n:]), 'max_prefix_k_nrmse': kv, 'tokens': ref.shape[1], 'target': n}, flush=True) return (out,) return split(*args, **kw) tf.forward = check with torch.inference_mode(): pipe(prompt=EVALUATION[8], width=1024, height=1024, generator=torch.Generator('cuda').manual_seed(1)) pipe(prompt=EVALUATION[1], width=2048, height=2048, generator=torch.Generator('cuda').manual_seed(1)) ref = Image.open(ROOT / 'evaluation' / 'bf16' / '03.png').resize((1024, 1024)) pipe(prompt='Turn this room into a warm evening scene', image=ref, width=1024, height=1024, generator=torch.Generator('cuda').manual_seed(1))