"""Find the config that hits 50k+ tok/s at 20M context on A6000. Physics: tok/s ~ TFLOPS / (S * hd * H * 4). Need hd*H <= ~19 for 50k at 20M. Test: d_model=16 (hd=16,H=1), d_model=32 (hd=32,H=1), d_model=64 (hd=32,H=2). """ import time import torch from config import Config from model import build_model from inference import InferenceEngine 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") CTX = 20_000_000 K = 4096 K1 = K + 1 print(f"\nTarget: {CTX:,} context, 50,000+ tok/s") print(f"Theoretical ceiling: tok/s ~ 77.5e12 / (CTX * hd * H * 4)") print(f"\n{'d':>4} {'H':>3} {'kv':>3} {'hd':>4} {'K':>5} {'params(M)':>10} " f"{'KV(GB)':>7} {'step(ms)':>9} {'tok/s':>10} {'us/tok':>8}") print("-" * 70) configs = [ # (d_model, n_heads, n_kv_heads, K) (16, 1, 1, 4096), # hd=16, H=1: hd*H=16 -> ~60k theoretical (32, 1, 1, 4096), # hd=32, H=1: hd*H=32 -> ~30k theoretical (32, 1, 1, 2048), # same model, smaller K (64, 2, 1, 4096), # hd=32, H=2: hd*H=64 -> ~15k theoretical (16, 1, 1, 8192), # tiny model, big K (16, 1, 1, 2048), # tiny model, small K (48, 1, 1, 4096), # hd=48, H=1: hd*H=48 -> ~20k theoretical ] for d, H, kv, k in configs: hd = d // H cfg = Config(d_model=d, n_heads=H, n_kv_heads=kv, medusa_heads=k, medusa_rank=1, window_size=0, max_seq_len=CTX + k + 100) kv_per_token = 2 * kv * hd * 2 * 1 # bytes (1 layer) kv_gb = CTX * kv_per_token / 1e9 est = cfg.estimate_vram_gb() if est > 45: print(f"{d:4d} {H:3d} {kv:3d} {hd:4d} {k:5d} --- est VRAM {est:.1f}GB too high ---") continue try: m = build_model(cfg, "cuda") eng = InferenceEngine(m) eng.alloc_cache() except Exception as e: print(f"{d:4d} {H:3d} {kv:3d} {hd:4d} {k:5d} --- ERROR: {e} ---") torch.cuda.empty_cache() continue # Fill cache to CTX tokens with random data for lk, lv in zip(eng._cache_k, eng._cache_v): lk[:, :CTX].normal_(0, 0.1) lv[:, :CTX].normal_(0, 0.1) input_ids = torch.randint(0, cfg.vocab_size, (1, k + 1), device="cuda") params = sum(p.numel() for p in m.parameters()) / 1e6 # Warmup try: for _ in range(2): with torch.no_grad(): h = eng.model(input_ids, eng._cache_k, eng._cache_v, start_pos=CTX) torch.cuda.synchronize() except torch.cuda.OutOfMemoryError: print(f"{d:4d} {H:3d} {kv:3d} {hd:4d} {k:5d} {params:10.1f} {kv_gb:7.2f} --- OOM ---") del m, eng 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 = eng.model(input_ids, eng._cache_k, eng._cache_v, start_pos=CTX) 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"{d:4d} {H:3d} {kv:3d} {hd:4d} {k:5d} {params:10.1f} {kv_gb:7.2f} " f"{ms:9.1f} {tps:10,.0f} {us:8.2f}") # Clear for lk, lv in zip(eng._cache_k, eng._cache_v): lk.zero_() lv.zero_() del m, eng torch.cuda.empty_cache() print(f"\n peak VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")