diff --git a/unsloth/kernels/moe/grouped_gemm/interface.py b/unsloth/kernels/moe/grouped_gemm/interface.py index 572b02702b..be669dc0be 100644 --- a/unsloth/kernels/moe/grouped_gemm/interface.py +++ b/unsloth/kernels/moe/grouped_gemm/interface.py @@ -40,25 +40,45 @@ logger.addHandler(ch) # TMA requires both: # 1. GPU capability >= 9 (Hopper+) # 2. Triton version with TMA API (make_tensor_descriptor or _experimental_make_tensor_descriptor) -def _check_tma_support(): +def _check_tma_support(device = None): import triton.language as tl - - gpu_supports_tma = torch.cuda.get_device_capability()[0] >= 9 - # Check for both old experimental and new stable API names - triton_has_tma_api = hasattr(tl, "make_tensor_descriptor") or hasattr( - tl, "_experimental_make_tensor_descriptor" - ) - return gpu_supports_tma and triton_has_tma_api + try: + if device is None: + if not torch.cuda.is_available(): + return False + device_index = torch.cuda.current_device() + else: + if isinstance(device, torch.device): + if device.type != "cuda": + return False + device_index = ( + torch.cuda.current_device() + if device.index is None + else device.index + ) + else: + device_index = int(device) + gpu_supports_tma = torch.cuda.get_device_capability(device_index)[0] >= 9 + # Require stable API since kernels call tl.make_tensor_descriptor directly. + triton_has_tma_api = hasattr(tl, "make_tensor_descriptor") + return gpu_supports_tma and triton_has_tma_api + except Exception: + return False -_SUPPORTS_TMA = _check_tma_support() +_SUPPORTS_TMA = {} # Check if triton.set_allocator is available (Triton 3.0+) _HAS_SET_ALLOCATOR = hasattr(triton, "set_allocator") -def supports_tma(): - return _SUPPORTS_TMA +def supports_tma(device = None): + global _SUPPORTS_TMA + key = device + if key in _SUPPORTS_TMA: + return _SUPPORTS_TMA[key] + _SUPPORTS_TMA[key] = _check_tma_support(device) + return _SUPPORTS_TMA[key] # Helper to support allow_in_graph @@ -192,7 +212,7 @@ def grouped_gemm_forward( assert not permute_x, "Cannot use both use_tma_load_x and permute_x" use_tma = use_tma_load_w or use_tma_load_x or use_tma_store - if not supports_tma() and use_tma: + if not supports_tma(X.device) and use_tma: warnings.warn("TMA not supported, tma_load will be set to False") use_tma_load_w = False use_tma_load_x = False @@ -270,9 +290,8 @@ def grouped_gemm_forward( return (NUM_SMS,) if not autotune: - # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) - # BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) - pass + BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) + BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) if debug: print( @@ -399,7 +418,7 @@ def grouped_gemm_dX( assert not (permute_x and use_tma_store), "Cannot use both TMA store and permute_x" use_tma = use_tma_load_dy or use_tma_load_w or use_tma_store - if not supports_tma() and use_tma: + if not supports_tma(dY.device) and use_tma: warnings.warn("TMA not supported, tma_load will be set to False") use_tma_load_w = False use_tma_load_dy = False @@ -453,9 +472,8 @@ def grouped_gemm_dX( return (NUM_SMS,) if not autotune: - # BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N) - # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) - pass + BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N) + BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) if debug: print( @@ -567,7 +585,7 @@ def grouped_gemm_dW( assert not (permute_x and use_tma_load_x), "Cannot use both TMA load and permute_x" use_tma = use_tma_load_dy or use_tma_load_x or use_tma_store - if not supports_tma() and use_tma: + if not supports_tma(X.device) and use_tma: warnings.warn("TMA not supported, tma_load will be set to False") use_tma_load_x = False use_tma_load_dy = False @@ -607,9 +625,8 @@ def grouped_gemm_dW( dW = torch.zeros((num_experts, N, K), device = X.device, dtype = X.dtype) if not autotune: - # BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) - # BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) - pass + BLOCK_SIZE_N = min(N, BLOCK_SIZE_N) + BLOCK_SIZE_K = min(K, BLOCK_SIZE_K) def grid(META): return (NUM_SMS,) diff --git a/unsloth/models/qwen3_moe_triton.py b/unsloth/models/qwen3_moe_triton.py index 68b0ba522d..4f2a671b43 100644 --- a/unsloth/models/qwen3_moe_triton.py +++ b/unsloth/models/qwen3_moe_triton.py @@ -175,30 +175,13 @@ def Qwen3MoeSparseMoeBlock_triton_forward( down_weights ) # [num_experts, hidden_dim, intermediate_dim] - # Compute token counts and gather indices without array operations - expert_mask = torch.nn.functional.one_hot( - selected_experts, num_classes = num_experts - ).permute(2, 1, 0) - - token_counts_by_expert = expert_mask.sum(dim = 1).int() - - # Create gather indices for routing - avoid complex array operations - total_tokens = num_tokens * top_k - gather_indices = torch.zeros( - total_tokens, dtype = torch.long, device = hidden_states.device - ) - - # Simple sequential assignment for gather indices - current_idx = 0 - for expert_idx in range(num_experts): - expert_tokens = expert_mask[expert_idx].sum(dim = 0).bool() - num_expert_tokens = expert_tokens.sum().item() - if num_expert_tokens > 0: - expert_indices = torch.where(expert_tokens)[0] - gather_indices[current_idx : current_idx + num_expert_tokens] = ( - expert_indices - ) - current_idx += num_expert_tokens + # Compute token counts + gather indices in flattened top-k space. + # grouped_gemm expects gather_indices over selected_experts.view(-1). + flat_experts = selected_experts.view(-1) + token_counts_by_expert = torch.bincount( + flat_experts, minlength = num_experts + ).to(torch.int32) + gather_indices = flat_experts.argsort(stable = True) # First grouped GEMM: gate_up projection intermediate_states = grouped_gemm(