"""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()