Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions csrc/apis/sm90_mega.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -148,6 +148,8 @@ static void fp8_mega_moe(
const int& num_experts, const int& num_topk,
const std::tuple<int, int, int>& recipe,
const std::string& activation,
const float& activation_alpha,
const float& activation_up_bias,
const std::optional<float>& activation_clamp_opt,
const bool& fast_math
) {
Expand All @@ -160,7 +162,7 @@ static void fp8_mega_moe(
const auto num_tokens = static_cast<int>(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<float>::infinity());
Expand Down Expand Up @@ -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<int>("DG_COMM_KERNEL_DEBUG"))
sym_buffer.zero_();
Expand Down
10 changes: 10 additions & 0 deletions csrc/jit_kernels/impls/sm90_fp8_mega_moe.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -87,6 +89,8 @@ static void __instantiate_kernel() {{
{},
{},
{},
{},
{},
{}, {}, {},
{}, {},
{},
Expand All @@ -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",
Expand Down Expand Up @@ -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<int>(sym_buffer_ptrs.size());
Expand Down Expand Up @@ -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,
Expand Down
8 changes: 6 additions & 2 deletions csrc/tvm_ffi_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<TensorView> cumulative_local_expert_recv_stats, TensorView sym_buffer, Array<int64_t> sym_buffer_ptrs,
int64_t rank_idx, int64_t num_max_tokens_per_rank, int64_t num_experts, int64_t num_topk,
Tuple<int64_t, int64_t, int64_t> recipe, std::string activation, Optional<double> activation_clamp_opt, bool fast_math) {
Tuple<int64_t, int64_t, int64_t> recipe, std::string activation,
double activation_alpha, double activation_up_bias,
Optional<double> activation_clamp_opt, bool fast_math) {
auto c_val = cumulative_local_expert_recv_stats.has_value()? std::optional<torch::Tensor>(convert_to_torch_tensor(cumulative_local_expert_recv_stats.value())) : std::nullopt;
auto act_clamp_opt_val = activation_clamp_opt.has_value()? std::optional<float>(static_cast<float>(activation_clamp_opt.value())) : std::nullopt;
std::vector<int64_t> sym_buffer_ptrs_val;
Expand All @@ -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<int>(rank_idx),
static_cast<int>(num_max_tokens_per_rank), static_cast<int>(num_experts),
static_cast<int>(num_topk), recipe_val, activation, act_clamp_opt_val, fast_math
static_cast<int>(num_topk), recipe_val, activation,
static_cast<float>(activation_alpha), static_cast<float>(activation_up_bias),
act_clamp_opt_val, fast_math
);
}

Expand Down
28 changes: 16 additions & 12 deletions deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -41,18 +41,20 @@ __forceinline__ __device__ float sm90_fp8_mega_moe_clamp_up(float x) {
return x;
}

template <bool kFastMath>
template <bool kFastMath, float kActivationAlpha>
__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 <bool kFastMath, float kActivationClamp>
template <bool kFastMath, float kActivationClamp, float kActivationAlpha,
float kActivationUpBias>
__forceinline__ __device__ float sm90_fp8_mega_moe_swiglu(float g, float u) {
g = sm90_fp8_mega_moe_clamp_gate<kActivationClamp>(g);
u = sm90_fp8_mega_moe_clamp_up<kActivationClamp>(u);
return sm90_fp8_mega_moe_silu<kFastMath>(g) * u;
return sm90_fp8_mega_moe_silu<kFastMath, kActivationAlpha>(g) *
(u + kActivationUpBias);
}

// Continuous FP32 activation scale. SM90 WGMMA has no hardware block-scale operand (the SF
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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;
};
Expand All @@ -1654,7 +1658,7 @@ sm90_fp8_mega_moe_impl(void* y,
.get_data_buffer(m_idx + token_0)
.get_base_ptr<float>();
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];
Expand All @@ -1665,7 +1669,7 @@ sm90_fp8_mega_moe_impl(void* y,
.get_data_buffer(m_idx + token_1)
.get_base_ptr<float>();
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;
}
};

Expand Down Expand Up @@ -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;
};
Expand All @@ -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;
Expand Down
5 changes: 4 additions & 1 deletion sgl_deep_gemm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
)

Expand Down
37 changes: 32 additions & 5 deletions sgl_deep_gemm/tests/test_mega_moe_hopper.py
Original file line number Diff line number Diff line change
Expand Up @@ -275,15 +275,20 @@ 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
gate, up = gate_up[..., :half], gate_up[..., half:]
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(
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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, (
Expand Down Expand Up @@ -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")

Expand All @@ -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,
)
Expand All @@ -506,6 +521,8 @@ def run_fused():
hidden,
intermediate_hidden,
activation_clamp,
activation_alpha,
activation_up_bias,
)

diff = calc_diff(y_fused, y_ref)
Expand Down Expand Up @@ -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]]]:
Expand Down