File size: 5,830 Bytes
f642d1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
"""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()