diff --git a/csrc/apis/sm90_mega.hpp b/csrc/apis/sm90_mega.hpp index 0bebfa1b2..8ce3f19bc 100644 --- a/csrc/apis/sm90_mega.hpp +++ b/csrc/apis/sm90_mega.hpp @@ -38,7 +38,7 @@ get_symm_buffer_size_for_sm90_mega_moe( const bool& use_fp8_dispatch, const std::string& activation) { DG_HOST_ASSERT(num_experts % num_ranks == 0); DG_HOST_ASSERT(use_fp8_dispatch); - DG_HOST_ASSERT(activation == "swiglu"); + DG_HOST_ASSERT(activation == "swiglu" or activation == "swigluoai"); const auto workspace = layout::SM90Workspace( nullptr, num_ranks, num_experts, num_max_tokens_per_rank, num_topk); @@ -148,6 +148,8 @@ static void fp8_mega_moe( const int& num_experts, const int& num_topk, const std::tuple& recipe, const std::string& activation, + const float& activation_alpha, + const float& activation_up_bias, const std::optional& activation_clamp_opt, const bool& fast_math ) { @@ -160,7 +162,7 @@ static void fp8_mega_moe( const auto num_tokens = static_cast(y.size(0)); const auto [rm, rn, rk] = recipe; DG_HOST_ASSERT(rm == 128 and rn == 128 and rk == 128); - DG_HOST_ASSERT(activation == "swiglu"); + DG_HOST_ASSERT(activation == "swiglu" or activation == "swigluoai"); const auto activation_clamp = activation_clamp_opt.value_or(std::numeric_limits::infinity()); @@ -215,7 +217,8 @@ static void fp8_mega_moe( num_experts_per_rank, num_tokens, num_topk, hidden, intermediate_hidden, - activation_clamp, fast_math); + activation_clamp, activation_alpha, activation_up_bias, + fast_math); if (deep_jit::get_env("DG_COMM_KERNEL_DEBUG")) sym_buffer.zero_(); diff --git a/csrc/jit_kernels/impls/sm90_fp8_mega_moe.hpp b/csrc/jit_kernels/impls/sm90_fp8_mega_moe.hpp index 0f28f2446..d269a8d55 100644 --- a/csrc/jit_kernels/impls/sm90_fp8_mega_moe.hpp +++ b/csrc/jit_kernels/impls/sm90_fp8_mega_moe.hpp @@ -39,6 +39,8 @@ class SM90FP8MegaMoERuntime final { int num_experts, num_topk; int num_ranks; float activation_clamp; + float activation_alpha; + float activation_up_bias; bool fast_math; int epilogue_registers; bool reuse_accum_as_final; @@ -87,6 +89,8 @@ static void __instantiate_kernel() {{ {}, {}, {}, + {}, + {}, {}, {}, {}, {}, {}, {}, @@ -111,6 +115,8 @@ static void __instantiate_kernel() {{ args.config.num_dispatch_threads, args.config.num_non_epilogue_threads, args.config.num_epilogue_threads, args.launch_args.grid_dim->x, args.num_ranks, to_string(args.activation_clamp), + to_string(args.activation_alpha), + to_string(args.activation_up_bias), args.fast_math ? "true" : "false", args.epilogue_registers, args.reuse_accum_as_final ? "true" : "false", @@ -152,6 +158,8 @@ static void sm90_fp8_mega_moe( const int& num_tokens, const int& num_topk, const int& hidden, const int& intermediate_hidden, const float& activation_clamp, + const float& activation_alpha, + const float& activation_up_bias, const bool& fast_math ) { const auto num_ranks = static_cast(sym_buffer_ptrs.size()); @@ -281,6 +289,8 @@ static void sm90_fp8_mega_moe( .num_experts = num_experts, .num_topk = num_topk, .num_ranks = num_ranks, .activation_clamp = activation_clamp, + .activation_alpha = activation_alpha, + .activation_up_bias = activation_up_bias, .fast_math = fast_math, .epilogue_registers = epilogue_registers, .reuse_accum_as_final = reuse_accum_as_final, diff --git a/csrc/tvm_ffi_api.cpp b/csrc/tvm_ffi_api.cpp index afde8f843..083356eaa 100644 --- a/csrc/tvm_ffi_api.cpp +++ b/csrc/tvm_ffi_api.cpp @@ -826,7 +826,9 @@ void dg_bf16_mega_moe(TensorView y, TensorView l1_weights, TensorView l2_weights void dg_fp8_mega_moe(TensorView y, TensorView l1_weights, TensorView l1_weights_sf, TensorView l2_weights, TensorView l2_weights_sf, Optional cumulative_local_expert_recv_stats, TensorView sym_buffer, Array sym_buffer_ptrs, int64_t rank_idx, int64_t num_max_tokens_per_rank, int64_t num_experts, int64_t num_topk, - Tuple recipe, std::string activation, Optional activation_clamp_opt, bool fast_math) { + Tuple recipe, std::string activation, + double activation_alpha, double activation_up_bias, + Optional activation_clamp_opt, bool fast_math) { auto c_val = cumulative_local_expert_recv_stats.has_value()? std::optional(convert_to_torch_tensor(cumulative_local_expert_recv_stats.value())) : std::nullopt; auto act_clamp_opt_val = activation_clamp_opt.has_value()? std::optional(static_cast(activation_clamp_opt.value())) : std::nullopt; std::vector sym_buffer_ptrs_val; @@ -844,7 +846,9 @@ void dg_fp8_mega_moe(TensorView y, TensorView l1_weights, TensorView l1_weights_ std::make_pair(convert_to_torch_tensor(l2_weights), convert_to_torch_tensor(l2_weights_sf)), c_val, convert_to_torch_tensor(sym_buffer), sym_buffer_ptrs_val, static_cast(rank_idx), static_cast(num_max_tokens_per_rank), static_cast(num_experts), - static_cast(num_topk), recipe_val, activation, act_clamp_opt_val, fast_math + static_cast(num_topk), recipe_val, activation, + static_cast(activation_alpha), static_cast(activation_up_bias), + act_clamp_opt_val, fast_math ); } diff --git a/deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh b/deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh index 02c623a2f..f7db06e29 100644 --- a/deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh @@ -41,18 +41,20 @@ __forceinline__ __device__ float sm90_fp8_mega_moe_clamp_up(float x) { return x; } -template +template __forceinline__ __device__ float sm90_fp8_mega_moe_silu(float x) { - const float e = kFastMath ? __expf(-x) : expf(-x); + const float e = kFastMath ? __expf(-kActivationAlpha * x) : expf(-kActivationAlpha * x); const float sig = kFastMath ? math::fast_rcp(1.0f + e) : 1.0f / (1.0f + e); return x * sig; } -template +template __forceinline__ __device__ float sm90_fp8_mega_moe_swiglu(float g, float u) { g = sm90_fp8_mega_moe_clamp_gate(g); u = sm90_fp8_mega_moe_clamp_up(u); - return sm90_fp8_mega_moe_silu(g) * u; + return sm90_fp8_mega_moe_silu(g) * + (u + kActivationUpBias); } // Continuous FP32 activation scale. SM90 WGMMA has no hardware block-scale operand (the SF @@ -140,6 +142,8 @@ template < uint32_t kNumEpilogueThreads, uint32_t kNumSMs, uint32_t kNumRanks, float kActivationClamp, + float kActivationAlpha, + float kActivationUpBias, bool kFastMath, uint32_t kEpilogueRegisterBudget, bool kReuseAccumAsFinal, @@ -1627,7 +1631,7 @@ sm90_fp8_mega_moe_impl(void* y, if (block_phase == sched::BlockPhase::Linear1) { if constexpr (kSwapABActive) { auto silu = [](float x) -> float { - const float e = kFastMath ? __expf(-x) : expf(-x); + const float e = kFastMath ? __expf(-kActivationAlpha * x) : expf(-kActivationAlpha * x); const float sig = kFastMath ? math::fast_rcp(1.0f + e) : 1.0f / (1.0f + e); return x * sig; }; @@ -1654,7 +1658,7 @@ sm90_fp8_mega_moe_impl(void* y, .get_data_buffer(m_idx + token_0) .get_base_ptr(); smem_cd_swap_l1_fp32[token_0 * L1_OUT_BLOCK_N + out_col_base] = - silu(g0) * u0 * weight_0; + silu(g0) * (u0 + kActivationUpBias) * weight_0; } if (token_1 < valid_m) { float g1 = final_accum[i * 4 + 1]; @@ -1665,7 +1669,7 @@ sm90_fp8_mega_moe_impl(void* y, .get_data_buffer(m_idx + token_1) .get_base_ptr(); smem_cd_swap_l1_fp32[token_1 * L1_OUT_BLOCK_N + out_col_base] = - silu(g1) * u1 * weight_1; + silu(g1) * (u1 + kActivationUpBias) * weight_1; } }; @@ -1774,7 +1778,7 @@ sm90_fp8_mega_moe_impl(void* y, x = cute::min(cute::max(x, -kActivationClamp), kActivationClamp); }; auto silu = [](float x) -> float { - const float e = kFastMath ? __expf(-x) : expf(-x); + const float e = kFastMath ? __expf(-kActivationAlpha * x) : expf(-kActivationAlpha * x); const float sig = kFastMath ? math::fast_rcp(1.0f + e) : 1.0f / (1.0f + e); return x * sig; }; @@ -1801,16 +1805,16 @@ sm90_fp8_mega_moe_impl(void* y, clamp_up(u_r1_c1); if (valid_r0) { - swiglu_r0[p][0] = silu(g_r0_c0) * u_r0_c0; - swiglu_r0[p][1] = silu(g_r0_c1) * u_r0_c1; + swiglu_r0[p][0] = silu(g_r0_c0) * (u_r0_c0 + kActivationUpBias); + swiglu_r0[p][1] = silu(g_r0_c1) * (u_r0_c1 + kActivationUpBias); amax_r0 = cute::max(amax_r0, cute::max(cute::abs(swiglu_r0[p][0]), cute::abs(swiglu_r0[p][1]))); } else { swiglu_r0[p][0] = 0.0f; swiglu_r0[p][1] = 0.0f; } if (valid_r1) { - swiglu_r1[p][0] = silu(g_r1_c0) * u_r1_c0; - swiglu_r1[p][1] = silu(g_r1_c1) * u_r1_c1; + swiglu_r1[p][0] = silu(g_r1_c0) * (u_r1_c0 + kActivationUpBias); + swiglu_r1[p][1] = silu(g_r1_c1) * (u_r1_c1 + kActivationUpBias); amax_r1 = cute::max(amax_r1, cute::max(cute::abs(swiglu_r1[p][0]), cute::abs(swiglu_r1[p][1]))); } else { swiglu_r1[p][0] = 0.0f; diff --git a/sgl_deep_gemm/__init__.py b/sgl_deep_gemm/__init__.py index 671caf418..1db60be04 100644 --- a/sgl_deep_gemm/__init__.py +++ b/sgl_deep_gemm/__init__.py @@ -515,6 +515,8 @@ def fp8_mega_moe(y: torch.Tensor, cumulative_local_expert_recv_stats: Optional[torch.Tensor] = None, recipe: Tuple[int, int, int] = (128, 128, 128), activation: str = 'swiglu', + activation_alpha: float = 1.0, + activation_up_bias: float = 0.0, activation_clamp: Optional[float] = None, fast_math: bool = True): (l1_weights_data, l1_weights_sf) = l1_weights @@ -529,7 +531,8 @@ def fp8_mega_moe(y: torch.Tensor, sym_buffer.num_max_tokens_per_rank, sym_buffer.num_experts, sym_buffer.num_topk, recipe, - activation, activation_clamp, + activation, float(activation_alpha), float(activation_up_bias), + activation_clamp, fast_math ) diff --git a/sgl_deep_gemm/tests/test_mega_moe_hopper.py b/sgl_deep_gemm/tests/test_mega_moe_hopper.py index 4b3ff8331..4f23f2940 100644 --- a/sgl_deep_gemm/tests/test_mega_moe_hopper.py +++ b/sgl_deep_gemm/tests/test_mega_moe_hopper.py @@ -275,7 +275,12 @@ def _dequant_per_token_per_128_k(x_fp8: torch.Tensor, sf: torch.Tensor) -> torch return (x_view * sf.unsqueeze(-1)).view(m, k) -def _swiglu_fp32(gate_up: torch.Tensor, clamp: float) -> torch.Tensor: +def _swiglu_fp32( + gate_up: torch.Tensor, + clamp: float, + alpha: float = 1.0, + up_bias: float = 0.0, +) -> torch.Tensor: """SwiGLU matching the fused SM90 path's clamp semantics.""" n2 = gate_up.size(-1) half = n2 // 2 @@ -283,7 +288,7 @@ def _swiglu_fp32(gate_up: torch.Tensor, clamp: float) -> torch.Tensor: if math.isfinite(clamp): gate = gate.clamp(max=clamp) up = up.clamp(min=-clamp, max=clamp) - return torch.nn.functional.silu(gate) * up + return gate * torch.sigmoid(alpha * gate) * (up + up_bias) def _reference_fused( @@ -303,6 +308,8 @@ def _reference_fused( hidden: int, intermediate_hidden: int, activation_clamp: float, + activation_alpha: float = 1.0, + activation_up_bias: float = 0.0, ) -> torch.Tensor: """PyTorch BF16/FP32 reference for this rank's fused output.""" num_experts_per_rank = num_experts // num_ranks @@ -357,7 +364,9 @@ def _reference_fused( l1_y = torch.einsum("sk,snk->sn", x_sel, l1_w_sel) del l1_w_sel - l1_y = _swiglu_fp32(l1_y, activation_clamp) * weights.unsqueeze(-1) + l1_y = _swiglu_fp32( + l1_y, activation_clamp, activation_alpha, activation_up_bias + ) * weights.unsqueeze(-1) s, ih = l1_y.shape assert ih == intermediate_hidden and ih % 64 == 0 l1_view = l1_y.view(s, ih // 64, 64) @@ -401,6 +410,9 @@ def _run_accuracy_scenario( num_topk = cfg["num_topk"] masked_ratio = cfg.get("masked_ratio", 0.0) activation_clamp = cfg.get("activation_clamp", 10.0) + activation = cfg.get("activation", "swiglu") + activation_alpha = cfg.get("activation_alpha", 1.0) + activation_up_bias = cfg.get("activation_up_bias", 0.0) fast_math = cfg.get("fast_math", True) assert num_experts % num_ranks == 0, ( @@ -460,6 +472,7 @@ def trace(stage: str): num_topk, hidden, intermediate_hidden, + activation=activation, ) cum_stats = torch.zeros((num_experts_per_rank,), dtype=torch.int, device="cuda") @@ -479,7 +492,9 @@ def run_fused(): buffer, cumulative_local_expert_recv_stats=cum_stats, recipe=(128, 128, 128), - activation="swiglu", + activation=activation, + activation_alpha=activation_alpha, + activation_up_bias=activation_up_bias, activation_clamp=activation_clamp if math.isfinite(activation_clamp) else None, fast_math=fast_math, ) @@ -506,6 +521,8 @@ def run_fused(): hidden, intermediate_hidden, activation_clamp, + activation_alpha, + activation_up_bias, ) diff = calc_diff(y_fused, y_ref) @@ -575,7 +592,17 @@ def check_reused_output(label: str): def _accuracy_layer1_smoke() -> List[Tuple[str, Dict[str, Any]]]: - return [("L1.smoke", dict(_ACCURACY_SMOKE))] + oai_swiglu = dict( + _ACCURACY_SMOKE, + activation="swigluoai", + activation_alpha=1.702, + activation_up_bias=1.0, + activation_clamp=7.0, + ) + return [ + ("L1.smoke", dict(_ACCURACY_SMOKE)), + ("L1.oai_swiglu", oai_swiglu), + ] def _accuracy_layer2_heuristic_branches(num_ranks: int) -> List[Tuple[str, Dict[str, Any]]]: