File size: 1,518 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
import os
import torch, json, random
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))
for l in list(open("data/rendered/hermes_glaive.eval.jsonl"))[:20]: samples.append(json.loads(l)["text"])
ids=[tok(s,return_tensors="pt").input_ids[:,:1024].cuda() for s in samples]
L=len(lm.layers); keep=set(range(L))
def hook(i):
    def h(m,args,kw,out):
        if i not in keep: return args[0] if args else kw["hidden_states"]
    return h
for i in range(L): lm.layers[i].register_forward_hook(hook(i),with_kwargs=True)
@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
cands={"A":[0,7,14,15,17,23],"B":[0,15,16,19,22,23],"C":[0,7,14,15,16,23],"D":[0,11,14,15,17,23],"E":[0,3,14,15,17,23],"F":[0,7,10,15,17,23],"G":[0,7,14,19,22,23],"H":[0,11,14,19,22,23],"I":[0,3,6,7,22,23],"J":[0,1,2,3,22,23],"8L":[0,7,14,15,16,19,22,23],"4L":[0,15,22,23]}
for k,v in cands.items():
    keep=set(v); print(k,v,"loss %.3f"%loss_all(),flush=True)