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