File size: 1,632 Bytes
5fa4f6b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""V5Engine + SAM smoke test at moderate context on the A40."""
import sys, time
import torch
import tiktoken

from config import Config
from model import build_model
from inference import V5Engine
from sam import build_from_ids


def main():
    ckpt_path = sys.argv[1] if len(sys.argv) > 1 else \
        "/root/checkpoints/spec_v7/best.pt"
    cfg = Config.v5_500m()
    cfg.max_seq_len = min(cfg.max_seq_len, 1 << 20)  # rope table size
    model = build_model(cfg, "cuda")
    ckpt = torch.load(ckpt_path, map_location="cuda")
    model.load_state_dict(ckpt["model"], strict=False)
    model.init_compressor_from_attn()
    model.eval()
    print(f"loaded {ckpt_path} (step {ckpt.get('step')})")

    enc = tiktoken.get_encoding("gpt2")

    from data_pipeline import load_tokens
    toks = load_tokens("mixture500m")
    sam = build_from_ids(toks[:1_000_000].tolist())
    print(f"SAM: {len(sam.text):,} tokens")

    eng = V5Engine(model, "cuda")
    eng.sam = sam

    prompt = enc.encode_ordinary(
        "The history of the Roman Empire begins in 753 BC with the "
        "founding of Rome. The Roman Republic was established in 509 BC.")

    for sam_ext in [0, 2048, 16384]:
        t0 = time.perf_counter()
        r = eng.generate_with_sam(prompt, 30000, sam_extend=sam_ext,
                                  min_len=2)
        print(f"sam={sam_ext:6d}: {r['tok/s']:>10,.0f} tok/s  "
              f"({r['model_tokens']} model + {r['sam_tokens']} sam, "
              f"{r['steps']} steps)")
        print(f"   text: {enc.decode(r['output'][:70])!r}")
    print("V5SAM_DONE")


if __name__ == "__main__":
    main()