"""50M context window benchmark on A6000. Tests two modes: 1. Sliding window (W=4097): generate 1M tokens, measure sustained speed (should be flat) 2. Growing cache: fill KV cache to N tokens, time a generation step at each context size (4097, 32K, 256K, 1M, 5M, 10M, 50M) — shows how speed degrades with context length The growing-cache test fills the KV cache with random data (skipping actual prefill) to isolate the attention cost at each context size. """ import time import torch import tiktoken from config import Config from model import build_model from inference import FastEngine, InferenceEngine enc = tiktoken.get_encoding("gpt2") prompt = enc.encode_ordinary("The quick brown fox jumps over the lazy dog. ") print(f"GPU: {torch.cuda.get_device_name(0)}") print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") print(f"free VRAM: {torch.cuda.mem_get_info()[0]/1e9:.2f} GB") print(f"torch: {torch.__version__}") # ============================================================ # Test 1: Sliding window (W=4097) — sustained speed # ============================================================ print("\n" + "="*60) print("TEST 1: Sliding window (W=4097) — sustained 1M tokens") print("="*60) cfg = Config() # default: d=384, L=1, K=4096, r=1, W=4097 print(f"Config: d={cfg.d_model} L={cfg.n_layers} K={cfg.medusa_heads} r={cfg.medusa_rank} W={cfg.window_size}") print(f" params: {cfg.count_params()['grand_total']/1e6:.1f}M") m = build_model(cfg, "cuda") eng = FastEngine(m) eng.capture_graph() K1 = cfg.medusa_heads + 1 # warmup for _ in range(3): eng.speculative(prompt, 5 * K1) torch.cuda.synchronize() # sustained 1M tokens nt = 1_000_000 t0 = time.perf_counter() eng.speculative(prompt, nt) torch.cuda.synchronize() elapsed = time.perf_counter() - t0 print(f" {nt:,} tokens in {elapsed:.2f}s = {nt/elapsed:,.0f} tok/s") print(f" VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB") del m, eng torch.cuda.empty_cache() # ============================================================ # Test 2: Growing cache — speed vs context size # ============================================================ print("\n" + "="*60) print("TEST 2: Growing cache — speed vs context size (up to 50M)") print("="*60) # For growing cache: window_size=0, max_seq_len=50M cfg2 = Config(window_size=0, max_seq_len=50_000_000) print(f"Config: d={cfg2.d_model} L={cfg2.n_layers} K={cfg2.medusa_heads} r={cfg2.medusa_rank} " f"W=growing max_seq={cfg2.max_seq_len:,}") # Estimate KV cache size kv_per_token = 2 * cfg2.n_kv_heads * cfg2.head_dim * 2 * cfg2.n_layers # bytes kv_50m = 50_000_000 * kv_per_token / 1e9 print(f" KV cache at 50M: {kv_50m:.1f} GB") print(f" est VRAM: {cfg2.estimate_vram_gb():.1f} GB") if cfg2.estimate_vram_gb() > 45: print(" WARNING: est VRAM close to 48GB limit, reducing max_seq_len") cfg2.max_seq_len = 40_000_000 m2 = build_model(cfg2, "cuda") eng2 = InferenceEngine(m2) eng2.alloc_cache() print(f" cache allocated, VRAM: {torch.cuda.memory_allocated()/1e9:.2f} GB") # Fill cache with random data to simulate N processed tokens # Then time a single generation step (K+1 tokens) at that context size context_sizes = [4097, 32768, 262144, 1_048_576, 5_242_880, 10_485_760, 50_033_164] print(f"\n{'context':>12} {'cache(GB)':>10} {'step(ms)':>10} {'tok/s':>12} {'us/tok':>8}") print("-" * 56) for ctx_size in context_sizes: if ctx_size > cfg2.max_seq_len: print(f"{ctx_size:12,} --- exceeds max_seq_len ---") continue # Fill KV cache with random data up to ctx_size for layer_k, layer_v in zip(eng2._cache_k, eng2._cache_v): layer_k[:, :ctx_size].normal_(0, 0.1) layer_v[:, :ctx_size].normal_(0, 0.1) # Prepare input tokens input_ids = torch.randint(0, cfg2.vocab_size, (1, K1), device="cuda") # Warmup try: for _ in range(2): with torch.no_grad(): h = eng2.model(input_ids, eng2._cache_k, eng2._cache_v, start_pos=ctx_size) torch.cuda.synchronize() except torch.cuda.OutOfMemoryError: print(f"{ctx_size:12,} --- OOM ---") torch.cuda.empty_cache() continue # Measure torch.cuda.synchronize() times = [] for _ in range(3): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): h = eng2.model(input_ids, eng2._cache_k, eng2._cache_v, start_pos=ctx_size) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() med = times[len(times)//2] ms = med * 1000 tps = K1 / med us = ms / K1 * 1000 cache_gb = ctx_size * kv_per_token / 1e9 vram = torch.cuda.max_memory_allocated() / 1e9 print(f"{ctx_size:12,} {cache_gb:10.2f} {ms:10.1f} {tps:12,.0f} {us:8.2f}") # Clear cache for next test for layer_k, layer_v in zip(eng2._cache_k, eng2._cache_v): layer_k.zero_() layer_v.zero_() torch.cuda.empty_cache() print(f"\n peak VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB") print(f"\n ANALYSIS:") print(f" - Sliding window (W=4097): speed is FLAT regardless of total tokens generated") print(f" - Growing cache: speed degrades as O(N) because each step reads all N KV entries") print(f" - At 50M context, reading 25.6GB KV cache at ~768GB/s = ~33ms minimum per step") print(f" - 4097 tokens / 33ms = ~124k tok/s theoretical floor at 50M context")