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()