small-test / src /prune_model.py
Serveurperso's picture
Serveurperso HF Staff
small-test: a 95M multimodal fixture for the llama.cpp server CI
4a393d1
Raw History Blame Contribute Delete
8.23 kB
import os
# build the pruned Qwen3.5 model: layer subset, GDN head pruning, FFN pruning, vocab pruning, vision block subset
import torch, json, os, sys, random, argparse, shutil, copy
import numpy as np
from safetensors.torch import save_file
from transformers import AutoTokenizer, Qwen3_5ForConditionalGeneration
from datasets import load_from_disk
from prune_tokenizer import prune_tokenizer
ap=argparse.ArgumentParser()
ap.add_argument("--layers",default="0,11,14,19,22,23")
ap.add_argument("--ffn",type=int,default=1024)
ap.add_argument("--lin-heads",type=int,default=8)
ap.add_argument("--min-count",type=int,default=20)
ap.add_argument("--vision-blocks",default="0,1,2,3,4,5")
ap.add_argument("--out",default="pruned")
ap.add_argument("--src",default=os.environ.get("PARENT","Qwen/Qwen3.5-0.8B"))
args=ap.parse_args()
P=args.src
keep_layers=[int(x) for x in args.layers.split(",")]
vis_blocks=[int(x) for x in args.vision_blocks.split(",")]
tok=AutoTokenizer.from_pretrained(P)
model=Qwen3_5ForConditionalGeneration.from_pretrained(P, dtype=torch.bfloat16).cuda().eval()
lm=model.model.language_model
cfg=model.config
tcfg=cfg.text_config
types=[tcfg.layer_types[i] for i in keep_layers]
# llama.cpp derives layer types from full_attention_interval, so the pattern must be uniform
interval=types.index("full_attention")+1
assert all((t=="full_attention")==((i+1)%interval==0) for i,t in enumerate(types)), types
print("kept layers",keep_layers,types,"interval",interval)
# ---- calibration data
random.seed(0); samples=[]
ds=load_from_disk("data/smol_smoltalk")["train"]
for i in random.sample(range(len(ds)),48): samples.append(tok.apply_chat_template(ds[i]["messages"],tokenize=False))
for f in ["apigen","hermes_glaive"]:
for l in list(open(f"data/rendered/{f}.eval.jsonl"))[:24]: samples.append(json.loads(l)["text"])
ids=[tok(s,return_tensors="pt").input_ids[:,:1024].cuda() for s in samples]
# ---- importance hooks: FFN neuron activation, GDN head output contribution
ffn_score={}; head_score={}
def ffn_hook(i):
def h(m,args):
x=args[0].float().abs().mean(dim=(0,1))
ffn_score[i]=ffn_score.get(i,0)+x
return h
def head_hook(i):
def h(m,args):
x=args[0].float() # [B,T,H*dv]
H=lm.layers[i].linear_attn.num_v_heads; dv=lm.layers[i].linear_attn.head_v_dim
W=m.weight.float() # [hidden, H*dv]
xs=x.view(-1,H,dv)
Ws=W.view(W.shape[0],H,dv)
contrib=torch.einsum("nhd,ohd->nho",xs,Ws).norm(dim=-1).mean(0) # [H]
head_score[i]=head_score.get(i,0)+contrib
return h
hs=[]
for i in keep_layers:
hs.append(lm.layers[i].mlp.down_proj.register_forward_pre_hook(ffn_hook(i)))
if tcfg.layer_types[i]=="linear_attention":
hs.append(lm.layers[i].linear_attn.out_proj.register_forward_pre_hook(head_hook(i)))
with torch.no_grad():
for x in ids: model(input_ids=x,use_cache=False)
for h in hs: h.remove()
sd={k:v.cpu() for k,v in model.state_dict().items()}
# safetensors file also has mtp.* which HF drops; read them raw
from safetensors import safe_open
with safe_open(os.path.join(P,"model.safetensors-00001-of-00001.safetensors"),"pt") as f:
for k in f.keys():
if k.startswith("mtp."): sd[k]=f.get_tensor(k)
# ---- tokenizer / vocab
old2new,n_vocab=prune_tokenizer(args.min_count,args.out)
V=(n_vocab+63)//64*64
old_ids=torch.tensor(sorted(old2new,key=lambda k:old2new[k]))
def map_id(i): return old2new[i]
new={}
def ffn_prune(prefix_in,prefix_out,score):
F=args.ffn
idx=torch.topk(score,F).indices.sort().values.cpu()
for n in ["gate_proj","up_proj"]: new[f"{prefix_out}.mlp.{n}.weight"]=sd[f"{prefix_in}.mlp.{n}.weight"][idx].contiguous()
new[f"{prefix_out}.mlp.down_proj.weight"]=sd[f"{prefix_in}.mlp.down_proj.weight"][:,idx].contiguous()
for new_i,old_i in enumerate(keep_layers):
pi=f"model.language_model.layers.{old_i}"; po=f"model.language_model.layers.{new_i}"
for n in ["input_layernorm","post_attention_layernorm"]: new[f"{po}.{n}.weight"]=sd[f"{pi}.{n}.weight"]
ffn_prune(pi,po,ffn_score[old_i])
if tcfg.layer_types[old_i]=="full_attention":
for k in sd:
if k.startswith(pi+".self_attn."): new[po+k[len(pi):]]=sd[k]
else:
H=tcfg.linear_num_value_heads; dk=tcfg.linear_key_head_dim; dv=tcfg.linear_value_head_dim
heads=torch.topk(head_score[old_i],args.lin_heads).indices.sort().values.cpu()
kd=H*dk
qi=torch.cat([torch.arange(h*dk,(h+1)*dk) for h in heads]); vi=torch.cat([torch.arange(h*dv,(h+1)*dv) for h in heads])
qkv_idx=torch.cat([qi, kd+qi, 2*kd+vi])
la=pi+".linear_attn."; lo=po+".linear_attn."
new[lo+"in_proj_qkv.weight"]=sd[la+"in_proj_qkv.weight"][qkv_idx].contiguous()
new[lo+"conv1d.weight"]=sd[la+"conv1d.weight"][qkv_idx].contiguous()
new[lo+"in_proj_z.weight"]=sd[la+"in_proj_z.weight"][vi].contiguous()
new[lo+"in_proj_a.weight"]=sd[la+"in_proj_a.weight"][heads].contiguous()
new[lo+"in_proj_b.weight"]=sd[la+"in_proj_b.weight"][heads].contiguous()
new[lo+"A_log"]=sd[la+"A_log"][heads].contiguous()
new[lo+"dt_bias"]=sd[la+"dt_bias"][heads].contiguous()
new[lo+"norm.weight"]=sd[la+"norm.weight"]
new[lo+"out_proj.weight"]=sd[la+"out_proj.weight"][:,vi].contiguous()
print(f"layer {old_i}: kept heads {heads.tolist()}")
new["model.language_model.norm.weight"]=sd["model.language_model.norm.weight"]
emb=torch.zeros(V,tcfg.hidden_size,dtype=sd["model.language_model.embed_tokens.weight"].dtype)
emb[:len(old_ids)]=sd["model.language_model.embed_tokens.weight"][old_ids]
emb[len(old_ids):]=emb[:len(old_ids)].float().mean(0).to(emb.dtype)
new["model.language_model.embed_tokens.weight"]=emb
# mtp: weight-norm heuristic for FFN (block is not run by HF)
g=sd["mtp.layers.0.mlp.gate_proj.weight"].float(); u=sd["mtp.layers.0.mlp.up_proj.weight"].float(); d=sd["mtp.layers.0.mlp.down_proj.weight"].float()
ffn_prune("mtp.layers.0","mtp.layers.0",(g.norm(dim=1)*u.norm(dim=1)*d.norm(dim=0)).cuda())
for k in sd:
if k.startswith("mtp.") and ".mlp." not in k: new[k]=sd[k]
# vision
for new_i,old_i in enumerate(vis_blocks):
for k in sd:
pi=f"model.visual.blocks.{old_i}."
if k.startswith(pi): new[f"model.visual.blocks.{new_i}."+k[len(pi):]]=sd[k]
for k in sd:
if k.startswith("model.visual.") and ".blocks." not in k: new[k]=sd[k]
new={k:v.contiguous() for k,v in new.items()}
n_text=sum(v.numel() for k,v in new.items() if k.startswith("model.language_model."))
n_mtp=sum(v.numel() for k,v in new.items() if k.startswith("mtp."))
n_vis=sum(v.numel() for k,v in new.items() if k.startswith("model.visual."))
print(f"params: text {n_text/1e6:.2f}M + mtp {n_mtp/1e6:.2f}M = {(n_text+n_mtp)/1e6:.2f}M ; vision {n_vis/1e6:.2f}M")
# config
c=json.load(open(os.path.join(P,"config.json")))
t=c["text_config"]
t["num_hidden_layers"]=len(keep_layers); t["layer_types"]=types; t["full_attention_interval"]=interval
t["intermediate_size"]=args.ffn; t["linear_num_key_heads"]=args.lin_heads; t["linear_num_value_heads"]=args.lin_heads
t["vocab_size"]=V; t["eos_token_id"]=map_id(t["eos_token_id"])
for k in ["image_token_id","video_token_id","vision_end_token_id","vision_start_token_id"]: c[k]=map_id(c[k])
c["vision_config"]["depth"]=len(vis_blocks)
c["transformers_version"]="5.16.1"
os.makedirs(args.out,exist_ok=True)
json.dump(c,open(os.path.join(args.out,"config.json"),"w"),indent=2)
gen={"bos_token_id":map_id(tok.convert_tokens_to_ids("<|endoftext|>")),"eos_token_id":[map_id(tok.convert_tokens_to_ids("<|im_end|>")),map_id(tok.convert_tokens_to_ids("<|endoftext|>"))],"pad_token_id":map_id(tok.convert_tokens_to_ids("<|endoftext|>")),"do_sample":False}
json.dump(gen,open(os.path.join(args.out,"generation_config.json"),"w"),indent=2)
for f in ["preprocessor_config.json","video_preprocessor_config.json"]: shutil.copy(os.path.join(P,f),args.out)
save_file(new,os.path.join(args.out,"model.safetensors"),metadata={"format":"pt"})
json.dump({"layers":keep_layers,"ffn":args.ffn,"lin_heads":args.lin_heads,"min_count":args.min_count,"vocab":V,"vision_blocks":vis_blocks},open(os.path.join(args.out,"prune_info.json"),"w"),indent=1)
print("saved to",args.out)