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:
co-authored by
Chenggang Zhao
parent
7f2a703ed5
commit
891d57b4db
@@ -55,38 +55,68 @@ struct MegaMoEConfig {
|
||||
}
|
||||
};
|
||||
|
||||
static int get_block_m_for_mega_moe(const int& num_ranks, const int& num_experts,
|
||||
const int& num_max_tokens_per_rank, const int& num_topk) {
|
||||
// TODO: compute based on configs
|
||||
return 192;
|
||||
static std::tuple<int, int, int, int> get_block_config_for_mega_moe(
|
||||
const int& num_ranks, const int& num_experts,
|
||||
const int& num_max_tokens_per_rank, const int& num_topk,
|
||||
const int& num_tokens) {
|
||||
const auto& [cluster_size, block_m, store_block_m, num_epilogue_warpgroups] = [&]() -> std::tuple<int, int, int, int> {
|
||||
float num_expected_tokens_per_expert = static_cast<float>(num_tokens) * num_ranks * num_topk / num_experts;
|
||||
if (num_expected_tokens_per_expert <= 8.5) {
|
||||
// Really small token-per-expert (e.g. RL long-tail rollout), use the smallest block_m
|
||||
return {2, 16, 8, 2};
|
||||
} else if (num_expected_tokens_per_expert <= 16.5) {
|
||||
// Small batch size, small EP, decoding, e.g. 6/384 experts, EP8, bsz 128
|
||||
return {2, 32, 16, 2};
|
||||
} else if (num_expected_tokens_per_expert <= 32.5) {
|
||||
// Medium batch size, small EP, decoding, e.g. 6/384 experts, EP8, bsz 256
|
||||
return {2, 64, 32, 1};
|
||||
} else if (num_expected_tokens_per_expert <= 64.5) {
|
||||
// Large batch size, small EP, decoding, e.g. 6/384 experts, EP8, bsz 512
|
||||
return {2, 96, 16, 2};
|
||||
} else if (num_expected_tokens_per_expert <= 96.5) {
|
||||
// Medium batch size, Medium EP, decoding, e.g. 6/384 experts, EP16, bsz 256, or EP32, bsz128
|
||||
return {2, 128, 32, 2};
|
||||
} else {
|
||||
// Prefill, or large EP decoding
|
||||
return {2, 192, 32, 2};
|
||||
}
|
||||
}();
|
||||
|
||||
// Check whether our `block_m` lies in `kCandidateBlockM`
|
||||
DG_HOST_ASSERT(std::any_of(
|
||||
layout::kCandidateBlockM, layout::kCandidateBlockM + layout::kNumCandidateBlockMs,
|
||||
[=](const auto& candidate) { return candidate == block_m; })
|
||||
);
|
||||
|
||||
// Return configs
|
||||
return {cluster_size, block_m, store_block_m, num_epilogue_warpgroups * 128};
|
||||
}
|
||||
|
||||
static int get_num_experts_per_wave_for_mega_moe(
|
||||
const int& num_experts_per_rank, const int& num_tokens, const int& num_topk,
|
||||
const int& intermediate_hidden, const int& block_m, const int& block_n, const int& num_sms) {
|
||||
|
||||
float expected_tokens_per_expert = static_cast<float>(num_tokens) * num_topk / num_experts_per_rank;
|
||||
if (expected_tokens_per_expert < 1) {
|
||||
// Most experts don't have tokens, calculate all experts at once
|
||||
return num_experts_per_rank;
|
||||
}
|
||||
|
||||
// Reduce per-expert block count by this factor since uneven routing leaves some experts with fewer tokens
|
||||
constexpr int kImbalanceFactor = 2;
|
||||
|
||||
// TODO: support num_experts_per_rank > 32
|
||||
// Find the largest divisor of num_experts_per_rank that fits in 32 as the upper bound
|
||||
int max_num_experts_per_wave = std::min(32, num_experts_per_rank);
|
||||
while (max_num_experts_per_wave > 1 and num_experts_per_rank % max_num_experts_per_wave != 0)
|
||||
-- max_num_experts_per_wave;
|
||||
|
||||
// Count L1 blocks per expert assuming tokens are evenly spread across experts
|
||||
const int expected_tokens_per_expert =
|
||||
num_tokens * num_topk / num_experts_per_rank + 1;
|
||||
const int num_m_blocks = ceil_div(expected_tokens_per_expert, block_m);
|
||||
const int num_n_blocks = intermediate_hidden / block_n;
|
||||
const int num_m_blocks = ceil_div(static_cast<int>(std::ceil(expected_tokens_per_expert)), block_m);
|
||||
const int num_n_blocks = (2 * intermediate_hidden) / block_n;
|
||||
const int num_l1_blocks_per_expert = num_m_blocks * num_n_blocks;
|
||||
|
||||
// Pick the smallest value whose total blocks (after imbalance reduction) can keep all SMs busy
|
||||
int num_experts_per_wave = num_l1_blocks_per_expert > 0
|
||||
? ceil_div(kImbalanceFactor * num_sms, num_l1_blocks_per_expert) : 1;
|
||||
num_experts_per_wave = std::min(num_experts_per_wave, max_num_experts_per_wave);
|
||||
num_experts_per_wave = std::min(num_experts_per_wave, num_experts_per_rank);
|
||||
|
||||
// Round up to the nearest divisor of num_experts_per_rank so every wave processes the same count
|
||||
while (num_experts_per_wave < max_num_experts_per_wave and num_experts_per_rank % num_experts_per_wave != 0)
|
||||
while (num_experts_per_wave < num_experts_per_rank and num_experts_per_rank % num_experts_per_wave != 0)
|
||||
++ num_experts_per_wave;
|
||||
|
||||
return num_experts_per_wave;
|
||||
@@ -148,18 +178,18 @@ static std::pair<int, int> get_pipeline_config_for_mega_moe(
|
||||
static MegaMoEConfig get_mega_moe_config(
|
||||
const int& num_ranks, const int& num_experts, const int& num_experts_per_rank,
|
||||
const int& num_max_tokens_per_rank, const int& num_tokens, const int& num_topk,
|
||||
const int& hidden, const int& intermediate_hidden) {
|
||||
// Block tiling
|
||||
const int block_m = get_block_m_for_mega_moe(num_ranks, num_experts, num_max_tokens_per_rank, num_topk);
|
||||
const int& hidden, const int& intermediate_hidden,
|
||||
const int& num_padded_sf_pool_tokens) {
|
||||
// Block config
|
||||
const auto [cluster_size, block_m, store_block_m, num_epilogue_threads] =
|
||||
get_block_config_for_mega_moe(num_ranks, num_experts, num_max_tokens_per_rank, num_topk, num_tokens);
|
||||
const int block_n = 128;
|
||||
const int block_k = 128;
|
||||
const int load_block_m = block_m / 2;
|
||||
const int load_block_n = block_n;
|
||||
const int store_block_m = 32;
|
||||
const auto [sf_block_m, sf_block_n] = SM100ArchSpec::get_sf_uttcp_aligned_block_sizes(block_m, block_n, MmaKind::MXFP8FP4);
|
||||
const int num_max_pool_tokens = layout::get_num_max_pool_tokens(
|
||||
num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank, block_m);
|
||||
const int num_padded_sf_pool_tokens = layout::get_num_padded_sf_pool_tokens(num_max_pool_tokens, block_m);
|
||||
num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank);
|
||||
// NOTES: FP8 activations and FP4 weights (unpacked to 8-bit in smem) both use 128B swizzle
|
||||
const int swizzle_acts_mode = 128;
|
||||
const int swizzle_weights_mode = 128;
|
||||
@@ -173,7 +203,6 @@ static MegaMoEConfig get_mega_moe_config(
|
||||
// Thread layout
|
||||
const int num_dispatch_threads = 128;
|
||||
const int num_non_epilogue_threads = 128;
|
||||
const int num_epilogue_threads = 256;
|
||||
|
||||
// Pipeline
|
||||
const auto [num_stages, smem_size] = get_pipeline_config_for_mega_moe(
|
||||
|
||||
Reference in New Issue
Block a user