small-test / src /layer_analysis.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
2.5 kB
import os
import torch, json, random, sys, time
from transformers import AutoTokenizer, Qwen3_5ForConditionalGeneration
from datasets import load_from_disk
P=os.environ.get("PARENT","Qwen/Qwen3.5-0.8B")
tok=AutoTokenizer.from_pretrained(P)
model=Qwen3_5ForConditionalGeneration.from_pretrained(P, dtype=torch.bfloat16).cuda().eval()
lm=model.model.language_model
random.seed(0)
samples=[]
ds=load_from_disk("data/smol_smoltalk")["train"]
for i in random.sample(range(len(ds)),40):
samples.append(tok.apply_chat_template(ds[i]["messages"],tokenize=False))
ds=load_from_disk("data/apigen")["train"]
for i in random.sample(range(len(ds)),20):
r=ds[i]
try: tools=json.loads(r["tools"]); ans=json.loads(r["answers"])
except Exception: continue
msgs=[{"role":"user","content":r["query"]},{"role":"assistant","content":"","tool_calls":[{"type":"function","function":a} for a in ans]}]
samples.append(tok.apply_chat_template(msgs,tools=tools,tokenize=False))
ids=[tok(s,return_tensors="pt").input_ids[:,:1024].cuda() for s in samples]
print("samples",len(ids),"tokens",sum(x.shape[1] for x in ids))
L=len(lm.layers)
# per-layer angular distance via hooks
dist=[[] for _ in range(L)]
def mk(i):
def hook(m,args,kw,out):
x=args[0] if args else kw["hidden_states"]; y=out
c=torch.nn.functional.cosine_similarity(x.float(),y.float(),dim=-1).clamp(-1,1)
dist[i].append((torch.arccos(c)/3.14159).mean().item())
return hook
hs=[lm.layers[i].register_forward_hook(mk(i),with_kwargs=True) for i in range(L)]
@torch.no_grad()
def loss_all():
tot=0;n=0
for x in ids:
out=model(input_ids=x,labels=x,use_cache=False); tot+=out.loss.item()*x.shape[1]; n+=x.shape[1]
return tot/n
base=loss_all()
for h in hs: h.remove()
print("base loss %.4f"%base)
for i in range(L): print("layer %2d %s angdist %.4f"%(i,lm.config.layer_types[i][:4],sum(dist[i])/len(dist[i])))
# leave-one-out: skip layer i
skip=set()
def skiphook(i):
def hook(m,args,kw,out):
if i in skip: return args[0] if args else kw["hidden_states"]
return hook
hs=[lm.layers[i].register_forward_hook(skiphook(i),with_kwargs=True) for i in range(L)]
res=[]
for i in range(L):
skip={i}; l=loss_all(); res.append(l); print("skip layer %2d %s loss %.4f (+%.4f)"%(i,lm.config.layer_types[i][:4],l,l-base),flush=True)
skip=set()
json.dump({"base":base,"angdist":[sum(d)/len(d) for d in dist],"loo":res,"types":lm.config.layer_types},open("layer_analysis.json","w"),indent=1)