"""Profile the v5 step to find the real bottleneck. Weight read floor: 480M * 2B / 768 GB/s = 1.25 ms Compute floor: ~936 GFLOP / 1552 TFLOP/s = 0.6 ms Actual: 33 ms Overhead ratio: 33 / 1.25 = 26x above weight-read floor If overhead-dominated, code fixes matter more than architecture changes. """ import time import torch from config import Config from model import build_model cfg = Config.v5_500m() model = build_model(cfg, "cuda") model.precompute_medusa_tokens() K1 = cfg.medusa_heads + 1 BS_comp = cfg.block_size_comp m = cfg.kv_compress_m dt = model.tok_emb.weight.dtype # Allocate cache n_blocks = 5000 cache_comp = [ torch.zeros(1, n_blocks, BS_comp, cfg.n_kv_heads, cfg.head_dim, device="cuda", dtype=torch.float8_e4m3fn) for _ in range(cfg.n_layers) ] k_bar_comp = [ torch.zeros(1, n_blocks, cfg.n_kv_heads, cfg.head_dim, device="cuda", dtype=dt) for _ in range(cfg.n_layers) ] # Fill with random for cc, kb in zip(cache_comp, k_bar_comp): cc[:, :100] = torch.randn_like(cc[:, :100], dtype=dt).to(torch.float8_e4m3fn) kb[:, :100] = torch.randn_like(kb[:, :100]) input_ids = torch.randint(0, cfg.vocab_size, (1, K1), device="cuda") # --- Profile each component --- torch.cuda.synchronize() # 1. Full forward_moba_v5 for _ in range(3): with torch.no_grad(): h = model.forward_moba_v5(input_ids, cache_comp, k_bar_comp, 100, 0) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): h = model.forward_moba_v5(input_ids, cache_comp, k_bar_comp, 100, 0) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() full_ms = times[5] * 1000 print(f"Full forward_moba_v5: {full_ms:.2f} ms") # 2. Just embedding + norm + FFN (no attention) class NoAttnFwd(torch.nn.Module): def __init__(self, model): super().__init__() self.model = model def forward(self, ids): x = self.model.tok_emb(ids) for blk in self.model.blocks: x = x + blk.ffn(blk.norm2(x)) return self.model.norm_f(x) no_attn = NoAttnFwd(model).cuda().to(dt) for _ in range(3): with torch.no_grad(): _ = no_attn(input_ids) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): _ = no_attn(input_ids) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() ffn_ms = times[5] * 1000 print(f"Embed + FFN only (8 layers): {ffn_ms:.2f} ms") # 3. Just Q/K/V/O projections (no SDPA, no compression) class JustProj(torch.nn.Module): def __init__(self, model): super().__init__() self.model = model def forward(self, ids): x = self.model.tok_emb(ids) for blk in self.model.blocks: h = blk.norm1(x) q = blk.attn.wq(h).view(1, K1, 16, 64) k = blk.attn.wk(h).view(1, K1, 1, 64) v = blk.attn.wv(h).view(1, K1, 1, 64) o = blk.attn.wo(q.reshape(1, K1, -1)) x = x + o return x just_proj = JustProj(model).cuda().to(dt) for _ in range(3): with torch.no_grad(): _ = just_proj(input_ids) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): _ = just_proj(input_ids) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() proj_ms = times[5] * 1000 print(f"Q/K/V/O projections only: {proj_ms:.2f} ms") # 4. Compression projections only class JustComp(torch.nn.Module): def __init__(self, model): super().__init__() self.model = model def forward(self, ids): x = self.model.tok_emb(ids) for blk in self.model.blocks: h = blk.norm1(x) comp_kv = blk.attn.w_kv_comp(h) comp_z = blk.attn.w_z_comp(h) x = x + h # dummy residual (use normed h, not comp output) return x just_comp = JustComp(model).cuda().to(dt) for _ in range(3): with torch.no_grad(): _ = just_comp(input_ids) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): _ = just_comp(input_ids) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() comp_ms = times[5] * 1000 print(f"Compression proj only: {comp_ms:.2f} ms") # 5. Medusa fast path h_last = torch.randn(1, cfg.d_model, device="cuda", dtype=dt) for _ in range(3): with torch.no_grad(): _ = model.medusa_argmax_fast(h_last) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): _ = model.medusa_argmax_fast(h_last) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() medusa_ms = times[5] * 1000 print(f"Medusa argmax fast: {medusa_ms:.2f} ms") # 6. SDPA calls (3 blocks, the loop) q_t = torch.randn(1, 16, K1, 64, device="cuda", dtype=dt) bk = torch.randn(1, BS_comp, 1, 64, device="cuda", dtype=dt).transpose(1, 2) bv = bk.clone() for _ in range(3): _ = torch.nn.functional.scaled_dot_product_attention(q_t, bk, bv, enable_gqa=True) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() for _ in range(3): # 3 blocks _ = torch.nn.functional.scaled_dot_product_attention(q_t, bk, bv, enable_gqa=True) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() sdpa_ms = times[5] * 1000 print(f"3x SDPA (per layer): {sdpa_ms:.2f} ms (x8 layers = {sdpa_ms*8:.2f} ms)") # 7. Fused SDPA (concat 3 blocks into 1) bk_fused = torch.randn(1, 1, 3*BS_comp, 64, device="cuda", dtype=dt) for _ in range(3): _ = torch.nn.functional.scaled_dot_product_attention(q_t, bk_fused, bk_fused, enable_gqa=True) torch.cuda.synchronize() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() _ = torch.nn.functional.scaled_dot_product_attention(q_t, bk_fused, bk_fused, enable_gqa=True) torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() sdpa_fused_ms = times[5] * 1000 print(f"1x fused SDPA (3 blocks): {sdpa_fused_ms:.2f} ms (x8 = {sdpa_fused_ms*8:.2f} ms)") print() print(f"=== Breakdown ===") print(f" FFN (8 layers): {ffn_ms:.2f} ms ({ffn_ms/full_ms*100:.0f}%)") print(f" QKVO proj (8 layers): {proj_ms:.2f} ms ({proj_ms/full_ms*100:.0f}%)") print(f" Compression (8 layers): {comp_ms:.2f} ms ({comp_ms/full_ms*100:.0f}%)") print(f" SDPA 3x loop (8 layers): {sdpa_ms*8:.2f} ms ({sdpa_ms*8/full_ms*100:.0f}%)") print(f" SDPA fused (8 layers): {sdpa_fused_ms*8:.2f} ms ({sdpa_fused_ms*8/full_ms*100:.0f}%)") print(f" Medusa: {medusa_ms:.2f} ms ({medusa_ms/full_ms*100:.0f}%)") print(f" Full step: {full_ms:.2f} ms") print(f" Sum of parts: {ffn_ms+proj_ms+comp_ms+sdpa_ms*8+medusa_ms:.2f} ms") print(f" Unaccounted (Python): {full_ms - (ffn_ms+proj_ms+comp_ms+sdpa_ms*8+medusa_ms):.2f} ms")