"""v5 benchmark: compressed MoBA + FP8 cache + within-block sparse + real chained prefill. Tests the full stack at 20M context on A6000: 1. Build 500M model (d=1024, L=8, MQA kv=1, ffn_mult=8) 2. Allocate compressed FP8 KV cache 3. Run real chained speculative prefill (20M tokens in K+1 chunks) 4. Measure throughput + VRAM """ import time import torch from config import Config from model import build_model from inference import V5Engine print(f"GPU: {torch.cuda.get_device_name(0)}") vram_free = torch.cuda.mem_get_info()[0] / 1e9 print(f"VRAM: {vram_free:.1f} GB free") print() # --- v5 500M config --- cfg = Config.v5_500m() params = cfg.count_params() print(f"Config: d={cfg.d_model} L={cfg.n_layers} H={cfg.n_heads} kv={cfg.n_kv_heads} " f"ffn={cfg.ffn_mult} K={cfg.medusa_heads}") print(f" base params : {params['base_total']/1e6:7.2f}M") print(f" medusa heads: {cfg.medusa_heads} x {params['medusa_per_head']/1e6:.3f}M = " f"{params['medusa_total']/1e6:6.2f}M") print(f" GRAND TOTAL : {params['grand_total']/1e6:7.2f}M") print(f" compression: m={cfg.kv_compress_m}, fp8={cfg.kv_fp8}, stride={cfg.sparse_stride}") print(f" block_size_comp: {cfg.block_size_comp} entries/block") print() # KV cache size estimate CTX = 20_000_000 m = cfg.kv_compress_m compressed_len = CTX // m kv_bytes = cfg.kv_cache_dtype_size # 1 for fp8, 2 for bf16 kv_gb = compressed_len * cfg.n_kv_heads * cfg.head_dim * kv_bytes * cfg.n_layers / 1e9 print(f"KV cache at {CTX/1e6:.0f}M context:") print(f" uncompressed: {CTX * cfg.n_kv_heads * cfg.head_dim * 4 * cfg.n_layers / 1e9:.2f} GB") print(f" compressed: {kv_gb:.2f} GB (m={m}, fp8={cfg.kv_fp8})") print(f" reduction: {CTX * cfg.n_kv_heads * cfg.head_dim * 4 * cfg.n_layers / 1e9 / kv_gb:.1f}x") print() # --- Build model --- print("Building model...") model = build_model(cfg, "cuda") print(f" model VRAM: {torch.cuda.memory_allocated()/1e9:.2f} GB") # --- Create engine --- engine = V5Engine(model, "cuda") print(f" engine ready (medusa_fast={engine._medusa_fast})") print() # --- Generate random prompt (simulating 20M tokens) --- print(f"Generating {CTX/1e6:.0f}M random tokens...") prompt_ids = torch.randint(0, cfg.vocab_size, (CTX,)).tolist() print(f" prompt: {len(prompt_ids)} tokens") print() # --- Run real chained speculative prefill --- print(f"=== Real Chained Speculative Prefill ({CTX/1e6:.0f}M tokens) ===") print(f" chunk size: {cfg.medusa_heads + 1} tokens") print(f" compression: {cfg.kv_compress_m} tokens -> 1 KV entry") print(f" cache: FP8={cfg.kv_fp8}, stride={cfg.sparse_stride}") print() # Warmup (small) warmup_prompt = prompt_ids[:8192] print("Warmup (8k tokens)...") result = engine.prefill_chained(warmup_prompt) print(f" warmup: {result['tokens']} tokens in {result['time']:.2f}s " f"({result['tok/s']:,.0f} tok/s)") print() # Full prefill print(f"Full prefill ({CTX/1e6:.0f}M tokens)...") torch.cuda.reset_peak_memory_stats() t0 = time.perf_counter() # Process in 1M-token segments to show progress SEGMENT = 1_000_000 total_tokens = 0 total_time = 0.0 n_segments = CTX // SEGMENT for seg in range(n_segments): seg_prompt = prompt_ids[seg * SEGMENT:(seg + 1) * SEGMENT] r = engine.prefill_chained(seg_prompt) total_tokens += r["tokens"] total_time += r["time"] print(f" segment {seg+1}/{n_segments}: {total_tokens/1e6:.1f}M tokens, " f"{total_tokens/total_time:,.0f} tok/s, " f"VRAM {torch.cuda.max_memory_allocated()/1e9:.2f} GB") peak_vram = torch.cuda.max_memory_allocated() / 1e9 print() print(f"=== Results ===") print(f" tokens: {total_tokens:,}") print(f" time: {total_time:.1f}s") print(f" speed: {total_tokens/total_time:,.0f} tok/s") print(f" steps: {total_tokens // (cfg.medusa_heads + 1):,}") print(f" per-step: {total_time / (total_tokens // (cfg.medusa_heads + 1)) * 1000:.2f} ms") print(f" peak VRAM: {peak_vram:.2f} GB") print(f" target: 50,000 tok/s {'PASS' if total_tokens/total_time > 50000 else 'FAIL'}")