ProCreations's picture
Accelerate full 40-step FP8 generation with native precision, measured quality and real-time demo
1081be0 verified
Raw History Blame Contribute Delete
3.46 kB
import sys,time,json,statistics,collections,argparse
from pathlib import Path
import torch
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from fp8_runtime import load_pipeline,CalibratedFP8Linear
from acceleration import accelerate_pipeline
from prompts import EVALUATION
ap=argparse.ArgumentParser();ap.add_argument('--baseline',action='store_true');a=ap.parse_args()
out=Path(__file__).parent/('final-baseline' if a.baseline else 'final-optimized');out.mkdir(exist_ok=True)
@torch.inference_mode()
def main():
start=time.perf_counter();p=load_pipeline('/home/user/models/qwen-image-2.1-b3179ad',ROOT/'release/transformer');load_sec=time.perf_counter()-start
if not a.baseline:accelerate_pipeline(p)
fps=[m for m in p.transformer.modules() if isinstance(m,CalibratedFP8Linear)]
assert len(fps)==224 and all(m.weight.dtype==torch.float8_e4m3fn and m.smooth.dtype==torch.float32 and m.weight_scale.dtype==torch.float32 for m in fps)
result={'load_seconds':load_sec,'torch':torch.__version__,'gpu':torch.cuda.get_device_name(),'fp8_linears':len(fps),'steps':40,'cfg':1,'extra_quantization':False,'approximate_cache':False,'timing':{},'protocol':'CUDA synchronized; batch1; full40steps; includes encoder, denoising and VAE; excludes model load, resolution warmup and file writes. Weights and prefixKVcache unchanged. Compiled mode emulates intermediate precision casts.'}
for size in [1024,2048]:
vals=[];n=2 if a.baseline else (5 if size==1024 else 3)
for j in range(n+1):
torch.cuda.reset_peak_memory_stats();torch.cuda.synchronize();t=time.perf_counter()
im=p(prompt=EVALUATION[15],width=size,height=size,num_inference_steps=40,generator=torch.Generator('cuda').manual_seed(30000+j)).images[0]
torch.cuda.synchronize();sec=time.perf_counter()-t
if j:vals.append(sec)
else:warm=sec
print(json.dumps({'size':size,'i':j,'seconds':sec,'baseline':a.baseline}),flush=True)
result['timing'][str(size)]={'seconds':vals,'mean':statistics.mean(vals),'warmup_seconds':warm,'peak_gb':torch.cuda.max_memory_allocated()/1e9}
im.save(out/f'{size}.png');(out/'benchmark.json').write_text(json.dumps(result,indent=2))
if not a.baseline:
profiler=torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU,torch.profiler.ProfilerActivity.CUDA])
count={'n':0}
def before(m,args,kw):
if count['n']==10:profiler.start()
def after(m,args,kw,r):
if count['n']==10:torch.cuda.synchronize();profiler.stop()
count['n']+=1
h=p.transformer.register_forward_pre_hook(before,with_kwargs=True);g=p.transformer.register_forward_hook(after,with_kwargs=True)
p(prompt=EVALUATION[15],width=1024,height=1024,num_inference_steps=40,generator=torch.Generator('cuda').manual_seed(40000));h.remove();g.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)
proof={'fp8_sm120_launches':sum(v for k,v in kernels.items() if 'Sm120' in k and 'float_e4m3' in k),'native_bf16_flash_attention':sum(v for k,v in kernels.items() if 'flash_fwd_kernel' in k and 'bfloat16' in k),'all_kernels':dict(kernels),'transformer_calls':count['n'],'full_denoising_steps':40}
(out/'kernel_evidence.json').write_text(json.dumps(proof,indent=2));assert proof['fp8_sm120_launches']==224 and proof['native_bf16_flash_attention']==32 and count['n']==40,proof
print('BENCHMARK_COMPLETE',flush=True)
main()