Download code/bench_v5_sam.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 1.63 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/bench_v5_sam.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/bench_v5_sam.py
-
curl -L -o bench_v5_sam.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/bench_v5_sam.py
1.63 kB
| """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() | |