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