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.83 kB
"""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<best[0]:best=(nmse,q.clone(),scales.clone(),smooth.clone(),item)
_,q,scales,smooth,item=best
ref_check=torch.nn.functional.linear(check,layer.weight).float()
a,asc=quantize_activation(check,smooth)
pred=torch._scaled_mm(a,q.T,scale_a=asc,scale_b=scales.T,out_dtype=torch.bfloat16,use_fast_accum=False).float()
nmse=((pred-ref_check).square().mean()/ref_check.square().mean().clamp_min(1e-12)).item()
# Sensitive outliers stay BF16 instead of forcing a precision conversion.
preserve=not math.isfinite(nmse) 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()