"""Just test 50M context (the previous run hit everything up to 10M).""" 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"free VRAM: {torch.cuda.mem_get_info()[0]/1e9:.2f} GB") cfg = Config(window_size=0, max_seq_len=50_000_000) K1 = cfg.medusa_heads + 1 kv_per_token = 2 * cfg.n_kv_heads * cfg.head_dim * 2 * cfg.n_layers print(f"KV cache at 50M: {50_000_000 * kv_per_token / 1e9:.1f} GB") m = build_model(cfg, "cuda") eng = InferenceEngine(m) eng.alloc_cache() print(f"cache allocated, VRAM: {torch.cuda.memory_allocated()/1e9:.2f} GB") # Fill cache to 49M tokens (just under max_seq_len) ctx = 49_000_000 print(f"\nFilling cache to {ctx:,} tokens...") for layer_k, layer_v in zip(eng._cache_k, eng._cache_v): layer_k[:, :ctx].normal_(0, 0.1) layer_v[:, :ctx].normal_(0, 0.1) print(f" VRAM after fill: {torch.cuda.memory_allocated()/1e9:.2f} GB") input_ids = torch.randint(0, cfg.vocab_size, (1, K1), device="cuda") # Warmup print(" warmup...") 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() # Measure print(" measuring...") 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 = K1 / med us = ms / K1 * 1000 cache_gb = ctx * kv_per_token / 1e9 print(f"\n context={ctx:,} cache={cache_gb:.1f}GB step={ms:.0f}ms " f"speed={tps:,.0f} tok/s {us:.1f} us/tok") print(f" peak VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")