-
Notifications
You must be signed in to change notification settings - Fork 1.3k
[Feat] add support sm90 fp8 mega moe #422
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: nv_dev
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -15,6 +15,7 @@ | |
| #include "../jit_kernels/impls/sm100_bf16_mega_moe.hpp" | ||
| #include "../jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp" | ||
| #include "../jit_kernels/impls/sm100_fp4_fp4_mega_moe.hpp" | ||
| #include "../jit_kernels/impls/sm90_fp8_mega_moe.hpp" | ||
|
|
||
| namespace deep_gemm::mega { | ||
|
|
||
|
|
@@ -514,6 +515,228 @@ static void fp4_fp4_mega_moe( | |
| sym_buffer.zero_(); | ||
| } | ||
|
|
||
|
|
||
| // ============================================================================ | ||
| // SM90 (Hopper) FP8 MegaMoE | ||
| // ---------------------------------------------------------------------------- | ||
| // Ported from upstream PR https://github.com/sgl-project/DeepGEMM/pull/36 | ||
| // (`csrc/apis/sm90_mega.hpp` on the old `dev` branch). | ||
| // | ||
| // Unlike `fp8_fp4_mega_moe` (SM100, FP4 weights, UE8M0 SF, ring-buffer workspace, optional | ||
| // shared experts), this path is FP8-only (both activations and weights), uses float SF at | ||
| // per-128 (L1)/per-64 (L2) K granularity, and uses a simpler pool-based workspace/scheduler | ||
| // (`layout::MegaMoESM90Workspace` / `sched::MegaMoESM90Scheduler`). It does not support shared | ||
| // experts. The symmetric buffer layout is therefore also different and is *not* | ||
| // interchangeable with `get_symm_buffer_size_for_mega_moe`'s buffer. | ||
| // ============================================================================ | ||
|
|
||
| static std::tuple<int64_t, std::function<std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>(const torch::Tensor&)>> | ||
| get_symm_buffer_size_for_sm90_mega_moe( | ||
| const int& num_ranks, const int& num_experts, | ||
| const int& num_max_tokens_per_rank, const int& num_topk, | ||
| const int& hidden, const int& intermediate_hidden, | ||
| 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"); | ||
| // `get_mega_moe_config_sm90` may pick `num_dispatch_threads == 64`, and the SM90 | ||
| // `nvlink_barrier` only has one signaling thread per rank (`thread_idx < kNumRanks`), | ||
| // so more than 64 ranks either fails the kernel's `kNumRanks <= kNumThreads` | ||
| // static_assert at JIT time or, if that were relaxed, would hang on missed signals. | ||
| DG_HOST_ASSERT(num_ranks <= 64); | ||
|
|
||
| const auto workspace = layout::MegaMoESM90Workspace(nullptr, num_ranks, num_experts, num_max_tokens_per_rank, num_topk); | ||
|
|
||
| const auto fp8_token_layout = layout::Data(hidden); | ||
| const auto bf16_token_layout = layout::Data(hidden * 2); | ||
| const auto fp8_intermediate_token_layout = layout::Data(intermediate_hidden); | ||
| const auto fp8_sf_layout = layout::Data(hidden / 32); | ||
| const auto fp8_intermediate_sf_layout = layout::Data(intermediate_hidden / 16); | ||
| const auto input_topk_idx_layout = layout::Data(num_topk * sizeof(int64_t), false); | ||
| const auto input_topk_weights_layout = layout::Data(num_topk * sizeof(float), false); | ||
| const auto l1_topk_weights_layout = layout::Data(sizeof(float), false); | ||
|
|
||
| const auto input_token_buffer = layout::Buffer( | ||
| fp8_token_layout, 1, num_max_tokens_per_rank, | ||
| workspace.get_end_ptr()); | ||
| const auto input_sf_buffer = layout::Buffer( | ||
| fp8_sf_layout, 1, num_max_tokens_per_rank, | ||
| input_token_buffer.get_end_ptr()); | ||
| const auto input_topk_idx_buffer = layout::Buffer( | ||
| input_topk_idx_layout, 1, num_max_tokens_per_rank, | ||
| input_sf_buffer.get_end_ptr()); | ||
| const auto input_topk_weights_buffer = layout::Buffer( | ||
| input_topk_weights_layout, 1, num_max_tokens_per_rank, | ||
| input_topk_idx_buffer.get_end_ptr()); | ||
|
|
||
| const auto num_max_pool_tokens = static_cast<int>(workspace.num_max_pool_tokens); | ||
| // Unlike SM100 (which can select any of `layout::kCandidateBlockM`), SM90's | ||
| // `get_block_config_for_mega_moe_sm90` only ever picks block_m in {64, 128}. Sizing the | ||
| // SF pool against the full shared candidate set (which includes block_m=8) would | ||
| // over-allocate the SF pool by ~8x. | ||
| constexpr int kSm90CandidateBlockM[] = {64, 128}; | ||
| int num_max_padded_sf_pool_tokens = 0; | ||
| for (int block_m: kSm90CandidateBlockM) { | ||
| num_max_padded_sf_pool_tokens = std::max( | ||
| num_max_padded_sf_pool_tokens, | ||
| layout::get_num_padded_sf_pool_tokens(num_max_pool_tokens, block_m) | ||
| ); | ||
| } | ||
|
|
||
| const auto l1_token_buffer = layout::Buffer( | ||
| fp8_token_layout, 1, num_max_pool_tokens, | ||
| input_topk_weights_buffer.get_end_ptr()); | ||
| const auto l1_sf_buffer = layout::Buffer( | ||
| fp8_sf_layout, 1, num_max_padded_sf_pool_tokens, | ||
| l1_token_buffer.get_end_ptr()); | ||
| const auto l1_topk_weights_buffer = layout::Buffer( | ||
| l1_topk_weights_layout, 1, num_max_pool_tokens, | ||
| l1_sf_buffer.get_end_ptr()); | ||
|
|
||
| const auto l2_token_buffer = layout::Buffer( | ||
| fp8_intermediate_token_layout, 1, num_max_pool_tokens, | ||
| l1_topk_weights_buffer.get_end_ptr()); | ||
| const auto l2_sf_buffer = layout::Buffer( | ||
| fp8_intermediate_sf_layout, 1, num_max_padded_sf_pool_tokens, | ||
| l2_token_buffer.get_end_ptr()); | ||
|
|
||
| const auto combine_token_buffer = layout::Buffer( | ||
| bf16_token_layout, num_topk, num_max_tokens_per_rank, | ||
| l2_sf_buffer.get_end_ptr()); | ||
|
|
||
| // Kept in sync with the stricter check in `fp8_mega_moe_sm90` (see comment there): | ||
| // hidden must be a multiple of 256 for the scheduler's BLOCK_N=256 case to compile. | ||
| DG_HOST_ASSERT(hidden % 256 == 0 and intermediate_hidden % 128 == 0); | ||
|
|
||
| auto slice_input_buffers = [=](const torch::Tensor& buffer) { | ||
| 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}, | ||
| torch::TensorOptions().dtype(torch::kFloat8_e4m3fn).device(buffer.device())); | ||
| auto x_sf = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(input_sf_buffer.base)), | ||
| {num_max_tokens_per_rank, hidden / 128}, | ||
| torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device())); | ||
| auto topk_idx = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(input_topk_idx_buffer.base)), | ||
| {num_max_tokens_per_rank, num_topk}, | ||
| torch::TensorOptions().dtype(torch::kInt64).device(buffer.device())); | ||
| auto topk_weights = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(input_topk_weights_buffer.base)), | ||
| {num_max_tokens_per_rank, num_topk}, | ||
| torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device())); | ||
| auto l1_acts = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l1_token_buffer.base)), | ||
| {num_max_pool_tokens, hidden}, | ||
| torch::TensorOptions().dtype(torch::kFloat8_e4m3fn).device(buffer.device())); | ||
| auto l1_acts_sf = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l1_sf_buffer.base)), | ||
| {num_max_padded_sf_pool_tokens, hidden / 128}, | ||
| {1, num_max_padded_sf_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}, | ||
| torch::TensorOptions().dtype(torch::kFloat8_e4m3fn).device(buffer.device())); | ||
| auto l2_acts_sf = torch::from_blob( | ||
| math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l2_sf_buffer.base)), | ||
| {num_max_padded_sf_pool_tokens, intermediate_hidden / 64}, | ||
| {1, num_max_padded_sf_pool_tokens}, | ||
| torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device())); | ||
| return std::make_tuple(x, x_sf, topk_idx, topk_weights, l1_acts, l1_acts_sf, l2_acts, l2_acts_sf); | ||
| }; | ||
| return {reinterpret_cast<int64_t>(combine_token_buffer.get_end_ptr()), slice_input_buffers}; | ||
| } | ||
|
|
||
| static void fp8_mega_moe_sm90( | ||
| const torch::Tensor& y, | ||
| const std::tuple<torch::Tensor, torch::Tensor>& l1_weights_tuple, | ||
| const std::tuple<torch::Tensor, torch::Tensor>& l2_weights_tuple, | ||
| const std::optional<torch::Tensor>& cumulative_local_expert_recv_stats, | ||
| const torch::Tensor& sym_buffer, | ||
| const std::vector<int64_t>& sym_buffer_ptrs, const int& rank_idx, | ||
| const int& num_max_tokens_per_rank, | ||
| const int& num_experts, const int& num_topk, | ||
| const std::tuple<int, int, int>& recipe, | ||
| const std::string& activation, | ||
| const std::optional<float>& activation_clamp_opt, | ||
| const bool& fast_math | ||
| ) { | ||
| const auto [l1_weights, l1_weights_sf] = l1_weights_tuple; | ||
| const auto [l2_weights, l2_weights_sf] = l2_weights_tuple; | ||
|
|
||
| const auto arch_major = device_runtime->get_arch_major(); | ||
| DG_HOST_ASSERT(arch_major == 9); | ||
|
|
||
| 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"); | ||
|
|
||
| const auto activation_clamp = | ||
| activation_clamp_opt.value_or(std::numeric_limits<float>::infinity()); | ||
| DG_HOST_ASSERT(activation_clamp >= 0); | ||
|
|
||
| DG_HOST_ASSERT(get_major_type_ab(l1_weights) == cute::UMMA::Major::K); | ||
| DG_HOST_ASSERT(get_major_type_ab(l2_weights) == cute::UMMA::Major::K); | ||
| DG_HOST_ASSERT(l1_weights.scalar_type() == torch::kFloat8_e4m3fn); | ||
| DG_HOST_ASSERT(l2_weights.scalar_type() == torch::kFloat8_e4m3fn); | ||
| const auto [num_experts_per_rank, intermediate_hidden_2, hidden] = get_shape<3>(l1_weights); | ||
| const auto [num_experts_per_rank_, hidden_, intermediate_hidden] = get_shape<3>(l2_weights); | ||
| DG_HOST_ASSERT(num_tokens <= num_max_tokens_per_rank); | ||
| DG_HOST_ASSERT(num_experts_per_rank == num_experts_per_rank_); | ||
| DG_HOST_ASSERT(hidden == hidden_); | ||
| DG_HOST_ASSERT(intermediate_hidden_2 == 2 * intermediate_hidden); | ||
| DG_HOST_ASSERT(l1_weights.is_contiguous() and l2_weights.is_contiguous()); | ||
| // `get_mega_moe_config_sm90` may pick BLOCK_N=256 for either the L1 (2 * intermediate_hidden) | ||
| // or L2 (hidden) GEMM depending on the runtime token distribution, and the scheduler | ||
| // requires L1_SHAPE_N/L2_SHAPE_N to be an exact multiple of BLOCK_N or the JIT compile | ||
| // fails. Require hidden % 256 == 0 so this holds regardless of which BLOCK_N is chosen; | ||
| // intermediate_hidden % 128 == 0 already implies (2 * intermediate_hidden) % 256 == 0. | ||
| DG_HOST_ASSERT(hidden % 256 == 0 and intermediate_hidden % 128 == 0); | ||
| DG_HOST_ASSERT(intermediate_hidden / 64 <= 64); | ||
|
Comment on lines
+690
to
+697
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🔴 critical: 拒绝内核实际不支持的 hidden 对齐: 这里仅要求 🤖 v6 |
||
|
|
||
| constexpr int kGranMN = 128, kGranK = 128; | ||
| check_sf_layout(l1_weights_sf, intermediate_hidden * 2, hidden, kGranMN, kGranK, | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🟡 warning: 🤖 v4p |
||
| num_experts_per_rank, false, true, torch::kFloat); | ||
| check_sf_layout(l2_weights_sf, hidden, intermediate_hidden, kGranMN, kGranK, | ||
| num_experts_per_rank, false, true, torch::kFloat); | ||
|
Comment on lines
+699
to
+703
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🔴 critical: 不要接受内核无法解释的转置权重缩放因子: 这里启用的 🤖 v6 |
||
|
|
||
| if (cumulative_local_expert_recv_stats.has_value()) { | ||
| DG_HOST_ASSERT(cumulative_local_expert_recv_stats->scalar_type() == torch::kInt); | ||
| DG_HOST_ASSERT(cumulative_local_expert_recv_stats->numel() == num_experts_per_rank); | ||
| DG_HOST_ASSERT(cumulative_local_expert_recv_stats->is_contiguous()); | ||
| } | ||
|
|
||
| const auto num_ranks = static_cast<int>(sym_buffer_ptrs.size()); | ||
| const auto num_experts_ = num_experts_per_rank * num_ranks; | ||
| const auto [num_required_bytes, slice] = get_symm_buffer_size_for_sm90_mega_moe( | ||
| num_ranks, num_experts, | ||
| num_max_tokens_per_rank, num_topk, | ||
| hidden, intermediate_hidden, | ||
| true, activation); | ||
| DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast<size_t>(num_required_bytes)); | ||
| DG_HOST_ASSERT(num_experts == num_experts_); | ||
|
|
||
| const auto [x, x_sf, topk_idx, topk_weights, l1_acts, l1_acts_sf, l2_acts, l2_acts_sf] = slice(sym_buffer); | ||
|
|
||
| sm90_fp8_mega_moe(y, | ||
| l1_acts, l1_acts_sf, | ||
| l2_acts, l2_acts_sf, | ||
| l1_weights, l2_weights, | ||
| l1_weights_sf, l2_weights_sf, | ||
| cumulative_local_expert_recv_stats, | ||
| sym_buffer_ptrs, | ||
| rank_idx, num_max_tokens_per_rank, | ||
| num_experts_per_rank, | ||
| num_tokens, num_topk, | ||
| hidden, intermediate_hidden, | ||
| activation_clamp, fast_math); | ||
|
|
||
| if (get_env<int>("DG_COMM_KERNEL_DEBUG")) | ||
| sym_buffer.zero_(); | ||
| } | ||
|
|
||
| static void bf16_mega_moe( | ||
| const torch::Tensor& y, | ||
| const torch::Tensor& l1_weights, | ||
|
|
@@ -630,6 +853,8 @@ static void register_apis(pybind11::module_& m) { | |
| m.def("fp8_fp4_mega_moe", &fp8_fp4_mega_moe); | ||
| m.def("fp4_fp4_mega_moe", &fp4_fp4_mega_moe); | ||
| m.def("bf16_mega_moe", &bf16_mega_moe); | ||
| m.def("get_symm_buffer_size_for_sm90_mega_moe", &get_symm_buffer_size_for_sm90_mega_moe); | ||
| m.def("fp8_mega_moe_sm90", &fp8_mega_moe_sm90); | ||
| #endif | ||
| } | ||
|
|
||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🟡 warning: 仅按 SM90 实际候选块大小分配 SF 池: 该循环遍历共享候选集合中的
block_m=8,因此把 SF 池扩展到16 * num_max_pool_tokens;然而本次新增的 SM90 heuristic 只会选择 64 或 128,实际最大需求是2 * num_max_pool_tokens。两个 SF 池因此被放大约 8 倍,在多卡、大 token/hidden 配置下会额外占用数 GB 对称显存并可能直接 OOM。🤖 v6