"""Activation-aware FP8 smoothing search; no evaluation prompts used for fitting.""" import json,time,shutil,math from pathlib import Path import torch from safetensors.torch import load_file,save_file from diffusers import QwenImage21Transformer2DModel from fp8_runtime import eligible_modules,CalibratedFP8Linear,replace_module,quantize_activation ROOT=Path(__file__).parent OUT=ROOT/'release'/'transformer'; OUT.mkdir(parents=True,exist_ok=True) BASE=Path('/home/user/models/qwen-image-2.1-b3179ad') @torch.inference_mode() def main(): model=QwenImage21Transformer2DModel.from_pretrained(BASE,subfolder='transformer',torch_dtype=torch.bfloat16).to('cuda').eval() modules=eligible_modules(model) results={}; quantized={}; held={} for i,(name,layer) in enumerate(modules.items()): x=load_file(str(ROOT/'calibration'/(name+'.safetensors')))['activations'].to('cuda') # Disjoint selection and diagnostic rows from every calibration trajectory. groups=torch.arange(len(x)//4,device=x.device) train=x[groups*4+groups%4].contiguous() check=x[groups*4+(groups+1)%4].contiguous() w=layer.weight.float() amax=train.float().abs().amax(0).clamp_min(1e-5) wmax=w.abs().amax(0).clamp_min(1e-5) ref=torch.nn.functional.linear(train,layer.weight).float() denom=ref.square().mean().clamp_min(1e-12) candidates=[]; best=None for alpha in [0.,.25,.5,.75,1.]: smooth=torch.ones_like(amax) if alpha==0 else (amax.pow(alpha)/wmax.pow(1-alpha)).clamp(.01,100.) # A global factor is immaterial with per-row quantization. Normalize to # keep multipliers near unity, preserving channel-to-channel smoothing. smooth=smooth/smooth.log().mean().exp() ws=w*smooth[None,:] a,asc=quantize_activation(train,smooth) for clip in [1.,.995,.99]: scales=(ws.abs().amax(1,keepdim=True)*clip/448).clamp_min(1e-12) q=(ws/scales).clamp(-448,448).to(torch.float8_e4m3fn) y=torch._scaled_mm(a,q.T,scale_a=asc,scale_b=scales.T,out_dtype=torch.bfloat16,use_fast_accum=False).float() nmse=((y-ref).square().mean()/denom).item() item={'alpha':alpha,'weight_amax_clip_factor':clip,'selection_nmse':nmse} candidates.append(item) if best is None or nmse.0025 results[name]={**item,'diagnostic_nmse':nmse,'diagnostic_nrmse':math.sqrt(nmse), 'shape':list(q.shape),'calibration_rows':len(x),'fit_rows':len(train), 'diagnostic_rows':len(check),'kept_bf16':preserve,'candidates':candidates} if not preserve: quantized[name]={'shape':list(q.shape),'alpha':item['alpha'], 'weight_amax_clip_factor':item['weight_amax_clip_factor']} replace_module(model,name,CalibratedFP8Linear(q,scales,smooth)) else: held[name]=results[name] print(json.dumps({'event':'quantize','layer':name,'done':i+1,'total':len(modules),'nrmse':math.sqrt(nmse),'alpha':item['alpha'],'bf16':preserve}),flush=True) del x,train,check,w,ref,ref_check,pred,ws,a,asc,y model.save_config(OUT) config={'format':'procreations-calibrated-fp8-v1','base_model':'Qwen/Qwen-Image-2.1', 'base_revision':'b3179ad355be050328e483a9dfdd9e60cd62adfa', 'weight_dtype':'float8_e4m3fn','activation_dtype':'float8_e4m3fn', 'weight_scaling':'per-output-channel FP32','activation_scaling':'dynamic per-token FP32', 'accumulation':'FP32, use_fast_accum=False','output_dtype':'bfloat16', 'kernel':'torch._scaled_mm native CUTLASS SM120, Triton fused smoothing and activation quantization', 'quantized_modules':quantized,'kept_bf16_modules':list(held), 'modified_notice':'Built with Qwen. ProCreations modified transformer weights through calibrated FP8 quantization on 2026-09-20.'} (OUT/'quantization_config.json').write_text(json.dumps(config,indent=2)) report=ROOT/'release'/'reports';report.mkdir(exist_ok=True) (report/'calibration_search.json').write_text(json.dumps(results,indent=2)) shutil.copy(ROOT/'calibration'/'manifest.json',report/'calibration_manifest.json') state=model.state_dict(); shards=[];shard={};size=0 for n,v in state.items(): nb=v.numel()*v.element_size() if shard and size+nb>3_500_000_000: shards.append(shard);shard={};size=0 shard[n]=v.detach().cpu().contiguous();size+=nb if shard:shards.append(shard) total=sum(v.numel()*v.element_size() for v in state.values()) for i,s in enumerate(shards): save_file(s,str(OUT/f'model-{i+1:05d}-of-{len(shards):05d}.safetensors'),metadata={'format':'pt','notice':config['modified_notice']}) (report/'quantization_summary.json').write_text(json.dumps({'quantized_modules':len(quantized),'bf16_outliers':list(held),'transformer_bytes':total, 'max_diagnostic_nrmse':max(v['diagnostic_nrmse'] for v in results.values()), 'mean_diagnostic_nrmse':sum(v['diagnostic_nrmse'] for v in results.values())/len(results)},indent=2)) print('QUANTIZATION_COMPLETE',total,flush=True) if __name__=='__main__':main()