"""Paired held-out evaluation: 16 disjoint generation prompts + 2 disjoint edits, matched seeds. precision: bf16 | fp8 | nvfp4. --accelerate uses the compiled runtime. BF16 must run first (edits use BF16 outputs as their reference inputs for every precision). Saves PNGs and final latents. """ import argparse, json, time, collections from pathlib import Path import torch from PIL import Image from safetensors.torch import save_file from common import ROOT, BASE, load_bf16_pipeline, sizes from prompts import EVALUATION EDITS = ['Replace the background with a blooming spring garden and preserve the animal.', 'Turn this room into a warm evening scene with lamps switched on, preserving its furniture.'] def load(precision, quant, offload=False): if precision == 'bf16': if offload: from diffusers import QwenImage21Pipeline pipe = QwenImage21Pipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16) pipe.set_progress_bar_config(disable=True) return pipe return load_bf16_pipeline() if precision == 'fp8': from fp8_runtime import load_pipeline else: from nvfp4_runtime import load_pipeline return load_pipeline(str(BASE), quant) @torch.inference_mode() def main(): ap = argparse.ArgumentParser() ap.add_argument('precision', choices=['bf16', 'fp8', 'nvfp4']) ap.add_argument('--quant') ap.add_argument('--name') ap.add_argument('--accelerate', action='store_true') ap.add_argument('--attention', default='bf16', choices=['bf16', 'sage']) ap.add_argument('--fp8-first-steps', type=int, default=0) ap.add_argument('--only', type=str, help='comma-separated subset of case ids, e.g. 00,03,edit-0') ap.add_argument('--offload', action='store_true', help='model CPU offload (same arithmetic, lower peak VRAM; timings not meaningful)') a = ap.parse_args() name = a.name or a.precision out = ROOT / 'evaluation' / name; out.mkdir(parents=True, exist_ok=True) pipe = load(a.precision, a.quant or str(ROOT / a.precision / 'release' / 'transformer'), a.offload) if a.offload: pipe.to('cpu'); pipe.enable_model_cpu_offload() if a.accelerate: from acceleration import accelerate_pipeline accelerate_pipeline(pipe, attention=a.attention) if a.fp8_first_steps: from nvfp4_runtime import enable_fp8_first_steps enable_fp8_first_steps(pipe, str(ROOT / 'fp8' / 'release' / 'transformer'), a.fp8_first_steps, accelerate=a.accelerate, attention=a.attention) only = set(a.only.split(',')) if a.only else None rows = [] for i, prompt in enumerate(EVALUATION): key = f'{i:02d}' if only and key not in only: continue w, h = sizes(i, 'evaluation') latest = {} def cb(p, step, t, kw): if step == p._num_timesteps - 1: latest['latents'] = kw['latents'].detach() return kw torch.cuda.synchronize(); start = time.perf_counter() im = pipe(prompt=prompt, width=w, height=h, generator=torch.Generator('cuda').manual_seed(20000 + i), callback_on_step_end=cb).images[0] torch.cuda.synchronize(); seconds = time.perf_counter() - start im.save(out / f'{key}.png') save_file({'latents': latest['latents'].cpu().contiguous()}, str(out / f'{key}-latents.safetensors')) rows.append({'id': key, 'prompt': prompt, 'seed': 20000 + i, 'width': w, 'height': h, 'steps': pipe._num_timesteps, 'seconds_including_first_use': seconds}) print(json.dumps({'event': 'evaluation', 'name': name, 'id': key, 'seconds': seconds}), flush=True) for j, prompt in enumerate(EDITS): key = f'edit-{j}' if only and key not in only: continue idx = [0, 3][j] ref = Image.open(ROOT / 'evaluation' / 'bf16' / f'{idx:02d}.png').resize((1024, 1024)) im = pipe(prompt=prompt, image=ref, width=1024, height=1024, generator=torch.Generator('cuda').manual_seed(21000 + j)).images[0] im.save(out / f'{key}.png') rows.append({'id': key, 'prompt': prompt, 'seed': 21000 + j, 'reference': f'bf16/{idx:02d}.png'}) (out / 'runs.json').write_text(json.dumps({'precision': a.precision, 'accelerated': a.accelerate, 'quant': a.quant, 'torch': torch.__version__, 'gpu': torch.cuda.get_device_name(), 'transformer_dtypes': dict(collections.Counter(str(t.dtype) for t in pipe.transformer.state_dict().values())), 'rows': rows}, indent=2, ensure_ascii=False)) print('EVALUATION_COMPLETE', name, flush=True) if __name__ == '__main__': main()