diff --git a/README.md b/README.md index 73e6c710ee..412390a00e 100644 --- a/README.md +++ b/README.md @@ -169,6 +169,37 @@ The input activation SF is FP32 with shape `[num_tokens, hidden / 128]`, using o For distributed correctness and performance drivers, refer to `tests/test_mega_moe_sm90.py` and `tests/bench_mega_moe_sm90.py`. +##### SM90 FP8xFP8, fused single kernel + +The fused SM90 implementation runs linear 1 and linear 2 in one kernel: one task stream, with linear 2 trailing linear 1 by a fixed schedule lag, and the linear-1 activation pool sized as a ring around that lag rather than around the token bound. It keeps the split path's tensor contract, weight transform and entry point; `fused=True` selects it when the buffer is allocated, and the buffer type then selects the kernel at launch: + +```python +# NOTES: requires PyTorch >= 2.9 +buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( + group, num_experts, num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + fused=True, +) + +transformed_l1, transformed_l2 = deep_gemm.transform_weights_for_mega_moe_sm90( + (l1_weight_fp8, l1_weight_sf), + (l2_weight_fp8, l2_weight_sf), +) + +buffer.x[:num_tokens].copy_(x_fp8) +buffer.x_sf[:num_tokens].copy_(x_sf) +buffer.topk_idx[:num_tokens].copy_(topk_idx) +buffer.topk_weights[:num_tokens].copy_(topk_weights) + +y = torch.empty((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') +# `num_tokens_bound` must be the same value on every rank, e.g. the max over ranks of `num_tokens` +deep_gemm.fp8_mega_moe(y, transformed_l1, transformed_l2, buffer, num_tokens_bound=num_tokens) +``` + +The activation and weight SF contract, the shape constraints on `hidden` and `intermediate_hidden`, and the 128-token alignment are the ones above; the fused combine stages whole rows where they fit and chunks of a row where they do not, so a large `hidden` can need a larger multiple than 256, which the host reports before compiling. Two buffer arguments are specific to this path, both optional: `num_experts_per_wave` chooses the activation-pool layout (`None` follows the schedule, `-1` sizes a ring for an automatically chosen expert wave, `N > 0` for a wave of `min(N, experts_per_rank)` experts, never larger than a full pool), and `l2_act_sf_gran_k` (64 or 128) is the K granularity of the FP8 scale on the intermediate activation, defaulting to 128 where the shape allows it (pinning 128 elsewhere is rejected at launch on decode-sized calls). `num_tokens_bound` lets a call smaller than the buffer be scheduled for its own token count; it must be identical on every rank and cover every rank's own count, and it is rejected above the buffer capacity. Each rank checks only its own count: a bound below a peer's count sizes this rank's single-wave pool for fewer rows than arrive, which reuses a ring slot inside one wave and hangs. `intermediate_hidden` must be at most 4096 on this path. + +`tests/test_mega_moe_sm90.py --fused` and `tests/test_mega_moe_sm90_fused_ring.py` drive it, and `tests/bench_mega_moe_sm90.py --arms split fused` benchmarks it against the two-kernel path. + #### Utilities The library provides some utility functions besides the above kernels: diff --git a/csrc/apis/sm90_fused_mega.hpp b/csrc/apis/sm90_fused_mega.hpp new file mode 100644 index 0000000000..ee70190d36 --- /dev/null +++ b/csrc/apis/sm90_fused_mega.hpp @@ -0,0 +1,299 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +#if DG_TENSORMAP_COMPATIBLE +#include "../jit/compiler.hpp" +#endif +#include "../jit/device_runtime.hpp" +#include "../jit_kernels/impls/sm90_fp8_fused_mega_moe.hpp" +#include "../utils/layout.hpp" +#include "../utils/system.hpp" + +namespace deep_gemm::mega { + +static int get_token_alignment_for_sm90_fused_mega_moe() { + return layout::kSM90FusedLCMBlockM; +} + +// Byte layout of a fused symm buffer whose data pools hold `num_ring_tokens` tokens (0 = the full pool) +static std::tuple(const torch::Tensor&)>> +get_symm_buffer_layout_for_sm90_fused_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, + const int& num_ring_tokens, const int& l2_act_sf_gran_k) { + DG_HOST_ASSERT(num_ranks > 0); + DG_HOST_ASSERT(num_experts % num_ranks == 0); + if (not use_fp8_dispatch) + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE supports FP8 dispatch only"); + if (activation != "swiglu") + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE supports the swiglu activation only"); + DG_HOST_ASSERT(num_max_tokens_per_rank > 0 and + num_max_tokens_per_rank % layout::kSM90FusedLCMBlockM == 0); + if (hidden <= 0 or hidden % 128 != 0 or intermediate_hidden <= 0 or intermediate_hidden % 128 != 0) + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE requires hidden and intermediate_hidden to be positive multiples of 128"); + // The kernel L2 arrival mask covers at most 64 per-64-K groups. Checked here as well as at launch so a shape this path + // cannot serve is reported before a symmetric buffer is allocated for it. + DG_HOST_ASSERT(intermediate_hidden / 64 <= 64); + DG_HOST_ASSERT(l2_act_sf_gran_k == 64 or l2_act_sf_gran_k == 128); + DG_HOST_ASSERT(num_ring_tokens >= 0); + + // Workspace bytes + const int num_experts_per_rank = num_experts / num_ranks; + const int num_max_pool_tokens = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + const auto workspace_end_ptr = layout::SM90FusedWorkspace( + nullptr, num_ranks, num_experts, num_max_tokens_per_rank, num_topk, num_ring_tokens).get_end_ptr(); + + // Layouts + 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); + // Slot-major L2 act SF pool (`[k_sf_idx][pool token]`): the per-token byte count only sizes it and need not be TMA-aligned + const auto fp8_intermediate_sf_layout = layout::Data(intermediate_hidden * 4 / l2_act_sf_gran_k, false); + 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); + + // Input buffers + const auto input_token_buffer = layout::Buffer( + fp8_token_layout, 1, num_max_tokens_per_rank, + workspace_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()); + + // Data pools hold the ring (the full pool when no ring); SF pools are padded for the worst-case BLOCK_M + const int num_data_pool_tokens = num_ring_tokens == 0 ? num_max_pool_tokens : num_ring_tokens; + const int num_max_padded_sf_pool_tokens = get_num_padded_sf_pool_tokens_sm90_fused(num_data_pool_tokens); + + // L1 input buffer + const auto l1_token_buffer = layout::Buffer( + fp8_token_layout, 1, num_data_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_data_pool_tokens, + l1_sf_buffer.get_end_ptr()); + + // L2 input buffer + const auto l2_token_buffer = layout::Buffer( + fp8_intermediate_token_layout, 1, num_data_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()); + + // Combine input buffer: BF16 tokens for cross-rank combine + const auto combine_token_buffer = layout::Buffer( + bf16_token_layout, num_topk, num_max_tokens_per_rank, + l2_sf_buffer.get_end_ptr()); + + // `x_sf` is K-major; pool scale factors are M-major. + auto slice_input_buffers = [=](const torch::Tensor& buffer) { + auto x = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), reinterpret_cast(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(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(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(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(l1_token_buffer.base)), + {num_data_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(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(l2_token_buffer.base)), + {num_data_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(l2_sf_buffer.base)), + {num_max_padded_sf_pool_tokens, intermediate_hidden / l2_act_sf_gran_k}, + {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(combine_token_buffer.get_end_ptr()), slice_input_buffers}; +} + +// Returns (num_bytes, slicer, num_ring_tokens, l2_lag_encoded). The ring capacity and the encoded L2-lag schedule are +// derived here from `num_experts_per_wave` (0 = full pool, -1 = auto wave size, N > 0 = fixed wave size) and must be +// passed back unchanged at every launch. +static std::tuple(const torch::Tensor&)>, int, int> +get_symm_buffer_size_for_sm90_fused_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, + const int& num_experts_per_wave, const int& l2_act_sf_gran_k) { + DG_HOST_ASSERT(num_ranks > 0); + DG_HOST_ASSERT(num_experts % num_ranks == 0); + const auto [num_ring_tokens, l2_lag_encoded] = get_num_ring_tokens_for_sm90_fused_mega_moe( + num_ranks, num_experts, num_experts / num_ranks, + num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + num_experts_per_wave, l2_act_sf_gran_k); + const auto [num_bytes, slice_input_buffers] = get_symm_buffer_layout_for_sm90_fused_mega_moe( + num_ranks, num_experts, + num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + use_fp8_dispatch, activation, + num_ring_tokens, l2_act_sf_gran_k); + return {num_bytes, slice_input_buffers, num_ring_tokens, l2_lag_encoded}; +} + +// SM90 (Hopper) fused FP8 MegaMoE entry point: the same contract as `fp8_mega_moe` (FP8 e4m3 weights with block +// (128, 128) float scale factors) on a buffer sized by `get_symm_buffer_size_for_sm90_fused_mega_moe`. +static void sm90_fused_fp8_mega_moe( + const torch::Tensor& y, + const std::tuple& l1_weights_tuple, + const std::tuple& l2_weights_tuple, + const std::optional& cumulative_local_expert_recv_stats, + const torch::Tensor& sym_buffer, + const std::vector& 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& recipe, + const std::string& activation, + const std::optional& activation_clamp_opt, + const bool& fast_math, + // `num_ring_tokens`, `l2_lag_encoded` and `l2_act_sf_gran_k` are the values `get_symm_buffer_size_for_sm90_fused_mega_moe` + // sized `sym_buffer` with; `num_tokens_bound` is the caller's bound on this call's per-rank token count on every rank (0 = capacity) + const int& num_ring_tokens, + const int& num_tokens_bound, + const int& l2_lag_encoded, + const int& l2_act_sf_gran_k +) { + const auto [l1_weights, l1_weights_sf] = l1_weights_tuple; + const auto [l2_weights, l2_weights_sf] = l2_weights_tuple; + + // Architecture check + if (device_runtime->get_arch_major() != 9) + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE requires a compute capability 9.x GPU"); + + // Config checks: block (128, 128) float SF for weights, per-token per-128-K float SF for activations + const auto num_tokens = static_cast(y.size(0)); + const auto [rm, rn, rk] = recipe; + if (rm != 128 or rn != 128 or rk != 128) + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE requires recipe=(128, 128, 128)"); + if (activation != "swiglu") + DG_HOST_UNREACHABLE("SM90 fused FP8 MegaMoE supports the swiglu activation only"); + + // Activation checks + const auto activation_clamp = + activation_clamp_opt.value_or(std::numeric_limits::infinity()); + DG_HOST_ASSERT(activation_clamp >= 0); + + // Tensor checks: weights must be FP8 e4m3, K-major + 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()); + DG_HOST_ASSERT(hidden % 128 == 0 and intermediate_hidden % 128 == 0); + // The kernel's L2 arrival mask covers at most 64 per-64-K groups + DG_HOST_ASSERT(intermediate_hidden / 64 <= 64); + + // Weight SFs are raw global-memory loads in natural MN-major order. + constexpr int kGranMN = 128, kGranK = 128; + check_sf_layout(l1_weights_sf, intermediate_hidden * 2, hidden, kGranMN, kGranK, + 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); + if (not l1_weights_sf.is_contiguous() or not l2_weights_sf.is_contiguous()) + DG_HOST_UNREACHABLE( + "SM90 fused FP8 MegaMoE weight scale factors must use contiguous natural layouts"); + + // 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 against the geometry it was sized with + const auto num_ranks = static_cast(sym_buffer_ptrs.size()); + DG_HOST_ASSERT(num_experts == num_experts_per_rank * num_ranks); + // 0 means the buffer capacity. A bound is the per-rank token count every rank's schedule is sized for, so it must + // be the same on every rank: this checks only the local count, and a bound below a peer's count sizes this rank's + // single-wave pool for fewer rows than arrive, which reuses a ring slot inside one wave and hangs. + DG_HOST_ASSERT(num_tokens_bound >= 0); + DG_HOST_ASSERT(num_tokens_bound == 0 or num_tokens_bound >= num_tokens); + DG_HOST_ASSERT(num_tokens_bound <= num_max_tokens_per_rank); + const auto [num_required_bytes, slice] = get_symm_buffer_layout_for_sm90_fused_mega_moe( + num_ranks, num_experts, + num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + true, activation, + num_ring_tokens, l2_act_sf_gran_k); + DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast(num_required_bytes)); + + // 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); + + sm90_fp8_fused_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, + num_ring_tokens, l2_act_sf_gran_k, + activation_clamp, fast_math, + num_tokens_bound, l2_lag_encoded); + + if (get_env("DG_COMM_KERNEL_DEBUG")) + sym_buffer.zero_(); +} + +static void register_sm90_fused_apis(pybind11::module_& m) { +#if DG_TENSORMAP_COMPATIBLE + m.def("get_token_alignment_for_sm90_fused_mega_moe", &get_token_alignment_for_sm90_fused_mega_moe); + m.def("get_symm_buffer_size_for_sm90_fused_mega_moe", &get_symm_buffer_size_for_sm90_fused_mega_moe); + m.def("sm90_fused_fp8_mega_moe", &sm90_fused_fp8_mega_moe); +#endif +} + +} // namespace deep_gemm::mega diff --git a/csrc/jit/kernel_runtime.hpp b/csrc/jit/kernel_runtime.hpp index 40597fb448..7846bdeb5c 100644 --- a/csrc/jit/kernel_runtime.hpp +++ b/csrc/jit/kernel_runtime.hpp @@ -17,12 +17,16 @@ struct LaunchArgs { int smem_size; int cluster_dim; bool enable_pdl; + // the kernel asked for programmatic dependent launch itself (matching griddepcontrol wait/trigger in its code); the attribute is then set even when the global `deep_gemm.set_pdl` flag is off + bool force_pdl; - LaunchArgs(const int& grid_dim_x, const int& num_threads, const int& smem_size = 0, const int& cluster_dim = 1, const bool& enable_pdl = true): - grid_dim({grid_dim_x, 1}), num_threads(num_threads), smem_size(smem_size), cluster_dim(cluster_dim), enable_pdl(enable_pdl) {} + LaunchArgs(const int& grid_dim_x, const int& num_threads, const int& smem_size = 0, const int& cluster_dim = 1, const bool& enable_pdl = true, + const bool& force_pdl = false): + grid_dim({grid_dim_x, 1}), num_threads(num_threads), smem_size(smem_size), cluster_dim(cluster_dim), enable_pdl(enable_pdl), force_pdl(force_pdl) {} - LaunchArgs(const std::pair& grid_dim, const int& num_threads, const int& smem_size = 0, const int& cluster_dim = 1, const bool& enable_pdl = true): - grid_dim(grid_dim), num_threads(num_threads), smem_size(smem_size), cluster_dim(cluster_dim), enable_pdl(enable_pdl) {} + LaunchArgs(const std::pair& grid_dim, const int& num_threads, const int& smem_size = 0, const int& cluster_dim = 1, const bool& enable_pdl = true, + const bool& force_pdl = false): + grid_dim(grid_dim), num_threads(num_threads), smem_size(smem_size), cluster_dim(cluster_dim), enable_pdl(enable_pdl), force_pdl(force_pdl) {} }; class KernelRuntime final { @@ -143,7 +147,8 @@ class LaunchRuntime { // Allow runtime override from Python. // NOTES: the default is enabled. - launch_args.enable_pdl = device_runtime->get_pdl(); + // a kernel compiled with griddepcontrol (force_pdl) keeps the attribute even with the global flag off + launch_args.enable_pdl = launch_args.force_pdl or device_runtime->get_pdl(); const dim3 grid_dim = {static_cast(launch_args.grid_dim.first), static_cast(launch_args.grid_dim.second), diff --git a/csrc/jit_kernels/heuristics/sm90_fused_mega_moe.hpp b/csrc/jit_kernels/heuristics/sm90_fused_mega_moe.hpp new file mode 100644 index 0000000000..fde49b989a --- /dev/null +++ b/csrc/jit_kernels/heuristics/sm90_fused_mega_moe.hpp @@ -0,0 +1,710 @@ +#pragma once + +#include "mega_moe.hpp" + +#include + +namespace deep_gemm { + +// ============================================================================ +// SM90 (Hopper) MegaMoE configuration +// ---------------------------------------------------------------------------- +// SM90 differs from SM100 in: +// - No tensor memory (TMEM): WGMMA accumulators live in registers. +// - No FP4: weights are FP8 e4m3 with per-128 channel float scales. +// - No 2-CTA cluster MMA: TMA multicast cluster=2 may still be used. +// - Activation SF is float, not UE8M0 int: L1 input uses per-128 K and the +// fused L1 epilogue writes L2 activation SF at per-128 or per-64 K granularity (the buffer's l2_act_sf_gran_k). +// The kernel implementation is in `deep_gemm/impls/sm90_fp8_fused_mega_moe.cuh`. +// ============================================================================ + +struct MegaMoESM90FusedConfig { + int block_m, block_n, block_k; + int cluster_size; + int num_max_pool_tokens; + int num_padded_sf_pool_tokens; + int num_ring_tokens; + int swizzle_acts_mode, swizzle_weights_mode; + int num_experts_per_wave; + int num_stages, smem_size; + int num_dispatch_threads, num_non_epilogue_threads, num_epilogue_threads; + bool half_l2_cd; + bool multicast_on_b; + // > 0: L2 tiles interleaved into the L1 phase at this lag (units of kSM90FusedLagUnitM m-blocks, encoded lag + 1000 x group); 0: wave schedule + int l2_lag_units; + // L2 BF16 staging passes (0 = derived from half_l2_cd; 4 = quarter-width buffer) + int l2_cd_passes; + // L2 BF16 staging on the quarter-width buffer: 0 = column passes, 2 = row passes, 4 = row passes through stmatrix + int l2_stage_mode; + // 1: combine the tokens whose experts have all published their done flag while waiting in the pre-combine barrier (needs hidden >= 512 x topk); 0: barrier, then combine + int early_combine; + // pull arrivals published per gpu-scope release fence (1 = one red.release per row; N > 1 ramps 1, 2, 4, ... every 8 rows of a warp) + int pull_publish_batch; + + friend std::ostream& operator << (std::ostream& os, const MegaMoESM90FusedConfig& config) { + os << "MegaMoESM90FusedConfig(" + << "block_m=" << config.block_m << ", block_n=" << config.block_n << ", block_k=" << config.block_k + << ", cluster_size=" << config.cluster_size + << ", num_max_pool_tokens=" << config.num_max_pool_tokens + << ", num_padded_sf_pool_tokens=" << config.num_padded_sf_pool_tokens + << ", num_ring_tokens=" << config.num_ring_tokens + << ", swizzle_acts_mode=" << config.swizzle_acts_mode << ", swizzle_weights_mode=" << config.swizzle_weights_mode + << ", num_experts_per_wave=" << config.num_experts_per_wave + << ", num_stages=" << config.num_stages << ", smem_size=" << config.smem_size + << ", num_dispatch_threads=" << config.num_dispatch_threads + << ", num_non_epilogue_threads=" << config.num_non_epilogue_threads + << ", num_epilogue_threads=" << config.num_epilogue_threads + << ", half_l2_cd=" << config.half_l2_cd << ", l2_cd_passes=" << config.l2_cd_passes + << ", multicast_on_b=" << config.multicast_on_b + << ", l2_lag_units=" << config.l2_lag_units + << ", l2_stage_mode=" << config.l2_stage_mode + << ", early_combine=" << config.early_combine + << ", pull_publish_batch=" << config.pull_publish_batch << ")"; + return os; + } +}; + +// L2-lag schedule constants (schedule selection: `get_buffer_schedule_sm90_fused` below) +static constexpr int kSm90FusedAutoLagMinTokens = 1024; +static constexpr double kSm90FusedAutoLagFraction = 0.45; +static constexpr int kSm90FusedAutoLagGroup = 4; +static constexpr int kSm90FusedAutoLagMinUnits = 8; +static constexpr int kSm90FusedAutoLagMaxUnits = 64; +static constexpr int kSm90FusedLagRingMargin = 4; + +static int get_auto_late_lag_sm90_fused(const int& num_max_tokens_per_rank, const int& num_topk, const int& num_experts_per_rank) { + const int units_per_expert = std::max(1, (num_max_tokens_per_rank * num_topk + num_experts_per_rank * 1024 - 1) / + (num_experts_per_rank * 1024)); + const int units = num_experts_per_rank * units_per_expert; + return std::clamp(static_cast(kSm90FusedAutoLagFraction * units + 0.5), kSm90FusedAutoLagMinUnits, kSm90FusedAutoLagMaxUnits); +} + +// Schedule a symm buffer is sized for. Decided once here from the wave-size knob (Python `num_experts_per_wave`, 0 when unset) and +// the capacity, returned to Python with the ring capacity and the encoded lag and passed back at every launch, so the ring capacity +// and the kernel schedule can never disagree. +// knob 0 (unset): kLag from kSm90FusedAutoLagMinTokens tokens per rank (a lag-sized ring running the L2-lag schedule), kFullPool below it +// knob -1: kWave, ring sized by the occupancy heuristic at the capacity +// knob N > 0: kWave, ring sized for a fixed wave of N experts (at call time the wave size is re-derived from the stored capacity, like SM100) +enum class Sm90FusedBufferSchedule { kFullPool, kWave, kLag }; + +static Sm90FusedBufferSchedule get_buffer_schedule_sm90_fused(const int& num_max_tokens_per_rank, const int& num_experts_per_wave_knob) { + DG_HOST_ASSERT(num_experts_per_wave_knob >= -1); + if (num_experts_per_wave_knob != 0) + return Sm90FusedBufferSchedule::kWave; + return num_max_tokens_per_rank >= kSm90FusedAutoLagMinTokens ? Sm90FusedBufferSchedule::kLag : Sm90FusedBufferSchedule::kFullPool; +} + +// Template encoding consumed by the scheduler: lag + 1000 * group +static int get_l2_lag_encoded_sm90_fused(const int& lag) { + return lag > 0 ? lag + 1000 * kSm90FusedAutoLagGroup : 0; +} + +static int sm90_fused_lag_units_of(const int& encoded) { return encoded % 1000; } +static int sm90_fused_lag_group_of(const int& encoded) { return std::max(encoded / 1000, 1); } + +// Ring capacity (in pool blocks) that the lag schedule needs: L2(p) is issued at most kSM90FusedLagUnitM * (lag + group) pool +// blocks after L1(p), plus the blocks that may be in flight on top. +static int get_lag_ring_blocks_sm90_fused(const int& lag, const int& group) { + return static_cast(layout::kSM90FusedLagUnitM) * (lag + std::max(group, 1)) + kSm90FusedLagRingMargin; +} + +// Wave-layout helpers. The pool sizer below is the worst-case token span of a wave of `num_experts_per_wave` +// local experts; the chooser after it is the occupancy heuristic of the two-kernel path +// (`get_generic_num_experts_per_wave_for_mega_moe_sm90`) with the ring capacity as an upper bound, which is +// what the fused path adds: a wave whose span does not fit the ring would make the schedule wait on a slot +// the same wave still owns. The bound also narrows the tail-ratio sweep, so the two do not share code. +static int get_num_wave_pool_tokens_sm90_fused( + const int& num_ranks, const int& num_topk, const int& num_max_tokens_per_rank, + const int& num_experts_per_wave, const int& block_m) { + DG_HOST_ASSERT(num_max_tokens_per_rank % block_m == 0); + const auto num_tokens_from_all_ranks = num_max_tokens_per_rank * num_ranks; + if (num_experts_per_wave == 1) + return num_tokens_from_all_ranks; + + return std::min( + num_tokens_from_all_ranks * num_experts_per_wave, + math::align( + num_tokens_from_all_ranks * num_topk + num_experts_per_wave * (block_m - 1), + block_m)); +} + +static int get_capped_num_experts_per_wave_sm90_fused( + 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, + const int& num_ring_tokens, const int& num_max_tokens_per_rank, const int& num_ranks) { + int num_max_experts_per_wave = num_experts_per_rank; + while (num_max_experts_per_wave > 0 and + get_num_wave_pool_tokens_sm90_fused( + num_ranks, num_topk, num_max_tokens_per_rank, + num_max_experts_per_wave, block_m) > num_ring_tokens) + --num_max_experts_per_wave; + DG_HOST_ASSERT(num_max_experts_per_wave > 0 and "Buffer size is too small"); + + constexpr int kImbalanceFactor = 2; + const float num_expected_tokens_per_expert = + static_cast(num_tokens * num_topk) / num_experts_per_rank; + const int num_expected_m_blocks = std::max( + ceil_div(static_cast(std::ceil(num_expected_tokens_per_expert)), block_m), 1); + const int num_l1_n_blocks = (2 * intermediate_hidden) / block_n; + const int num_expected_l1_blocks_per_expert = num_expected_m_blocks * num_l1_n_blocks; + int num_min_expected_experts_to_fill_sms = + ceil_div(kImbalanceFactor * num_sms, num_expected_l1_blocks_per_expert); + + if (num_expected_tokens_per_expert < 1) + num_min_expected_experts_to_fill_sms = num_experts_per_rank; + if (num_min_expected_experts_to_fill_sms >= num_max_experts_per_wave) + return num_max_experts_per_wave; + if (num_expected_l1_blocks_per_expert >= num_sms) + return num_min_expected_experts_to_fill_sms; + + const int num_sweep_max_experts_per_wave = std::min( + num_max_experts_per_wave, num_min_expected_experts_to_fill_sms * 2); + int best_num_experts_per_wave = num_min_expected_experts_to_fill_sms; + float best_tail_ratio = -1.0f; + for (int num_experts_per_wave = num_min_expected_experts_to_fill_sms; + num_experts_per_wave <= num_sweep_max_experts_per_wave; + ++num_experts_per_wave) { + const int remainder = num_experts_per_rank % num_experts_per_wave; + const float tail_ratio = remainder == 0 ? + 1.0f : static_cast(remainder) / num_experts_per_wave; + if (tail_ratio > best_tail_ratio) { + best_tail_ratio = tail_ratio; + best_num_experts_per_wave = num_experts_per_wave; + } + } + return best_num_experts_per_wave; +} + +// Decode single wave: the BLOCK_M-64 split-N topology under this many tokens per rank runs one wave of all local experts when its worst-case pool fits the ring. +static constexpr int kSm90FusedSingleWaveMaxTokens = 256; + +// Per-call wave schedule on lag-ring buffers: when the call's worst-case pool (num_ranks x per-rank token bound x topk rows plus one +// padding block per local expert) fits the ring, every pool block of the call maps to ring generation 0, so no slot is reused inside +// the call and the wave schedule cannot wait on a slot release. Without a caller bound the capacity is tested, which never fits. +static constexpr int kSm90FusedCallWaveMaxTokens = 256; + +static bool sm90_fused_single_wave_pool_fits(const int& num_ranks, const int& num_topk, const int& num_wave_bound_tokens_per_rank, + const int& num_experts_per_rank, const int& block_m, const int& num_ring_tokens) { + return get_num_wave_pool_tokens_sm90_fused(num_ranks, num_topk, num_wave_bound_tokens_per_rank, num_experts_per_rank, block_m) <= num_ring_tokens; +} + +static bool sm90_fused_is_decode_single_wave_call(const int& block_m, const int& num_tokens) { + return block_m == 64 and num_tokens <= kSm90FusedSingleWaveMaxTokens; +} + +// Per-rank token count the M-keyed host rules are evaluated at: the caller's global bound when given (the same on every rank), else this rank's own count. +static int sm90_fused_rule_tokens_per_rank(const int& num_tokens, const int& num_tokens_bound) { + return num_tokens_bound > 0 ? num_tokens_bound : num_tokens; +} + +static bool sm90_fused_call_wave_class(const int& block_m, const int& num_tokens, const int& num_tokens_bound) { + if (block_m == 64) + return num_tokens <= kSm90FusedSingleWaveMaxTokens; + if (block_m == 128) + return sm90_fused_rule_tokens_per_rank(num_tokens, num_tokens_bound) <= kSm90FusedCallWaveMaxTokens; + return false; +} + +// Publish-batch rule: calls with at least this many tokens per rank publish kSm90FusedPullPublishBatch pull rows per release fence. +static constexpr int kSm90FusedPullPublishMinTokens = 8192; +static constexpr int kSm90FusedPullPublishBatch = 16; + +static constexpr int kSm90FusedPullEagerMaxTokens = 256; +static constexpr int kSm90FusedPdlMaxTokens = 256; + +static int get_num_max_pool_tokens_sm90_fused( + const int& num_ranks, const int& num_max_tokens_per_rank, const int& num_topk, + const int& num_experts_per_rank) { + return layout::get_num_max_pool_tokens_sm90(num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); +} + +static int get_num_padded_sf_pool_tokens_sm90_fused(const int& num_data_pool_tokens) { + int num_padded = 0; + for (const int& block_m: layout::kSM90FusedCandidateBlockM) + num_padded = std::max(num_padded, layout::get_num_sf_ring_tokens(num_data_pool_tokens, block_m)); + return num_padded; +} + +static bool is_candidate_block_m_sm90_fused(const int& block_m) { + return std::any_of(layout::kSM90FusedCandidateBlockM, layout::kSM90FusedCandidateBlockM + layout::kNumSM90CandidateBlockMs, + [=](const auto& candidate) { return candidate == block_m; }); +} + +static std::tuple get_block_config_sm90_fused( + const int& num_ranks, const int& num_experts, + const int& num_topk, const int& num_tokens) { + const float expected_tokens_per_expert = + static_cast(num_tokens) * num_ranks * num_topk / num_experts; + const bool auto_split_mn = expected_tokens_per_expert >= 64.0f; + if (auto_split_mn) // 2 math warpgroups on the 128x256 tile (64x256 each, spill-free) + return {128, 256}; + + const int block_m = 64; + const int num_epilogue_warpgroups = 2; + + DG_HOST_ASSERT(is_candidate_block_m_sm90_fused(block_m)); + return {block_m, num_epilogue_warpgroups * 128}; +} + +static int get_num_experts_per_wave_sm90_fused( + 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, + const int& num_ring_tokens, const int& num_max_tokens_per_rank, const int& num_ranks, + // Per-rank token count the ring-capacity clamp is evaluated at (<= num_max_tokens_per_rank) + const int& num_wave_bound_tokens_per_rank, + const int& num_tokens_bound = 0) { + // Ring mode must derive the wave size from the ring capacity first and + // foremost: a wave whose pool exceeds the ring would reuse slots inside a + // single wave (L1 epilogue of generation g+1 waits on L2 consumption that + // only happens after the wave's whole L1 phase) and deadlock. The + // occupancy-driven early returns below are only safe when the ring covers + // the full pool. + const int num_max_pool_tokens = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + const bool call_pool_fits_ring = num_ring_tokens < num_max_pool_tokens and + sm90_fused_call_wave_class(block_m, num_tokens, num_tokens_bound) and + sm90_fused_single_wave_pool_fits(num_ranks, num_topk, num_wave_bound_tokens_per_rank, num_experts_per_rank, block_m, num_ring_tokens); + if (num_ring_tokens < num_max_pool_tokens and not call_pool_fits_ring) + return get_capped_num_experts_per_wave_sm90_fused( + num_experts_per_rank, num_tokens, num_topk, + intermediate_hidden, block_m, block_n, num_sms, + num_ring_tokens, num_wave_bound_tokens_per_rank, num_ranks); + + if (sm90_fused_is_decode_single_wave_call(block_m, num_tokens)) + return num_experts_per_rank; + + const float expected_tokens_per_expert = + static_cast(num_tokens) * num_topk / num_experts_per_rank; + if (block_m == 64 and expected_tokens_per_expert > 4.0f) { + const int num_n_blocks_per_expert = (2 * intermediate_hidden) / block_n; + const int wave = std::max(1, num_sms / std::max(1, num_n_blocks_per_expert)); + return std::min(wave, num_experts_per_rank); + } + if (expected_tokens_per_expert < 1.0f or expected_tokens_per_expert > 4.0f) + return num_experts_per_rank; + + if (block_m == 64 and intermediate_hidden >= 3072) { + const int num_n_blocks_per_expert = (2 * intermediate_hidden) / block_n; + const int single_wave_blocks = + num_experts_per_rank * num_n_blocks_per_expert; + if (single_wave_blocks >= 4 * num_sms) + return num_experts_per_rank; + } + return get_capped_num_experts_per_wave_sm90_fused( + num_experts_per_rank, num_tokens, num_topk, + intermediate_hidden, block_m, block_n, num_sms, + num_ring_tokens, num_max_tokens_per_rank, num_ranks); +} + +static bool should_use_swap_ab_sm90_fused( + const int& num_experts_per_rank, const int& num_tokens, const int& num_topk, + const int& block_m, const int& num_epilogue_threads, const int& l2_act_sf_gran_k) { + const float expected_tokens_per_expert = + static_cast(num_tokens) * num_topk / num_experts_per_rank; + const bool decode_split_n_path = + block_m == 64 and num_epilogue_threads == 256; + // The swapAB L1 epilogue hands post-SwiGLU values across warpgroups before quantization, which is not bitwise stable run to run, + // so swapAB is selected only where one warpgroup owns a whole L2 activation-SF group of the tile (the kernel asserts the same). + constexpr int kSwapABWarpgroupOutputColumns = (128 / 2) / 2; + if (decode_split_n_path and kSwapABWarpgroupOutputColumns < l2_act_sf_gran_k) + return false; + return decode_split_n_path and num_tokens <= 128 and expected_tokens_per_expert > 0.0f; +} + +static std::tuple get_block_mn_config_sm90_fused( + const int& num_ranks, const int& num_experts, const int& num_experts_per_rank, + const int& num_tokens, const int& num_topk, + const int& hidden, const int& intermediate_hidden, const int& l2_act_sf_gran_k) { + const auto [block_m, num_epilogue_threads] = get_block_config_sm90_fused( + num_ranks, num_experts, num_topk, num_tokens); + const float expected_tokens_per_expert = + static_cast(num_tokens) * num_ranks * num_topk / num_experts; + const bool split_m2 = + block_m == 128 and num_epilogue_threads == 256; + const bool decode_split_n_path = + block_m == 64 and num_epilogue_threads == 256; + // The split-N decode tile's post-SwiGLU output (BLOCK_N / 2 columns) must be exactly one L2 activation-SF group (kernel + // static_assert): with the per-128 SF only the 256-wide tile qualifies; the 128-wide tile is reserved for the per-64 recipe. + // The 2-CTA cluster pairs adjacent N blocks, so both N block counts must be even at the chosen + // BLOCK_N (scheduler/mega_moe.cuh kNumL1BlockNs/kNumL2BlockNs): 256-wide needs multiples of 512. + const bool decode_tile_n_256_fits = (2 * intermediate_hidden) % 512 == 0 and hidden % 512 == 0; + const bool decode_use_block_n_256 = decode_split_n_path and decode_tile_n_256_fits and + (l2_act_sf_gran_k == 128 or (intermediate_hidden >= 2048 and expected_tokens_per_expert >= 0.25f)); + DG_HOST_ASSERT((not decode_split_n_path or decode_use_block_n_256 or l2_act_sf_gran_k == 64) && + "the 128-wide decode tile needs the per-64 L2 activation scale: with the per-128 scale hidden and " + "2 x intermediate_hidden must be multiples of 512 (256-wide tile, even N block counts)"); + const bool use_swap_ab = (not decode_use_block_n_256) and + should_use_swap_ab_sm90_fused( + num_experts_per_rank, num_tokens, num_topk, + block_m, num_epilogue_threads, l2_act_sf_gran_k); + const int block_n = use_swap_ab ? 128 + : (split_m2 ? 256 : + (decode_use_block_n_256 ? 256 : 128)); + return {block_m, block_n, num_epilogue_threads, use_swap_ab}; +} + +// Lag schedule: the ring holds the lag + group units plus the in-flight blocks (get_lag_ring_blocks_sm90_fused), sized at BLOCK_M 128. +static std::pair get_lag_ring_tokens_sm90_fused( + const int& num_ranks, const int& num_max_tokens_per_rank, const int& num_topk, const int& num_experts_per_rank) { + const int lag = get_auto_late_lag_sm90_fused(num_max_tokens_per_rank, num_topk, num_experts_per_rank); + const int lag_encoded = get_l2_lag_encoded_sm90_fused(lag); + const int num_max_pool_tokens = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + const int tokens = get_lag_ring_blocks_sm90_fused(lag, sm90_fused_lag_group_of(lag_encoded)) * 128; + // Clamping to the full pool leaves the encoded lag describing more ring than is allocated, which is + // safe: a pool-sized ring never wraps, so the lag schedule has no capacity constraint to violate (and the + // launch skips the wrap assertion for it). + return {std::min(math::align(tokens, layout::kSM90FusedLCMBlockM), num_max_pool_tokens), lag_encoded}; +} + +// Wave schedule: the ring capacity is derived from the wave size, mirroring SM100's causality +// (E_wave is decided first, the pool capacity is derived from it) instead of +// deriving E_wave from a user-supplied capacity. -1 = auto wave size from the occupancy heuristic, +// N > 0 = fixed wave size for the derivation only. +static int get_wave_ring_tokens_sm90_fused( + const int& num_ranks, const int& num_experts, const int& num_experts_per_rank, + const int& num_max_tokens_per_rank, const int& num_topk, + const int& hidden, const int& intermediate_hidden, + const int& num_experts_per_wave_knob, const int& l2_act_sf_gran_k) { + // Size with the worst-case token count: block_m is monotone in num_tokens, + // so the sizing block_m >= any call-time block_m and the call-time capacity + // clamp never fires for buffers derived here. + const auto [block_m, block_n, num_epilogue_threads, use_swap_ab] = get_block_mn_config_sm90_fused( + num_ranks, num_experts, num_experts_per_rank, + num_max_tokens_per_rank, num_topk, hidden, intermediate_hidden, l2_act_sf_gran_k); + const int num_max_pool_tokens = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + int num_experts_per_wave; + if (num_experts_per_wave_knob > 0) { + num_experts_per_wave = std::min(num_experts_per_wave_knob, num_experts_per_rank); + } else { + // Auto: the capacity is what this call is about to derive, so the wave comes from the occupancy + // heuristic with the full pool as the bound -- the cap is then a no-op, since a wave of every local + // expert spans at most a full pool. The per-call chooser cannot be used here: its gates read the + // stored ring capacity, which does not exist yet. + num_experts_per_wave = get_capped_num_experts_per_wave_sm90_fused( + num_experts_per_rank, num_max_tokens_per_rank, num_topk, + intermediate_hidden, block_m, block_n, device_runtime->get_num_sms(), + num_max_pool_tokens, num_max_tokens_per_rank, num_ranks); + // Floor at 2: a wave-1 ring holds exactly the recv working set, leaving no + // headroom between producer and consumer. + num_experts_per_wave = std::min(std::max(num_experts_per_wave, 2), num_experts_per_rank); + } + return std::min(math::align(get_num_wave_pool_tokens_sm90_fused( + num_ranks, num_topk, num_max_tokens_per_rank, num_experts_per_wave, block_m), + layout::kSM90FusedLCMBlockM), + num_max_pool_tokens); +} + +// Ring capacity (0 = full-pool buffers, no ring) and encoded L2 lag (0 = wave schedule) of a symm buffer. +static std::pair get_num_ring_tokens_for_sm90_fused_mega_moe( + const int& num_ranks, const int& num_experts, const int& num_experts_per_rank, + const int& num_max_tokens_per_rank, const int& num_topk, + const int& hidden, const int& intermediate_hidden, + const int& num_experts_per_wave_knob, const int& l2_act_sf_gran_k) { + const auto schedule = get_buffer_schedule_sm90_fused(num_max_tokens_per_rank, num_experts_per_wave_knob); + if (schedule == Sm90FusedBufferSchedule::kFullPool) + return {0, 0}; + DG_HOST_ASSERT(num_max_tokens_per_rank % layout::kSM90FusedLCMBlockM == 0 and + "num_max_tokens_per_rank must be token-aligned before deriving a ring capacity"); + if (schedule == Sm90FusedBufferSchedule::kLag) + return get_lag_ring_tokens_sm90_fused(num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + return {get_wave_ring_tokens_sm90_fused( + num_ranks, num_experts, num_experts_per_rank, + num_max_tokens_per_rank, num_topk, hidden, intermediate_hidden, + num_experts_per_wave_knob, l2_act_sf_gran_k), + 0}; +} + +// Weight-SF staging slot of one math warpgroup (floats; kernel `kNumWeightSFFloatsPerWG`); `wg_block_n` is the warpgroup's N extent. +static int get_weight_sf_floats_per_warpgroup_sm90_fused(const int& hidden, const int& intermediate_hidden, const int& wg_block_n) { + const int num_sf_groups_per_wg = wg_block_n >= 128 ? wg_block_n / 128 : 1; + return std::max(2 * (hidden / 128), num_sf_groups_per_wg * (intermediate_hidden / 128)); +} + +// L2 stage mode 4 parks the four unwanted rows of every stmatrix.x4 in the warpgroup's weight-SF slot, which must hold at least this many bytes (kernel static_assert). +static constexpr int kSm90FusedStmatrixJunkSlotBytes = 256; + +static std::pair get_pipeline_config_sm90_fused( + const int& smem_capacity, + const int& num_experts, const int& hidden, const int& intermediate_hidden, + const int& block_m, const int& block_n, const int& block_k, + const int& num_dispatch_warps, const int& num_epilogue_warps, + const bool& use_swap_ab = false, const bool& half_l2_cd = false, + const int& l2_cd_passes = 0, + const int& tile_table_entries = 0, + const int& l2_act_sf_gran_k = 64, + const bool& early_combine = false, + const int& pull_publish_batch = 1) { + constexpr int kSmemAlignment = 1024; + const int num_l2_cd_passes = l2_cd_passes != 0 ? l2_cd_passes : (half_l2_cd ? 2 : 1); + + const int smem_expert_count_bytes = num_experts * static_cast(sizeof(uint32_t)); + const int smem_send_buffers_size = align( + static_cast(layout::Buffer(layout::Data(hidden), num_dispatch_warps, 1).get_num_bytes()), + kSmemAlignment); + const int smem_dispatch_size = smem_send_buffers_size; + + const int smem_cd_l1 = block_m * (block_n / 2); + const int smem_cd_l2 = block_m * (block_n / num_l2_cd_passes) * static_cast(sizeof(nv_bfloat16)); + const int smem_cd_swap_l1 = use_swap_ab + ? block_m * (block_n / 2) * + (static_cast(sizeof(float)) + static_cast(sizeof(uint8_t))) + : 0; + const int smem_cd = align( + std::max(std::max(smem_cd_l1, smem_cd_l2), smem_cd_swap_l1), + kSmemAlignment); + + const int smem_sfa_per_stage = + align((l2_act_sf_gran_k == block_k ? 1 : 2) * block_m * static_cast(sizeof(float)), 128); + const int smem_sfb_per_stage = 0; + const int smem_per_stage = block_m * block_k + block_n * block_k + + smem_sfa_per_stage + smem_sfb_per_stage; + + const int num_epilogue_warpgroups = num_epilogue_warps / 4; + const int wg_block_n = (block_m == 64 and num_epilogue_warpgroups > 1) ? block_n / num_epilogue_warpgroups : + ((block_m == 128 and block_n == 256 and num_epilogue_warpgroups == 4) ? 128 : block_n); + const int weight_sf_floats_per_wg = get_weight_sf_floats_per_warpgroup_sm90_fused(hidden, intermediate_hidden, wg_block_n); + const int smem_weight_sf = align( + num_epilogue_warpgroups * weight_sf_floats_per_wg * static_cast(sizeof(float)), 128); + + const int smem_barriers_fixed = (num_dispatch_warps + 2 * num_epilogue_warps) * 8; + const int smem_barriers_per_stage = 2 * 8; + const int smem_early_combine = early_combine ? num_dispatch_warps * 8 + 64 + 8 : 0; + const int smem_tile_table = tile_table_entries * 8; + const int smem_tile_table_region = std::max(smem_tile_table, smem_expert_count_bytes) + 8; + const int smem_pull_pending = pull_publish_batch > 1 ? + align(num_dispatch_warps * pull_publish_batch * static_cast(sizeof(uint32_t)), 16) + 8 : 0; + const int smem_fixed = smem_dispatch_size + smem_cd + smem_weight_sf + smem_barriers_fixed + smem_early_combine + smem_tile_table_region + smem_pull_pending; + + const int num_stages = (smem_capacity - smem_fixed) / + (smem_per_stage + smem_barriers_per_stage); + DG_HOST_ASSERT(num_stages >= 2); + const int smem_size = smem_fixed + num_stages * (smem_per_stage + smem_barriers_per_stage); + DG_HOST_ASSERT(smem_size <= smem_capacity); + return {num_stages, smem_size}; +} + +// Mirror the kernel's pre-barrier region (send buffers, CD staging and the GEMM stages) so a hidden size the +// combine vectorization cannot serve is rejected before NVRTC sees the specialization. +static uint32_t get_sm90_fused_pre_barrier_smem_size_for_combine( + const int& hidden, const MegaMoESM90FusedConfig& config, const bool& use_swap_ab) { + constexpr int kSmemAlignment = 1024; + const int num_dispatch_warps = config.num_dispatch_threads / 32; + const int num_l2_cd_passes = config.l2_cd_passes != 0 ? config.l2_cd_passes : (config.half_l2_cd ? 2 : 1); + const int smem_send_buffers = align( + static_cast(layout::Buffer(layout::Data(hidden), num_dispatch_warps, 1).get_num_bytes()), + kSmemAlignment); + const int smem_cd_l1 = config.block_m * (config.block_n / 2); + const int smem_cd_l2 = config.block_m * (config.block_n / num_l2_cd_passes) * + static_cast(sizeof(nv_bfloat16)); + const int smem_cd_swap_l1 = use_swap_ab ? + config.block_m * (config.block_n / 2) * + (static_cast(sizeof(float)) + static_cast(sizeof(uint8_t))) : 0; + const int smem_cd = align( + std::max(std::max(smem_cd_l1, smem_cd_l2), smem_cd_swap_l1), kSmemAlignment); + const int smem_gemm = config.num_stages * + (config.block_m * config.block_k + config.block_n * config.block_k); + return static_cast(smem_send_buffers + smem_cd + smem_gemm); +} + +static MegaMoESM90FusedConfig get_mega_moe_config_sm90_fused( + 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, + const int& num_padded_sf_pool_tokens, + const int& l2_act_sf_gran_k, + const int& num_ring_tokens = 0, + // > 0: the caller's bound on this call's token count on EVERY rank (the same value on every rank); 0: use the capacity + const int& num_tokens_bound = 0, + // encoded L2-lag schedule the symm buffer was sized for (the value the sizing returned); -1 = none, i.e. the wave schedule + const int& l2_lag_encoded_in = -1) { + const auto [block_m, block_n, num_epilogue_threads, use_swap_ab] = get_block_mn_config_sm90_fused( + num_ranks, num_experts, num_experts_per_rank, + num_tokens, num_topk, hidden, intermediate_hidden, l2_act_sf_gran_k); + const int l2_lag_encoded = l2_lag_encoded_in >= 0 ? l2_lag_encoded_in : 0; + const int block_k = 128; + const int num_sms = device_runtime->get_num_sms(); + const bool cluster_pairing_valid = + num_sms % 2 == 0 and + ((2 * intermediate_hidden) / block_n) % 2 == 0 and + (hidden / block_n) % 2 == 0 and + (2 * intermediate_hidden) % block_n == 0 and hidden % block_n == 0; + const int cluster_size = (block_m == 128 and block_n == 256 and cluster_pairing_valid) ? 2 : 1; + const bool multicast_on_b = cluster_size == 2; + const int num_max_pool_tokens = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + // 0 means "no ring": data pools sized by the full pool (legacy behavior). + const int num_effective_ring_tokens = num_ring_tokens == 0 ? num_max_pool_tokens : num_ring_tokens; + DG_HOST_ASSERT(num_tokens_bound == 0 or num_tokens_bound >= num_tokens); + // the M-keyed rules below read the bound as given, so a bound above the capacity would pick a schedule for traffic this buffer cannot receive + DG_HOST_ASSERT(num_tokens_bound <= num_max_tokens_per_rank); + const int call_bound = num_tokens_bound > 0 ? num_tokens_bound : num_max_tokens_per_rank; + const int num_wave_bound_tokens_per_rank = + std::min(num_max_tokens_per_rank, align(std::max(call_bound, 1), layout::kSM90FusedLCMBlockM)); + // a qualifying call on a lag ring whose worst-case single-wave pool fits runs the wave schedule (kL2LagUnits 0) on the same ring + const int l2_lag_units_buffer = sm90_fused_lag_units_of(l2_lag_encoded); + const bool call_wave_schedule = l2_lag_units_buffer > 0 and num_ring_tokens != 0 and + sm90_fused_call_wave_class(block_m, num_tokens, num_tokens_bound) and + sm90_fused_single_wave_pool_fits(num_ranks, num_topk, num_wave_bound_tokens_per_rank, num_experts_per_rank, block_m, num_ring_tokens); + // every other call on a lag ring runs the lag of its own per-rank token bound, never above the buffer's (the ring demand grows with the lag) + const int l2_lag_units_call = std::min(l2_lag_units_buffer, + get_auto_late_lag_sm90_fused(num_wave_bound_tokens_per_rank, num_topk, num_experts_per_rank)); + DG_HOST_ASSERT(get_lag_ring_blocks_sm90_fused(l2_lag_units_call, sm90_fused_lag_group_of(l2_lag_encoded)) <= + get_lag_ring_blocks_sm90_fused(l2_lag_units_buffer, sm90_fused_lag_group_of(l2_lag_encoded)) && + "the buffer's ring must hold the call's lag schedule"); + const int l2_lag_encoded_launch = call_wave_schedule ? 0 : get_l2_lag_encoded_sm90_fused(l2_lag_units_call); + const int l2_lag_units = sm90_fused_lag_units_of(l2_lag_encoded_launch); + DG_HOST_ASSERT((l2_lag_units == 0 or num_ring_tokens != 0) && "the L2-lag schedule needs a ring-sized buffer"); + if (num_ring_tokens != 0) { + DG_HOST_ASSERT(num_ring_tokens % layout::kSM90FusedLCMBlockM == 0); + DG_HOST_ASSERT(num_ring_tokens <= num_max_pool_tokens); + if (l2_lag_units > 0) { + DG_HOST_ASSERT(num_ring_tokens % block_m == 0); + // the wrap constraint is the allocation's own block count (get_lag_ring_blocks_sm90_fused, sized at BLOCK_M 128 for the buffer's + // lag; the call's lag and BLOCK_M are at most those). A ring clamped to the full pool never reuses a slot, so it is exempt + if (num_ring_tokens < num_max_pool_tokens) { + const int group = sm90_fused_lag_group_of(l2_lag_encoded_launch); + DG_HOST_ASSERT(num_ring_tokens / block_m >= get_lag_ring_blocks_sm90_fused(l2_lag_units, group) && + "Lag schedule: the ring must hold the lag + group units plus the in-flight blocks"); + } + } else { + const int num_min_ring_tokens = get_num_wave_pool_tokens_sm90_fused( + num_ranks, num_topk, call_wave_schedule ? num_wave_bound_tokens_per_rank : num_max_tokens_per_rank, 1, block_m); + DG_HOST_ASSERT(num_ring_tokens >= num_min_ring_tokens && + "Ring capacity must be within [tokens from all ranks, full pool]"); + } + } + const int swizzle_acts_mode = 128; + const int swizzle_weights_mode = 128; + + // The wave size is re-derived from the (derived-at-allocation) capacity via + // the occupancy heuristic, with the capacity clamp evaluated at this call's per-rank token bound (num_tokens_bound). + const int num_experts_per_wave = l2_lag_units > 0 ? num_experts_per_rank : + get_num_experts_per_wave_sm90_fused( + num_experts_per_rank, num_tokens, num_topk, + intermediate_hidden, block_m, block_n, num_sms, + num_effective_ring_tokens, num_max_tokens_per_rank, num_ranks, + num_wave_bound_tokens_per_rank, num_tokens_bound); + DG_HOST_ASSERT((not call_wave_schedule or block_m != 128 or num_experts_per_wave == num_experts_per_rank) && + "the BLOCK_M-128 per-call wave schedule expects one wave of all local experts"); + + const bool reduce_decode_threads = num_epilogue_threads == 128; + const bool decode_split_n = + block_m == 64 and num_epilogue_threads == 256; + const bool split_m2 = + block_m == 128 and num_epilogue_threads == 256; + const bool shrink_non_epilogue = reduce_decode_threads or decode_split_n or split_m2; + const int num_dispatch_threads = + (num_epilogue_threads == 512 or shrink_non_epilogue) ? 64 : 128; + const bool split_sfa_loader_warp = false; + const int num_non_epilogue_threads = + split_sfa_loader_warp ? 128 : + ((num_epilogue_threads == 512 or shrink_non_epilogue) ? 64 : 128); + DG_HOST_ASSERT((num_dispatch_threads + num_non_epilogue_threads) % 128 == 0); + + // Early combine (kernel kEarlyCombineMode 1) is selected from hidden >= 512 x topk (the staging bound of the original dispatch-warp + // receiver, kept as the selection threshold) + const int early_combine_fits_smem = hidden >= 512 * num_topk ? 1 : 0; + const int tile_table_entries = layout::get_sm90_tile_table_entries_compact( + num_max_pool_tokens, block_m, (2 * intermediate_hidden) / block_n, hidden / block_n, num_sms); + const bool ring_mode_pull = num_ring_tokens != 0 and num_effective_ring_tokens < num_max_pool_tokens; + const auto pull_publish_batch_fits_ring = [&](const int& n) { + return not ring_mode_pull or + n * num_sms * (num_dispatch_threads / 32) < (num_effective_ring_tokens / block_m - num_experts_per_rank - 1) * block_m; + }; + const bool pull_publish_rule_on = sm90_fused_rule_tokens_per_rank(num_tokens, num_tokens_bound) >= kSm90FusedPullPublishMinTokens and + pull_publish_batch_fits_ring(kSm90FusedPullPublishBatch); + const int pull_publish_batch = pull_publish_rule_on ? kSm90FusedPullPublishBatch : 1; + // Publish-batch bound (ring mode): a dispatch warp at pool block p waits for the L1 consumers of block p - R (R = ring blocks), + // which need every row of that block published. The kernel publishes a warp's pending rows before it blocks on a slot; this + // bound additionally keeps the batch's unpublished span (N x 264 / block_m + one partial tail block per expert) inside the ring. + if (pull_publish_batch > 1 and ring_mode_pull) { + DG_HOST_ASSERT(pull_publish_batch_fits_ring(pull_publish_batch) && + "publish batch too large for the ring: a warp's pending rows could be needed by the ring slot it waits for"); + } + auto [num_stages, smem_size] = get_pipeline_config_sm90_fused( + SM90ArchSpec::smem_capacity, + num_experts, hidden, intermediate_hidden, + block_m, block_n, block_k, + num_dispatch_threads / 32, num_epilogue_threads / 32, + use_swap_ab, /*half_l2_cd=*/false, + 0, tile_table_entries, l2_act_sf_gran_k, early_combine_fits_smem != 0, pull_publish_batch); + + const bool split_phase_prefill = + block_m == 128 and block_n == 256 and hidden >= 4096; + bool half_l2_cd = false; + int l2_cd_passes = 0; + int l2_stage_mode = 0; + if (split_m2) { + // row passes on the quarter-width buffer: stmatrix (mode 4) when the weight-SF slot holds kSm90FusedStmatrixJunkSlotBytes, else plain stores (mode 2) + const int weight_sf_slot_bytes = + get_weight_sf_floats_per_warpgroup_sm90_fused(hidden, intermediate_hidden, block_n) * static_cast(sizeof(float)); + l2_stage_mode = use_swap_ab ? 0 : (weight_sf_slot_bytes >= kSm90FusedStmatrixJunkSlotBytes ? 4 : 2); + l2_cd_passes = l2_stage_mode != 0 ? 4 : 2; + const auto [ns_half, sz_half] = get_pipeline_config_sm90_fused( + SM90ArchSpec::smem_capacity, + num_experts, hidden, intermediate_hidden, + block_m, block_n, block_k, + num_dispatch_threads / 32, num_epilogue_threads / 32, + use_swap_ab, /*half_l2_cd=*/true, + l2_cd_passes, tile_table_entries, l2_act_sf_gran_k, early_combine_fits_smem != 0, pull_publish_batch); + half_l2_cd = true; + num_stages = ns_half; + smem_size = sz_half; + } else if (split_phase_prefill and num_stages < 3) { + const auto [ns_half, sz_half] = get_pipeline_config_sm90_fused( + SM90ArchSpec::smem_capacity, + num_experts, hidden, intermediate_hidden, + block_m, block_n, block_k, + num_dispatch_threads / 32, num_epilogue_threads / 32, + use_swap_ab, /*half_l2_cd=*/true, + 0, tile_table_entries, l2_act_sf_gran_k, early_combine_fits_smem != 0, pull_publish_batch); + if (ns_half > num_stages) { + half_l2_cd = true; + num_stages = ns_half; + smem_size = sz_half; + } + } + + // The early-combine signal reads the k-block kNumStages ahead of the current tile, so both GEMMs need + // strictly more k-blocks than stages (impls/sm90_fp8_fused_mega_moe.cuh, kEarlyCombine). num_stages is only + // final here; the smem above was sized with early combine on, so demoting now over-allocates, never under. + const int early_combine = (early_combine_fits_smem != 0 and + hidden / block_k > num_stages and + intermediate_hidden / block_k > num_stages) ? 1 : 0; + + const auto config = MegaMoESM90FusedConfig { + block_m, block_n, block_k, + cluster_size, + num_max_pool_tokens, num_padded_sf_pool_tokens, + num_effective_ring_tokens, + swizzle_acts_mode, swizzle_weights_mode, + num_experts_per_wave, + num_stages, smem_size, + num_dispatch_threads, num_non_epilogue_threads, num_epilogue_threads, + half_l2_cd, + multicast_on_b, + l2_lag_encoded_launch, + l2_cd_passes, + l2_stage_mode, + early_combine, + pull_publish_batch + }; + + + if (get_env("DG_JIT_DEBUG") or get_env("DG_PRINT_CONFIGS")) { + const auto key = fmt::format( + "MegaMoESM90FusedConfig(num_ranks={}, num_experts={}, hidden={}, intermediate_hidden={}, num_max_tokens_per_rank={}, num_tokens={}, num_tokens_bound={}, num_topk={}, swap_ab={})", + num_ranks, num_experts, hidden, intermediate_hidden, num_max_tokens_per_rank, num_tokens, + num_tokens_bound, num_topk, use_swap_ab); + static std::unordered_set printed; + if (printed.count(key) == 0) { + std::cout << key << ": " << config << std::endl; + printed.insert(key); + } + } + return config; +} + +} // namespace deep_gemm diff --git a/csrc/jit_kernels/impls/sm90_fp8_fused_mega_moe.hpp b/csrc/jit_kernels/impls/sm90_fp8_fused_mega_moe.hpp new file mode 100644 index 0000000000..805b82352b --- /dev/null +++ b/csrc/jit_kernels/impls/sm90_fp8_fused_mega_moe.hpp @@ -0,0 +1,400 @@ +#pragma once + +#include +#include "../../jit/compiler.hpp" +#include "../../jit/kernel_runtime.hpp" +#include "../../utils/exception.hpp" +#include "../../utils/format.hpp" +#include "runtime_utils.hpp" + +#include +#include +#include + +#include "../heuristics/sm90_fused_mega_moe.hpp" + +namespace deep_gemm { + +// ============================================================================ +// SM90 (Hopper) FP8 MegaMoE host runtime +// ---------------------------------------------------------------------------- +// This is the SM90 counterpart of `SM100FP8FP4MegaMoERuntime`. The kernel +// itself lives in `deep_gemm/impls/sm90_fp8_fused_mega_moe.cuh`. +// +// Differences from SM100 path: +// * Activations and weights are both FP8 (e4m3); no FP4. +// * Activation/weight scale factors (SF) are float, not UE8M0 int + per-32 +// UTCCP layout. L1 activation SF and weight SF are per-128 K; the fused L1 +// epilogue writes the L2 activation SF at per-128 or per-64 K granularity (kL2ActSFK). +// * No tensor memory: WGMMA accumulators are register-resident. +// * Cluster size is at most 2 (TMA multicast on A); no 2-CTA UMMA. +// ============================================================================ + +class SM90FP8FusedMegaMoERuntime final : public LaunchRuntime { +public: + struct Args { + // Templated arguments + int num_max_tokens_per_rank; + int hidden, intermediate_hidden; + int num_experts, num_topk; + int num_ranks; + float activation_clamp; + bool fast_math; + int epilogue_registers; + bool reuse_accum_as_final; + bool l2_arrival_counter; + bool l2_epilogue_requires_full_sync; + bool use_swap_ab; + bool half_l2_cd; + int num_ring_tokens; + MegaMoESM90FusedConfig config; + + // Runtime arguments + void* y; + int* cumulative_local_expert_recv_stats; + int num_tokens; + layout::SymBuffer<> sym_buffer_ptrs; + // in-kernel event trace buffer; nullptr = tracing compiled out (kTrace = false) + void* trace_ptr; + bool trace_epi; + // L2 activation SF K granularity, 128 or 64 + int l2_act_sf_k; + bool pdl; + bool pull_eager_publish; + + // Tensormaps for activations and weights. Weight scale factors use + // block (128, 128) quantization and are loaded by the math warpgroup + // directly from global memory (no TMA descriptor required). + CUtensorMap tensor_map_l1_acts; + CUtensorMap tensor_map_l1_acts_sf; + CUtensorMap tensor_map_l1_weights; + const float* l1_weights_sf; + CUtensorMap tensor_map_l1_output; + CUtensorMap tensor_map_l2_acts; + CUtensorMap tensor_map_l2_acts_sf; + CUtensorMap tensor_map_l2_weights; + const float* l2_weights_sf; + + // Launch configs + LaunchArgs launch_args; + }; + + static std::string generate_impl(const Args& args) { + return fmt::format(R"( +#include + +using namespace deep_gemm; + +static void __instantiate_kernel() {{ + auto ptr = reinterpret_cast(&sm90_fp8_fused_mega_moe_impl< + {}, + {}, {}, + {}, {}, + {}, + {}, {}, {}, + {}, + {}, + {}, + {}, + {}, {}, {}, + {}, {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {}, + {} + >); +}}; +)", + args.num_max_tokens_per_rank, + args.hidden, args.intermediate_hidden, + args.num_experts, args.num_topk, + args.config.num_experts_per_wave, + args.config.block_m, args.config.block_n, args.config.block_k, + args.config.num_max_pool_tokens, + args.config.num_padded_sf_pool_tokens, + // Must be the *effective* capacity (full pool when the caller passed 0): + // the kernel derives `kRingCoversFullPool` and `kNumRingBlocks` from it. + args.config.num_ring_tokens, + args.config.num_stages, + args.config.num_dispatch_threads, args.config.num_non_epilogue_threads, args.config.num_epilogue_threads, + args.launch_args.grid_dim.first, args.num_ranks, + to_string(args.activation_clamp), + args.fast_math ? "true" : "false", + args.epilogue_registers, + args.reuse_accum_as_final ? "true" : "false", + args.l2_arrival_counter ? "true" : "false", + args.l2_epilogue_requires_full_sync ? "true" : "false", + args.use_swap_ab ? "true" : "false", + args.half_l2_cd ? "true" : "false", + args.config.cluster_size, + args.config.multicast_on_b ? "true" : "false", + args.config.l2_lag_units, + args.config.l2_cd_passes, + args.trace_ptr != nullptr ? "true" : "false", + args.trace_epi ? "true" : "false", + args.l2_act_sf_k, + args.config.l2_stage_mode, + args.pdl ? "true" : "false", + args.config.early_combine, + args.config.pull_publish_batch, + args.pull_eager_publish ? "true" : "false"); + } + + static void launch_impl(const KernelHandle& kernel, const LaunchConfigHandle& config, Args args) { + DG_CUDA_UNIFIED_CHECK(launch_kernel(kernel, config, + args.y, + args.cumulative_local_expert_recv_stats, + args.num_tokens, + args.sym_buffer_ptrs, + args.tensor_map_l1_acts, + args.tensor_map_l1_acts_sf, + args.tensor_map_l1_weights, + args.l1_weights_sf, + args.tensor_map_l1_output, + args.tensor_map_l2_acts, + args.tensor_map_l2_acts_sf, + args.tensor_map_l2_weights, + args.l2_weights_sf, + args.trace_ptr + )); + } +}; + +static void sm90_fp8_fused_mega_moe( + const torch::Tensor& y, + const torch::Tensor& l1_acts, const torch::Tensor& l1_acts_sf, + const torch::Tensor& l2_acts, const torch::Tensor& l2_acts_sf, + const torch::Tensor& l1_weights, const torch::Tensor& l2_weights, + const torch::Tensor& l1_weights_sf, const torch::Tensor& l2_weights_sf, + const std::optional cumulative_local_expert_recv_stats, + const std::vector& sym_buffer_ptrs, + const int& rank_idx, const int& num_max_tokens_per_rank, + const int& num_experts_per_rank, + const int& num_tokens, const int& num_topk, + const int& hidden, const int& intermediate_hidden, + const int& num_ring_tokens, + const int& l2_act_sf_gran_k, + const float& activation_clamp, + const bool& fast_math, + const int& num_tokens_bound = 0, + const int& l2_lag_encoded = -1 +) { + const auto num_ranks = static_cast(sym_buffer_ptrs.size()); + const auto num_experts = num_experts_per_rank * num_ranks; + const auto num_max_pool_tokens_h = get_num_max_pool_tokens_sm90_fused( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + // SF pools are sized by the data pool (ring capacity when ringing), so the + // padded token count must be derived from the same base as the allocation + // side in `get_symm_buffer_size_for_sm90_fused_mega_moe`. + const int num_data_pool_tokens = + num_ring_tokens == 0 ? num_max_pool_tokens_h : num_ring_tokens; + const int num_padded_sf_pool_tokens = get_num_padded_sf_pool_tokens_sm90_fused(num_data_pool_tokens); + DG_HOST_ASSERT(static_cast(l1_acts_sf.size(0)) == num_padded_sf_pool_tokens); + + // Heuristics + const auto config = get_mega_moe_config_sm90_fused( + num_ranks, num_experts, num_experts_per_rank, + num_max_tokens_per_rank, num_tokens, num_topk, + hidden, intermediate_hidden, num_padded_sf_pool_tokens, l2_act_sf_gran_k, + num_ring_tokens, num_tokens_bound, l2_lag_encoded); + // a BLOCK_M outside the SM90 candidate set would index SF rows past the pool sized above + DG_HOST_ASSERT(is_candidate_block_m_sm90_fused(config.block_m) and + config.num_padded_sf_pool_tokens >= (num_data_pool_tokens / config.block_m) * align(config.block_m, 128)); + const int default_epilogue_registers = + config.num_epilogue_threads == 512 ? 112 : 0; + const int epilogue_registers = default_epilogue_registers; + if (epilogue_registers > 0) { + const int dispatch_registers = + config.num_epilogue_threads == 512 ? 32 : 48; + const int non_epilogue_registers = + config.num_epilogue_threads == 512 ? 24 : 40; + DG_HOST_ASSERT(dispatch_registers * config.num_dispatch_threads + + non_epilogue_registers * config.num_non_epilogue_threads + + epilogue_registers * config.num_epilogue_threads <= 64512); + } + const bool reuse_accum_as_final = config.block_m == 128; + const bool default_split_mn_barrier_opt = + config.block_m == 128 and config.block_n == 256 and + (config.num_epilogue_threads == 512 or config.num_epilogue_threads == 256); + const bool decode_split_n_path = + config.block_m == 64 and config.num_epilogue_threads == 256; + const bool decode_split_n_bn256 = + decode_split_n_path and config.block_n == 256; + const bool decode_l2_counter = + decode_split_n_bn256 and num_tokens >= 4 and num_tokens <= 128; + const bool l2_arrival_counter = + default_split_mn_barrier_opt or decode_l2_counter; + const bool l2_epilogue_requires_full_sync = + not l2_arrival_counter; + // the swapAB epilogues assume per-64 L2 activation scales (kernel static_assert); per-128 takes the non-swap path + const bool use_swap_ab = config.block_n == 128 and l2_act_sf_gran_k == 64 and + should_use_swap_ab_sm90_fused( + num_experts_per_rank, num_tokens, num_topk, + config.block_m, config.num_epilogue_threads, l2_act_sf_gran_k); + + // Tensormap construction + // Acts/weights: standard 2D TMA descriptors (FP8 K-major). + // Activation SF: per-128 channel float for L1, per-128 or per-64 K for L2 (MN-major, no swizzle). Weight SF: raw float pointer. + constexpr int kGranK = 128; + DG_HOST_ASSERT(l2_act_sf_gran_k == 64 or l2_act_sf_gran_k == 128); + const int kL2ActsSFGranK = l2_act_sf_gran_k; + DG_HOST_ASSERT(static_cast(l2_acts_sf.size(1)) == intermediate_hidden / kL2ActsSFGranK); + DG_HOST_ASSERT(static_cast(l1_acts.size(0)) == num_data_pool_tokens); + DG_HOST_ASSERT(num_data_pool_tokens == config.num_ring_tokens); + const auto tensor_map_l1_acts = make_tma_2d_desc(l1_acts, + hidden, num_data_pool_tokens, + config.block_k, config.block_m, + static_cast(l1_acts.stride(-2)), + config.swizzle_acts_mode); + const auto tensor_map_l1_acts_sf = make_tma_sf_desc(cute::UMMA::Major::MN, l1_acts_sf, + config.num_padded_sf_pool_tokens, hidden, + config.block_m, kGranK, + 1, 0); + const int weight_tma_block_n = config.block_n > 256 ? 256 : config.block_n; + const auto tensor_map_l1_weights = make_tma_2d_desc(l1_weights, + hidden, num_experts_per_rank * intermediate_hidden * 2, + config.block_k, weight_tma_block_n, + static_cast(l1_weights.stride(-2)), + config.swizzle_weights_mode); + // L1 output (post-SwiGLU FP8): N is halved. The correctness path stages + // this tile in plain row-major SMEM before the TMA store. Later L2 TMA + // loads may still swizzle from this row-major global buffer into their own + // SMEM tile. + // The usual TMA store is issued per warpgroup, each writing a `WG_BLOCK_M` + // row tile from its own SMEM offset. The m64n128 2-WG split-N decode path is + // different: both warpgroups stage one joint 64-column L1-output tile and a + // single warpgroup issues the combined store, so the descriptor must cover + // the full block_m x (block_n / 2) tile. + const int num_epilogue_warpgroups_h = config.num_epilogue_threads / 128; + const bool split_n_warpgroups = + config.block_m == 64 and num_epilogue_warpgroups_h > 1 and + config.block_n % num_epilogue_warpgroups_h == 0 and + (config.block_n / num_epilogue_warpgroups_h == 64 or + config.block_n / num_epilogue_warpgroups_h == 128); + const bool split_mn_warpgroups = + config.block_m == 128 and config.block_n == 256 and num_epilogue_warpgroups_h == 4; + const int wg_split_m = split_n_warpgroups ? 1 : + (split_mn_warpgroups ? 2 : num_epilogue_warpgroups_h); + const int wg_split_n = split_n_warpgroups ? num_epilogue_warpgroups_h : + (split_mn_warpgroups ? 2 : 1); + DG_HOST_ASSERT(wg_split_m * wg_split_n == num_epilogue_warpgroups_h); + if (not layout::is_sm90_fused_moe_combine_vectorization_legal( + static_cast(hidden), + static_cast(config.num_epilogue_threads / 32), + get_sm90_fused_pre_barrier_smem_size_for_combine(hidden, config, use_swap_ab), + split_mn_warpgroups)) + DG_HOST_UNREACHABLE( + "SM90 fused FP8 MegaMoE hidden size is incompatible with the selected combine vectorization"); + const int wg_block_m = config.block_m / wg_split_m; + const int wg_block_n = config.block_n / wg_split_n; + const int wg_l1_out_block_n = wg_block_n / 2; + const bool split_n_shares_sf = + split_n_warpgroups and wg_l1_out_block_n < kL2ActsSFGranK; + // The L1 fp8 output tile is staged in the TMA SWIZZLE_128B layout when a warpgroup's staging row is exactly 128 bytes; + // the descriptor must use the layout the kernel stages (the kernel derives the same condition, kL1OutSwizzled). + const int l1_output_swizzle_mode = + (wg_l1_out_block_n == 128 and not split_n_shares_sf and not use_swap_ab) ? 128 : 0; + const int l1_output_box_n = split_n_shares_sf ? config.block_n / 2 : wg_l1_out_block_n; + const int l1_output_box_m = split_n_shares_sf ? config.block_m : wg_block_m; + const auto tensor_map_l1_output = make_tma_2d_desc(l2_acts, + intermediate_hidden, num_data_pool_tokens, + l1_output_box_n, l1_output_box_m, + static_cast(l2_acts.stride(-2)), + l1_output_swizzle_mode); + const auto tensor_map_l2_acts = make_tma_2d_desc(l2_acts, + intermediate_hidden, num_data_pool_tokens, + config.block_k, config.block_m, + static_cast(l2_acts.stride(-2)), + config.swizzle_acts_mode); + const auto tensor_map_l2_acts_sf = make_tma_sf_desc(cute::UMMA::Major::MN, l2_acts_sf, + config.num_padded_sf_pool_tokens, intermediate_hidden, + config.block_m, kL2ActsSFGranK, + 1, 0); + const auto tensor_map_l2_weights = make_tma_2d_desc(l2_weights, + intermediate_hidden, num_experts_per_rank * hidden, + config.block_k, weight_tma_block_n, + static_cast(l2_weights.stride(-2)), + config.swizzle_weights_mode); + + // Stats can be optional + int* cumulative_local_expert_recv_stats_ptr = nullptr; + if (cumulative_local_expert_recv_stats.has_value()) + cumulative_local_expert_recv_stats_ptr = cumulative_local_expert_recv_stats->data_ptr(); + + // opt-in in-kernel event trace: device address of the caller's [num_sms][4][kTraceSlots][2] uint64 buffer; 0 = off (kTrace false) + void* trace_ptr = nullptr; + if (const auto trace_env = get_env("DG_SM90_TRACE_PTR"); not trace_env.empty()) { + const auto trace_addr = std::stoull(trace_env); + trace_ptr = trace_addr != 0 ? reinterpret_cast(static_cast(trace_addr)) : nullptr; + } + // epilogue sub-events in the trace; a role's slots overflow silently above kTraceSlots events + const bool trace_epi = trace_ptr != nullptr and get_env("DG_SM90_TRACE_EPI", 0) != 0; + + // Programmatic dependent launch: a launch that carries the attribute must run the instantiation with the matching griddepcontrol + // wait + trigger (the wait-less cubin could read the pre-dispatch outputs before they are flushed); calls of at most + // kSm90FusedPdlMaxTokens tokens per rank take the attribute on their own. + const int rule_tokens_per_rank = sm90_fused_rule_tokens_per_rank(num_tokens, num_tokens_bound); + const bool pdl = device_runtime->get_pdl() or rule_tokens_per_rank <= kSm90FusedPdlMaxTokens; + + const bool pull_eager_publish = config.pull_publish_batch == 1 and rule_tokens_per_rank <= kSm90FusedPullEagerMaxTokens; + + // Launch + const auto num_sms = device_runtime->get_num_sms(); + const SM90FP8FusedMegaMoERuntime::Args args = { + .num_max_tokens_per_rank = num_max_tokens_per_rank, + .hidden = hidden, .intermediate_hidden = intermediate_hidden, + .num_experts = num_experts, .num_topk = num_topk, + .num_ranks = num_ranks, + .activation_clamp = activation_clamp, + .fast_math = fast_math, + .epilogue_registers = epilogue_registers, + .reuse_accum_as_final = reuse_accum_as_final, + .l2_arrival_counter = l2_arrival_counter, + .l2_epilogue_requires_full_sync = l2_epilogue_requires_full_sync, + .use_swap_ab = use_swap_ab, + .half_l2_cd = config.half_l2_cd, + .num_ring_tokens = num_ring_tokens, + .config = config, + .y = y.data_ptr(), + .cumulative_local_expert_recv_stats = cumulative_local_expert_recv_stats_ptr, + .num_tokens = num_tokens, + .sym_buffer_ptrs = layout::SymBuffer<>(sym_buffer_ptrs, rank_idx), + .trace_ptr = trace_ptr, + .trace_epi = trace_epi, + .l2_act_sf_k = kL2ActsSFGranK, + .pdl = pdl, + .pull_eager_publish = pull_eager_publish, + .tensor_map_l1_acts = tensor_map_l1_acts, + .tensor_map_l1_acts_sf = tensor_map_l1_acts_sf, + .tensor_map_l1_weights = tensor_map_l1_weights, + .l1_weights_sf = l1_weights_sf.data_ptr(), + .tensor_map_l1_output = tensor_map_l1_output, + .tensor_map_l2_acts = tensor_map_l2_acts, + .tensor_map_l2_acts_sf = tensor_map_l2_acts_sf, + .tensor_map_l2_weights = tensor_map_l2_weights, + .l2_weights_sf = l2_weights_sf.data_ptr(), + .launch_args = LaunchArgs(num_sms, config.num_dispatch_threads + config.num_non_epilogue_threads + config.num_epilogue_threads, + config.smem_size, config.cluster_size, /*enable_pdl=*/true, /*force_pdl=*/pdl) + }; + const auto code = SM90FP8FusedMegaMoERuntime::generate(args); + const auto runtime = compiler->build("sm90_fp8_fused_mega_moe", code); + SM90FP8FusedMegaMoERuntime::launch(runtime, args); +} + +} // namespace deep_gemm diff --git a/csrc/python_api.cpp b/csrc/python_api.cpp index 55c0fa2b33..ec9796a886 100644 --- a/csrc/python_api.cpp +++ b/csrc/python_api.cpp @@ -8,6 +8,7 @@ #include "apis/layout.hpp" #include "apis/mega.hpp" #include "apis/sm90_mega.hpp" +#include "apis/sm90_fused_mega.hpp" #include "apis/runtime.hpp" #ifndef TORCH_EXTENSION_NAME @@ -26,5 +27,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { deep_gemm::layout::register_apis(m); deep_gemm::mega::register_apis(m); deep_gemm::mega::register_sm90_apis(m); + deep_gemm::mega::register_sm90_fused_apis(m); deep_gemm::runtime::register_apis(m); } diff --git a/deep_gemm/__init__.py b/deep_gemm/__init__.py index 48819ca4ee..0328dbb05c 100644 --- a/deep_gemm/__init__.py +++ b/deep_gemm/__init__.py @@ -86,6 +86,7 @@ from .mega import ( SymmBuffer, SM90SymmBuffer, + SM90FusedSymmBuffer, get_symm_buffer_for_mega_moe, get_symm_buffer_for_sm90_mega_moe, transform_weights_for_mega_moe, diff --git a/deep_gemm/include/deep_gemm/comm/barrier.cuh b/deep_gemm/include/deep_gemm/comm/barrier.cuh index e2ef54a946..92406cfffa 100644 --- a/deep_gemm/include/deep_gemm/comm/barrier.cuh +++ b/deep_gemm/include/deep_gemm/comm/barrier.cuh @@ -19,29 +19,50 @@ CUTLASS_DEVICE void cluster_sync_with_relaxed_arrive() { cute::cluster_wait(); } +// `WorkspaceT` is any workspace exposing `get_grid_sync_count_ptr<>`, `get_nvl_barrier_counter_ptr()` and +// `get_nvl_barrier_signal_ptr()`. Under `TrapOnly` both waits keep the same timeout and call the same handler, but +// leave the loop first: a `trap;` inside the wait loop makes ptxas allocate the registers of the region containing +// the loop against the kernel's launch bound and ignore the region's `setmaxnreg.inc` (spills in the SM90 fused +// kernel). `Diagnostic` keeps the loop it had, since its handler pulls in a `printf`. template -CUTLASS_DEVICE void grid_sync(const layout::Workspace& workspace, + typename WorkspaceT, typename sync_scope_t> +CUTLASS_DEVICE void grid_sync(const WorkspaceT& workspace, const uint32_t& sm_idx, const uint32_t& thread_idx, const sync_scope_t& sync_scope) { // NOTES: the implementation idea is from `cooperative_groups::this_grid().sync()` static constexpr uint32_t kFinishSumTag = 0x80000000u; sync_scope(); if (thread_idx == 0) { - const auto count_ptr = workspace.get_grid_sync_count_ptr(); + const auto count_ptr = workspace.template get_grid_sync_count_ptr(); const auto old_value = ptx::atomic_add_rel( count_ptr, sm_idx == 0 ? (kFinishSumTag - (kNumSMs - 1)) : 1); uint32_t new_value; - const auto start_clock = clock64(); - do { - new_value = ptx::ld_acq(count_ptr); - if (clock64() - start_clock >= kNumTimeoutCycles) { + if constexpr (kTimeoutPolicy == BarrierTimeoutPolicy::Diagnostic) { + const auto start_clock = clock64(); + do { + new_value = ptx::ld_acq(count_ptr); + if (clock64() - start_clock >= kNumTimeoutCycles) { + handle_grid_sync_timeout( + sm_idx, thread_idx, kGridSyncIndex, old_value, new_value, + old_value ^ kFinishSumTag); + } + } while (((new_value ^ old_value) & kFinishSumTag) == 0); + } else { + const auto start_clock = clock64(); + bool timed_out = false; + do { + new_value = ptx::ld_acq(count_ptr); + if (clock64() - start_clock >= kNumTimeoutCycles) { + timed_out = true; + break; + } + } while (((new_value ^ old_value) & kFinishSumTag) == 0); + if (timed_out) handle_grid_sync_timeout( sm_idx, thread_idx, kGridSyncIndex, old_value, new_value, old_value ^ kFinishSumTag); - } - } while (((new_value ^ old_value) & kFinishSumTag) == 0); + } } sync_scope(); } @@ -49,8 +70,8 @@ CUTLASS_DEVICE void grid_sync(const layout::Workspace& workspace, template -CUTLASS_DEVICE void nvlink_barrier(const layout::Workspace& workspace, + typename WorkspaceT, typename sync_scope_t> +CUTLASS_DEVICE void nvlink_barrier(const WorkspaceT& workspace, const layout::SymBuffer& sym_buffer, const uint32_t& sm_idx, const uint32_t& thread_idx, const sync_scope_t& sync_scope, @@ -80,13 +101,28 @@ CUTLASS_DEVICE void nvlink_barrier(const layout::Workspace& workspace, ptx::red_add(counter_ptr, 1); const int target = signal_sign ? 0 : static_cast(kNumRanks); const auto start_clock = clock64(); - while (ptx::ld_acq_sys(signal_ptr) != target) { - if (clock64() - start_clock >= kNumTimeoutCycles) { + if constexpr (kTimeoutPolicy == BarrierTimeoutPolicy::Diagnostic) { + while (ptx::ld_acq_sys(signal_ptr) != target) { + if (clock64() - start_clock >= kNumTimeoutCycles) { + handle_nvlink_barrier_timeout( + sym_buffer.rank_idx, *counter_ptr, + ptx::ld_acq_sys(signal_ptr), target, + signal_phase, signal_sign, kTag); + } + } + } else { + bool timed_out = false; + while (ptx::ld_acq_sys(signal_ptr) != target) { + if (clock64() - start_clock >= kNumTimeoutCycles) { + timed_out = true; + break; + } + } + if (timed_out) handle_nvlink_barrier_timeout( sym_buffer.rank_idx, *counter_ptr, ptx::ld_acq_sys(signal_ptr), target, signal_phase, signal_sign, kTag); - } } } } diff --git a/deep_gemm/include/deep_gemm/impls/sm90_fp8_fused_mega_moe.cuh b/deep_gemm/include/deep_gemm/impls/sm90_fp8_fused_mega_moe.cuh new file mode 100644 index 0000000000..eeaccc5a18 --- /dev/null +++ b/deep_gemm/include/deep_gemm/impls/sm90_fp8_fused_mega_moe.cuh @@ -0,0 +1,3836 @@ +#pragma once + +#pragma clang diagnostic push +#pragma clang diagnostic ignored "-Wunknown-attributes" + +#include +#include +#include +#include + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace deep_gemm { + +// kFastMath SiLU with `ex2.approx.ftz.f32` instead of `__expf`'s non-ftz form: the two only differ when exp(-x) < 2^-126, +// where `1.0f + e` is 1.0f either way, so the result is bit-identical. Same operation order: e = ex2(x * -log2e), rcp.approx.ftz(1 + e), x * sig. +__forceinline__ __device__ float sm90_fp8_fused_mega_moe_silu_ftz_exp(float x) { + float e; + asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(e) : "f"(__fmul_rn(-x, 1.4426950408889634f))); + return x * math::fast_rcp(1.0f + e); +} + +// Continuous FP32 activation scale. SM90 WGMMA has no hardware block-scale operand (the SF +// is a plain FFMA in the epilogue), so the previous UE8M0 (power-of-two) scale bought nothing +// on SM90 and only cost precision; the SF pool is already fp32, so this is byte/layout neutral. +// clamp amax before the reciprocal: padded rows have amax==0, and 448/0=inf -> 0*inf=NaN. +__forceinline__ __device__ void sm90_fp8_fused_mega_moe_get_e4m3_sf_and_sf_inv( + const float2& amax, float2& sf, float2& sf_inv) { + constexpr float kScale = 1.0f / 448.0f; + const auto ax = fmaxf(amax.x, 1e-10f); + const auto ay = fmaxf(amax.y, 1e-10f); + sf.x = __fmul_rn(ax, kScale), sf_inv.x = 1.0f / sf.x; + sf.y = __fmul_rn(ay, kScale), sf_inv.y = 1.0f / sf.y; +} + +// ============================================================================ +// SM90 (Hopper) FP8 MegaMoE: dispatch warps pull FP8 tokens + per-128 channel float SF from remote ranks over NVLink +// into the local pool; TMA loader warps (A+SFA, B+SFB) feed the stages; math warpgroups run WGMMA and then either the +// L1 epilogue (SwiGLU on the gate/up gran-8 interleaved layout, per-row amax per output-SF group, FP8 e4m3 quantize, +// TMA store; the row SF is written as a float at per-kL2ActSFK granularity, no cross-CTA amax) or the L2 epilogue +// (BF16 cast, NVLink scatter to the remote combine buffers); after all tiles the math warps run the top-k COMBINE. +// ============================================================================ + +template < + uint32_t kNumMaxTokensPerRank, + uint32_t kHidden, uint32_t kIntermediateHidden, + uint32_t kNumExperts, uint32_t kNumTopk, + uint32_t kNumExpertsPerWave, + uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K, + uint32_t kNumMaxPoolTokens, + uint32_t kNumPaddedSFPoolTokens, + uint32_t kNumRingTokens, + uint32_t kNumStages, + uint32_t kNumDispatchThreads, uint32_t kNumNonEpilogueThreads, + uint32_t kNumEpilogueThreads, + uint32_t kNumSMs, uint32_t kNumRanks, + float kActivationClamp, + bool kFastMath, + uint32_t kEpilogueRegisterBudget, + bool kReuseAccumAsFinal, + bool kL2ArrivalCounter, + bool kL2EpilogueRequiresFullSync, + bool kFP8SwapAB = false, + bool kHalfL2CD = false, + // 2-CTA cluster with TMA multicast on A + uint32_t kClusterSize = 1, + bool kMulticastOnB = false, + // L2-lag interleaved schedule (scheduler/sm90_fused_mega_moe.cuh); 0 = wave schedule + uint32_t kL2LagUnits = 0, + // L2 BF16 staging passes per tile (0 = 2 with kHalfL2CD, else 1); 4 = quarter-width buffer + uint32_t kL2CDPasses = 0, + // opt-in in-kernel event trace (host: DG_SM90_TRACE_PTR); false compiles every trace call out + bool kTrace = false, + // epilogue sub-events in the trace (events 50-54 L1, 60-62 L2; host: DG_SM90_TRACE_EPI); needs kTrace + bool kTraceEpi = false, + // K granularity of the L2 activation SF the L1 epilogue writes (the symm buffer's `l2_act_sf_gran_k`): 64 = two wgmma groups + // + two rescales per 128-K block in L2; 128 = the L1-style single-group k-block and half the L2 act-SF pool + uint32_t kL2ActSFK = 64, + // L2 staging passes: 0 = column passes; 2 = 4 passes over ROWS (4 rows x 256 columns) through plain shared stores; 4 = the + // same via stmatrix.m8n8.x4 (junk rows parked in the weight-SF slot). Needs + // kL2CDPasses == 4 and the 64x256 tile + uint32_t kL2StageMode = 0, + // programmatic dependent launch (host: forced for calls of at most kSm90FusedPdlMaxTokens tokens per rank, else `deep_gemm.set_pdl`): + // every thread executes griddepcontrol.wait after the CTA-local prologue and before its first global memory access; the + // trigger is issued per CTA at TILES_DONE + bool kPDL = false, + // early combine: 0 off (host: the smem gate hidden < 512 x topk, or either GEMM having no more k-blocks than num_stages -- so it + // follows this rank's token count, not the shapes alone); 1 = while the CTA waits in the pre-combine all-rank barrier, math warp + // 0 runs the barrier protocol and warps 1.. combine their tokens whose top-k experts are all flagged (fixed slot-order fp32 sum, + // output bit-identical). Contract of the done flag of expert e (set on every rank after the expert's last L2 tile, + // fence.acq_rel.sys + relaxed sys REDs): every combine row written by e is final + uint32_t kEarlyCombineMode = 0, + // pull publish batch N: the arrivals of up to N landed rows are published behind ONE gpu-scope release fence + N relaxed reds; + // 1 = one red.release.gpu per row. A warp's pending rows are in different pool blocks and all are flushed after the + // loop, so the per-block counts are unchanged + uint32_t kPullPublishBatch = 1, + // eager pull publish (per-row publish path only): a warp's first-iteration rows are published right after the local store + bool kPullEagerPublish = false, + uint32_t L1_SHAPE_N = kIntermediateHidden * 2, + uint32_t L1_SHAPE_K = kHidden, + uint32_t L2_SHAPE_N = kHidden, + uint32_t L2_SHAPE_K = kIntermediateHidden, + uint32_t kNumDispatchWarps = kNumDispatchThreads / 32, + uint32_t kNumMMANonEpilogueWarps = kNumNonEpilogueThreads / 32, + uint32_t kNumEpilogueWarps = kNumEpilogueThreads / 32, + uint32_t kNumEpilogueWarpgroups = kNumEpilogueWarps / 4, + uint32_t kNumThreads = kNumDispatchThreads + kNumNonEpilogueThreads + kNumEpilogueThreads, + uint32_t kNumTokensPerWarp = 32 / kNumTopk, + uint32_t kNumExpertsPerRank = kNumExperts / kNumRanks +> +CUTLASS_GLOBAL __launch_bounds__(kNumThreads, 1) void +sm90_fp8_fused_mega_moe_impl(void* y, + int* cumulative_local_expert_recv_stats, + const uint32_t num_tokens, + const __grid_constant__ layout::SymBuffer sym_buffer, + const __grid_constant__ cute::TmaDescriptor tensor_map_l1_acts, + const __grid_constant__ cute::TmaDescriptor tensor_map_l1_acts_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_l1_weights, + const float* __restrict__ l1_weights_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_l1_output, + const __grid_constant__ cute::TmaDescriptor tensor_map_l2_acts, + const __grid_constant__ cute::TmaDescriptor tensor_map_l2_acts_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_l2_weights, + const float* __restrict__ l2_weights_sf, + uint64_t* trace) { +#if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 900) and (__CUDA_ARCH__ < 1000)) or defined(__CLION_IDE__) + using Barrier = cutlass::arch::ClusterTransactionBarrier; + + // ===================================================================== + // Template checks + // ===================================================================== + DG_STATIC_ASSERT(kNumDispatchThreads >= 64 and kNumDispatchThreads % 64 == 0, + "Invalid number of dispatch threads"); + DG_STATIC_ASSERT(kNumNonEpilogueThreads == 64 or kNumNonEpilogueThreads == 128, + "Invalid number of GEMM TMA warps"); + DG_STATIC_ASSERT((kNumDispatchThreads + kNumNonEpilogueThreads) % 128 == 0, + "Math warpgroup start must be 128-thread aligned"); + DG_STATIC_ASSERT(kNumEpilogueThreads % 128 == 0, "Invalid number of math/epilogue threads"); + DG_STATIC_ASSERT(kNumExperts % kNumRanks == 0, "Invalid number of experts or ranks"); + DG_STATIC_ASSERT(BLOCK_M % 64 == 0, "BLOCK_M must be a multiple of WGMMA::M (64)"); + DG_STATIC_ASSERT(BLOCK_N == 128 or BLOCK_N == 256 or BLOCK_N == 512, + "SM90 MegaMoE supports CTA BLOCK_N=128/256/512"); + DG_STATIC_ASSERT(BLOCK_K == 128, "BLOCK_K is fixed to 128 (per-128 SF)"); + DG_STATIC_ASSERT(kClusterSize == 1 or kClusterSize == 2, + "Only 1- or 2-CTA clusters are supported"); + DG_STATIC_ASSERT(not kMulticastOnB or kClusterSize == 2, + "B multicast requires a 2-CTA cluster"); + DG_STATIC_ASSERT(kL2ActSFK == 64 or kL2ActSFK == 128, "L2 activation SF granularity must be 64 or 128 K"); + DG_STATIC_ASSERT(kL2StageMode == 0 or kL2StageMode == 2 or kL2StageMode == 4, + "L2 staging mode: 0 column passes, 2 row passes (STS), 4 row passes (stmatrix)"); + DG_STATIC_ASSERT(kEarlyCombineMode <= 1, "early combine: 0 off, 1 on"); + DG_STATIC_ASSERT(kPullPublishBatch == 1 or not kPullEagerPublish, "the eager publish is defined for the per-row publish path (the batch rule wins)"); + DG_STATIC_ASSERT(kNumRanks <= 32 and kNumRanks <= kNumDispatchThreads, "low-latency head: one count flag per lane / dispatch thread"); + + // ===================================================================== + // Thread / warp identification + // ===================================================================== + const uint32_t sm_idx = blockIdx.x; + const uint32_t thread_idx = threadIdx.x; + const uint32_t warp_idx = cutlass::canonical_warp_idx_sync(); + const uint32_t lane_idx = ptx::get_lane_idx(); + + // Event trace (kTrace only): `trace` is [kNumSMs][4 roles][kTraceSlots] pairs of + // (event_id << 56 | aux, %globaltimer ns), written with plain stores by ONE thread per role: + // role 0 = dispatch warp 0 lane 0, 1 = A/SFA loader lane 0, 2 + g = math warpgroup g warp 0 lane 0. + // Slots past kTraceSlots are dropped, never written. + constexpr uint32_t kTraceSlots = 2048; + uint32_t trace_slot = 0; + const auto trace_event_at = [&](const uint32_t& role, const uint32_t& event_id, const uint64_t& aux, const uint64_t& t) { + if constexpr (kTrace) { + if (trace_slot < kTraceSlots) { + auto ptr = trace + ((static_cast(sm_idx) * 4 + role) * kTraceSlots + trace_slot) * 2; + ptr[0] = (static_cast(event_id) << 56) | (aux & 0x00ffffffffffffffull); + ptr[1] = t; + } + ++ trace_slot; + } + }; + const auto trace_event = [&](const uint32_t& role, const uint32_t& event_id, const uint64_t& aux) { + if constexpr (kTrace) { + uint64_t t; + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t) :: "memory"); + trace_event_at(role, event_id, aux, t); + } + }; + const auto trace_dispatch = [&](const uint32_t& event_id, const uint64_t& aux) { + if constexpr (kTrace) { + if (warp_idx == 0 and lane_idx == 0) + trace_event(0, event_id, aux); + } + }; + // with PDL nothing may be stored to global memory before griddepcontrol.wait (the trace buffer included: the + // harness zeroes it with a memset right before the traced launch), so KERNEL_START only takes its timestamp here and + // is stored after the wait + uint64_t trace_t_kernel_start = 0; + if constexpr (kTrace and kPDL) { + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(trace_t_kernel_start) :: "memory"); + } else { + trace_dispatch(1, 0); // KERNEL_START + } + + // Prefetch all TMA descriptors at the very beginning + if (warp_idx == 0 and cute::elect_one_sync()) { + cute::prefetch_tma_descriptor(&tensor_map_l1_acts); + cute::prefetch_tma_descriptor(&tensor_map_l1_weights); + cute::prefetch_tma_descriptor(&tensor_map_l1_output); + cute::prefetch_tma_descriptor(&tensor_map_l2_acts); + cute::prefetch_tma_descriptor(&tensor_map_l2_weights); + cute::prefetch_tma_descriptor(&tensor_map_l1_acts_sf); + cute::prefetch_tma_descriptor(&tensor_map_l2_acts_sf); + } + + // ===================================================================== + // Workspaces and symmetric buffer slicing + // ===================================================================== + constexpr uint32_t SF_BLOCK_M = math::constexpr_align(BLOCK_M, 128u); + DG_STATIC_ASSERT(kNumMaxPoolTokens % BLOCK_M == 0, "Invalid SM90 MegaMoE pool size"); + // the host sizes the pool with the same helper the workspace uses, so the template value and the workspace agree + DG_STATIC_ASSERT(kNumMaxPoolTokens == layout::get_num_max_pool_tokens_sm90(kNumRanks, kNumMaxTokensPerRank, kNumTopk, kNumExpertsPerRank), + "SM90 MegaMoE pool size does not match the workspace layout"); + + // Fixed-capacity ring pool: data pools are sized by + // `kNumRingTokens` while all metadata (token source table, scheduler + // offsets) keeps full-pool absolute addressing. + constexpr bool kRingCoversFullPool = kNumRingTokens >= kNumMaxPoolTokens; + // early combine (see the template parameter) + constexpr bool kEarlyCombine = kEarlyCombineMode != 0; + // mode 1: the source side publishes the flags and the math warps combine their flagged tokens while they wait in the tag-2 + // all-rank barrier (warp 0 drives the barrier alone; the dispatch warps' pre-cleanup rendezvous is the stop flag), see COMBINE + constexpr bool kEarlyCombinePublish = kEarlyCombineMode == 1; + constexpr bool kEarlyCombineMathWait = kEarlyCombineMode == 1; + constexpr uint32_t kNumDataPoolTokens = kRingCoversFullPool ? kNumMaxPoolTokens : kNumRingTokens; + DG_STATIC_ASSERT(kRingCoversFullPool or kNumRingTokens % BLOCK_M == 0, + "Ring capacity must be BLOCK_M aligned"); + // SF pool is addressed by (ring) pool block only, and the host sizes it by + // the data pool (`get_num_sf_ring_tokens` maximized over candidate BLOCK_M). + constexpr uint32_t kNumDataPoolBlocks = kNumDataPoolTokens / BLOCK_M; + DG_STATIC_ASSERT(kNumPaddedSFPoolTokens >= kNumDataPoolBlocks * SF_BLOCK_M, + "Invalid SM90 MegaMoE SF pool capacity"); + + // The words remote ranks write into this workspace (recv counts, done flags, count flags) are double-buffered by launch + // parity (SM90FusedWorkspace::t2_bank, non-const: selected after the PDL wait; every role's scheduler holds a reference to + // this thread's copy), so the kernel needs no exit all-rank barrier before its cleanup, and the head's count exchange + // completes on per-source release flags instead of an all-rank barrier. + auto workspace = layout::SM90FusedWorkspace( + sym_buffer.get_base_ptr(), kNumRanks, kNumExperts, kNumMaxTokensPerRank, kNumTopk, + kRingCoversFullPool ? 0 : kNumRingTokens); + + // Ring translation layer: the scheduler produces absolute pool indices; ring mode maps them onto a fixed-capacity set of + // physical slots. `kRingCoversFullPool` degrades to the identity mapping. + constexpr uint32_t kNumRingBlocks = kNumRingTokens / BLOCK_M; + const auto get_ring_block_idx = [](const uint32_t& pool_block_idx) { + if constexpr (kRingCoversFullPool) + return pool_block_idx; + else + return pool_block_idx % kNumRingBlocks; + }; + const auto get_ring_token_idx = [](const uint32_t& pool_token_idx) { + if constexpr (kRingCoversFullPool) + return pool_token_idx; + else + return pool_token_idx % kNumRingTokens; + }; + const auto get_ring_wave_idx = [](const uint32_t& pool_block_idx) { + if constexpr (kRingCoversFullPool) + return 0u; + else + return pool_block_idx / kNumRingBlocks; + }; + + constexpr auto fp8_token_layout = layout::Data(kHidden); + constexpr auto bf16_token_layout = layout::Data(kHidden * sizeof(nv_bfloat16)); + constexpr auto fp8_intermediate_token_layout = layout::Data(kIntermediateHidden); + // Per-128 K float SF: 4 bytes per per-128 group => `kHidden / 32` bytes/token + constexpr auto fp8_sf_layout = layout::Data(kHidden / 32); + // L2 acts float SF: 4 bytes per per-kL2ActSFK group. The pool is slot-major (`[k_sf_idx][pool token]`, row pitch = padded + // pool token count), so the per-token byte count only sizes it and need not be a multiple of 16 + constexpr auto fp8_intermediate_sf_layout = layout::Data(kIntermediateHidden * 4 / kL2ActSFK, false); + constexpr auto input_topk_idx_layout = layout::Data(kNumTopk * sizeof(int64_t), false); + constexpr auto input_topk_weights_layout = layout::Data(kNumTopk * sizeof(float), false); + constexpr auto l1_topk_weights_layout = layout::Data(sizeof(float), false); + + // Registered input area + const auto input_token_buffer = layout::Buffer(fp8_token_layout, 1, kNumMaxTokensPerRank, workspace.get_end_ptr()); + const auto input_sf_buffer = layout::Buffer(fp8_sf_layout, 1, kNumMaxTokensPerRank, input_token_buffer.get_end_ptr()); + const auto input_topk_idx_buffer = layout::Buffer(input_topk_idx_layout, 1, kNumMaxTokensPerRank, input_sf_buffer.get_end_ptr()); + const auto input_topk_weights_buffer = layout::Buffer(input_topk_weights_layout, 1, kNumMaxTokensPerRank, input_topk_idx_buffer.get_end_ptr()); + + // L1 input area (ring-sized data pools in ring mode) + const auto l1_token_buffer = layout::Buffer(fp8_token_layout, 1, kNumDataPoolTokens, input_topk_weights_buffer.get_end_ptr()); + const auto l1_sf_buffer = layout::Buffer(fp8_sf_layout, 1, kNumPaddedSFPoolTokens, l1_token_buffer.get_end_ptr()); + const auto l1_topk_weights_buffer = layout::Buffer(l1_topk_weights_layout, 1, kNumDataPoolTokens, l1_sf_buffer.get_end_ptr()); + + // L2 input area + const auto l2_token_buffer = layout::Buffer(fp8_intermediate_token_layout, 1, kNumDataPoolTokens, l1_topk_weights_buffer.get_end_ptr()); + const auto l2_sf_buffer = layout::Buffer(fp8_intermediate_sf_layout, 1, kNumPaddedSFPoolTokens, l2_token_buffer.get_end_ptr()); + + // Combine input area + const auto combine_token_buffer = layout::Buffer(bf16_token_layout, kNumTopk, kNumMaxTokensPerRank, l2_sf_buffer.get_end_ptr()); + + // ===================================================================== + // GEMM data types and shape constants + // ===================================================================== + using a_dtype_t = cutlass::float_e4m3_t; + using b_dtype_t = cutlass::float_e4m3_t; + // under ping-pong a warpgroup owns whole tiles, so the N split (and everything it implied: shared SF group, joint staging + // tile, cross-warpgroup amax) is off + constexpr bool kSplitNWarpgroups = + BLOCK_M == 64 and kNumEpilogueWarpgroups > 1 and + BLOCK_N % kNumEpilogueWarpgroups == 0 and + ((BLOCK_N / kNumEpilogueWarpgroups == 64) or (BLOCK_N / kNumEpilogueWarpgroups == 128)); + constexpr bool kSplitMNWarpgroups = + BLOCK_M == 128 and BLOCK_N == 256 and kNumEpilogueWarpgroups == 4; + // the decode topology: BLOCK_M-64 tiles split along N over two math warpgroups (the host picks it below ~56 expected + // tokens per expert); the levers below that only pay there key on it + constexpr bool kDecodeTopology = BLOCK_M == 64 and kNumEpilogueWarpgroups == 2; + // deferred ring fill (decode topology under the early combine): the B loader issues stage 0 of its first tile at once and the + // rest of the ring fill only after both dispatch warps published their first pulled row (or found none). Scheduling hint only + constexpr bool kBFillDefer = kDecodeTopology and kEarlyCombineMode == 1; + constexpr uint32_t kWarpgroupSplitM = kSplitNWarpgroups ? 1 : + (kSplitMNWarpgroups ? 2 : kNumEpilogueWarpgroups); + constexpr uint32_t kWarpgroupSplitN = kSplitNWarpgroups ? kNumEpilogueWarpgroups : + (kSplitMNWarpgroups ? 2 : 1); + constexpr uint32_t WG_BLOCK_M = BLOCK_M / kWarpgroupSplitM; + constexpr uint32_t WG_BLOCK_N = BLOCK_N / kWarpgroupSplitN; + constexpr uint32_t kNumCombineWarps = kNumEpilogueWarps; + using L1WGMMA = typename mma::sm90::FP8MMASelector::type; // M=64, N=WG_BLOCK_N, K=32 + using L2WGMMA = typename mma::sm90::FP8MMASelector::type; + constexpr uint32_t kL1OutputArrivalParts = 1; + static_assert(L1WGMMA::M == 64 and L1WGMMA::N == WG_BLOCK_N and L1WGMMA::K == 32, + "Unexpected WGMMA shape"); + DG_STATIC_ASSERT(kWarpgroupSplitM * kWarpgroupSplitN == kNumEpilogueWarpgroups, + "Invalid warpgroup split"); + DG_STATIC_ASSERT(WG_BLOCK_M == L1WGMMA::M, + "Each warpgroup must run exactly one WGMMA-M tile"); + DG_STATIC_ASSERT(kNumCombineWarps <= kNumEpilogueWarps, + "Combine warp count must fit in epilogue warps"); + + // Cluster=1 -> no multicast, A/B are loaded full-sized + constexpr uint32_t LOAD_BLOCK_M = BLOCK_M; + constexpr uint32_t LOAD_BLOCK_N = BLOCK_N; + constexpr uint32_t L1_OUT_BLOCK_N = BLOCK_N / 2; // post-SwiGLU + constexpr uint32_t WG_L1_OUT_BLOCK_N = WG_BLOCK_N / 2; + // 128-column weight-SF blocks covered by one warpgroup tile (L2: lo/hi weight SF per k-block; with per-64 L2 act SF + // this is also the number of per-64 post-SwiGLU SF groups of the L1 epilogue) + constexpr uint32_t kNumSFGroupsPerWG = WG_BLOCK_N >= 128 ? WG_BLOCK_N / 128 : 1; + DG_STATIC_ASSERT(kNumSFGroupsPerWG <= 2, "At most two weight-SF blocks per warpgroup tile"); + constexpr uint32_t kGranK = 128; // L1 acts SF, weights SF + constexpr uint32_t kL2ActsSFGranK = kL2ActSFK; // L2 acts SF (per-64 or per-128 K) + // one L2 act SF per 128 K -> the L2 mainloop consumes one act SF per k-block like L1 (single wgmma group) + constexpr bool kL2ActSFPerBlockK = kL2ActsSFGranK == BLOCK_K; + // When WG_L1_OUT_BLOCK_N < kL2ActsSFGranK the two N-split warpgroups jointly own one L2-acts SF group: they publish ONE + // shared SF slot (k_sf_idx == n_block_idx) and the amax feeding it is reduced across both warpgroups. + constexpr bool kSplitNSharesSF = kSplitNWarpgroups and (WG_L1_OUT_BLOCK_N < kL2ActsSFGranK); + // Both N-split warpgroups sit inside one 128-column weight-SF block (they stage the same weight-SF row) + constexpr bool kSplitNSharesWeightSF = kSplitNWarpgroups and (WG_BLOCK_N < 128); + DG_STATIC_ASSERT(kL2ActsSFGranK != 64 or kSplitNSharesWeightSF == kSplitNSharesSF, + "per-64 L2 act SF: the shared act-SF split is exactly the shared weight-SF split"); + // L2 act-SF groups (post-SwiGLU columns / kL2ActsSFGranK) produced by one warpgroup's L1 epilogue: 2 per-64 groups + // for the 128-column production tile, 1 per-128 group; the shared-SF split publishes one joint group + constexpr uint32_t kNumL1OutSFGroups = kSplitNSharesSF ? 1u : + (WG_L1_OUT_BLOCK_N >= kL2ActsSFGranK ? WG_L1_OUT_BLOCK_N / kL2ActsSFGranK : 1u); + DG_STATIC_ASSERT(kL2ActsSFGranK != 64 or kNumL1OutSFGroups == kNumSFGroupsPerWG, + "per-64 L2 act SF: one post-SwiGLU SF group per 128-column weight-SF block"); + DG_STATIC_ASSERT(kSplitNSharesSF or (WG_L1_OUT_BLOCK_N % kL2ActsSFGranK == 0), + "A warpgroup's L1 output must cover whole L2 act-SF groups"); + // the shared-SF split publishes ONE SF slot per tile (slot n_block): the joint post-SwiGLU output of the two warpgroups must + // be exactly one L2 act-SF group (the host picks the tile width so) + DG_STATIC_ASSERT(not kSplitNSharesSF or L1_OUT_BLOCK_N == kL2ActsSFGranK, + "shared-SF split: the tile's post-SwiGLU output must be exactly one L2 act-SF group"); + constexpr bool kSwapABEligible = + kFP8SwapAB and kSplitNWarpgroups and (BLOCK_M == 64) and (BLOCK_N == 128) and + (kWarpgroupSplitN == 2); + constexpr bool kSwapABActive = kSwapABEligible; + // The L1 fp8 output tile is staged in the TMA SWIZZLE_128B layout when a warpgroup's staging row is exactly 128 bytes + // (the 2-warpgroup split-M 128x256 tile; neither shared-SF split nor swapAB); the host builds the descriptor with the same rule + constexpr bool kL1OutSwizzled = WG_L1_OUT_BLOCK_N == 128 and not kSplitNSharesSF and not kSwapABActive; + constexpr uint32_t kSwapABTokenChunks = BLOCK_M / 8; + // Ring mode reuses the mask-mode arrival flow (one CTA-wide barrier per + // (m, n) L1 block, then a single elected publish), because per-WG counter + // arrivals are `valid_m`-dependent and cannot carry ring generations. + constexpr bool kL2ArrivalNeedsFullSync = (not kL2ArrivalCounter) or (not kRingCoversFullPool); + DG_STATIC_ASSERT(not kSwapABEligible or (BLOCK_M % 8 == 0), + "swapAB epilogue token chunks assume BLOCK_M is a multiple of 8"); + constexpr uint32_t kSwizzleAMode = BLOCK_K * sizeof(a_dtype_t); // 128 + constexpr uint32_t kSwizzleBMode = BLOCK_K * sizeof(b_dtype_t); // 128 + constexpr uint32_t kSwizzleCDMode = 128; + DG_STATIC_ASSERT(not kSwapABActive or kL2ActsSFGranK == 64, "the swapAB epilogues assume per-64 L2 activation scales"); + // The swapAB L1 epilogue lets the first N-split warpgroup quantize and store the other warpgroup's half too; that + // cross-warpgroup hand-off is not bitwise stable run to run, so a warpgroup must own a whole activation-SF group of the + // tile. The non-swap epilogue combines the warpgroups only through a commutative amax reduction (host: should_use_swap_ab_sm90_fused). + DG_STATIC_ASSERT(not kSwapABActive or not kSplitNSharesSF, + "swapAB needs a tile whose post-SwiGLU columns one warpgroup owns end to end"); + + // ===================================================================== + // Shared memory layout + // ===================================================================== + constexpr uint32_t kSharedMemoryAlignment = 1024; + extern __shared__ __align__(kSharedMemoryAlignment) uint8_t smem_buffer[]; + + // The per-expert routing counts (kNumExperts x u32) are dead after the routing grid sync; they live at the head of the + // tile-table region at the END of the layout (see SMEM_TILE_TABLE_OFFSET / smem_expert_count) + constexpr uint32_t SMEM_EXPERT_COUNT_BYTES = kNumExperts * sizeof(uint32_t); + constexpr uint32_t SMEM_SEND_BUFFER_SIZE = + math::constexpr_align(fp8_token_layout.get_num_bytes() * kNumDispatchWarps, kSharedMemoryAlignment); + constexpr uint32_t SMEM_A_SIZE_PER_STAGE = LOAD_BLOCK_M * BLOCK_K * sizeof(a_dtype_t); + constexpr uint32_t SMEM_B_SIZE_PER_STAGE = LOAD_BLOCK_N * BLOCK_K * sizeof(b_dtype_t); + // SFA per stage: BLOCK_M floats for L1 and for the per-128 L2 act SF (one (BLOCK_M, 1) TMA per k-block), 2*BLOCK_M + // floats (lo/hi halves) with the per-64 L2 act SF + constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = + math::constexpr_align((kL2ActSFPerBlockK ? 1u : 2u) * BLOCK_M * sizeof(float), 128u); + constexpr uint32_t SMEM_SFB_SIZE_PER_STAGE = 0; + constexpr uint32_t kNumL1WeightSFFloatsPerWG = 2 * (kHidden / 128); + constexpr uint32_t kNumL2WeightSFFloatsPerWG = kNumSFGroupsPerWG * (kIntermediateHidden / 128); + constexpr uint32_t kNumWeightSFFloatsPerWG = + kNumL1WeightSFFloatsPerWG > kNumL2WeightSFFloatsPerWG ? + kNumL1WeightSFFloatsPerWG : kNumL2WeightSFFloatsPerWG; + constexpr uint32_t SMEM_WEIGHT_SF_SIZE = + math::constexpr_align( + kNumEpilogueWarpgroups * kNumWeightSFFloatsPerWG * sizeof(float), 128u); + + // CD output: max of L1 FP8 (BLOCK_M * BLOCK_N/2), L2 BF16 (BLOCK_M * BLOCK_N * 2) and the swapAB L1 FP32+FP8 staging; + // split-M warpgroups own disjoint row slices, shared-SF split-N warpgroups disjoint column slices of one CTA tile. + constexpr uint32_t SMEM_CD_L1_SIZE = BLOCK_M * L1_OUT_BLOCK_N * sizeof(cutlass::float_e4m3_t); + // kHalfL2CD stages the L2 BF16 output one N-half at a time (two-pass scatter), halving this buffer + constexpr uint32_t kNumL2CDPasses = kL2CDPasses != 0 ? kL2CDPasses : (kHalfL2CD ? 2u : 1u); + constexpr uint32_t L2_CD_STAGE_N = BLOCK_N / kNumL2CDPasses; + constexpr uint32_t SMEM_CD_L2_SIZE = BLOCK_M * L2_CD_STAGE_N * sizeof(nv_bfloat16); + // row-pass staging: 2 KiB per math warp == the 4-column-pass buffer (BLOCK_M x BLOCK_N/4 bf16) + DG_STATIC_ASSERT(kL2StageMode == 0 or (kNumL2CDPasses == 4 and not kFP8SwapAB), + "L2 row-pass staging needs the quarter-width (4-pass) buffer"); + constexpr uint32_t SMEM_CD_SWAP_L1_FP32_SIZE = + kSwapABActive ? BLOCK_M * L1_OUT_BLOCK_N * sizeof(float) : 0; + constexpr uint32_t SMEM_CD_SWAP_L1_FP8_SIZE = + kSwapABActive ? BLOCK_M * L1_OUT_BLOCK_N * sizeof(cutlass::float_e4m3_t) : 0; + constexpr uint32_t SMEM_CD_SWAP_L1_SIZE = + kSwapABActive ? (SMEM_CD_SWAP_L1_FP32_SIZE + SMEM_CD_SWAP_L1_FP8_SIZE) : 0; + constexpr uint32_t SMEM_CD_BASE_SIZE = + SMEM_CD_L1_SIZE > SMEM_CD_L2_SIZE ? SMEM_CD_L1_SIZE : SMEM_CD_L2_SIZE; + constexpr uint32_t SMEM_CD_SIZE = math::constexpr_align( + SMEM_CD_BASE_SIZE > SMEM_CD_SWAP_L1_SIZE ? SMEM_CD_BASE_SIZE : SMEM_CD_SWAP_L1_SIZE, + kSharedMemoryAlignment); + + constexpr uint32_t SMEM_BEFORE_BARRIER_SIZE = + SMEM_SEND_BUFFER_SIZE + SMEM_CD_SIZE + + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE); + + // SMEM pointers (the layout starts with the dispatch send buffers; the routing counts are at the end, see below) + const auto smem_send_buffers = layout::Buffer(fp8_token_layout, kNumDispatchWarps, 1, smem_buffer); + + auto smem_gemm_base = math::advance_ptr(smem_buffer, SMEM_SEND_BUFFER_SIZE); + + // CD output is shared by L1 (FP8) and L2 (BF16); reinterpret-cast as needed. + auto smem_cd_l1 = reinterpret_cast(smem_gemm_base); + auto smem_cd_l2 = reinterpret_cast(smem_gemm_base); + auto smem_cd_swap_l1_fp32 = reinterpret_cast(smem_gemm_base); + auto smem_cd_swap_l1_fp8 = reinterpret_cast( + math::advance_ptr(smem_gemm_base, SMEM_CD_SWAP_L1_FP32_SIZE)); + + auto smem_a = utils::PatternVisitor([=](const uint32_t& i) { + return math::advance_ptr(smem_gemm_base, SMEM_CD_SIZE + i * SMEM_A_SIZE_PER_STAGE); + }); + auto smem_b = utils::PatternVisitor([=](const uint32_t& i) { + return math::advance_ptr(smem_gemm_base, SMEM_CD_SIZE + kNumStages * SMEM_A_SIZE_PER_STAGE + i * SMEM_B_SIZE_PER_STAGE); + }); + auto sf_start_ptr = math::advance_ptr(smem_gemm_base, + SMEM_CD_SIZE + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE)); + auto smem_sfa = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast(sf_start_ptr + i * SMEM_SFA_SIZE_PER_STAGE); + }); + + // Per-warpgroup weight-SF staging slots (see SMEM_WEIGHT_SF_SIZE above) + auto smem_weight_sf = reinterpret_cast( + sf_start_ptr + kNumStages * SMEM_SFA_SIZE_PER_STAGE); + + // Barriers live after the weight-SF staging area + // A (+ per-group SFA) is multicast only when the pairing is on m + // (n-inner order); with B multicast the pair's A tiles differ. + constexpr uint32_t kNumMulticastA = kMulticastOnB ? 1u : kClusterSize; + + auto barrier_start_ptr = reinterpret_cast( + sf_start_ptr + kNumStages * SMEM_SFA_SIZE_PER_STAGE + SMEM_WEIGHT_SF_SIZE); + auto dispatch_barriers = utils::PatternVisitor([=](const uint32_t& i) { return barrier_start_ptr + i; }); + auto full_barriers = utils::PatternVisitor([=](const uint32_t& i) { return barrier_start_ptr + kNumDispatchWarps + i; }); + auto empty_barriers = utils::PatternVisitor([=](const uint32_t& i) { return barrier_start_ptr + kNumDispatchWarps + kNumStages + i; }); + auto combine_barriers = utils::PatternVisitor([=](const uint32_t& i) { return barrier_start_ptr + kNumDispatchWarps + kNumStages * 2 + i; }); + + // Tile table right after the barriers (host: `smem_tile_table` mirrors this). Its head holds the routing counts + // (`smem_expert_count`, kNumExperts x u32), dead after the dispatch warps' second `read_topk_idx`; the table is initialised by + // the dispatch warps after that (past a bar.sync, no routing atomic pending) and handed to the loaders with + // kDispatchWithLoadersBarrierIdx, to the math warpgroups with kDispatchWithEpilogueBarrierIdx (both after the routing grid + // sync); the cluster partner's DSMEM tail entries are written after the leader's `fetch_expert_recv_count`, i.e. past that sync. + constexpr uint32_t kNumBarriers = kNumDispatchWarps + kNumStages * 2 + kNumCombineWarps * 2; + // Early-combine region AFTER the tile table (kEarlyCombine only; host `smem_early_combine`): one more transaction barrier per + // dispatch warp, then a 64 B control block: word 0 = L2 tiles whose scatter stores are issued (the B loader adds 1 for tile T-1 + // once tile T's k-block kNumStages passed the empty-barrier wait), word 1 = stop flag, word 2 = done flag (every tile finished), + // words 4/5 = trace counters of the dispatch warps (0), words 8.. = the bitmap of tokens the math warps combined while + // waiting in the barrier, bit (stripe * kNumCombineWarps + math warp) + constexpr uint32_t SMEM_EC_BYTES = kEarlyCombine ? (kNumDispatchWarps * static_cast(sizeof(Barrier)) + 64u) : 0u; + constexpr uint32_t kNumECTokenStripes = math::constexpr_ceil_div(kNumMaxTokensPerRank, kNumSMs * kNumCombineWarps); + constexpr uint32_t kNumECClaimWords = math::constexpr_ceil_div(kNumECTokenStripes * kNumCombineWarps, 32u); + DG_STATIC_ASSERT(not kEarlyCombine or 32u + kNumECClaimWords * 4u <= 64u, "early combine: too many tokens per CTA for the combined-token bitmap"); + DG_STATIC_ASSERT(not kEarlyCombine or kNumECTokenStripes <= 32, "early combine: the math warps' wait-combine keeps one bit per token stripe"); + // (8-byte entries, one more spare entry for the bounded dynamic tail, see the scheduler) + constexpr uint32_t kNumTileTableEntries = layout::get_sm90_tile_table_entries_compact( + kNumMaxPoolTokens, BLOCK_M, L1_SHAPE_N / BLOCK_N, L2_SHAPE_N / BLOCK_N, kNumSMs); + constexpr uint32_t SMEM_TILE_TABLE_OFFSET = math::constexpr_align( + SMEM_SEND_BUFFER_SIZE + SMEM_CD_SIZE + + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE) + + SMEM_WEIGHT_SF_SIZE + kNumBarriers * static_cast(sizeof(Barrier)), 16u); + using TileEntry = uint2; + constexpr uint32_t SMEM_TILE_TABLE_BYTES = kNumTileTableEntries * static_cast(sizeof(TileEntry)); + constexpr uint32_t SMEM_TILE_TABLE_REGION_SIZE = + SMEM_TILE_TABLE_BYTES > SMEM_EXPERT_COUNT_BYTES ? SMEM_TILE_TABLE_BYTES : SMEM_EXPERT_COUNT_BYTES; + // Early-combine region (see above): 8-byte aligned right after the tile-table region (whose size is a multiple of 8) + constexpr uint32_t SMEM_EC_OFFSET = math::constexpr_align(SMEM_TILE_TABLE_OFFSET + SMEM_TILE_TABLE_REGION_SIZE, 8u); + // SM90 dynamic shared memory capacity (227 KiB); the host derives kNumStages from the same accounting + DG_STATIC_ASSERT(SMEM_EC_OFFSET + SMEM_EC_BYTES <= 232448u, "SM90 MegaMoE smem layout overflows"); + auto tile_table = reinterpret_cast(smem_buffer + SMEM_TILE_TABLE_OFFSET); + auto smem_expert_count = reinterpret_cast(smem_buffer + SMEM_TILE_TABLE_OFFSET); + auto ec_barriers = utils::PatternVisitor([=](const uint32_t& i) { return reinterpret_cast(smem_buffer + SMEM_EC_OFFSET) + i; }); + auto smem_ec = reinterpret_cast(smem_buffer + SMEM_EC_OFFSET + kNumDispatchWarps * static_cast(sizeof(Barrier))); + auto smem_ec_tiles = smem_ec, smem_ec_stop = smem_ec + 1, smem_ec_done = smem_ec + 2, + smem_ec_count = smem_ec + 4, smem_ec_gs_old = smem_ec + 6; + // CTA-wide token bitmap, bit (stripe * kNumCombineWarps + warp), words 8.. of the early-combine control block; the static + // sliced combine records the tokens the wait-combine already finished in it + auto smem_ec_claim = smem_ec + 8; + + // Pull pending list (kPullPublishBatch > 1): per dispatch warp, the landed-but-unpublished pull rows (u32: pool token | + // is_last_of_expert << 31), 16-byte aligned after the early-combine region; 0 bytes with per-row publish (host: `smem_pull_pending`) + DG_STATIC_ASSERT(kPullPublishBatch >= 1 and kPullPublishBatch <= 256, "Invalid pull publish batch"); + constexpr uint32_t SMEM_PULL_PENDING_OFFSET = math::constexpr_align(SMEM_EC_OFFSET + SMEM_EC_BYTES, 16u); + constexpr uint32_t SMEM_PULL_PENDING_SIZE = kPullPublishBatch > 1 ? + math::constexpr_align(kNumDispatchWarps * kPullPublishBatch * sizeof(uint32_t), 16u) : 0u; + DG_STATIC_ASSERT(SMEM_PULL_PENDING_OFFSET + SMEM_PULL_PENDING_SIZE <= 232448u, "SM90 MegaMoE smem layout overflows (pull pending list)"); + auto smem_pull_pending = reinterpret_cast(smem_buffer + SMEM_PULL_PENDING_OFFSET); + + // Release a pipeline stage toward every CTA whose producer writes into it. + auto release_empty_stage = [&](const uint32_t& stage_idx) { + if constexpr (kClusterSize == 1) { + if (lane_idx == 0) + empty_barriers[stage_idx]->arrive(); + } else { + if (lane_idx < kClusterSize) + empty_barriers[stage_idx]->arrive(lane_idx); + } + }; + + // ===================================================================== + // Initialization + // ===================================================================== + if (warp_idx == 0) { + // Clean expert-count shared memory (it is the head of the tile-table region; the table itself is + // initialised by the dispatch warps once routing is done, see kDispatchWithLoadersBarrierIdx) + #pragma unroll + for (uint32_t i = lane_idx; i < kNumExperts; i += 32) + ptx::st_shared(smem_expert_count + i, 0u); + if constexpr (kEarlyCombine) { + if (lane_idx < 16) + ptx::st_shared(smem_ec + lane_idx, 0u); + } + } else if (warp_idx == 1) { + // Init dispatch m-barriers + #pragma unroll + for (uint32_t i = lane_idx; i < kNumDispatchWarps; i += 32) + dispatch_barriers[i]->init(1); + if constexpr (kEarlyCombine) { + #pragma unroll + for (uint32_t i = lane_idx; i < kNumDispatchWarps; i += 32) + ec_barriers[i]->init(1); + } + cutlass::arch::fence_barrier_init(); + } else if (warp_idx == 2) { + // Init GEMM full/empty barriers and combine barriers + if (cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kNumStages; ++ i) { + // Two producer warps (A+SFA loader, B+SFB loader) each call + // `arrive_and_expect_tx` per stage, so init count must be 2. + full_barriers[i]->init(2); + // Each math warp arrives once per stage release + empty_barriers[i]->init(kNumEpilogueWarps * kClusterSize); + } + #pragma unroll + for (uint32_t i = 0; i < kNumCombineWarps * 2; ++ i) + combine_barriers[i]->init(1); + } + cutlass::arch::fence_barrier_init(); + } + if constexpr (kClusterSize > 1) + comm::cluster_sync_with_relaxed_arrive(); + else + __syncthreads(); + + // PDL: everything above is CTA-local; everything below reads the previous grid's outputs (topk buffers, the workspace + // counters and recv counts zeroed by its cleanup), so every thread waits here for the prerequisite grids. + if constexpr (kPDL) { + cudaGridDependencySynchronize(); + if constexpr (kTrace) { + if (warp_idx == 0 and lane_idx == 0) + trace_event_at(0, 1, 0, trace_t_kernel_start); // KERNEL_START (timestamp taken before the prologue) + } + } + // launch parity -> bank of the remotely written words (see SM90FusedWorkspace::t2_bank). Every thread reads the word here, + // before any role touches the workspace; SM0's flip in the tail is ordered behind every read of this launch (loaders and + // math warps sync with their dispatch warps before those pass the routing / tag-1 grid syncs that precede SM0's tail) + workspace.t2_bank = ptx::ld_relaxed_sys(workspace.get_t2_parity_ptr()) & 1u; + // trace: 23 PROLOGUE_DONE = the CTA-local prologue is over (with PDL: the wait has returned, so KERNEL_START -> 23 also + // holds the time this CTA sat waiting for the previous grid) + trace_dispatch(23, 0); + + // ===================================================================== + // Scheduler (cluster=1 / 2) + // ===================================================================== + constexpr uint32_t kNumExpertsPerLane = math::constexpr_ceil_div(kNumExpertsPerRank, 32u); + constexpr uint32_t kNumL1BlockNs = L1_SHAPE_N / BLOCK_N; + constexpr uint32_t kNumL2BlockNs = L2_SHAPE_N / BLOCK_N; + constexpr uint32_t kNumL1BlockKs = L1_SHAPE_K / BLOCK_K; + constexpr uint32_t kNumL2BlockKs = L2_SHAPE_K / BLOCK_K; + auto scheduler = sched::SM90FusedMegaMoEScheduler< + BLOCK_M, BLOCK_N, BLOCK_K, + L1_SHAPE_N, L1_SHAPE_K, + L2_SHAPE_N, L2_SHAPE_K, + kNumExpertsPerRank, kNumExpertsPerWave, + kNumSMs, kNumRanks, + kMulticastOnB, + kNumExpertsPerLane, kNumL1BlockNs, kNumL2BlockNs, + kNumL1BlockKs, kNumL2BlockKs, + layout::SM90FusedWorkspace, kL2LagUnits, + kNumMaxPoolTokens / BLOCK_M, kNumRanks * kNumMaxTokensPerRank, kClusterSize>(workspace); + + // Pipeline state shared by TMA loaders and math warpgroups + uint32_t stage_idx = 0, phase = 0; + auto advance_pipeline = [&](uint32_t& k_block_idx) { + ++ k_block_idx; + stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + }; + + // Intra-SM barrier indices + constexpr uint32_t kDispatchBarrierIdx = 0; + constexpr uint32_t kDispatchWithEpilogueBarrierIdx = 1; + constexpr uint32_t kEpilogueFullBarrierIdx = 2; + constexpr uint32_t kEpilogueWGBarrierStartIdx = 3; + // dispatch warps + the A and B loader warps, once per launch, after the dispatch warps initialised the tile + // table in the region the routing counts occupied + constexpr uint32_t kDispatchWithLoadersBarrierIdx = kEpilogueWGBarrierStartIdx + kNumEpilogueWarpgroups; + constexpr uint32_t kNumDispatchWithLoadersThreads = kNumDispatchThreads + 2 * 32; + DG_STATIC_ASSERT(kDispatchWithLoadersBarrierIdx < 16, "Out of named barriers"); + + // Cross-rank NVLink barrier tag (the head and the tail need no all-rank barrier: the head completes on per-source + // release flags, the tail on the launch-parity banks) + constexpr uint32_t kBeforeCombineReduceBarrierTag = 2; + + // Register reconfiguration counts (64512 budget). 256-epilogue-thread split-N decode: 64*48 + 64*40 + 256*168 = 48640; + // kNumThreads <= 256 decode: 64*48 + 64*40 + 128*256 = 38400 (launch-bounds ceiling 256, the accumulator double-buffer fits); + // 2 math warpgroups on the 128x256 split-M tile: 64*104 + 64*104 + 256*200 = 64512 (setmaxnreg is warpgroup-collective, + // so the dispatch warps and the TMA loader warps share one count); the 512-epilogue-thread split-MN path trims both roles. + constexpr bool kSplitMWide = kNumEpilogueThreads == 256 and BLOCK_M == 128 and BLOCK_N == 256 and + kNumDispatchThreads == 64 and kNumNonEpilogueThreads == 64; + constexpr uint32_t kNumEpilogueRegisters = + kEpilogueRegisterBudget == 0 ? + (kNumEpilogueThreads == 512 ? 112 : + (kSplitMWide ? 200 : + (kNumEpilogueThreads == 256 ? 168 : + (kNumThreads <= 256u ? 256 : 208)))) : + kEpilogueRegisterBudget; + constexpr uint32_t kNumDispatchRegisters = + kNumEpilogueThreads == 512 ? 32 : (kSplitMWide ? 104 : 48); + constexpr uint32_t kNumNonEpilogueRegisters = + kNumEpilogueThreads == 512 ? 24 : (kSplitMWide ? 104 : 40); + DG_STATIC_ASSERT(kNumDispatchRegisters * kNumDispatchThreads + + kNumNonEpilogueRegisters * kNumNonEpilogueThreads + + kNumEpilogueRegisters * kNumEpilogueThreads <= 64512, + "Too many registers"); + DG_STATIC_ASSERT(kNumEpilogueRegisters % 8 == 0 and kNumEpilogueRegisters >= 24 and kNumEpilogueRegisters <= 256 and + kNumDispatchRegisters % 8 == 0 and kNumDispatchRegisters >= 24 and kNumDispatchRegisters <= 256 and + kNumNonEpilogueRegisters % 8 == 0 and kNumNonEpilogueRegisters >= 24 and kNumNonEpilogueRegisters <= 256, + "setmaxnreg counts must be multiples of 8 in [24, 256]"); + + constexpr uint32_t kDispatchGridSyncIndex = 0; + constexpr uint32_t kEpilogueGridSyncIndex = 1; + // dynamic tail: pair ticket counter (a spare grid-sync slot; zeroed by the dispatch cleanup every launch) + constexpr uint32_t kTailCounterGridSyncIndex = 2; + + // ===================================================================== + // ROLE 1: DISPATCH WARPS + // SF is per-128 channel float, stored straight into the local L1 SF buffer in MN-major layout + // `local_sf[k_chunk * num_padded_sf_pool_tokens + token_idx]`; token_idx_in_expert -> SF token index is the per-block linear mapping. + // ===================================================================== + if (warp_idx < kNumDispatchWarps) { + cutlass::arch::warpgroup_reg_dealloc(); + + DG_STATIC_ASSERT(kNumTopk <= 32, "Invalid number of topk"); + constexpr uint32_t kNumActivateLanes = kNumTokensPerWarp * kNumTopk; + const auto read_topk_idx = [&](const auto& process) { + #pragma unroll + for (uint32_t i = (sm_idx * kNumDispatchWarps + warp_idx) * kNumTokensPerWarp; + i < num_tokens; + i += kNumSMs * kNumDispatchWarps * kNumTokensPerWarp) { + int expert_idx = -1; + if (i + (lane_idx / kNumTopk) < num_tokens and lane_idx < kNumActivateLanes) { + expert_idx = static_cast( + __ldg(input_topk_idx_buffer.get_base_ptr() + i * kNumTopk + lane_idx)); + if (expert_idx >= 0) + process(i * kNumTopk + lane_idx, expert_idx); + } + __syncwarp(); + } + }; + + // Count tokens per expert + read_topk_idx([&](const uint32_t& token_topk_idx, const int& expert_idx) { + atomicAdd_block(smem_expert_count + expert_idx, 1); + }); + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + + // Stake out per-expert SM offsets via global atomic + #pragma unroll + for (uint32_t i = thread_idx; i < kNumExperts; i += kNumDispatchThreads) { + const uint64_t send_value = (1ull << 32) | static_cast(smem_expert_count[i]); + smem_expert_count[i] = static_cast( + ptx::atomic_add(workspace.get_expert_send_count_ptr(i), send_value)); + } + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + + // Write source token-topk indices to remote ranks + read_topk_idx([&](const uint32_t& token_topk_idx, const int& expert_idx) { + const auto dst_rank_idx = expert_idx / kNumExpertsPerRank; + const auto dst_slot_idx = atomicAdd_block(smem_expert_count + expert_idx, 1); + const auto dst_ptr = workspace.get_src_token_topk_idx_ptr( + expert_idx % kNumExpertsPerRank, sym_buffer.rank_idx, dst_slot_idx); + *sym_buffer.map(dst_ptr, dst_rank_idx) = token_topk_idx; + }); + + // The routing counts are dead from here on (last readers: the slot allocations above, complete in both dispatch warps past + // this barrier). Their region is the head of the tile table: mark every entry "not published" and release the A/B loader + // warps, which wait on this barrier before their first table access; the math warpgroups are ordered behind it by + // kDispatchWithEpilogueBarrierIdx. + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + for (uint32_t i = thread_idx; i < kNumTileTableEntries; i += kNumDispatchThreads) + ptx::st_shared(tile_table + i, 0xffffffffu, 0xffffffffu); + ptx::sync_aligned(kNumDispatchWithLoadersThreads, kDispatchWithLoadersBarrierIdx); + + comm::grid_sync( + workspace, sm_idx, thread_idx, + [=]() { ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); } + ); + trace_dispatch(2, 0); // ROUTING_DONE (local grid sync: every CTA of this rank has staked out its slots) + + if (sm_idx == 0) { + { + // loads first, then the + // recv-count stores; no sum atomics -- the release flag below publishes the counts + constexpr uint32_t kNumPublishPerThread = math::constexpr_ceil_div(kNumExperts, kNumDispatchThreads); + uint64_t expert_status[kNumPublishPerThread]; + #pragma unroll + for (uint32_t k = 0; k < kNumPublishPerThread; ++ k) { + const uint32_t i = thread_idx + k * kNumDispatchThreads; + expert_status[k] = i < kNumExperts ? *workspace.get_expert_send_count_ptr(i) : 0ull; + } + #pragma unroll + for (uint32_t k = 0; k < kNumPublishPerThread; ++ k) { + const uint32_t i = thread_idx + k * kNumDispatchThreads; + if (i < kNumExperts) { + *sym_buffer.map(workspace.get_expert_recv_count_ptr(sym_buffer.rank_idx, i % kNumExpertsPerRank), + i / kNumExpertsPerRank) = expert_status[k] & 0xffffffff; + *workspace.get_expert_send_count_ptr(i) = 0; // (its only reader was the load above) + } + } + } + } + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + if (sm_idx == 0) + trace_dispatch(10, 0); // PUBLISH_DONE (SM0: every thread's recv-count stores issued) + + { + // SM0 publishes "my counts for you are stored" to every rank (release: cumulative over this CTA's stores + // through the barrier above and over every CTA's pass-2 slot stores through the routing grid sync); then every CTA + // waits for the num_ranks flags of this generation itself (one lane per source) -- no all-rank barrier, no grid sync + if (sm_idx == 0) { + if (thread_idx < kNumRanks) + ptx::st_release_sys_u32(sym_buffer.map(workspace.get_hll_count_flag_ptr(sym_buffer.rank_idx), thread_idx), 1u); + trace_dispatch(11, 0); // SIGNAL_SENT (the count flags issued = the count stores drained) + } + if (thread_idx < kNumRanks) + while (ptx::ld_acq_sys(workspace.get_hll_count_flag_ptr(thread_idx)) == 0u); + trace_dispatch(12, 0); // SIGNALS_SEEN (thread 0 saw... its own source's flag; 8 below = all of them) + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + } + trace_dispatch(8, 0); // NVLINK1_DONE (all ranks' expert counts published) + + { + // Zero the previous generation's bank (SM90FusedWorkspace::t2_bank): every remote generation-(g - 1) write into it precedes that + // rank's launch-g count flag, acquired above; the remote generation-(g + 1) writes follow our tag-2 signal, ordered behind + // these stores by the DISPATCH_SYNC_DONE barrier (math warps) -> tag-2 grid sync 1 (release). Bank layout (u32 words): + // [recv_count u64 x kNumExperts | recv_count_sum u64 x kNumExpertsPerRank] then [flags | tile counts | count flags] + constexpr uint32_t kNumT2CountWords = 2 * (kNumExperts + kNumExpertsPerRank); + constexpr uint32_t kNumT2Words = kNumT2CountWords + kNumExperts + kNumExpertsPerRank + kNumRanks; + constexpr uint32_t kNumT2WordsPerCTA = math::constexpr_ceil_div(kNumT2Words, kNumSMs); + DG_STATIC_ASSERT(kNumT2WordsPerCTA <= kNumDispatchThreads, "one bank word per dispatch thread"); + const uint32_t prev_bank = workspace.t2_bank ^ 1u; + const uint32_t w = sm_idx * kNumT2WordsPerCTA + thread_idx; + if (thread_idx < kNumT2WordsPerCTA and w < kNumT2Words) { + if (w < kNumT2CountWords) + reinterpret_cast(workspace.get_t2_recv_count_bank_ptr(prev_bank))[w] = 0u; + else + workspace.get_t2_flag_bank_ptr(prev_bank)[w - kNumT2CountWords] = 0u; + } + } + + // Sync with epilogue warps before pulling tokens. + ptx::sync_unaligned(kNumDispatchThreads + kNumEpilogueThreads, kDispatchWithEpilogueBarrierIdx); + trace_dispatch(4, 0); // DISPATCH_SYNC_DONE (math warps released; pull loop starts) + + // Token / SF pull loop + uint32_t pull_mbarrier_phase = 0; + const auto pull_buffer = smem_send_buffers.get_rank_buffer(warp_idx).get_data_buffer(0); + const auto pull_mbarrier = dispatch_barriers[warp_idx]; + // Pull pipeline: the local TMA store of row t is not waited for in iteration t; its arrival + // count is published in iteration t+1 (or after the loop) once the store has completed, so the + // store drain overlaps the next row's remote loads. The single pull buffer is reused only after + // the store has finished reading it (`.read` wait right before the next remote TMA load). + uint32_t* pending_arrival_ptr = nullptr; + uint32_t pending_arrival_add = 0; + bool bfd_signalled = false; // deferred ring fill (lane 0): this warp's first-row publish has been signalled to the B loader + + // Batched publish (kPullPublishBatch > 1): lane 0 publishes the warp's landed-but-unpublished rows together: one + // fence.acq_rel.gpu, then one relaxed red per row (the release pattern the loaders' ld.acquire pairs with; every pending + // row's store is complete and its SF / weight / metadata stores are ordered before by the __syncwarp that ended their + // iteration). Pending rows are flushed before the warp blocks on a ring slot (see the ring wait), so no slot release this + // warp waits for can depend on them. + constexpr bool kPullBatched = kPullPublishBatch > 1; + uint32_t num_pending = 0; // lane 0 (<= kPullPublishBatch) + const auto pending_list = smem_pull_pending + warp_idx * kPullPublishBatch; + const auto flush_pending = [&](const uint32_t& trace_aux) { + if constexpr (kPullBatched) { + asm volatile("fence.acq_rel.gpu;" ::: "memory"); + for (uint32_t i = 0; i < num_pending; ++ i) { + const uint32_t word = ptx::ld_shared(pending_list + i); + const uint32_t p_token_idx = word & 0x7fffffffu; + const uint32_t p_block_idx = p_token_idx / BLOCK_M; + if constexpr (kRingCoversFullPool) { + asm volatile("red.relaxed.gpu.global.add.u32 [%0], %1;" + :: "l"(workspace.get_l1_arrival_count_ptr(p_block_idx)), "r"(1u) : "memory"); + } else { + // the expert's last token pads its tail block to BLOCK_M arrivals (see the per-row path) + const uint32_t add = (word >> 31) ? BLOCK_M - (p_token_idx % BLOCK_M) : 1u; + asm volatile("red.relaxed.gpu.global.add.u32 [%0], %1;" + :: "l"(workspace.get_l1_full_count_ptr(get_ring_block_idx(p_block_idx))), "r"(add) : "memory"); + } + } + trace_dispatch(9, static_cast(num_pending) | (static_cast(trace_aux) << 8)); // PUBLISH (aux = rows | token_idx << 8) + num_pending = 0; + } + }; + + scheduler.fetch_expert_recv_count(); + + constexpr uint32_t kNumRanksPerLane = math::constexpr_ceil_div(kNumRanks, 32u); + int current_expert_idx = -1; + uint32_t stored_rank_count[kNumRanksPerLane] = {}; + uint32_t expert_start_idx = 0, expert_end_idx = 0; + uint32_t expert_pool_block_offset = 0; + + constexpr uint32_t kNumGlobalWarps = kNumSMs * kNumDispatchWarps; + for (uint32_t token_idx = sm_idx * kNumDispatchWarps + warp_idx; ; token_idx += kNumGlobalWarps) { + int old_expert_idx = current_expert_idx; + while (token_idx >= expert_end_idx) { + if (++ current_expert_idx >= kNumExpertsPerRank) + break; + expert_pool_block_offset += math::ceil_div(expert_end_idx - expert_start_idx, BLOCK_M); + expert_start_idx = expert_end_idx; + expert_end_idx += scheduler.get_num_tokens(current_expert_idx); + } + if (current_expert_idx >= kNumExpertsPerRank) + break; + + if (old_expert_idx != current_expert_idx) { + old_expert_idx = current_expert_idx; + #pragma unroll + for (uint32_t i = 0; i < kNumRanksPerLane; ++ i) { + const uint32_t j = i * 32 + lane_idx; + stored_rank_count[i] = j < kNumRanks ? + static_cast(*workspace.get_expert_recv_count_ptr(j, current_expert_idx)) : 0; + } + } + + // Round-robin rank selection + uint32_t current_rank_in_expert_idx; + uint32_t remaining[kNumRanksPerLane]; + #pragma unroll + for (uint32_t i = 0; i < kNumRanksPerLane; ++ i) + remaining[i] = stored_rank_count[i]; + uint32_t offset = 0; + uint32_t token_idx_in_expert = token_idx - expert_start_idx; + uint32_t slot_idx = token_idx_in_expert; + uint32_t token_idx_in_rank; + while (true) { + uint32_t num_actives_in_lane = 0; + uint32_t min_in_lane = 0xffffffff; + #pragma unroll + for (uint32_t i = 0; i < kNumRanksPerLane; ++ i) { + num_actives_in_lane += remaining[i] > 0; + if (remaining[i] > 0) + min_in_lane = cute::min(min_in_lane, remaining[i]); + } + const uint32_t num_active_ranks = __reduce_add_sync(0xffffffff, num_actives_in_lane); + const uint32_t length = __reduce_min_sync(0xffffffff, min_in_lane); + + const uint32_t num_round_tokens = length * num_active_ranks; + if (slot_idx < num_round_tokens) { + const uint32_t slot_idx_in_round = slot_idx % num_active_ranks; + uint32_t num_seen_ranks = 0; + current_rank_in_expert_idx = 0; + #pragma unroll + for (uint32_t i = 0; i < kNumRanksPerLane; ++ i) { + const uint32_t mask = __ballot_sync(0xffffffff, remaining[i] > 0); + const uint32_t num_active_lanes = __popc(mask); + if (slot_idx_in_round >= num_seen_ranks and slot_idx_in_round < num_seen_ranks + num_active_lanes) + current_rank_in_expert_idx = i * 32 + __fns(mask, 0, slot_idx_in_round - num_seen_ranks + 1); + num_seen_ranks += num_active_lanes; + } + token_idx_in_rank = offset + (slot_idx / num_active_ranks); + break; + } + slot_idx -= num_round_tokens; + offset += length; + #pragma unroll + for (uint32_t i = 0; i < kNumRanksPerLane; ++ i) + remaining[i] -= cute::min(remaining[i], length); + } + + const uint32_t src_token_topk_idx = *workspace.get_src_token_topk_idx_ptr( + current_expert_idx, current_rank_in_expert_idx, token_idx_in_rank); + const uint32_t src_token_idx = src_token_topk_idx / kNumTopk; + const uint32_t src_topk_idx = src_token_topk_idx % kNumTopk; + + const uint32_t pool_token_idx = expert_pool_block_offset * BLOCK_M + token_idx_in_expert; + const uint32_t pool_block_idx = expert_pool_block_offset + token_idx_in_expert / BLOCK_M; + + // Ring mode: wait until the previous generation of consumers + // (all L1 N-blocks of the pool block mapped to this slot) has + // released the physical slot before overwriting it. + if constexpr (not kRingCoversFullPool) { + constexpr uint32_t kNumL1BlockNs = L1_SHAPE_N / BLOCK_N; + const auto l1_empty_target = get_ring_wave_idx(pool_block_idx) * kNumL1BlockNs; + if (l1_empty_target > 0) { + const auto empty_ptr = workspace.get_l1_empty_count_ptr(get_ring_block_idx(pool_block_idx)); + if (ptx::ld_acq(empty_ptr) < l1_empty_target) { + // the slot is still held by the L1 tiles of block p - R, which complete only once every row of that block + // is published: publish this warp's pending rows before blocking, so the wait can never depend on a row + // this warp holds back (their stores have had the remote round trip to land) + if (lane_idx == 0) { + ptx::tma_store_wait<0>(); + if constexpr (kPullBatched) { + if (num_pending > 0) + flush_pending(token_idx); + } else if (pending_arrival_ptr != nullptr) { + ptx::red_add_rel(pending_arrival_ptr, pending_arrival_add); + pending_arrival_ptr = nullptr; + } + } + trace_dispatch(6, pool_block_idx); // RING_SLOT_WAIT_START (aux = pool block) + while (ptx::ld_acq(empty_ptr) < l1_empty_target); + trace_dispatch(7, pool_block_idx); // RING_SLOT_WAIT_END + } + } + } + + // Pull token data. Overlap a remote TMA load with SF copy and + // then use TMA store to materialize the local L1 input. + if (lane_idx == 0) { + // previous row's store must have finished *reading* the pull buffer (not necessarily + // landed in global) before it is overwritten by the next remote load + cute::tma_store_wait<0>(); + ptx::tma_load_1d( + pull_buffer.get_base_ptr(), + sym_buffer.map(input_token_buffer.get_data_buffer(src_token_idx).get_base_ptr(), + current_rank_in_expert_idx), + pull_mbarrier, kHidden); + } + __syncwarp(); + + // Copy SF: per-128 K floats, written linearly (no UTCCP transpose). + constexpr uint32_t kNumSFFloats = kHidden / 128; + DG_STATIC_ASSERT(kNumSFFloats > 0 and kHidden % 128 == 0, "Invalid SF"); + const auto remote_sf_ptr = sym_buffer.map( + input_sf_buffer.get_data_buffer(src_token_idx).get_base_ptr(), + current_rank_in_expert_idx); + const auto local_sf_ptr = l1_sf_buffer.get_base_ptr(); + const uint32_t token_idx_in_block = token_idx_in_expert % BLOCK_M; + const uint32_t sf_pool_token_idx = get_ring_block_idx(pool_block_idx) * SF_BLOCK_M + token_idx_in_block; + // weight (lane 0) and SF (all lanes) remote loads are issued together + float weight_val = 0.f; + if (lane_idx == 0) { + weight_val = *sym_buffer.map( + input_topk_weights_buffer.get_base_ptr() + src_token_topk_idx, + current_rank_in_expert_idx); + } + #pragma unroll + for (uint32_t i = 0; i < math::constexpr_ceil_div(kNumSFFloats, 32u); ++ i) { + const uint32_t j = i * 32 + lane_idx; + if (j < kNumSFFloats) + local_sf_ptr[j * kNumPaddedSFPoolTokens + sf_pool_token_idx] = remote_sf_ptr[j]; + } + if (lane_idx == 0) { + *l1_topk_weights_buffer.get_data_buffer(get_ring_token_idx(pool_token_idx)).template get_base_ptr() = weight_val; + if constexpr (not kPullBatched) { + // previous row: its local store has had the whole remote round trip to land + if (pending_arrival_ptr != nullptr) { + ptx::tma_store_wait<0>(); + ptx::red_add_rel(pending_arrival_ptr, pending_arrival_add); + pending_arrival_ptr = nullptr; + } + } else { + // batched publish: the pending rows (the latest stored one iteration ago, with the whole remote + // round trip to land) are published once the list reaches the batch limit; the ramp limit is + // min(N, 1 << (rows pulled so far by this warp / 8)) + const uint32_t limit = cute::min(kPullPublishBatch, 1u << cute::min(token_idx / (kNumGlobalWarps * 8u), 8u)); + if (num_pending >= limit) { + ptx::tma_store_wait<0>(); + flush_pending(token_idx); + } + } + } + __syncwarp(); + + if (lane_idx == 0) { + ptx::mbarrier_arrive_and_set_tx(pull_mbarrier, kHidden); + ptx::mbarrier_wait_and_flip_phase(pull_mbarrier, pull_mbarrier_phase); + + ptx::tma_store_1d( + l1_token_buffer.get_data_buffer(get_ring_token_idx(pool_token_idx)).get_base_ptr(), + pull_buffer.get_base_ptr(), pull_buffer.get_num_bytes()); + + *workspace.get_token_src_metadata_ptr(pool_token_idx) = + {current_rank_in_expert_idx, src_token_idx, src_topk_idx}; + + cute::tma_store_arrive(); + if constexpr (kPullBatched) { + // batched publish: queue this row (the list has room: it was flushed above once it reached the limit) + const bool is_last_token = (token_idx == expert_end_idx - 1); + ptx::st_shared(pending_list + num_pending, pool_token_idx | (is_last_token ? 0x80000000u : 0u)); + ++ num_pending; + } else if constexpr (kRingCoversFullPool) { + pending_arrival_ptr = workspace.get_l1_arrival_count_ptr(pool_block_idx); + pending_arrival_add = 1; + } else { + // Pad the tail of the expert's last m-block so that + // every pass of a ring slot contributes exactly + // BLOCK_M arrivals regardless of `valid_m`. + const bool is_last_token = (token_idx == expert_end_idx - 1); + pending_arrival_ptr = workspace.get_l1_full_count_ptr(get_ring_block_idx(pool_block_idx)); + pending_arrival_add = is_last_token ? BLOCK_M - (token_idx_in_expert % BLOCK_M) : 1u; + } + if constexpr (kPullEagerPublish) { + // eager publish: publish this row now (the SF / weight / metadata stores above are ordered before the + // release by the __syncwarp + program order, exactly as in the deferred form) + if (token_idx < kNumGlobalWarps) { + ptx::tma_store_wait<0>(); + ptx::red_add_rel(pending_arrival_ptr, pending_arrival_add); + pending_arrival_ptr = nullptr; + trace_dispatch(9, 1ull | (static_cast(token_idx) << 8)); // PUBLISH (aux = rows | token_idx << 8) + if constexpr (kBFillDefer) { + if (token_idx < kNumGlobalWarps and not bfd_signalled) { + ptx::red_release_cta_shared_add(smem_ec_gs_old, 1u); // deferred ring fill: this warp's first row is published + bfd_signalled = true; + } + } + } + } else if constexpr (kBFillDefer) { + // deferred ring fill without the eager publish: the first row is published at the top of the next iteration + // (or in the drain) -- signal when that publish is issued (pending_arrival_ptr == nullptr again after a first row) + if (token_idx >= kNumGlobalWarps and not bfd_signalled) { + ptx::red_release_cta_shared_add(smem_ec_gs_old, 1u); + bfd_signalled = true; + } + } + } + __syncwarp(); + } + // drain: the last row's store and arrival / the pending list (batched publish) + if (lane_idx == 0) { + if constexpr (kPullBatched) { + if (num_pending > 0) { + ptx::tma_store_wait<0>(); + flush_pending(0xffffffu); + } + } else if (pending_arrival_ptr != nullptr) { + ptx::tma_store_wait<0>(); + ptx::red_add_rel(pending_arrival_ptr, pending_arrival_add); + pending_arrival_ptr = nullptr; + } + } + if constexpr (kBFillDefer) { + // deferred ring fill: a warp with no row (or whose first row was published by the drain) signals here, so the B loader never waits forever + if (lane_idx == 0 and not bfd_signalled) + ptx::red_release_cta_shared_add(smem_ec_gs_old, 1u); + } + __syncwarp(); + trace_dispatch(3, 0); // PULL_DONE (this warp's last pull stored and its arrival published) + + // ================= Early combine on the dispatch warps (see the template parameter) ================= + // Source (warp 0, kEarlyCombinePublish): the math warps bump smem_ec_tiles (release.cta) per warp per L2 tile after their + // scatter stores; warp 0 walks the tile table in step: one fence.acq_rel.gpu per batch of finished tiles, one relaxed + // gpu-scope count add per tile, and for an expert's last tile fence.acq_rel.sys + relaxed sys-scope adds to its flag on + // every rank. Handover: the math warps set smem_ec_stop after the tag-2 barrier; both warps poll it and join the sync below. + uint32_t ec_tiles_published = 0; + if constexpr (kEarlyCombine) { + // source-side publish (warp 0) + uint32_t ec_table_pos = 0; + bool ec_table_done = false; + const auto ec_publish = [&]() { + if constexpr (kEarlyCombinePublish) { + const uint32_t tiles_done = ptx::ld_acquire_cta_shared(smem_ec_tiles); + const bool all_done = ec_table_done or ptx::ld_acquire_cta_shared(smem_ec_done) != 0; + if (ec_table_done or (not all_done and tiles_done == ec_tiles_published)) + return; + // one gpu-scope fence per batch: the scatter stores of these tiles (ordered before the counter / done flag by + // the math warps' release arrives and the loader's release add) become visible at gpu scope before the + // relaxed count adds + ptx::fence_acq_rel_gpu(); + while (all_done or ec_tiles_published < tiles_done) { + const uint32_t tag = scheduler.load_tile_entry(tile_table + ec_table_pos); + if (tag == 0u) { // end marker: every tile is accounted + ec_table_done = true; + break; + } + ++ ec_table_pos; + if (tag != 2u) + continue; + ++ ec_tiles_published; + const uint32_t e = scheduler.current_local_expert_idx; + // the expert's L2 tile count: from the entry's token count, or from the per-expert recv counts this + // warp fetched for the pull when the entry carries the tile's own row count (scheduler kTilePayloadByRows) + uint32_t expert_num_tokens = scheduler.current_num_tokens; + if constexpr (std::remove_reference_t::kTilePayloadByRows) + expert_num_tokens = scheduler.get_num_tokens(e); + const uint32_t total = math::ceil_div(expert_num_tokens, BLOCK_M) * kNumL2BlockNs; + if (lane_idx == 0) { + const uint32_t old = ptx::atom_add_relaxed_gpu(workspace.get_l2_tile_done_count_ptr(e), 1u); + if (old + 1u == total) { + // this CTA scattered the expert's last tile: the sys fence is the acquire side of the count (the + // other CTAs' releases) and the release side of the flags every rank's dispatch warps poll + ptx::fence_acq_rel_sys(); + const auto flag_ptr = workspace.get_expert_done_flag_ptr(sym_buffer.rank_idx * kNumExpertsPerRank + e); + #pragma unroll + for (uint32_t r = 0; r < kNumRanks; ++ r) + ptx::red_add_relaxed_sys(sym_buffer.map(flag_ptr, r), 1u); + trace_dispatch(71, e); // EC_EXPERT_PUBLISHED (aux = local expert) + } + } + } + __syncwarp(); + } + }; + + while (true) { + if (warp_idx == 0) + ec_publish(); + if (ptx::ld_volatile_shared(smem_ec_stop) != 0) + break; + __nanosleep(1000); + } + // handover: every store of this warp is performed before the math warps' combine (nothing of this SM is outstanding at + // this point, so the fence is cheap) + ptx::fence_acq_rel_gpu(); + if (lane_idx == 0) + ptx::st_shared(smem_ec_count + warp_idx, 0u); // trace counter (tokens combined by this warp: none) + } + + // Cleanup workspace, overlapping with combine. + if constexpr (kEarlyCombine) { + // early combine: the math warps do not rendezvous here; their thread 0 sets the stop flag once the all-rank barrier completed + while (ptx::ld_acquire_cta_shared(smem_ec_stop) == 0) + __nanosleep(500); + } else { + ptx::sync_unaligned(kNumDispatchThreads + kNumEpilogueThreads, kDispatchWithEpilogueBarrierIdx); + } + trace_dispatch(5, 0); // CLEAN_START (math warps passed the pre-combine all-rank barrier) + if constexpr (kEarlyCombine and kTrace) { + // EC_STOP: aux = tokens combined by dispatch warp 0 | warp 1 << 16 | L2 tiles this CTA published << 32 + trace_dispatch(72, static_cast(ptx::ld_shared(smem_ec_count)) | (static_cast(ptx::ld_shared(smem_ec_count + 1)) << 16) | + (static_cast(ec_tiles_published) << 32)); + } + + DG_STATIC_ASSERT(kNumSMs > 1, "Invalid SM count"); + if (sm_idx == 0) { + // (the send counts were zeroed in the head, right after the publish loop) + // every B loader fetched its end-of-tail ticket before its math warps reached the pre-combine barrier + if (thread_idx == 0) + *workspace.template get_grid_sync_count_ptr() = 0; + // flip the launch parity for the next launch (every read of this launch precedes this store: the loaders / math + // warps of each CTA sync with their dispatch warps before those reach the routing grid sync, and the math warps' + // tag-2 grid sync 1 precedes this CTA's CLEAN_START); the next launch reads it after the kernel boundary + if (thread_idx == 0) + *workspace.get_t2_parity_ptr() = workspace.t2_bank ^ 1u; + } else { + for (uint32_t i = sm_idx - 1; i < kNumExpertsPerRank; i += kNumSMs - 1) { + // the expert's count is the sum of the per-rank recv counts (all landed before the head's flags) + uint32_t num_recv_tokens = 0; + #pragma unroll + for (uint32_t r = 0; r < kNumRanks; ++ r) + num_recv_tokens += static_cast(*workspace.get_expert_recv_count_ptr(r, i)); + const auto num_recv_m_blocks = math::ceil_div(num_recv_tokens, BLOCK_M); + + expert_pool_block_offset = scheduler.get_pool_block_offset(i); + + ptx::sync_aligned(kNumDispatchThreads, kDispatchBarrierIdx); + + DG_STATIC_ASSERT(kNumDispatchWarps >= 2, "Not enough dispatch warps"); + if (warp_idx == 1) { + if (cute::elect_one_sync() and cumulative_local_expert_recv_stats != nullptr) + ptx::red_add(cumulative_local_expert_recv_stats + i, static_cast(num_recv_tokens)); + __syncwarp(); + } + + for (uint32_t j = thread_idx; j < num_recv_m_blocks; j += kNumDispatchThreads) { + if constexpr (kRingCoversFullPool) { + *workspace.get_l1_arrival_count_ptr(expert_pool_block_offset + j) = 0; + *workspace.get_l2_arrival_mask_ptr(expert_pool_block_offset + j) = 0; + } else { + const auto ring_block_idx = get_ring_block_idx(expert_pool_block_offset + j); + *workspace.get_l1_full_count_ptr(ring_block_idx) = 0; + *workspace.get_l2_full_count_ptr(ring_block_idx) = 0; + *workspace.get_l1_empty_count_ptr(ring_block_idx) = 0; + *workspace.get_l2_empty_count_ptr(ring_block_idx) = 0; + } + } + __syncwarp(); + } + } + + // No exit barrier. The arrival / ring counts zeroed above have no user left on this rank (every CTA is past + // TILES_DONE = tag-2 grid sync 1) and their next users are in the next launch, behind the kernel boundary; the + // remotely written words are bank words, zeroed by the next launch after its head. Nothing to wait for: the + // dispatch warps exit while the math warps combine. + trace_dispatch(22, 0); // KERNEL_END (dispatch warp's last statement) + + // ===================================================================== + // ROLE 2: GEMM TMA LOAD warps (warp 0 of `kNumNonEpilogueThreads` loads A + SFA, warp 1 loads B + SFB) + // ===================================================================== + } else if (warp_idx == kNumDispatchWarps) { + cutlass::arch::warpgroup_reg_dealloc(); + // the tile table is initialised by the dispatch warps after routing (its head is the routing-count region) + ptx::sync_aligned(kNumDispatchWithLoadersThreads, kDispatchWithLoadersBarrierIdx); + + // trace role 1: tile ordinal (same schedule as the math warps) and per-tile arrival wait / issue done + uint32_t trace_tile = 0; + const auto trace_loader = [&](const uint32_t& event_id, const uint64_t& aux) { + if constexpr (kTrace) { + if (lane_idx == 0) + trace_event(1, event_id, aux); + } + }; + + auto process_a_sfa_block = [&](const auto& block_phase, + const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + const auto tensor_map_a_ptr = block_phase == sched::BlockPhase::Linear2 + ? &tensor_map_l2_acts : &tensor_map_l1_acts; + const auto tensor_map_sfa_ptr = block_phase == sched::BlockPhase::Linear2 + ? &tensor_map_l2_acts_sf : &tensor_map_l1_acts_sf; + + const uint32_t pool_block_idx = scheduler.get_current_pool_block_offset() + m_block_idx; + const uint32_t ring_block_idx = get_ring_block_idx(pool_block_idx); + + const uint64_t trace_aux = static_cast(trace_tile) | + (static_cast(block_phase == sched::BlockPhase::Linear1 ? 1u : 2u) << 32); + if constexpr (kTrace) + ++ trace_tile; + trace_loader(30, trace_aux); // WAIT_ARRIVAL_START + + // Wait for the pool to be ready + if (block_phase == sched::BlockPhase::Linear1) { + if constexpr (kRingCoversFullPool) { + const auto ptr = workspace.get_l1_arrival_count_ptr(pool_block_idx); + const auto expected = scheduler.template get_valid_m(); + while (ptx::ld_acq(ptr) != expected); + } else { + // Ring mode: every pass of this slot contributed exactly + // BLOCK_M arrivals (dispatch pads the tail block). + const auto ptr = workspace.get_l1_full_count_ptr(ring_block_idx); + const uint32_t expected = BLOCK_M * (get_ring_wave_idx(pool_block_idx) + 1); + while (ptx::ld_acq(ptr) != expected); + } + } else { + constexpr uint32_t kNumL1BlockNs = L1_SHAPE_N / BLOCK_N; + if constexpr (not kRingCoversFullPool) { + // Ring mode: one arrival per (m, n) L1 block after its + // output TMA store drained (pass-constant accounting). + const auto ptr = workspace.get_l2_full_count_ptr(ring_block_idx); + const uint32_t expected = kNumL1BlockNs * (get_ring_wave_idx(pool_block_idx) + 1); + while (ptx::ld_acq(ptr) != expected); + } else if constexpr (kL2ArrivalCounter) { + const auto ptr = reinterpret_cast( + workspace.get_l2_arrival_mask_ptr(pool_block_idx)); + const uint32_t active_m_wgs = math::ceil_div( + scheduler.template get_valid_m(), WG_BLOCK_M); + const uint32_t expected = + kNumL1BlockNs * active_m_wgs * kWarpgroupSplitN * kL1OutputArrivalParts; + while (ptx::ld_acq(ptr) != expected); + } else { + const auto ptr = workspace.get_l2_arrival_mask_ptr(pool_block_idx); + const uint64_t expected = (kNumL1BlockNs >= 64) + ? ~0ull : ((1ull << kNumL1BlockNs) - 1ull); + while (ptx::ld_acq_gpu(ptr) != expected); + } + } + trace_loader(31, trace_aux); // WAIT_ARRIVAL_END + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; advance_pipeline(k_block_idx)) { + empty_barriers[stage_idx]->wait(phase ^ 1); + + if (cute::elect_one_sync()) { + const uint32_t m_idx = ring_block_idx * BLOCK_M; + const uint32_t sfa_m_idx = ring_block_idx * SF_BLOCK_M; + const uint32_t k_idx = k_block_idx * BLOCK_K; + + // TMA load A + tma::copy( + tensor_map_a_ptr, full_barriers[stage_idx], smem_a[stage_idx], + k_idx, m_idx, kNumMulticastA); + + // TMA load SFA + if (kL2ActSFPerBlockK or block_phase == sched::BlockPhase::Linear1) { + // L1 SFA per-128 (and the per-128 L2 SFA): load (BLOCK_M, 1) at K=k_block_idx + tma::copy( + tensor_map_sfa_ptr, full_barriers[stage_idx], smem_sfa[stage_idx], + sfa_m_idx, k_block_idx, kNumMulticastA); + full_barriers[stage_idx]->arrive_and_expect_tx( + SMEM_A_SIZE_PER_STAGE + BLOCK_M * sizeof(float)); + } else { + // L2 SFA per-64: descriptor box is (block_mn, 1) (see make_tma_sf_desc), + // so we must issue two single-group TMAs and place them at smem offsets + // 0 and BLOCK_M to match math's load offsets (`+ 0 * BLOCK_M` / `+ 1 * BLOCK_M`). + tma::copy( + tensor_map_sfa_ptr, full_barriers[stage_idx], smem_sfa[stage_idx], + sfa_m_idx, k_block_idx * 2, kNumMulticastA); + tma::copy( + tensor_map_sfa_ptr, full_barriers[stage_idx], + smem_sfa[stage_idx] + BLOCK_M, + sfa_m_idx, k_block_idx * 2 + 1, kNumMulticastA); + full_barriers[stage_idx]->arrive_and_expect_tx( + SMEM_A_SIZE_PER_STAGE + 2 * BLOCK_M * sizeof(float)); + } + } + __syncwarp(); + } + trace_loader(32, trace_aux); // TILE_LOADS_DONE (last k-block's TMA issued) + }; + + scheduler.for_each_block_replay(tile_table, + [&](const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + process_a_sfa_block( + std::integral_constant{}, + local_expert_idx, num_k_blocks, m_block_idx, n_block_idx); + }, + [&](const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + process_a_sfa_block( + std::integral_constant{}, + local_expert_idx, num_k_blocks, m_block_idx, n_block_idx); + }); + + } else if (warp_idx == kNumDispatchWarps + 1) { + cutlass::arch::warpgroup_reg_dealloc(); + // publish only into the initialised table (see the A loader / the dispatch warps) + ptx::sync_aligned(kNumDispatchWithLoadersThreads, kDispatchWithLoadersBarrierIdx); + bool bfd_pending = kBFillDefer; // deferred ring fill: the ring fill waits for the CTA's dispatch warps' first-row publishes + + // early combine: the previous tile was an L2 tile whose "scatter issued" signal is still owed (see smem_ec_tiles) + bool ec_prev_l2 = false; + auto process_b_block = [&](const sched::BlockPhase& block_phase, + const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + const auto tensor_map_b_ptr = + block_phase == sched::BlockPhase::Linear2 ? &tensor_map_l2_weights : &tensor_map_l1_weights; + + const uint32_t shape_n = block_phase == sched::BlockPhase::Linear2 ? L2_SHAPE_N : L1_SHAPE_N; + + // B multicast validity. + uint32_t num_multicast_b = 1; + if constexpr (kMulticastOnB and kClusterSize > 1) { + const bool pair_valid = scheduler.is_pair_valid(cute::block_rank_in_cluster() == 0); + num_multicast_b = pair_valid ? kClusterSize : 1; + } + + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; advance_pipeline(k_block_idx)) { + if constexpr (kBFillDefer) { + // deferred ring fill: the first stage goes out at once, the rest of the ring fill after the first pulled rows + if (bfd_pending and k_block_idx == 1) { + if (lane_idx == 0) + while (ptx::ld_acquire_cta_shared(smem_ec_gs_old) < kNumDispatchWarps); + __syncwarp(); + bfd_pending = false; + } + } + empty_barriers[stage_idx]->wait(phase ^ 1); + // early combine: this wait returned once every consumer warp released this tile's k-block 0 (the stage kNumStages + // k-blocks back), so all of them have finished the previous tile's epilogue: signal that L2 tile's scatter as + // issued (release.cta: the math warps' arrives are release.cta, this wait is acquire.cta -> the stores are ordered) + if constexpr (kEarlyCombine) { + if (k_block_idx == kNumStages and ec_prev_l2) { + if (lane_idx == 0) + ptx::red_release_cta_shared_add(smem_ec_tiles, 1u); + ec_prev_l2 = false; + } + } + + if (cute::elect_one_sync()) { + const uint32_t n_idx = local_expert_idx * shape_n + n_block_idx * BLOCK_N; + const uint32_t k_idx = k_block_idx * BLOCK_K; + + // TMA load B (weight SF is now loaded directly by math warps from global) + if constexpr (LOAD_BLOCK_N <= 256) { + tma::copy( + tensor_map_b_ptr, full_barriers[stage_idx], smem_b[stage_idx], + k_idx, n_idx, num_multicast_b); + } else { + DG_STATIC_ASSERT(LOAD_BLOCK_N % 256 == 0, + "Large B tiles are loaded as 256-column TMA slices"); + #pragma unroll + for (uint32_t b_slice_idx = 0; b_slice_idx < LOAD_BLOCK_N / 256; ++ b_slice_idx) { + tma::copy( + tensor_map_b_ptr, full_barriers[stage_idx], + smem_b[stage_idx] + b_slice_idx * 256 * BLOCK_K, + k_idx, n_idx + b_slice_idx * 256, num_multicast_b); + } + } + + full_barriers[stage_idx]->arrive_and_expect_tx(SMEM_B_SIZE_PER_STAGE); + } + __syncwarp(); + } + if constexpr (kEarlyCombine) { + DG_STATIC_ASSERT(kNumL2BlockKs > kNumStages and kNumL1BlockKs > kNumStages, "the early-combine signal needs k-block kNumStages of the next tile"); + ec_prev_l2 = block_phase == sched::BlockPhase::Linear2; + } + }; + scheduler.template for_each_block_publish_dynamic_tail( + tile_table, workspace.template get_grid_sync_count_ptr(), + kClusterSize > 1 ? cute::block_rank_in_cluster() : 0u, kNumTileTableEntries, process_b_block); + + } else if (warp_idx < kNumDispatchWarps + kNumMMANonEpilogueWarps) { + // Idle non-epilogue warps (kNumDispatchWarps+2, +3). They must still + // participate in the warpgroup-collective `setmaxnreg.dec.sync.aligned` + // so that the math warpgroup's `warpgroup_reg_alloc` can succeed. + cutlass::arch::warpgroup_reg_dealloc(); + + } else if (warp_idx >= kNumDispatchWarps + kNumMMANonEpilogueWarps) { + // ===================================================================== + // ROLE 3: MATH WARPGROUPS (WGMMA + epilogue + combine) + // ===================================================================== + cutlass::arch::warpgroup_reg_alloc(); + + const uint32_t epilogue_warp_idx = warp_idx - (kNumDispatchWarps + kNumMMANonEpilogueWarps); + const uint32_t epilogue_wg_idx = epilogue_warp_idx / 4; + const uint32_t epilogue_thread_idx = epilogue_warp_idx * 32 + lane_idx; + const uint32_t warp_idx_in_wg = epilogue_warp_idx % 4; + + // WGMMA-output register layout helpers + const uint32_t row_idx = lane_idx / 4; + const uint32_t col_idx = lane_idx % 4; + const uint32_t r_0 = warp_idx_in_wg * 16 + row_idx; + const uint32_t r_1 = r_0 + 8; + + // When the two N-split warpgroups share a single per-64 SF group they + // also stage into ONE shared row-major L1-output tile (stride + // L1_OUT_BLOCK_N), each writing its own WG_L1_OUT_BLOCK_N-column half, + // so a single combined TMA store matches the host descriptor box. + constexpr uint32_t WG_SMEM_CD_L1_STRIDE_N = + kSplitNSharesSF ? L1_OUT_BLOCK_N : WG_L1_OUT_BLOCK_N; + constexpr uint32_t WG_SMEM_CD_L2_STRIDE_N = WG_BLOCK_N; + + // Sync with dispatch in the full communication path. + ptx::sync_unaligned(kNumDispatchThreads + kNumEpilogueThreads, kDispatchWithEpilogueBarrierIdx); + + // trace roles 2/3: warp 0 lane 0 of math warpgroups 0/1 (further warpgroups are not traced) + uint32_t trace_tile = 0; + const auto trace_math = [&](const uint32_t& event_id, const uint64_t& aux) { + if constexpr (kTrace) { + if (warp_idx_in_wg == 0 and lane_idx == 0 and epilogue_wg_idx < 2) + trace_event(2 + epilogue_wg_idx, event_id, aux); + } + }; + // epilogue sub-events (kTraceEpi): L1 50 EPI_MATH_DONE 51 EPI_STAGED 52 EPI_STORE_ISSUED 53 EPI_STORE_WAITED + // 54 EPI_PUBLISHED; L2 60 EPI_CVT_DONE / 61 EPI_SCATTER_DONE (aux = pass) 62 EPI_FULL_SYNC_DONE + const auto trace_epi = [&](const uint32_t& event_id, const uint64_t& aux) { + if constexpr (kTrace and kTraceEpi) + trace_math(event_id, aux); + }; + DG_STATIC_ASSERT(kTrace or not kTraceEpi, "epilogue sub-events need the trace"); + + // Weight-SF prologue prefetch (decode topology; hidden <= 4096 and intermediate <= 2048 so one + // warp's lanes stage the whole weight-SF row): math warp 0 peeks the NEXT tile-table entry right after the mainloop + // (non-blocking; an unpublished entry falls back to the plain prologue) and loads that tile's weight SF row into + // registers, which the next prologue stores: the same values, the same smem layout. + constexpr bool kSFProloguePrefetch = kDecodeTopology and kHidden / 128 <= 32 and 2 * (kIntermediateHidden / 128) <= 32; + // the next tile's weight SF, loaded by math warp 0 (lanes j < 32: L1 gate / up SF of + // k-block j, L2 lo / hi SF) under the drain of the current tile's k-block num_k - 2 and stored in the next prologue. + // sf_pf_key identifies the tile the values belong to (1 | is_l1 << 1 | sf_n_block << 2 | expert << 10; 0 = none). + uint32_t sf_pf_key = 0; + float sf_pf_v0 = 0.0f, sf_pf_v1 = 0.0f; + + // CTA-wide barrier of the math warps between tile phases + const auto sync_tile_math = [&]() { + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + }; + const bool tile_lead_warp = epilogue_warp_idx == 0; + + auto process_math_block = [&](const auto& block_phase, + const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + const uint32_t valid_m = scheduler.template get_valid_m(); + const uint32_t pool_block_idx = scheduler.get_current_pool_block_offset() + m_block_idx; + const uint32_t ring_block_idx = get_ring_block_idx(pool_block_idx); + const uint32_t m_idx = pool_block_idx * BLOCK_M; // Full-pool offset for metadata + const uint32_t ring_m_idx = ring_block_idx * BLOCK_M; // Ring offset for data buffers + const uint32_t n_idx = n_block_idx * BLOCK_N; + // TILE_START: aux = phase | expert << 8 | m_block << 24 | n_block << 32 | tile ordinal << 40 + trace_math(40, static_cast(block_phase == sched::BlockPhase::Linear1 ? 1u : 2u) | + (static_cast(local_expert_idx) << 8) | + (static_cast(m_block_idx) << 24) | + (static_cast(n_block_idx) << 32) | + (static_cast(trace_tile) << 40)); + if constexpr (kTrace) + ++ trace_tile; + const uint32_t epilogue_wg_m_idx = epilogue_wg_idx / kWarpgroupSplitN; + const uint32_t epilogue_wg_n_idx = epilogue_wg_idx - epilogue_wg_m_idx * kWarpgroupSplitN; + const uint32_t wg_n_offset = epilogue_wg_n_idx * WG_BLOCK_N; + const uint32_t wg_l1_out_n_offset = epilogue_wg_n_idx * WG_L1_OUT_BLOCK_N; + const uint32_t row_base = epilogue_wg_m_idx * WG_BLOCK_M; + const uint32_t row_offset_r0 = row_base + r_0; + const uint32_t row_offset_r1 = row_base + r_1; + // 128-column weight-SF block of the warpgroup's first column (L1: gate/up block pair; L2: lo block) + const uint32_t sf_n_block_idx = kSplitNSharesWeightSF ? n_block_idx + : ((n_block_idx * kWarpgroupSplitN + epilogue_wg_n_idx) * kNumSFGroupsPerWG); + // L2 act-SF group (post-SwiGLU columns / kL2ActsSFGranK) of the warpgroup's first L1 output column + const uint32_t l2_sf_group_base = kSplitNSharesSF ? n_block_idx + : ((n_block_idx * kWarpgroupSplitN + epilogue_wg_n_idx) * kNumL1OutSFGroups); + const uint32_t smem_a_wg_offset = epilogue_wg_m_idx * WG_BLOCK_M * BLOCK_K; + const uint32_t smem_b_wg_offset = epilogue_wg_n_idx * WG_BLOCK_N * BLOCK_K; + // In the shared-tile case the WG stages into the joint L1-output tile + // at its own column offset (row stride L1_OUT_BLOCK_N); otherwise each + // WG owns a disjoint contiguous WG_BLOCK_M x WG_L1_OUT_BLOCK_N slice. + const uint32_t smem_cd_l1_wg_offset = + kSplitNSharesSF ? wg_l1_out_n_offset : (epilogue_wg_idx * WG_BLOCK_M * WG_L1_OUT_BLOCK_N); + // With kHalfL2CD the L2 BF16 tile is staged one N-half at a time, so + // each WG's slot is only WG_BLOCK_N/2 wide (see the L2 epilogue). + const uint32_t smem_cd_l2_wg_offset = + epilogue_wg_idx * WG_BLOCK_M * (WG_BLOCK_N / kNumL2CDPasses); + const bool valid_r0 = row_offset_r0 < valid_m; + const bool valid_r1 = row_offset_r1 < valid_m; + + constexpr uint32_t kL1SFKBlocks = kHidden / 128; + constexpr uint32_t kL2SFKBlocks = kIntermediateHidden / 128; + constexpr uint32_t kL1SFGateBlks = kIntermediateHidden / 128; + constexpr uint32_t kL1SFPerExpert = (kIntermediateHidden * 2 / 128) * kL1SFKBlocks; + constexpr uint32_t kL2SFPerExpert = (kHidden / 128) * kL2SFKBlocks; + float* smem_weight_sf_wg = smem_weight_sf + epilogue_wg_idx * kNumWeightSFFloatsPerWG; + const uint32_t thread_idx_in_wg = warp_idx_in_wg * 32 + lane_idx; + + ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + + // prologue prefetch: warp 0 holds this tile's SF in registers when its peek at the previous tile's k-block num_k - 2 saw this + // entry published (warp-uniform: the key is lane-uniform in warp 0 and 0 in the other warps, whose staging loop + // is empty anyway). The stores below are the loads of the else branch evaluated earlier: same values. + bool sf_pf_hit = false; + if constexpr (kSFProloguePrefetch) { + const uint32_t my_key = 1u | ((block_phase == sched::BlockPhase::Linear1 ? 1u : 0u) << 1) | (sf_n_block_idx << 2) | + (local_expert_idx << 10); + sf_pf_hit = sf_pf_key == my_key; + sf_pf_key = 0; + } + if (kSFProloguePrefetch and sf_pf_hit) { + if (block_phase == sched::BlockPhase::Linear1) { + if (thread_idx_in_wg < kL1SFKBlocks) { + smem_weight_sf_wg[thread_idx_in_wg] = sf_pf_v0; + smem_weight_sf_wg[kL1SFKBlocks + thread_idx_in_wg] = sf_pf_v1; + } + } else { + if (thread_idx_in_wg < kNumSFGroupsPerWG * kL2SFKBlocks) + smem_weight_sf_wg[thread_idx_in_wg] = sf_pf_v0; + } + } else if (block_phase == sched::BlockPhase::Linear1) { + const uint32_t gate_n = sf_n_block_idx / 2u; + const uint32_t up_n = kL1SFGateBlks + gate_n; + const float* expert_base = l1_weights_sf + local_expert_idx * kL1SFPerExpert; + #pragma unroll + for (uint32_t j = thread_idx_in_wg; j < kL1SFKBlocks; j += 128) { + smem_weight_sf_wg[j] = __ldg(expert_base + gate_n * kL1SFKBlocks + j); + smem_weight_sf_wg[kL1SFKBlocks + j] = __ldg(expert_base + up_n * kL1SFKBlocks + j); + } + } else { + const float* sf_row = l2_weights_sf + local_expert_idx * kL2SFPerExpert + sf_n_block_idx * kL2SFKBlocks; + #pragma unroll + for (uint32_t j = thread_idx_in_wg; j < kNumSFGroupsPerWG * kL2SFKBlocks; j += 128) + smem_weight_sf_wg[j] = __ldg(sf_row + j); + } + ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + trace_math(41, kSFProloguePrefetch ? (sf_pf_hit ? 1u : 0u) : 0u); // SF_READY (weight SF staged; mainloop about to start; aux 1 = prefetched) + + // prologue prefetch: warp 0 peeks the tile table's next entry (a not-ready tag (3) or the end marker (0) leaves + // sf_pf_key at 0 and the next prologue loads as before) and issues the next tile's weight-SF loads into registers: + // lane j takes k-block j's (gate, up) SF for an L1 tile, the lo / hi SF for an L2 tile (event 48 SF_PEEK in the trace). + const auto sf_peek = [&]() { + if constexpr (kSFProloguePrefetch) { + using Sched = std::remove_reference_t; + const TileEntry* next_entry = tile_table + scheduler.replay_next_idx; + const uint32_t tag_word = Sched::ld_acquire_tile_tag(next_entry); + const uint32_t tag = tag_word >> 30; + trace_math(48, tag | (2u << 4)); // SF_PEEK (aux = tag | 2 << 4: the peek after the mainloop) + if (tag == 1u or tag == 2u) { + uint32_t nx_expert, nx_n_block; + nx_expert = tag_word & ((1u << Sched::kEBits) - 1u); + nx_n_block = (tag_word >> (Sched::kEBits + Sched::kMBits)) & ((1u << Sched::kNBits) - 1u); + const uint32_t nx_sf_n = kSplitNSharesWeightSF ? nx_n_block + : ((nx_n_block * kWarpgroupSplitN + epilogue_wg_n_idx) * kNumSFGroupsPerWG); + if (tag == 1u) { + const uint32_t gate_n = nx_sf_n / 2u; + const uint32_t up_n = kL1SFGateBlks + gate_n; + const float* expert_base = l1_weights_sf + nx_expert * kL1SFPerExpert; + if (lane_idx < kL1SFKBlocks) { + sf_pf_v0 = __ldg(expert_base + gate_n * kL1SFKBlocks + lane_idx); + sf_pf_v1 = __ldg(expert_base + up_n * kL1SFKBlocks + lane_idx); + } + } else { + const float* sf_row = l2_weights_sf + nx_expert * kL2SFPerExpert + nx_sf_n * kL2SFKBlocks; + if (lane_idx < kNumSFGroupsPerWG * kL2SFKBlocks) + sf_pf_v0 = __ldg(sf_row + lane_idx); + } + sf_pf_key = 1u | ((tag == 1u ? 1u : 0u) << 1) | (nx_sf_n << 2) | (nx_expert << 10); + } + } + }; + + // ---------------- GEMM ---------------- + using WGMMA = L1WGMMA; + constexpr uint32_t kAccumPerThread = WGMMA::kNumAccum; + + DG_STATIC_ASSERT(SMEM_A_SIZE_PER_STAGE % 16 == 0 and SMEM_B_SIZE_PER_STAGE % 16 == 0, + "Descriptor folding needs 16B-aligned stage strides"); + const auto desc_a_stage0 = mma::sm90::make_smem_desc(smem_a[0], 1); + const auto desc_b_stage0 = mma::sm90::make_smem_desc(smem_b[0], 1); + auto smem_desc_a = [&](const uint32_t& stage, const uint32_t& byte_offset) { + return mma::sm90::advance_smem_desc(desc_a_stage0, stage * SMEM_A_SIZE_PER_STAGE + byte_offset); + }; + auto smem_desc_b = [&](const uint32_t& stage, const uint32_t& byte_offset) { + return mma::sm90::advance_smem_desc(desc_b_stage0, stage * SMEM_B_SIZE_PER_STAGE + byte_offset); + }; + + float final_accum[kAccumPerThread] = {}; + + if constexpr (kReuseAccumAsFinal) { + auto prescale_l1_final = [&](const float& scale_a_0, const float& scale_a_1, + const float& gate_sf, const float& up_sf) { + const float inv_s0_gate = kFastMath ? math::fast_rcp(scale_a_0 * gate_sf) : 1.0f / (scale_a_0 * gate_sf); + const float inv_s1_gate = kFastMath ? math::fast_rcp(scale_a_1 * gate_sf) : 1.0f / (scale_a_1 * gate_sf); + const float inv_s0_up = kFastMath ? math::fast_rcp(scale_a_0 * up_sf) : 1.0f / (scale_a_0 * up_sf); + const float inv_s1_up = kFastMath ? math::fast_rcp(scale_a_1 * up_sf) : 1.0f / (scale_a_1 * up_sf); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const float inv_s0 = (i & 1u) ? inv_s0_up : inv_s0_gate; + const float inv_s1 = (i & 1u) ? inv_s1_up : inv_s1_gate; + final_accum[i*4+0] *= inv_s0; + final_accum[i*4+1] *= inv_s0; + final_accum[i*4+2] *= inv_s1; + final_accum[i*4+3] *= inv_s1; + } + }; + auto postscale_l1_final = [&](const float& scale_a_0, const float& scale_a_1, + const float& gate_sf, const float& up_sf) { + const float s0_gate = scale_a_0 * gate_sf; + const float s1_gate = scale_a_1 * gate_sf; + const float s0_up = scale_a_0 * up_sf; + const float s1_up = scale_a_1 * up_sf; + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const float s0 = (i & 1u) ? s0_up : s0_gate; + const float s1 = (i & 1u) ? s1_up : s1_gate; + final_accum[i*4+0] *= s0; + final_accum[i*4+1] *= s0; + final_accum[i*4+2] *= s1; + final_accum[i*4+3] *= s1; + } + }; + // Chunk i (8 columns) of a >=256-wide warpgroup tile falls in weight-SF block i / 16. + constexpr uint32_t kChunksPerSFGroup = kAccumPerThread / 4 / kNumSFGroupsPerWG; + auto prescale_l2_final = [&](const float& scale_a_0, const float& scale_a_1, + const float& l2_sf, const float& l2_sf_hi) { + const float inv_s0 = kFastMath ? math::fast_rcp(scale_a_0 * l2_sf) : 1.0f / (scale_a_0 * l2_sf); + const float inv_s1 = kFastMath ? math::fast_rcp(scale_a_1 * l2_sf) : 1.0f / (scale_a_1 * l2_sf); + const float inv_s0_hi = kFastMath ? math::fast_rcp(scale_a_0 * l2_sf_hi) : 1.0f / (scale_a_0 * l2_sf_hi); + const float inv_s1_hi = kFastMath ? math::fast_rcp(scale_a_1 * l2_sf_hi) : 1.0f / (scale_a_1 * l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const bool hi = kNumSFGroupsPerWG > 1 and i >= kChunksPerSFGroup; + final_accum[i*4+0] *= hi ? inv_s0_hi : inv_s0; + final_accum[i*4+1] *= hi ? inv_s0_hi : inv_s0; + final_accum[i*4+2] *= hi ? inv_s1_hi : inv_s1; + final_accum[i*4+3] *= hi ? inv_s1_hi : inv_s1; + } + }; + auto postscale_l2_final = [&](const float& scale_a_0, const float& scale_a_1, + const float& l2_sf, const float& l2_sf_hi) { + const float s0 = scale_a_0 * l2_sf; + const float s1 = scale_a_1 * l2_sf; + const float s0_hi = scale_a_0 * l2_sf_hi; + const float s1_hi = scale_a_1 * l2_sf_hi; + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const bool hi = kNumSFGroupsPerWG > 1 and i >= kChunksPerSFGroup; + final_accum[i*4+0] *= hi ? s0_hi : s0; + final_accum[i*4+1] *= hi ? s0_hi : s0; + final_accum[i*4+2] *= hi ? s1_hi : s1; + final_accum[i*4+3] *= hi ? s1_hi : s1; + } + }; + auto rescale_l1_final = [&](const float& prev_scale_a_0, const float& prev_scale_a_1, + const float& prev_gate_sf, const float& prev_up_sf, + const float& scale_a_0, const float& scale_a_1, + const float& gate_sf, const float& up_sf) { + const float r0_gate = (prev_scale_a_0 * prev_gate_sf) * + (kFastMath ? math::fast_rcp(scale_a_0 * gate_sf) : 1.0f / (scale_a_0 * gate_sf)); + const float r1_gate = (prev_scale_a_1 * prev_gate_sf) * + (kFastMath ? math::fast_rcp(scale_a_1 * gate_sf) : 1.0f / (scale_a_1 * gate_sf)); + const float r0_up = (prev_scale_a_0 * prev_up_sf) * + (kFastMath ? math::fast_rcp(scale_a_0 * up_sf) : 1.0f / (scale_a_0 * up_sf)); + const float r1_up = (prev_scale_a_1 * prev_up_sf) * + (kFastMath ? math::fast_rcp(scale_a_1 * up_sf) : 1.0f / (scale_a_1 * up_sf)); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const float r0 = (i & 1u) ? r0_up : r0_gate; + const float r1 = (i & 1u) ? r1_up : r1_gate; + final_accum[i*4+0] *= r0; + final_accum[i*4+1] *= r0; + final_accum[i*4+2] *= r1; + final_accum[i*4+3] *= r1; + } + }; + auto rescale_l2_final = [&](const float& prev_scale_a_0, const float& prev_scale_a_1, + const float& prev_l2_sf, const float& prev_l2_sf_hi, + const float& scale_a_0, const float& scale_a_1, + const float& l2_sf, const float& l2_sf_hi) { + const float r0 = (prev_scale_a_0 * prev_l2_sf) * + (kFastMath ? math::fast_rcp(scale_a_0 * l2_sf) : 1.0f / (scale_a_0 * l2_sf)); + const float r1 = (prev_scale_a_1 * prev_l2_sf) * + (kFastMath ? math::fast_rcp(scale_a_1 * l2_sf) : 1.0f / (scale_a_1 * l2_sf)); + const float r0_hi = (prev_scale_a_0 * prev_l2_sf_hi) * + (kFastMath ? math::fast_rcp(scale_a_0 * l2_sf_hi) : 1.0f / (scale_a_0 * l2_sf_hi)); + const float r1_hi = (prev_scale_a_1 * prev_l2_sf_hi) * + (kFastMath ? math::fast_rcp(scale_a_1 * l2_sf_hi) : 1.0f / (scale_a_1 * l2_sf_hi)); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const bool hi = kNumSFGroupsPerWG > 1 and i >= kChunksPerSFGroup; + final_accum[i*4+0] *= hi ? r0_hi : r0; + final_accum[i*4+1] *= hi ? r0_hi : r0; + final_accum[i*4+2] *= hi ? r1_hi : r1; + final_accum[i*4+3] *= hi ? r1_hi : r1; + } + }; + auto rescale_l2_act_final = [&](const float& prev_scale_a_0, const float& prev_scale_a_1, + const float& scale_a_0, const float& scale_a_1) { + const float r0 = prev_scale_a_0 * (kFastMath ? math::fast_rcp(scale_a_0) : 1.0f / scale_a_0); + const float r1 = prev_scale_a_1 * (kFastMath ? math::fast_rcp(scale_a_1) : 1.0f / scale_a_1); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + final_accum[i*4+0] *= r0; + final_accum[i*4+1] *= r0; + final_accum[i*4+2] *= r1; + final_accum[i*4+3] *= r1; + } + }; + + auto load_weight_sf = [&](const uint32_t& k, float& g, float& u, float& l, float& l_hi) { + if (block_phase == sched::BlockPhase::Linear1) { + g = ptx::ld_shared(smem_weight_sf_wg + k); + u = ptx::ld_shared(smem_weight_sf_wg + kL1SFKBlocks + k); + } else { + l = ptx::ld_shared(smem_weight_sf_wg + k); + if constexpr (kNumSFGroupsPerWG > 1) + l_hi = ptx::ld_shared(smem_weight_sf_wg + kL2SFKBlocks + k); + } + }; + float gate_sf = 0.0f, up_sf = 0.0f, l2_sf = 0.0f, l2_sf_hi = 0.0f; + if (num_k_blocks != 0) + load_weight_sf(0, gate_sf, up_sf, l2_sf, l2_sf_hi); + + if constexpr (kHidden >= 4096) { + float prev_scale_a_0 = 1.0f, prev_scale_a_1 = 1.0f; + float prev_gate_sf = 1.0f, prev_up_sf = 1.0f, prev_l2_sf = 1.0f, prev_l2_sf_hi = 1.0f; + // Hoisted k-loop (single-group k-block: L1, and L2 under the per-128 act SF): stage-static wgmma descriptors, an + // mbarrier.test_wait on the NEXT k-block's full barrier right after the commit and, when it passed, its activation-SF + // loads issued under the drain. The pipeline handshake (full wait per k-block by every thread, empty arrive after the + // drain) is unchanged; the per-64 L2 k-block (two groups, act SF hi/lo) keeps the original loop. + constexpr bool kHoistedMainloop = kReuseAccumAsFinal and kHidden >= 4096; + const bool hoisted = kHoistedMainloop and (block_phase == sched::BlockPhase::Linear1 or kL2ActSFPerBlockK); + if constexpr (kHoistedMainloop) { + if (hoisted) { + const bool is_l1 = block_phase == sched::BlockPhase::Linear1; + // per-warpgroup base descriptors (stage 0 + the warpgroup's A / B offset), opaque so that ptxas keeps them; the + // stage / k advance below is the same low-word add as advance_smem_desc (no field can overflow: smem addresses < 2^18 B) + uint64_t desc_a_wg = mma::sm90::make_smem_desc(smem_a[0] + smem_a_wg_offset, 1).desc_; + uint64_t desc_b_wg = mma::sm90::make_smem_desc(smem_b[0] + smem_b_wg_offset, 1).desc_; + asm volatile("" : "+l"(desc_a_wg)); + asm volatile("" : "+l"(desc_b_wg)); + const auto stage_desc = [](const uint64_t& base, const uint32_t& lo_advance) { + cute::GmmaDescriptor d; + d.desc_ = base; + d.reg32_[0] += lo_advance; + return d.desc_; + }; + // the four ratios prev -> cur of the rescale chain, exactly rescale_l1_final / rescale_l2_final's expressions + const auto ratio = [](const float& prev_a, const float& prev_w, const float& cur_a, const float& cur_w) { + return (prev_a * prev_w) * (kFastMath ? math::fast_rcp(cur_a * cur_w) : 1.0f / (cur_a * cur_w)); + }; + const auto ratios = [&](const float& pa0, const float& pa1, const float& pw0, const float& pw1, + const float& a0, const float& a1, const float& w0, const float& w1, + float& r0, float& r1, float& r2, float& r3) { + r0 = ratio(pa0, pw0, a0, w0); // L1: r0_gate, L2: r0 + r1 = ratio(pa1, pw0, a1, w0); // L1: r1_gate, L2: r1 + r2 = ratio(pa0, pw1, a0, w1); // L1: r0_up, L2: r0_hi + r3 = ratio(pa1, pw1, a1, w1); // L1: r1_up, L2: r1_hi + }; + const auto apply_ratios = [&](const float& r0, const float& r1, const float& r2, const float& r3) { + if (is_l1) { + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const float ra = (i & 1u) ? r2 : r0; + const float rb = (i & 1u) ? r3 : r1; + final_accum[i*4+0] *= ra; + final_accum[i*4+1] *= ra; + final_accum[i*4+2] *= rb; + final_accum[i*4+3] *= rb; + } + } else { + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const bool hi = kNumSFGroupsPerWG > 1 and i >= kChunksPerSFGroup; + final_accum[i*4+0] *= hi ? r2 : r0; + final_accum[i*4+1] *= hi ? r2 : r0; + final_accum[i*4+2] *= hi ? r3 : r1; + final_accum[i*4+3] *= hi ? r3 : r1; + } + } + }; + const auto ld_act_sf = [&](const uint32_t& stage, float& s0, float& s1) { + s0 = ptx::ld_shared(smem_sfa[stage] + row_offset_r0); // L2 per-128: the lo half at offset 0 + s1 = ptx::ld_shared(smem_sfa[stage] + row_offset_r1); + }; + // weight SF of this k-block: (w0, w1) = (gate, up) for L1, (l2, l2_hi) for L2 + float w0 = is_l1 ? gate_sf : l2_sf, w1 = is_l1 ? up_sf : l2_sf_hi; + float pw0 = 1.0f, pw1 = 1.0f; // previous k-block's weight SF + float cur_a0 = 1.0f, cur_a1 = 1.0f; // this k-block's activation SF + float r0 = 1.0f, r1 = 1.0f, r2 = 1.0f, r3 = 1.0f; // ratios previous -> this k-block + bool ready = false; // this k-block's full barrier known complete + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; ++ k_block_idx) { + if (__builtin_expect(not ready, 0)) { + // k == 0, or the early test of this stage did not see the phase complete + full_barriers[stage_idx]->wait(phase); + ld_act_sf(stage_idx, cur_a0, cur_a1); + } + if constexpr (kTrace) { + if (k_block_idx == 0) + trace_math(42, 0); // FIRST_STAGE_READY + } + if (k_block_idx != 0) { + ratios(prev_scale_a_0, prev_scale_a_1, pw0, pw1, cur_a0, cur_a1, w0, w1, r0, r1, r2, r3); + apply_ratios(r0, r1, r2, r3); + } + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + const uint64_t desc_a = stage_desc(desc_a_wg, stage_idx * (SMEM_A_SIZE_PER_STAGE >> 4) + k * (WGMMA::K >> 4)); + const uint64_t desc_b = stage_desc(desc_b_wg, stage_idx * (SMEM_B_SIZE_PER_STAGE >> 4) + k * (WGMMA::K >> 4)); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + + // under the drain: k+1's weight SF, the test of k+1's full barrier and, when it passed, its + // activation SF + const bool has_next_k = k_block_idx + 1 < num_k_blocks; + float next_gate_sf = gate_sf, next_up_sf = up_sf, next_l2_sf = l2_sf, next_l2_sf_hi = l2_sf_hi; + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + const float next_w0 = is_l1 ? next_gate_sf : next_l2_sf, next_w1 = is_l1 ? next_up_sf : next_l2_sf_hi; + float next_a0 = cur_a0, next_a1 = cur_a1; + bool next_ready = false; + const uint32_t next_stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1; + const uint32_t next_phase = phase ^ (next_stage_idx == 0 ? 1u : 0u); + if (has_next_k) { + next_ready = full_barriers[next_stage_idx]->test_wait(next_phase); + if (next_ready) + ld_act_sf(next_stage_idx, next_a0, next_a1); + } + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + // roll the chain state (the postscale after the loop reads prev_* of the last k-block) + prev_scale_a_0 = cur_a0; + prev_scale_a_1 = cur_a1; + pw0 = w0, pw1 = w1; + if (is_l1) { + prev_gate_sf = gate_sf, prev_up_sf = up_sf; + } else { + prev_l2_sf = l2_sf, prev_l2_sf_hi = l2_sf_hi; + } + cur_a0 = next_a0, cur_a1 = next_a1; + w0 = next_w0, w1 = next_w1; + gate_sf = next_gate_sf; + up_sf = next_up_sf; + l2_sf = next_l2_sf; + l2_sf_hi = next_l2_sf_hi; + ready = next_ready; + stage_idx = next_stage_idx; + phase = next_phase; + } + } + } + if (not hoisted) + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; advance_pipeline(k_block_idx)) { + full_barriers[stage_idx]->wait(phase); + if constexpr (kTrace) { + if (k_block_idx == 0) + trace_math(42, 0); // FIRST_STAGE_READY + } + + float scale_a_0_lo, scale_a_1_lo; + float scale_a_0_hi, scale_a_1_hi; + if (block_phase == sched::BlockPhase::Linear1) { + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r1); + } else { + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r1); + if constexpr (kL2ActSFPerBlockK) { + // one act SF per k-block, the hi half does not exist + scale_a_0_hi = scale_a_0_lo, scale_a_1_hi = scale_a_1_lo; + } else { + scale_a_0_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r0); + scale_a_1_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r1); + } + } + + float next_gate_sf = gate_sf, next_up_sf = up_sf, next_l2_sf = l2_sf, next_l2_sf_hi = l2_sf_hi; + const bool has_next_k = k_block_idx + 1 < num_k_blocks; + + if (block_phase == sched::BlockPhase::Linear1) { + if (k_block_idx != 0) + rescale_l1_final(prev_scale_a_0, prev_scale_a_1, + prev_gate_sf, prev_up_sf, + scale_a_0_lo, scale_a_1_lo, + gate_sf, up_sf); + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + prev_scale_a_0 = scale_a_0_lo; + prev_scale_a_1 = scale_a_1_lo; + prev_gate_sf = gate_sf; + prev_up_sf = up_sf; + } else { + if (k_block_idx != 0) + rescale_l2_final(prev_scale_a_0, prev_scale_a_1, prev_l2_sf, prev_l2_sf_hi, + scale_a_0_lo, scale_a_1_lo, l2_sf, l2_sf_hi); + + if constexpr (kL2ActSFPerBlockK) { + // The per-128 act SF: no per-64 split — one full-BLOCK_K + // WGMMA group, one drain, half the fences (the L1 shape). + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + } else { + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + rescale_l2_act_final(scale_a_0_lo, scale_a_1_lo, + scale_a_0_hi, scale_a_1_hi); + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + const uint32_t k_off = (BLOCK_K / 2) + k * WGMMA::K; + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k_off); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k_off); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + } + + release_empty_stage(stage_idx); + + prev_scale_a_0 = scale_a_0_hi; + prev_scale_a_1 = scale_a_1_hi; + prev_l2_sf = l2_sf; + prev_l2_sf_hi = l2_sf_hi; + } + + gate_sf = next_gate_sf; + up_sf = next_up_sf; + l2_sf = next_l2_sf; + l2_sf_hi = next_l2_sf_hi; + } + + if (num_k_blocks != 0) { + if (block_phase == sched::BlockPhase::Linear1) { + postscale_l1_final(prev_scale_a_0, prev_scale_a_1, + prev_gate_sf, prev_up_sf); + } else { + postscale_l2_final(prev_scale_a_0, prev_scale_a_1, prev_l2_sf, prev_l2_sf_hi); + } + } + } else { + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; advance_pipeline(k_block_idx)) { + full_barriers[stage_idx]->wait(phase); + if constexpr (kTrace) { + if (k_block_idx == 0) + trace_math(42, 0); // FIRST_STAGE_READY + } + + float scale_a_0_lo, scale_a_1_lo; + float scale_a_0_hi, scale_a_1_hi; + if (block_phase == sched::BlockPhase::Linear1) { + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r1); + } else { + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r1); + if constexpr (kL2ActSFPerBlockK) { + // one act SF per k-block, the hi half does not exist + scale_a_0_hi = scale_a_0_lo, scale_a_1_hi = scale_a_1_lo; + } else { + scale_a_0_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r0); + scale_a_1_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r1); + } + } + + // Weight SF for this k was prefetched into registers + float next_gate_sf = gate_sf, next_up_sf = up_sf, next_l2_sf = l2_sf, next_l2_sf_hi = l2_sf_hi; + const bool has_next_k = k_block_idx + 1 < num_k_blocks; + + if (block_phase == sched::BlockPhase::Linear1) { + if (k_block_idx != 0) + prescale_l1_final(scale_a_0_lo, scale_a_1_lo, gate_sf, up_sf); + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + postscale_l1_final(scale_a_0_lo, scale_a_1_lo, gate_sf, up_sf); + } else { + if (k_block_idx != 0) + prescale_l2_final(scale_a_0_lo, scale_a_1_lo, l2_sf, l2_sf_hi); + + if constexpr (kL2ActSFPerBlockK) { + // one act SF per k-block -> one full-BLOCK_K WGMMA group (the L1 shape) + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + postscale_l2_final(scale_a_0_lo, scale_a_1_lo, l2_sf, l2_sf_hi); + } else { + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + // Prefetch k+1's weight SF while the WGMMAs drain. + if (has_next_k) + load_weight_sf(k_block_idx + 1, next_gate_sf, next_up_sf, next_l2_sf, next_l2_sf_hi); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + postscale_l2_final(scale_a_0_lo, scale_a_1_lo, l2_sf, l2_sf_hi); + prescale_l2_final(scale_a_0_hi, scale_a_1_hi, l2_sf, l2_sf_hi); + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + const uint32_t k_off = (BLOCK_K / 2) + k * WGMMA::K; + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k_off); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k_off); + WGMMA::wgmma(desc_a, desc_b, final_accum, true); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(final_accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + postscale_l2_final(scale_a_0_hi, scale_a_1_hi, l2_sf, l2_sf_hi); + } + } + + gate_sf = next_gate_sf; + up_sf = next_up_sf; + l2_sf = next_l2_sf; + l2_sf_hi = next_l2_sf_hi; + } + } + } else { + DG_STATIC_ASSERT(kNumSFGroupsPerWG == 1, "The promotion path assumes one weight-SF block per warpgroup tile"); + float accum[kAccumPerThread]; + + for (uint32_t k_block_idx = 0; k_block_idx < num_k_blocks; advance_pipeline(k_block_idx)) { + full_barriers[stage_idx]->wait(phase); + if constexpr (kTrace) { + if (k_block_idx == 0) + trace_math(42, 0); // FIRST_STAGE_READY + } + + // Read SF (must precede warpgroup_arrive) + float scale_a_0_lo, scale_a_1_lo; + float scale_a_0_hi, scale_a_1_hi; // Only used in L2 (per-64 K) + if (block_phase == sched::BlockPhase::Linear1) { + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + row_offset_r1); + } else { + // L2: SFA layout is (K=2, M=BLOCK_M) MN-major; first half SF at offset 0, second at BLOCK_M + scale_a_0_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r0); + scale_a_1_lo = ptx::ld_shared(smem_sfa[stage_idx] + 0 * BLOCK_M + row_offset_r1); + if constexpr (kL2ActSFPerBlockK) { + // one act SF per k-block, the hi half does not exist + scale_a_0_hi = scale_a_0_lo, scale_a_1_hi = scale_a_1_lo; + } else { + scale_a_0_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r0); + scale_a_1_hi = ptx::ld_shared(smem_sfa[stage_idx] + 1 * BLOCK_M + row_offset_r1); + } + } + + // ----- Block (128, 128) weight SF (staged in SMEM per block) ----- + // L1 weight SF: (E, 2*IH/128, H/128) MN-major, N axis = [gate(IH/128), up(IH/128)]; with the gate/up gran-8 interleave a + // logical 128-wide N tile covers 64 gate + 64 up rows of the same original 128-row block, so gate_sf_n = sf_n_block_idx / 2, + // up_sf_n = IH/128 + sf_n_block_idx / 2. L2 weight SF: (E, H/128, IH/128) MN-major, one scalar per logical 128x128 tile. + float gate_sf = 0.0f, up_sf = 0.0f, l2_sf = 0.0f; + if (block_phase == sched::BlockPhase::Linear1) { + gate_sf = ptx::ld_shared(smem_weight_sf_wg + k_block_idx); + up_sf = ptx::ld_shared(smem_weight_sf_wg + kL1SFKBlocks + k_block_idx); + } else { + l2_sf = ptx::ld_shared(smem_weight_sf_wg + k_block_idx); + } + + if (block_phase == sched::BlockPhase::Linear1) { + if constexpr (kSwapABActive) { + auto run_swap_ab_l1 = [&]() { + using SwapWGMMA = typename mma::sm90::FP8MMASelector::type; + constexpr uint32_t kSwapAccum = SwapWGMMA::kNumAccum; + float swap_accum[kSwapAccum]; + + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / SwapWGMMA::K; ++ k) { + auto desc_a = smem_desc_b(stage_idx, smem_b_wg_offset + k * SwapWGMMA::K); + auto desc_b = smem_desc_a(stage_idx, k * SwapWGMMA::K); + SwapWGMMA::wgmma(desc_a, desc_b, swap_accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_wait<0>(); + + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum / 4; ++ i) { + const uint32_t token_0 = i * 8 + col_idx * 2; + const uint32_t token_1 = token_0 + 1; + const float scale_0 = token_0 < valid_m ? + ptx::ld_shared(smem_sfa[stage_idx] + token_0) : 0.0f; + const float scale_1 = token_1 < valid_m ? + ptx::ld_shared(smem_sfa[stage_idx] + token_1) : 0.0f; + final_accum[i * 4 + 0] += scale_0 * gate_sf * swap_accum[i * 4 + 0]; + final_accum[i * 4 + 2] += scale_0 * up_sf * swap_accum[i * 4 + 2]; + final_accum[i * 4 + 1] += scale_1 * gate_sf * swap_accum[i * 4 + 1]; + final_accum[i * 4 + 3] += scale_1 * up_sf * swap_accum[i * 4 + 3]; + } + + release_empty_stage(stage_idx); + }; + + const uint32_t n_swap = ((valid_m + 7u) / 8u) * 8u; + if constexpr (kIntermediateHidden <= 2048) { + if (n_swap <= 8) { + run_swap_ab_l1.template operator()<8>(); + } else if (n_swap <= 16) { + run_swap_ab_l1.template operator()<16>(); + } else if (n_swap <= 32) { + run_swap_ab_l1.template operator()<32>(); + } else { + run_swap_ab_l1.template operator()<64>(); + } + } else { + switch (n_swap) { + case 8: run_swap_ab_l1.template operator()<8>(); break; + case 16: run_swap_ab_l1.template operator()<16>(); break; + case 24: run_swap_ab_l1.template operator()<24>(); break; + case 32: run_swap_ab_l1.template operator()<32>(); break; + case 40: run_swap_ab_l1.template operator()<40>(); break; + case 48: run_swap_ab_l1.template operator()<48>(); break; + case 56: run_swap_ab_l1.template operator()<56>(); break; + default: run_swap_ab_l1.template operator()<64>(); break; + } + } + } else { + // Single per-128 K-block WGMMA group + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + // L1: gate/up alternate at gran=8 along N; each `i` block of 8 + // cols belongs entirely to one of {gate, up}, so .x and .y + // share the same scalar. + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + const float sb = (i & 1u) ? up_sf : gate_sf; + final_accum[i*4+0] += scale_a_0_lo * sb * accum[i*4+0]; + final_accum[i*4+1] += scale_a_0_lo * sb * accum[i*4+1]; + final_accum[i*4+2] += scale_a_1_lo * sb * accum[i*4+2]; + final_accum[i*4+3] += scale_a_1_lo * sb * accum[i*4+3]; + } + } + } else { + if constexpr (kSwapABActive) { + DG_STATIC_ASSERT(kL2ActsSFGranK == 64, + "L2 swapAB assumes per-64 activation scales"); + auto run_swap_ab_l2 = [&]() { + using SwapWGMMA = typename mma::sm90::FP8MMASelector::type; + constexpr uint32_t kSwapAccum = SwapWGMMA::kNumAccum; + float swap_accum[kSwapAccum]; + + auto promote_swap_accum = [&](const uint32_t& sf_group) { + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum / 4; ++ i) { + const uint32_t token_0 = i * 8 + col_idx * 2; + const uint32_t token_1 = token_0 + 1; + const float scale_0 = token_0 < valid_m ? + ptx::ld_shared(smem_sfa[stage_idx] + sf_group * BLOCK_M + token_0) : 0.0f; + const float scale_1 = token_1 < valid_m ? + ptx::ld_shared(smem_sfa[stage_idx] + sf_group * BLOCK_M + token_1) : 0.0f; + final_accum[i * 4 + 0] += scale_0 * l2_sf * swap_accum[i * 4 + 0]; + final_accum[i * 4 + 2] += scale_0 * l2_sf * swap_accum[i * 4 + 2]; + final_accum[i * 4 + 1] += scale_1 * l2_sf * swap_accum[i * 4 + 1]; + final_accum[i * 4 + 3] += scale_1 * l2_sf * swap_accum[i * 4 + 3]; + } + }; + + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / SwapWGMMA::K; ++ k) { + auto desc_a = smem_desc_b(stage_idx, smem_b_wg_offset + k * SwapWGMMA::K); + auto desc_b = smem_desc_a(stage_idx, k * SwapWGMMA::K); + SwapWGMMA::wgmma(desc_a, desc_b, swap_accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_wait<0>(); + promote_swap_accum(0); + + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / SwapWGMMA::K; ++ k) { + const uint32_t k_off = (BLOCK_K / 2) + k * SwapWGMMA::K; + auto desc_a = smem_desc_b(stage_idx, smem_b_wg_offset + k_off); + auto desc_b = smem_desc_a(stage_idx, k_off); + SwapWGMMA::wgmma(desc_a, desc_b, swap_accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kSwapAccum; ++ i) + ptx::warpgroup_fence_operand(swap_accum[i]); + ptx::warpgroup_wait<0>(); + promote_swap_accum(1); + + release_empty_stage(stage_idx); + }; + + const uint32_t n_swap = ((valid_m + 7u) / 8u) * 8u; + if constexpr (kIntermediateHidden <= 2048) { + if (n_swap <= 8) { + run_swap_ab_l2.template operator()<8>(); + } else if (n_swap <= 16) { + run_swap_ab_l2.template operator()<16>(); + } else if (n_swap <= 32) { + run_swap_ab_l2.template operator()<32>(); + } else { + run_swap_ab_l2.template operator()<64>(); + } + } else { + switch (n_swap) { + case 8: run_swap_ab_l2.template operator()<8>(); break; + case 16: run_swap_ab_l2.template operator()<16>(); break; + case 24: run_swap_ab_l2.template operator()<24>(); break; + case 32: run_swap_ab_l2.template operator()<32>(); break; + case 40: run_swap_ab_l2.template operator()<40>(); break; + case 48: run_swap_ab_l2.template operator()<48>(); break; + case 56: run_swap_ab_l2.template operator()<56>(); break; + default: run_swap_ab_l2.template operator()<64>(); break; + } + } + } else if constexpr (kL2ActSFPerBlockK) { + // The per-128 act SF: no per-64 split — one full-BLOCK_K + // WGMMA group, single promotion with the k-block's per-row scale. + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + final_accum[i*4+0] += scale_a_0_lo * l2_sf * accum[i*4+0]; + final_accum[i*4+1] += scale_a_0_lo * l2_sf * accum[i*4+1]; + final_accum[i*4+2] += scale_a_1_lo * l2_sf * accum[i*4+2]; + final_accum[i*4+3] += scale_a_1_lo * l2_sf * accum[i*4+3]; + } + } else { + // L2: split BLOCK_K=128 into two halves (per-64 SFA), each 2 WGMMAs. + // First half: K=0..63, SFA = scale_a_*_lo + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k * WGMMA::K); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k * WGMMA::K); + WGMMA::wgmma(desc_a, desc_b, accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_wait<0>(); + + // L2 first half: single scalar `l2_sf` broadcast across N. + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + final_accum[i*4+0] += scale_a_0_lo * l2_sf * accum[i*4+0]; + final_accum[i*4+1] += scale_a_0_lo * l2_sf * accum[i*4+1]; + final_accum[i*4+2] += scale_a_1_lo * l2_sf * accum[i*4+2]; + final_accum[i*4+3] += scale_a_1_lo * l2_sf * accum[i*4+3]; + } + + // Second half: K=64..127, SFA = scale_a_*_hi + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_arrive(); + #pragma unroll + for (uint32_t k = 0; k < (BLOCK_K / 2) / WGMMA::K; ++ k) { + const uint32_t k_off = (BLOCK_K / 2) + k * WGMMA::K; + auto desc_a = smem_desc_a(stage_idx, smem_a_wg_offset + k_off); + auto desc_b = smem_desc_b(stage_idx, smem_b_wg_offset + k_off); + WGMMA::wgmma(desc_a, desc_b, accum, k); + } + ptx::warpgroup_commit_batch(); + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread; ++ i) ptx::warpgroup_fence_operand(accum[i]); + ptx::warpgroup_wait<0>(); + + release_empty_stage(stage_idx); + + // L2 second half: same broadcast scalar `l2_sf`. + #pragma unroll + for (uint32_t i = 0; i < kAccumPerThread / 4; ++ i) { + final_accum[i*4+0] += scale_a_0_hi * l2_sf * accum[i*4+0]; + final_accum[i*4+1] += scale_a_0_hi * l2_sf * accum[i*4+1]; + final_accum[i*4+2] += scale_a_1_hi * l2_sf * accum[i*4+2]; + final_accum[i*4+3] += scale_a_1_hi * l2_sf * accum[i*4+3]; + } + } + } + } + + } + trace_math(43, 0); // MAINLOOP_END (last warpgroup_wait<0> of the tile done) + // prologue prefetch: peek after the mainloop; the epilogue hides the load + if constexpr (kSFProloguePrefetch) { + if (warp_idx_in_wg == 0 and sf_pf_key == 0) + sf_peek(); + } + + // Skip epilogue when block is past valid M (still must release via empty) + if (row_base >= valid_m) { + if (block_phase == sched::BlockPhase::Linear1) { + if constexpr (kL2ArrivalNeedsFullSync) + sync_tile_math(); + } else { + if constexpr (kL2EpilogueRequiresFullSync) + sync_tile_math(); + } + trace_math(44, 1); // EPILOGUE_END (aux 1: no valid rows, epilogue skipped) + return; + } + + if (block_phase == sched::BlockPhase::Linear1) { + // Ring mode: wait until the previous generation of L2 blocks + // (all N blocks of the earlier pool block mapped to this slot) + // has consumed the L2-acts slot before overwriting it. + if constexpr (not kRingCoversFullPool) { + const auto l2_empty_ptr = workspace.get_l2_empty_count_ptr(ring_block_idx); + const uint32_t l2_empty_target = kNumL2BlockNs * get_ring_wave_idx(pool_block_idx); + // generation 0 has no previous consumers (the counter is zeroed by the last launch's cleanup and no L2 tile of this + // slot can complete before its L1 tiles publish), so the acquire round trip is skipped; later generations wait + if (l2_empty_target > 0) + while (ptx::ld_acq(l2_empty_ptr) != l2_empty_target); + trace_math(45, 0); // RING_L2_SLOT_FREE (L1 epilogue may overwrite the L2-acts slot) + } + + if constexpr (kSwapABActive) { + auto silu = [](float x) -> float { + const float e = kFastMath ? __expf(-x) : expf(-x); + const float sig = kFastMath ? math::fast_rcp(1.0f + e) : 1.0f / (1.0f + e); + return x * sig; + }; + auto clamp_gate = [](float& x) { + if constexpr (kActivationClamp != cute::numeric_limits::infinity()) + x = cute::min(x, kActivationClamp); + }; + auto clamp_up = [](float& x) { + if constexpr (kActivationClamp != cute::numeric_limits::infinity()) + x = cute::min(cute::max(x, -kActivationClamp), kActivationClamp); + }; + + const uint32_t out_col_base = + wg_l1_out_n_offset + warp_idx_in_wg * 8 + row_idx; + // XOR-swizzle of the FP32 staging tile: the 16 B column chunk is permuted by the token index ((c ^ t) & 7), a bijection + // within each row that keeps 4-float vector accesses contiguous; the FP8 tile feeding the TMA store stays linear. + DG_STATIC_ASSERT((L1_OUT_BLOCK_N & (L1_OUT_BLOCK_N - 1)) == 0 and + L1_OUT_BLOCK_N >= 4, + "swapAB FP32 staging swizzle needs a power-of-2 row"); + constexpr uint32_t kSwapL1FP32SwizzleMask = L1_OUT_BLOCK_N / 4 - 1; + auto swap_l1_fp32_idx = [](const uint32_t& token, const uint32_t& col) { + return token * L1_OUT_BLOCK_N + + (col ^ ((token & kSwapL1FP32SwizzleMask) << 2)); + }; + auto store_l1_swap_chunk = [&](const uint32_t& i) { + const uint32_t token_0 = i * 8 + col_idx * 2; + const uint32_t token_1 = token_0 + 1; + if (token_0 < valid_m) { + float g0 = final_accum[i * 4 + 0]; + float u0 = final_accum[i * 4 + 2]; + clamp_gate(g0); + clamp_up(u0); + smem_cd_swap_l1_fp32[swap_l1_fp32_idx(token_0, out_col_base)] = + silu(g0) * u0; + } + if (token_1 < valid_m) { + float g1 = final_accum[i * 4 + 1]; + float u1 = final_accum[i * 4 + 3]; + clamp_gate(g1); + clamp_up(u1); + smem_cd_swap_l1_fp32[swap_l1_fp32_idx(token_1, out_col_base)] = + silu(g1) * u1; + } + }; + + const uint32_t num_swap_token_chunks = (valid_m + 7u) / 8u; + store_l1_swap_chunk(0); + if (valid_m > 8) { + #pragma unroll + for (uint32_t i = 1; i < kSwapABTokenChunks; ++ i) { + if (i < num_swap_token_chunks) + store_l1_swap_chunk(i); + } + } + + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + + for (uint32_t token = epilogue_thread_idx; token < valid_m; token += kNumEpilogueThreads) { + const float weight = *l1_topk_weights_buffer + .get_data_buffer(ring_m_idx + token) + .get_base_ptr(); + float weight_sf_inv; + float amax = 0.0f; + #pragma unroll + for (uint32_t col = 0; col < L1_OUT_BLOCK_N; ++ col) { + const float v = smem_cd_swap_l1_fp32[swap_l1_fp32_idx(token, col)]; + amax = cute::max(amax, cute::abs(v)); + } + amax *= cute::abs(weight); + float2 amax_pair = {amax, amax}; + float2 sf_pair, sf_inv_pair; + sm90_fp8_fused_mega_moe_get_e4m3_sf_and_sf_inv(amax_pair, sf_pair, sf_inv_pair); + const float sf = sf_pair.x; + weight_sf_inv = weight * sf_inv_pair.x; + + auto sf_base_ptr = l2_sf_buffer.get_base_ptr(); + // The L2-activation SF pool is strided by SF_BLOCK_M (= align(BLOCK_M, 128)): the L2 producer reads it as + // sfa_m_idx = ring_block_idx * SF_BLOCK_M and the non-swap L1 writes it so; this path must use the same stride. + const uint32_t token_idx = ring_block_idx * SF_BLOCK_M + token; + sf_base_ptr[n_block_idx * kNumPaddedSFPoolTokens + token_idx] = sf; + + #pragma unroll + for (uint32_t col = 0; col < L1_OUT_BLOCK_N; col += 2) { + // col is even, so col and col+1 share one 16B chunk and + // stay adjacent under the swizzle. + const uint32_t fp32_idx = swap_l1_fp32_idx(token, col); + const float v0 = smem_cd_swap_l1_fp32[fp32_idx + 0] * weight_sf_inv; + const float v1 = smem_cd_swap_l1_fp32[fp32_idx + 1] * weight_sf_inv; + const __nv_fp8x2_e4m3 pair(make_float2(v0, v1)); + auto* ptr = reinterpret_cast( + smem_cd_swap_l1_fp8 + token * L1_OUT_BLOCK_N + col); + *ptr = pair.__x; + } + } + + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + + if (epilogue_wg_n_idx == 0 and warp_idx_in_wg == 0 and cute::elect_one_sync()) { + cute::tma_store_fence(); + cute::SM90_TMA_STORE_2D::copy( + &tensor_map_l1_output, + smem_cd_swap_l1_fp8, + n_block_idx * L1_OUT_BLOCK_N, + ring_m_idx); + cute::tma_store_arrive(); + } + __syncwarp(); + ptx::tma_store_wait<0>(); + + if constexpr (kL2ArrivalCounter and kRingCoversFullPool) { + if (epilogue_wg_n_idx == 0 and warp_idx_in_wg == 0 and cute::elect_one_sync()) { + ptx::red_add_rel( + reinterpret_cast(workspace.get_l2_arrival_mask_ptr(pool_block_idx)), + kWarpgroupSplitN); + } + } else if constexpr (kRingCoversFullPool) { + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + if (epilogue_warp_idx == 0 and cute::elect_one_sync()) { + ptx::red_or_rel_gpu( + workspace.get_l2_arrival_mask_ptr(pool_block_idx), + 1ull << n_block_idx); + } + } else { + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + if (epilogue_warp_idx == 0 and cute::elect_one_sync()) { + ptx::red_add_rel( + workspace.get_l2_full_count_ptr(ring_block_idx), 1u); + ptx::red_add( + workspace.get_l1_empty_count_ptr(ring_block_idx), 1u); + } + } + __syncwarp(); + if constexpr (kL2ArrivalCounter and kRingCoversFullPool) + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + } else { + + // ---------------- L1 EPILOGUE: activation + FP8 quantize + TMA store ---------------- + // `final_accum`: kAccumPerThread/4 chunks of 4 floats per thread = (r0c0, r0c1, r1c0, r1c1); gate and up chunks + // alternate, pair `p` uses chunks 2p and 2p+1 and yields 4 post-SwiGLU floats at output cols p*8 + col_idx*2 + {0,1}. + + constexpr uint32_t kNumPairs = kAccumPerThread / 8; + // Output columns [kL2ActsSFGranK*g, kL2ActsSFGranK*(g+1)) of this warpgroup's tile form L2 act-SF group g + // (per-64: two groups on the 128-column production tile; per-128: one group over the whole tile). + constexpr uint32_t kNumSFGroups = kNumL1OutSFGroups; + constexpr uint32_t kPairsPerSFGroup = kNumPairs / kNumSFGroups; + DG_STATIC_ASSERT(kNumPairs % kNumSFGroups == 0, "L2 act-SF groups must cover whole 8-column pairs"); + DG_STATIC_ASSERT(not kSplitNSharesSF or kNumSFGroups == 1, "Shared-SF split implies one SF group per warpgroup"); + float sf_r0[kNumSFGroups], sf_inv_r0[kNumSFGroups]; + float sf_r1[kNumSFGroups], sf_inv_r1[kNumSFGroups]; + + float swiglu_r0[kNumPairs][2]; + float swiglu_r1[kNumPairs][2]; + float amax_r0[kNumSFGroups], amax_r1[kNumSFGroups]; + #pragma unroll + for (uint32_t g = 0; g < kNumSFGroups; ++ g) + amax_r0[g] = amax_r1[g] = 0.0f; + + auto clamp_gate = [](float& x) { + if constexpr (kActivationClamp != cute::numeric_limits::infinity()) + x = cute::min(x, kActivationClamp); + }; + auto clamp_up = [](float& x) { + if constexpr (kActivationClamp != cute::numeric_limits::infinity()) + x = cute::min(cute::max(x, -kActivationClamp), kActivationClamp); + }; + auto silu = [](float x) -> float { + if constexpr (kFastMath) + return sm90_fp8_fused_mega_moe_silu_ftz_exp(x); + return x * (1.0f / (1.0f + expf(-x))); + }; + + // Rows at or beyond `valid_m` are computed like the others; nothing of theirs is observable (the staging and SF stores + // stay predicated on the row's validity). fmaxf/fabsf agree bit-for-bit with cute::max/abs for non-NaN values (amax + // starts at +0 and never adopts a NaN in either form). + #pragma unroll + for (uint32_t p = 0; p < kNumPairs; ++ p) { + const uint32_t gate = 2 * p, up = 2 * p + 1; + + float g_r0_c0 = final_accum[gate*4 + 0]; + float g_r0_c1 = final_accum[gate*4 + 1]; + float g_r1_c0 = final_accum[gate*4 + 2]; + float g_r1_c1 = final_accum[gate*4 + 3]; + float u_r0_c0 = final_accum[up*4 + 0]; + float u_r0_c1 = final_accum[up*4 + 1]; + float u_r1_c0 = final_accum[up*4 + 2]; + float u_r1_c1 = final_accum[up*4 + 3]; + clamp_gate(g_r0_c0); + clamp_gate(g_r0_c1); + clamp_gate(g_r1_c0); + clamp_gate(g_r1_c1); + clamp_up(u_r0_c0); + clamp_up(u_r0_c1); + clamp_up(u_r1_c0); + clamp_up(u_r1_c1); + + const uint32_t g = p / kPairsPerSFGroup; + swiglu_r0[p][0] = silu(g_r0_c0) * u_r0_c0; + swiglu_r0[p][1] = silu(g_r0_c1) * u_r0_c1; + amax_r0[g] = fmaxf(amax_r0[g], fmaxf(fabsf(swiglu_r0[p][0]), fabsf(swiglu_r0[p][1]))); + swiglu_r1[p][0] = silu(g_r1_c0) * u_r1_c0; + swiglu_r1[p][1] = silu(g_r1_c1) * u_r1_c1; + amax_r1[g] = fmaxf(amax_r1[g], fmaxf(fabsf(swiglu_r1[p][0]), fabsf(swiglu_r1[p][1]))); + } + + // Apply token weight: SwiGLU * topk_weight (single load per row) + const float weight_r0 = valid_r0 ? *l1_topk_weights_buffer + .get_data_buffer(ring_m_idx + row_offset_r0) + .template get_base_ptr() : 0.0f; + const float weight_r1 = valid_r1 ? *l1_topk_weights_buffer + .get_data_buffer(ring_m_idx + row_offset_r1) + .template get_base_ptr() : 0.0f; + #pragma unroll + for (uint32_t p = 0; p < kNumPairs; ++ p) { + swiglu_r0[p][0] *= weight_r0; + swiglu_r0[p][1] *= weight_r0; + swiglu_r1[p][0] *= weight_r1; + swiglu_r1[p][1] *= weight_r1; + } + + #pragma unroll + for (uint32_t g = 0; g < kNumSFGroups; ++ g) { + amax_r0[g] *= cute::abs(weight_r0); + amax_r1[g] *= cute::abs(weight_r1); + + // Reduce amax across the 4 col-lanes that share a row (same `lane_idx >> 2`, different `lane_idx & 3` partition the + // WG-owned columns of the same r_0/r_1): an INTRA-group reduction (xor 2, xor 1); an inter-group one would merge 8 rows. + amax_r0[g] = fmaxf(amax_r0[g], __shfl_xor_sync(0xffffffffu, amax_r0[g], 2)); + amax_r1[g] = fmaxf(amax_r1[g], __shfl_xor_sync(0xffffffffu, amax_r1[g], 2)); + amax_r0[g] = fmaxf(amax_r0[g], __shfl_xor_sync(0xffffffffu, amax_r0[g], 1)); + amax_r1[g] = fmaxf(amax_r1[g], __shfl_xor_sync(0xffffffffu, amax_r1[g], 1)); + + // Phase 2: cross-WG amax. When two N-split warpgroups share one per-64 SF group each has only seen its own + // WG_L1_OUT_BLOCK_N columns; reduce across both through a small smem scratch (upper, currently unused half of the + // CD staging region) so BOTH WGs quantize with the SAME SF. + if constexpr (kSplitNSharesSF) { + float* amax_scratch = reinterpret_cast( + reinterpret_cast(smem_cd_l1) + SMEM_CD_SIZE / 2); + #pragma unroll + for (uint32_t i = epilogue_thread_idx; i < BLOCK_M; i += kNumEpilogueThreads) + amax_scratch[i] = 0.0f; + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + if (col_idx == 0) { + atomicMax(reinterpret_cast(&amax_scratch[r_0]), __float_as_uint(amax_r0[g])); + atomicMax(reinterpret_cast(&amax_scratch[r_1]), __float_as_uint(amax_r1[g])); + } + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + amax_r0[g] = amax_scratch[r_0]; + amax_r1[g] = amax_scratch[r_1]; + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + } + + // Compute SF and inverse SF for each row + float2 amax_pair = {amax_r0[g], amax_r1[g]}; + float2 sf_pair, sf_inv_pair; + sm90_fp8_fused_mega_moe_get_e4m3_sf_and_sf_inv(amax_pair, sf_pair, sf_inv_pair); + sf_r0[g] = sf_pair.x; sf_inv_r0[g] = sf_inv_pair.x; + sf_r1[g] = sf_pair.y; sf_inv_r1[g] = sf_inv_pair.y; + } + trace_epi(50, 0); // EPI_MATH_DONE (SwiGLU, weights, amax, SF; fp8 conversion is fused into the staging loop) + + // Quantize into the staging tile through predicated 16-bit stores. With kL1OutSwizzled the tile is staged in the TMA + // SWIZZLE_128B layout: the 16 B chunk index of a 128 B row is XORed with (row & 7); r_0 and r_1 = r_0 + 8 share + // row & 7 == row_idx, so the XOR term is a per-thread constant. + auto* smem_cd_l1_wg = smem_cd_l1 + smem_cd_l1_wg_offset; + auto l1_stage_offset = [&](const uint32_t& row, const uint32_t& col) { + if constexpr (kL1OutSwizzled) + return row * WG_SMEM_CD_L1_STRIDE_N + (((col >> 4) ^ (row & 7u)) << 4) + (col & 15u); + else + return row * WG_SMEM_CD_L1_STRIDE_N + col; + }; + #pragma unroll + for (uint32_t p = 0; p < kNumPairs; ++ p) { + const uint32_t g = p / kPairsPerSFGroup; + const float v00 = swiglu_r0[p][0] * sf_inv_r0[g]; + const float v01 = swiglu_r0[p][1] * sf_inv_r0[g]; + const float v10 = swiglu_r1[p][0] * sf_inv_r1[g]; + const float v11 = swiglu_r1[p][1] * sf_inv_r1[g]; + const __nv_fp8x2_e4m3 r0_pair(make_float2(v00, v01)); + const __nv_fp8x2_e4m3 r1_pair(make_float2(v10, v11)); + + const uint32_t col = p * 8 + col_idx * 2; + auto* p0 = reinterpret_cast(smem_cd_l1_wg + l1_stage_offset(r_0, col)); + auto* p1 = reinterpret_cast(smem_cd_l1_wg + l1_stage_offset(r_1, col)); + if (valid_r0) + *p0 = r0_pair.__x; + if (valid_r1) + *p1 = r1_pair.__x; + } + + // Write SF as float at `[token, group]` in the L2 acts SF buffer (per-kL2ActsSFGranK layout): only col_idx == 0 writes, + // and in the shared-SF split only the first N-split warpgroup publishes (both own the same slot and rows). + if (col_idx == 0 and (not kSplitNSharesSF or epilogue_wg_n_idx == 0)) { + auto sf_base_ptr = l2_sf_buffer.get_base_ptr(); + // SF buffer is (kNumPaddedSFPoolTokens x kIntermediateHidden/kL2ActsSFGranK), MN-major: + // addr[k_idx * num_padded_sf_pool_tokens + token_idx] + const uint32_t token_r0 = ring_block_idx * SF_BLOCK_M + row_offset_r0; + const uint32_t token_r1 = ring_block_idx * SF_BLOCK_M + row_offset_r1; + #pragma unroll + for (uint32_t g = 0; g < kNumSFGroups; ++ g) { + const uint32_t k_sf_idx = l2_sf_group_base + g; // post-SwiGLU SF group + if (valid_r0) + sf_base_ptr[k_sf_idx * kNumPaddedSFPoolTokens + token_r0] = sf_r0[g]; + if (valid_r1) + sf_base_ptr[k_sf_idx * kNumPaddedSFPoolTokens + token_r1] = sf_r1[g]; + } + } + + // Sync the warpgroup before TMA store. In the shared-tile split + // both N-split warpgroups must finish writing their halves of the + // joint L1-output tile, so sync across all epilogue threads. + if constexpr (kSplitNSharesSF) + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + else + ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + trace_epi(51, 0); // EPI_STAGED (fp8 tile in smem, warpgroup barrier passed) + + // TMA store of the entire tile. Padding rows beyond `valid_m` hold stale FP8 / SF but are never consumed: the L2 tile + // loads them, but its NVLink-scatter epilogue is gated by `m_idx_in_block >= valid_m` and NaN accumulators stay in + // registers (only valid rows are converted to BF16 and STSM'd into smem). + if constexpr (kSplitNSharesSF) { + // One combined store of the joint L1_OUT_BLOCK_N tile, issued + // by the first N-split warpgroup once both halves are staged. + if (epilogue_wg_n_idx == 0 and warp_idx_in_wg == 0 and cute::elect_one_sync()) { + const uint32_t out_n_idx = n_block_idx * L1_OUT_BLOCK_N; + cute::tma_store_fence(); + cute::SM90_TMA_STORE_2D::copy( + &tensor_map_l1_output, + smem_cd_l1, + out_n_idx, + ring_m_idx + row_base); + cute::tma_store_arrive(); + } + } else { + if (warp_idx_in_wg == 0 and cute::elect_one_sync()) { + const uint32_t out_n_idx = n_block_idx * L1_OUT_BLOCK_N + wg_l1_out_n_offset; + cute::tma_store_fence(); + cute::SM90_TMA_STORE_2D::copy( + &tensor_map_l1_output, + smem_cd_l1 + smem_cd_l1_wg_offset, + out_n_idx, + ring_m_idx + row_base); + cute::tma_store_arrive(); + } + } + __syncwarp(); + trace_epi(52, 0); // EPI_STORE_ISSUED + ptx::tma_store_wait<0>(); + trace_epi(53, 0); // EPI_STORE_WAITED + + // Notify L2 that this L1 output (and SF) is ready. Counter mode: independent WG tiles publish arrivals without a + // CTA-wide barrier. Ring mode: the mask-mode barrier flow with counting publishes, so every (m, n) L1 block + // contributes exactly one arrival and one L1-slot release, independent of `valid_m`. + if constexpr (kL2ArrivalCounter and kRingCoversFullPool) { + if constexpr (kSplitNSharesSF) { + // The combined tile counts for both N-split warpgroups; the + // storing warpgroup publishes all kWarpgroupSplitN arrivals + // after its TMA store has drained. + if (epilogue_wg_n_idx == 0 and warp_idx_in_wg == 0 and cute::elect_one_sync()) { + ptx::red_add_rel( + reinterpret_cast(workspace.get_l2_arrival_mask_ptr(pool_block_idx)), + kWarpgroupSplitN); + } + } else if (warp_idx_in_wg == 0 and cute::elect_one_sync()) { + ptx::red_add_rel( + reinterpret_cast(workspace.get_l2_arrival_mask_ptr(pool_block_idx)), 1); + } + } else if constexpr (kRingCoversFullPool) { + sync_tile_math(); + if (tile_lead_warp and cute::elect_one_sync()) { + ptx::red_or_rel_gpu( + workspace.get_l2_arrival_mask_ptr(pool_block_idx), + 1ull << n_block_idx); + } + } else { + sync_tile_math(); + if (tile_lead_warp and cute::elect_one_sync()) { + ptx::red_add_rel( + workspace.get_l2_full_count_ptr(ring_block_idx), 1u); + ptx::red_add( + workspace.get_l1_empty_count_ptr(ring_block_idx), 1u); + } + } + __syncwarp(); + trace_epi(54, 0); // EPI_PUBLISHED + // In the shared-tile split only the first warpgroup issues and drains the combined TMA store; gate the other + // warpgroup so it cannot overwrite the joint smem tile in the next block until that store has drained. + if constexpr (kSplitNSharesSF) + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + } + } else { + // ---------------- L2 EPILOGUE: BF16 cast + NVLink scatter ---------------- + constexpr bool kFullRowL2Stage = kL2StageMode == 2 or kL2StageMode == 4; + constexpr bool kL2MetaPrefetch = kFullRowL2Stage, kL2StageSwizzle = kFullRowL2Stage; + constexpr bool kL2EpiSparse = BLOCK_M == 64 and not kFP8SwapAB and kL2StageMode == 0; + // row-metadata prefetch (kL2MetaPrefetch): lane l < 16 loads the TokenSrcMetadata of the warp's row l (row-pass staging + // emits the warp's 16 rows as h * 8 + g * 4 + rr, see below); each row's store takes it from that lane with shfl + uint32_t pf_meta_rank = 0, pf_meta_token = 0, pf_meta_topk = 0; + if constexpr (kL2MetaPrefetch) { + const uint32_t pf_row_in_wg = warp_idx_in_wg * 16u + (lane_idx & 15u); + if (lane_idx < 16 and row_base + pf_row_in_wg < valid_m) { + const auto* pf_meta = workspace.get_token_src_metadata_ptr(m_idx + row_base + pf_row_in_wg); + pf_meta_rank = pf_meta->rank_idx; + pf_meta_token = pf_meta->token_idx; + pf_meta_topk = pf_meta->topk_idx; + } + } + // Ring mode: the L2-acts input of this slot was fully consumed + // into SMEM by the time any warpgroup reaches the epilogue, so + // release the slot for the L1 epilogue of the next generation + // (one increment per N block). + if constexpr (not kRingCoversFullPool) { + if (tile_lead_warp and cute::elect_one_sync()) { + ptx::red_add( + workspace.get_l2_empty_count_ptr(ring_block_idx), 1u); + } + } + + constexpr uint32_t kNumRowsPerWarp = WG_BLOCK_M / 8; + + const uint32_t row_in_warp_block = lane_idx / 16; // 0 or 1 + const uint32_t lane_in_row = lane_idx % 16; + const uint32_t cols_per_lane = WG_BLOCK_N / 16; + + DG_STATIC_ASSERT(not kSwapABActive or WG_BLOCK_N == 64, + "swapAB BF16 staging swizzle assumes WG_BLOCK_N == 64"); + auto smem_cd_l2_token_idx = [](const uint32_t& token, const uint32_t& col) { + if constexpr (kSwapABActive) { + constexpr uint32_t kSwizzleMask = WG_BLOCK_N / 4 - 1; + return token * WG_BLOCK_N + (col ^ ((token & kSwizzleMask) << 2)); + } else { + return token * WG_BLOCK_N + col; + } + }; + + if constexpr (kSwapABActive) { + auto store_bf16 = [&](const uint32_t& token, const uint32_t& col, float value) { + smem_cd_l2[smem_cd_l2_wg_offset + smem_cd_l2_token_idx(token, col)] = + __float2bfloat16_rn(value); + }; + + auto store_l2_swap_chunk = [&](const uint32_t& i) { + const uint32_t token_0 = i * 8 + col_idx * 2; + const uint32_t token_1 = token_0 + 1; + if (token_0 < valid_m) { + store_bf16(token_0, r_0, final_accum[i * 4 + 0]); + store_bf16(token_0, r_1, final_accum[i * 4 + 2]); + } + if (token_1 < valid_m) { + store_bf16(token_1, r_0, final_accum[i * 4 + 1]); + store_bf16(token_1, r_1, final_accum[i * 4 + 3]); + } + }; + + const uint32_t num_swap_token_chunks = (valid_m + 7u) / 8u; + store_l2_swap_chunk(0); + if (valid_m > 8) { + #pragma unroll + for (uint32_t i = 1; i < kSwapABTokenChunks; ++ i) { + if (i < num_swap_token_chunks) + store_l2_swap_chunk(i); + } + } + // swapAB: the whole warpgroup produced the tile; sync then + // scatter the full WG_BLOCK_N-wide staged rows. + ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + using ScatterVec = std::conditional_t<(WG_BLOCK_N <= 64), uint2, uint4>; + DG_STATIC_ASSERT(cols_per_lane * sizeof(nv_bfloat16) == sizeof(ScatterVec), + "Scatter vector width must match cols_per_lane"); + #pragma unroll + for (uint32_t j = 0; j < kNumRowsPerWarp; ++ j) { + const uint32_t row_in_wg = warp_idx_in_wg * 16 + j * 2 + row_in_warp_block; + const uint32_t m_idx_in_block = row_base + row_in_wg; + if (m_idx_in_block >= valid_m) break; + auto smem_ptr = smem_cd_l2 + smem_cd_l2_wg_offset + + smem_cd_l2_token_idx(row_in_wg, lane_in_row * cols_per_lane); + const auto packed = *reinterpret_cast(smem_ptr); + const auto src_metadata = *workspace.get_token_src_metadata_ptr(m_idx + m_idx_in_block); + const auto dst_token = combine_token_buffer.get_rank_buffer(src_metadata.topk_idx) + .get_data_buffer(src_metadata.token_idx); + auto dst_ptr = math::advance_ptr(dst_token.get_base_ptr(), + (n_idx + wg_n_offset) * sizeof(nv_bfloat16) + lane_in_row * sizeof(ScatterVec)); + *sym_buffer.map(dst_ptr, src_metadata.rank_idx) = packed; + } + } else if constexpr (kL2StageMode != 0) { + // Row-pass staging: every math warp owns a 2 KiB slot and emits its 16 rows in 4 passes over ROWS (mode 1: 8 rows x + // 128 columns per pass, two rows per store; mode 2: 4 rows x 256 columns, one whole row per store). Passes are + // warp-local like the column passes: STS, __syncwarp, LDS + remote store, __syncwarp. + DG_STATIC_ASSERT(WG_BLOCK_M == 64 and WG_BLOCK_N == 256 and kAccumPerThread == 128, + "L2 row-pass staging assumes the 64x256 warpgroup tile"); + constexpr uint32_t kWarpStageElems = 4u * WG_BLOCK_N; // 2 KiB of bf16 per warp + DG_STATIC_ASSERT(kWarpStageElems * 4u == WG_BLOCK_M * (WG_BLOCK_N / kNumL2CDPasses), + "row-pass staging must fit the warpgroup's quarter-width slot"); + nv_bfloat16* warp_stage = smem_cd_l2 + smem_cd_l2_wg_offset + warp_idx_in_wg * kWarpStageElems; + constexpr bool kStageSTSM = kL2StageMode == 4; + // stmatrix.m8n8.x4: matrix i (0..3) of an instruction is chunk 4q+i of the accumulator row group h + // (rows 8h..8h+7 of the warp); lane t provides the address of row t%8 of matrix t/8, and its + // register i holds the pair (row t/4, columns 8(4q+i) + 2(t%4), +1) = exactly the wgmma fragment + const auto stsm_x4 = [](const uint32_t& addr, const uint32_t& r0, const uint32_t& r1, + const uint32_t& r2, const uint32_t& r3) { + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};" + :: "r"(addr), "r"(r0), "r"(r1), "r"(r2), "r"(r3) : "memory"); + }; + const uint32_t warp_stage_addr = static_cast(__cvta_generic_to_shared(warp_stage)); + const uint32_t stsm_mat = lane_idx >> 3; // matrix index inside an x4 instruction + const uint32_t stsm_row = lane_idx & 7u; // row of that matrix this lane addresses + constexpr uint32_t kRowChunks = WG_BLOCK_N / 8; // 32 chunks of 16 B per row + // chunk XOR f(row): without kL2StageSwizzle f = row & 3 (the 4 rows of a pass on distinct banks for the STS); with it + // f = (row & 1) << 2 | row >> 1, so that in a stmatrix.x4 matrix the 4 wanted rows and the 4 junk rows land on 8 distinct + // bank groups (512-byte rows start on bank 0, group of a 16 B chunk = (chunk ^ f) & 7). The LDS.128 of a row reads + // chunk lane_idx: any XOR below 8 keeps a quarter-warp on 8 distinct groups. + const auto swz_f = [](const uint32_t& row) { + if constexpr (kL2StageSwizzle) + return ((row & 1u) << 2) | (row >> 1); + else + return row & 3u; + }; + auto stage_elem = [&](const uint32_t& row, const uint32_t& chunk) { + return row * WG_BLOCK_N + ((chunk ^ swz_f(row)) << 3); + }; + // stmatrix: only matrix rows (row & 4) == 4g belong to pass g; the other four rows of every matrix go to this + // warpgroup's weight-SF slot (256 B, dead between the mainloop's last SF read and the next tile's prologue). Wanted + // address = row * 512 + 16 * ((4q + mat) ^ f(row)) = base + 16 * (mat ^ (row >> 1)) + 64 * (q ^ (row & 1)), quad q is + // base ^ (q << 6); junk slot s = (mat & 1) << 3 | (mat ^ 2 ^ f(row)) & 7: 16 distinct 16 B slots per instruction. + const uint32_t stsm_r = stsm_row & 3u; + const uint32_t stsm_base = kL2StageSwizzle + ? warp_stage_addr + stsm_r * (WG_BLOCK_N * 2u) + ((stsm_mat ^ (stsm_r >> 1)) << 4) + ((stsm_r & 1u) << 6) + : warp_stage_addr + stsm_r * (WG_BLOCK_N * 2u) + ((stsm_mat ^ stsm_r) << 4); + const uint32_t stsm_junk = static_cast(__cvta_generic_to_shared(smem_weight_sf_wg)) + + (kL2StageSwizzle ? ((((stsm_mat & 1u) << 3) | ((stsm_mat ^ 2u ^ swz_f(stsm_r)) & 7u)) << 4) + : ((stsm_mat * 4u + stsm_r) << 4)); + // the slot must exist in the layout, not just in the float count + DG_STATIC_ASSERT(not kStageSTSM or (SMEM_WEIGHT_SF_SIZE >= kNumEpilogueWarpgroups * 256u and + kNumWeightSFFloatsPerWG * sizeof(float) >= 256u), + "stmatrix junk rows need a 256 B weight-SF scratch slot per warpgroup (use L2 stage mode 2)"); + DG_STATIC_ASSERT(not kL2StageSwizzle or WG_BLOCK_N * 2u == 512u, "the staging swizzle assumes 512-byte staging rows"); + #pragma unroll + for (uint32_t pass = 0; pass < 4; ++ pass) { + const uint32_t h = pass >> 1; // accumulator row r_0 / r_1 + const uint32_t g = pass & 1; // which 4 of the 8 accumulator rows (lanes 0-15 / 16-31) + const uint32_t row_in_pass = row_idx & 3u; + if constexpr (kStageSTSM) { + const bool wanted = (stsm_row >> 2) == g; + #pragma unroll + for (uint32_t q = 0; q < kRowChunks / 4; ++ q) { + const uint32_t c0 = 4 * q; + // junk rows always land on the same 16 B slot of the scratch (no per-q offset) + const uint32_t stsm_addr = kL2StageSwizzle ? (stsm_base ^ (q << 6)) : (stsm_base + (q << 6)); + stsm_x4(wanted ? stsm_addr : stsm_junk, + math::cast_into_bf16_and_pack(final_accum[(c0 + 0) * 4 + 2 * h], final_accum[(c0 + 0) * 4 + 2 * h + 1]), + math::cast_into_bf16_and_pack(final_accum[(c0 + 1) * 4 + 2 * h], final_accum[(c0 + 1) * 4 + 2 * h + 1]), + math::cast_into_bf16_and_pack(final_accum[(c0 + 2) * 4 + 2 * h], final_accum[(c0 + 2) * 4 + 2 * h + 1]), + math::cast_into_bf16_and_pack(final_accum[(c0 + 3) * 4 + 2 * h], final_accum[(c0 + 3) * 4 + 2 * h + 1])); + } + } else if ((row_idx >> 2) == g and (h == 0 ? valid_r0 : valid_r1)) { + #pragma unroll + for (uint32_t c = 0; c < kRowChunks; ++ c) { + const uint32_t packed = math::cast_into_bf16_and_pack( + final_accum[c * 4 + 2 * h], final_accum[c * 4 + 2 * h + 1]); + *reinterpret_cast(warp_stage + stage_elem(row_in_pass, c) + col_idx * 2) = packed; + } + } + __syncwarp(); + trace_epi(60, pass); // EPI_CVT_DONE + #pragma unroll + for (uint32_t rr = 0; rr < 4; ++ rr) { + const uint32_t row_in_wg = warp_idx_in_wg * 16 + h * 8 + g * 4 + rr; + const uint32_t m_idx_in_block = row_base + row_in_wg; + if (m_idx_in_block >= valid_m) break; + const auto packed = *reinterpret_cast(warp_stage + stage_elem(rr, lane_idx)); + uint32_t dst_rank, dst_token_idx, dst_topk_idx; + if constexpr (kL2MetaPrefetch) { + // row-metadata prefetch: the row's record sits in lane h * 8 + g * 4 + rr (loaded at the epilogue entry) + dst_rank = __shfl_sync(0xffffffffu, pf_meta_rank, h * 8 + g * 4 + rr); + dst_token_idx = __shfl_sync(0xffffffffu, pf_meta_token, h * 8 + g * 4 + rr); + dst_topk_idx = __shfl_sync(0xffffffffu, pf_meta_topk, h * 8 + g * 4 + rr); + } else { + const auto src_metadata = *workspace.get_token_src_metadata_ptr(m_idx + m_idx_in_block); + dst_rank = src_metadata.rank_idx; + dst_token_idx = src_metadata.token_idx; + dst_topk_idx = src_metadata.topk_idx; + } + const auto dst_token = combine_token_buffer.get_rank_buffer(dst_topk_idx) + .get_data_buffer(dst_token_idx); + auto dst_ptr = math::advance_ptr(dst_token.get_base_ptr(), + (n_idx + wg_n_offset) * sizeof(nv_bfloat16) + lane_idx * sizeof(uint4)); + *sym_buffer.map(dst_ptr, dst_rank) = packed; + } + trace_epi(61, pass); // EPI_SCATTER_DONE + __syncwarp(); + } + } else { + // Non-swap L2: STSM the BF16 tile into SMEM then NVLink-scatter; with kHalfL2CD the tile is emitted in two N-halves + // through a half-width SMEM buffer, otherwise in one full-width pass. + constexpr uint32_t kL2Passes = kNumL2CDPasses; + constexpr uint32_t WG_L2_STAGE_N = WG_BLOCK_N / kL2Passes; + constexpr uint32_t kIterPerPass = (kAccumPerThread / 8) / kL2Passes; + using ScatterVec = std::conditional_t<(WG_L2_STAGE_N <= 64), uint2, uint4>; + const uint32_t l2_cols_per_lane = WG_L2_STAGE_N / 16; + DG_STATIC_ASSERT(l2_cols_per_lane * sizeof(nv_bfloat16) == sizeof(ScatterVec), + "Scatter vector width must match cols_per_lane"); + DG_STATIC_ASSERT(WG_L2_STAGE_N >= 64, + "L2 staging swizzle needs >= 8 16B chunks per row"); + auto l2_stage_idx = [](const uint32_t& row, const uint32_t& col) { + return row * WG_L2_STAGE_N + (col ^ ((row & 7u) << 3)); + }; + // sparse-tile epilogue (kL2EpiSparse): the half-warp's source metadata for all of its rows is loaded once per tile + // BEFORE the conversion / staging stores (predicated on the row bound), so the scatter loop is a register-fed chain + layout::TokenSrcMetadata sparse_meta[kNumRowsPerWarp]; + if constexpr (kL2EpiSparse) { + #pragma unroll + for (uint32_t j = 0; j < kNumRowsPerWarp; ++ j) { + const uint32_t row_in_wg = warp_idx_in_wg * 16 + j * 2 + row_in_warp_block; + const uint32_t m_idx_in_block = row_base + row_in_wg; + sparse_meta[j] = {}; + if (m_idx_in_block < valid_m) + sparse_meta[j] = *workspace.get_token_src_metadata_ptr(m_idx + m_idx_in_block); + } + } + #pragma unroll + for (uint32_t pass = 0; pass < kL2Passes; ++ pass) { + const uint32_t col_base = pass * WG_L2_STAGE_N; + // STSM this pass's columns [col_base, col_base+WG_L2_STAGE_N) + // into the (half-)width buffer at (global col - col_base). + #pragma unroll + for (uint32_t t = 0; t < kIterPerPass; ++ t) { + const uint32_t i = pass * kIterPerPass + t; + const uint32_t chunk_lo = 2 * i, chunk_hi = 2 * i + 1; + auto write_pair = [&](uint32_t row, uint32_t gcol, uint32_t packed) { + *reinterpret_cast( + smem_cd_l2 + smem_cd_l2_wg_offset + + l2_stage_idx(row, gcol - col_base)) = packed; + }; + if (valid_r0) { + const uint32_t r0_lo = math::cast_into_bf16_and_pack( + final_accum[chunk_lo*4 + 0], final_accum[chunk_lo*4 + 1]); + const uint32_t r0_hi = math::cast_into_bf16_and_pack( + final_accum[chunk_hi*4 + 0], final_accum[chunk_hi*4 + 1]); + write_pair(r_0, chunk_lo * 8 + col_idx * 2, r0_lo); + write_pair(r_0, chunk_hi * 8 + col_idx * 2, r0_hi); + } + if (valid_r1) { + const uint32_t r1_lo = math::cast_into_bf16_and_pack( + final_accum[chunk_lo*4 + 2], final_accum[chunk_lo*4 + 3]); + const uint32_t r1_hi = math::cast_into_bf16_and_pack( + final_accum[chunk_hi*4 + 2], final_accum[chunk_hi*4 + 3]); + write_pair(r_1, chunk_lo * 8 + col_idx * 2, r1_lo); + write_pair(r_1, chunk_hi * 8 + col_idx * 2, r1_hi); + } + } + __syncwarp(); + trace_epi(60, pass); // EPI_CVT_DONE (this pass's bf16 columns staged in smem) + // Scatter this pass's columns to remote ranks via NVLink. + #pragma unroll + for (uint32_t j = 0; j < kNumRowsPerWarp; ++ j) { + const uint32_t row_in_wg = warp_idx_in_wg * 16 + j * 2 + row_in_warp_block; + const uint32_t m_idx_in_block = row_base + row_in_wg; + if (m_idx_in_block >= valid_m) break; + const auto packed = *reinterpret_cast( + smem_cd_l2 + smem_cd_l2_wg_offset + + l2_stage_idx(row_in_wg, lane_in_row * l2_cols_per_lane)); + layout::TokenSrcMetadata src_metadata; + if constexpr (kL2EpiSparse) + src_metadata = sparse_meta[j]; + else + src_metadata = *workspace.get_token_src_metadata_ptr(m_idx + m_idx_in_block); + const auto dst_token = combine_token_buffer.get_rank_buffer(src_metadata.topk_idx) + .get_data_buffer(src_metadata.token_idx); + auto dst_ptr = math::advance_ptr(dst_token.get_base_ptr(), + (n_idx + wg_n_offset + col_base) * sizeof(nv_bfloat16) + + lane_in_row * sizeof(ScatterVec)); + *sym_buffer.map(dst_ptr, src_metadata.rank_idx) = packed; + } + trace_epi(61, pass); // EPI_SCATTER_DONE (this pass's remote stores issued) + // Pass 0's scatter must finish reading before pass 1's STSM + // overwrites the shared half-width buffer. + if constexpr (kL2Passes > 1) + __syncwarp(); + } + } + + if constexpr (kL2EpilogueRequiresFullSync) { + sync_tile_math(); + trace_epi(62, 0); // EPI_FULL_SYNC_DONE + } + } + trace_math(44, 0); // EPILOGUE_END + }; + + scheduler.for_each_block_replay(tile_table, + [&](const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + process_math_block( + std::integral_constant{}, + local_expert_idx, num_k_blocks, m_block_idx, n_block_idx); + }, + [&](const uint32_t& local_expert_idx, + const uint32_t& num_k_blocks, + const uint32_t& m_block_idx, const uint32_t& n_block_idx) { + process_math_block( + std::integral_constant{}, + local_expert_idx, num_k_blocks, m_block_idx, n_block_idx); + }); + trace_math(19, trace_tile); // TILES_DONE (aux = number of tiles this CTA ran) + // Combine staging: whole rows (kNumChunks == 1) when three whole-row slots per warp fit the smem in front of the + // barriers, else the row is combined in chunks + constexpr uint32_t kNumHiddenBytes = kHidden * sizeof(nv_bfloat16); + constexpr uint32_t kNumElemsPerUint4 = sizeof(uint4) / sizeof(nv_bfloat162); + + // the ladder itself lives in `layout/sm90_fused_mega_moe.cuh`, where the host reads it too + constexpr uint32_t kNumChunkSlots = layout::kSM90FusedCombineChunkSlots; + constexpr uint32_t kNumChunks = layout::get_sm90_fused_moe_combine_num_chunks( + kHidden, kNumCombineWarps, SMEM_BEFORE_BARRIER_SIZE, kSplitMNWarpgroups); + constexpr uint32_t kNumChunkBytes = kNumHiddenBytes / kNumChunks; + constexpr uint32_t kNumChunkUint4 = kNumChunkBytes / sizeof(uint4); + constexpr uint32_t kNumUint4PerLane = kNumChunkUint4 / 32; + + // Sliced post-barrier combine of the decode topology (whole-row slots, one 512 B slice per warp; each element still sums + // the slots in ascending order in fp32 -> bitwise): 2 = dynamic work-list form, 1 = static form (one CTA barrier, every + // warp walks the same remaining set); 0 = the plain per-warp loop of the prefill topology and of the chunked combine + constexpr uint32_t kCombineSlicedForm = 2u; + // staging capacity (the static_asserts of the two forms below): the dynamic form stages one item's kNumTopk slices in the + // warp's whole-row stage-0 buffer, the static form kNumTopk + 1 slices per warp in the two whole-row stages + constexpr bool kCombineSlicedFits = kCombineSlicedForm == 2 ? (kNumTopk <= kNumCombineWarps) + : (kNumTopk + 1 <= 2 * kNumCombineWarps); + constexpr uint32_t kCombineSliced = + (kDecodeTopology and kNumChunks == 1 and kHidden % (kNumCombineWarps * 256) == 0 and kCombineSlicedFits) ? kCombineSlicedForm : 0u; + // dynamic sliced combine: the work list lives in the SFA stage area once that is dead: word 0 = item counter, words 1.. = one + // entry per (warp, stripe): 0xffffffff unknown (owner not yet past its wait-combine), 0xfffffffe nothing left, else the token + // index. Initialised below, behind a CTA-wide sync of the math warps + constexpr uint32_t kNumWLEntries = kNumCombineWarps * kNumECTokenStripes; + auto wl_words = reinterpret_cast(sf_start_ptr); + DG_STATIC_ASSERT(kCombineSliced != 2 or (SMEM_SFA_SIZE_PER_STAGE > 0 and (1u + kNumWLEntries) * 4u <= kNumStages * SMEM_SFA_SIZE_PER_STAGE), + "the sliced-combine work list needs the block-scaled SFA stage area"); + // PDL: this warpgroup's last tile is published (L2 scatter issued, arrivals released); let the runtime schedule the + // next grid. Its CTAs only get an SM once this CTA exits and may not touch memory before their own wait, so the + // combine and the workspace cleanup below are never observed early. + if constexpr (kPDL) + cudaTriggerProgrammaticLaunchCompletion(); + + DG_STATIC_ASSERT(kHidden % kNumChunks == 0, "Hidden must be divisible by number of chunks"); + DG_STATIC_ASSERT(kNumChunkSlots * kNumCombineWarps * kNumHiddenBytes / kNumChunks <= SMEM_BEFORE_BARRIER_SIZE, "Hidden is too large"); + DG_STATIC_ASSERT(kNumChunkBytes % 16 == 0, "Combine chunk must be TMA-aligned (16 bytes)"); + DG_STATIC_ASSERT(kNumChunkBytes % sizeof(uint4) == 0, "Combine chunk must be divisible by 16 bytes"); + DG_STATIC_ASSERT(kNumChunkUint4 % 32 == 0, "Combine chunk must be a multiple of 32 16-byte elements"); + DG_STATIC_ASSERT(kNumTopk <= 32, "Top-k must fit in a single warp"); + + const auto combine_load_buffer = utils::PatternVisitor([&](const uint32_t& i) { + return math::advance_ptr(smem_buffer, (epilogue_warp_idx + i * kNumCombineWarps) * kNumChunkBytes); + }); + const auto combine_store_buffer = math::advance_ptr( + smem_buffer, (epilogue_warp_idx + kNumCombineWarps * 2) * kNumChunkBytes); + + auto combine_load_barriers = utils::PatternVisitor([&](const uint32_t& i) { + return combine_barriers[i + epilogue_warp_idx * 2]; + }); + + uint32_t combine_phase = 0; + uint32_t load_stage_idx = 0; + + // Combine of one token (also used by the early-combine wait): the valid slots' rows are TMA-loaded chunk by chunk into + // the two staging stages, summed in fp32 in ascending slot order (ptx::accumulate) and stored once through smem + TMA. + // `abort_check()` (warp-uniform) is consulted before every slot load; when it fires the loads in flight are drained, + // nothing is stored and false is returned (the token is combined again later: bit-identical, same fixed-order sum). + const auto combine_one_token = [&](const uint32_t& token_idx, const uint32_t& total_mask, const auto& abort_check) -> bool { + for (uint32_t chunk = 0; chunk < kNumChunks; ++ chunk) { + const uint32_t chunk_byte_offset = chunk * kNumChunkBytes; + uint32_t mask = total_mask; + const auto move_mask_and_load = [&](const uint32_t& i) { + if (mask) { + const uint32_t slot_idx = __ffs(mask) - 1; + mask ^= 1 << slot_idx; + if (cute::elect_one_sync()) { + const auto src_ptr = math::advance_ptr( + combine_token_buffer.get_rank_buffer(slot_idx) + .get_data_buffer(token_idx).get_base_ptr(), + chunk_byte_offset); + ptx::tma_load_1d(combine_load_buffer[i], src_ptr, combine_load_barriers[i], kNumChunkBytes); + ptx::mbarrier_arrive_and_set_tx(combine_load_barriers[i], kNumChunkBytes); + } + __syncwarp(); + return true; + } + return false; + }; + bool do_reduce = move_mask_and_load(load_stage_idx); + float2 reduced[kNumUint4PerLane * kNumElemsPerUint4] = {}; + while (do_reduce) { + if (abort_check()) { + // drain the load in flight (stage load_stage_idx), keep the stage / parity bookkeeping consistent + combine_load_barriers[load_stage_idx]->wait(combine_phase); + combine_phase ^= load_stage_idx; + load_stage_idx ^= 1; + return false; + } + do_reduce = move_mask_and_load(load_stage_idx ^ 1); + combine_load_barriers[load_stage_idx]->wait(combine_phase); + #pragma unroll + for (uint32_t j = 0; j < kNumUint4PerLane; ++ j) { + const auto uint4_values = combine_load_buffer[load_stage_idx][j * 32 + lane_idx]; + const auto bf16_values = reinterpret_cast(&uint4_values); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + ptx::accumulate(reduced[j * kNumElemsPerUint4 + l], bf16_values[l]); + } + combine_phase ^= load_stage_idx; + load_stage_idx ^= 1; + } + #pragma unroll + for (uint32_t j = 0; j < kNumUint4PerLane; ++ j) { + uint4 casted; + auto casted_bf16 = reinterpret_cast(&casted); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + casted_bf16[l] = __float22bfloat162_rn(reduced[j * kNumElemsPerUint4 + l]); + if (j == 0) { + ptx::tma_store_wait<0>(); + __syncwarp(); + } + ptx::st_shared(combine_store_buffer + j * 32 + lane_idx, + casted.x, casted.y, casted.z, casted.w); + } + __syncwarp(); + if (cute::elect_one_sync()) { + cute::tma_store_fence(); + ptx::tma_store_1d( + math::advance_ptr(y, static_cast(token_idx) * kNumHiddenBytes + chunk_byte_offset), + combine_store_buffer, kNumChunkBytes); + cute::tma_store_arrive(); + } + __syncwarp(); + } + return true; + }; + + uint32_t ec_wait_done = 0; // early-combine wait: bit j = this warp's token of stripe j is already combined + // ---------------- COMBINE ---------------- + // early combine: every math warp of the CTA is past its last epilogue -> the remaining L2 tiles of the table (the last one has + // no successor tile for the B loader's signal) are complete + if constexpr (kEarlyCombine) { + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + if (epilogue_thread_idx == 0) + ptx::st_release_cta_shared(smem_ec_done, 1u); + } + if constexpr (kCombineSliced == 2) { + // the SFA stage area is dead only once EVERY math warp is past its last mainloop (a warpgroup may trail the other by up + // to kNumStages k-blocks, still reading stage SFA): the work list is initialised behind a CTA-wide sync of the math warps; + // the barriers before the combine (arrive sync of the wait-combine, or the tag-2 barrier's own syncs) order these stores + // before any warp's post / grab + if constexpr (not kEarlyCombine) + sync_tile_math(); + if (epilogue_warp_idx < kNumCombineWarps and lane_idx < kNumECTokenStripes) + ptx::st_shared(wl_words + 1 + epilogue_warp_idx * kNumECTokenStripes + lane_idx, 0xffffffffu); + if (epilogue_thread_idx == 0) + ptx::st_shared(wl_words, 0u); + } + if constexpr (kEarlyCombineMathWait) { + // early-combine wait: the same all-rank barrier as comm::nvlink_barrier (grid sync, SM0's sys signal, grid sync), but + // while a CTA waits for a grid sync its math warps combine their own tokens whose top-k experts have all published + // their done flag. A token in flight when the sync completes is dropped (nothing stored) and redone by the ordinary + // combine below; tokens combined here are remembered per warp (one bit per stripe) and skipped below. + static constexpr uint32_t kFinishSumTag = 0x80000000u; + const auto count_ptr = workspace.template get_grid_sync_count_ptr(); + const auto grid_flipped = [&](const uint32_t& old_value) { + uint32_t v = 0; + if (lane_idx == 0) + v = ptx::ld_acq(count_ptr); + v = __shfl_sync(0xffffffffu, v, 0); + return ((v ^ old_value) & kFinishSumTag) != 0; + }; + // this warp's tokens: stripe j <-> token j * kNumSMs * kNumCombineWarps + sm_idx * kNumCombineWarps + warp + const auto my_token = [&](const uint32_t& j) { return j * kNumSMs * kNumCombineWarps + sm_idx * kNumCombineWarps + epilogue_warp_idx; }; + uint32_t wc_pending = 0; // bit j: stripe j still to do (and not combined by a dispatch receiver) + #pragma unroll + for (uint32_t j = 0; j < kNumECTokenStripes; ++ j) { + bool p = my_token(j) < num_tokens; + wc_pending |= p ? (1u << j) : 0u; + } + // token j: relaxed check of its experts' flags, then the acquire loads (8 lanes) once they are all set + const auto token_flags_ready = [&](const uint32_t& token_idx, uint32_t& total_mask) { + const int slot = lane_idx < kNumTopk ? + static_cast(__ldg(input_topk_idx_buffer.get_base_ptr() + token_idx * kNumTopk + lane_idx)) : -1; + total_mask = __ballot_sync(0xffffffffu, slot >= 0); + const auto flag_ptr = workspace.get_expert_done_flag_ptr(slot >= 0 ? static_cast(slot) : 0u); + const bool set = slot < 0 or ptx::ld_relaxed_sys(flag_ptr) != 0; + if (not __all_sync(0xffffffffu, set)) + return false; + if (slot >= 0) { + const uint32_t v = ptx::ld_acq_sys(flag_ptr); + asm volatile("" :: "r"(v)); + } + __syncwarp(); + return true; + }; + uint32_t num_wait_combined = 0; + { + // arrive (release covers every warp's scatter stores: bar.sync first), then warp 0 runs the whole barrier + // protocol while warps 1.. combine their ready tokens; thread 0 publishes the completion in smem + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + if (epilogue_warp_idx == 0) { + uint32_t old_value = 0; + if (lane_idx == 0) + old_value = ptx::atomic_add_rel(count_ptr, sm_idx == 0 ? (kFinishSumTag - (kNumSMs - 1)) : 1); + old_value = __shfl_sync(0xffffffffu, old_value, 0); + while (not grid_flipped(old_value)); + if (sm_idx == 0) { + auto* counter_ptr = workspace.get_nvl_barrier_counter_ptr(); + const auto status = (*counter_ptr) & 3; + const auto signal_phase = status & 1, signal_sign = status >> 1; + auto* signal_ptr = workspace.get_nvl_barrier_signal_ptr(signal_phase); + if (lane_idx < kNumRanks) + ptx::red_add_rel_sys(sym_buffer.map(signal_ptr, lane_idx), signal_sign ? -1 : 1); + __syncwarp(); + if (lane_idx == 0) { + ptx::red_add(counter_ptr, 1); + const int target = signal_sign ? 0 : static_cast(kNumRanks); + const auto start_clock = clock64(); + // (the timeout trap sits after the loop: a trap inside a loop body of the math region makes + // ptxas cap the region at the 168-register launch bound) + bool timed_out = false; + while (ptx::ld_acq_sys(signal_ptr) != target) { + if (clock64() - start_clock >= comm::kNumTimeoutCycles) { + timed_out = true; + break; + } + } + DG_TRAP_ONLY_DEVICE_ASSERT(not timed_out); + } + __syncwarp(); + } + if (lane_idx == 0) + old_value = ptx::atomic_add_rel(count_ptr, sm_idx == 0 ? (kFinishSumTag - (kNumSMs - 1)) : 1); + old_value = __shfl_sync(0xffffffffu, old_value, 0); + while (not grid_flipped(old_value)); + // every rank's scatter is complete and visible: release the other warps and the dispatch warps + if (lane_idx == 0) + ptx::st_release_cta_shared(smem_ec_stop, 1u); + } else { + const auto never = [&]() { return false; }; + while (ptx::ld_acquire_cta_shared(smem_ec_stop) == 0) { + bool progressed = false; + uint32_t bits = wc_pending; + while (bits != 0) { + const uint32_t j = __ffs(bits) - 1; + bits &= bits - 1; + const uint32_t token_idx = my_token(j); + uint32_t total_mask; + if (not token_flags_ready(token_idx, total_mask)) + continue; + combine_one_token(token_idx, total_mask, never); + wc_pending &= ~(1u << j); + ++ num_wait_combined; + progressed = true; + } + if (not progressed) + __nanosleep(500); + } + } + trace_math(25, num_wait_combined | (static_cast(__popc(wc_pending)) << 16)); // WAIT_COMBINE_DONE + ec_wait_done = ~wc_pending; + } + } else { + // NVLink barrier first: signals remote ranks that this rank's GEMM + // outputs (NVLink scatter targets) are fully written. + comm::nvlink_barrier( + workspace, sym_buffer, sm_idx, epilogue_thread_idx, + [&]() { ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); } + ); + } + + // Sync with dispatch (paired with dispatch's pre-cleanup sync) so that + // dispatch may now safely clean workspace state. (early combine: the dispatch warps wait on the stop flag instead) + if constexpr (not kEarlyCombine) + ptx::sync_unaligned(kNumDispatchThreads + kNumEpilogueThreads, kDispatchWithEpilogueBarrierIdx); + trace_math(20, 0); // COMBINE_START (all-rank barrier passed; combine runs on the math warps) + + if (epilogue_warp_idx >= kNumCombineWarps) + return; + + DG_TRAP_ONLY_DEVICE_ASSERT(kNumChunkSlots * kNumCombineWarps * kNumChunkBytes <= static_cast( + reinterpret_cast(barrier_start_ptr) - smem_buffer)); + + uint32_t num_combined = 0, num_ec_skipped = 0; // trace (COMBINE_END aux) + if constexpr (kCombineSliced == 2) { + // dynamic sliced combine (see the template parameter). This warp is free: post its leftover tokens, then take items. + constexpr uint32_t kSliceBytes = kNumHiddenBytes / kNumCombineWarps; + constexpr uint32_t kSliceUint4PerLane = kSliceBytes / sizeof(uint4) / 32; + DG_STATIC_ASSERT(kNumChunks == 1 and kNumHiddenBytes % kNumCombineWarps == 0 and kSliceBytes % (32 * sizeof(uint4)) == 0, + "dynamic sliced combine: whole-row slots, hidden slices of a multiple of 512 bytes per warp"); + DG_STATIC_ASSERT(kNumTopk * kSliceBytes <= kNumChunkBytes, "the slot slices of one item fit the warp's stage-0 buffer"); + const auto wl_entries = wl_words + 1; + if (lane_idx < kNumECTokenStripes) { + const uint32_t token_idx = lane_idx * kNumSMs * kNumCombineWarps + sm_idx * kNumCombineWarps + epilogue_warp_idx; + bool pending = token_idx < num_tokens; + if constexpr (kEarlyCombine) { + pending = pending and not ((ec_wait_done >> lane_idx) & 1u); + if (not pending and token_idx < num_tokens) + ++ num_ec_skipped; + } + ptx::st_shared(wl_entries + epilogue_warp_idx * kNumECTokenStripes + lane_idx, pending ? token_idx : 0xfffffffeu); + } + __syncwarp(); + const auto sl_load_buffer = [&](const uint32_t& s) { return math::advance_ptr(combine_load_buffer[0], s * kSliceBytes); }; + const auto sl_store_buffer = combine_store_buffer; + constexpr uint32_t kNumWLItems = kNumWLEntries * kNumCombineWarps; // (warp, stripe) x slice + #pragma unroll 1 + while (true) { + uint32_t idx = 0; + if (lane_idx == 0) + idx = atomicAdd(wl_words, 1u); + idx = __shfl_sync(0xffffffffu, idx, 0); + if (idx >= kNumWLItems) + break; + const uint32_t slice = idx % kNumCombineWarps, entry = idx / kNumCombineWarps; + uint32_t token_idx; + while ((token_idx = ptx::ld_volatile_shared(wl_entries + entry)) == 0xffffffffu) + __nanosleep(64); + if (token_idx == 0xfffffffeu) + continue; + ++ num_combined; + const int slot_of_lane = lane_idx < kNumTopk ? + static_cast(__ldg(input_topk_idx_buffer.get_base_ptr() + token_idx * kNumTopk + lane_idx)) : -1; + const uint32_t mask = __ballot_sync(0xffffffffu, slot_of_lane >= 0); + float2 reduced[kSliceUint4PerLane * kNumElemsPerUint4] = {}; + if (mask) { + const uint32_t num_slots = __popc(mask); + if (cute::elect_one_sync()) { + uint32_t issue_mask = mask; + uint32_t s = 0; + while (issue_mask) { + const uint32_t slot_idx = __ffs(issue_mask) - 1; + issue_mask ^= 1u << slot_idx; + const auto src_ptr = math::advance_ptr( + combine_token_buffer.get_rank_buffer(slot_idx).get_data_buffer(token_idx).get_base_ptr(), slice * kSliceBytes); + ptx::tma_load_1d(sl_load_buffer(s), src_ptr, combine_load_barriers[load_stage_idx], kSliceBytes); + ++ s; + } + ptx::mbarrier_arrive_and_set_tx(combine_load_barriers[load_stage_idx], num_slots * kSliceBytes); + } + __syncwarp(); + combine_load_barriers[load_stage_idx]->wait(combine_phase); + #pragma unroll 1 + for (uint32_t s = 0; s < num_slots; ++ s) { + #pragma unroll + for (uint32_t j = 0; j < kSliceUint4PerLane; ++ j) { + const auto uint4_values = sl_load_buffer(s)[j * 32 + lane_idx]; + const auto bf16_values = reinterpret_cast(&uint4_values); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + ptx::accumulate(reduced[j * kNumElemsPerUint4 + l], bf16_values[l]); + } + } + combine_phase ^= load_stage_idx; + load_stage_idx ^= 1; + } + #pragma unroll + for (uint32_t j = 0; j < kSliceUint4PerLane; ++ j) { + uint4 casted; + auto casted_bf16 = reinterpret_cast(&casted); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + casted_bf16[l] = __float22bfloat162_rn(reduced[j * kNumElemsPerUint4 + l]); + if (j == 0) { + ptx::tma_store_wait<0>(); + __syncwarp(); + } + ptx::st_shared(sl_store_buffer + j * 32 + lane_idx, casted.x, casted.y, casted.z, casted.w); + } + __syncwarp(); + if (cute::elect_one_sync()) { + cute::tma_store_fence(); + ptx::tma_store_1d(math::advance_ptr(y, static_cast(token_idx) * kNumHiddenBytes + slice * kSliceBytes), + sl_store_buffer, kSliceBytes); + cute::tma_store_arrive(); + } + __syncwarp(); + } + } else if constexpr (kCombineSliced == 1) { + // static sliced combine: the CTA's remaining tokens, one at a time by all warps; warp w reduces and stores hidden slice w. Layout in the + // two whole-row load stages' region (free once every warp is past the barrier below; the in-flight stores read the + // store buffers behind it): warp w = kNumTopk staging slices + its store slice at w x (kNumTopk + 1) slices. + constexpr uint32_t kSliceBytes = kNumHiddenBytes / kNumCombineWarps; + constexpr uint32_t kSliceUint4PerLane = kSliceBytes / sizeof(uint4) / 32; + DG_STATIC_ASSERT(kNumChunks == 1 and kNumHiddenBytes % kNumCombineWarps == 0 and kSliceBytes % (32 * sizeof(uint4)) == 0, + "static sliced combine: whole-row slots, hidden slices of a multiple of 512 bytes per warp"); + DG_STATIC_ASSERT(kNumCombineWarps * (kNumTopk + 1) * kSliceBytes <= 2 * kNumCombineWarps * kNumChunkBytes, + "the sliced staging must fit the whole-row load stages' region"); + const auto sl_load_buffer = [&](const uint32_t& s) { + return math::advance_ptr(smem_buffer, (epilogue_warp_idx * (kNumTopk + 1) + s) * kSliceBytes); + }; + const auto sl_store_buffer = math::advance_ptr(smem_buffer, (epilogue_warp_idx * (kNumTopk + 1) + kNumTopk) * kSliceBytes); + // tokens already combined during the wait (per-warp ec_wait_done) -> bit (stripe x warps + warp) of the CTA-wide bitmap + // (the claim words of the early-combine control block), so every warp walks the same remaining set + if constexpr (kEarlyCombine) { + DG_STATIC_ASSERT(kNumECTokenStripes * kNumCombineWarps <= 8u * 32u, + "the done bitmap needs stripes x warps <= 256 bits (words 8..15 of the early-combine control block)"); + if (lane_idx < kNumECTokenStripes and ((ec_wait_done >> lane_idx) & 1u)) { + const uint32_t bit = lane_idx * kNumCombineWarps + epilogue_warp_idx; + atomicOr(smem_ec_claim + (bit >> 5), 1u << (bit & 31u)); + } + ptx::sync_aligned(kNumEpilogueThreads, kEpilogueFullBarrierIdx); + } + #pragma unroll 1 + for (uint32_t token_base = sm_idx * kNumCombineWarps, stripe = 0; token_base < num_tokens; + token_base += kNumSMs * kNumCombineWarps, ++ stripe) { + #pragma unroll 1 + for (uint32_t w = 0; w < kNumCombineWarps; ++ w) { + const uint32_t token_idx = token_base + w; + if (token_idx >= num_tokens) + break; + if constexpr (kEarlyCombine) { + const uint32_t bit = stripe * kNumCombineWarps + w; + if ((ptx::ld_shared(smem_ec_claim + (bit >> 5)) >> (bit & 31u)) & 1u) { + ++ num_ec_skipped; + continue; + } + } + if constexpr (kTrace) + ++ num_combined; + const int slot_of_lane = lane_idx < kNumTopk ? + static_cast(__ldg(input_topk_idx_buffer.get_base_ptr() + token_idx * kNumTopk + lane_idx)) : -1; + uint32_t mask = __ballot_sync(0xffffffffu, slot_of_lane >= 0); + float2 reduced[kSliceUint4PerLane * kNumElemsPerUint4] = {}; + if (mask) { + const uint32_t num_slots = __popc(mask); + // all valid slots' slices in flight together, one arrive with the total byte count + if (cute::elect_one_sync()) { + uint32_t issue_mask = mask; + uint32_t s = 0; + while (issue_mask) { + const uint32_t slot_idx = __ffs(issue_mask) - 1; + issue_mask ^= 1u << slot_idx; + const auto src_ptr = math::advance_ptr( + combine_token_buffer.get_rank_buffer(slot_idx).get_data_buffer(token_idx).get_base_ptr(), + epilogue_warp_idx * kSliceBytes); + ptx::tma_load_1d(sl_load_buffer(s), src_ptr, combine_load_barriers[load_stage_idx], kSliceBytes); + ++ s; + } + ptx::mbarrier_arrive_and_set_tx(combine_load_barriers[load_stage_idx], num_slots * kSliceBytes); + } + __syncwarp(); + combine_load_barriers[load_stage_idx]->wait(combine_phase); + // ascending slot order per element, as the whole-row combine + #pragma unroll 1 + for (uint32_t s = 0; s < num_slots; ++ s) { + #pragma unroll + for (uint32_t j = 0; j < kSliceUint4PerLane; ++ j) { + const auto uint4_values = sl_load_buffer(s)[j * 32 + lane_idx]; + const auto bf16_values = reinterpret_cast(&uint4_values); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + ptx::accumulate(reduced[j * kNumElemsPerUint4 + l], bf16_values[l]); + } + } + // one wait on the current stage's barrier: same bookkeeping as the whole-row loop + combine_phase ^= load_stage_idx; + load_stage_idx ^= 1; + } + #pragma unroll + for (uint32_t j = 0; j < kSliceUint4PerLane; ++ j) { + uint4 casted; + auto casted_bf16 = reinterpret_cast(&casted); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + casted_bf16[l] = __float22bfloat162_rn(reduced[j * kNumElemsPerUint4 + l]); + if (j == 0) { + ptx::tma_store_wait<0>(); + __syncwarp(); + } + ptx::st_shared(sl_store_buffer + j * 32 + lane_idx, casted.x, casted.y, casted.z, casted.w); + } + __syncwarp(); + if (cute::elect_one_sync()) { + cute::tma_store_fence(); + ptx::tma_store_1d( + math::advance_ptr(y, static_cast(token_idx) * kNumHiddenBytes + epilogue_warp_idx * kSliceBytes), + sl_store_buffer, kSliceBytes); + cute::tma_store_arrive(); + } + __syncwarp(); + } + } + } else + for (uint32_t token_idx = sm_idx * kNumCombineWarps + epilogue_warp_idx; + token_idx < num_tokens; + token_idx += kNumSMs * kNumCombineWarps) { + // early combine: this warp combined the token while waiting in the barrier + if constexpr (kEarlyCombine) { + const uint32_t stripe = token_idx / (kNumSMs * kNumCombineWarps); + if ((ec_wait_done >> stripe) & 1u) { + ++ num_ec_skipped; + continue; + } + } + if constexpr (kTrace) + ++ num_combined; + const int stored_topk_slot_idx = lane_idx < kNumTopk ? + static_cast(__ldg(input_topk_idx_buffer.get_base_ptr() + token_idx * kNumTopk + lane_idx)) : -1; + const uint32_t total_mask = __ballot_sync(0xffffffff, stored_topk_slot_idx >= 0); + + for (uint32_t chunk = 0; chunk < kNumChunks; ++ chunk) { + const uint32_t chunk_byte_offset = chunk * kNumChunkBytes; + + uint32_t mask = total_mask; + const auto move_mask_and_load = [&](const uint32_t& i) { + if (mask) { + const uint32_t slot_idx = __ffs(mask) - 1; + mask ^= 1 << slot_idx; + if (cute::elect_one_sync()) { + const auto src_ptr = math::advance_ptr( + combine_token_buffer.get_rank_buffer(slot_idx) + .get_data_buffer(token_idx).get_base_ptr(), + chunk_byte_offset); + ptx::tma_load_1d(combine_load_buffer[i], src_ptr, combine_load_barriers[i], kNumChunkBytes); + ptx::mbarrier_arrive_and_set_tx(combine_load_barriers[i], kNumChunkBytes); + } + __syncwarp(); + return true; + } + return false; + }; + + bool do_reduce = move_mask_and_load(load_stage_idx); + + float2 reduced[kNumUint4PerLane * kNumElemsPerUint4] = {}; + while (do_reduce) { + do_reduce = move_mask_and_load(load_stage_idx ^ 1); + combine_load_barriers[load_stage_idx]->wait(combine_phase); + #pragma unroll + for (uint32_t j = 0; j < kNumUint4PerLane; ++ j) { + const auto uint4_values = combine_load_buffer[load_stage_idx][j * 32 + lane_idx]; + const auto bf16_values = reinterpret_cast(&uint4_values); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + ptx::accumulate(reduced[j * kNumElemsPerUint4 + l], bf16_values[l]); + } + combine_phase ^= load_stage_idx; + load_stage_idx ^= 1; + } + + #pragma unroll + for (uint32_t j = 0; j < kNumUint4PerLane; ++ j) { + uint4 casted; + auto casted_bf16 = reinterpret_cast(&casted); + #pragma unroll + for (uint32_t l = 0; l < kNumElemsPerUint4; ++ l) + casted_bf16[l] = __float22bfloat162_rn(reduced[j * kNumElemsPerUint4 + l]); + + if (j == 0) { + ptx::tma_store_wait<0>(); + __syncwarp(); + } + ptx::st_shared(combine_store_buffer + j * 32 + lane_idx, + casted.x, casted.y, casted.z, casted.w); + } + __syncwarp(); + + if (cute::elect_one_sync()) { + cute::tma_store_fence(); + ptx::tma_store_1d( + math::advance_ptr(y, static_cast(token_idx) * kNumHiddenBytes + chunk_byte_offset), + combine_store_buffer, kNumChunkBytes); + cute::tma_store_arrive(); + } + __syncwarp(); + } + } + trace_math(21, num_combined | (static_cast(num_ec_skipped) << 16)); // COMBINE_END (aux: tokens combined | skipped (early combine) << 16) + } +#else + if (blockIdx.x == 0 and threadIdx.x == 0) + DG_TRAP_ONLY_DEVICE_ASSERT(false and "This kernel only supports sm_90"); +#endif +} + +} // namespace deep_gemm + +#pragma clang diagnostic pop diff --git a/deep_gemm/include/deep_gemm/layout/sm90_fused_mega_moe.cuh b/deep_gemm/include/deep_gemm/layout/sm90_fused_mega_moe.cuh new file mode 100644 index 0000000000..7e94d4e346 --- /dev/null +++ b/deep_gemm/include/deep_gemm/layout/sm90_fused_mega_moe.cuh @@ -0,0 +1,294 @@ +#pragma once + +#include +#include +#include + +namespace deep_gemm::layout { + +// SM90 (`impls/sm90_fp8_fused_mega_moe.cuh`) only instantiates BLOCK_M in {64, 128}, so its pool/SF sizing uses this set +static constexpr int kNumSM90CandidateBlockMs = 2; +static constexpr int kSM90FusedCandidateBlockM[kNumSM90CandidateBlockMs] = {64, 128}; +static constexpr int kSM90FusedMaxCandidateBlockM = 128; +static constexpr int kSM90FusedLCMBlockM = 128; +static constexpr uint32_t kSM90FusedLagUnitM = 8; + +// Combine staging ladder of `impls/sm90_fp8_fused_mega_moe.cuh`, shared with the host so a hidden size the +// vectorization cannot serve is reported before NVRTC sees the specialization. Three whole-row slots per combine +// warp; the row is combined in chunks when they do not fit the shared memory in front of the barriers. +static constexpr uint32_t kSM90FusedCombineChunkSlots = 3; +static constexpr uint32_t kSM90FusedCombineMaxRegistersPerLane = 128; +static constexpr uint32_t kSM90FusedCombineElemBytes = 2; // bf16 + +CUTLASS_HOST_DEVICE constexpr uint32_t get_sm90_fused_moe_combine_num_chunks( + const uint32_t hidden, const uint32_t num_combine_warps, + const uint32_t smem_before_barriers, const bool split_mn_warpgroups) { + // hidden = 7 * 1024 is combined in 7 chunks, which keeps the 32-lane uint4 mapping + if (split_mn_warpgroups) + return hidden % 7 == 0 ? 7u : (hidden >= 1024 ? 4u : 1u); + const uint64_t one_chunk_bytes = static_cast(kSM90FusedCombineChunkSlots) * + num_combine_warps * hidden * kSM90FusedCombineElemBytes; + // 2 chunks do not always fit either: hidden 7168 at 3 stages needs 3 * 8 * 14336 / 2 = 172032 B of staging and + // the region in front of the barriers is short of that, so the ladder goes on to 4 + const uint32_t num_chunks_if_chunked = one_chunk_bytes / 2 <= smem_before_barriers ? 2u : 4u; + return (one_chunk_bytes <= smem_before_barriers and + hidden <= 32u * kSM90FusedCombineMaxRegistersPerLane) ? 1u : num_chunks_if_chunked; +} + +CUTLASS_HOST_DEVICE constexpr bool is_sm90_fused_moe_combine_vectorization_legal( + const uint32_t hidden, const uint32_t num_combine_warps, + const uint32_t smem_before_barriers, const bool split_mn_warpgroups) { + const auto num_chunks = get_sm90_fused_moe_combine_num_chunks( + hidden, num_combine_warps, smem_before_barriers, split_mn_warpgroups); + const uint64_t selected_chunk_bytes = static_cast(kSM90FusedCombineChunkSlots) * + num_combine_warps * hidden * kSM90FusedCombineElemBytes / num_chunks; + // every lane moves whole 16-byte elements, so a chunk must be a whole number of 32-lane rounds + return hidden > 0u and hidden % num_chunks == 0u and + selected_chunk_bytes <= smem_before_barriers and + (hidden * kSM90FusedCombineElemBytes / num_chunks) % (32u * sizeof(uint4)) == 0u; +} + +template +CUTLASS_HOST_DEVICE constexpr T get_num_max_pool_tokens_sm90(T num_ranks, T num_max_tokens_per_rank, T num_topk, + T num_experts_per_rank) { + const auto num_max_recv_tokens = num_ranks * num_max_tokens_per_rank; + const auto num_max_experts_per_token = math::constexpr_min(num_topk, num_experts_per_rank); + return math::constexpr_align( + num_max_recv_tokens * num_max_experts_per_token + num_experts_per_rank * (static_cast(kSM90FusedMaxCandidateBlockM) - 1), + static_cast(kSM90FusedLCMBlockM)); +} + +// SM90 tile table: at most ceil(total tiles / num_sms) entries per CTA plus the end marker, 16 bytes each (host smem sizing and the kernel layout both use this) +template +CUTLASS_HOST_DEVICE constexpr T get_sm90_tile_table_entries(T num_max_pool_tokens, T block_m, T num_l1_block_ns, + T num_l2_block_ns, T num_sms) { + return math::constexpr_ceil_div((num_max_pool_tokens / block_m) * (num_l1_block_ns + num_l2_block_ns), num_sms) + 1; +} + +// Compact table (8-byte entries): one spare entry more than the round-robin share, so the CTAs together absorb the whole bounded dynamic tail +template +CUTLASS_HOST_DEVICE constexpr T get_sm90_tile_table_entries_compact(T num_max_pool_tokens, T block_m, T num_l1_block_ns, + T num_l2_block_ns, T num_sms) { + return get_sm90_tile_table_entries(num_max_pool_tokens, block_m, num_l1_block_ns, num_l2_block_ns, num_sms) + 1; +} + +// SM90 MegaMoE predates the reusable ring workspace used by the SM100 kernels. +// Keep its compact layout explicit so Hopper codegen and buffer slicing stay +// identical to the tuned SM90 implementation while SM100 can use Workspace. +// The words remote ranks write into this workspace are double-buffered by launch parity (`t2_bank`); the per-expert recv counts are plain stores + one release flag per source rank, summed by the receivers. +struct SM90FusedWorkspace { + static constexpr bool kHeadLL = true; + void* base; + uint32_t num_ranks, num_experts; + uint32_t num_experts_per_rank; + uint32_t num_max_tokens_per_rank; + uint32_t num_max_recv_tokens_per_expert; + + uint32_t num_max_pool_tokens; + uint32_t num_max_pool_blocks; + + // Ring-buffer capacity (defaults to the full pool): only the reusable + // full/empty semaphores are sized by it, sized conservatively at + // `kMinCandidateBlockM` granularity so any compiled BLOCK_M fits. + uint32_t num_ring_tokens; + uint32_t num_ring_blocks; + + // bank of the words remote ranks write into THIS rank's workspace (recv counts, done flags, L2 tile counts): launch g + // uses bank g & 1 and zeroes bank (g + 1) & 1 after its tag-1 barrier, so nothing of generation g is zeroed while a remote rank may still write it + uint32_t t2_bank = 0; + + static constexpr uint64_t kNumBarrierSignalBytes = 32; + + CUTLASS_HOST_DEVICE + SM90FusedWorkspace(void* base, + const uint32_t& num_ranks, + const uint32_t& num_experts, + const uint32_t& num_max_tokens_per_rank, + const uint32_t& num_topk, + const uint32_t& num_ring_tokens = 0): + base(base), + num_ranks(num_ranks), num_experts(num_experts), + num_max_tokens_per_rank(num_max_tokens_per_rank), + num_ring_tokens(num_ring_tokens) { + num_experts_per_rank = num_experts / num_ranks; + num_max_recv_tokens_per_expert = num_ranks * num_max_tokens_per_rank; + num_max_pool_tokens = get_num_max_pool_tokens_sm90( + num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank); + num_max_pool_blocks = num_max_pool_tokens / kMinCandidateBlockM; + num_ring_blocks = (num_ring_tokens == 0 ? num_max_pool_tokens : num_ring_tokens) / kMinCandidateBlockM; + } + + CUTLASS_HOST_DEVICE + uint64_t get_num_bytes() const { + uint64_t num_bytes = 0; + num_bytes += kNumBarrierSignalBytes; + num_bytes += num_experts * sizeof(uint64_t) * 2; + num_bytes += num_experts_per_rank * sizeof(uint64_t); + num_bytes += math::align(num_max_pool_blocks, 2u) * sizeof(uint32_t); + num_bytes += num_max_pool_blocks * sizeof(uint64_t); + // Ring full/empty semaphores (alias the arrival arrays for `*_full`) + num_bytes += num_ring_blocks * sizeof(uint32_t) * 2; + num_bytes += num_experts_per_rank * num_ranks * num_max_recv_tokens_per_expert * sizeof(int); + num_bytes += num_max_pool_tokens * sizeof(TokenSrcMetadata); + // early-combine flags, tile counts and the launch-parity second bank; padded to 1 KiB so the data pools that follow keep their alignment + num_bytes += math::align(get_t2_bank1_bytes() + get_t2_bank1_offset_in_flags(), 1024); + return math::align(num_bytes, 16); + } + + // Bank geometry (bank 1 = the copy behind bank 0, 8-byte aligned): u64 part [recv_count x num_experts | unused x num_experts_per_rank], + // u32 part [done_flag x num_experts | l2_tile_done_count x num_experts_per_rank | head count flags x num_ranks] + CUTLASS_HOST_DEVICE + uint32_t get_t2_bank_flag_words() const { + return num_experts + num_experts_per_rank + num_ranks; + } + + CUTLASS_HOST_DEVICE + uint64_t get_t2_bank1_offset_in_flags() const { + return math::align(get_t2_bank_flag_words() * sizeof(uint32_t), 8); + } + + CUTLASS_HOST_DEVICE + uint64_t get_t2_bank1_bytes() const { + return (num_experts + num_experts_per_rank) * sizeof(uint64_t) + get_t2_bank_flag_words() * sizeof(uint32_t); + } + + CUTLASS_HOST_DEVICE + void* get_end_ptr() const { + return math::advance_ptr(base, get_num_bytes()); + } + + static constexpr uint32_t kNumMaxGridSyncCounters = 4; + + template + CUTLASS_DEVICE + uint32_t* get_grid_sync_count_ptr() const { + DG_STATIC_ASSERT(kIndex < kNumMaxGridSyncCounters, "Grid sync index out of bounds"); + return static_cast(base) + kIndex; + } + + CUTLASS_DEVICE + uint32_t* get_nvl_barrier_counter_ptr() const { + return static_cast(base) + kNumMaxGridSyncCounters; + } + + CUTLASS_DEVICE + int* get_nvl_barrier_signal_ptr(const uint32_t& phase) const { + return math::advance_ptr( + base, (kNumMaxGridSyncCounters + 1) * sizeof(uint32_t) + phase * sizeof(int)); + } + + // Launch parity word (last word of the 32-byte signal block; flipped by SM0 in the tail of every launch, read after the PDL wait) + CUTLASS_DEVICE + uint32_t* get_t2_parity_ptr() const { + return static_cast(base) + kNumMaxGridSyncCounters + 3; + } + + CUTLASS_DEVICE + uint64_t* get_expert_send_count_ptr(const uint32_t& expert_idx = 0) const { + return math::advance_ptr(base, kNumBarrierSignalBytes) + expert_idx; + } + + CUTLASS_DEVICE + uint64_t* get_bank0_u64_end_ptr() const { + return get_expert_send_count_ptr(num_experts * 2) + num_experts_per_rank; + } + + CUTLASS_DEVICE + uint64_t* get_t2_recv_count_bank_ptr(const uint32_t& bank) const { + return bank ? get_t2_bank1_recv_count_ptr() : get_expert_send_count_ptr(num_experts); + } + + CUTLASS_DEVICE + uint64_t* get_expert_recv_count_ptr( + const uint32_t& rank_idx = 0, const uint32_t& expert_idx = 0) const { + return get_t2_recv_count_bank_ptr(t2_bank) + rank_idx * num_experts_per_rank + expert_idx; + } + + CUTLASS_DEVICE + uint32_t* get_l1_arrival_count_ptr(const uint32_t& pool_block_idx = 0) const { + return reinterpret_cast(get_bank0_u64_end_ptr()) + pool_block_idx; + } + + CUTLASS_DEVICE + uint64_t* get_l2_arrival_mask_ptr(const uint32_t& pool_block_idx = 0) const { + const auto base = get_l1_arrival_count_ptr(math::align(num_max_pool_blocks, 2u)); + return reinterpret_cast(base) + pool_block_idx; + } + + // Ring-mode counting semaphores. `*_full` alias the legacy arrival arrays + // (u32 view for L2), `*_empty` live in the dedicated ring area. + CUTLASS_DEVICE + uint32_t* get_l1_full_count_ptr(const uint32_t& ring_block_idx = 0) const { + return get_l1_arrival_count_ptr(ring_block_idx); + } + + CUTLASS_DEVICE + uint32_t* get_l2_full_count_ptr(const uint32_t& ring_block_idx = 0) const { + return reinterpret_cast(get_l2_arrival_mask_ptr(ring_block_idx)); + } + + CUTLASS_DEVICE + uint32_t* get_l1_empty_count_ptr(const uint32_t& ring_block_idx = 0) const { + const auto base = reinterpret_cast( + get_l2_arrival_mask_ptr(num_max_pool_blocks)); + return base + ring_block_idx; + } + + CUTLASS_DEVICE + uint32_t* get_l2_empty_count_ptr(const uint32_t& ring_block_idx = 0) const { + return get_l1_empty_count_ptr(num_ring_blocks) + ring_block_idx; + } + + CUTLASS_DEVICE + uint32_t* get_src_token_topk_idx_ptr( + const uint32_t& expert_idx = 0, const uint32_t& rank_idx = 0, const uint32_t& token_idx = 0) const { + const auto base = get_l2_empty_count_ptr(num_ring_blocks); + return reinterpret_cast(base) + + expert_idx * (num_ranks * num_max_recv_tokens_per_expert) + + rank_idx * num_max_recv_tokens_per_expert + token_idx; + } + + CUTLASS_DEVICE + TokenSrcMetadata* get_token_src_metadata_ptr(const uint32_t& pool_token_idx = 0) const { + const auto base = reinterpret_cast(get_src_token_topk_idx_ptr(num_experts_per_rank)); + return base + pool_token_idx; + } + + // early-combine done flags: set by the owning rank once all L2 tiles of a global expert have scattered. Warps 1.. combine + // a token as soon as its top-k experts' flags are all set, overlapping that with the same all-rank barrier the other + // mode also takes; tokens whose flags are not all set are combined after it + CUTLASS_DEVICE + uint32_t* get_expert_done_flag_ptr_bank0(const uint32_t& expert_idx = 0) const { + return reinterpret_cast(get_token_src_metadata_ptr(num_max_pool_tokens)) + expert_idx; + } + + CUTLASS_DEVICE + uint64_t* get_t2_bank1_recv_count_ptr() const { + return reinterpret_cast( + reinterpret_cast(get_expert_done_flag_ptr_bank0()) + get_t2_bank1_offset_in_flags()); + } + + CUTLASS_DEVICE + uint32_t* get_t2_flag_bank_ptr(const uint32_t& bank) const { + return bank ? reinterpret_cast(get_t2_bank1_recv_count_ptr() + num_experts + num_experts_per_rank) + : get_expert_done_flag_ptr_bank0(); + } + + CUTLASS_DEVICE + uint32_t* get_expert_done_flag_ptr(const uint32_t& expert_idx = 0) const { + return get_t2_flag_bank_ptr(t2_bank) + expert_idx; + } + + CUTLASS_DEVICE + uint32_t* get_l2_tile_done_count_ptr(const uint32_t& local_expert_idx = 0) const { + return get_t2_flag_bank_ptr(t2_bank) + num_experts + local_expert_idx; + } + + // per source rank "my recv counts for you are stored" flag (st.release.sys by the source's SM0 after its count stores; receivers poll) + CUTLASS_DEVICE + uint32_t* get_hll_count_flag_ptr(const uint32_t& src_rank_idx = 0) const { + return get_t2_flag_bank_ptr(t2_bank) + num_experts + num_experts_per_rank + src_rank_idx; + } +}; + +} // namespace deep_gemm::layout diff --git a/deep_gemm/include/deep_gemm/mma/sm90.cuh b/deep_gemm/include/deep_gemm/mma/sm90.cuh index 2c061940de..00747f76ba 100644 --- a/deep_gemm/include/deep_gemm/mma/sm90.cuh +++ b/deep_gemm/include/deep_gemm/mma/sm90.cuh @@ -208,6 +208,14 @@ make_smem_desc(PointerType smem_ptr, const int& layout_type, return desc; } +CUTLASS_DEVICE cute::GmmaDescriptor +advance_smem_desc(const cute::GmmaDescriptor& desc, const uint32_t& byte_offset) { + cute::GmmaDescriptor advanced; + advanced.desc_ = desc.desc_; + advanced.reg32_[0] += byte_offset >> 4; + return advanced; +} + template constexpr uint32_t get_inner_block_atom_size() { return kSwizzleMode == 0 ? BLOCK_INNER : kSwizzleMode / sizeof(dtype_t); diff --git a/deep_gemm/include/deep_gemm/ptx/ld_st.cuh b/deep_gemm/include/deep_gemm/ptx/ld_st.cuh index 806a4c2e20..cd1bf6d123 100644 --- a/deep_gemm/include/deep_gemm/ptx/ld_st.cuh +++ b/deep_gemm/include/deep_gemm/ptx/ld_st.cuh @@ -179,6 +179,71 @@ CUTLASS_DEVICE void st_async_cluster(T* dst, const T& src, const uint32_t& dst_c } } +// early combine: CTA-scope message passing through shared memory and the wider-scope publish +CUTLASS_DEVICE uint2 ld_shared(const uint2* ptr) { + uint2 ret; + asm volatile("ld.shared.v2.u32 {%0, %1}, [%2];" : "=r"(ret.x), "=r"(ret.y) : "l"(__cvta_generic_to_shared(ptr))); + return ret; +} + +CUTLASS_DEVICE uint32_t ld_volatile_shared(const uint32_t* ptr) { + uint32_t ret; + asm volatile("ld.volatile.shared.u32 %0, [%1];" : "=r"(ret) : "l"(__cvta_generic_to_shared(ptr)) : "memory"); + return ret; +} + +CUTLASS_DEVICE uint32_t ld_acquire_cta_shared(const uint32_t* ptr) { + uint32_t ret; + asm volatile("ld.acquire.cta.shared::cta.u32 %0, [%1];" : "=r"(ret) : "l"(__cvta_generic_to_shared(ptr)) : "memory"); + return ret; +} + +CUTLASS_DEVICE void st_release_cta_shared(const uint32_t* ptr, const uint32_t& value) { + asm volatile("st.release.cta.shared::cta.u32 [%0], %1;" :: "l"(__cvta_generic_to_shared(ptr)), "r"(value) : "memory"); +} + +CUTLASS_DEVICE void red_release_cta_shared_add(const uint32_t* ptr, const uint32_t& value) { + asm volatile("red.release.cta.shared::cta.add.u32 [%0], %1;" :: "l"(__cvta_generic_to_shared(ptr)), "r"(value) : "memory"); +} + +CUTLASS_DEVICE void fence_acq_rel_gpu() { + asm volatile("fence.acq_rel.gpu;" ::: "memory"); +} + +CUTLASS_DEVICE void fence_acq_rel_sys() { + asm volatile("fence.acq_rel.sys;" ::: "memory"); +} + +CUTLASS_DEVICE uint32_t atom_add_relaxed_gpu(const uint32_t* ptr, const uint32_t& value) { + uint32_t ret; + asm volatile("atom.relaxed.gpu.global.add.u32 %0, [%1], %2;" : "=r"(ret) : "l"(ptr), "r"(value) : "memory"); + return ret; +} + +CUTLASS_DEVICE void red_add_relaxed_sys(const uint32_t* ptr, const uint32_t& value) { + asm volatile("red.relaxed.sys.global.add.u32 [%0], %1;" :: "l"(ptr), "r"(value) : "memory"); +} + +// count flag store of the low-latency head (release: orders the thread's earlier stores, and by cumulativity those of the +// threads it synchronised with, before the flag at system scope) +CUTLASS_DEVICE void st_release_sys_u32(uint32_t* ptr, const uint32_t& value) { + asm volatile("st.release.sys.global.u32 [%0], %1;" :: "l"(ptr), "r"(value) : "memory"); +} + +CUTLASS_DEVICE uint32_t ld_relaxed_sys(const uint32_t* ptr) { + uint32_t ret; + asm volatile("ld.relaxed.sys.global.u32 %0, [%1];" : "=r"(ret) : "l"(ptr) : "memory"); + return ret; +} + +CUTLASS_DEVICE void st_global_v2(void* ptr, const uint32_t& x, const uint32_t& y) { + asm volatile("st.global.v2.u32 [%0], {%1, %2};" :: "l"(ptr), "r"(x), "r"(y) : "memory"); +} + +CUTLASS_DEVICE void prefetch_bulk_l2(const void* ptr, const uint32_t& num_bytes) { + asm volatile("cp.async.bulk.prefetch.L2.global [%0], %1;" :: "l"(ptr), "r"(num_bytes) : "memory"); +} + CUTLASS_DEVICE void st_shared_bulk(void* smem_ptr, const uint32_t& num_bytes) { // `size` must be 64-bit before PTX ISA 9.0 DG_DEVICE_ASSERT(num_bytes % 8 == 0); diff --git a/deep_gemm/include/deep_gemm/scheduler/sm90_fused_mega_moe.cuh b/deep_gemm/include/deep_gemm/scheduler/sm90_fused_mega_moe.cuh new file mode 100644 index 0000000000..4995f0bb47 --- /dev/null +++ b/deep_gemm/include/deep_gemm/scheduler/sm90_fused_mega_moe.cuh @@ -0,0 +1,655 @@ +#pragma once + +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace deep_gemm::sched { + +template 0: L2-lag interleaved schedule (see get_next_block_lag); 0: wave schedule + uint32_t kL2LagUnits = 0, + // bounds of the bit-packed tile-table entry fields (pool blocks of the rank, tokens one expert can receive); tile-table users (SM90) only + uint32_t kNumMaxPoolBlocks = 0, + uint32_t kNumMaxTokensPerExpert = 0, + // CTAs per cluster: with 2 the leader publishes the partner's dynamic-tail entries over DSMEM, so the tag loads + // acquire at cluster scope + uint32_t kTileTableClusterSize = 1> +struct SM90FusedMegaMoEScheduler { + DG_STATIC_ASSERT(L1_SHAPE_N % BLOCK_N == 0, "Invalid shape"); + DG_STATIC_ASSERT(L2_SHAPE_N % BLOCK_N == 0, "Invalid shape"); + DG_STATIC_ASSERT(L1_SHAPE_K % BLOCK_K == 0, "Invalid shape"); + DG_STATIC_ASSERT(L2_SHAPE_K % BLOCK_K == 0, "Invalid shape"); + DG_STATIC_ASSERT(kNumExpertsPerWave > 0 and kNumExpertsPerWave <= kNumExpertsPerRank, "Invalid wave config"); + + // NOTES: with a 2-CTA cluster the N block counts must be even so that 2 adjacent CTAs always land on + // the same m_block_idx with n_block_idx differing by 1; a shape (or SM count) that cannot be paired runs + // 1-CTA clusters instead, which the host selects + DG_STATIC_ASSERT(kTileTableClusterSize == 1 or kNumSMs % 2 == 0, "Number of SMs must be even for 2-CTA cluster"); + DG_STATIC_ASSERT(kTileTableClusterSize == 1 or kNumL1BlockNs % 2 == 0, "L1 N block count must be even for 2-CTA cluster"); + DG_STATIC_ASSERT(kTileTableClusterSize == 1 or kNumL2BlockNs % 2 == 0, "L2 N block count must be even for 2-CTA cluster"); + + // Arrival counts + const WorkspaceT& workspace; + + // Scheduler state + BlockPhase next_phase = BlockPhase::Linear1; + + // Current expert and block indices + uint32_t current_local_expert_idx = 0; + uint32_t current_num_tokens = 0; + uint32_t current_pool_block_offset = 0; + uint32_t block_idx = 0; + uint32_t m_block_idx = 0; + uint32_t n_block_idx = 0; + + // Pre-cached per-expert token counts (filled during `for_each_block` init) + // Layout: `stored_num_tokens_per_expert[i]` holds expert (i * 32 + lane_idx)'s count + uint32_t stored_num_tokens_per_expert[kNumExpertsPerLane] = {}; + + // L2-lag schedule state (kL2LagUnits > 0). A unit is up to layout::kSM90FusedLagUnitM consecutive m-blocks of one expert + // times all N blocks; inside a unit the order is m-inner and pair validity is decided per position (lag_decode). + // The global sequence is + // L1(U0) .. L1(U_lag) L2(U0) L1(U_lag+1) L2(U1) ... then the remaining L2 units, + // i.e. L2 of a pool block is issued kL2LagUnits units after its L1 and the whole rank is one wave. L2(p) only ever + // waits on L1(p), which precedes it and never waits on any L2, so the smallest pending L2 always makes progress. + uint32_t l1_expert = 0, l1_num_tokens = 0, l1_pool_offset = 0, l1_m0 = 0; + uint32_t l2_expert = 0, l2_num_tokens = 0, l2_pool_offset = 0, l2_m0 = 0; + uint32_t num_l1_units_done = 0; + uint32_t num_l2_units_pending = 0; // L2 units still to emit in the current group + bool current_pair_valid = true; + // kL2LagUnits encodes lag + 1000 * group: after every `group` L1 units past the lag the same number of L2 units is emitted back to back + static constexpr uint32_t kLagUnitsOnly = kL2LagUnits % 1000u; + static constexpr uint32_t kLagGroupUnits = (kL2LagUnits / 1000u) > 0 ? (kL2LagUnits / 1000u) : 1u; + + CUTLASS_DEVICE explicit SM90FusedMegaMoEScheduler(const WorkspaceT& workspace): workspace(workspace) { + block_idx = blockIdx.x; + } + + CUTLASS_DEVICE uint32_t get_wave_expert_end_idx() const { + // Align up to wave boundary, clamped for the last partial wave + const auto aligned = math::align(current_local_expert_idx + 1, kNumExpertsPerWave); + return cute::min(aligned, kNumExpertsPerRank); + } + + CUTLASS_DEVICE uint32_t get_num_tokens(const uint32_t& expert_idx) const { + uint32_t valid_value; + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++ i) { + valid_value = (expert_idx == i * 32 + ptx::get_lane_idx()) ? + stored_num_tokens_per_expert[i] : valid_value; + } + return ptx::exchange(valid_value, expert_idx % 32); + } + + // Get pool block offset for a given expert index from a per-lane token count array + CUTLASS_DEVICE uint32_t get_pool_block_offset(const uint32_t& expert_idx) { + uint32_t num_blocks = 0; + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++ i) { + if (i * 32 + ptx::get_lane_idx() < expert_idx) + num_blocks += math::ceil_div(stored_num_tokens_per_expert[i], BLOCK_M); + } + return __reduce_add_sync(0xffffffff, num_blocks); + } + + CUTLASS_DEVICE void advance_expert_idx() { + current_pool_block_offset += get_current_num_m_blocks(); + current_local_expert_idx += 1; + current_num_tokens = get_num_tokens(current_local_expert_idx); + } + + CUTLASS_DEVICE void set_expert_idx(const uint32_t& expert_idx) { + current_local_expert_idx = expert_idx; + current_num_tokens = get_num_tokens(expert_idx); + current_pool_block_offset = get_pool_block_offset(expert_idx); + if constexpr (kL2LagUnits > 0) { + l1_expert = l2_expert = expert_idx; + l1_num_tokens = l2_num_tokens = current_num_tokens; + l1_pool_offset = l2_pool_offset = current_pool_block_offset; + l1_m0 = l2_m0 = 0; + num_l1_units_done = 0; + num_l2_units_pending = 0; + lag_skip_empty(l1_expert, l1_num_tokens, l1_pool_offset, l1_m0); + lag_skip_empty(l2_expert, l2_num_tokens, l2_pool_offset, l2_m0); + } + } + + // Cluster pairing: positions 2j (leader) and 2j+1 (partner) must share n and differ in m by 1 for B multicast. Wave + // schedule: local check (chunks are even-sized). Lag schedule: decided per unit so both CTAs agree on 1-wide tail units. + CUTLASS_DEVICE bool is_pair_valid(const bool& is_leader) const { + if constexpr (kL2LagUnits > 0) { + return current_pair_valid; + } else { + const auto num_m_blocks = get_current_num_m_blocks(); + return is_leader ? (m_block_idx + 1 < num_m_blocks) : (m_block_idx > 0); + } + } + + // Advance a lag cursor past experts that have no m-blocks left + CUTLASS_DEVICE void lag_skip_empty(uint32_t& expert, uint32_t& num_tokens, uint32_t& pool_offset, uint32_t& m0) { + while (expert < kNumExpertsPerRank and m0 >= math::ceil_div(num_tokens, BLOCK_M)) { + pool_offset += math::ceil_div(num_tokens, BLOCK_M); + expert += 1; + m0 = 0; + num_tokens = expert < kNumExpertsPerRank ? get_num_tokens(expert) : 0u; + } + } + + // Decode the CTA's offset inside the current unit chunk into (m, n) + CUTLASS_DEVICE void lag_decode(const uint32_t& q, const uint32_t& um, const uint32_t& num_n_blocks, const uint32_t& m0) { + if constexpr (kMulticastOnB) { + const uint32_t m_in_unit = q % um; + m_block_idx = m0 + m_in_unit; + n_block_idx = q / um; + current_pair_valid = (q % 2 == 0) ? (m_in_unit + 1 < um) : (m_in_unit > 0); + } else { + m_block_idx = m0 + q / num_n_blocks; + n_block_idx = q % num_n_blocks; + current_pair_valid = false; + } + } + + // Dynamic tail (kStopAtTail): the walk stops at the trailing L2-only units with `tail_reached` set and the L2 cursor on the first tail unit + bool tail_reached = false; + + template + CUTLASS_DEVICE cute::tuple get_next_block_lag() { + constexpr uint32_t kUnitM = layout::kSM90FusedLagUnitM; + while (true) { + if (l2_expert >= kNumExpertsPerRank) + return {BlockPhase::None, 0, 0, 0}; + const bool walk_l2 = (num_l2_units_pending > 0) or (l1_expert >= kNumExpertsPerRank); + if (not walk_l2) { + const uint32_t um = cute::min(kUnitM, math::ceil_div(l1_num_tokens, BLOCK_M) - l1_m0); + const uint32_t chunk = um * kNumL1BlockNs; + if (block_idx < chunk) { + current_local_expert_idx = l1_expert; + current_num_tokens = l1_num_tokens; + current_pool_block_offset = l1_pool_offset; + lag_decode(block_idx, um, kNumL1BlockNs, l1_m0); + block_idx += kNumSMs; + return {BlockPhase::Linear1, l1_expert, m_block_idx, n_block_idx}; + } + block_idx -= chunk; + l1_m0 += um; + lag_skip_empty(l1_expert, l1_num_tokens, l1_pool_offset, l1_m0); + num_l1_units_done += 1; + if (num_l1_units_done > kLagUnitsOnly and (num_l1_units_done - kLagUnitsOnly) % kLagGroupUnits == 0) + num_l2_units_pending = kLagGroupUnits; + } else { + if constexpr (kStopAtTail) { + if (l1_expert >= kNumExpertsPerRank) { + tail_reached = true; + return {BlockPhase::None, 0, 0, 0}; + } + } + const uint32_t um = cute::min(kUnitM, math::ceil_div(l2_num_tokens, BLOCK_M) - l2_m0); + const uint32_t chunk = um * kNumL2BlockNs; + if (block_idx < chunk) { + current_local_expert_idx = l2_expert; + current_num_tokens = l2_num_tokens; + current_pool_block_offset = l2_pool_offset; + lag_decode(block_idx, um, kNumL2BlockNs, l2_m0); + block_idx += kNumSMs; + return {BlockPhase::Linear2, l2_expert, m_block_idx, n_block_idx}; + } + block_idx -= chunk; + l2_m0 += um; + lag_skip_empty(l2_expert, l2_num_tokens, l2_pool_offset, l2_m0); + if (num_l2_units_pending > 0) + num_l2_units_pending -= 1; + } + } + } + + CUTLASS_DEVICE uint32_t get_current_pool_block_offset() const { + return current_pool_block_offset; + } + + CUTLASS_DEVICE uint32_t get_current_num_m_blocks() const { + return math::ceil_div(current_num_tokens, BLOCK_M); + } + + template + CUTLASS_DEVICE uint32_t get_valid_m() const { + const auto m = cute::min(current_num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + return kDoUMMAAligned ? math::align(m, 16u) : m; + } + + CUTLASS_DEVICE bool fetch_next_l1_block() { + const auto wave_end_expert_idx = get_wave_expert_end_idx(); + while (current_local_expert_idx < wave_end_expert_idx) { + const auto num_m_blocks = get_current_num_m_blocks(); + if constexpr (kMulticastOnB) { + if (block_idx < num_m_blocks * kNumL1BlockNs) { + m_block_idx = block_idx % num_m_blocks; + n_block_idx = block_idx / num_m_blocks; + return true; + } + } else { + m_block_idx = block_idx / kNumL1BlockNs; + if (m_block_idx < num_m_blocks) { + n_block_idx = block_idx - m_block_idx * kNumL1BlockNs; + return true; + } + } + + // Current expert is fully assigned, move to the next + block_idx -= num_m_blocks * kNumL1BlockNs; + advance_expert_idx(); + } + return false; + } + + CUTLASS_DEVICE bool fetch_next_l2_block() { + const auto wave_end_expert_idx = get_wave_expert_end_idx(); + while (current_local_expert_idx < wave_end_expert_idx) { + const auto num_m_blocks = get_current_num_m_blocks(); + if (block_idx < num_m_blocks * kNumL2BlockNs) { + if constexpr (kMulticastOnB) { + m_block_idx = block_idx % num_m_blocks; + n_block_idx = block_idx / num_m_blocks; + } else { + m_block_idx = block_idx / kNumL2BlockNs; + n_block_idx = block_idx - m_block_idx * kNumL2BlockNs; + } + return true; + } + + // Current expert is fully assigned, move to the next + block_idx -= num_m_blocks * kNumL2BlockNs; + advance_expert_idx(); + } + return false; + } + + // Core state machine. kStopAtTail: the wave schedule stops when the L2 phase of the last wave begins, the lag schedule at its trailing L2-only units + template + CUTLASS_DEVICE cute::tuple get_next_block() { + if constexpr (kL2LagUnits > 0) + return get_next_block_lag(); + while (true) { + if (current_local_expert_idx >= kNumExpertsPerRank) + break; + + if (next_phase == BlockPhase::Linear1) { + if (fetch_next_l1_block()) { + // Found a new L1 block (m/n set by the fetcher); jump to next + block_idx += kNumSMs; + return {BlockPhase::Linear1, current_local_expert_idx, m_block_idx, n_block_idx}; + } else { + // L1 for the current wave is complete, transition to L2 + next_phase = BlockPhase::Linear2; + set_expert_idx(math::align(current_local_expert_idx - 1, kNumExpertsPerWave)); + if constexpr (kStopAtTail) { + if (get_wave_expert_end_idx() >= kNumExpertsPerRank) { + tail_reached = true; + return {BlockPhase::None, 0, 0, 0}; + } + } + } + } else { + if (fetch_next_l2_block()) { + // Found a new L2 block (m/n set by the fetcher); jump to next + block_idx += kNumSMs; + return {BlockPhase::Linear2, current_local_expert_idx, m_block_idx, n_block_idx}; + } else { + // Move to L1 of the next wave + next_phase = BlockPhase::Linear1; + } + } + } + + // All waves and experts are fully processed + return {BlockPhase::None, 0, 0, 0}; + } + + // does the workspace publish the recv counts through per-source flags (SM90FusedWorkspace) or through the summed count (layout::Workspace)? + template struct HasHeadLL : std::false_type {}; + template struct HasHeadLL> : std::bool_constant {}; + static constexpr bool kHeadLL = HasHeadLL::value; + + CUTLASS_DEVICE void fetch_expert_recv_count() { + if constexpr (kHeadLL) { + // low-latency head: lane r waits for source rank r's release flag (its count stores are visible once the acquire + // returns), the __syncwarp orders every lane behind all the acquires, then each lane sums its experts' counts over the ranks + DG_STATIC_ASSERT(kNumRanks <= 32, "low-latency head: one flag per lane"); + if (ptx::get_lane_idx() < kNumRanks) + while (ptx::ld_acq_sys(workspace.get_hll_count_flag_ptr(ptx::get_lane_idx())) == 0u); + __syncwarp(); + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++ i) { + const auto expert_idx = i * 32 + ptx::get_lane_idx(); + uint32_t sum = 0; + if (expert_idx < kNumExpertsPerRank) { + #pragma unroll + for (uint32_t r = 0; r < kNumRanks; ++ r) + sum += static_cast(ptx::ld_volatile(workspace.get_expert_recv_count_ptr(r, expert_idx))); + } + stored_num_tokens_per_expert[i] = sum; + } + } else { + // NOTES: each lane caches experts at indices (i * 32 + lane_idx) + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++ i) { + const auto expert_idx = i * 32 + ptx::get_lane_idx(); + uint64_t value = 0; + if (expert_idx < kNumExpertsPerRank) { + do { + value = ptx::ld_volatile(workspace.get_expert_recv_count_sum_ptr(expert_idx)); + } while (static_cast(value >> 32) != kNumSMs * kNumRanks); + } + stored_num_tokens_per_expert[i] = static_cast(value); + } + } + __syncwarp(); + } + + template + CUTLASS_DEVICE void for_each_block(Func&& func) { + // Wait for all expert counters to be finalized + fetch_expert_recv_count(); + + // Initialize current expert with 0 + set_expert_idx(0); + + // Iterate over all blocks + while (true) { + CUTE_TIE_DECL(get_next_block(), block_phase, current_local_expert_idx, m_block_idx, n_block_idx); + if (block_phase == BlockPhase::None) + break; + + func(block_phase, current_local_expert_idx, + block_phase == BlockPhase::Linear2 ? kNumL2BlockKs : kNumL1BlockKs, + m_block_idx, n_block_idx); + } + } + + // Tile table (SM90): one role walks the schedule with `for_each_block_publish` and stores every tile as a shared-memory + // entry; the other roles replay the entries with `for_each_block_replay`. The publisher must run at least one tile ahead + // of every replayer (the B loader does: it only waits on pipeline slots, never on data), and the table must hold + // ceil(total tiles / kNumSMs) + 1 entries, all initialised to kTileTableNotReady before any replay starts. + // 8-byte entry, bit-packed with widths derived from the template constants (static_asserts below): + // x = expert | m_block << kEBits | n_block << (kEBits + kMBits) | pair_valid << 29 | tag << 30 + // (tag 0: end of schedule, 1: Linear1, 2: Linear2, 3: not ready) + // y = num_tokens | pool_block_offset << kTBits (kTBits + kPBits <= 32) + // y = valid_m | last_m_block << kVBits | pool_block_offset << kVBits + 1 (otherwise, see kTilePayloadByRows) + // Publish protocol: payload words first, then the tag word with a release store (CTA scope locally, cluster scope for + // the DSMEM publish of the dynamic tail); replayers spin on the 32-bit tag word with an acquire load of the matching + // scope and read the payload afterwards. A single 16 B store / 16 B load pair is not single-copy atomic, hence the + // two-step protocol. + static constexpr uint32_t kTileTableNotReady = 0xffffffffu; + using TileEntry = uint2; + + // Bit widths of the compact entry (the value ranges are template constants) + CUTLASS_HOST_DEVICE static constexpr uint32_t bits_for(uint32_t max_value) { + uint32_t bits = 1; + while (bits < 32 and (max_value >> bits) != 0) + ++ bits; + return bits; + } + static constexpr uint32_t kEBits = bits_for(kNumExpertsPerRank); + static constexpr uint32_t kMBits = bits_for(kNumMaxPoolBlocks); + static constexpr uint32_t kNBits = bits_for(kNumL1BlockNs > kNumL2BlockNs ? kNumL1BlockNs : kNumL2BlockNs); + static constexpr uint32_t kTBits = bits_for(kNumMaxTokensPerExpert); + static constexpr uint32_t kPBits = bits_for(kNumMaxPoolBlocks); + DG_STATIC_ASSERT(kNumMaxPoolBlocks == 0 or (kEBits + kMBits + kNBits <= 29u), "tile table: tag word overflow"); + // Payload word format: the expert's token count and its pool-block offset share the word when they fit; otherwise + // (kTilePayloadByRows) the entry carries the tile's valid row count, a "last m-block of the expert" flag and the pool offset, + // and the replayer rebuilds a token count that yields the same `get_valid_m` and the same pair validity as the expert's count. + // Only those two derived quantities and the pool offset are read after a replay. + static constexpr bool kTilePayloadByRows = kNumMaxTokensPerExpert != 0 and (kTBits + kPBits > 32u); + static constexpr uint32_t kVBits = bits_for(BLOCK_M); + DG_STATIC_ASSERT(not kTilePayloadByRows or (kVBits + 1u + kPBits <= 32u), "tile table: payload word overflow"); + + CUTLASS_DEVICE static uint32_t ld_acquire_tile_tag(const TileEntry* entry) { + uint32_t ret; + if constexpr (kTileTableClusterSize > 1) { + asm volatile("ld.acquire.cluster.shared::cta.u32 %0, [%1];" : "=r"(ret) : "l"(__cvta_generic_to_shared(entry)) : "memory"); + } else { + asm volatile("ld.acquire.cta.shared::cta.u32 %0, [%1];" : "=r"(ret) : "l"(__cvta_generic_to_shared(entry)) : "memory"); + } + return ret; + } + + CUTLASS_DEVICE static void st_release_tile_tag(const TileEntry* entry, const uint32_t& tag_word) { + asm volatile("st.release.cta.shared.u32 [%0], %1;" :: "l"(__cvta_generic_to_shared(entry)), "r"(tag_word) : "memory"); + } + + // The same two stores into the shared memory of another CTA of the cluster + CUTLASS_DEVICE static void st_remote_tile_entry(const TileEntry* entry, const uint32_t& cta_rank, + const uint32_t& tag_word, const uint32_t& y) { + const auto local_addr = static_cast(__cvta_generic_to_shared(entry)); + uint32_t remote_addr; + asm volatile("mapa.shared::cluster.u32 %0, %1, %2;" : "=r"(remote_addr) : "r"(local_addr), "r"(cta_rank)); + asm volatile("st.shared::cluster.v2.u32 [%0], {%1, %2};" + :: "r"(remote_addr), "r"(kTileTableNotReady), "r"(y) : "memory"); + asm volatile("st.release.cluster.shared::cluster.u32 [%0], %1;" :: "r"(remote_addr), "r"(tag_word) : "memory"); + } + + CUTLASS_DEVICE static uint32_t make_compact_tag_word(const uint32_t& expert, const uint32_t& m_block, const uint32_t& n_block, + const bool& pair_valid, const uint32_t& tag) { + return expert | (m_block << kEBits) | (n_block << (kEBits + kMBits)) | + (static_cast(pair_valid) << 29) | (tag << 30); + } + CUTLASS_DEVICE static uint32_t make_compact_payload_word(const uint32_t& num_tokens, const uint32_t& pool_block_offset, + const uint32_t& m_block) { + if constexpr (kTilePayloadByRows) { + const uint32_t valid_m = cute::min(num_tokens - m_block * BLOCK_M, BLOCK_M); + const uint32_t last_m_block = (m_block + 1 == math::ceil_div(num_tokens, BLOCK_M)) ? 1u : 0u; + return valid_m | (last_m_block << kVBits) | (pool_block_offset << (kVBits + 1u)); + } else { + return num_tokens | (pool_block_offset << kTBits); + } + } + + CUTLASS_DEVICE void publish_tile_entry(TileEntry* entry, const uint32_t& tag) const { + ptx::st_shared(entry, kTileTableNotReady, make_compact_payload_word(current_num_tokens, current_pool_block_offset, m_block_idx)); + st_release_tile_tag(entry, make_compact_tag_word(current_local_expert_idx, m_block_idx, n_block_idx, current_pair_valid, tag)); + } + + // Decodes a tile-table entry into the scheduler state; returns the tag + CUTLASS_DEVICE uint32_t load_tile_entry(const TileEntry* entry) { + uint32_t tag_word = ld_acquire_tile_tag(entry); + while ((tag_word >> 30) == 3u) + tag_word = ld_acquire_tile_tag(entry); + const uint32_t tag = tag_word >> 30; + if (tag == 0u) + return 0u; + const uint32_t payload = ld_tile_table_payload(entry); + current_local_expert_idx = tag_word & ((1u << kEBits) - 1u); + m_block_idx = (tag_word >> kEBits) & ((1u << kMBits) - 1u); + n_block_idx = (tag_word >> (kEBits + kMBits)) & ((1u << kNBits) - 1u); + current_pair_valid = ((tag_word >> 29) & 1u) != 0; + if constexpr (kTilePayloadByRows) { + const uint32_t valid_m = payload & ((1u << kVBits) - 1u); + const bool last_m_block = ((payload >> kVBits) & 1u) != 0; + current_num_tokens = m_block_idx * BLOCK_M + valid_m + (last_m_block ? 0u : BLOCK_M); + current_pool_block_offset = payload >> (kVBits + 1u); + } else { + current_num_tokens = payload & ((1u << kTBits) - 1u); + current_pool_block_offset = payload >> kTBits; + } + return tag; + } + + CUTLASS_DEVICE static uint32_t ld_tile_table_payload(const uint2* entry) { + uint32_t ret; + asm volatile("ld.shared.u32 %0, [%1];" + : "=r"(ret) : "l"(__cvta_generic_to_shared(entry) + 4) : "memory"); + return ret; + } + + // Lag schedule with a dynamic tail. The static prefix is walked round robin as in `for_each_block`; the trailing L2-only + // units are handed out in cluster pairs through `counter`, a per-rank global atomic that the dispatch warps zero after the + // pre-combine barrier. Only the cluster leader fetches; it publishes its own entry and the partner's (into the partner's table + // over DSMEM, same index: the static prefix gives both CTAs of a pair the same tile count). Tail tiles wait only on L1 + // outputs, all issued in the static prefix, so any order of tail tiles is deadlock-free. Wave schedule: the tail is the L2 + // phase of the last wave. Table bound: the leader stops fetching tickets once its table cannot hold one more tile plus the + // end marker and the other clusters take the rest (together they hold kNumSMs * (capacity - 1) - P >= T - P tail slots). + template + CUTLASS_DEVICE void for_each_block_publish_dynamic_tail(TileEntry* table, uint32_t* counter, const uint32_t& cta_rank_in_cluster, + const uint32_t& capacity, Func&& func) { + DG_STATIC_ASSERT(kClusterSize == 1 or kClusterSize == 2, "dynamic tail: 1- or 2-CTA clusters"); + DG_STATIC_ASSERT(kClusterSize == kTileTableClusterSize, "dynamic tail: the tag-load scope must cover the publishing CTA"); + constexpr uint32_t kUnitM = layout::kSM90FusedLagUnitM; + fetch_expert_recv_count(); + set_expert_idx(0); + uint32_t num_published = 0; + while (true) { + CUTE_TIE_DECL(get_next_block(), block_phase, current_local_expert_idx, m_block_idx, n_block_idx); + if (block_phase == BlockPhase::None) + break; + if (ptx::get_lane_idx() == 0) + publish_tile_entry(table + num_published, block_phase == BlockPhase::Linear2 ? 2u : 1u); + ++ num_published; + func(block_phase, current_local_expert_idx, + block_phase == BlockPhase::Linear2 ? kNumL2BlockKs : kNumL1BlockKs, + m_block_idx, n_block_idx); + } + if (not tail_reached) { + if (ptx::get_lane_idx() == 0) + publish_tile_entry(table + num_published, 0u); + return; + } + if (cta_rank_in_cluster != 0) { + // partner: the leader publishes our tail entries + while (true) { + const uint32_t tag = load_tile_entry(table + num_published); + if (tag == 0u) + return; + ++ num_published; + func(BlockPhase::Linear2, current_local_expert_idx, kNumL2BlockKs, m_block_idx, n_block_idx); + } + } + // leader: the cursor is on the first tail unit (lag) / the last wave's first expert (wave) + uint32_t chunk_start = 0; // tail-relative position of the current chunk's first tile + while (true) { + // table bound (both entry formats): stop before fetching a ticket the table cannot hold together with the end marker + if (num_published + 2 > capacity) { + if (ptx::get_lane_idx() == 0) { + publish_tile_entry(table + num_published, 0u); + if constexpr (kClusterSize > 1) + st_remote_tile_entry(table + num_published, 1u, 0u, 0u); // end marker: tag word 0 + } + return; + } + uint32_t fetched = 0; + if (ptx::get_lane_idx() == 0) + fetched = atomicAdd(counter, 1u); + fetched = __shfl_sync(0xffffffffu, fetched, 0); + const uint32_t g = fetched * kClusterSize; + // advance the cursor to the chunk holding g; `um` = m-blocks of that chunk, `m0` its first m-block + uint32_t um = 0, m0 = 0; + bool at_end; + if constexpr (kL2LagUnits > 0) { + while (l2_expert < kNumExpertsPerRank) { + um = cute::min(kUnitM, math::ceil_div(l2_num_tokens, BLOCK_M) - l2_m0); + if (g < chunk_start + um * kNumL2BlockNs) + break; + chunk_start += um * kNumL2BlockNs; + l2_m0 += um; + lag_skip_empty(l2_expert, l2_num_tokens, l2_pool_offset, l2_m0); + } + at_end = l2_expert >= kNumExpertsPerRank; + current_local_expert_idx = l2_expert; + current_num_tokens = l2_num_tokens; + current_pool_block_offset = l2_pool_offset; + m0 = l2_m0; + } else { + while (current_local_expert_idx < kNumExpertsPerRank) { + um = get_current_num_m_blocks(); + if (g < chunk_start + um * kNumL2BlockNs) + break; + chunk_start += um * kNumL2BlockNs; + advance_expert_idx(); + } + at_end = current_local_expert_idx >= kNumExpertsPerRank; + m0 = 0; + } + if (at_end) { + if (ptx::get_lane_idx() == 0) { + publish_tile_entry(table + num_published, 0u); + if constexpr (kClusterSize > 1) + st_remote_tile_entry(table + num_published, 1u, 0u, 0u); // end marker: tag word 0 + } + return; + } + const uint32_t q = g - chunk_start; + // own tile (position q, even) and the partner's (q + 1), decoded with the chunk's m-inner / n-inner rule + uint32_t partner_m, partner_n; + bool partner_pair_valid; + if constexpr (kMulticastOnB) { + m_block_idx = m0 + q % um; + n_block_idx = q / um; + current_pair_valid = (q % um) + 1 < um; + partner_m = m0 + (q + 1) % um; + partner_n = (q + 1) / um; + partner_pair_valid = ((q + 1) % um) > 0; + } else { + m_block_idx = m0 + q / kNumL2BlockNs; + n_block_idx = q % kNumL2BlockNs; + current_pair_valid = false; + partner_m = m0 + (q + 1) / kNumL2BlockNs; + partner_n = (q + 1) % kNumL2BlockNs; + partner_pair_valid = false; + } + if (ptx::get_lane_idx() == 0) { + publish_tile_entry(table + num_published, 2u); + if constexpr (kClusterSize > 1) { + st_remote_tile_entry(table + num_published, 1u, + make_compact_tag_word(current_local_expert_idx, partner_m, partner_n, partner_pair_valid, 2u), + make_compact_payload_word(current_num_tokens, current_pool_block_offset, partner_m)); + } + } + ++ num_published; + func(BlockPhase::Linear2, current_local_expert_idx, kNumL2BlockKs, m_block_idx, n_block_idx); + } + } + + // Replays the published sequence: `l1_func(expert, num_k_blocks, m_block, n_block)` / `l2_func(...)` for every tile in order, + // with the getters valid inside. Returns the number of tiles replayed; `replay_next_idx` = the index of the entry after the current tile's. + uint32_t replay_next_idx = 0; + + template + CUTLASS_DEVICE uint32_t for_each_block_replay(const TileEntry* table, L1Func&& l1_func, L2Func&& l2_func) { + uint32_t i = 0; + while (true) { + const uint32_t tag = load_tile_entry(table + i); + if (tag == 0u) + return i; + ++ i; + replay_next_idx = i; + if (tag == 2u) { + l2_func(current_local_expert_idx, kNumL2BlockKs, m_block_idx, n_block_idx); + } else { + l1_func(current_local_expert_idx, kNumL1BlockKs, m_block_idx, n_block_idx); + } + } + } +}; + +} // namespace deep_gemm::sched diff --git a/deep_gemm/mega/__init__.py b/deep_gemm/mega/__init__.py index 3da1297cf4..5149f078d9 100644 --- a/deep_gemm/mega/__init__.py +++ b/deep_gemm/mega/__init__.py @@ -117,6 +117,75 @@ def destroy(self): self.group = None +# K granularity of the fused kernel's L2 activation scale factor (`l2_act_sf_gran_k`, 64 or 128); fixed when the buffer is +# sized. Default (None): per-128 K where the 256-wide decode tile fits, i.e. hidden and 2 x intermediate_hidden are multiples +# of 512 (the 2-CTA pairing needs even N block counts); otherwise per-64 K, whose 128-wide decode tile fits every +# hidden % 256 == 0. +def _default_l2_act_sf_gran_k_sm90_fused(hidden: int, intermediate_hidden: int) -> int: + return 128 if hidden % 512 == 0 and (2 * intermediate_hidden) % 512 == 0 else 64 + + +class SM90FusedSymmBuffer: + def __init__(self, group: dist.ProcessGroup, + num_experts: int, + num_max_tokens_per_rank: int, num_topk: int, + hidden: int, intermediate_hidden: int, + use_fp8_dispatch: bool = True, + activation: str = 'swiglu', + num_experts_per_wave: Optional[int] = None, + l2_act_sf_gran_k: Optional[int] = None): + self.group = group + self.num_experts = num_experts + self.num_max_tokens_per_rank = num_max_tokens_per_rank + self.num_topk = num_topk + self.hidden = hidden + self.intermediate_hidden = intermediate_hidden + # Wave-size knob; it selects the schedule the buffer is sized for (get_buffer_schedule_sm90_fused in + # csrc/jit_kernels/heuristics/sm90_fused_mega_moe.hpp): None (0) = the default, a lag-sized ring running the L2-lag schedule + # from 1024 tokens per rank and full-pool buffers below that; -1 = wave schedule, ring sized by the occupancy heuristic; + # N > 0 = wave schedule, ring sized for a fixed wave of N experts + self.num_experts_per_wave = 0 if num_experts_per_wave is None else num_experts_per_wave + self.l2_act_sf_gran_k = l2_act_sf_gran_k if l2_act_sf_gran_k is not None else \ + _default_l2_act_sf_gran_k_sm90_fused(hidden, intermediate_hidden) + + num_bytes, slice_input_buffers, num_ring_tokens, l2_lag_encoded = \ + _C.get_symm_buffer_size_for_sm90_fused_mega_moe( + group.size(), num_experts, + num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + use_fp8_dispatch, activation, + self.num_experts_per_wave, self.l2_act_sf_gran_k, + ) + # Ring capacity (0 = full pool) and encoded L2-lag schedule the buffer was sized for; passed back at every launch + self.num_ring_tokens = num_ring_tokens + self.l2_lag_encoded = l2_lag_encoded + allocator = torch if group.size() == 1 else symm_mem + self.buffer = allocator.empty(num_bytes, dtype=torch.int8, device='cuda') + self.handle = ( + types.SimpleNamespace(buffer_ptrs=[self.buffer.data_ptr()]) + if group.size() == 1 + else symm_mem.rendezvous(self.buffer, group=group) + ) + self.buffer.zero_() + self.group.barrier() + torch.cuda.synchronize() + + (self.x, self.x_sf, + self.topk_idx, self.topk_weights, + self.l1_acts, self.l1_acts_sf, + self.l2_acts, self.l2_acts_sf) = slice_input_buffers(self.buffer) + + def destroy(self): + self.handle = None + for name in ( + 'x', 'x_sf', 'topk_idx', 'topk_weights', + 'l1_acts', 'l1_acts_sf', 'l2_acts', 'l2_acts_sf', + ): + setattr(self, name, None) + self.buffer = None + self.group = None + + def get_symm_buffer_for_mega_moe(group: dist.ProcessGroup, num_experts: int, num_max_tokens_per_rank: int, num_topk: int, @@ -150,7 +219,23 @@ def get_symm_buffer_for_sm90_mega_moe(group: dist.ProcessGroup, num_max_tokens_per_rank: int, num_topk: int, hidden: int, intermediate_hidden: int, use_fp8_dispatch: bool = True, - activation: str = 'swiglu') -> SM90SymmBuffer: + activation: str = 'swiglu', + fused: bool = False, + num_experts_per_wave: Optional[int] = None, + l2_act_sf_gran_k: Optional[int] = None + ) -> Union[SM90SymmBuffer, SM90FusedSymmBuffer]: + if fused: + num_max_tokens_per_rank = align( + num_max_tokens_per_rank, _C.get_token_alignment_for_sm90_fused_mega_moe()) + return SM90FusedSymmBuffer( + group, num_experts, + num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, + use_fp8_dispatch, activation, + num_experts_per_wave, l2_act_sf_gran_k, + ) + if num_experts_per_wave is not None or l2_act_sf_gran_k is not None: + raise ValueError('`num_experts_per_wave` and `l2_act_sf_gran_k` apply to the fused buffer only') num_max_tokens_per_rank = align( num_max_tokens_per_rank, _C.get_token_alignment_for_sm90_mega_moe()) return SM90SymmBuffer( @@ -389,19 +474,44 @@ def bf16_mega_moe(y: torch.Tensor, def fp8_mega_moe(y: torch.Tensor, l1_weights: Tuple[torch.Tensor, torch.Tensor], l2_weights: Tuple[torch.Tensor, torch.Tensor], - sym_buffer: SM90SymmBuffer, + sym_buffer: Union[SM90SymmBuffer, SM90FusedSymmBuffer], cumulative_local_expert_recv_stats: Optional[torch.Tensor] = None, recipe: Tuple[int, int, int] = (128, 128, 128), activation: str = 'swiglu', activation_clamp: Optional[float] = None, - fast_math: bool = True): + fast_math: bool = True, + *, + num_tokens_bound: Optional[int] = None): """SM90 (Hopper) MegaMoE entry point. Expects FP8 e4m3 weights and block-(128, 128) float scale factors. The weight SF layout matches the convention used by ``DeepSeekV4FlashFp8`` / DeepEP, so the same SF tensors can be physically shared between the DeepEP path and this kernel. + + An ``SM90FusedSymmBuffer`` selects the single-kernel implementation. + ``num_tokens_bound`` (fused only) is an upper bound on this call's + per-rank token count on every rank; ``None`` uses the buffer capacity. """ + if isinstance(sym_buffer, SM90FusedSymmBuffer): + _C.sm90_fused_fp8_mega_moe( + y, + l1_weights, l2_weights, + cumulative_local_expert_recv_stats, + sym_buffer.buffer, + sym_buffer.handle.buffer_ptrs, sym_buffer.group.rank(), + sym_buffer.num_max_tokens_per_rank, + sym_buffer.num_experts, sym_buffer.num_topk, + recipe, + activation, activation_clamp, + fast_math, + sym_buffer.num_ring_tokens, + 0 if num_tokens_bound is None else num_tokens_bound, + sym_buffer.l2_lag_encoded, sym_buffer.l2_act_sf_gran_k + ) + return + if num_tokens_bound is not None: + raise ValueError('`num_tokens_bound` applies to the fused buffer only') _C.fp8_mega_moe( y, l1_weights, l2_weights, diff --git a/tests/bench_mega_moe_sm90.py b/tests/bench_mega_moe_sm90.py index 90e9ed4dfe..74ccef0396 100644 --- a/tests/bench_mega_moe_sm90.py +++ b/tests/bench_mega_moe_sm90.py @@ -1,4 +1,4 @@ -"""Benchmark the SM90 FP8 MegaMoE kernel on Flash and Pro model shapes.""" +"""Benchmark the SM90 FP8 MegaMoE kernels (split L1/L2, fused) on Flash and Pro model shapes.""" import argparse import json @@ -40,6 +40,11 @@ 'sm90_fp8_mega_moe_l1_impl', 'sm90_fp8_mega_moe_l2_impl', ) +FUSED_KERNEL_NAMES = ('sm90_fp8_fused_mega_moe_impl',) +ARM_KERNEL_NAMES = { + 'split': PHASE_KERNEL_NAMES, + 'fused': FUSED_KERNEL_NAMES, +} def _stable_seed(name: str) -> int: @@ -96,14 +101,6 @@ def _benchmark_case( + _stable_seed(f'{model_name}:{num_tokens}') ) torch.manual_seed(case_seed) - buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( - group, - num_experts, - args.num_max_tokens_per_rank, - num_topk, - hidden, - intermediate_hidden, - ) x_bf16 = torch.randn( (num_tokens, hidden), dtype=torch.bfloat16, device='cuda', @@ -145,95 +142,139 @@ def _benchmark_case( (num_tokens, hidden), dtype=torch.bfloat16, device='cuda', ) - def run_sm90() -> torch.Tensor: - buffer.x[:num_tokens].copy_(x_fp8) - buffer.x_sf[:num_tokens].copy_(x_sf) - buffer.topk_idx[:num_tokens].copy_(topk_idx) - buffer.topk_weights[:num_tokens].copy_(topk_weights) - deep_gemm.fp8_mega_moe( - y, - transformed_l1, - transformed_l2, - buffer, - cumulative_local_expert_recv_stats=cumulative_recv_stats, - recipe=(128, 128, 128), - activation='swiglu', - activation_clamp=args.activation_clamp, - fast_math=bool(args.fast_math), + # The arms share inputs and weights; each gets its own symmetric buffer. + medians: Dict[str, float] = {} + symm_bytes: Dict[str, int] = {} + for arm in args.arms: + # None keeps the package default granularity + buffer_kwargs = dict(fused=True) if arm == 'fused' else {} + if arm == 'fused' and args.l2_act_sf_gran_k is not None: + buffer_kwargs['l2_act_sf_gran_k'] = args.l2_act_sf_gran_k + buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( + group, + num_experts, + args.num_max_tokens_per_rank, + num_topk, + hidden, + intermediate_hidden, + **buffer_kwargs, ) - return y + symm_bytes[arm] = buffer.buffer.numel() + + # Every rank here runs the same count, so that count is also the bound a caller would + # pass as the maximum over its step. + launch_kwargs = dict(num_tokens_bound=num_tokens) if arm == 'fused' else {} + + def run_sm90() -> torch.Tensor: + buffer.x[:num_tokens].copy_(x_fp8) + buffer.x_sf[:num_tokens].copy_(x_sf) + buffer.topk_idx[:num_tokens].copy_(topk_idx) + buffer.topk_weights[:num_tokens].copy_(topk_weights) + deep_gemm.fp8_mega_moe( + y, + transformed_l1, + transformed_l2, + buffer, + cumulative_local_expert_recv_stats=cumulative_recv_stats, + recipe=(128, 128, 128), + activation='swiglu', + activation_clamp=args.activation_clamp, + fast_math=bool(args.fast_math), + **launch_kwargs, + ) + return y + + if args.ncu_profile_only: + dist_print( + f'[NCU] model={model_name} M={num_tokens} arm={arm}', once_in_node=True, + ) + run_sm90() + torch.cuda.synchronize() + dist.barrier(group=group) + buffer.destroy() + continue + + repeats = args.repeats + if repeats is None: + repeats = args.small_repeats if num_tokens <= 128 else args.large_repeats - if args.ncu_profile_only: - dist_print( - f'[NCU] model={model_name} M={num_tokens}', once_in_node=True, - ) run_sm90() torch.cuda.synchronize() dist.barrier(group=group) - buffer.destroy() - return - repeats = args.repeats - if repeats is None: - repeats = args.small_repeats if num_tokens <= 128 else args.large_repeats - - run_sm90() - torch.cuda.synchronize() - dist.barrier(group=group) - - rank0_observations = [] - max_rank_observations = [] - for repeat in range(repeats): - phase_times = bench_kineto( - run_sm90, - PHASE_KERNEL_NAMES, - barrier=lambda: dist.barrier(group=group), - num_tests=args.num_tests, - suppress_kineto_output=True, - ) - local_time = sum(phase_times) - max_rank_time = torch.tensor(local_time, dtype=torch.float64, device='cuda') - dist.all_reduce(max_rank_time, op=dist.ReduceOp.MAX, group=group) + rank0_observations = [] + max_rank_observations = [] + for repeat in range(repeats): + phase_times = bench_kineto( + run_sm90, + ARM_KERNEL_NAMES[arm], + barrier=lambda: dist.barrier(group=group), + num_tests=args.num_tests, + suppress_kineto_output=True, + ) + local_time = sum(phase_times) + max_rank_time = torch.tensor(local_time, dtype=torch.float64, device='cuda') + dist.all_reduce(max_rank_time, op=dist.ReduceOp.MAX, group=group) + + rank0_observations.append(local_time) + max_rank_observations.append(max_rank_time.item()) + if rank_idx == 0: + observation = { + 'model': model_name, + 'm': num_tokens, + 'arm': arm, + 'repeat': repeat, + 'rank0_us': local_time * 1e6, + 'max_rank_us': max_rank_time.item() * 1e6, + 'num_tests': args.num_tests, + 'num_max_tokens_per_rank': args.num_max_tokens_per_rank, + 'seed': args.seed, + } + if arm == 'split': + observation['l1_rank0_us'] = phase_times[0] * 1e6 + observation['l2_rank0_us'] = phase_times[1] * 1e6 + else: + observation['fused_rank0_us'] = phase_times[0] * 1e6 + print('BENCH_OBS_JSON ' + json.dumps(observation, sort_keys=True), flush=True) - rank0_observations.append(local_time) - max_rank_observations.append(max_rank_time.item()) if rank_idx == 0: - print('BENCH_OBS_JSON ' + json.dumps({ + median_time = statistics.median(max_rank_observations) + medians[arm] = median_time + arm_label = f' {arm}' if len(args.arms) > 1 else '' + print( + f'[{model_name:5s}] M={num_tokens:4d}{arm_label} obs={repeats:2d} ' + f'max-rank median={median_time * 1e6:8.1f} us ' + f'range={min(max_rank_observations) * 1e6:.1f}-' + f'{max(max_rank_observations) * 1e6:.1f} us', + flush=True, + ) + print('BENCH_SUMMARY_JSON ' + json.dumps({ 'model': model_name, 'm': num_tokens, - 'repeat': repeat, - 'rank0_us': local_time * 1e6, - 'max_rank_us': max_rank_time.item() * 1e6, - 'l1_rank0_us': phase_times[0] * 1e6, - 'l2_rank0_us': phase_times[1] * 1e6, + 'arm': arm, + 'observations': repeats, + 'rank0_median_us': statistics.median(rank0_observations) * 1e6, + 'max_rank_median_us': median_time * 1e6, + 'max_rank_min_us': min(max_rank_observations) * 1e6, + 'max_rank_max_us': max(max_rank_observations) * 1e6, 'num_tests': args.num_tests, 'num_max_tokens_per_rank': args.num_max_tokens_per_rank, - 'seed': args.seed, + 'num_tokens_bound': launch_kwargs.get('num_tokens_bound', 0), + 'symm_buffer_bytes': symm_bytes[arm], }, sort_keys=True), flush=True) - if rank_idx == 0: - median_time = statistics.median(max_rank_observations) + dist.barrier(group=group) + buffer.destroy() + + if rank_idx == 0 and len(medians) > 1: + split_us = medians['split'] * 1e6 + fused_us = medians['fused'] * 1e6 print( - f'[{model_name:5s}] M={num_tokens:4d} obs={repeats:2d} ' - f'max-rank median={median_time * 1e6:8.1f} us ' - f'range={min(max_rank_observations) * 1e6:.1f}-' - f'{max(max_rank_observations) * 1e6:.1f} us', + f'[{model_name:5s}] M={num_tokens:4d} split={split_us:8.1f} us fused={fused_us:8.1f} us ' + f'fused/split={fused_us / split_us:.3f} ' + f'symm MiB split={symm_bytes["split"] / 2 ** 20:.1f} fused={symm_bytes["fused"] / 2 ** 20:.1f}', flush=True, ) - print('BENCH_SUMMARY_JSON ' + json.dumps({ - 'model': model_name, - 'm': num_tokens, - 'observations': repeats, - 'rank0_median_us': statistics.median(rank0_observations) * 1e6, - 'max_rank_median_us': median_time * 1e6, - 'max_rank_min_us': min(max_rank_observations) * 1e6, - 'max_rank_max_us': max(max_rank_observations) * 1e6, - 'num_tests': args.num_tests, - 'num_max_tokens_per_rank': args.num_max_tokens_per_rank, - }, sort_keys=True), flush=True) - - dist.barrier(group=group) - buffer.destroy() def _benchmark_worker( @@ -292,6 +333,14 @@ def _parse_args() -> argparse.Namespace: parser.add_argument('--activation-clamp', type=float, default=10.0) parser.add_argument('--fast-math', type=int, choices=[0, 1], default=1) parser.add_argument('--ncu-profile-only', action='store_true') + parser.add_argument( + '--arms', nargs='+', choices=sorted(ARM_KERNEL_NAMES), default=['split'], + help='Implementations to time per case; both print side by side.', + ) + parser.add_argument( + '--l2-act-sf-gran-k', type=int, choices=[64, 128], default=None, + help='L2 activation scale-factor K granularity of the fused arm (default: the package default).', + ) args = parser.parse_args() assert args.num_processes > 0 @@ -301,6 +350,7 @@ def _parse_args() -> argparse.Namespace: assert args.repeats is None or args.repeats > 0 assert args.num_tests > 0 assert 0 <= args.masked_ratio <= 1 + assert len(set(args.arms)) == len(args.arms) return args diff --git a/tests/test_mega_moe_sm90.py b/tests/test_mega_moe_sm90.py index 52e5f448af..af16c9154c 100644 --- a/tests/test_mega_moe_sm90.py +++ b/tests/test_mega_moe_sm90.py @@ -33,7 +33,7 @@ import sys import torch import torch.distributed as dist -from typing import Tuple, List, Dict, Any +from typing import Tuple, List, Dict, Any, Optional REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) if REPO_ROOT not in sys.path: @@ -115,6 +115,7 @@ def _reference_fused( num_experts: int, num_topk: int, hidden: int, intermediate_hidden: int, activation_clamp: float, + l2_act_sf_gran_k: int, ) -> torch.Tensor: """Reference: returns (num_tokens, hidden) bf16 result for *this* rank. @@ -192,11 +193,11 @@ def _reference_fused( # SwiGLU + clamp + multiply by topk weight l1_y = _swiglu_fp32(l1_y, activation_clamp) * weights.unsqueeze(-1) # (S, IH) - # Per-row, per-64-col FP8 quantize -> dequantize + # Per-row, per-`l2_act_sf_gran_k`-col FP8 quantize -> dequantize s_, ih = l1_y.shape - assert ih == intermediate_hidden and ih % 64 == 0 - l1_view = l1_y.view(s_, ih // 64, 64) - amax = l1_view.abs().amax(dim=-1).clamp(1e-4) # (S, IH/64) + assert ih == intermediate_hidden and ih % l2_act_sf_gran_k == 0 + l1_view = l1_y.view(s_, ih // l2_act_sf_gran_k, l2_act_sf_gran_k) + amax = l1_view.abs().amax(dim=-1).clamp(1e-4) # (S, IH/gran_k) sf2 = amax / 448.0 l1_q = (l1_view / sf2.unsqueeze(-1)).to(torch.float8_e4m3fn).float() l2_in = (l1_q * sf2.unsqueeze(-1)).view(s_, ih) # (S, IH) fp32 @@ -226,6 +227,8 @@ def _run_scenario( cfg: Dict[str, Any], rank_idx: int, num_ranks: int, group: dist.ProcessGroup, diff_tol: float, + fused: bool = False, + l2_act_sf_gran_k: Optional[int] = None, ): num_max = cfg['num_max_tokens_per_rank'] num_tokens = cfg.get('num_tokens', num_max) @@ -282,10 +285,15 @@ def _trace(stage: str): # ---- Allocate symm buffer ----------------------------------------------- _trace('alloc_symm_buffer') + # None keeps the package default; the reference reads the granularity back from the buffer + buffer_kwargs = dict(fused=True) if fused else {} + if fused and l2_act_sf_gran_k is not None: + buffer_kwargs['l2_act_sf_gran_k'] = l2_act_sf_gran_k buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( group, num_experts, num_max, num_topk, hidden, intermediate_hidden, + **buffer_kwargs, ) cum_stats = torch.zeros(num_experts_per_rank, dtype=torch.int, device='cuda') @@ -322,6 +330,8 @@ def _trace(stage: str): num_experts, num_topk, hidden, intermediate_hidden, activation_clamp, + # split kernel: per-64-K L2 activation scales; fused kernel: the granularity the buffer was sized with + buffer.l2_act_sf_gran_k if fused else 64, ) diff = calc_diff(y_fused, y_ref) @@ -498,14 +508,25 @@ def _test_worker(local_rank: int, num_local_ranks: int, args: argparse.Namespace if args.filter: layers = [(n, c) for n, c in layers if args.filter in n] + if args.fused: + # The fused kernel tiles the intermediate dimension in at most 64 N blocks. + limit = 64 * 64 + skipped = [n for n, c in layers if c['intermediate_hidden'] > limit] + layers = [(n, c) for n, c in layers if c['intermediate_hidden'] <= limit] + if skipped: + dist_print(f' [SKIP] beyond the fused intermediate_hidden limit: ' + f'{", ".join(skipped)}', once_in_node=True) + + impl = f'fused (l2_act_sf_gran_k={args.l2_act_sf_gran_k or "package default"})' if args.fused else 'split' dist_print(f'SM90 MegaMoE test plan: {len(layers)} scenarios across ' - f'layers {sorted(args.layers)} on {num_ranks} ranks', + f'layers {sorted(args.layers)} on {num_ranks} ranks, {impl} kernels', once_in_node=True) failures: List[str] = [] for name, cfg in layers: try: - _run_scenario(name, cfg, rank_idx, num_ranks, group, diff_tol) + _run_scenario(name, cfg, rank_idx, num_ranks, group, diff_tol, + args.fused, args.l2_act_sf_gran_k) except AssertionError as ex: dist_print(f' [{name}] FAIL: {ex}', once_in_node=True) failures.append(name) @@ -540,6 +561,10 @@ def _test_worker(local_rank: int, num_local_ranks: int, args: argparse.Namespace help='calc_diff tolerance (default: 0.01)') parser.add_argument('--fail-fast', action='store_true', help='Stop on first failing scenario') + parser.add_argument('--fused', action='store_true', + help='Run the single-kernel SM90 fused implementation') + parser.add_argument('--l2-act-sf-gran-k', type=int, choices=[64, 128], default=None, + help='L2 activation scale-factor K granularity of the fused kernel (default: the package default)') args = parser.parse_args() np = args.num_processes diff --git a/tests/test_mega_moe_sm90_fused_ring.py b/tests/test_mega_moe_sm90_fused_ring.py new file mode 100644 index 0000000000..fd8d8816ed --- /dev/null +++ b/tests/test_mega_moe_sm90_fused_ring.py @@ -0,0 +1,233 @@ +"""SM90 fused MegaMoE ring-buffer check. + +For every shape the symm buffers of three layouts must produce the same output bit for bit -- the default layout (a lag-scheduled +ring from 1024 tokens per rank, the full pool below that), a ring sized for one automatically chosen expert wave +(``num_experts_per_wave=-1``) and, where the default is a ring, the full pool under the wave schedule +(``num_experts_per_wave=num_experts_per_rank``, the legacy layout) -- and the first launch must match the reference of +test_mega_moe_sm90.py within its tolerance. Every buffer runs four launches: the full routing, a smaller token count with a +skewed routing, the full routing again, and a decode-sized one, so both launch-parity banks and the ring wrap are exercised on one +allocation. The two-rank shape is the configuration whose requested lag ring exceeds the full pool (the buffer is the pool itself +under the lag schedule). A fourth buffer of the default layout then repeats all four launches with each call's own +``num_tokens_bound`` instead of the buffer capacity, which must not change a single bit either: on a lag ring that holds one wave +of that bound, the decode-sized launch runs the wave order under the bound where it runs the lag order without one, and on the +top-6 shape the smaller launch runs a shorter lag under its bound. A bound above the capacity must be rejected instead. +""" + +import argparse +import os +import sys + +import torch +import torch.distributed as dist + +REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +if REPO_ROOT not in sys.path: + sys.path.insert(0, REPO_ROOT) + +import deep_gemm +from deep_gemm.utils import per_token_cast_to_fp8 +from deep_gemm.utils.dist import dist_print, init_dist +from deep_gemm.testing import calc_diff, get_arch_major + +from test_mega_moe_sm90 import _quantize_grouped_fp8_block_128_128, _reference_fused + +ACTIVATION_CLAMP = 10.0 +# tokens per rank from which the default layout is a lag ring (heuristics kSm90FusedAutoLagMinTokens) +AUTO_LAG_MIN_TOKENS = 1024 +# per-rank token bound up to which a call on a lag ring takes the wave schedule instead (heuristics kSm90FusedCallWaveMaxTokens) +CALL_WAVE_MAX_TOKENS = 256 +TOKEN_ALIGNMENT = 128 + + +def _shapes(num_ranks): + out = [ + (f'ring.h1024.t{tokens}', dict(hidden=1024, intermediate_hidden=1024, num_experts=8 * num_ranks, num_topk=2, num_tokens=tokens)) + for tokens in (64, 256, 512, 1024, 2048) + ] + out.append(('ring.h4096.experts288.t2048', dict(hidden=4096, intermediate_hidden=2048, num_experts=288, num_topk=8, num_tokens=2048))) + # top-6 at a 5120-token capacity: the buffer's lag spans four 1024-token batches per local expert and the smaller launch's + # own bound spans three, so here a bound shortens the lag instead of switching the order, which no other shape's bound does + out.append(('ring.h1024.topk6.t5120', dict(hidden=1024, intermediate_hidden=1024, num_experts=8 * num_ranks, num_topk=6, num_tokens=5120))) + # two ranks, 8 local experts: the sizer asks for a 12800-token lag ring, the full pool is 9216 tokens -> clamped to the pool + out.append(('ring.h1024.t2048.ep2', dict(hidden=1024, intermediate_hidden=1024, num_experts=16, num_topk=2, num_tokens=2048, world=2))) + return out + + +def _align(x, a): + return (x + a - 1) // a * a + + +def _full_pool_tokens(num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank): + # layout::get_num_max_pool_tokens_sm90 (one partial block of kSM90FusedMaxCandidateBlockM 128 per expert, kSM90FusedLCMBlockM 128) + return _align(num_ranks * num_max_tokens_per_rank * min(num_topk, num_experts_per_rank) + num_experts_per_rank * 127, 128) + + +def _routing(num_tokens, num_experts, num_topk, seed, num_hot_experts=0): + gen = torch.Generator(device='cuda') + gen.manual_seed(seed) + scores = torch.randn((num_tokens, num_experts), dtype=torch.float, device='cuda', generator=gen) + if num_hot_experts: + # skewed routing: a few experts take most rows, the others a handful + scores[:, :num_hot_experts] += 4.0 + topk_weights, topk_idx = torch.topk(scores, num_topk, dim=-1, largest=True, sorted=False) + return topk_idx, topk_weights + + +def _run_shape(name, cfg, rank_idx, num_ranks, group, diff_tol, l2_act_sf_gran_k): + """Runs one shape on `group`; returns (ok, message). No assertion and no collective outside `group`.""" + hidden, intermediate_hidden = cfg['hidden'], cfg['intermediate_hidden'] + num_experts, num_topk, num_tokens = cfg['num_experts'], cfg['num_topk'], cfg['num_tokens'] + num_experts_per_rank = num_experts // num_ranks + num_max_tokens_per_rank = _align(num_tokens, TOKEN_ALIGNMENT) + seed = rank_idx * 1000 + sum(map(ord, name)) + torch.manual_seed(seed) + + x_bf16 = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device='cuda') + l1_weights_bf16 = torch.randn((num_experts_per_rank, intermediate_hidden * 2, hidden), dtype=torch.bfloat16, device='cuda') * 0.05 + l2_weights_bf16 = torch.randn((num_experts_per_rank, hidden, intermediate_hidden), dtype=torch.bfloat16, device='cuda') * 0.05 + x_fp8 = per_token_cast_to_fp8(x_bf16, use_ue8m0=False, gran_k=128, use_packed_ue8m0=False) + l1_weights = _quantize_grouped_fp8_block_128_128(l1_weights_bf16) + l2_weights = _quantize_grouped_fp8_block_128_128(l2_weights_bf16) + transformed_l1, transformed_l2 = deep_gemm.transform_weights_for_mega_moe_sm90(l1_weights, l2_weights) + del l1_weights_bf16, l2_weights_bf16 + + # launch 1 and 3: every token, the same routing (3 runs on the parity bank of launch 1 again); launch 2: fewer tokens, skewed + # routing (the smaller call may select a different tile class on the same buffer; the checks do not depend on it) + launches = [(num_tokens,) + _routing(num_tokens, num_experts, num_topk, seed)] + num_tokens_2 = max(1, num_tokens * 3 // 4) + launches.append((num_tokens_2,) + _routing(num_tokens_2, num_experts, num_topk, seed + 1, num_hot_experts=2)) + launches.append(launches[0]) + num_tokens_4 = min(CALL_WAVE_MAX_TOKENS, num_tokens) + launches.append((num_tokens_4,) + _routing(num_tokens_4, num_experts, num_topk, seed + 2)) + + # full pool -> the default layout is a lag ring (or the pool itself when the ring would exceed it), else the pool + full_pool = _full_pool_tokens(num_ranks, num_max_tokens_per_rank, num_topk, num_experts_per_rank) + default_is_ring = num_max_tokens_per_rank >= AUTO_LAG_MIN_TOKENS + arms = (None, -1, num_experts_per_rank) if default_is_ring else (None, -1) + + def run_launches(buffer, bounds): + ys = [] + for (n, topk_idx, topk_weights), bound in zip(launches, bounds): + buffer.x[:n].copy_(x_fp8[0][:n]) + buffer.x_sf[:n].copy_(x_fp8[1][:n]) + buffer.topk_idx[:n].copy_(topk_idx) + buffer.topk_weights[:n].copy_(topk_weights) + y = torch.empty((n, hidden), dtype=torch.bfloat16, device='cuda') + deep_gemm.fp8_mega_moe(y, transformed_l1, transformed_l2, buffer, recipe=(128, 128, 128), activation='swiglu', + activation_clamp=ACTIVATION_CLAMP, fast_math=True, num_tokens_bound=bound) + torch.cuda.synchronize() + ys.append(y) + return ys + + outputs, caps = {}, {} + for num_experts_per_wave in arms: + buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( + group, num_experts, num_tokens, num_topk, hidden, intermediate_hidden, + fused=True, num_experts_per_wave=num_experts_per_wave, l2_act_sf_gran_k=l2_act_sf_gran_k) + outputs[num_experts_per_wave] = run_launches(buffer, [None] * len(launches)) + # the buffer reports the full pool as a ring capacity of 0 + caps[num_experts_per_wave] = (buffer.num_ring_tokens or None, buffer.l2_lag_encoded, buffer.l2_act_sf_gran_k) + buffer.destroy() + dist.barrier(group=group) + + # the default layout again, this time telling the kernel each call's own bound instead of letting it assume the capacity + buffer = deep_gemm.get_symm_buffer_for_sm90_mega_moe( + group, num_experts, num_tokens, num_topk, hidden, intermediate_hidden, + fused=True, l2_act_sf_gran_k=l2_act_sf_gran_k) + bounds = [_align(n, TOKEN_ALIGNMENT) for n, _, _ in launches] + bounded_outputs = run_launches(buffer, bounds) + # a bound the buffer cannot receive is rejected on the host, before any rank launches + try: + run_launches(buffer, [num_max_tokens_per_rank + TOKEN_ALIGNMENT] * len(launches)) + rejected = False + except Exception: # noqa: BLE001 -- the host assertion arrives as whatever the FFI layer wraps it in + rejected = True + buffer.destroy() + dist.barrier(group=group) + + # capacities: the default arm is a lag ring (at most the pool) or the pool, the auto-wave arm a ring, the control arm the pool + ring_default, lag_default, gran_k = caps[None] + ring_auto = caps[-1][0] + problems = [] + if default_is_ring: + if ring_default is None or ring_default > full_pool or lag_default <= 0: + problems.append(f'default layout ring={ring_default} lag={lag_default} (expected a lag ring <= pool {full_pool})') + if caps[num_experts_per_rank][0] != full_pool or caps[num_experts_per_rank][1] != 0: + problems.append(f'full-pool control ring={caps[num_experts_per_rank][0]} lag={caps[num_experts_per_rank][1]} (expected pool {full_pool}, no lag)') + elif ring_default is not None: + problems.append(f'default layout ring={ring_default} (expected the full pool below {AUTO_LAG_MIN_TOKENS} tokens)') + if ring_auto is None or ring_auto > full_pool: + problems.append(f'auto-wave ring={ring_auto} (expected a ring <= pool {full_pool})') + + # every arm bit-identical to the default arm on every launch; launch 3 bit-identical to launch 1 within every arm + bitwise = all(torch.equal(outputs[arm][i], outputs[None][i]) for arm in arms[1:] for i in range(len(launches))) + replay = all(torch.equal(outputs[arm][0], outputs[arm][2]) for arm in arms) + bounded = all(torch.equal(bounded_outputs[i], outputs[None][i]) for i in range(len(launches))) + topk_idx_1, topk_weights_1 = launches[0][1], launches[0][2] + y_ref = _reference_fused( + x_fp8[0], x_fp8[1], topk_idx_1, topk_weights_1, + l1_weights[0], l1_weights[1], l2_weights[0], l2_weights[1], + rank_idx, num_ranks, group, num_experts, num_topk, hidden, intermediate_hidden, ACTIVATION_CLAMP, gran_k) + diff = calc_diff(outputs[None][0], y_ref) + ok = bitwise and replay and bounded and rejected and diff < diff_tol and not problems + control = f' pool={caps[num_experts_per_rank][0]}' if default_is_ring else '' + msg = (f"ring={ring_auto} default={ring_default or 'full'}{control} lag={lag_default} sf_gran_k={gran_k} " + f"bitwise={'OK' if bitwise else 'FAIL'} replay={'OK' if replay else 'FAIL'} " + f"bounds={bounds}<=cap{num_max_tokens_per_rank}:{'OK' if bounded else 'FAIL'} " + f"over_cap_rejected={'OK' if rejected else 'FAIL'} diff={diff:.4f} (tol={diff_tol:.2f})" + + (' ' + '; '.join(problems) if problems else '')) + return ok, msg + + +def test(local_rank, num_local_ranks, args): + rank_idx, num_ranks, group = init_dist(local_rank, num_local_ranks) + if get_arch_major() != 9: + dist_print(f'[SKIP] test_mega_moe_sm90_fused_ring requires SM90; got SM{get_arch_major()}0', once_in_node=True) + dist.destroy_process_group() + return + + # the two-rank shapes run on ranks 0 and 1 (every rank takes part in creating the group) + group_2 = dist.new_group(ranks=[0, 1]) if num_ranks > 2 else None + + failures = [] + for name, cfg in _shapes(num_ranks): + world = cfg.get('world', num_ranks) + if world > num_ranks or cfg['num_experts'] % world != 0: + dist_print(f" [{name:<24}] SKIP ({cfg['num_experts']} experts on {world} ranks, {num_ranks} available)", once_in_node=True) + continue + try: + if world == num_ranks: + ok, msg = _run_shape(name, cfg, rank_idx, num_ranks, group, args.diff_tol, args.l2_act_sf_gran_k) + elif rank_idx < world: + ok, msg = _run_shape(name, cfg, rank_idx, world, group_2, args.diff_tol, args.l2_act_sf_gran_k) + else: + ok, msg = True, '' + except Exception as ex: # noqa: BLE001 -- a shape that raises on every rank is reported and the loop goes on; a raise on one rank alone still hangs on the collectives it skipped, as it would without the handler + ok, msg = False, f'exception: {ex!r}' + # one verdict for all ranks before anyone prints or moves on: a rank-local failure must not skip a collective + flag = torch.tensor([0 if ok else 1], dtype=torch.int32, device='cuda') + dist.all_reduce(flag, op=dist.ReduceOp.MAX) + messages = [None] * num_ranks + dist.all_gather_object(messages, msg if not ok else '') + ok_all = flag.item() == 0 + failed_ranks = [r for r, m in enumerate(messages) if m] + dist_print(f" [{name:<24}] {msg if rank_idx == 0 and msg else ''} {'OK' if ok_all else 'FAIL'}" + + (f' ranks {failed_ranks}: {messages[failed_ranks[0]]}' if failed_ranks else ''), once_in_node=True) + if not ok_all: + failures.append(name) + dist.barrier() + dist.destroy_process_group() + if failures: + sys.exit(1) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='SM90 fused MegaMoE ring-buffer bitwise check') + parser.add_argument('--num-processes', type=int, default=8, help='Number of spawned processes, one per GPU') + parser.add_argument('--diff-tol', type=float, default=0.07, help='calc_diff tolerance against the reference; default: 0.07') + parser.add_argument( + '--l2-act-sf-gran-k', type=int, choices=[64, 128], default=None, + help='L2 activation-scale K granularity the symm buffers are built with; default: the package default', + ) + args = parser.parse_args() + torch.multiprocessing.spawn(test, args=(args.num_processes, args), nprocs=args.num_processes)