feat: implement sm90 megamoe phase2 dispatch-only
This commit is contained in:
+30
-4
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user