File size: 1,700 Bytes
a8f07a3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Sampled seq-draft text quality: emit a draft chain, decode it.
python3 eval_seq_text.py <ckpt> [prompt]"""
import sys, torch
from config import Config
from model import build_model
from data_pipeline import load_tokens
import tiktoken

enc = tiktoken.get_encoding("gpt2")
ckpt_path = sys.argv[1]
prompt = sys.argv[2] if len(sys.argv) > 2 else "The history of the Roman Empire"

cfg = Config.v5_500m()
model = build_model(cfg, "cuda")
ckpt = torch.load(ckpt_path, map_location="cuda")
msd = model.state_dict()
sd = {k: v for k, v in ckpt["model"].items()
      if k in msd and msd[k].shape == v.shape}
model.load_state_dict(sd, strict=False)
model.eval()

G = cfg.medusa_cond_group
ids = torch.tensor([enc.encode(prompt)], device="cuda")
with torch.no_grad():
    h = model(ids)
    h_a = h[:, -1]                                    # [1,d]
    pend = (h_a @ model.tok_emb.weight.T).argmax(-1, keepdim=True)
    hist = ids[:, -G:]                                # last G committed toks
    print(f"base argmax next: {enc.decode(pend[0].tolist())!r}")
    # argmax chain
    d1 = model.spec_draft_seq(h_a, prefix_ids=pend, hist_ids=hist, n=32)
    print(f"\nARGMAX draft (32): {enc.decode(d1[0].tolist())!r}")
    # sampled chains
    for i in range(3):
        d2 = model.spec_draft_seq(h_a, prefix_ids=pend, hist_ids=hist, n=32,
                                  sample=True, temperature=0.9, top_p=0.9)
        print(f"SAMPLED draft {i}: {enc.decode(d2[0].tolist())!r}")
    # base reference: sample from base at same anchor
    lg = (h_a @ model.tok_emb.weight.T).float() / 0.9
    b_tok = torch.multinomial(lg.softmax(-1), 1)
    print(f"\nbase sample next: {enc.decode(b_tok[0].tolist())!r}")