"""500M model with DeepSeek-inspired KV compression + MQA + FP8 cache. Techniques from DeepSeek V4.1: 1. MQA (kv=1): one KV head, shared as both K and V. 4x KV reduction vs kv=4. 2. KV compression (m=4): compress every 4 tokens into 1 KV entry via learned weighted pooling. 4x KV reduction. 3. FP8 KV cache: store compressed KV in fp8 (E4M3). 2x reduction vs bf16. 4. Sliding window branch: keep recent n_win tokens uncompressed for local detail. Combined: 4 (MQA) × 4 (compression) × 2 (fp8) = 32x KV reduction. At 20M context, 8 layers, d=1024: 40 GB -> 1.25 GB. Target: 500M base + Medusa heads, ~1M tok/s at 20M context. """ import time import torch import torch.nn as nn import torch.nn.functional as F import math from config import Config from model import build_model, SpecModel 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") CTX = 20_000_000 TOP_K = 3 # --- Configs to test --- # Key insight: with MQA (kv=1) + compression (m=4) + fp8, KV cache is tiny. # So we can afford many layers + large d_model for 500M base. # # KV cache = CTX / m * kv * hd * 2 (K+V) * 1 (fp8) * L # With m=4, kv=1, hd=64, L=8: 20M/4 * 1 * 64 * 2 * 1 * 8 = 5.12 GB # With m=4, kv=1, hd=128, L=8: 20M/4 * 1 * 128 * 2 * 1 * 8 = 10.24 GB # # Base params: emb(50257*d) + L*(attn + ffn + norm) # d=1024, L=8, ffn_mult=4: emb=51.5M, attn~8*16.8M=134M, ffn=8*12.6M=100M -> ~286M # d=1024, L=8, ffn_mult=8: ffn=8*25.2M=201M -> ~387M # d=1024, L=12, ffn_mult=4: ~429M # d=1280, L=8, ffn_mult=4: emb=64.3M, attn~8*26.2M=210M, ffn=8*19.7M=157M -> ~432M # d=1280, L=10, ffn_mult=4: ~540M configs = [ # (name, d, L, H, kv, ffn_mult, K, compress_m, use_fp8) # Baseline (current 228M, no compression) ("baseline_228M", 384, 1, 12, 4, 2.667, 4096, 1, False), # 500M with MQA + compression ("500M_d1024_L8_m4", 1024, 8, 16, 1, 4.0, 4096, 4, True), ("500M_d1024_L8_m4_nofp8", 1024, 8, 16, 1, 4.0, 4096, 4, False), ("500M_d1024_L8_m8", 1024, 8, 16, 1, 4.0, 4096, 8, True), ("500M_d1024_L12_m4", 1024, 12, 16, 1, 4.0, 4096, 4, True), ("500M_d1280_L8_m4", 1280, 8, 20, 1, 4.0, 4096, 4, True), ("500M_d1280_L10_m4", 1280, 10, 20, 1, 4.0, 4096, 4, True), # Smaller K for less Medusa overhead ("500M_d1024_L8_m4_K2048", 1024, 8, 16, 1, 4.0, 2048, 4, True), ("500M_d1024_L8_m4_K1024", 1024, 8, 16, 1, 4.0, 1024, 4, True), # No compression (just MQA + fp8) for comparison ("500M_d1024_L8_m1_fp8", 1024, 8, 16, 1, 4.0, 4096, 1, True), ] print(f"\n{'name':>28} {'d':>5} {'L':>3} {'K':>5} {'m':>3} " f"{'base(M)':>8} {'total(M)':>9} {'KV(GB)':>7} " f"{'step(ms)':>9} {'tok/s':>10} {'us/tok':>8}") print("-" * 100) for name, d, L, H, kv, ffn, K, m, use_fp8 in configs: cfg = Config( d_model=d, n_layers=L, n_heads=H, n_kv_heads=kv, ffn_mult=ffn, medusa_heads=K, medusa_rank=1, medusa_hidden=0, window_size=0, max_seq_len=CTX + K + 100, use_alibi=False, moba_block_size=K + 1, moba_top_k=TOP_K, dtype="bfloat16", ) dt = getattr(torch, cfg.dtype) dev = "cuda" hd = d // H params = cfg.count_params() base_m = params["base_total"] / 1e6 total_m = params["grand_total"] / 1e6 # KV cache with compression + fp8 kv_bytes = 2 if use_fp8 else 4 # fp8=1 byte * 2 (K+V), bf16=2 bytes * 2 compressed_len = CTX // m kv_gb = compressed_len * kv * hd * kv_bytes * L / 1e9 if kv_gb + total_m * 2 / 1e9 > vram_free * 0.85: print(f"{name:>28} {d:5d} {L:3d} {K:5d} {m:3d} " f"{base_m:8.1f} {total_m:9.1f} {kv_gb:7.2f} --- VRAM too high ---") continue BS = K + 1 n_blocks = (compressed_len // BS) + 2 try: m_model = build_model(cfg, dev) m_model.precompute_medusa_tokens() # Block-sparse KV cache (compressed) cache_dt = torch.float8_e4m3fn if use_fp8 else dt # For fp8, we store in fp8 but attention needs fp16/bf16 # Simple approach: store in bf16 for now, measure with compression only cache_k = [torch.zeros(1, n_blocks, BS, kv, hd, device=dev, dtype=dt) for _ in range(L)] cache_v = [torch.zeros(1, n_blocks, BS, kv, hd, device=dev, dtype=dt) for _ in range(L)] k_bar = [torch.zeros(1, n_blocks, kv, hd, device=dev, dtype=dt) for _ in range(L)] nf = CTX // BS // m # filled blocks (compressed) if nf >= n_blocks: nf = n_blocks - 1 # Fill cache for lk, lv, lb in zip(cache_k, cache_v, k_bar): lk[:, :nf].normal_(0, 0.1) lv[:, :nf].normal_(0, 0.1) lb[:, :nf] = lk[:, :nf].mean(dim=2) input_ids = torch.randint(0, cfg.vocab_size, (1, K + 1), device=dev) # Warmup for _ in range(2): with torch.no_grad(): h = m_model.forward_moba(input_ids, cache_k, cache_v, k_bar, nf, 0) _ = m_model.medusa_argmax_fast(h[:, -1]) torch.cuda.synchronize() # Measure times = [] for _ in range(5): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): h = m_model.forward_moba(input_ids, cache_k, cache_v, k_bar, nf, 0) _ = m_model.medusa_argmax_fast(h[:, -1]) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() med = times[len(times) // 2] ms = med * 1000 tps = (K + 1) / med us = ms / (K + 1) * 1000 print(f"{name:>28} {d:5d} {L:3d} {K:5d} {m:3d} " f"{base_m:8.1f} {total_m:9.1f} {kv_gb:7.2f} " f"{ms:9.1f} {tps:10,.0f} {us:8.2f}") except Exception as e: err = str(e)[:60] print(f"{name:>28} {d:5d} {L:3d} {K:5d} {m:3d} " f"{base_m:8.1f} {total_m:9.1f} {kv_gb:7.2f} --- ERROR: {err} ---") torch.cuda.empty_cache() try: del m_model, cache_k, cache_v, k_bar except: pass torch.cuda.empty_cache() print(f"\n peak VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")