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)