diff --git a/vllm/nvfp4_cutedsl.py b/vllm/nvfp4_cutedsl.py index 33a39052..a049398e 100644 --- a/vllm/nvfp4_cutedsl.py +++ b/vllm/nvfp4_cutedsl.py @@ -79,11 +79,14 @@ class CuTeDSLMoERunner: self._buffers_allocated = False def _fill_token_indices(self): - """Fill _token_indices with [0,0,..0, 1,1,..1, ...] (each token repeated top_k times).""" - src = torch.arange(self.max_num_tokens, dtype=torch.int32, device=self.device) - self._token_indices.copy_( - src.unsqueeze(1).expand(-1, self.top_k).contiguous().view(-1) - ) + """Fill _token_indices with [0,0,..0, 1,1,..1, ...] (each token repeated top_k times). + + Builds on CPU first, then copies to GPU, to ensure correctness + regardless of CuTeDSL JIT GPU memory corruption. + """ + src = torch.arange(self.max_num_tokens, dtype=torch.int32) + cpu_indices = src.unsqueeze(1).expand(-1, self.top_k).contiguous().view(-1) + self._token_indices.copy_(cpu_indices) def _allocate_buffers(self): """Pre-allocate scale buffers at max size for cudagraph compatibility.""" @@ -146,6 +149,10 @@ class CuTeDSLMoERunner: self.max_num_tokens * self.top_k, dtype=torch.int32, device=self.device ) self._fill_token_indices() + # Verify the fill worked (CuTeDSL JIT may have corrupted GPU memory) + torch.cuda.synchronize() + assert self._token_indices.max().item() == self.max_num_tokens - 1, \ + f"token_indices corrupted after fill! max={self._token_indices.max().item()}" self._needs_token_refill = True # GEMM JIT may corrupt; refill after first run self._expert_id_range = torch.arange( @@ -341,8 +348,18 @@ class CuTeDSLMoERunner: # -- Gather hidden states into slot order -- # DEBUG: Check for OOB indices before indexing - if sorted_token_ids.max() >= num_tokens or sorted_token_ids.min() < 0: - print(f"[CLAWMINE BUG] sorted_token_ids OOB! max={sorted_token_ids.max().item()} min={sorted_token_ids.min().item()} num_tokens={num_tokens} _token_indices_max={self._token_indices[:num_slots].max().item()}", flush=True) + ti_max = self._token_indices[:num_slots].max().item() + ti_min = self._token_indices[:num_slots].min().item() + si_max = sorted_token_ids.max().item() + si_min = sorted_token_ids.min().item() + print(f"[CLAWMINE] num_tokens={num_tokens} top_k={top_k} num_slots={num_slots} " + f"token_indices max={ti_max} min={ti_min} " + f"sorted_token_ids max={si_max} min={si_min} " + f"experts_start_idx={self.experts_start_idx} num_experts={self.num_experts} " + f"topk_ids max={topk_ids.max().item()} min={topk_ids.min().item()}", flush=True) + if si_max >= num_tokens or si_min < 0: + print(f"[CLAWMINE BUG] sorted_token_ids OOB! " + f"max={si_max} min={si_min} num_tokens={num_tokens}", flush=True) slot_hidden = hidden_states[sorted_token_ids] # === L1: gate + up ===