feat: implement sm90 megamoe phase3 l1 wgmma
This commit is contained in:
@@ -15,7 +15,8 @@ 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;
|
||||
@@ -85,10 +86,16 @@ get_symm_buffer_size_for_sm90_mega_moe(
|
||||
fp8_intermediate_sf_layout, 1, num_max_padded_sf_pool_tokens,
|
||||
l2_token_buffer.get_end_ptr());
|
||||
|
||||
// Phase 3 verification output: one scaled FP32 L1 accumulator tile.
|
||||
const auto l1_accum_debug_layout = layout::Data(128 * static_cast<int>(sizeof(float)), false);
|
||||
const auto l1_accum_debug_buffer = layout::Buffer(
|
||||
l1_accum_debug_layout, 1, 128,
|
||||
l2_sf_buffer.get_end_ptr());
|
||||
|
||||
// Combine input buffer
|
||||
const auto combine_token_buffer = layout::Buffer(
|
||||
bf16_token_layout, num_topk, num_max_tokens_per_rank,
|
||||
l2_sf_buffer.get_end_ptr());
|
||||
l1_accum_debug_buffer.get_end_ptr());
|
||||
|
||||
// Check SF buffer requirements
|
||||
DG_HOST_ASSERT(hidden % 128 == 0 and intermediate_hidden % 128 == 0);
|
||||
@@ -149,10 +156,15 @@ get_symm_buffer_size_for_sm90_mega_moe(
|
||||
reinterpret_cast<int32_t*>(runtime_workspace.get_token_src_metadata_ptr()),
|
||||
{num_max_pool_tokens, 3},
|
||||
torch::TensorOptions().dtype(torch::kInt).device(buffer.device()));
|
||||
auto l1_accum_debug = torch::from_blob(
|
||||
math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l1_accum_debug_buffer.base)),
|
||||
{128, 128},
|
||||
torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device()));
|
||||
return std::make_tuple(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);
|
||||
expert_recv_count_sum, l1_arrival_count, token_src_metadata,
|
||||
l1_accum_debug);
|
||||
};
|
||||
return {reinterpret_cast<int64_t>(combine_token_buffer.get_end_ptr()), slice_input_buffers};
|
||||
}
|
||||
@@ -228,7 +240,8 @@ static void fp8_mega_moe(
|
||||
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);
|
||||
expert_recv_count_sum, l1_arrival_count, token_src_metadata,
|
||||
l1_accum_debug] = slice(sym_buffer);
|
||||
|
||||
// Dispatch into SM90 path
|
||||
DG_HOST_ASSERT(arch_major == 9);
|
||||
@@ -237,6 +250,7 @@ static void fp8_mega_moe(
|
||||
l2_acts, l2_acts_sf,
|
||||
l1_weights, l2_weights,
|
||||
l1_weights_sf, l2_weights_sf,
|
||||
l1_accum_debug,
|
||||
cumulative_local_expert_recv_stats,
|
||||
sym_buffer_ptrs,
|
||||
rank_idx, num_max_tokens_per_rank,
|
||||
|
||||
Reference in New Issue
Block a user