patdev's picture
Fix v7 ModelOpt FP8 and AutoQuantize API for 2026 release
d0cff17 verified
Raw
History Blame Contribute Delete
3 kB
from __future__ import annotations
import argparse,json,re,torch
import torch.nn as nn
@torch.no_grad()
def prune_2to4_(model:nn.Module,include:str,exclude:str):
inc,rex=re.compile(include),re.compile(exclude); touched=[]
for name,m in model.named_modules():
if not isinstance(m,nn.Linear) or not inc.search(name) or rex.search(name): continue
w=m.weight.data; k=w.shape[-1]
if k%4: continue
flat=w.reshape(-1,k); g=flat.reshape(-1,k//4,4); idx=g.abs().argsort(dim=-1)[...,:2]
mask=torch.ones_like(g,dtype=torch.bool); mask.scatter_(-1,idx,False); g.mul_(mask); touched.append(name)
return touched
def enforce_2to4_masks(model,masks):
with torch.no_grad():
for n,p in model.named_parameters():
if n in masks:p.mul_(masks[n])
def capture_masks(model): return {n:(p!=0).to(p.dtype) for n,p in model.named_parameters() if p.ndim==2 and (p==0).any()}
def modelopt_fp8(model,forward_loop,include,exclude):
import copy
import modelopt.torch.quantization as mtq
cfg=copy.deepcopy(mtq.FP8_DEFAULT_CFG)
# ModelOpt quant_cfg is an ordered list; later rules override earlier ones.
sensitive=[
'*geo_head*','*skin_head*','*skl_head*','*out_layer*','*mesh*',
'*sparse*','*conv*','*norm*','*router*','*gate*'
]
for pat in sensitive: cfg['quant_cfg'].append({'quantizer_name':pat,'enable':False})
def calibrate(m): forward_loop(m)
mtq.quantize(model, cfg, forward_loop=calibrate)
return model
def modelopt_auto_fp8(model,data_loader,forward_step,loss_func=None,bits=8.0):
import modelopt.torch.quantization as mtq
return mtq.auto_quantize(model,constraints={'effective_bits':bits},quantization_formats=[mtq.FP8_DEFAULT_CFG],data_loader=data_loader,forward_step=forward_step,loss_func=loss_func,disabled_layers=['*geo_head*','*skin_head*','*skl_head*','*out_layer*','*mesh*','*sparse*','*conv*','*norm*','*router*','*gate*'],num_calib_steps=256,num_score_steps=64,verbose=True)
def selective_report(model,include,exclude):
rows=[]
for n,m in model.named_modules():
if isinstance(m,nn.Linear): rows.append({'name':n,'shape':list(m.weight.shape),'fp8_candidate':bool(re.search(include,n) and not re.search(exclude,n))})
return rows
def main():
ap=argparse.ArgumentParser();ap.add_argument('--model',required=True);ap.add_argument('--out',required=True);ap.add_argument('--include',default='(qkv|to_q|to_k|to_v|proj|fc1|fc2|mlp|linear)');ap.add_argument('--exclude',default='(geo_head|skin_head|skl_head|out_layer|mesh|sparse|conv|norm)');ap.add_argument('--sparsity',action='store_true');a=ap.parse_args()
model=torch.load(a.model,map_location='cpu',weights_only=False)
touched=prune_2to4_(model,a.include,a.exclude) if a.sparsity else []
torch.save(model,a.out);open(a.out+'.report.json','w').write(json.dumps({'sparse_modules':touched,'layers':selective_report(model,a.include,a.exclude)},indent=2))
if __name__=='__main__':main()