feat: implement sm90 megamoe phase2 dispatch-only

This commit is contained in:
Xinyi Liu
2026-06-18 00:00:36 +08:00
parent 74eb59dfaa
commit 540e5aeadc
4 changed files with 622 additions and 13 deletions
+30 -4
View File
@@ -13,7 +13,9 @@ namespace deep_gemm::mega {
using SM90MegaMoEBufferViews = std::tuple<
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>;
torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor, torch::Tensor>;
static int get_token_alignment_for_sm90_mega_moe() {
return layout::kLCMCandidateBlockM;
@@ -92,8 +94,11 @@ get_symm_buffer_size_for_sm90_mega_moe(
DG_HOST_ASSERT(hidden % 128 == 0 and intermediate_hidden % 128 == 0);
DG_HOST_ASSERT(num_max_padded_sf_pool_tokens % 4 == 0);
// Slice function: creates input and L1/L2 pool tensor views.
// Slice function: creates input, L1/L2 pool, and Phase 2 dispatch-result tensor views.
auto slice_input_buffers = [=](const torch::Tensor& buffer) {
static_assert(sizeof(layout::TokenSrcMetadata) == 3 * sizeof(uint32_t));
const auto runtime_workspace = layout::Workspace(
buffer.data_ptr(), num_ranks, num_experts, num_max_tokens_per_rank, num_topk);
auto x = torch::from_blob(
math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(input_token_buffer.base)),
{num_max_tokens_per_rank, hidden},
@@ -119,6 +124,10 @@ get_symm_buffer_size_for_sm90_mega_moe(
{num_max_padded_sf_pool_tokens, hidden / 128},
{1, num_max_padded_sf_pool_tokens},
torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device()));
auto l1_topk_weights = torch::from_blob(
math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l1_topk_weights_buffer.base)),
{num_max_pool_tokens},
torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device()));
auto l2_acts = torch::from_blob(
math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l2_token_buffer.base)),
{num_max_pool_tokens, intermediate_hidden},
@@ -128,8 +137,22 @@ get_symm_buffer_size_for_sm90_mega_moe(
{num_max_padded_sf_pool_tokens, intermediate_hidden / 128},
{1, num_max_padded_sf_pool_tokens},
torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device()));
auto expert_recv_count_sum = torch::from_blob(
runtime_workspace.get_expert_recv_count_sum_ptr(),
{num_experts / num_ranks},
torch::TensorOptions().dtype(torch::kInt64).device(buffer.device()));
auto l1_arrival_count = torch::from_blob(
runtime_workspace.get_l1_arrival_count_ptr(),
{static_cast<int>(runtime_workspace.num_max_pool_blocks)},
torch::TensorOptions().dtype(torch::kInt).device(buffer.device()));
auto token_src_metadata = torch::from_blob(
reinterpret_cast<int32_t*>(runtime_workspace.get_token_src_metadata_ptr()),
{num_max_pool_tokens, 3},
torch::TensorOptions().dtype(torch::kInt).device(buffer.device()));
return std::make_tuple(x, x_sf, topk_idx, topk_weights,
l1_acts, l1_acts_sf, l2_acts, l2_acts_sf);
l1_acts, l1_acts_sf, l1_topk_weights,
l2_acts, l2_acts_sf,
expert_recv_count_sum, l1_arrival_count, token_src_metadata);
};
return {reinterpret_cast<int64_t>(combine_token_buffer.get_end_ptr()), slice_input_buffers};
}
@@ -202,7 +225,10 @@ static void fp8_mega_moe(
DG_HOST_ASSERT(num_experts == num_experts_);
// Already registered tensors
const auto [x, x_sf, topk_idx, topk_weights, l1_acts, l1_acts_sf, l2_acts, l2_acts_sf] = slice(sym_buffer);
const auto [x, x_sf, topk_idx, topk_weights,
l1_acts, l1_acts_sf, l1_topk_weights,
l2_acts, l2_acts_sf,
expert_recv_count_sum, l1_arrival_count, token_src_metadata] = slice(sym_buffer);
// Dispatch into SM90 path
DG_HOST_ASSERT(arch_major == 9);