ProCreations's picture
Release calibrated Image2.1 FP8 transformer, native SM120 runtime, quality evidence and real-time demo
f642d1f verified
Raw History Blame Contribute Delete
5.85 kB
import argparse,json,time,statistics,collections,platform,subprocess
from pathlib import Path
import torch
from PIL import Image
from safetensors.torch import save_file
from fp8_runtime import load_pipeline,CalibratedFP8Linear
from prompts import EVALUATION,EDIT_CALIBRATION
ROOT=Path(__file__).parent
@torch.inference_mode()
def main():
mode=argparse.ArgumentParser();mode.add_argument('precision',choices=['bf16','fp8']);args=mode.parse_args()
name=args.precision
out=ROOT/'evaluation'/name;out.mkdir(parents=True,exist_ok=True)
pipe=load_pipeline('/home/user/models/qwen-image-2.1-b3179ad', ROOT/'release'/'transformer' if name=='fp8' else None)
fps=[m for m in pipe.transformer.modules() if isinstance(m,CalibratedFP8Linear)]
if name=='fp8':
assert len(fps)>0
assert all(m.weight.dtype==torch.float8_e4m3fn and m.weight_scale.dtype==torch.float32 and m.smooth.dtype==torch.float32 for m in fps)
evidence={'torch':torch.__version__,'gpu':torch.cuda.get_device_name(),'capability':torch.cuda.get_device_capability(),
'fp8_linears':len(fps),'transformer_dtypes':dict(collections.Counter(str(t.dtype) for t in pipe.transformer.state_dict().values())),
'resident_gb':torch.cuda.memory_allocated()/1e9,'precision':name}
rows=[]
for i,prompt in enumerate(EVALUATION):
w=h=2048 if i%4==0 else 1024
if i==12:w,h=1536,864
if i==13:w,h=864,1536
latest={}
def cb(p,step,t,kw):
if step==39:latest['latents']=kw['latents'].detach()
return kw
torch.cuda.reset_peak_memory_stats();torch.cuda.synchronize();start=time.perf_counter()
im=pipe(prompt=prompt,width=w,height=h,num_inference_steps=40,
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'{i:02d}.png')
save_file({'latents':latest['latents'].cpu().contiguous()},str(out/f'{i:02d}-latents.safetensors'))
row={'index':i,'prompt':prompt,'seed':20000+i,'width':w,'height':h,'steps':40,'seconds':seconds,'peak_gb':torch.cuda.max_memory_allocated()/1e9}
rows.append(row);(out/'quality_runs.json').write_text(json.dumps(rows,indent=2,ensure_ascii=False))
print(json.dumps({'event':'evaluation','precision':name,'done':i+1,'seconds':seconds}),flush=True)
# Paired editing checks use the same BF16 reference image in both precisions.
for j,prompt in enumerate(['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.']):
idx=[0,3][j];im=Image.open(ROOT/'evaluation'/'bf16'/f'{idx:02d}.png').resize((1024,1024))
result=pipe(prompt=prompt,image=im,width=1024,height=1024,num_inference_steps=40,generator=torch.Generator('cuda').manual_seed(21000+j)).images[0]
result.save(out/f'edit-{j}.png')
timing={}
for size,n in [(1024,5),(2048,3)]:
# Resolution-specific warmup is excluded. Prompt encoding and VAE decode
# are included in each end-to-end measurement; no cached prompt embeds.
timings=[]
for j in range(n+1):
torch.cuda.synchronize();start=time.perf_counter()
pipe(prompt=EVALUATION[15],width=size,height=size,num_inference_steps=40,
generator=torch.Generator('cuda').manual_seed(30000+j))
torch.cuda.synchronize();seconds=time.perf_counter()-start
if j:timings.append(seconds)
print(json.dumps({'event':'benchmark','precision':name,'size':size,'iteration':j,'seconds':seconds,'warmup':j==0}),flush=True)
timing[str(size)]={'seconds':timings,'mean':statistics.mean(timings),'median':statistics.median(timings),'min':min(timings),'max':max(timings)}
evidence.update(timing=timing,protocol='Batch1;40steps;CFG1;prefixKVcache;fully GPU resident;warmup excluded;CUDA synchronized;includes text encoding, denoising and VAE decode;excludes load and PNG write.')
(out/'benchmark.json').write_text(json.dumps(evidence,indent=2))
if name=='fp8':
# Capture an actual denoising invocation after warmup, keeping trace bounded.
profiler=torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU,torch.profiler.ProfilerActivity.CUDA],record_shapes=True)
counter={'n':0}
def before(m,args,kwargs):
if counter['n']==10:profiler.start()
def after(m,args,kwargs,result):
if counter['n']==10:torch.cuda.synchronize();profiler.stop()
counter['n']+=1
a=pipe.transformer.register_forward_pre_hook(before,with_kwargs=True)
b=pipe.transformer.register_forward_hook(after,with_kwargs=True)
pipe(prompt=EVALUATION[0],width=1024,height=1024,num_inference_steps=40,generator=torch.Generator('cuda').manual_seed(40000))
a.remove();b.remove()
profiler.export_chrome_trace(str(out/'denoising-trace.json'))
kernels=collections.Counter(e.name for e in profiler.events() if e.device_type==torch.autograd.DeviceType.CUDA)
scaled=[e for e in profiler.events() if e.name=='aten::_scaled_mm']
proof={'scaled_mm_calls_in_one_denoising_step':len(scaled),'expected_fp8_linears':len(fps),'kernels':dict(kernels),
'native_fp8_sm120':any('Sm120' in n and 'float_e4m3' in n for n in kernels),
'fp8_kernel_launches':sum(c for n,c in kernels.items() if 'Sm120' in n and 'float_e4m3' in n)}
(out/'kernel_evidence.json').write_text(json.dumps(proof,indent=2))
assert proof['native_fp8_sm120'] and len(scaled)==len(fps),proof
print('EVALUATION_COMPLETE',name,flush=True)
if __name__=='__main__':main()