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)