ProCreations's picture
Release calibrated Image 2.1 Turbo FP8 transformer, accelerated SM120 runtime, quality evidence and real-time demo
ffc54ec verified
Raw History Blame Contribute Delete
4.85 kB
"""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()