Add various optimizations and Mega MoE benchmarks (#316)
* Merge with private repo * Add Mega MoE Benchmark * Minor fix * Update --------- Co-authored-by: Chenggang Zhao <chenggangz@deepseek.com>
This commit is contained in:
@@ -190,12 +190,13 @@ static torch::Tensor fp8_fp4_mqa_logits(const std::tuple<torch::Tensor, std::opt
|
||||
return logits;
|
||||
}
|
||||
|
||||
static torch::Tensor get_paged_mqa_logits_metadata(const torch::Tensor& context_lens, int block_kv, int num_sms) {
|
||||
static torch::Tensor get_paged_mqa_logits_metadata(const torch::Tensor& context_lens, int block_kv, int num_sms, const std::optional<torch::Tensor>& indices) {
|
||||
// NOTES: Only 2D context lens is supported for now
|
||||
DG_HOST_ASSERT(context_lens.dim() == 2);
|
||||
const bool is_context_lens_2d = true;
|
||||
const int batch_size = context_lens.size(0);
|
||||
const int next_n = context_lens.size(1);
|
||||
const bool is_varlen = indices.has_value();
|
||||
DG_HOST_ASSERT(context_lens.scalar_type() == torch::kInt);
|
||||
DG_HOST_ASSERT(context_lens.is_contiguous());
|
||||
|
||||
@@ -204,9 +205,16 @@ static torch::Tensor get_paged_mqa_logits_metadata(const torch::Tensor& context_
|
||||
|
||||
// Dispatch implementation
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9 or arch_major == 10) {
|
||||
if (is_varlen) {
|
||||
const auto& indices_tensor = indices.value();
|
||||
DG_HOST_ASSERT(arch_major == 10 and next_n == 1 and (block_kv == 64 or block_kv == 32));
|
||||
DG_HOST_ASSERT(indices_tensor.dim() == 1 and indices_tensor.size(0) == batch_size);
|
||||
DG_HOST_ASSERT(indices_tensor.is_contiguous());
|
||||
DG_HOST_ASSERT(indices_tensor.scalar_type() == torch::kInt);
|
||||
smxx_paged_mqa_logits_metadata(context_lens, schedule_metadata, batch_size, next_n, block_kv, num_sms, is_context_lens_2d, true, indices_tensor.data_ptr<int>());
|
||||
} else if (arch_major == 9 or arch_major == 10) {
|
||||
DG_HOST_ASSERT(block_kv == 64 or (arch_major == 10 and block_kv == 32));
|
||||
smxx_paged_mqa_logits_metadata(context_lens, schedule_metadata, batch_size, next_n, block_kv, num_sms, is_context_lens_2d);
|
||||
smxx_paged_mqa_logits_metadata(context_lens, schedule_metadata, batch_size, next_n, block_kv, num_sms, is_context_lens_2d, false, nullptr);
|
||||
} else {
|
||||
DG_HOST_UNREACHABLE("Unsupported architecture");
|
||||
}
|
||||
@@ -222,7 +230,8 @@ static torch::Tensor fp8_fp4_paged_mqa_logits(const std::tuple<torch::Tensor, st
|
||||
const torch::Tensor& schedule_meta,
|
||||
const int& max_context_len,
|
||||
const bool& clean_logits,
|
||||
const at::ScalarType& logits_dtype) {
|
||||
const at::ScalarType& logits_dtype,
|
||||
const std::optional<torch::Tensor>& indices) {
|
||||
const auto [q_fp, q_sf] = q;
|
||||
const bool is_fp4 = q_sf.has_value();
|
||||
|
||||
@@ -321,6 +330,17 @@ static torch::Tensor fp8_fp4_paged_mqa_logits(const std::tuple<torch::Tensor, st
|
||||
DG_HOST_ASSERT(block_table.stride(1) == 1);
|
||||
DG_HOST_ASSERT(block_table.scalar_type() == torch::kInt);
|
||||
|
||||
// Check indices
|
||||
const bool is_varlen = indices.has_value();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
const auto indices_tensor = indices.value_or(torch::Tensor());
|
||||
if (is_varlen) {
|
||||
DG_HOST_ASSERT(arch_major == 10 and next_n == 1);
|
||||
DG_HOST_ASSERT(indices_tensor.dim() == 1 and indices_tensor.size(0) == batch_size);
|
||||
DG_HOST_ASSERT(indices_tensor.is_contiguous());
|
||||
DG_HOST_ASSERT(indices_tensor.scalar_type() == torch::kInt);
|
||||
}
|
||||
|
||||
// Check schedule metadata
|
||||
auto [_schedule_meta_size, _meta_info_size] = get_shape<2>(schedule_meta);
|
||||
DG_HOST_ASSERT(_schedule_meta_size == num_sms + 1 and _meta_info_size == 2);
|
||||
@@ -344,15 +364,14 @@ static torch::Tensor fp8_fp4_paged_mqa_logits(const std::tuple<torch::Tensor, st
|
||||
DG_HOST_ASSERT(logits_dtype == torch::kFloat32 or logits_dtype == torch::kBFloat16);
|
||||
|
||||
// Dispatch implementation
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (is_fp4 and arch_major == 10) {
|
||||
sm100_fp4_paged_mqa_logits(q_fp, q_sf.value(), kv_cache, kv_cache_sf, weights, context_lens, logits, block_table, schedule_meta,
|
||||
sm100_fp4_paged_mqa_logits(q_fp, q_sf.value(), kv_cache, kv_cache_sf, weights, context_lens, logits, block_table, indices_tensor, schedule_meta,
|
||||
logits_dtype, batch_size, next_n, num_heads, head_dim, num_kv_blocks, block_kv, is_context_lens_2d,
|
||||
aligned_max_context_len, block_table_stride, num_sms, split_kv);
|
||||
is_varlen, aligned_max_context_len, block_table_stride, num_sms, split_kv);
|
||||
} else if (not is_fp4 and (arch_major == 9 or arch_major == 10)) {
|
||||
smxx_fp8_paged_mqa_logits(q_fp, kv_cache, kv_cache_sf, weights, context_lens, logits, block_table, schedule_meta,
|
||||
smxx_fp8_paged_mqa_logits(q_fp, kv_cache, kv_cache_sf, weights, context_lens, logits, block_table, indices_tensor, schedule_meta,
|
||||
logits_dtype, batch_size, next_n, num_heads, head_dim, num_kv_blocks, block_kv, is_context_lens_2d,
|
||||
aligned_max_context_len, block_table_stride, num_sms, split_kv);
|
||||
is_varlen, aligned_max_context_len, block_table_stride, num_sms, split_kv);
|
||||
} else {
|
||||
DG_HOST_UNREACHABLE("Unsupported architecture");
|
||||
}
|
||||
@@ -386,10 +405,11 @@ static torch::Tensor fp8_paged_mqa_logits(const torch::Tensor& q,
|
||||
const torch::Tensor& block_table,
|
||||
const torch::Tensor& schedule_meta,
|
||||
const int& max_context_len,
|
||||
const bool& clean_logits) {
|
||||
const bool& clean_logits,
|
||||
const std::optional<torch::Tensor>& indices) {
|
||||
return fp8_fp4_paged_mqa_logits(std::make_tuple(q, std::nullopt), fused_kv_cache, weights,
|
||||
context_lens, block_table, schedule_meta,
|
||||
max_context_len, clean_logits, torch::kFloat);
|
||||
max_context_len, clean_logits, torch::kFloat, indices);
|
||||
}
|
||||
#endif
|
||||
|
||||
@@ -407,13 +427,15 @@ static void register_apis(pybind11::module_& m) {
|
||||
py::arg("max_seqlen_k") = 0,
|
||||
py::arg("logits_dtype") = torch::kFloat32);
|
||||
m.def("get_paged_mqa_logits_metadata", &get_paged_mqa_logits_metadata,
|
||||
py::arg("context_lens"), py::arg("block_kv"), py::arg("num_sms"));
|
||||
py::arg("context_lens"), py::arg("block_kv"), py::arg("num_sms"),
|
||||
py::arg("indices") = std::nullopt);
|
||||
m.def("fp8_fp4_paged_mqa_logits", &fp8_fp4_paged_mqa_logits,
|
||||
py::arg("q"), py::arg("kv_cache"), py::arg("weights"),
|
||||
py::arg("context_lens"), py::arg("block_table"), py::arg("schedule_meta"),
|
||||
py::arg("max_context_len"),
|
||||
py::arg("clean_logits") = false,
|
||||
py::arg("logits_dtype") = torch::kFloat32);
|
||||
py::arg("logits_dtype") = torch::kFloat32,
|
||||
py::arg("indices") = std::nullopt);
|
||||
// Legacy API
|
||||
m.def("fp8_mqa_logits", &fp8_mqa_logits,
|
||||
py::arg("q"), py::arg("kv"), py::arg("weights"),
|
||||
@@ -423,7 +445,8 @@ static void register_apis(pybind11::module_& m) {
|
||||
m.def("fp8_paged_mqa_logits", &fp8_paged_mqa_logits,
|
||||
py::arg("q"), py::arg("kv_cache"), py::arg("weights"),
|
||||
py::arg("context_lens"), py::arg("block_table"), py::arg("schedule_meta"),
|
||||
py::arg("max_context_len"), py::arg("clean_logits") = false);
|
||||
py::arg("max_context_len"), py::arg("clean_logits") = false,
|
||||
py::arg("indices") = std::nullopt);
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,10 @@
|
||||
|
||||
namespace deep_gemm::mega {
|
||||
|
||||
static int get_token_alignment_for_mega_moe() {
|
||||
return layout::kLCMCandidateBlockM;
|
||||
}
|
||||
|
||||
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_mega_moe(
|
||||
const int& num_ranks, const int& num_experts,
|
||||
@@ -20,8 +24,7 @@ get_symm_buffer_size_for_mega_moe(
|
||||
DG_HOST_ASSERT(num_experts % num_ranks == 0);
|
||||
|
||||
// Workspace bytes
|
||||
const auto block_m = get_block_m_for_mega_moe(num_ranks, num_experts, num_max_tokens_per_rank, num_topk);
|
||||
const auto workspace = layout::Workspace(nullptr, num_ranks, num_experts, num_max_tokens_per_rank, num_topk, block_m);
|
||||
const auto workspace = layout::Workspace(nullptr, num_ranks, num_experts, num_max_tokens_per_rank, num_topk);
|
||||
|
||||
// Layouts
|
||||
const auto fp8_token_layout = layout::Data(hidden);
|
||||
@@ -49,14 +52,20 @@ get_symm_buffer_size_for_mega_moe(
|
||||
|
||||
// Buffer configs
|
||||
const auto num_max_pool_tokens = static_cast<int>(workspace.num_max_pool_tokens);
|
||||
const auto num_padded_sf_pool_tokens = layout::get_num_padded_sf_pool_tokens(num_max_pool_tokens, block_m);
|
||||
int num_max_padded_sf_pool_tokens = 0;
|
||||
for (int block_m: layout::kCandidateBlockM) {
|
||||
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)
|
||||
);
|
||||
}
|
||||
|
||||
// L1 input buffer
|
||||
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_padded_sf_pool_tokens,
|
||||
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,
|
||||
@@ -67,7 +76,7 @@ get_symm_buffer_size_for_mega_moe(
|
||||
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_padded_sf_pool_tokens,
|
||||
fp8_intermediate_sf_layout, 1, num_max_padded_sf_pool_tokens,
|
||||
l2_token_buffer.get_end_ptr());
|
||||
|
||||
// Combine input buffer: BF16 tokens for cross-rank combine
|
||||
@@ -77,7 +86,7 @@ get_symm_buffer_size_for_mega_moe(
|
||||
|
||||
// Check SF buffer requirements
|
||||
DG_HOST_ASSERT(hidden % 128 == 0 and intermediate_hidden % 128 == 0);
|
||||
DG_HOST_ASSERT(num_padded_sf_pool_tokens % 4 == 0);
|
||||
DG_HOST_ASSERT(num_max_padded_sf_pool_tokens % 4 == 0);
|
||||
|
||||
// Slice function: creates `(x, x_sf, topk_weights, topk_idx, l1_acts, l1_acts_sf, l2_acts, l2_acts_sf)` tensor views from the raw buffer
|
||||
// NOTES: `x_sf` is K-major, while `l1_acts_sf` and `l2_acts_sf` are M-major
|
||||
@@ -104,8 +113,8 @@ get_symm_buffer_size_for_mega_moe(
|
||||
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_padded_sf_pool_tokens, hidden / 128},
|
||||
{1, num_padded_sf_pool_tokens},
|
||||
{num_max_padded_sf_pool_tokens, hidden / 128},
|
||||
{1, num_max_padded_sf_pool_tokens},
|
||||
torch::TensorOptions().dtype(torch::kInt).device(buffer.device()));
|
||||
auto l2_acts = torch::from_blob(
|
||||
math::advance_ptr(buffer.data_ptr(), reinterpret_cast<int64_t>(l2_token_buffer.base)),
|
||||
@@ -113,8 +122,8 @@ get_symm_buffer_size_for_mega_moe(
|
||||
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_padded_sf_pool_tokens, intermediate_hidden / 128},
|
||||
{1, num_padded_sf_pool_tokens},
|
||||
{num_max_padded_sf_pool_tokens, intermediate_hidden / 128},
|
||||
{1, num_max_padded_sf_pool_tokens},
|
||||
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);
|
||||
};
|
||||
@@ -123,8 +132,9 @@ get_symm_buffer_size_for_mega_moe(
|
||||
|
||||
static void fp8_fp4_mega_moe(
|
||||
const torch::Tensor& y,
|
||||
const std::tuple<torch::Tensor, torch::Tensor>& l1_weights_,
|
||||
const std::tuple<torch::Tensor, torch::Tensor>& l2_weights_,
|
||||
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,
|
||||
@@ -132,9 +142,10 @@ static void fp8_fp4_mega_moe(
|
||||
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_;
|
||||
const auto [l2_weights, l2_weights_sf] = l2_weights_;
|
||||
const bool& fast_math
|
||||
) {
|
||||
const auto [l1_weights, l1_weights_sf] = l1_weights_tuple;
|
||||
const auto [l2_weights, l2_weights_sf] = l2_weights_tuple;
|
||||
|
||||
// Config checks
|
||||
const auto num_tokens = static_cast<int>(y.size(0));
|
||||
@@ -161,13 +172,20 @@ static void fp8_fp4_mega_moe(
|
||||
DG_HOST_ASSERT(intermediate_hidden_2 == 2 * intermediate_hidden);
|
||||
DG_HOST_ASSERT(l1_weights.is_contiguous() and l2_weights.is_contiguous());
|
||||
|
||||
// Check weight SF layout for UE8M0 packing, MN-major, and TMA alignment
|
||||
// Check weight SF layout for UE8M0 packing, MN-major, and TMA alignment
|
||||
constexpr int kGranMN = 1, kGranK = 32;
|
||||
check_sf_layout(l1_weights_sf, intermediate_hidden * 2, hidden, kGranMN, kGranK,
|
||||
num_experts_per_rank, true, false, torch::kInt);
|
||||
check_sf_layout(l2_weights_sf, hidden, intermediate_hidden, kGranMN, kGranK,
|
||||
num_experts_per_rank, true, false, torch::kInt);
|
||||
|
||||
// Check stats counter
|
||||
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());
|
||||
}
|
||||
|
||||
// Check buffer bytes
|
||||
const auto num_ranks = static_cast<int>(sym_buffer_ptrs.size());
|
||||
const auto num_experts_ = num_experts_per_rank * num_ranks;
|
||||
@@ -175,7 +193,7 @@ static void fp8_fp4_mega_moe(
|
||||
num_ranks, num_experts,
|
||||
num_max_tokens_per_rank, num_topk,
|
||||
hidden, intermediate_hidden,
|
||||
true, "swiglu");
|
||||
true, activation);
|
||||
DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast<size_t>(num_required_bytes));
|
||||
DG_HOST_ASSERT(num_experts == num_experts_);
|
||||
|
||||
@@ -189,6 +207,7 @@ static void fp8_fp4_mega_moe(
|
||||
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,
|
||||
@@ -207,7 +226,7 @@ static void fp8_fp4_mega_moe(
|
||||
|
||||
static void register_apis(pybind11::module_& m) {
|
||||
#if DG_TENSORMAP_COMPATIBLE
|
||||
m.def("get_block_m_for_mega_moe", &get_block_m_for_mega_moe);
|
||||
m.def("get_token_alignment_for_mega_moe", &get_token_alignment_for_mega_moe);
|
||||
m.def("get_symm_buffer_size_for_mega_moe", &get_symm_buffer_size_for_mega_moe);
|
||||
m.def("fp8_fp4_mega_moe", &fp8_fp4_mega_moe);
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user