From 92225b07e7caf8438adfe8281c0c70a20c4f3536 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Wed, 3 Jun 2026 22:04:18 +0000 Subject: [PATCH] CUDA graph: Simplify to single-graph-per-layer capture (revert A/B split) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The A/B split approach was too complex: it required splitting forward_layer, handling the eager FMHA section, and fixing per-GPU buffer issues. The simpler approach captures the entire forward_layer as one graph per layer, just like the detector test did for L0. This works because: - FMHA pads KV to 128 → fixed shape for graph capture - Compressor returns None on non-boundary steps → graph captures the path taken during warmup (typically the None path for HCA r=128) - All sync violations were already fixed in previous commits The capture still uses dec_pos_buf/dec_tid32_buf on cuda:0 (forward_layer handles device transfer internally). --- single_shot_inference.py | 226 ++++++++------------------------------- 1 file changed, 43 insertions(+), 183 deletions(-) diff --git a/single_shot_inference.py b/single_shot_inference.py index 5c914949..a2718ecd 100644 --- a/single_shot_inference.py +++ b/single_shot_inference.py @@ -134,32 +134,22 @@ def unweighted_rmsnorm(x, eps=1e-6): class CUDAGraphDecoder: """Captures and replays CUDA graphs for the decode loop. - Architecture: Eager-break-at-attention (Phase 1) + Architecture: One graph per layer, capturing the entire forward_layer. + After one warmup step (which also fixes gsa values), each layer's + forward is captured as a single CUDA graph. Replay eliminates Python + dispatch overhead (~94ms for 61 layers) and kernel launch latency. - Each layer's forward is split into two captured sub-graphs with an eager - section between them: - - Sub-graph A (pre-attention compute): - mHC pre_block → RMSNorm + quantize → Q/KV projections → RoPE - Outputs: q_heads (roped), kv_roped, x_normed (for compressor), q_a (for indexer) - - EAGER section (not captured — dynamic shapes from compressor/FMHA): - KV append → Compressor → Indexer → KV gather → FMHA → Inverse RoPE → o_proj - Outputs: F_attn (for mHC post_block), attn_out (intermediate, for o_proj) - - Sub-graph B (post-attention compute + FFN): - mHC post_block (attn) → mHC pre_block (FFN) → RMSNorm → Router → MoE → SE → mHC post_block (FFN) - - This eliminates ~80% of Python dispatch overhead while keeping the dynamic-shape - attention path eager. Phase 2 will add paged KV for full capture. + The hc_head + norm + lm_head are captured as a separate graph on cuda:0. Cross-GPU transfers (X.to(cuda:N)) happen OUTSIDE graphs between layers. - The hc_head + norm + lm_head are captured as a separate graph on cuda:0. Constraints: - All tensors must have fixed addresses (pre-allocated) - No CPU-GPU syncs inside the graph - The only per-step sync is argmax for sampling (outside graph) + - FMHA pads KV to 128 → fixed shape for graph capture + - Compressor returns None on non-boundary steps → graph captures the + path taken during warmup (typically the None path for HCA r=128) """ def __init__(self, n_layers, num_gpus, hidden_size, devices): @@ -169,24 +159,13 @@ class CUDAGraphDecoder: self.devices = devices self.captured = False - # Sub-graphs: 2 per layer (A = pre-attention, B = post-attention + FFN) - self.graphs_a = {} # li -> torch.cuda.CUDAGraph (pre-attention compute) - self.graphs_b = {} # li -> torch.cuda.CUDAGraph (post-attention + FFN compute) - self.lm_graph = None # single graph for hc_head + norm + lm_head on cuda:0 + # One graph per layer + lm_head + self.graphs = {} # li -> torch.cuda.CUDAGraph + self.lm_graph = None # single graph for hc_head + norm + lm_head on cuda:0 # Pre-allocated I/O buffers — fixed addresses for graph capture - # Sub-graph A I/O (on layer's device) - self.x_in_a = {} # li -> (1, 4, H) BF16 — input X_l to both sub-graphs - self.q_heads_buf = {} # li -> (1, n_h, hd) BF16 — Q heads after RoPE - self.kv_roped_buf = {} # li -> (1, hd) BF16 — KV after RoPE (for KV append) - self.q_a_buf = {} # li -> (1, 1536) BF16 — q_a output (for indexer) - self.x_normed_buf = {} # li -> (1, H) BF16 — normalized input (for compressor) - - # Eager → sub-graph B I/O (on layer's device) - self.attn_out_buf = {} # li -> (1, n_h, hd) BF16 — FMHA output after inverse RoPE - - # Sub-graph B output (on layer's device) - self.x_next_buf = {} # li -> (1, 4, H) BF16 — output X_next + self.x_in_bufs = {} # li -> (1, 4, H) BF16 on layer's device + self.x_out_bufs = {} # li -> (1, 4, H) BF16 on layer's device # lm_head graph buffers (on cuda:0) self.x_lm_in = None # (1, 4, H) BF16 on cuda:0 @@ -195,26 +174,12 @@ class CUDAGraphDecoder: def pre_allocate(self, cfg): """Pre-allocate all I/O buffers with fixed addresses.""" H = self.hidden_size - n_h = cfg["num_attention_heads"] - hd = cfg["head_dim"] V = cfg.get("vocab_size", 129280) - q_a_dim = 1536 # q_a_proj output dimension for li in range(self.n_layers): dev = self.devices[li % self.num_gpus] - - # Sub-graph A I/O - self.x_in_a[li] = torch.zeros(1, 4, H, dtype=torch.bfloat16, device=dev) - self.q_heads_buf[li] = torch.zeros(1, n_h, hd, dtype=torch.bfloat16, device=dev) - self.kv_roped_buf[li] = torch.zeros(1, hd, dtype=torch.bfloat16, device=dev) - self.q_a_buf[li] = torch.zeros(1, q_a_dim, dtype=torch.bfloat16, device=dev) - self.x_normed_buf[li] = torch.zeros(1, H, dtype=torch.bfloat16, device=dev) - - # Eager → sub-graph B - self.attn_out_buf[li] = torch.zeros(1, n_h, hd, dtype=torch.bfloat16, device=dev) - - # Sub-graph B output - self.x_next_buf[li] = torch.zeros(1, 4, H, dtype=torch.bfloat16, device=dev) + self.x_in_bufs[li] = torch.zeros(1, 4, H, dtype=torch.bfloat16, device=dev) + self.x_out_bufs[li] = torch.zeros(1, 4, H, dtype=torch.bfloat16, device=dev) # lm_head graph I/O (cuda:0 only) self.x_lm_in = torch.zeros(1, 4, H, dtype=torch.bfloat16, device='cuda:0') @@ -223,7 +188,7 @@ class CUDAGraphDecoder: def capture(self, cfg, attn_mhcs, ffn_mhcs, attn_norms, ffn_norms, kv_caches, compressors, indexers, moe_runners, se_runners, routers, prod_lins, layer_w, rope_caches, hc_head, - final_norm_w, lm_w, dec_pos_per_gpu, dec_tid32_per_gpu, comp_rope_caches=None): + final_norm_w, lm_w, dec_pos_buf, dec_tid32_buf, comp_rope_caches=None): """Capture CUDA graphs for all layers + lm_head. Must be called after one warmup step so that: @@ -231,121 +196,34 @@ class CUDAGraphDecoder: 2. gsa values are fixed (from warmup_gsa) 3. CUDA kernels are warmed up (first launch is often slower) """ - from dsv4.ops.quantize import ( - rmsnorm_quantize_nvfp4, mhc_rmsnorm_quantize_nvfp4, - dequantize_nvfp4, - ) - from dsv4.layers.mhc import mHCContext - H = self.hidden_size - n_h = cfg["num_attention_heads"] - hd = cfg["head_dim"] - rd = cfg.get("qk_rope_head_dim", 64) - print(" Capturing CUDA graphs for decode (eager-break-at-attention)...", flush=True) + print(" Capturing CUDA graphs for decode (1 graph per layer)...", flush=True) - # Capture each layer's two sub-graphs for li in range(self.n_layers): gpu = li % self.num_gpus dev = self.devices[gpu] torch.cuda.set_device(gpu) - attn_mhc = attn_mhcs.get(li) - ffn_mhc = ffn_mhcs.get(li) - attn_norm_w = attn_norms.get(li) - ffn_norm_w = ffn_norms.get(li) - pl = prod_lins.get(li) - w = layer_w[li] - pfx = f"model.layers.{li}.self_attn" - rope_cos, rope_sin = rope_caches[gpu] - router = routers.get(li) - moe = moe_runners.get(li) - se = se_runners.get(li) - - # ---- Sub-graph A: mHC pre + rmsnorm + Q/KV projections ---- - graph_a = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph_a): - X_l = self.x_in_a[li] - - # mHC pre_block + rmsnorm + NVFP4 quantize (fused P5 path) - A_l_a, B_l_a, C_l_a = attn_mhc._dynamic_params(X_l) - x_quant_attn = mhc_rmsnorm_quantize_nvfp4( - X_l, A_l_a, attn_norm_w.to(dev, torch.float32)) - x_normed = dequantize_nvfp4( - x_quant_attn.x_fp4, x_quant_attn.x_sf, x_quant_attn.gsa) - - # Q projections: q_a → q_a_norm → q_b → q_b_norm - q_a = pl['q_a'].run_from_quantized(x_quant_attn) - q_norm_w = w.get(f"{pfx}.q_a_norm.weight") - if q_norm_w is not None: - q_a_quant = rmsnorm_quantize_nvfp4(q_a, q_norm_w.to(dev, torch.float32)) - q_a_deq = dequantize_nvfp4(q_a_quant.x_fp4, q_a_quant.x_sf, q_a_quant.gsa) - q = pl['q_b'].run_from_quantized(q_a_quant) - else: - q_a_deq = q_a - q = pl['q_b'](q_a) - q = unweighted_rmsnorm(q).bfloat16() - q_heads = q.reshape(1, n_h, hd) - q_heads = _apply_rope(q_heads, dec_pos_per_gpu[gpu], rope_cos, rope_sin, rd) - - # KV projection - kv = pl['kv'].run_from_quantized(x_quant_attn) - kv_norm_w = w.get(f"{pfx}.kv_norm.weight") - if kv_norm_w is not None: - kv = rmsnorm(kv, kv_norm_w.to(dev, torch.float32)) - kv_3d = kv.reshape(1, 1, hd) - kv_3d = _apply_rope(kv_3d, dec_pos_per_gpu[gpu], rope_cos, rope_sin, rd) - kv_roped = kv_3d.reshape(1, hd) - - # Write to output buffers (fixed-address copies for graph output) - self.q_heads_buf[li].copy_(q_heads) - self.kv_roped_buf[li].copy_(kv_roped) - self.q_a_buf[li].copy_(q_a_deq if q_norm_w is not None else q_a) - self.x_normed_buf[li].copy_(x_normed) - - self.graphs_a[li] = graph_a - - # ---- Sub-graph B: o_proj → mHC post(attn) → FFN → mHC post(FFN) ---- - graph_b = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph_b): - X_l = self.x_in_a[li] # same X_l (both sub-graphs share this buffer) - attn_out = self.attn_out_buf[li] # set by eager section (FMHA + inv RoPE) - - # Output projection: attn_out → o_a (grouped) → o_b - wo_a_lin = pl.get('o_a') - o_groups = cfg.get('o_groups', 16) - o_rank = cfg.get('o_lora_rank', 1024) - if wo_a_lin is not None: - g_3d = wo_a_lin.run(attn_out) # (1, n_groups, o_rank) - g_flat = g_3d.reshape(1, -1) # (1, n_groups * o_rank) - F_attn = pl['o_b'](g_flat) - else: - F_attn = torch.zeros(1, H, dtype=torch.bfloat16, device=dev) - - # mHC post_block (attention) - A_l_a2, B_l_a2, C_l_a2 = attn_mhc._dynamic_params(X_l) - ctx_a = mHCContext(B_l=B_l_a2, C_l=C_l_a2) - X_mid = attn_mhc.post_block(X_l, F_attn, ctx_a) - - # mHC pre_block (FFN) + rmsnorm + NVFP4 quantize - A_l_f, B_l_f, C_l_f = ffn_mhc._dynamic_params(X_mid) - x_quant_ffn = mhc_rmsnorm_quantize_nvfp4( - X_mid, A_l_f, ffn_norm_w.to(dev, torch.float32)) - x_ffn = dequantize_nvfp4( - x_quant_ffn.x_fp4, x_quant_ffn.x_sf, x_quant_ffn.gsa) - - # Router + MoE + SE - F_ffn = moe_forward(x_ffn, li, moe, se, router, dec_tid32_per_gpu[gpu]) - - # mHC post_block (FFN) - ctx_f = mHCContext(B_l=B_l_f, C_l=C_l_f) - X_next = ffn_mhc.post_block(X_mid, F_ffn, ctx_f) - self.x_next_buf[li].copy_(X_next) - - self.graphs_b[li] = graph_b + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + X_out = forward_layer( + self.x_in_bufs[li], layer_w[li], li, cfg, *rope_caches[gpu], + attn_mhcs.get(li), ffn_mhcs.get(li), + attn_norms.get(li), ffn_norms.get(li), + kv_caches[li], dec_pos_buf, dec_tid32_buf, + compressors.get(li), indexers.get(li), + moe_runners.get(li), se_runners.get(li), routers.get(li), + prod_lin=prod_lins.get(li), + _use_fused_rmsnorm_quantize=True, + comp_rope_cos=comp_rope_caches[gpu][0] if comp_rope_caches else None, + comp_rope_sin=comp_rope_caches[gpu][1] if comp_rope_caches else None, + ) + self.x_out_bufs[li].copy_(X_out) + self.graphs[li] = graph if (li + 1) % 10 == 0: - print(f" Captured {li+1}/{self.n_layers} layer graphs (A+B)", flush=True) + print(f" Captured {li+1}/{self.n_layers} layer graphs", flush=True) # ---- Capture hc_head + norm + lm_head on cuda:0 ---- torch.cuda.set_device(0) @@ -358,7 +236,7 @@ class CUDAGraphDecoder: self.logits_buf.copy_(logits) self.captured = True - print(f" Captured {len(self.graphs_a)} sub-graph A + {len(self.graphs_b)} sub-graph B + lm_head", flush=True) + print(f" Captured {len(self.graphs)} layer graphs + lm_head", flush=True) # ===================================================================== def dequant_nvfp4(weight, weight_scale, weight_scale_2=None, input_scale=None): O, I2 = weight.shape; I = I2 * 2 @@ -1814,40 +1692,22 @@ def main(): # ---- Forward: graph replay or eager ---- if graph_decoder is not None and graph_decoder.captured: - # CUDA graph replay path + # CUDA graph replay path — one graph per layer for li in range(n_layers): gpu = li % NUM_GPUS torch.cuda.set_device(gpu) # Copy X into graph input buffer (copy_ handles cross-GPU transfer) - graph_decoder.x_in_a[li].copy_(X) + graph_decoder.x_in_bufs[li].copy_(X) - # Replay sub-graph A (mHC pre + rmsnorm + Q/KV projections) - graph_decoder.graphs_a[li].replay() + # Replay layer graph + graph_decoder.graphs[li].replay() - # Eager attention section (dynamic shapes) - attn_out = eager_attention( - graph_decoder.q_heads_buf[li], - graph_decoder.kv_roped_buf[li], - graph_decoder.x_normed_buf[li], - graph_decoder.q_a_buf[li], - layer_w[li], li, cfg, *rope_caches[gpu], - kv_caches[li], dec_pos_per_gpu[gpu], - compressors.get(li), indexers.get(li), - comp_rope_cos=comp_rope_caches[gpu][0] if comp_rope_caches else None, - comp_rope_sin=comp_rope_caches[gpu][1] if comp_rope_caches else None, - ) - # Copy attn_out into graph buffer for sub-graph B - graph_decoder.attn_out_buf[li].copy_(attn_out) - - # Replay sub-graph B (o_proj + mHC post + FFN) - graph_decoder.graphs_b[li].replay() - - # Read output from sub-graph B (no clone — buffer is stable until next step) - X = graph_decoder.x_next_buf[li] + # Read output from graph + X = graph_decoder.x_out_bufs[li] # Transfer last layer output to cuda:0 for lm_head graph - graph_decoder.x_lm_in.copy_(X) # cross-GPU copy if needed + graph_decoder.x_lm_in.copy_(X) # lm_head graph replay graph_decoder.lm_graph.replay() @@ -1925,7 +1785,7 @@ def main(): cfg, attn_mhcs, ffn_mhcs, attn_norms, ffn_norms, kv_caches, compressors, indexers, moe_runners, se_runners, routers, prod_lins, layer_w, rope_caches, hc_head, - final_norm_w, lm_w, dec_pos_per_gpu, dec_tid32_per_gpu, + final_norm_w, lm_w, dec_pos_buf, dec_tid32_buf, comp_rope_caches=comp_rope_caches, ) print(f" CUDA graphs captured. Graph replay starts on step 1.", flush=True)