File size: 2,500 Bytes
4a393d1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
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)