diff --git a/sonicmoe/functional/__init__.py b/sonicmoe/functional/__init__.py index eebead21..d69dfff0 100644 --- a/sonicmoe/functional/__init__.py +++ b/sonicmoe/functional/__init__.py @@ -6,12 +6,14 @@ import torch import torch.nn.functional as F -from quack.gemm_interface import gemm, gemm_dgated, gemm_gated +from quack.gemm_interface import gemm, gemm_add_inplace, gemm_dgated, gemm_gated from ..enums import ActivationType, is_glu from .backward import ( _down_projection_backward_act, _token_broadcast_backward, + _sigmoid_over_topk_bwd, + _topk_over_sigmoid_bwd, _topk_softmax_bwd, _up_projection_backward_act, ) @@ -19,11 +21,63 @@ from .triton_kernels import TC_topk_router_metadata_triton, general_routing_router_metadata_triton +def _base_parameter_from_view(tensor: torch.Tensor) -> torch.nn.Parameter | None: + base = getattr(tensor, "_base", None) + if isinstance(base, torch.nn.Parameter): + return base + if isinstance(tensor, torch.nn.Parameter): + return tensor + return None + + +def _megatron_main_grad_buffer(weight_view: torch.Tensor) -> torch.Tensor | None: + if torch.compiler.is_compiling(): + return None + base = _base_parameter_from_view(weight_view) + if base is None: + return None + if not hasattr(base, "grad_added_to_main_grad"): + return None + main_grad = getattr(base, "main_grad", None) + if not isinstance(main_grad, torch.Tensor): + return None + if main_grad.shape != base.shape or main_grad.device != base.device: + return None + return main_grad + + +def _megatron_wgrad_return(weight_view: torch.Tensor) -> torch.Tensor | None: + base = _base_parameter_from_view(weight_view) + if base is None: + return None + if not hasattr(base, "grad_added_to_main_grad"): + return None + base.grad_added_to_main_grad = True + + # Megatron's DDP hook needs param.grad to be non-None to mark the bucket + # ready, but the real dW has already been accumulated into main_grad. + dummy = ( + torch.zeros((), dtype=base.dtype, device=base.device) + if getattr(base, "zero_out_wgrad", False) + else torch.empty((), dtype=base.dtype, device=base.device) + ) + base.grad = dummy.expand_as(base) + return None + + class TC_Softmax_Topk_Router_Function(torch.autograd.Function): @staticmethod def forward( - ctx, router_logits: torch.Tensor, E: int, K: int, is_softmax_over_topk: bool, norm_topk_probs: bool + ctx, + router_logits: torch.Tensor, + E: int, + K: int, + is_softmax_over_topk: bool, + norm_topk_probs: bool, + router_score_function: str, ) -> tuple[torch.Tensor, torch.Tensor]: + if router_score_function not in {"softmax", "sigmoid"}: + raise ValueError(f"unexpected router_score_function ({router_score_function})") T = router_logits.size(0) topk_router_score = torch.empty(T, K, dtype=torch.float32, device=router_logits.device) @@ -37,6 +91,7 @@ def forward( K, is_softmax_over_topk=is_softmax_over_topk, norm_topk_probs=norm_topk_probs, + is_sigmoid_router=(router_score_function == "sigmoid"), ) # Save router_logits for topk(softmax()) backward (recompute full softmax). @@ -46,6 +101,7 @@ def forward( ctx.dtype = router_logits.dtype ctx.is_softmax_over_topk = is_softmax_over_topk ctx.norm_topk_probs = norm_topk_probs + ctx.router_score_function = router_score_function return topk_router_score, topk_router_indices @@ -56,20 +112,24 @@ def backward(ctx, dtopk_score: torch.Tensor, _: torch.Tensor): topk_router_score, topk_router_indices, router_logits = ctx.saved_tensors dlogits = torch.zeros(T, ctx.E, dtype=ctx.dtype, device=topk_router_score.device) - _topk_softmax_bwd( - router_logits, - dlogits, - None, - dtopk_score, - topk_router_score, - topk_router_indices, - E, - K, - is_softmax_over_topk=ctx.is_softmax_over_topk, - norm_topk_probs=ctx.norm_topk_probs, - ) - - return dlogits, None, None, None, None + if ctx.router_score_function == "sigmoid": + sigmoid_bwd = _sigmoid_over_topk_bwd if ctx.is_softmax_over_topk else _topk_over_sigmoid_bwd + sigmoid_bwd(dlogits, dtopk_score, topk_router_score, topk_router_indices, K) + else: + _topk_softmax_bwd( + router_logits, + dlogits, + None, + dtopk_score, + topk_router_score, + topk_router_indices, + E, + K, + is_softmax_over_topk=ctx.is_softmax_over_topk, + norm_topk_probs=ctx.norm_topk_probs, + ) + + return dlogits, None, None, None, None, None class _UpProjection(torch.autograd.Function): @@ -90,6 +150,7 @@ def forward( activation_type: ActivationType, is_inference_mode_enabled: bool, concat_layout: bool = False, + accumulate_wgrad_into_main_grad: bool = False, ) -> torch.Tensor: T, H = x.shape I, H, E = w1.shape @@ -131,6 +192,7 @@ def forward( ctx.is_each_token_has_variable_activated_experts = is_each_token_has_variable_activated_experts ctx.is_glu_activation = is_glu_activation ctx.concat_layout = concat_layout + ctx.accumulate_wgrad_into_main_grad = accumulate_wgrad_into_main_grad ctx.save_for_backward( x, @@ -158,6 +220,7 @@ def backward(ctx, _: None, dh: torch.Tensor): is_glu_activation = ctx.is_glu_activation is_each_token_has_variable_activated_experts = ctx.is_each_token_has_variable_activated_experts concat_layout = ctx.concat_layout + accumulate_wgrad_into_main_grad = ctx.accumulate_wgrad_into_main_grad ( x, @@ -171,7 +234,8 @@ def backward(ctx, _: None, dh: torch.Tensor): ) = ctx.saved_tensors dx_expanded = torch.empty(TK, H, dtype=dh.dtype, device=dh.device) - dw1 = torch.empty_like(w1) + dw1_main_grad = _megatron_main_grad_buffer(w1) if accumulate_wgrad_into_main_grad else None + dw1 = None if dw1_main_grad is not None else torch.empty_like(w1) db1 = None if b1 is None else torch.empty_like(b1) _up_projection_backward_act( @@ -184,16 +248,30 @@ def backward(ctx, _: None, dh: torch.Tensor): concat_layout=concat_layout, ) - gemm( - x.T, - dh, - out=dw1.permute(2, 1, 0), - cu_seqlens_k=expert_frequency_offset, - A_idx=x_gather_idx, - batch_idx_permute=None, - dynamic_scheduler=False, - concat_layout=(("out",) if concat_layout else None), - ) + if dw1_main_grad is None: + gemm( + x.T, + dh, + out=dw1.permute(2, 1, 0), + cu_seqlens_k=expert_frequency_offset, + A_idx=x_gather_idx, + batch_idx_permute=None, + dynamic_scheduler=False, + concat_layout=(("out",) if concat_layout else None), + ) + dw1_return = dw1 + else: + gemm_add_inplace( + x.T, + dh, + dw1_main_grad.permute(0, 2, 1), + cu_seqlens_k=expert_frequency_offset, + A_idx=x_gather_idx, + batch_idx_permute=None, + dynamic_scheduler=False, + concat_layout=(("out",) if concat_layout else None), + ) + dw1_return = _megatron_wgrad_return(w1) dx_reduced = torch.empty(T, H, dtype=dh.dtype, device=dh.device) @@ -207,7 +285,7 @@ def backward(ctx, _: None, dh: torch.Tensor): is_varlen_K=is_each_token_has_variable_activated_experts, ) - return dx_reduced, dw1, db1, *[None] * 13 + return dx_reduced, dw1_return, db1, *[None] * 13 class _DownProjection(torch.autograd.Function): @@ -228,6 +306,7 @@ def forward( num_activated_expert_per_token_offset: torch.Tensor, is_varlen_K: bool, activation_type: ActivationType, + accumulate_wgrad_into_main_grad: bool = False, ) -> torch.Tensor: TK = a.size(0) H, I, E = w2.shape @@ -254,6 +333,7 @@ def forward( ctx.K = K ctx.is_varlen_K = is_varlen_K ctx.activation_type = activation_type + ctx.accumulate_wgrad_into_main_grad = accumulate_wgrad_into_main_grad ctx.save_for_backward( h, @@ -273,6 +353,7 @@ def backward(ctx, dout: torch.Tensor): K = ctx.K is_varlen_K = ctx.is_varlen_K activation_type = ctx.activation_type + accumulate_wgrad_into_main_grad = ctx.accumulate_wgrad_into_main_grad ( h, @@ -284,7 +365,8 @@ def backward(ctx, dout: torch.Tensor): s_scatter_idx, ) = ctx.saved_tensors - dw2 = torch.empty_like(w2) + dw2_main_grad = _megatron_main_grad_buffer(w2) if accumulate_wgrad_into_main_grad else None + dw2 = None if dw2_main_grad is not None else torch.empty_like(w2) db2 = None if b2 is None else torch.empty_like(b2) dh = torch.empty_like(h) @@ -310,21 +392,34 @@ def backward(ctx, dout: torch.Tensor): activation_type=activation_type.value, ) - gemm( - dout.T, - a_prime, - out=dw2.permute(2, 0, 1), - cu_seqlens_k=expert_frequency_offset, - A_idx=x_gather_idx, - batch_idx_permute=None, - dynamic_scheduler=False, - ) + if dw2_main_grad is None: + gemm( + dout.T, + a_prime, + out=dw2.permute(2, 0, 1), + cu_seqlens_k=expert_frequency_offset, + A_idx=x_gather_idx, + batch_idx_permute=None, + dynamic_scheduler=False, + ) + dw2_return = dw2 + else: + gemm_add_inplace( + dout.T, + a_prime, + dw2_main_grad, + cu_seqlens_k=expert_frequency_offset, + A_idx=x_gather_idx, + batch_idx_permute=None, + dynamic_scheduler=False, + ) + dw2_return = _megatron_wgrad_return(w2) # TC top-K routing if not is_varlen_K: ds = ds.view(T, K) - return None, dh, dw2, db2, ds, *[None] * 10 + return None, dh, dw2_return, db2, ds, *[None] * 10 def moe_TC_softmax_topk_layer( @@ -340,7 +435,9 @@ def moe_TC_softmax_topk_layer( is_inference_mode_enabled: bool = False, is_softmax_over_topk: bool = True, norm_topk_probs: bool = False, + router_score_function: str = "softmax", concat_layout: bool = False, + accumulate_wgrad_into_main_grad: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: assert ((b1 is None) and (b2 is None)) or ( (b1 is not None) and (b2 is not None) @@ -348,7 +445,7 @@ def moe_TC_softmax_topk_layer( E = router_w.size(0) router_logits = F.linear(x, router_w) topk_scores, topk_indices = TC_Softmax_Topk_Router_Function.apply( - router_logits, E, K, is_softmax_over_topk, norm_topk_probs + router_logits, E, K, is_softmax_over_topk, norm_topk_probs, router_score_function ) T, K = topk_indices.size() @@ -386,6 +483,7 @@ def moe_TC_softmax_topk_layer( activation_type, is_inference_mode_enabled, concat_layout, + accumulate_wgrad_into_main_grad, ) o = _DownProjection.apply( @@ -403,6 +501,7 @@ def moe_TC_softmax_topk_layer( None, False, # is_each_token_has_variable_activated_expert activation_type, + accumulate_wgrad_into_main_grad, ) return o, router_logits, expert_frequency @@ -433,6 +532,7 @@ def moe_general_routing_inputs( activation_type: ActivationType, is_inference_mode_enabled: bool = False, concat_layout: bool = False, + accumulate_wgrad_into_main_grad: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: assert ((b1 is None) and (b2 is None)) or ( (b1 is not None) and (b2 is not None) @@ -484,6 +584,7 @@ def moe_general_routing_inputs( activation_type, is_inference_mode_enabled, concat_layout, + accumulate_wgrad_into_main_grad, ) o = _DownProjection.apply( @@ -501,6 +602,7 @@ def moe_general_routing_inputs( num_activated_expert_per_token_offset, True, # is_each_token_has_variable_activated_expert activation_type, + accumulate_wgrad_into_main_grad, ) return o, expert_frequency diff --git a/sonicmoe/functional/backward.py b/sonicmoe/functional/backward.py index ef5e483a..2ca87e5c 100644 --- a/sonicmoe/functional/backward.py +++ b/sonicmoe/functional/backward.py @@ -5,10 +5,16 @@ from typing import Optional import cuda.bindings.driver as cuda +import cutlass import cutlass.cute as cute +import math import torch import triton import triton.language as tl +from quack.cache.jit import jit_cache +from quack.compile_utils import make_fake_tensor as fake_tensor +from quack.cute_dsl_utils import torch2cute_dtype_map +from quack.dsl.torch_library_op import cute_op from quack.gemm_interface import gemm, gemm_dgated from ..enums import LIBRARY_NAME @@ -518,6 +524,198 @@ def _topk_softmax_bwd( ) +@triton.jit +def _sigmoid_topk_bwd_triton_kernel( + dlogits_full_ptr, + dscore_ptr, + score_ptr, + idx_ptr, + stride_dm: tl.constexpr, + stride_dn: tl.constexpr, + stride_gm: tl.constexpr, + stride_gk: tl.constexpr, + stride_sm: tl.constexpr, + stride_sk: tl.constexpr, + stride_im: tl.constexpr, + stride_ik: tl.constexpr, + K: tl.constexpr, + BLOCK_K: tl.constexpr, +): + row = tl.program_id(axis=0) + + k_offs = tl.arange(0, BLOCK_K) + k_mask = k_offs < K + + idx = tl.load(idx_ptr + row * stride_im + k_offs * stride_ik, mask=k_mask, other=0).to(tl.int32) + dscore = tl.load(dscore_ptr + row * stride_gm + k_offs * stride_gk, mask=k_mask, other=0).to(tl.float32) + score = tl.load(score_ptr + row * stride_sm + k_offs * stride_sk, mask=k_mask, other=0).to(tl.float32) + + dlogit = dscore * score * (1.0 - score) + tl.store(dlogits_full_ptr + row * stride_dm + idx * stride_dn, dlogit, mask=k_mask) + + +def _run_sigmoid_topk_bwd_triton( + dlogits_full: torch.Tensor, + dtopk_score: torch.Tensor, + topk_router_score: torch.Tensor, + topk_router_indices: torch.Tensor, + K: int, +) -> None: + T = dtopk_score.shape[0] + + _sigmoid_topk_bwd_triton_kernel[T,]( + dlogits_full, + dtopk_score, + topk_router_score, + topk_router_indices, + dlogits_full.stride(0), + dlogits_full.stride(1), + dtopk_score.stride(0), + dtopk_score.stride(1), + topk_router_score.stride(0), + topk_router_score.stride(1), + topk_router_indices.stride(0), + topk_router_indices.stride(1), + K, + triton.next_power_of_2(K), + ) + + +class _SigmoidTopKBackwardCute: + def __init__(self, K: int) -> None: + self.K = K + self.num_threads = 128 + self.rows_per_block = self.num_threads + + @cute.jit + def __call__( + self, + mDLogits: cute.Tensor, + mDScore: cute.Tensor, + mScore: cute.Tensor, + mIdx: cute.Tensor, + stream: cuda.CUstream, + ): + assert mDScore.element_type == mScore.element_type + assert mIdx.element_type == cutlass.Int32 + self.kernel(mDLogits, mDScore, mScore, mIdx).launch( + grid=[cute.ceil_div(mDScore.shape[0], self.rows_per_block), 1, 1], + block=[self.num_threads, 1, 1], + stream=stream, + ) + + @cute.kernel + def kernel( + self, + mDLogits: cute.Tensor, + mDScore: cute.Tensor, + mScore: cute.Tensor, + mIdx: cute.Tensor, + ): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + row = bidx * self.rows_per_block + tidx + + if row < mDScore.shape[0]: + for kidx in cutlass.range_constexpr(self.K): + idx = cutlass.Int32(mIdx[row, kidx]) + dscore = cutlass.Float32(mDScore[row, kidx]) + score = cutlass.Float32(mScore[row, kidx]) + dlogit = dscore * score * (cutlass.Float32(1.0) - score) + mDLogits[row, idx] = dlogit.to(mDLogits.element_type) + + @staticmethod + @jit_cache + def compile(op_cls, dlogits_dtype, dscore_dtype, score_dtype, idx_dtype, E: int, K: int): + batch_sym = cute.sym_int() + dlogits_div = math.gcd(128 // dlogits_dtype.width, E) + dscore_div = math.gcd(128 // dscore_dtype.width, K) + score_div = math.gcd(128 // score_dtype.width, K) + idx_div = math.gcd(128 // idx_dtype.width, K) + dlogits_cute = fake_tensor(dlogits_dtype, (batch_sym, E), dlogits_div) + dscore_cute = fake_tensor(dscore_dtype, (batch_sym, K), dscore_div) + score_cute = fake_tensor(score_dtype, (batch_sym, K), score_div) + idx_cute = fake_tensor(idx_dtype, (batch_sym, K), idx_div) + return cute.compile( + op_cls(K), + dlogits_cute, + dscore_cute, + score_cute, + idx_cute, + cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True), + options="--enable-tvm-ffi", + ) + + +class SigmoidOverTopKBackwardCute(_SigmoidTopKBackwardCute): + pass + + +class TopKOverSigmoidBackwardCute(_SigmoidTopKBackwardCute): + pass + + +def _run_sigmoid_topk_bwd_cute( + op_cls, + dlogits_full: torch.Tensor, + dtopk_score: torch.Tensor, + topk_router_score: torch.Tensor, + topk_router_indices: torch.Tensor, + K: int, +) -> None: + compile_args = ( + op_cls, + torch2cute_dtype_map[dlogits_full.dtype], + torch2cute_dtype_map[dtopk_score.dtype], + torch2cute_dtype_map[topk_router_score.dtype], + torch2cute_dtype_map[topk_router_indices.dtype], + dlogits_full.size(1), + K, + ) + _SigmoidTopKBackwardCute.compile(*compile_args)( + dlogits_full, + dtopk_score, + topk_router_score, + topk_router_indices, + ) + + +@cute_op(f"{LIBRARY_NAME}::_sigmoid_over_topk_bwd", mutates_args={"dlogits_full"}) +def _sigmoid_over_topk_bwd( + dlogits_full: torch.Tensor, + dtopk_score: torch.Tensor, + topk_router_score: torch.Tensor, + topk_router_indices: torch.Tensor, + K: int, +) -> None: + _run_sigmoid_topk_bwd_cute( + SigmoidOverTopKBackwardCute, + dlogits_full, + dtopk_score, + topk_router_score, + topk_router_indices, + K, + ) + + +@cute_op(f"{LIBRARY_NAME}::_topk_over_sigmoid_bwd", mutates_args={"dlogits_full"}) +def _topk_over_sigmoid_bwd( + dlogits_full: torch.Tensor, + dtopk_score: torch.Tensor, + topk_router_score: torch.Tensor, + topk_router_indices: torch.Tensor, + K: int, +) -> None: + _run_sigmoid_topk_bwd_cute( + TopKOverSigmoidBackwardCute, + dlogits_full, + dtopk_score, + topk_router_score, + topk_router_indices, + K, + ) + + @triton.jit def _topk_bwd_scatter_small_kernel( dlogits_full_ptr, diff --git a/sonicmoe/functional/forward.py b/sonicmoe/functional/forward.py index 95f11b2e..b812e01f 100644 --- a/sonicmoe/functional/forward.py +++ b/sonicmoe/functional/forward.py @@ -12,7 +12,7 @@ from ..enums import LIBRARY_NAME from .reduction_over_k_gather import token_gather_and_sum_varlen_K_triton -from .topk import Softmax_Over_TopK, TopK_Over_Softmax +from .topk import Sigmoid_Over_TopK, Softmax_Over_TopK, TopK_Over_Sigmoid, TopK_Over_Softmax @torch.library.custom_op(f"{LIBRARY_NAME}::_topk_fwd", mutates_args={"values", "indices"}) @@ -23,6 +23,7 @@ def _topk_fwd( indices: torch.Tensor, is_softmax_over_topk: bool, norm_topk_probs: bool, + is_sigmoid_router: bool, ) -> None: """Top-k forward pass. Args: @@ -41,13 +42,21 @@ def _topk_fwd( x_tensor, values_tensor, indices_tensor = [convert_from_dlpack(tensor) for tensor in (x, values, indices)] current_stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) - if is_softmax_over_topk: + if is_sigmoid_router: + compile_key = (input_dtype, output_dtype, N, k, "sigmoid", is_softmax_over_topk) + elif is_softmax_over_topk: compile_key = (input_dtype, output_dtype, N, k, True) else: compile_key = (input_dtype, output_dtype, N, k, False, norm_topk_probs) if compile_key not in _topk_fwd.compile_cache: - if is_softmax_over_topk: + if is_sigmoid_router: + topk_op = ( + Sigmoid_Over_TopK(input_dtype, output_dtype, N, k) + if is_softmax_over_topk + else TopK_Over_Sigmoid(input_dtype, output_dtype, N, k) + ) + elif is_softmax_over_topk: topk_op = Softmax_Over_TopK(input_dtype, output_dtype, N, k) else: topk_op = TopK_Over_Softmax(input_dtype, output_dtype, N, k, norm_topk_probs) @@ -115,6 +124,7 @@ def _topk_softmax_fwd( K: int, is_softmax_over_topk: bool, norm_topk_probs: bool, + is_sigmoid_router: bool = False, ) -> None: if E <= 4096 and K <= 16 and E % 8 == 0: _topk_fwd( @@ -124,9 +134,19 @@ def _topk_softmax_fwd( topk_router_indices, is_softmax_over_topk=is_softmax_over_topk, norm_topk_probs=norm_topk_probs, + is_sigmoid_router=is_sigmoid_router, ) else: - if is_softmax_over_topk: + if is_sigmoid_router: + if is_softmax_over_topk: + topk_results = router_logits.topk(K, dim=-1) + vals = topk_results.values.sigmoid() + else: + topk_results = router_logits.sigmoid().topk(K, dim=-1) + vals = topk_results.values + topk_router_score.copy_(vals.to(topk_router_score.dtype)) + topk_router_indices.copy_(topk_results.indices.to(topk_router_indices.dtype)) + elif is_softmax_over_topk: topk_results = router_logits.topk(K, dim=-1) vals = topk_results.values.softmax(dim=-1, dtype=torch.float32) topk_router_score.copy_(vals.to(topk_router_score.dtype)) diff --git a/sonicmoe/functional/topk.py b/sonicmoe/functional/topk.py index 043d0f2f..d8144d82 100644 --- a/sonicmoe/functional/topk.py +++ b/sonicmoe/functional/topk.py @@ -22,6 +22,8 @@ class _TopKMode(Enum): SOFTMAX_OVER_TOPK = "softmax_over_topk" # most common choice: softmax(topk(x)) TOPK_OVER_SOFTMAX = "topk_over_softmax" # Qwen3: topk(softmax(x)) + SIGMOID_OVER_TOPK = "sigmoid_over_topk" # sigmoid(topk(x)) + TOPK_OVER_SIGMOID = "topk_over_sigmoid" # topk(sigmoid(x)) TOPK_NO_FUSION = "topk" @@ -197,6 +199,13 @@ def kernel( for i in cutlass.range_constexpr(cute.size(tXrX_f32)): tXrX_f32[i] = tXrX_f32[i] - log_normalizer + # Sigmoid-then-TopK: full-row sigmoid before top-k. + if const_expr(self.mode == _TopKMode.TOPK_OVER_SIGMOID): + if const_expr((not is_even_N) or (self.N != self.next_power_of_2_N)): + utils.fill_oob(tXrX_f32, tXpX, -tXrX_f32.element_type.inf) + for i in cutlass.range_constexpr(cute.size(tXrX_f32)): + tXrX_f32[i] = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-tXrX_f32[i])) + # Encode indices into mantissa low bits. log_N = int(math.log2(self.next_power_of_2_N)) idx_mask = const_expr((1 << log_N) - 1) @@ -259,6 +268,11 @@ def kernel( for i in cutlass.range_constexpr(self.k): topk_vals[i] = topk_vals[i] / topk_sum + # TopK-then-Sigmoid: select on raw logits, then transform selected gates. + if const_expr(self.mode == _TopKMode.SIGMOID_OVER_TOPK): + for i in cutlass.range_constexpr(self.k): + topk_vals[i] = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-topk_vals[i])) + topk_vals_out = cute.make_rmem_tensor_like(topk_indices, mValues.element_type) for i in cutlass.range_constexpr(self.k): topk_vals_out[i] = topk_vals[i].to(mValues.element_type) @@ -338,3 +352,41 @@ def __init__( k=k, mode=_TopKMode.TOPK_NO_FUSION, ) + + +class Sigmoid_Over_TopK(_TopK): + """sigmoid(topk(x)) without a separate sigmoid kernel.""" + + def __init__( + self, + input_dtype: Type[cutlass.Numeric], + output_dtype: Type[cutlass.Numeric], + N: int, + k: int, + ): + super().__init__( + input_dtype=input_dtype, + output_dtype=output_dtype, + N=N, + k=k, + mode=_TopKMode.SIGMOID_OVER_TOPK, + ) + + +class TopK_Over_Sigmoid(_TopK): + """topk(sigmoid(x)) without a separate sigmoid kernel.""" + + def __init__( + self, + input_dtype: Type[cutlass.Numeric], + output_dtype: Type[cutlass.Numeric], + N: int, + k: int, + ): + super().__init__( + input_dtype=input_dtype, + output_dtype=output_dtype, + N=N, + k=k, + mode=_TopKMode.TOPK_OVER_SIGMOID, + ) diff --git a/sonicmoe/moe.py b/sonicmoe/moe.py index b6b3da11..8accb015 100644 --- a/sonicmoe/moe.py +++ b/sonicmoe/moe.py @@ -2,7 +2,7 @@ # Copyright (c) 2025, Wentao Guo, Mayank Mishra, Xinle Cheng, Ion Stoica, Tri Dao # ******************************************************************************** -from typing import Callable +from typing import Callable, Literal import torch import torch.nn as nn @@ -12,6 +12,9 @@ from .functional import moe_TC_softmax_topk_layer +RouterScoreFunction = Literal["softmax", "sigmoid"] + + try: from xma.modules.moe import scattered_experts @@ -173,11 +176,17 @@ def __init__( activation_function: ActivationType, add_bias: bool, std: float, + router_score_function: RouterScoreFunction = "softmax", + router_score_over_topk: bool = True, ) -> None: super().__init__() + if router_score_function not in {"softmax", "sigmoid"}: + raise ValueError(f"unexpected router_score_function ({router_score_function})") self.num_experts = num_experts self.top_k = num_experts_per_tok + self.router_score_function = router_score_function + self.router_score_over_topk = router_score_over_topk self.hidden_size = hidden_size self.intermediate_size = intermediate_size @@ -209,7 +218,7 @@ def forward( hidden_states: torch.Tensor, kernel_backend_moe: KernelBackendMoE = KernelBackendMoE.sonicmoe, is_inference_mode: bool = False, - ) -> tuple[torch.Tensor, torch.Tensor]: + ) -> tuple[torch.Tensor, torch.Tensor | None]: original_shape = hidden_states.shape # hidden_states -> (batch_size, query_length, hidden_size) @@ -227,6 +236,9 @@ def forward( self.stream_id, self.activation_function, is_inference_mode or not self.training, + is_softmax_over_topk=self.router_score_over_topk, + router_score_function=self.router_score_function, + accumulate_wgrad_into_main_grad=getattr(self, "accumulate_wgrad_into_main_grad", False), ) else: # hidden_states -> (total_q, hidden_size) @@ -250,9 +262,14 @@ def forward( if is_inference_mode: aux_loss = None else: + aux_probs = ( + torch.sigmoid(router_logits.float()) + if self.router_score_function == "sigmoid" + else F.softmax(router_logits, dim=-1, dtype=torch.float32) + ) aux_loss = self._compute_switch_loss( logits=router_logits, - probs=F.softmax(router_logits, dim=-1, dtype=torch.float32), + probs=aux_probs, expert_frequency=expert_frequency, ) @@ -280,12 +297,26 @@ def _compute_routing_weights(self, hidden_states: torch.Tensor) -> tuple[torch.T router_logits = self.router(hidden_states) # router_logits -> (total_q, num_experts) - router_weights, selected_experts = self._get_topk(router_logits) + if self.router_score_over_topk: + router_weights, selected_experts = self._get_topk(router_logits) + elif self.router_score_function == "sigmoid": + router_weights, selected_experts = self._get_topk(torch.sigmoid(router_logits.float())) + elif self.router_score_function == "softmax": + router_weights, selected_experts = self._get_topk(F.softmax(router_logits, dim=-1, dtype=torch.float32)) + else: + raise ValueError(f"unexpected router_score_function ({self.router_score_function})") # router_weights -> (total_q, top_k) # selected_experts -> (total_q, top_k) - router_weights = F.softmax(router_weights.float(), dim=-1) + if not self.router_score_over_topk: + pass + elif self.router_score_function == "sigmoid": + router_weights = torch.sigmoid(router_weights.float()) + elif self.router_score_function == "softmax": + router_weights = F.softmax(router_weights.float(), dim=-1) + else: + raise ValueError(f"unexpected router_score_function ({self.router_score_function})") router_weights = router_weights.type_as(hidden_states) return router_logits, router_weights, selected_experts diff --git a/tests/megatron_moe_test.py b/tests/megatron_moe_test.py new file mode 100644 index 00000000..e8201270 --- /dev/null +++ b/tests/megatron_moe_test.py @@ -0,0 +1,104 @@ +# ******************************************************************************** +# Copyright (c) 2025, Wentao Guo, Mayank Mishra, Xinle Cheng, Ion Stoica, Tri Dao +# ******************************************************************************** + +import unittest + +import torch + +from sonicmoe import KernelBackendMoE, MoE +from sonicmoe.enums import ActivationType + +from .test_commons import TestCommons + + +_SEED = 123 +_SHAPE = (128, 256, 128, 8, 2) # T, H, I, E, K + + +class _MainGradAccumulatingMoE(MoE): + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + self.accumulate_wgrad_into_main_grad = True + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for SonicMoE kernels") +class MegatronMoETest(TestCommons): + def _make_moe(self, cls: type[MoE]) -> MoE: + _, H, I, E, K = _SHAPE + return cls( + num_experts=E, + num_experts_per_tok=K, + hidden_size=H, + intermediate_size=I, + activation_function=ActivationType.SWIGLU, + add_bias=False, + std=0.02, + ).to(device="cuda", dtype=torch.bfloat16) + + def test_accumulates_expert_wgrad_into_main_grad(self) -> None: + self.set_seed(_SEED) + T, H, _, _, _ = _SHAPE + + reference = self._make_moe(MoE) + megatron_moe = self._make_moe(_MainGradAccumulatingMoE) + megatron_moe.load_state_dict(reference.state_dict()) + + x = torch.randn(T, H, device="cuda", dtype=torch.bfloat16, requires_grad=True) + dy = torch.randn_like(x) + + # Normal MoE should keep the standard autograd contract even if a caller + # attaches Megatron-like fields. + for param in (reference.c_fc.weight, reference.c_proj.weight): + param.main_grad = torch.zeros_like(param) + param.grad_added_to_main_grad = False + + y_ref, _ = reference(x, kernel_backend_moe=KernelBackendMoE.sonicmoe) + ref_grads = torch.autograd.grad(y_ref, [x] + list(reference.parameters()), grad_outputs=dy) + ref_c_fc_wgrad = ref_grads[2] + ref_c_proj_wgrad = ref_grads[3] + + self.assertIsNotNone(ref_c_fc_wgrad) + self.assertIsNotNone(ref_c_proj_wgrad) + self.assertEqual(float(reference.c_fc.weight.main_grad.abs().max()), 0.0) + self.assertEqual(float(reference.c_proj.weight.main_grad.abs().max()), 0.0) + self.assertFalse(reference.c_fc.weight.grad_added_to_main_grad) + self.assertFalse(reference.c_proj.weight.grad_added_to_main_grad) + + # The Megatron-side subclass should write the real expert dW into the bucket-backed + # main_grad buffer and only leave a tiny dummy param.grad for the DDP hook. + for param in (megatron_moe.c_fc.weight, megatron_moe.c_proj.weight): + param.main_grad = torch.zeros_like(param) + param.grad_added_to_main_grad = False + + x_megatron = x.detach().clone().requires_grad_() + y_megatron, _ = megatron_moe(x_megatron, kernel_backend_moe=KernelBackendMoE.sonicmoe) + y_megatron.backward(dy) + torch.cuda.synchronize() + + self.assert_equal_tensors(megatron_moe.c_fc.weight.main_grad, ref_c_fc_wgrad, exact_match=True) + self.assert_equal_tensors(megatron_moe.c_proj.weight.main_grad, ref_c_proj_wgrad, exact_match=True) + self.assertTrue(megatron_moe.c_fc.weight.grad_added_to_main_grad) + self.assertTrue(megatron_moe.c_proj.weight.grad_added_to_main_grad) + self.assertLessEqual(megatron_moe.c_fc.weight.grad.untyped_storage().nbytes(), 4) + self.assertLessEqual(megatron_moe.c_proj.weight.grad.untyped_storage().nbytes(), 4) + + def test_falls_back_without_megatron_ddp_flag(self) -> None: + self.set_seed(_SEED) + T, H, _, _, _ = _SHAPE + + megatron_moe = self._make_moe(_MainGradAccumulatingMoE) + for param in (megatron_moe.c_fc.weight, megatron_moe.c_proj.weight): + param.main_grad = torch.zeros_like(param) + if hasattr(param, "grad_added_to_main_grad"): + delattr(param, "grad_added_to_main_grad") + + x = torch.randn(T, H, device="cuda", dtype=torch.bfloat16, requires_grad=True) + dy = torch.randn_like(x) + y, _ = megatron_moe(x, kernel_backend_moe=KernelBackendMoE.sonicmoe) + grads = torch.autograd.grad(y, [x] + list(megatron_moe.parameters()), grad_outputs=dy) + + self.assertIsNotNone(grads[2]) + self.assertIsNotNone(grads[3]) + self.assertEqual(float(megatron_moe.c_fc.weight.main_grad.abs().max()), 0.0) + self.assertEqual(float(megatron_moe.c_proj.weight.main_grad.abs().max()), 0.0)