File size: 4,971 Bytes
1961af5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Reproducible BF16 -> calibrated NVFP4 with rank-128 BF16 correction."""
import os,json,time,shutil,statistics
from pathlib import Path
import torch
from safetensors.torch import load_file,save_file
from probe_dynamic import BASE,OLD,INDEX,weight,split,prepare,run,error
ROOT=Path(os.environ['IMAGE21_WORK_DIR'])
OUT=ROOT/'release'/'transformer';OUT.mkdir(parents=True,exist_ok=True)
REPORT=ROOT/'release'/'reports';REPORT.mkdir(exist_ok=True)
CACHE=ROOT/'calibrated-layers-dynamic';CACHE.mkdir(exist_ok=True)

@torch.inference_mode()
def main():
 names=json.loads((OLD/'calibration/manifest.json').read_text())['modules'];results={};specs={}
 for i,name in enumerate(names):
  record=CACHE/(name+'.json');saved=CACHE/(name+'.safetensors')
  if record.exists() and saved.exists():result=json.loads(record.read_text())
  else:
   w=weight(name);x=load_file(str(OLD/'calibration'/f'{name}.safetensors'))['activations'].cuda();fit,check=split(x)
   rf=fit@w.T;rc=check@w.T
   naive=prepare(w,fit,0.,0);naive_error=error(run(check,naive),rc)
   best=None;candidates=[]
   for smoothing in [0.,.25,.5,.75,1.]:
    for mse in [False,True]:
     p=prepare(w,fit,smoothing,128,method='svd',mse=mse)
     err=error(run(fit,p),rf)
     item={'smoothing_alpha':smoothing,'weight_scale_search':mse,'fit_nrmse':err};candidates.append(item)
     if best is None or err<best[0]:best=(err,{k:v.clone() for k,v in p.items()},item)
   err,p,item=best;check_error=error(run(check,p),rc)
   result={**item,'diagnostic_nrmse':check_error,'naive_nvfp4_diagnostic_nrmse':naive_error,'rank':128,'shape':list(w.shape),'sample_rows':len(x),'fit_rows':len(fit),'diagnostic_rows':len(check),'candidates':candidates}
   assert check_error<.2 and all(torch.isfinite(v.float()).all() for v in p.values())
   save_file({k:v.cpu().contiguous() for k,v in p.items()},str(saved));record.write_text(json.dumps(result,indent=2))
  results[name]=result;specs[name]={k:result[k] for k in ['shape','rank','smoothing_alpha','weight_scale_search']}
  print(json.dumps({'event':'quantize','done':i+1,'total':len(names),'name':name,'nrmse':result['diagnostic_nrmse'],'naive':result['naive_nvfp4_diagnostic_nrmse']}),flush=True)
 state={}
 for shard in sorted(set(INDEX.values())):
  for name,v in load_file(str(BASE/shard)).items():
   if name.removesuffix('.weight') not in specs:state[name]=v
 for name in names:
  for k,v in load_file(str(CACHE/(name+'.safetensors'))).items():state[name+'.'+k]=v
 shutil.copy(BASE/'config.json',OUT/'config.json')
 cfg={'format':'procreations-calibrated-nvfp4-svd-v1','base_model':'Qwen/Qwen-Image-2.1','base_revision':'b3179ad355be050328e483a9dfdd9e60cd62adfa','quantized_modules':specs,
 'weight_dtype':'NVFP4 E2M1 packed uint8','activation_dtype':'NVFP4 E2M1','scale_format':'E4M3 per 16 elements, 128x4 swizzle, FP32 global multiplier','low_rank_correction':'BF16 rank128, native fused FP32 accumulator epilogue','activation_scale':'dynamic tensor-wide activation amax at every call; dynamic per16 block scales; no static activation clipping','flashinfer_revision':'975f90583d9ac8896db14cf0f26e99a853c2f136','kernel':'SM120 b12x CuTe DSL native block-scaled FP4; BF16 rank correction','modified_notice':'Built with Qwen. ProCreations modified transformer weights with calibrated NVFP4 quantization, 2026-09-20.'}
 (OUT/'quantization_config.json').write_text(json.dumps(cfg,indent=2))
 shards=[];s={};size=0
 for n,v in state.items():
  nb=v.numel()*v.element_size()
  if s and size+nb>3_000_000_000:shards.append(s);s={};size=0
  s[n]=v.contiguous();size+=nb
 if s:shards.append(s)
 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':cfg['modified_notice']})
 total=sum(v.numel()*v.element_size() for v in state.values())
 summary={'modules':len(names),'transformer_bytes':total,'mean_diagnostic_nrmse':statistics.mean(r['diagnostic_nrmse'] for r in results.values()),'max_diagnostic_nrmse':max(r['diagnostic_nrmse'] for r in results.values()),'naive_mean_diagnostic_nrmse':statistics.mean(r['naive_nvfp4_diagnostic_nrmse'] for r in results.values()),'protocol':'64 original BF16 trajectories;384 fit and384 disjoint diagnostic rows per layer, stratified across prompts and timesteps. Five smoothing exponents, two weight-scale methods. Rank128 randomized SVD (seed1234,q160,niter3); optional activation-weighted15-point block-scale MSE search. Native W4A4 kernels with dynamic activation amax used for all calibrated selection/diagnostic output errors. Plain reference uses original static range. No held-out image evaluation prompts used for calibration.'}
 (REPORT/'calibration_search.json').write_text(json.dumps(results,indent=2));(REPORT/'quantization_summary.json').write_text(json.dumps(summary,indent=2));shutil.copy(OLD/'calibration/manifest.json',REPORT/'calibration_manifest.json')
 print('QUANTIZATION_COMPLETE',json.dumps(summary),flush=True)
if __name__=='__main__':main()