spec100m / sweep_500m_v2.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
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")