From 518a1d3f9535f2f50c0fea784e556d1ce19a92a8 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Wed, 3 Jun 2026 17:53:40 +0000 Subject: [PATCH] CUDA graph: Fix MoE scatter_add_ index dtype + fix second bincount MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. scatter_add_ requires int64 indices — ensure sorted_ids is .long() 2. Fixed the SECOND torch.bincount call (line 590) — same scatter_add_ pattern 3. Both code paths now use pre-allocated _tokens_per_expert_buf --- dsv4/layers/moe.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/dsv4/layers/moe.py b/dsv4/layers/moe.py index b9b327d7..3d1b33fb 100644 --- a/dsv4/layers/moe.py +++ b/dsv4/layers/moe.py @@ -474,11 +474,12 @@ class Nvfp4MoE: # torch.bincount produces data-dependent shapes (violates graph capture). # Instead, use scatter_add_ into a pre-allocated buffer (fixed shape, GPU-only). self._tokens_per_expert_buf.zero_() - # Pre-allocated ones buffer — avoids per-call torch.ones() allocation - n_slots = sorted_ids.shape[0] + # scatter_add_ requires int64 indices — ensure sorted_ids is int64 + sorted_ids_i64 = sorted_ids.long() + n_slots = sorted_ids_i64.shape[0] if not hasattr(self, '_ones_buf') or self._ones_buf.shape[0] < n_slots: - self._ones_buf = torch.ones(self.max_num_tokens * self.top_k, dtype=torch.int32, device=sorted_ids.device) - self._tokens_per_expert_buf.scatter_add_(0, sorted_ids, self._ones_buf[:n_slots]) + self._ones_buf = torch.ones(self.max_num_tokens * self.top_k, dtype=self._tokens_per_expert_buf.dtype, device=sorted_ids_i64.device) + self._tokens_per_expert_buf.scatter_add_(0, sorted_ids_i64, self._ones_buf[:n_slots]) tokens_per_expert = self._tokens_per_expert_buf[:self.num_experts] expert_offsets = self._expert_offsets_buf expert_offsets.zero_() @@ -586,7 +587,14 @@ class Nvfp4MoE: sorted_token_ids = token_indices[sort_idx] # Expert offsets (real token counts) - tokens_per_expert = torch.bincount(sorted_ids, minlength=self.num_experts)[:self.num_experts].int() + # CUDA-graph-safe: scatter_add_ instead of bincount (fixed shape, GPU-only) + self._tokens_per_expert_buf.zero_() + sorted_ids_i64 = sorted_ids.long() + n_slots = sorted_ids_i64.shape[0] + if not hasattr(self, '_ones_buf') or self._ones_buf.shape[0] < n_slots: + self._ones_buf = torch.ones(self.max_num_tokens * self.top_k, dtype=self._tokens_per_expert_buf.dtype, device=sorted_ids_i64.device) + self._tokens_per_expert_buf.scatter_add_(0, sorted_ids_i64, self._ones_buf[:n_slots]) + tokens_per_expert = self._tokens_per_expert_buf[:self.num_experts] expert_offsets = self._expert_offsets_buf expert_offsets.zero_() expert_offsets[1:self.num_experts + 1] = tokens_per_expert.cumsum(0)