diff --git a/README.md b/README.md index 6ef705ffce..0d65932888 100644 --- a/README.md +++ b/README.md @@ -26,7 +26,7 @@ Despite its lightweight design, DeepGEMM's performance matches or exceeds expert ### Requirements -- NVIDIA SM90 or SM100 architecture GPU +- NVIDIA SM90, SM100, or SM103 architecture GPU - Python 3.8 or higher - Compilers with C++20 support - CUDA Toolkit: @@ -139,6 +139,38 @@ deep_gemm.fp8_fp4_mega_moe(y, transformed_l1, transformed_l2, buffer) For the full example with multi-process setup and benchmarking, please refer to `tests/test_mega_moe.py`. +#### SM103 FP8-block128 Mega MoE training + +`fp8_block128_mega_moe` is a separate SM103-only training backend. It accepts +compact BF16 source tokens, global top-k expert IDs and FP32 route scores, and +GLM-style E4M3 expert weights with FP32 inverse scales on exact 128 x 128 +blocks in canonical interleaved `[gate, up]` storage. The W13 launch presents +`[up; gate]` to fused SwiGLU without copying those weights, and accepts the +existing BF16 gate/down/up masters as three separate tensors. The operation +owns activation quantization, expert-parallel transport, W13/SwiGLU/W2, +post-down route scaling and combine, and the complete routed backward. It +returns BF16 input and master-weight gradients and FP32 route-score gradients +through autograd. + +Multi-rank execution uses one internally owned expanded +`deep_ep.ElasticBuffer` route. Expert segments are aligned to 128 rows while +the device PSUM retains each expert's real end. The SM103 companion zeros only +undefined padding before grouped BF16-master reduction, explicitly zeros empty +expert gradients, and uses deterministic routing with one combine reduction +for bitwise activation-checkpoint replay. There is no outer dispatch/combine, +dynamic `torch.distributed` split exchange, expert-weight dequantization, +UE8M0 requantization, or MXFP4 transcode. The single-rank specialization is +only for focused kernel validation and is not a distributed transport +fallback. + +The backend requires CUDA compute capability exactly 10.3. It has no SM100, +SM90, generic, or compatibility fallback. Use +`get_fp8_block128_mega_moe_capabilities()` for a non-launching capability +manifest and `deep_gemm.__git_commit__` for the full source OID embedded in the +native extension. Reference, adversarial, distributed, and performance +examples live in `tests/test_fp8_block128_*` and +`tests/benchmark_fp8_block128_mega_moe.py`. + #### Utilities The library provides some utility functions besides the above kernels: diff --git a/csrc/apis/gemm.hpp b/csrc/apis/gemm.hpp index 991eabca11..2b8511cf77 100644 --- a/csrc/apis/gemm.hpp +++ b/csrc/apis/gemm.hpp @@ -586,7 +586,7 @@ static void k_grouped_bf16_gemm_tn_contiguous(const torch::Tensor& a, DG_HOST_ASSERT(a.is_contiguous()); DG_HOST_ASSERT(b.is_contiguous()); DG_HOST_ASSERT(d.is_contiguous()); - DG_HOST_ASSERT(c.has_value() and c.value().is_contiguous()); + DG_HOST_ASSERT(not c.has_value() or c.value().is_contiguous()); // Early return for trivial cases if (early_return(m, n, sum_k, d, c)) diff --git a/csrc/apis/sm103_fp8_block128.hpp b/csrc/apis/sm103_fp8_block128.hpp new file mode 100644 index 0000000000..6557f2e36c --- /dev/null +++ b/csrc/apis/sm103_fp8_block128.hpp @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace deep_gemm::sm103_fp8_block128 { + +void register_apis(pybind11::module_& m); + +} // namespace deep_gemm::sm103_fp8_block128 diff --git a/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp b/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp index a5cc98fc08..b70c52fb0b 100644 --- a/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp +++ b/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp @@ -13,8 +13,18 @@ #include "../heuristics/mega_moe.hpp" +#include +#include + namespace deep_gemm { +static uint32_t get_fp32_bits(const float value) { + uint32_t bits; + static_assert(sizeof(bits) == sizeof(value)); + std::memcpy(&bits, &value, sizeof(bits)); + return bits; +} + // Map an activation name to its `deep_gemm::ActivationType` enumerator token // (resolved inside the JIT-generated translation unit via `using namespace deep_gemm`). static std::string get_activation_type_name(const std::string& activation) { @@ -41,6 +51,7 @@ class SM100FP8FP4MegaMoERuntime final : public LaunchRuntime sym_buffer_ptrs; @@ -67,7 +78,6 @@ using namespace deep_gemm; static void __instantiate_kernel() {{ auto ptr = reinterpret_cast(&sm100_fp8_fp4_mega_moe_impl< - {}, {}, {}, {}, {}, {}, @@ -76,8 +86,6 @@ static void __instantiate_kernel() {{ {}, {}, {}, {}, - {}, - {}, {}, {}, {}, {}, {}, {}, @@ -85,20 +93,17 @@ static void __instantiate_kernel() {{ {} >); }}; -)", args.num_max_tokens_per_rank, - args.hidden, args.intermediate_hidden, +)", 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.store_block_m, args.config.sf_block_m, args.config.sf_block_n, - args.config.num_ring_tokens, - args.config.num_sf_ring_tokens, args.config.num_stages, args.config.num_bytes_per_pull, 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), + fmt::format("{}u", get_fp32_bits(args.activation_clamp)), args.fast_math ? "true" : "false", get_activation_type_name(args.activation)); } @@ -108,17 +113,32 @@ static void __instantiate_kernel() {{ DG_CUDA_UNIFIED_CHECK(launch_kernel(kernel, config, args.y, args.cumulative_local_expert_recv_stats, + args.saved_token_src_metadata, args.num_tokens, + args.num_max_tokens_per_rank, + args.config.num_ring_tokens, + args.config.num_sf_ring_tokens, args.sym_buffer_ptrs, args.tensor_map_l1_acts, args.tensor_map_l1_acts_sf, args.tensor_map_l1_weights, args.tensor_map_l1_weights_sf, args.tensor_map_l1_output, + args.tensor_map_l1_output, args.tensor_map_l2_acts, args.tensor_map_l2_acts_sf, args.tensor_map_l2_weights, - args.tensor_map_l2_weights_sf + args.tensor_map_l2_weights_sf, + args.tensor_map_l1_output, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr )); } }; @@ -221,6 +241,7 @@ static void sm100_fp8_fp4_mega_moe( .config = config, .y = y.data_ptr(), .cumulative_local_expert_recv_stats = cumulative_local_expert_recv_stats_ptr, + .saved_token_src_metadata = nullptr, .num_tokens = num_tokens, .sym_buffer_ptrs = layout::SymBuffer<>(sym_buffer_ptrs, rank_idx), .tensor_map_l1_acts = tensor_map_l1_acts, diff --git a/csrc/jit_kernels/impls/smxx_layout.hpp b/csrc/jit_kernels/impls/smxx_layout.hpp index 9942e221da..b50079df89 100644 --- a/csrc/jit_kernels/impls/smxx_layout.hpp +++ b/csrc/jit_kernels/impls/smxx_layout.hpp @@ -46,7 +46,7 @@ class TransposeAndPackFP32IntoUE8M0Runtime final: public LaunchRuntime(&transpose_and_pack_fp32_into_ue8m0< - {}, {}, {}, {}, {} + {}, {}, {}, {}, {}, {} >); }}; )", args.launch_args.num_threads, args.block_mn, args.sf_k, - args.num_psum_groups, args.use_psum_layout ? "true" : "false"); + args.num_psum_groups, args.use_psum_layout ? "true" : "false", + args.sf_column_major ? "true" : "false"); } static void launch_impl(const KernelHandle& kernel, const LaunchConfigHandle& config, Args args) { @@ -185,10 +186,15 @@ static torch::Tensor get_mn_major_tma_aligned_packed_ue8m0_tensor(const torch::T {packed_sf_k * tma_aligned_mn, 1, tma_aligned_mn}, at::TensorOptions().device(batched_sf.device()).dtype(torch::kInt)); - // PSUM layout (always 2D contiguous) lets the pack kernel skip uninitialized MN gap rows + // PSUM layout lets the pack kernel skip uninitialized MN gap rows. Expanded + // DeepEP may supply its 2D SF matrix in an already TMA-aligned column-major + // layout, which the pack kernel consumes directly without an intermediate copy. const auto use_psum_layout = psum_layout.has_value(); + const auto sf_is_contiguous = batched_sf.is_contiguous(); + const auto sf_is_column_major = num_sf_batches == 1 and + batched_sf.stride(1) == 1 and batched_sf.stride(2) == tma_aligned_mn; if (use_psum_layout) { - DG_HOST_ASSERT(num_sf_batches == 1 and batched_sf.is_contiguous()); + DG_HOST_ASSERT(num_sf_batches == 1 and (sf_is_contiguous or sf_is_column_major)); DG_HOST_ASSERT(psum_layout->scalar_type() == torch::kInt and psum_layout->is_contiguous()); DG_HOST_ASSERT(psum_layout->numel() > 0); } @@ -196,7 +202,7 @@ static torch::Tensor get_mn_major_tma_aligned_packed_ue8m0_tensor(const torch::T const auto num_psum_groups = use_psum_layout ? static_cast(psum_layout->numel()) : 1; // Launch the kernel - if (batched_sf.is_contiguous()) { + if (sf_is_contiguous or (use_psum_layout and sf_is_column_major)) { if ((mn * sf_k) % 4 != 0 and num_sf_batches > 1) return get_mn_major_tma_aligned_packed_ue8m0_tensor_torch(sf); @@ -210,6 +216,7 @@ static torch::Tensor get_mn_major_tma_aligned_packed_ue8m0_tensor(const torch::T .num_psum_groups = num_psum_groups, .m_alignment = m_alignment, .use_psum_layout = use_psum_layout, + .sf_column_major = not sf_is_contiguous, .block_mn = block_mn, .sf = batched_sf.data_ptr(), .out = out.data_ptr(), diff --git a/csrc/python_api.cpp b/csrc/python_api.cpp index a966afe1ed..3a8e2495de 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/runtime.hpp" +#include "apis/sm103_fp8_block128.hpp" #ifndef TORCH_EXTENSION_NAME #define TORCH_EXTENSION_NAME _C @@ -25,4 +26,12 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { deep_gemm::layout::register_apis(m); deep_gemm::mega::register_apis(m); deep_gemm::runtime::register_apis(m); + deep_gemm::sm103_fp8_block128::register_apis(m); + +#define DG_STRINGIFY_IMPL(value) #value +#define DG_STRINGIFY(value) DG_STRINGIFY_IMPL(value) +#ifndef DEEP_GEMM_GIT_COMMIT_TOKEN +#define DEEP_GEMM_GIT_COMMIT_TOKEN unknown +#endif + m.attr("__git_commit__") = DG_STRINGIFY(DEEP_GEMM_GIT_COMMIT_TOKEN); } diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu new file mode 100644 index 0000000000..3178ba7480 --- /dev/null +++ b/csrc/sm103_fp8_block128.cu @@ -0,0 +1,3164 @@ +#include "apis/sm103_fp8_block128.hpp" + +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +#include "utils/system.hpp" +#include "jit_kernels/impls/runtime_utils.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include + +// CUTLASS 4.2 exposes an SM103 epilogue-builder alias but omits the equivalent +// alias for its software (FP32) blockwise-scale mainloop. The underlying UMMA +// implementation is ISA-compatible; this bridge lets the public kernel use an +// SM103 ArchTag while retaining CUTLASS's software-scale collective. The +// translation unit itself contains only sm_103a code and every host entry point +// checks compute capability 10.3 exactly. +namespace cutlass::gemm { +struct KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm103 final + : KernelSchedule2Sm, KernelScheduleSm100PtrArrayBlockwise {}; +} // namespace cutlass::gemm + +namespace cutlass::gemm::collective { +template < + class ElementA, + class GmemLayoutA, + int AlignmentA, + class ElementB, + class GmemLayoutB, + int AlignmentB, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + class StageCountType, + class KernelScheduleType> +struct CollectiveBuilder< + arch::Sm103, + arch::OpClassTensorOp, + ElementA, + GmemLayoutA, + AlignmentA, + ElementB, + GmemLayoutB, + AlignmentB, + ElementAccumulator, + TileShape_MNK, + ClusterShape_MNK, + StageCountType, + KernelScheduleType, + cute::enable_if_t>> + : CollectiveBuilder< + arch::Sm100, + arch::OpClassTensorOp, + ElementA, + GmemLayoutA, + AlignmentA, + ElementB, + GmemLayoutB, + AlignmentB, + ElementAccumulator, + TileShape_MNK, + ClusterShape_MNK, + StageCountType, + KernelScheduleType> {}; +} // namespace cutlass::gemm::collective + +namespace deep_gemm::sm103_fp8_block128 { +namespace { + +constexpr int kRequiredMajor = 10; +constexpr int kRequiredMinor = 3; +constexpr int kBlockK = 128; +constexpr float kE4M3Max = 448.0f; + +constexpr uint32_t kPersistentHidden = 6144; +constexpr uint32_t kPersistentIntermediate = 2048; +constexpr uint32_t kPersistentExperts = 256; +constexpr uint32_t kPersistentTopK = 8; +constexpr uint32_t kPersistentBlockM = 192; +constexpr uint32_t kPersistentBlockN = 128; +constexpr uint32_t kPersistentBlockK = 128; +constexpr uint32_t kPersistentStoreBlockM = 32; +constexpr uint32_t kPersistentSFBlockM = 256; +constexpr uint32_t kPersistentSFBlockN = 128; +constexpr uint32_t kPersistentStages = 6; +constexpr uint32_t kPersistentPullBytes = 3072; +constexpr uint32_t kPersistentDispatchThreads = 128; +constexpr uint32_t kPersistentNonEpilogueThreads = 128; +constexpr uint32_t kPersistentEpilogueThreads = 256; +constexpr uint32_t kPersistentThreads = + kPersistentDispatchThreads + kPersistentNonEpilogueThreads + + kPersistentEpilogueThreads; +constexpr uint32_t kPersistentLocalSMs = 148; +constexpr uint32_t kPersistentProductionSMs = 152; +// Upstream's persistent tile occupies 212,260 bytes through its last shared +// control word. Canonical FP8 W13 adds one BF16 warp-pair exchange slot per +// epilogue warpgroup. Exact block128 activation scaling also adds one peer-amax +// slot (16 float2 values) and one cluster barrier per warpgroup. +// Keep the launch extent derived from that private implementation detail; this +// is not a caller capacity or tuning option. +constexpr uint32_t kPersistentUpstreamSmemBytes = 212260; +constexpr uint32_t kPersistentW13PairExchangeBytes = + 2 * 4 * 32 * sizeof(uint2); +constexpr uint32_t kPersistentBlock128ScaleExchangeBytes = + 2 * (16 * sizeof(float2) + sizeof(uint64_t)); +constexpr uint32_t kPersistentSmemBytes = + kPersistentUpstreamSmemBytes + kPersistentW13PairExchangeBytes + + kPersistentBlock128ScaleExchangeBytes; +constexpr uint32_t kWorkspaceAlignment = + deep_gemm::layout::kLCMCandidateBlockM; + +uint32_t get_persistent_sm_count(const torch::Tensor& tensor) { + cudaDeviceProp properties{}; + C10_CUDA_CHECK(cudaGetDeviceProperties( + &properties, tensor.get_device())); + TORCH_CHECK( + properties.multiProcessorCount == kPersistentLocalSMs || + properties.multiProcessorCount == kPersistentProductionSMs, + "persistent GLM MegaMoE supports the 148-SM and 152-SM SM103 " + "topologies, got ", + properties.multiProcessorCount, + " SMs"); + return static_cast(properties.multiProcessorCount); +} + +constexpr uint32_t align_workspace_tokens(const uint32_t value) { + return (value + kWorkspaceAlignment - 1) / kWorkspaceAlignment * + kWorkspaceAlignment; +} + +struct PersistentWorkspaceLayout { + uint32_t num_ranks; + uint32_t capacity; + uint32_t ring_tokens; + uint32_t sf_ring_tokens; + deep_gemm::layout::Workspace workspace; + deep_gemm::layout::Buffer input_tokens; + deep_gemm::layout::Buffer input_scales; + deep_gemm::layout::Buffer input_topk_ids; + deep_gemm::layout::Buffer input_topk_scores; + deep_gemm::layout::Buffer l1_tokens; + deep_gemm::layout::Buffer l1_scales; + deep_gemm::layout::Buffer l1_scores; + deep_gemm::layout::Buffer l2_tokens; + deep_gemm::layout::Buffer l2_scales; + deep_gemm::layout::Buffer combine_tokens; + deep_gemm::layout::Buffer backward_grad_y_tokens; + deep_gemm::layout::Buffer backward_grad_y_scales; + deep_gemm::layout::Buffer backward_grad_scores; + deep_gemm::layout::Buffer backward_ring_grad_y; + deep_gemm::layout::Buffer backward_ring_grad_y_scales; + deep_gemm::layout::Buffer backward_ring_grad_preact; + deep_gemm::layout::Buffer backward_ring_grad_preact_scales; + deep_gemm::layout::Buffer backward_ring_bf16; + deep_gemm::layout::Buffer backward_ring_dscore; + deep_gemm::layout::Buffer backward_full_x; + deep_gemm::layout::Buffer backward_full_x_scales; + deep_gemm::layout::Buffer backward_full_grad_y; + deep_gemm::layout::Buffer backward_full_grad_y_scales; + deep_gemm::layout::Buffer backward_full_scores; + deep_gemm::layout::Buffer backward_full_h; + deep_gemm::layout::Buffer backward_full_h_scales; + deep_gemm::layout::Buffer backward_full_grad_preact; + deep_gemm::layout::Buffer backward_full_grad_preact_scales; + deep_gemm::layout::Buffer saved_l1_preact; + deep_gemm::layout::Buffer saved_down_unweighted; + deep_gemm::layout::Buffer backward_symmetric_grad_y; + + PersistentWorkspaceLayout( + void* base, + const uint32_t ranks, + const uint32_t context_tokens_per_rank + ) : num_ranks(ranks), + capacity(align_workspace_tokens(context_tokens_per_rank)), + ring_tokens(align_workspace_tokens(ranks * capacity)), + sf_ring_tokens(deep_gemm::layout::get_num_sf_ring_tokens( + ring_tokens, kPersistentBlockM)), + workspace( + base, + ranks, + kPersistentExperts, + capacity, + kPersistentTopK, + ring_tokens), + input_tokens( + deep_gemm::layout::Data(kPersistentHidden), + 1, + capacity, + workspace.get_end_ptr()), + input_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + capacity, + input_tokens.get_end_ptr()), + input_topk_ids( + deep_gemm::layout::Data( + kPersistentTopK * sizeof(int64_t), false), + 1, + capacity, + input_scales.get_end_ptr()), + input_topk_scores( + deep_gemm::layout::Data( + kPersistentTopK * sizeof(float), false), + 1, + capacity, + input_topk_ids.get_end_ptr()), + l1_tokens( + deep_gemm::layout::Data(kPersistentHidden), + 1, + ring_tokens, + input_topk_scores.get_end_ptr()), + l1_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + sf_ring_tokens, + l1_tokens.get_end_ptr()), + l1_scores( + deep_gemm::layout::Data(sizeof(float), false), + 1, + ring_tokens, + l1_scales.get_end_ptr()), + l2_tokens( + deep_gemm::layout::Data(kPersistentIntermediate), + 1, + ring_tokens, + l1_scores.get_end_ptr()), + l2_scales( + deep_gemm::layout::Data(kPersistentIntermediate / 32), + 1, + sf_ring_tokens, + l2_tokens.get_end_ptr()), + combine_tokens( + deep_gemm::layout::Data( + kPersistentHidden * sizeof(__nv_bfloat16)), + kPersistentTopK, + capacity, + l2_scales.get_end_ptr()), + backward_grad_y_tokens( + deep_gemm::layout::Data(kPersistentHidden), + 1, + capacity, + combine_tokens.get_end_ptr()), + backward_grad_y_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + capacity, + backward_grad_y_tokens.get_end_ptr()), + backward_grad_scores( + deep_gemm::layout::Data( + kPersistentTopK * sizeof(float), false), + 1, + capacity, + backward_grad_y_scales.get_end_ptr()), + backward_ring_grad_y( + deep_gemm::layout::Data(kPersistentHidden), + 1, + ring_tokens, + backward_grad_scores.get_end_ptr()), + backward_ring_grad_y_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + sf_ring_tokens, + backward_ring_grad_y.get_end_ptr()), + backward_ring_grad_preact( + deep_gemm::layout::Data(2 * kPersistentIntermediate), + 1, + ring_tokens, + backward_ring_grad_y_scales.get_end_ptr()), + backward_ring_grad_preact_scales( + deep_gemm::layout::Data( + 2 * kPersistentIntermediate / 32), + 1, + sf_ring_tokens, + backward_ring_grad_preact.get_end_ptr()), + backward_ring_bf16( + deep_gemm::layout::Data( + kPersistentHidden * sizeof(__nv_bfloat16)), + 1, + ring_tokens, + backward_ring_grad_preact_scales.get_end_ptr()), + backward_ring_dscore( + deep_gemm::layout::Data(sizeof(float), false), + 1, + ring_tokens, + backward_ring_bf16.get_end_ptr()), + backward_full_x( + deep_gemm::layout::Data(kPersistentHidden), + 1, + workspace.num_max_pool_tokens, + backward_ring_dscore.get_end_ptr()), + backward_full_x_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + workspace.num_max_pool_tokens, + backward_full_x.get_end_ptr()), + backward_full_grad_y( + deep_gemm::layout::Data(kPersistentHidden), + 1, + workspace.num_max_pool_tokens, + backward_full_x_scales.get_end_ptr()), + backward_full_grad_y_scales( + deep_gemm::layout::Data(kPersistentHidden / 32), + 1, + workspace.num_max_pool_tokens, + backward_full_grad_y.get_end_ptr()), + backward_full_scores( + deep_gemm::layout::Data(sizeof(float), false), + 1, + workspace.num_max_pool_tokens, + backward_full_grad_y_scales.get_end_ptr()), + backward_full_h( + deep_gemm::layout::Data(kPersistentIntermediate), + 1, + workspace.num_max_pool_tokens, + backward_full_scores.get_end_ptr()), + backward_full_h_scales( + deep_gemm::layout::Data(kPersistentIntermediate / 32), + 1, + workspace.num_max_pool_tokens, + backward_full_h.get_end_ptr()), + backward_full_grad_preact( + deep_gemm::layout::Data(2 * kPersistentIntermediate), + 1, + workspace.num_max_pool_tokens, + backward_full_h_scales.get_end_ptr()), + backward_full_grad_preact_scales( + deep_gemm::layout::Data( + 2 * kPersistentIntermediate / 32), + 1, + workspace.num_max_pool_tokens, + backward_full_grad_preact.get_end_ptr()), + saved_l1_preact( + deep_gemm::layout::Data( + 2 * kPersistentIntermediate * + sizeof(__nv_bfloat16)), + 1, + workspace.num_max_pool_tokens, + backward_full_grad_preact_scales.get_end_ptr()), + saved_down_unweighted( + deep_gemm::layout::Data( + kPersistentHidden * sizeof(__nv_bfloat16)), + 1, + workspace.num_max_pool_tokens, + saved_l1_preact.get_end_ptr()), + // Upstream reverse transport needs an immutable symmetric input plane + // distinct from its remote dX combine plane. It is compact-token + // storage, not a routed-pool wgrad materialization, so its capacity is + // the same context/CP-derived per-rank envelope as the forward input. + backward_symmetric_grad_y( + deep_gemm::layout::Data( + kPersistentHidden * sizeof(__nv_bfloat16)), + 1, + capacity, + saved_down_unweighted.get_end_ptr()) {} + + int64_t num_bytes() const { + return reinterpret_cast( + backward_symmetric_grad_y.get_end_ptr()) - + reinterpret_cast(workspace.base); + } +}; + +constexpr int64_t align_rows(const int64_t rows) { + return (rows + kBlockK - 1) / kBlockK * kBlockK; +} + +#define DG_CHECK_CUDA(tensor) \ + TORCH_CHECK((tensor).is_cuda(), #tensor " must be a CUDA tensor") +#define DG_CHECK_CONTIGUOUS(tensor) \ + TORCH_CHECK((tensor).is_contiguous(), #tensor " must be contiguous") + +void check_sm103_device(const torch::Tensor& tensor) { + DG_CHECK_CUDA(tensor); + c10::cuda::CUDAGuard guard(tensor.device()); + cudaDeviceProp properties{}; + C10_CUDA_CHECK(cudaGetDeviceProperties(&properties, tensor.get_device())); + TORCH_CHECK( + properties.major == kRequiredMajor && properties.minor == kRequiredMinor, + "FP8-block128 MegaMoE is SM103-only; CUDA runtime reported compute capability ", + properties.major, ".", properties.minor, " for device ", tensor.get_device(), + ". No fallback is available." + ); +} + +void check_bf16_matrix(const torch::Tensor& tensor, const char* name) { + check_sm103_device(tensor); + DG_CHECK_CONTIGUOUS(tensor); + TORCH_CHECK(tensor.dim() == 2, name, " must be rank 2"); + TORCH_CHECK(tensor.scalar_type() == torch::kBFloat16, name, " must be bfloat16"); + TORCH_CHECK(tensor.size(1) % kBlockK == 0, name, " K dimension must be divisible by 128"); +} + +void check_fp8_matrix_and_scales( + const torch::Tensor& tensor, + const torch::Tensor& scales, + const char* name +) { + check_sm103_device(tensor); + DG_CHECK_CONTIGUOUS(tensor); + DG_CHECK_CUDA(scales); + DG_CHECK_CONTIGUOUS(scales); + TORCH_CHECK(tensor.dim() == 2, name, " must be rank 2"); + TORCH_CHECK(tensor.scalar_type() == torch::kFloat8_e4m3fn, name, " must be float8_e4m3fn"); + TORCH_CHECK(tensor.size(1) % kBlockK == 0, name, " K dimension must be divisible by 128"); + TORCH_CHECK(scales.scalar_type() == torch::kFloat32, name, " scales must be float32"); + TORCH_CHECK(scales.dim() == 2, name, " scales must be rank 2"); + TORCH_CHECK(scales.size(0) == tensor.size(0), name, " scales row count mismatch"); + TORCH_CHECK(scales.size(1) == tensor.size(1) / kBlockK, name, " scales block count mismatch"); + TORCH_CHECK(scales.device() == tensor.device(), name, " and scales must be on the same device"); +} + +__device__ __forceinline__ float warp_max(float value) { + #pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + value = fmaxf(value, __shfl_down_sync(0xffffffff, value, offset)); + } + return value; +} + +__device__ __forceinline__ float block_max_128(float value, float* warp_values) { + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + value = warp_max(value); + if (lane == 0) { + warp_values[warp] = value; + } + __syncthreads(); + value = threadIdx.x < 4 ? warp_values[threadIdx.x] : 0.0f; + if (warp == 0) { + value = warp_max(value); + } + if (threadIdx.x == 0) { + warp_values[0] = value; + } + __syncthreads(); + return warp_values[0]; +} + +__global__ void sm103_quantize_bf16_e4m3_group128_kernel( + const __nv_bfloat16* input, + __nv_fp8_e4m3* output, + float* scales, + int64_t rows, + int64_t columns +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t num_blocks_k = columns / kBlockK; + const int64_t work_idx = blockIdx.x; + const int64_t row = work_idx / num_blocks_k; + const int64_t block_k = work_idx - row * num_blocks_k; + if (row >= rows) { + return; + } + + const int64_t column = block_k * kBlockK + threadIdx.x; + const int64_t offset = row * columns + column; + const float value = __bfloat162float(input[offset]); + __shared__ float warp_values[4]; + const float amax = block_max_128(fabsf(value), warp_values); + const float raw_scale = fmaxf(amax / kE4M3Max, 0x1p-127f); + const float scale = exp2f(ceilf(log2f(raw_scale))); + if (threadIdx.x == 0) { + scales[row * num_blocks_k + block_k] = scale; + } + output[offset] = __nv_fp8_e4m3(value / scale); +#endif +} + +__global__ void sm103_prepare_persistent_inputs_kernel( + const __nv_bfloat16* input, + const int64_t* topk_ids, + const float* topk_scores, + __nv_fp8_e4m3* output, + uint32_t* packed_scales, + int64_t* output_topk_ids, + float* output_topk_scores, + const int64_t rows +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t num_blocks_k = kPersistentHidden / kBlockK; + const int64_t work_idx = blockIdx.x; + const int64_t row = work_idx / num_blocks_k; + const int64_t block_k = work_idx - row * num_blocks_k; + if (row >= rows) + return; + + const int64_t column = block_k * kBlockK + threadIdx.x; + const int64_t offset = row * kPersistentHidden + column; + const float value = __bfloat162float(input[offset]); + __shared__ float warp_values[4]; + const float amax = block_max_128(fabsf(value), warp_values); + const float raw_scale = fmaxf(amax / kE4M3Max, 0x1p-127f); + const float scale = exp2f(ceilf(log2f(raw_scale))); + if (threadIdx.x == 0) { + const uint32_t exponent = __float_as_uint(scale) >> 23; + packed_scales[row * num_blocks_k + block_k] = + exponent * 0x01010101u; + } + output[offset] = __nv_fp8_e4m3(value / scale); + + // The first activation block also installs the route metadata into the + // registered symmetric input plane. This keeps preparation to one launch + // and leaves dispatch/transport to the persistent kernel. + if (block_k == 0 && threadIdx.x < kPersistentTopK) { + const int64_t route = row * kPersistentTopK + threadIdx.x; + output_topk_ids[route] = topk_ids[route]; + output_topk_scores[route] = topk_scores[route]; + } +#endif +} + +__global__ void sm103_dequantize_e4m3_group128_kernel( + const __nv_fp8_e4m3* input, + const float* scales, + __nv_bfloat16* output, + int64_t rows, + int64_t columns +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = rows * columns; + if (linear_idx >= numel) { + return; + } + const int64_t row = linear_idx / columns; + const int64_t column = linear_idx - row * columns; + const float scale = scales[row * (columns / kBlockK) + column / kBlockK]; + output[linear_idx] = __float2bfloat16_rn(static_cast(input[linear_idx]) * scale); +#endif +} + +__global__ void sm103_swiglu_quantize_group128_kernel( + const __nv_bfloat16* preactivation, + __nv_fp8_e4m3* output, + float* scales, + int64_t rows, + int64_t hidden +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t num_blocks_k = hidden / kBlockK; + const int64_t work_idx = blockIdx.x; + const int64_t row = work_idx / num_blocks_k; + const int64_t block_k = work_idx - row * num_blocks_k; + if (row >= rows) { + return; + } + + const int64_t column = block_k * kBlockK + threadIdx.x; + const int64_t pre_row_offset = row * hidden * 2; + const float up = __bfloat162float(preactivation[pre_row_offset + column]); + const float gate = __bfloat162float(preactivation[pre_row_offset + hidden + column]); + const float sigmoid_gate = 1.0f / (1.0f + expf(-gate)); + const float value = up * gate * sigmoid_gate; + __shared__ float warp_values[4]; + const float amax = block_max_128(fabsf(value), warp_values); + const float raw_scale = fmaxf(amax / kE4M3Max, 0x1p-127f); + const float scale = exp2f(ceilf(log2f(raw_scale))); + if (threadIdx.x == 0) { + scales[row * num_blocks_k + block_k] = scale; + } + output[row * hidden + column] = __nv_fp8_e4m3(value / scale); +#endif +} + +__global__ void sm103_swiglu_backward_kernel( + const __nv_bfloat16* grad_output, + const __nv_bfloat16* preactivation, + __nv_bfloat16* grad_preactivation, + int64_t rows, + int64_t hidden +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = rows * hidden; + if (linear_idx >= numel) { + return; + } + const int64_t row = linear_idx / hidden; + const int64_t column = linear_idx - row * hidden; + const int64_t pre_row_offset = row * hidden * 2; + const float up = __bfloat162float(preactivation[pre_row_offset + column]); + const float gate = __bfloat162float(preactivation[pre_row_offset + hidden + column]); + const float grad = __bfloat162float(grad_output[linear_idx]); + const float sigmoid_gate = 1.0f / (1.0f + expf(-gate)); + const float silu_gate = gate * sigmoid_gate; + const float silu_grad = sigmoid_gate * (1.0f + gate * (1.0f - sigmoid_gate)); + grad_preactivation[pre_row_offset + column] = __float2bfloat16_rn(grad * silu_gate); + grad_preactivation[pre_row_offset + hidden + column] = __float2bfloat16_rn(grad * up * silu_grad); +#endif +} + +// The forward W13 ABI is [up; gate], while FireTitan's canonical gradient +// ownership is the interleaved [gate, up] expert layout. Emit the activation +// gradient in canonical order so both dgrad and wgrad can consume the resident +// canonical W13 buffers directly without allocating a reordered weight copy. +__global__ void sm103_swiglu_backward_canonical_kernel( + const __nv_bfloat16* grad_output, + const __nv_bfloat16* preactivation, + __nv_bfloat16* grad_preactivation, + int64_t rows, + int64_t hidden +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = rows * hidden; + if (linear_idx >= numel) { + return; + } + const int64_t row = linear_idx / hidden; + const int64_t column = linear_idx - row * hidden; + const int64_t pre_row_offset = row * hidden * 2; + const float up = __bfloat162float(preactivation[pre_row_offset + column]); + const float gate = __bfloat162float(preactivation[pre_row_offset + hidden + column]); + const float grad = __bfloat162float(grad_output[linear_idx]); + const float sigmoid_gate = 1.0f / (1.0f + expf(-gate)); + const float silu_gate = gate * sigmoid_gate; + const float silu_grad = sigmoid_gate * (1.0f + gate * (1.0f - sigmoid_gate)); + // Canonical W13 order: gate, then up. + grad_preactivation[pre_row_offset + column] = __float2bfloat16_rn(grad * up * silu_grad); + grad_preactivation[pre_row_offset + hidden + column] = __float2bfloat16_rn(grad * silu_gate); +#endif +} + +__global__ void sm103_route_scale_quantize_group128_kernel( + const __nv_bfloat16* grad_output, + const float* route_scores, + const int64_t* route_order, + __nv_fp8_e4m3* output, + float* scales, + int64_t num_routes, + int64_t hidden, + int64_t topk +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t num_blocks_k = hidden / kBlockK; + const int64_t work_idx = blockIdx.x; + const int64_t route = work_idx / num_blocks_k; + const int64_t block_k = work_idx - route * num_blocks_k; + if (route >= num_routes) { + return; + } + + const int64_t source_route = route_order[route]; + const int64_t token = source_route / topk; + const float score = route_scores[source_route]; + const int64_t column = block_k * kBlockK + threadIdx.x; + const float value = __bfloat162float(grad_output[token * hidden + column]) * score; + __shared__ float warp_values[4]; + const float amax = block_max_128(fabsf(value), warp_values); + const float scale = amax == 0.0f ? 1.0f : amax / kE4M3Max; + if (threadIdx.x == 0) { + scales[route * num_blocks_k + block_k] = scale; + } + output[route * hidden + column] = __nv_fp8_e4m3(value / scale); +#endif +} + +__global__ void sm103_post_down_combine_kernel( + const __nv_bfloat16* route_output, + const float* route_scores, + __nv_bfloat16* output, + int64_t num_tokens, + int64_t hidden, + int64_t topk +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = num_tokens * hidden; + if (linear_idx >= numel) { + return; + } + const int64_t token = linear_idx / hidden; + const int64_t column = linear_idx - token * hidden; + float sum = 0.0f; + #pragma unroll 1 + for (int64_t route = 0; route < topk; ++route) { + const int64_t route_idx = token * topk + route; + sum += __bfloat162float(route_output[route_idx * hidden + column]) * route_scores[route_idx]; + } + output[linear_idx] = __float2bfloat16_rn(sum); +#endif +} + +__global__ void sm103_route_sum_kernel( + const __nv_bfloat16* route_grad, + __nv_bfloat16* output, + int64_t num_tokens, + int64_t hidden, + int64_t topk +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = num_tokens * hidden; + if (linear_idx >= numel) { + return; + } + const int64_t token = linear_idx / hidden; + const int64_t column = linear_idx - token * hidden; + float sum = 0.0f; + #pragma unroll 1 + for (int64_t route = 0; route < topk; ++route) { + sum += __bfloat162float(route_grad[(token * topk + route) * hidden + column]); + } + output[linear_idx] = __float2bfloat16_rn(sum); +#endif +} + +__global__ void sm103_post_down_score_grad_kernel( + const __nv_bfloat16* route_output, + const __nv_bfloat16* grad_output, + float* grad_scores, + int64_t num_routes, + int64_t hidden, + int64_t topk +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t route = blockIdx.x; + if (route >= num_routes) { + return; + } + const int64_t token = route / topk; + float partial = 0.0f; + for (int64_t column = threadIdx.x; column < hidden; column += blockDim.x) { + partial += __bfloat162float(route_output[route * hidden + column]) * + __bfloat162float(grad_output[token * hidden + column]); + } + #pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + partial += __shfl_down_sync(0xffffffff, partial, offset); + } + __shared__ float warp_sums[8]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + if (lane == 0) { + warp_sums[warp] = partial; + } + __syncthreads(); + if (warp == 0) { + float total = lane < 8 ? warp_sums[lane] : 0.0f; + #pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + total += __shfl_down_sync(0xffffffff, total, offset); + } + if (lane == 0) { + grad_scores[route] = total; + } + } +#endif +} + +__global__ void sm103_expanded_post_down_scale_kernel( + const __nv_bfloat16* route_output, + const float* route_scores, + __nv_bfloat16* scaled_output, + int64_t rows, + int64_t hidden +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t linear_idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t numel = rows * hidden; + if (linear_idx >= numel) { + return; + } + const int64_t row = linear_idx / hidden; + const float value = __bfloat162float(route_output[linear_idx]) * route_scores[row]; + scaled_output[linear_idx] = __float2bfloat16_rn(value); +#endif +} + +template +__global__ void sm103_expand_compact_routes_kernel( + const __nv_bfloat16* compact, + const RouteIndex* routes, + __nv_bfloat16* expanded, + int64_t compact_rows, + int64_t topk, + int64_t hidden, + int64_t expanded_rows, + int64_t route_stride_0, + int64_t route_stride_1 +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t route_linear = blockIdx.x; + if (route_linear >= compact_rows * topk) { + return; + } + const int64_t compact_row = route_linear / topk; + const int64_t slot = route_linear - compact_row * topk; + const int64_t expanded_row = static_cast( + routes[compact_row * route_stride_0 + slot * route_stride_1]); + if (expanded_row < 0 || expanded_row >= expanded_rows) { + return; + } + for (int64_t column = threadIdx.x; column < hidden; column += blockDim.x) { + expanded[expanded_row * hidden + column] = compact[compact_row * hidden + column]; + } +#endif +} + +template +__global__ void sm103_collapse_expanded_routes_kernel( + const __nv_bfloat16* expanded, + const RouteIndex* routes, + __nv_bfloat16* compact, + int64_t compact_rows, + int64_t topk, + int64_t hidden, + int64_t expanded_rows, + int64_t route_stride_0, + int64_t route_stride_1 +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t compact_row = blockIdx.x; + if (compact_row >= compact_rows) { + return; + } + for (int64_t column = threadIdx.x; column < hidden; column += blockDim.x) { + float sum = 0.0f; + #pragma unroll 1 + for (int64_t slot = 0; slot < topk; ++slot) { + const int64_t expanded_row = static_cast( + routes[compact_row * route_stride_0 + slot * route_stride_1]); + if (expanded_row >= 0 && expanded_row < expanded_rows) { + sum += __bfloat162float(expanded[expanded_row * hidden + column]); + } + } + compact[compact_row * hidden + column] = __float2bfloat16_rn(sum); + } +#endif +} + +template +__global__ void sm103_expanded_post_down_score_grad_kernel( + const __nv_bfloat16* route_output, + const __nv_bfloat16* grad_output, + const RouteIndex* routes, + float* grad_scores, + int64_t compact_rows, + int64_t topk, + int64_t hidden, + int64_t expanded_rows, + int64_t route_stride_0, + int64_t route_stride_1 +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t route_linear = blockIdx.x; + if (route_linear >= compact_rows * topk) { + return; + } + const int64_t compact_row = route_linear / topk; + const int64_t slot = route_linear - compact_row * topk; + const int64_t expanded_row = static_cast( + routes[compact_row * route_stride_0 + slot * route_stride_1]); + if (expanded_row < 0 || expanded_row >= expanded_rows) { + if (threadIdx.x == 0) { + grad_scores[route_linear] = 0.0f; + } + return; + } + float partial = 0.0f; + for (int64_t column = threadIdx.x; column < hidden; column += blockDim.x) { + partial += __bfloat162float(route_output[expanded_row * hidden + column]) * + __bfloat162float(grad_output[expanded_row * hidden + column]); + } + #pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + partial += __shfl_down_sync(0xffffffff, partial, offset); + } + __shared__ float warp_sums[8]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + if (lane == 0) { + warp_sums[warp] = partial; + } + __syncthreads(); + if (warp == 0) { + float total = lane < 8 ? warp_sums[lane] : 0.0f; + #pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + total += __shfl_down_sync(0xffffffff, total, offset); + } + if (lane == 0) { + grad_scores[route_linear] = total; + } + } +#endif +} + +__global__ void sm103_expanded_route_scale_quantize_group128_kernel( + const __nv_bfloat16* grad_output, + const float* route_scores, + __nv_fp8_e4m3* output, + float* scales, + int64_t rows, + int64_t hidden +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t num_blocks_k = hidden / kBlockK; + const int64_t work_idx = blockIdx.x; + const int64_t row = work_idx / num_blocks_k; + const int64_t block_k = work_idx - row * num_blocks_k; + if (row >= rows) { + return; + } + const int64_t column = block_k * kBlockK + threadIdx.x; + const float value = __bfloat162float(grad_output[row * hidden + column]) * route_scores[row]; + __shared__ float warp_values[4]; + const float amax = block_max_128(fabsf(value), warp_values); + const float scale = amax == 0.0f ? 1.0f : amax / kE4M3Max; + if (threadIdx.x == 0) { + scales[row * num_blocks_k + block_k] = scale; + } + output[row * hidden + column] = __nv_fp8_e4m3(value / scale); +#endif +} + +__global__ void sm103_zero_expanded_wgrad_padding_kernel( + __nv_bfloat16* left, + __nv_bfloat16* right, + const int32_t* group_counts, + const int32_t* padded_offsets, + int64_t groups, + int64_t left_columns, + int64_t right_columns +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t expert = blockIdx.x; + if (expert >= groups) { + return; + } + const int64_t start = expert == 0 ? 0 : padded_offsets[expert - 1]; + const int64_t valid_end = start + group_counts[expert]; + const int64_t padded_end = padded_offsets[expert]; + const int64_t padding_rows = padded_end - valid_end; + const int64_t left_elements = padding_rows * left_columns; + const int64_t total_elements = left_elements + padding_rows * right_columns; + for (int64_t index = threadIdx.x; index < total_elements; index += blockDim.x) { + if (index < left_elements) { + const int64_t row = valid_end + index / left_columns; + const int64_t column = index - (row - valid_end) * left_columns; + left[row * left_columns + column] = __float2bfloat16(0.0f); + } else { + const int64_t right_index = index - left_elements; + const int64_t row = valid_end + right_index / right_columns; + const int64_t column = right_index - (row - valid_end) * right_columns; + right[row * right_columns + column] = __float2bfloat16(0.0f); + } + } +#endif +} + +__global__ void sm103_prepare_expanded_wgrad_metadata_kernel( + const int32_t* psum, + int32_t* padded_offsets, + int32_t* group_counts, + int64_t groups, + int64_t storage_rows +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t expert = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (expert >= groups) { + return; + } + const int64_t previous_end = expert == 0 ? 0 : psum[expert - 1]; + const int64_t start = (previous_end + kBlockK - 1) / kBlockK * kBlockK; + const int64_t end = psum[expert]; + const int64_t padded_end = (end + kBlockK - 1) / kBlockK * kBlockK; + if (end < start || padded_end > storage_rows || + (expert == groups - 1 && padded_end != storage_rows)) { + asm volatile("trap;"); + return; + } + padded_offsets[expert] = static_cast(padded_end); + group_counts[expert] = static_cast(end - start); +#endif +} + +__global__ void sm103_zero_empty_wgrad_groups_kernel( + __nv_bfloat16* output, + const int32_t* group_counts, + int64_t groups, + int64_t elements_per_group +) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + const int64_t expert = blockIdx.x; + if (expert >= groups || group_counts[expert] != 0) { + return; + } + for (int64_t index = threadIdx.x; index < elements_per_group; index += blockDim.x) { + output[expert * elements_per_group + index] = __float2bfloat16(0.0f); + } +#endif +} + +template +struct SM103GroupedBlockwiseGemm { + using ProblemShape = cutlass::gemm::GroupProblemShape>; + using ElementA = cutlass::float_e4m3_t; + using ElementB = cutlass::float_e4m3_t; + using ElementC = cutlass::bfloat16_t; + using ElementD = cutlass::bfloat16_t; + using ElementAccumulator = float; + using ElementCompute = float; + using LayoutA = cutlass::layout::RowMajor; + using LayoutB = std::conditional_t< + kWeightIsKByN, + cutlass::layout::RowMajor, + cutlass::layout::ColumnMajor>; + using LayoutC = cutlass::layout::RowMajor; + using LayoutD = cutlass::layout::RowMajor; + static constexpr int AlignmentA = 16; + static constexpr int AlignmentB = 16; + static constexpr int AlignmentC = 8; + static constexpr int AlignmentD = 8; + using MmaTileShape = cute::Shape; + using ClusterShape = cute::Shape; + + // A scales are native row-major [M, K/128]. For the ordinary NT path, + // B scales are [N/128, K/128]; for the transposed-weight path the original + // [K/128, N/128] storage is consumed without a copy. + using ScaleConfig = cutlass::detail::Sm1xxBlockwiseScaleConfig< + 1, + kBlockK, + kBlockK, + cute::UMMA::Major::K, + kWeightIsKByN ? cute::UMMA::Major::MN : cute::UMMA::Major::K>; + using LayoutSFA = decltype(ScaleConfig::deduce_layoutSFA()); + using LayoutSFB = decltype(ScaleConfig::deduce_layoutSFB()); + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm103, + cutlass::arch::OpClassTensorOp, + MmaTileShape, + ClusterShape, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, + ElementCompute, + ElementC, + LayoutC*, + AlignmentC, + ElementD, + LayoutD*, + AlignmentD, + cutlass::epilogue::PtrArrayTmaWarpSpecialized2Sm>::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm103, + cutlass::arch::OpClassTensorOp, + ElementA, + cute::tuple, + AlignmentA, + ElementB, + cute::tuple, + AlignmentB, + ElementAccumulator, + MmaTileShape, + ClusterShape, + cutlass::gemm::collective::StageCountAutoCarveout< + static_cast(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm103>::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + ProblemShape, + CollectiveMainloop, + CollectiveEpilogue, + void>; + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + using StrideA = typename GemmKernel::InternalStrideA; + using StrideB = typename GemmKernel::InternalStrideB; + using StrideC = typename GemmKernel::InternalStrideC; + using StrideD = typename GemmKernel::InternalStrideD; +}; + +struct MetadataBlob { + std::vector bytes; + + template + size_t append(const std::vector& values) { + constexpr size_t kMinimumAlignment = 16; + const size_t alignment = std::max(kMinimumAlignment, alignof(T)); + const size_t offset = (bytes.size() + alignment - 1) / alignment * alignment; + const size_t value_bytes = values.size() * sizeof(T); + bytes.resize(offset + value_bytes); + if (value_bytes != 0) { + std::memcpy(bytes.data() + offset, values.data(), value_bytes); + } + return offset; + } + + torch::Tensor copy_to_device( + const torch::TensorOptions& options, + cudaStream_t stream + ) const { + auto storage = torch::empty( + {std::max(static_cast(bytes.size()), 1)}, + options.dtype(torch::kUInt8)); + if (!bytes.empty()) { + C10_CUDA_CHECK(cudaMemcpyAsync( + storage.data_ptr(), bytes.data(), bytes.size(), + cudaMemcpyHostToDevice, stream)); + } + return storage; + } +}; + +void check_grouped_bf16_wgrad_inputs( + const torch::Tensor& left, + const torch::Tensor& right, + const int64_t groups +) { + check_bf16_matrix(left, "left"); + check_bf16_matrix(right, "right"); + TORCH_CHECK(left.device() == right.device(), "left and right must share a device"); + TORCH_CHECK(left.size(0) == right.size(0), "left and right row counts differ"); + TORCH_CHECK(groups > 0 && groups < 1024, + "SM103 grouped BF16 wgrad requires between 1 and 1023 groups"); +} + +torch::Tensor launch_grouped_bf16_wgrad( + const torch::Tensor& left, + const torch::Tensor& right, + const torch::Tensor& metadata_device, + const bool zero_expanded_padding +) { + const int64_t groups = metadata_device.numel() / 2; + if (left.size(0) == 0) { + return torch::zeros( + {groups, left.size(1), right.size(1)}, left.options()); + } + + c10::cuda::CUDAGuard guard(left.device()); + const auto stream = at::cuda::getCurrentCUDAStream(left.get_device()); + auto padded_offsets = metadata_device.narrow(0, 0, groups); + const auto* counts_device = metadata_device.data_ptr() + groups; + if (zero_expanded_padding) { + constexpr int threads = 256; + sm103_zero_expanded_wgrad_padding_kernel<<>>( + reinterpret_cast<__nv_bfloat16*>(left.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(right.data_ptr()), + counts_device, + padded_offsets.data_ptr(), + groups, left.size(1), right.size(1)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + + // Exact SM103 validation above makes PyTorch's BF16 grouped tensor-core + // implementation the only reachable dispatch. There is no architecture + // fallback in this companion entry point. + auto output = at::_grouped_mm( + left.transpose(0, 1), right, padded_offsets, std::nullopt, std::nullopt); + constexpr int threads = 256; + sm103_zero_empty_wgrad_groups_kernel<<>>( + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + counts_device, + groups, left.size(1) * right.size(1)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return output; +} + +torch::Tensor grouped_bf16_wgrad( + const torch::Tensor& left, + const torch::Tensor& right, + const std::vector& padded_group_counts +) { + check_grouped_bf16_wgrad_inputs(left, right, padded_group_counts.size()); + + std::vector metadata(padded_group_counts.size() * 2); + int64_t storage_rows = 0; + for (size_t expert = 0; expert < padded_group_counts.size(); ++expert) { + const int64_t count = padded_group_counts[expert]; + TORCH_CHECK(count >= 0 && count <= std::numeric_limits::max(), + "group counts must fit non-negative int32"); + TORCH_CHECK(count == 0 || count % kBlockK == 0, + "BF16 wgrad group storage must be padded to 128 rows"); + storage_rows += count; + TORCH_CHECK(storage_rows <= std::numeric_limits::max(), + "BF16 wgrad storage rows must fit int32"); + metadata[expert] = static_cast(storage_rows); + metadata[padded_group_counts.size() + expert] = static_cast(count); + } + TORCH_CHECK(storage_rows == left.size(0), + "group counts do not match BF16 wgrad storage rows"); + + c10::cuda::CUDAGuard guard(left.device()); + const auto stream = at::cuda::getCurrentCUDAStream(left.get_device()); + auto metadata_device = torch::empty( + {static_cast(metadata.size())}, + left.options().dtype(torch::kInt32)); + C10_CUDA_CHECK(cudaMemcpyAsync( + metadata_device.data_ptr(), metadata.data(), metadata.size() * sizeof(int32_t), + cudaMemcpyHostToDevice, stream)); + return launch_grouped_bf16_wgrad(left, right, metadata_device, false); +} + +torch::Tensor grouped_bf16_wgrad_expanded( + const torch::Tensor& left, + const torch::Tensor& right, + const torch::Tensor& psum +) { + check_sm103_device(psum); + DG_CHECK_CONTIGUOUS(psum); + TORCH_CHECK(psum.dim() == 1 && psum.scalar_type() == torch::kInt32, + "expanded PSUM must be contiguous rank-1 int32"); + TORCH_CHECK(psum.device() == left.device(), "expanded PSUM device mismatch"); + check_grouped_bf16_wgrad_inputs(left, right, psum.numel()); + TORCH_CHECK(left.size(0) <= std::numeric_limits::max(), + "BF16 wgrad storage rows must fit int32"); + + auto metadata_device = torch::empty( + {psum.numel() * 2}, left.options().dtype(torch::kInt32)); + auto padded_offsets = metadata_device.narrow(0, 0, psum.numel()); + auto group_counts = metadata_device.narrow(0, psum.numel(), psum.numel()); + c10::cuda::CUDAGuard guard(left.device()); + const auto stream = at::cuda::getCurrentCUDAStream(left.get_device()); + constexpr int threads = 256; + const int64_t blocks = (psum.numel() + threads - 1) / threads; + sm103_prepare_expanded_wgrad_metadata_kernel<<>>( + psum.data_ptr(), + padded_offsets.data_ptr(), + group_counts.data_ptr(), + psum.numel(), left.size(0)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return launch_grouped_bf16_wgrad(left, right, metadata_device, true); +} + +template +torch::Tensor grouped_fp8_block128_gemm_impl( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& weights, + const torch::Tensor& weight_scales, + const std::vector& group_counts, + const bool expanded_layout +) { + using Config = SM103GroupedBlockwiseGemm; + using Problem = typename Config::ProblemShape::UnderlyingProblemShape; + using ElementA = typename Config::ElementA; + using ElementB = typename Config::ElementB; + using ElementC = typename Config::ElementC; + using ElementD = typename Config::ElementD; + + check_fp8_matrix_and_scales(activations, activation_scales, "activations"); + check_sm103_device(weights); + DG_CHECK_CONTIGUOUS(weights); + DG_CHECK_CUDA(weight_scales); + DG_CHECK_CONTIGUOUS(weight_scales); + TORCH_CHECK(weights.scalar_type() == torch::kFloat8_e4m3fn, "weights must be float8_e4m3fn"); + TORCH_CHECK(weights.dim() == 3, "weights must be rank 3"); + TORCH_CHECK(weight_scales.scalar_type() == torch::kFloat32 && weight_scales.dim() == 3, + "weight_scales must be contiguous rank-3 float32"); + TORCH_CHECK(weights.device() == activations.device() && weight_scales.device() == activations.device(), + "activations, weights, and scales must share a device"); + TORCH_CHECK(static_cast(group_counts.size()) == weights.size(0), + "group_counts must contain one entry per local expert"); + + const int64_t groups = weights.size(0); + const int64_t k = kWeightIsKByN ? weights.size(1) : weights.size(2); + const int64_t n = kWeightIsKByN ? weights.size(2) : weights.size(1); + TORCH_CHECK(k > 0 && n > 0 && k % kBlockK == 0 && n % kBlockK == 0, + "GEMM N and K dimensions must be positive multiples of 128"); + TORCH_CHECK(activations.size(1) == k, "activation K dimension does not match weights"); + if constexpr (kWeightIsKByN) { + TORCH_CHECK(weight_scales.sizes() == torch::IntArrayRef({groups, k / kBlockK, n / kBlockK}), + "transposed weight scales must have shape [G, K/128, N/128]"); + } else { + TORCH_CHECK(weight_scales.sizes() == torch::IntArrayRef({groups, n / kBlockK, k / kBlockK}), + "weight scales must have shape [G, N/128, K/128]"); + } + + int64_t storage_rows = 0; + int64_t active_groups = 0; + for (const int64_t count : group_counts) { + TORCH_CHECK(count >= 0, "group counts must be non-negative"); + TORCH_CHECK(expanded_layout || count == 0 || count % 4 == 0, + "packed active group counts must be padded to a multiple of four"); + storage_rows += expanded_layout ? align_rows(count) : count; + active_groups += count != 0; + } + TORCH_CHECK(storage_rows == activations.size(0), + "group counts do not match the activation storage rows"); + + auto output = torch::empty({storage_rows, n}, activations.options().dtype(torch::kBFloat16)); + if (active_groups == 0) { + return output; + } + + c10::cuda::CUDAGuard guard(activations.device()); + const auto stream = at::cuda::getCurrentCUDAStream(activations.get_device()); + + std::vector problems; + std::vector ptr_a; + std::vector ptr_b; + std::vector ptr_c; + std::vector ptr_d; + std::vector ptr_sfa; + std::vector ptr_sfb; + std::vector stride_a; + std::vector stride_b; + std::vector stride_c; + std::vector stride_d; + std::vector layout_sfa; + std::vector layout_sfb; + problems.reserve(active_groups); + ptr_a.reserve(active_groups); + ptr_b.reserve(active_groups); + ptr_c.reserve(active_groups); + ptr_d.reserve(active_groups); + ptr_sfa.reserve(active_groups); + ptr_sfb.reserve(active_groups); + stride_a.reserve(active_groups); + stride_b.reserve(active_groups); + stride_c.reserve(active_groups); + stride_d.reserve(active_groups); + layout_sfa.reserve(active_groups); + layout_sfb.reserve(active_groups); + + auto* activation_ptr = reinterpret_cast(activations.data_ptr()); + auto* activation_scale_ptr = activation_scales.data_ptr(); + auto* weight_ptr = reinterpret_cast(weights.data_ptr()); + auto* weight_scale_ptr = weight_scales.data_ptr(); + auto* output_ptr = reinterpret_cast(output.data_ptr()); + int64_t row_offset = 0; + const int64_t weight_elements = k * n; + const int64_t weight_scale_elements = (k / kBlockK) * (n / kBlockK); + for (int64_t expert = 0; expert < groups; ++expert) { + const int64_t m = group_counts[expert]; + if (m == 0) { + continue; + } + problems.emplace_back(cute::make_shape(static_cast(m), static_cast(n), static_cast(k))); + ptr_a.push_back(activation_ptr + row_offset * k); + ptr_b.push_back(weight_ptr + expert * weight_elements); + ptr_c.push_back(reinterpret_cast(output_ptr + row_offset * n)); + ptr_d.push_back(output_ptr + row_offset * n); + ptr_sfa.push_back(activation_scale_ptr + row_offset * (k / kBlockK)); + ptr_sfb.push_back(weight_scale_ptr + expert * weight_scale_elements); + stride_a.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideA{}, cute::make_shape(static_cast(m), static_cast(k), 1))); + stride_b.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideB{}, cute::make_shape(static_cast(n), static_cast(k), 1))); + stride_c.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideC{}, cute::make_shape(static_cast(m), static_cast(n), 1))); + stride_d.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideD{}, cute::make_shape(static_cast(m), static_cast(n), 1))); + layout_sfa.push_back(Config::ScaleConfig::tile_atom_to_shape_SFA( + cute::make_shape(static_cast(m), static_cast(n), static_cast(k), 1))); + layout_sfb.push_back(Config::ScaleConfig::tile_atom_to_shape_SFB( + cute::make_shape(static_cast(m), static_cast(n), static_cast(k), 1))); + row_offset += expanded_layout ? align_rows(m) : m; + } + + const auto metadata_options = activations.options().dtype(torch::kUInt8); + MetadataBlob metadata_blob; + const auto problems_offset = metadata_blob.append(problems); + const auto ptr_a_offset = metadata_blob.append(ptr_a); + const auto ptr_b_offset = metadata_blob.append(ptr_b); + const auto ptr_c_offset = metadata_blob.append(ptr_c); + const auto ptr_d_offset = metadata_blob.append(ptr_d); + const auto ptr_sfa_offset = metadata_blob.append(ptr_sfa); + const auto ptr_sfb_offset = metadata_blob.append(ptr_sfb); + const auto stride_a_offset = metadata_blob.append(stride_a); + const auto stride_b_offset = metadata_blob.append(stride_b); + const auto stride_c_offset = metadata_blob.append(stride_c); + const auto stride_d_offset = metadata_blob.append(stride_d); + const auto layout_sfa_offset = metadata_blob.append(layout_sfa); + const auto layout_sfb_offset = metadata_blob.append(layout_sfb); + auto metadata_device = metadata_blob.copy_to_device(metadata_options, stream); + auto* metadata_base = metadata_device.data_ptr(); + + cutlass::KernelHardwareInfo hardware_info; + hardware_info.device_id = activations.get_device(); + hardware_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count( + hardware_info.device_id); + + typename Config::Gemm::Arguments arguments{ + cutlass::gemm::GemmUniversalMode::kGrouped, + {static_cast(active_groups), + reinterpret_cast(metadata_base + problems_offset), + problems.data()}, + {reinterpret_cast(metadata_base + ptr_a_offset), + reinterpret_cast(metadata_base + stride_a_offset), + reinterpret_cast(metadata_base + ptr_b_offset), + reinterpret_cast(metadata_base + stride_b_offset), + reinterpret_cast(metadata_base + ptr_sfa_offset), + reinterpret_cast(metadata_base + layout_sfa_offset), + reinterpret_cast(metadata_base + ptr_sfb_offset), + reinterpret_cast(metadata_base + layout_sfb_offset)}, + {{}, + reinterpret_cast(metadata_base + ptr_c_offset), + reinterpret_cast(metadata_base + stride_c_offset), + reinterpret_cast(metadata_base + ptr_d_offset), + reinterpret_cast(metadata_base + stride_d_offset)}, + hardware_info}; + arguments.epilogue.thread.alpha = 1.0f; + arguments.epilogue.thread.beta = 0.0f; + + typename Config::Gemm gemm; + const auto implement_status = gemm.can_implement(arguments); + TORCH_CHECK(implement_status == cutlass::Status::kSuccess, + "SM103 grouped FP8-block128 GEMM cannot implement the requested problem: ", + cutlassGetStatusString(implement_status)); + const int64_t workspace_bytes = static_cast(gemm.get_workspace_size(arguments)); + auto workspace = torch::empty( + {std::max(workspace_bytes, 1)}, metadata_options); + const auto initialize_status = gemm.initialize(arguments, workspace.data_ptr(), stream); + TORCH_CHECK(initialize_status == cutlass::Status::kSuccess, + "SM103 grouped FP8-block128 GEMM initialization failed: ", + cutlassGetStatusString(initialize_status)); + const auto run_status = gemm.run(stream); + TORCH_CHECK(run_status == cutlass::Status::kSuccess, + "SM103 grouped FP8-block128 GEMM launch failed: ", + cutlassGetStatusString(run_status)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return output; +} + +template +torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_impl( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& canonical_weights, + const torch::Tensor& canonical_weight_scales, + const std::vector& group_counts, + const bool expanded_layout +) { + // FireTitan stores [gate_0, up_0, gate_1, up_1, ...]. Submit two + // independent N=H problems per active expert, pointing the first at the up + // weight and output columns [0,H), and the second at the gate weight and + // output columns [H,2H). This preserves the fused [up; gate] ABI without + // materializing a multi-gigabyte reordered weight tensor. + static_assert(!kWeightIsKByN, "canonical W13 forward requires NT weights"); + using Config = SM103GroupedBlockwiseGemm; + using Problem = typename Config::ProblemShape::UnderlyingProblemShape; + using ElementA = typename Config::ElementA; + using ElementB = typename Config::ElementB; + using ElementC = typename Config::ElementC; + using ElementD = typename Config::ElementD; + + check_fp8_matrix_and_scales(activations, activation_scales, "activations"); + check_sm103_device(canonical_weights); + DG_CHECK_CONTIGUOUS(canonical_weights); + DG_CHECK_CUDA(canonical_weight_scales); + DG_CHECK_CONTIGUOUS(canonical_weight_scales); + TORCH_CHECK(canonical_weights.scalar_type() == torch::kFloat8_e4m3fn, + "canonical W13 weights must be float8_e4m3fn"); + TORCH_CHECK(canonical_weights.dim() == 3 && canonical_weights.size(0) % 2 == 0, + "canonical W13 weights must have shape [2E, H, D]"); + TORCH_CHECK(canonical_weight_scales.scalar_type() == torch::kFloat32 && + canonical_weight_scales.dim() == 3, + "canonical W13 scales must be contiguous rank-3 float32"); + TORCH_CHECK(canonical_weights.device() == activations.device() && + canonical_weight_scales.device() == activations.device(), + "activations, canonical W13 weights, and scales must share a device"); + + const int64_t groups = canonical_weights.size(0) / 2; + const int64_t hidden = canonical_weights.size(1); + const int64_t k = canonical_weights.size(2); + const int64_t doubled_hidden = hidden * 2; + TORCH_CHECK(static_cast(group_counts.size()) == groups, + "group_counts must contain one entry per local expert"); + TORCH_CHECK(hidden > 0 && k > 0 && hidden % kBlockK == 0 && k % kBlockK == 0, + "canonical W13 H and D dimensions must be positive multiples of 128"); + TORCH_CHECK(activations.size(1) == k, + "activation K dimension does not match canonical W13 weights"); + TORCH_CHECK(canonical_weight_scales.sizes() == + torch::IntArrayRef({groups * 2, hidden / kBlockK, k / kBlockK}), + "canonical W13 scales must have shape [2E, H/128, D/128]"); + + int64_t storage_rows = 0; + int64_t active_groups = 0; + for (const int64_t count : group_counts) { + TORCH_CHECK(count >= 0, "group counts must be non-negative"); + TORCH_CHECK(expanded_layout || count == 0 || count % 4 == 0, + "packed active group counts must be padded to a multiple of four"); + storage_rows += expanded_layout ? align_rows(count) : count; + active_groups += count != 0; + } + TORCH_CHECK(storage_rows == activations.size(0), + "group counts do not match the activation storage rows"); + + auto output = torch::empty( + {storage_rows, doubled_hidden}, activations.options().dtype(torch::kBFloat16)); + if (active_groups == 0) { + return output; + } + + c10::cuda::CUDAGuard guard(activations.device()); + const auto stream = at::cuda::getCurrentCUDAStream(activations.get_device()); + const int64_t active_problems = active_groups * 2; + + std::vector problems; + std::vector ptr_a; + std::vector ptr_b; + std::vector ptr_c; + std::vector ptr_d; + std::vector ptr_sfa; + std::vector ptr_sfb; + std::vector stride_a; + std::vector stride_b; + std::vector stride_c; + std::vector stride_d; + std::vector layout_sfa; + std::vector layout_sfb; + problems.reserve(active_problems); + ptr_a.reserve(active_problems); + ptr_b.reserve(active_problems); + ptr_c.reserve(active_problems); + ptr_d.reserve(active_problems); + ptr_sfa.reserve(active_problems); + ptr_sfb.reserve(active_problems); + stride_a.reserve(active_problems); + stride_b.reserve(active_problems); + stride_c.reserve(active_problems); + stride_d.reserve(active_problems); + layout_sfa.reserve(active_problems); + layout_sfb.reserve(active_problems); + + auto* activation_ptr = reinterpret_cast(activations.data_ptr()); + auto* activation_scale_ptr = activation_scales.data_ptr(); + auto* weight_ptr = reinterpret_cast(canonical_weights.data_ptr()); + auto* weight_scale_ptr = canonical_weight_scales.data_ptr(); + auto* output_ptr = reinterpret_cast(output.data_ptr()); + const int64_t weight_elements = hidden * k; + const int64_t weight_scale_elements = (hidden / kBlockK) * (k / kBlockK); + int64_t row_offset = 0; + constexpr int64_t canonical_pair_for_output_half[2] = {1, 0}; // up, gate + for (int64_t expert = 0; expert < groups; ++expert) { + const int64_t m = group_counts[expert]; + if (m == 0) { + continue; + } + for (int64_t output_half = 0; output_half < 2; ++output_half) { + const int64_t canonical_pair = canonical_pair_for_output_half[output_half]; + const int64_t canonical_index = expert * 2 + canonical_pair; + problems.emplace_back(cute::make_shape( + static_cast(m), static_cast(hidden), static_cast(k))); + ptr_a.push_back(activation_ptr + row_offset * k); + ptr_b.push_back(weight_ptr + canonical_index * weight_elements); + auto* output_half_ptr = + output_ptr + row_offset * doubled_hidden + output_half * hidden; + ptr_c.push_back(reinterpret_cast(output_half_ptr)); + ptr_d.push_back(output_half_ptr); + ptr_sfa.push_back(activation_scale_ptr + row_offset * (k / kBlockK)); + ptr_sfb.push_back(weight_scale_ptr + canonical_index * weight_scale_elements); + stride_a.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideA{}, + cute::make_shape(static_cast(m), static_cast(k), 1))); + stride_b.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideB{}, + cute::make_shape(static_cast(hidden), static_cast(k), 1))); + // Use the full 2H row extent for C/D while each problem writes only + // one H-wide half. The two problems therefore never overlap. + stride_c.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideC{}, + cute::make_shape(static_cast(m), static_cast(doubled_hidden), 1))); + stride_d.push_back(cutlass::make_cute_packed_stride( + typename Config::StrideD{}, + cute::make_shape(static_cast(m), static_cast(doubled_hidden), 1))); + layout_sfa.push_back(Config::ScaleConfig::tile_atom_to_shape_SFA( + cute::make_shape(static_cast(m), static_cast(hidden), static_cast(k), 1))); + layout_sfb.push_back(Config::ScaleConfig::tile_atom_to_shape_SFB( + cute::make_shape(static_cast(m), static_cast(hidden), static_cast(k), 1))); + } + row_offset += expanded_layout ? align_rows(m) : m; + } + + const auto metadata_options = activations.options().dtype(torch::kUInt8); + MetadataBlob metadata_blob; + const auto problems_offset = metadata_blob.append(problems); + const auto ptr_a_offset = metadata_blob.append(ptr_a); + const auto ptr_b_offset = metadata_blob.append(ptr_b); + const auto ptr_c_offset = metadata_blob.append(ptr_c); + const auto ptr_d_offset = metadata_blob.append(ptr_d); + const auto ptr_sfa_offset = metadata_blob.append(ptr_sfa); + const auto ptr_sfb_offset = metadata_blob.append(ptr_sfb); + const auto stride_a_offset = metadata_blob.append(stride_a); + const auto stride_b_offset = metadata_blob.append(stride_b); + const auto stride_c_offset = metadata_blob.append(stride_c); + const auto stride_d_offset = metadata_blob.append(stride_d); + const auto layout_sfa_offset = metadata_blob.append(layout_sfa); + const auto layout_sfb_offset = metadata_blob.append(layout_sfb); + auto metadata_device = metadata_blob.copy_to_device(metadata_options, stream); + auto* metadata_base = metadata_device.data_ptr(); + + cutlass::KernelHardwareInfo hardware_info; + hardware_info.device_id = activations.get_device(); + hardware_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count( + hardware_info.device_id); + + typename Config::Gemm::Arguments arguments{ + cutlass::gemm::GemmUniversalMode::kGrouped, + {static_cast(active_problems), + reinterpret_cast(metadata_base + problems_offset), + problems.data()}, + {reinterpret_cast(metadata_base + ptr_a_offset), + reinterpret_cast(metadata_base + stride_a_offset), + reinterpret_cast(metadata_base + ptr_b_offset), + reinterpret_cast(metadata_base + stride_b_offset), + reinterpret_cast(metadata_base + ptr_sfa_offset), + reinterpret_cast(metadata_base + layout_sfa_offset), + reinterpret_cast(metadata_base + ptr_sfb_offset), + reinterpret_cast(metadata_base + layout_sfb_offset)}, + {{}, + reinterpret_cast(metadata_base + ptr_c_offset), + reinterpret_cast(metadata_base + stride_c_offset), + reinterpret_cast(metadata_base + ptr_d_offset), + reinterpret_cast(metadata_base + stride_d_offset)}, + hardware_info}; + arguments.epilogue.thread.alpha = 1.0f; + arguments.epilogue.thread.beta = 0.0f; + + typename Config::Gemm gemm; + const auto implement_status = gemm.can_implement(arguments); + TORCH_CHECK(implement_status == cutlass::Status::kSuccess, + "SM103 canonical W13 grouped GEMM cannot implement the requested problem: ", + cutlassGetStatusString(implement_status)); + const int64_t workspace_bytes = static_cast(gemm.get_workspace_size(arguments)); + auto workspace = torch::empty( + {std::max(workspace_bytes, 1)}, metadata_options); + const auto initialize_status = gemm.initialize(arguments, workspace.data_ptr(), stream); + TORCH_CHECK(initialize_status == cutlass::Status::kSuccess, + "SM103 canonical W13 grouped GEMM initialization failed: ", + cutlassGetStatusString(initialize_status)); + const auto run_status = gemm.run(stream); + TORCH_CHECK(run_status == cutlass::Status::kSuccess, + "SM103 canonical W13 grouped GEMM launch failed: ", + cutlassGetStatusString(run_status)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return output; +} + +torch::Tensor grouped_fp8_block128_gemm_nt( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& weights, + const torch::Tensor& weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_gemm_impl( + activations, activation_scales, weights, weight_scales, group_counts, false); +} + +torch::Tensor grouped_fp8_block128_gemm_nn( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& weights, + const torch::Tensor& weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_gemm_impl( + activations, activation_scales, weights, weight_scales, group_counts, false); +} + +torch::Tensor grouped_fp8_block128_gemm_nt_expanded( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& weights, + const torch::Tensor& weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_gemm_impl( + activations, activation_scales, weights, weight_scales, group_counts, true); +} + +torch::Tensor grouped_fp8_block128_gemm_nn_expanded( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& weights, + const torch::Tensor& weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_gemm_impl( + activations, activation_scales, weights, weight_scales, group_counts, true); +} + +torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& canonical_weights, + const torch::Tensor& canonical_weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_w13_gemm_nt_canonical_impl( + activations, + activation_scales, + canonical_weights, + canonical_weight_scales, + group_counts, + false); +} + +torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_expanded( + const torch::Tensor& activations, + const torch::Tensor& activation_scales, + const torch::Tensor& canonical_weights, + const torch::Tensor& canonical_weight_scales, + const std::vector& group_counts +) { + return grouped_fp8_block128_w13_gemm_nt_canonical_impl( + activations, + activation_scales, + canonical_weights, + canonical_weight_scales, + group_counts, + true); +} + +pybind11::dict persistent_workspace_info( + const int64_t num_ranks, + const int64_t context_tokens_per_rank +) { + TORCH_CHECK(num_ranks == 2 || num_ranks == 16, + "GLM MegaMoE supports only the target EP2 and EP16 topologies"); + TORCH_CHECK(context_tokens_per_rank > 0 && + context_tokens_per_rank <= std::numeric_limits::max(), + "context_tokens_per_rank must be a positive uint32 value"); + const PersistentWorkspaceLayout layout( + nullptr, + static_cast(num_ranks), + static_cast(context_tokens_per_rank)); + pybind11::dict result; + result["num_bytes"] = layout.num_bytes(); + result["capacity"] = layout.capacity; + result["ring_tokens"] = layout.ring_tokens; + result["sf_ring_tokens"] = layout.sf_ring_tokens; + result["block_m"] = kPersistentBlockM; + result["block_n"] = kPersistentBlockN; + result["block_k"] = kPersistentBlockK; + result["supported_num_sms"] = pybind11::make_tuple( + kPersistentLocalSMs, kPersistentProductionSMs); + return result; +} + +void prepare_persistent_inputs( + const torch::Tensor& buffer, + const torch::Tensor& input, + const torch::Tensor& topk_ids, + const torch::Tensor& topk_scores, + const int64_t num_ranks, + const int64_t context_tokens_per_rank +) { + check_bf16_matrix(input, "input"); + TORCH_CHECK(input.size(1) == kPersistentHidden, + "persistent GLM MegaMoE input width must be 6144"); + check_sm103_device(buffer); + DG_CHECK_CONTIGUOUS(buffer); + TORCH_CHECK(buffer.scalar_type() == torch::kInt8 && buffer.dim() == 1, + "persistent workspace must be a contiguous int8 vector"); + TORCH_CHECK(topk_ids.is_cuda() && topk_ids.is_contiguous() && + topk_ids.scalar_type() == torch::kInt64 && + topk_ids.sizes() == torch::IntArrayRef({input.size(0), kPersistentTopK}), + "persistent topk_ids must be contiguous CUDA int64 [tokens, 8]"); + TORCH_CHECK(topk_scores.is_cuda() && topk_scores.is_contiguous() && + topk_scores.scalar_type() == torch::kFloat32 && + topk_scores.sizes() == torch::IntArrayRef({input.size(0), kPersistentTopK}), + "persistent topk_scores must be contiguous CUDA float32 [tokens, 8]"); + TORCH_CHECK(input.device() == buffer.device() && + topk_ids.device() == buffer.device() && + topk_scores.device() == buffer.device(), + "persistent inputs and workspace must share a device"); + + const PersistentWorkspaceLayout layout( + buffer.data_ptr(), + static_cast(num_ranks), + static_cast(context_tokens_per_rank)); + TORCH_CHECK(buffer.nbytes() >= static_cast(layout.num_bytes()), + "persistent workspace is smaller than its derived context/CP layout"); + TORCH_CHECK(input.size(0) <= layout.capacity, + "input exceeds the private context/CP workspace capacity"); + + if (input.size(0) == 0) + return; + c10::cuda::CUDAGuard guard(input.device()); + const auto stream = at::cuda::getCurrentCUDAStream(input.get_device()); + sm103_prepare_persistent_inputs_kernel<<< + input.size(0) * (kPersistentHidden / kBlockK), + kBlockK, + 0, + stream>>>( + reinterpret_cast(input.data_ptr()), + topk_ids.data_ptr(), + topk_scores.data_ptr(), + layout.input_tokens.get_base_ptr<__nv_fp8_e4m3>(), + layout.input_scales.get_base_ptr(), + layout.input_topk_ids.get_base_ptr(), + layout.input_topk_scores.get_base_ptr(), + input.size(0)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + +template +void launch_persistent_forward( + const torch::Tensor& output, + const torch::Tensor& expert_counts, + const torch::Tensor& token_src_metadata, + const torch::Tensor& buffer, + const std::vector& buffer_ptrs, + const int64_t rank, + const PersistentWorkspaceLayout& layout, + const torch::Tensor& w13_weight, + const torch::Tensor& w13_scale, + const torch::Tensor& w2_weight, + const torch::Tensor& w2_scale, + const int64_t num_tokens +) { + const auto device = buffer.device(); + const auto fp8_options = torch::TensorOptions() + .dtype(torch::kFloat8_e4m3fn) + .device(device); + const auto int_options = torch::TensorOptions() + .dtype(torch::kInt) + .device(device); + const auto bf16_options = torch::TensorOptions() + .dtype(torch::kBFloat16) + .device(device); + + auto l1_acts = torch::from_blob( + layout.l1_tokens.base, + {layout.ring_tokens, kPersistentHidden}, + fp8_options); + auto l1_acts_sf = torch::from_blob( + layout.l1_scales.base, + {layout.sf_ring_tokens, kPersistentHidden / 128}, + {1, static_cast(layout.sf_ring_tokens)}, + int_options); + auto l2_acts = torch::from_blob( + layout.l2_tokens.base, + {layout.ring_tokens, kPersistentIntermediate}, + fp8_options); + auto l2_acts_sf = torch::from_blob( + layout.l2_scales.base, + {layout.sf_ring_tokens, kPersistentIntermediate / 128}, + {1, static_cast(layout.sf_ring_tokens)}, + int_options); + auto saved_l2_acts = torch::from_blob( + layout.backward_full_h.base, + {static_cast(layout.workspace.num_max_pool_tokens), + kPersistentIntermediate}, + fp8_options); + auto saved_down_unweighted = torch::from_blob( + layout.saved_down_unweighted.base, + {static_cast(layout.workspace.num_max_pool_tokens), + kPersistentHidden}, + bf16_options); + + const auto tensor_map_l1_acts = deep_gemm::make_tma_2d_desc( + l1_acts, + kPersistentHidden, + layout.ring_tokens, + kPersistentBlockK, + kPersistentBlockM / 2, + static_cast(l1_acts.stride(-2)), + 128); + const auto tensor_map_l1_acts_sf = deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, + l1_acts_sf, + layout.sf_ring_tokens, + kPersistentHidden, + kPersistentSFBlockM, + 32, + 1, + 0, + 0, + false, + 1); + // Canonical [2E,H,D] is addressed as a 2-D [2E*H,D] plane. The + // persistent kernel issues separate 64-row TMA offsets for up and gate. + const auto tensor_map_l1_weights = deep_gemm::make_tma_2d_desc( + w13_weight, + kPersistentHidden, + static_cast(w13_weight.size(0) * w13_weight.size(1)), + kPersistentBlockK, + kPersistentBlockN / 2, + static_cast(w13_weight.stride(-2)), + 128); + const auto tensor_map_l1_output = deep_gemm::make_tma_2d_desc( + l2_acts, + kPersistentIntermediate, + layout.ring_tokens, + kPersistentBlockN / 2, + kPersistentStoreBlockM, + static_cast(l2_acts.stride(-2)), + 64); + const auto tensor_map_l1_saved_output = + deep_gemm::make_tma_2d_desc( + saved_l2_acts, + kPersistentIntermediate, + static_cast(layout.workspace.num_max_pool_tokens), + kPersistentBlockN / 2, + kPersistentStoreBlockM, + kPersistentIntermediate, + 64); + const auto tensor_map_l2_acts = deep_gemm::make_tma_2d_desc( + l2_acts, + kPersistentIntermediate, + layout.ring_tokens, + kPersistentBlockK, + kPersistentBlockM / 2, + static_cast(l2_acts.stride(-2)), + 128); + const auto tensor_map_l2_acts_sf = deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, + l2_acts_sf, + layout.sf_ring_tokens, + kPersistentIntermediate, + kPersistentSFBlockM, + 32, + 1, + 0, + 0, + false, + 1); + const auto tensor_map_l2_weights = deep_gemm::make_tma_2d_desc( + w2_weight, + kPersistentIntermediate, + static_cast(w2_weight.size(0) * w2_weight.size(1)), + kPersistentBlockK, + kPersistentBlockN, + static_cast(w2_weight.stride(-2)), + 128); + const auto tensor_map_l2_saved_output = + deep_gemm::make_tma_2d_desc( + saved_down_unweighted, + kPersistentHidden, + static_cast(layout.workspace.num_max_pool_tokens), + kPersistentBlockN, + kPersistentStoreBlockM, + kPersistentHidden, + 128); + + using Kernel = decltype(&deep_gemm::sm100_fp8_fp4_mega_moe_impl< + kPersistentHidden, + kPersistentIntermediate, + kPersistentExperts, + kPersistentTopK, + 1, + kPersistentBlockM, + kPersistentBlockN, + kPersistentBlockK, + kPersistentStoreBlockM, + kPersistentSFBlockM, + kPersistentSFBlockN, + kPersistentStages, + kPersistentPullBytes, + kPersistentDispatchThreads, + kPersistentNonEpilogueThreads, + kPersistentEpilogueThreads, + kNumSMs, + kNumRanks, + 0x7f800000u, + true, + deep_gemm::ActivationType::SwiGLU, + true>); + Kernel kernel = &deep_gemm::sm100_fp8_fp4_mega_moe_impl< + kPersistentHidden, + kPersistentIntermediate, + kPersistentExperts, + kPersistentTopK, + 1, + kPersistentBlockM, + kPersistentBlockN, + kPersistentBlockK, + kPersistentStoreBlockM, + kPersistentSFBlockM, + kPersistentSFBlockN, + kPersistentStages, + kPersistentPullBytes, + kPersistentDispatchThreads, + kPersistentNonEpilogueThreads, + kPersistentEpilogueThreads, + kNumSMs, + kNumRanks, + 0x7f800000u, + true, + deep_gemm::ActivationType::SwiGLU, + true>; + + C10_CUDA_CHECK(cudaFuncSetAttribute( + kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, + kPersistentSmemBytes)); + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeClusterDimension; + attribute.val.clusterDim = {2, 1, 1}; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(kNumSMs, 1, 1); + config.blockDim = dim3(kPersistentThreads, 1, 1); + config.dynamicSmemBytes = kPersistentSmemBytes; + config.stream = at::cuda::getCurrentCUDAStream(buffer.get_device()); + config.attrs = &attribute; + config.numAttrs = 1; + + const auto sym_buffer = deep_gemm::layout::SymBuffer( + buffer_ptrs, static_cast(rank)); + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, + kernel, + output.data_ptr(), + expert_counts.data_ptr(), + reinterpret_cast( + token_src_metadata.data_ptr()), + static_cast(num_tokens), + layout.capacity, + layout.ring_tokens, + layout.sf_ring_tokens, + sym_buffer, + tensor_map_l1_acts, + tensor_map_l1_acts_sf, + tensor_map_l1_weights, + tensor_map_l1_acts_sf, + tensor_map_l1_output, + tensor_map_l1_saved_output, + tensor_map_l2_acts, + tensor_map_l2_acts_sf, + tensor_map_l2_weights, + tensor_map_l2_acts_sf, + tensor_map_l2_saved_output, + w13_scale.data_ptr(), + w2_scale.data_ptr(), + layout.backward_full_x + .get_base_ptr(), + layout.backward_full_x_scales.get_base_ptr(), + layout.backward_full_scores.get_base_ptr(), + layout.saved_l1_preact.get_base_ptr<__nv_bfloat16>(), + layout.backward_full_h + .get_base_ptr(), + layout.backward_full_h_scales.get_base_ptr(), + layout.saved_down_unweighted.get_base_ptr<__nv_bfloat16>())); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + +std::tuple persistent_forward( + const torch::Tensor& buffer, + const std::vector& buffer_ptrs, + const int64_t rank, + const int64_t context_tokens_per_rank, + const int64_t num_tokens, + const torch::Tensor& w13_weight, + const torch::Tensor& w13_scale, + const torch::Tensor& w2_weight, + const torch::Tensor& w2_scale +) { + check_sm103_device(buffer); + c10::cuda::CUDAGuard guard(buffer.device()); + const uint32_t num_sms = get_persistent_sm_count(buffer); + TORCH_CHECK(buffer_ptrs.size() == 2 || buffer_ptrs.size() == 16, + "persistent GLM MegaMoE supports EP2 or EP16 only"); + TORCH_CHECK(rank >= 0 && rank < static_cast(buffer_ptrs.size()), + "persistent workspace rank is out of range"); + TORCH_CHECK(num_tokens >= 0, + "persistent num_tokens must be nonnegative"); + + const auto local_experts = kPersistentExperts / buffer_ptrs.size(); + TORCH_CHECK(w13_weight.is_cuda() && w13_weight.is_contiguous() && + w13_weight.scalar_type() == torch::kFloat8_e4m3fn && + w13_weight.sizes() == torch::IntArrayRef( + {static_cast(2 * local_experts), + kPersistentIntermediate, + kPersistentHidden}), + "canonical W13 must be contiguous E4M3 [2E_local,2048,6144]"); + TORCH_CHECK(w13_scale.is_cuda() && w13_scale.is_contiguous() && + w13_scale.scalar_type() == torch::kFloat32 && + w13_scale.sizes() == torch::IntArrayRef( + {static_cast(2 * local_experts), 16, 48}), + "canonical W13 scales must be FP32 [2E_local,16,48]"); + TORCH_CHECK(w2_weight.is_cuda() && w2_weight.is_contiguous() && + w2_weight.scalar_type() == torch::kFloat8_e4m3fn && + w2_weight.sizes() == torch::IntArrayRef( + {static_cast(local_experts), + kPersistentHidden, + kPersistentIntermediate}), + "W2 must be contiguous E4M3 [E_local,6144,2048]"); + TORCH_CHECK(w2_scale.is_cuda() && w2_scale.is_contiguous() && + w2_scale.scalar_type() == torch::kFloat32 && + w2_scale.sizes() == torch::IntArrayRef( + {static_cast(local_experts), 48, 16}), + "W2 scales must be FP32 [E_local,48,16]"); + TORCH_CHECK(w13_weight.device() == buffer.device() && + w13_scale.device() == buffer.device() && + w2_weight.device() == buffer.device() && + w2_scale.device() == buffer.device(), + "persistent weights and workspace must share a device"); + + const PersistentWorkspaceLayout layout( + buffer.data_ptr(), + static_cast(buffer_ptrs.size()), + static_cast(context_tokens_per_rank)); + TORCH_CHECK(buffer.nbytes() >= static_cast(layout.num_bytes()), + "persistent workspace is smaller than its derived layout"); + TORCH_CHECK(num_tokens <= layout.capacity, + "persistent num_tokens exceeds the private context/CP capacity"); + auto output = torch::empty( + {num_tokens, kPersistentHidden}, + buffer.options().dtype(torch::kBFloat16)); + auto expert_counts = torch::zeros( + {static_cast(local_experts)}, + buffer.options().dtype(torch::kInt)); + auto token_src_metadata = torch::empty( + {static_cast(layout.workspace.num_max_pool_tokens), 3}, + buffer.options().dtype(torch::kInt)); + if (num_tokens == 0) + return {output, expert_counts, token_src_metadata}; + + if (num_sms == kPersistentLocalSMs) { + if (buffer_ptrs.size() == 2) { + launch_persistent_forward<2, kPersistentLocalSMs>( + output, expert_counts, token_src_metadata, + buffer, buffer_ptrs, rank, layout, + w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + } else { + launch_persistent_forward<16, kPersistentLocalSMs>( + output, expert_counts, token_src_metadata, + buffer, buffer_ptrs, rank, layout, + w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + } + } else { + if (buffer_ptrs.size() == 2) { + launch_persistent_forward<2, kPersistentProductionSMs>( + output, expert_counts, token_src_metadata, + buffer, buffer_ptrs, rank, layout, + w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + } else { + launch_persistent_forward<16, kPersistentProductionSMs>( + output, expert_counts, token_src_metadata, + buffer, buffer_ptrs, rank, layout, + w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + } + } + return {output, expert_counts, token_src_metadata}; +} + +template +void launch_persistent_backward_activation( + const torch::Tensor& grad_x, + const torch::Tensor& grad_scores, + const torch::Tensor& buffer, + const std::vector& buffer_ptrs, + const int64_t rank, + const PersistentWorkspaceLayout& layout, + const torch::Tensor& x, + const torch::Tensor& grad_output, + const torch::Tensor& topk_scores, + const torch::Tensor& expert_counts, + const torch::Tensor& token_src_metadata, + const torch::Tensor& w13_weight, + const torch::Tensor& w13_scale, + const torch::Tensor& w2_weight, + const torch::Tensor& w2_scale +) { + const auto device = buffer.device(); + const auto fp8_options = torch::TensorOptions() + .dtype(torch::kFloat8_e4m3fn) + .device(device); + const auto int_options = torch::TensorOptions() + .dtype(torch::kInt) + .device(device); + + auto ring_grad_y = torch::from_blob( + layout.backward_ring_grad_y.base, + {layout.ring_tokens, kPersistentHidden}, fp8_options); + auto ring_grad_y_sf = torch::from_blob( + layout.backward_ring_grad_y_scales.base, + {layout.sf_ring_tokens, kPersistentHidden / 128}, + {1, static_cast(layout.sf_ring_tokens)}, + int_options); + auto ring_grad_preact = torch::from_blob( + layout.backward_ring_grad_preact.base, + {layout.ring_tokens, 2 * kPersistentIntermediate}, fp8_options); + auto ring_grad_preact_sf = torch::from_blob( + layout.backward_ring_grad_preact_scales.base, + {layout.sf_ring_tokens, 2 * kPersistentIntermediate / 128}, + {1, static_cast(layout.sf_ring_tokens)}, + int_options); + const auto tensor_map_ring_grad_y = deep_gemm::make_tma_2d_desc( + ring_grad_y, kPersistentHidden, layout.ring_tokens, + kPersistentBlockK, kPersistentBlockM / 2, + kPersistentHidden, 128); + const auto tensor_map_ring_grad_y_sf = deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, ring_grad_y_sf, + layout.sf_ring_tokens, kPersistentHidden, + kPersistentSFBlockM, 32, 1, 0, 0, false, 1); + const auto tensor_map_ring_grad_preact = + deep_gemm::make_tma_2d_desc( + ring_grad_preact, 2 * kPersistentIntermediate, + layout.ring_tokens, kPersistentBlockK, + kPersistentBlockM / 2, + 2 * kPersistentIntermediate, 128); + const auto tensor_map_ring_grad_preact_sf = + deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, ring_grad_preact_sf, + layout.sf_ring_tokens, 2 * kPersistentIntermediate, + kPersistentSFBlockM, 32, 1, 0, 0, false, 1); + const auto tensor_map_w2_dgrad = deep_gemm::make_tma_b_desc( + cute::UMMA::Major::MN, w2_weight, + kPersistentIntermediate, kPersistentHidden, + kPersistentBlockN, kPersistentBlockK, + static_cast(w2_weight.stride(-2)), + static_cast(w2_weight.size(0)), 128); + const auto tensor_map_w13_dgrad = deep_gemm::make_tma_b_desc( + cute::UMMA::Major::MN, w13_weight, + kPersistentHidden, 2 * kPersistentIntermediate, + kPersistentBlockN, kPersistentBlockK, + static_cast(w13_weight.stride(-2)), + static_cast(w13_weight.size(0) / 2), 128); + using Kernel = decltype( + &deep_gemm::sm103_block128_backward:: + sm103_fp8_block128_mega_moe_backward_persistent_impl< + kNumRanks, kNumSMs>); + Kernel kernel = + &deep_gemm::sm103_block128_backward:: + sm103_fp8_block128_mega_moe_backward_persistent_impl< + kNumRanks, kNumSMs>; + constexpr uint32_t smem_bytes = sizeof( + deep_gemm::sm103_block128_backward::SharedStorage); + C10_CUDA_CHECK(cudaFuncSetAttribute( + kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes)); + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeClusterDimension; + attribute.val.clusterDim = {2, 1, 1}; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(kNumSMs, 1, 1); + config.blockDim = dim3( + deep_gemm::sm103_block128_backward::kThreads, 1, 1); + config.dynamicSmemBytes = smem_bytes; + config.stream = at::cuda::getCurrentCUDAStream(buffer.get_device()); + config.attrs = &attribute; + config.numAttrs = 1; + const auto sym_buffer = deep_gemm::layout::SymBuffer( + buffer_ptrs, static_cast(rank)); + + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, kernel, + expert_counts.data_ptr(), + reinterpret_cast( + token_src_metadata.data_ptr()), + static_cast(x.size(0)), layout.capacity, + layout.ring_tokens, layout.sf_ring_tokens, + layout.workspace.num_max_pool_tokens, + sym_buffer, layout.workspace, + reinterpret_cast( + grad_output.data_ptr()), + // Reverse transport retains upstream's independent immutable input + // and remote-combine output lifetimes. + layout.backward_symmetric_grad_y + .get_base_ptr(), + layout.combine_tokens.get_base_ptr(), + layout.backward_ring_grad_y + .get_base_ptr(), + layout.backward_ring_grad_y_scales.get_base_ptr(), + layout.backward_ring_grad_preact + .get_base_ptr(), + layout.backward_ring_grad_preact_scales.get_base_ptr(), + layout.backward_full_grad_y.get_base_ptr(), + layout.backward_full_grad_y_scales.get_base_ptr(), + layout.backward_full_scores.get_base_ptr(), + layout.saved_l1_preact.get_base_ptr(), + layout.saved_down_unweighted.get_base_ptr(), + layout.backward_full_grad_preact + .get_base_ptr(), + layout.backward_full_grad_preact_scales.get_base_ptr(), + reinterpret_cast(grad_x.data_ptr()), + layout.backward_grad_scores.get_base_ptr(), + grad_scores.data_ptr(), + tensor_map_ring_grad_y, tensor_map_ring_grad_y_sf, + tensor_map_ring_grad_preact, tensor_map_ring_grad_preact_sf, + tensor_map_w2_dgrad, tensor_map_w13_dgrad, + w13_scale.data_ptr(), w2_scale.data_ptr())); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + +template +void launch_persistent_wgrad( + const torch::Tensor& output_0, + const torch::Tensor& output_1, + const torch::Tensor& buffer, + const PersistentWorkspaceLayout& layout, + const torch::Tensor& expert_counts +) { + constexpr int64_t output_rows = + kW2 ? kPersistentHidden : kPersistentIntermediate; + constexpr int64_t output_columns = + kW2 ? kPersistentIntermediate : kPersistentHidden; + const int64_t local_experts = kPersistentExperts / kNumRanks; + void* full_a_base = kW2 + ? layout.backward_full_grad_y.base + : layout.backward_full_grad_preact.base; + void* full_b_base = kW2 + ? layout.backward_full_h.base + : layout.backward_full_x.base; + const uint32_t* full_a_sf = kW2 + ? layout.backward_full_grad_y_scales.get_base_ptr() + : layout.backward_full_grad_preact_scales.get_base_ptr(); + const uint32_t* full_b_sf = kW2 + ? layout.backward_full_h_scales.get_base_ptr() + : layout.backward_full_x_scales.get_base_ptr(); + const auto output_0_flat = output_0.view( + {local_experts * output_rows, output_columns}); + const auto output_1_flat = output_1.view( + {local_experts * output_rows, output_columns}); + const auto tensor_map_output_0 = deep_gemm::make_tma_cd_desc( + output_0_flat, + static_cast(local_experts * output_rows), + static_cast(output_columns), + deep_gemm::sm103_block128_wgrad::kStoreBlockM, + deep_gemm::sm103_block128_wgrad::kStoreBlockN, + static_cast(output_columns), 1, 128); + const auto tensor_map_output_1 = deep_gemm::make_tma_cd_desc( + output_1_flat, + static_cast(local_experts * output_rows), + static_cast(output_columns), + deep_gemm::sm103_block128_wgrad::kStoreBlockM, + deep_gemm::sm103_block128_wgrad::kStoreBlockN, + static_cast(output_columns), 1, 128); + + using Kernel = decltype( + &deep_gemm::sm103_block128_wgrad:: + sm103_fp8_block128_mega_moe_wgrad_impl< + kW2, kNumRanks, kNumSMs>); + Kernel kernel = + &deep_gemm::sm103_block128_wgrad:: + sm103_fp8_block128_mega_moe_wgrad_impl< + kW2, kNumRanks, kNumSMs>; + constexpr uint32_t smem_bytes = sizeof( + deep_gemm::sm103_block128_wgrad::SharedStorage); + C10_CUDA_CHECK(cudaFuncSetAttribute( + kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes)); + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeClusterDimension; + attribute.val.clusterDim = {2, 1, 1}; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(kNumSMs, 1, 1); + config.blockDim = dim3( + deep_gemm::sm103_block128_wgrad::kThreads, 1, 1); + config.dynamicSmemBytes = smem_bytes; + config.stream = at::cuda::getCurrentCUDAStream(buffer.get_device()); + config.attrs = &attribute; + config.numAttrs = 1; + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, kernel, + expert_counts.data_ptr(), + layout.workspace.num_max_pool_tokens, + static_cast(full_a_base), + static_cast(full_b_base), + full_a_sf, + full_b_sf, + tensor_map_output_0, tensor_map_output_1)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + +std::tuple +persistent_backward_activation( + const torch::Tensor& buffer, + const std::vector& buffer_ptrs, + const int64_t rank, + const int64_t context_tokens_per_rank, + const torch::Tensor& x, + const torch::Tensor& grad_output, + const torch::Tensor& topk_scores, + const torch::Tensor& expert_counts, + const torch::Tensor& token_src_metadata, + const torch::Tensor& w13_weight, + const torch::Tensor& w13_scale, + const torch::Tensor& w2_weight, + const torch::Tensor& w2_scale +) { + check_sm103_device(buffer); + DG_CHECK_CONTIGUOUS(buffer); + TORCH_CHECK(buffer.scalar_type() == torch::kInt8 && buffer.dim() == 1, + "persistent workspace must be a contiguous int8 vector"); + TORCH_CHECK(buffer_ptrs.size() == 2 || buffer_ptrs.size() == 16, + "persistent backward supports EP2 or EP16 only"); + TORCH_CHECK(rank >= 0 && rank < static_cast(buffer_ptrs.size()), + "persistent backward workspace rank is out of range"); + TORCH_CHECK(context_tokens_per_rank > 0 && + context_tokens_per_rank <= + std::numeric_limits::max(), + "persistent backward context/CP envelope must be positive uint32"); + c10::cuda::CUDAGuard guard(buffer.device()); + const uint32_t num_sms = get_persistent_sm_count(buffer); + check_bf16_matrix(x, "x"); + check_bf16_matrix(grad_output, "grad_output"); + TORCH_CHECK(x.sizes() == grad_output.sizes() && + x.size(1) == kPersistentHidden, + "persistent backward x/grad_output must match [tokens,6144]"); + TORCH_CHECK(topk_scores.is_cuda() && topk_scores.is_contiguous() && + topk_scores.scalar_type() == torch::kFloat32 && + topk_scores.sizes() == torch::IntArrayRef( + {x.size(0), kPersistentTopK}), + "persistent backward scores must be float32 [tokens,8]"); + TORCH_CHECK(x.device() == buffer.device() && + grad_output.device() == buffer.device() && + topk_scores.device() == buffer.device(), + "persistent backward activations and workspace must share a device"); + const auto local_experts = + kPersistentExperts / buffer_ptrs.size(); + TORCH_CHECK(expert_counts.is_cuda() && expert_counts.is_contiguous() && + expert_counts.scalar_type() == torch::kInt && + expert_counts.numel() == + static_cast(local_experts), + "saved expert counts mismatch"); + TORCH_CHECK(token_src_metadata.is_cuda() && + token_src_metadata.is_contiguous() && + token_src_metadata.scalar_type() == torch::kInt && + token_src_metadata.dim() == 2 && + token_src_metadata.size(1) == 3, + "saved source metadata must be int32 [pool,3]"); + TORCH_CHECK(expert_counts.device() == buffer.device() && + token_src_metadata.device() == buffer.device(), + "saved routing state and workspace must share a device"); + TORCH_CHECK(w13_weight.is_cuda() && w13_weight.is_contiguous() && + w13_weight.scalar_type() == torch::kFloat8_e4m3fn && + w13_weight.sizes() == torch::IntArrayRef( + {static_cast(2 * local_experts), + kPersistentIntermediate, kPersistentHidden}), + "canonical W13 must be contiguous E4M3 [2E_local,2048,6144]"); + TORCH_CHECK(w13_scale.is_cuda() && w13_scale.is_contiguous() && + w13_scale.scalar_type() == torch::kFloat32 && + w13_scale.sizes() == torch::IntArrayRef( + {static_cast(2 * local_experts), 16, 48}), + "canonical W13 scales must be FP32 [2E_local,16,48]"); + TORCH_CHECK(w2_weight.is_cuda() && w2_weight.is_contiguous() && + w2_weight.scalar_type() == torch::kFloat8_e4m3fn && + w2_weight.sizes() == torch::IntArrayRef( + {static_cast(local_experts), + kPersistentHidden, kPersistentIntermediate}), + "W2 must be contiguous E4M3 [E_local,6144,2048]"); + TORCH_CHECK(w2_scale.is_cuda() && w2_scale.is_contiguous() && + w2_scale.scalar_type() == torch::kFloat32 && + w2_scale.sizes() == torch::IntArrayRef( + {static_cast(local_experts), 48, 16}), + "W2 scales must be FP32 [E_local,48,16]"); + TORCH_CHECK(w13_weight.device() == buffer.device() && + w13_scale.device() == buffer.device() && + w2_weight.device() == buffer.device() && + w2_scale.device() == buffer.device(), + "persistent backward weights and workspace must share a device"); + + const PersistentWorkspaceLayout layout( + buffer.data_ptr(), static_cast(buffer_ptrs.size()), + static_cast(context_tokens_per_rank)); + TORCH_CHECK(buffer.nbytes() >= static_cast(layout.num_bytes()), + "persistent backward workspace is smaller than its derived layout"); + TORCH_CHECK(x.size(0) <= layout.capacity, + "persistent backward input exceeds context/CP capacity"); + TORCH_CHECK(token_src_metadata.size(0) >= + static_cast(layout.workspace.num_max_pool_tokens), + "saved source metadata does not cover the full route pool"); + auto grad_x = torch::empty_like(x); + auto grad_scores = torch::empty_like(topk_scores); + if (x.size(0) == 0) { + return {grad_x, grad_scores}; + } + if (num_sms == kPersistentLocalSMs) { + if (buffer_ptrs.size() == 2) { + launch_persistent_backward_activation<2, kPersistentLocalSMs>( + grad_x, grad_scores, buffer, buffer_ptrs, rank, layout, + x, grad_output, topk_scores, expert_counts, + token_src_metadata, w13_weight, w13_scale, + w2_weight, w2_scale); + } else { + launch_persistent_backward_activation<16, kPersistentLocalSMs>( + grad_x, grad_scores, buffer, buffer_ptrs, rank, layout, + x, grad_output, topk_scores, expert_counts, + token_src_metadata, w13_weight, w13_scale, + w2_weight, w2_scale); + } + } else { + if (buffer_ptrs.size() == 2) { + launch_persistent_backward_activation<2, kPersistentProductionSMs>( + grad_x, grad_scores, buffer, buffer_ptrs, rank, layout, + x, grad_output, topk_scores, expert_counts, + token_src_metadata, w13_weight, w13_scale, + w2_weight, w2_scale); + } else { + launch_persistent_backward_activation<16, kPersistentProductionSMs>( + grad_x, grad_scores, buffer, buffer_ptrs, rank, layout, + x, grad_output, topk_scores, expert_counts, + token_src_metadata, w13_weight, w13_scale, + w2_weight, w2_scale); + } + } + return {grad_x, grad_scores}; +} + +std::tuple +persistent_backward( + const torch::Tensor& buffer, + const std::vector& buffer_ptrs, + const int64_t rank, + const int64_t context_tokens_per_rank, + const torch::Tensor& x, + const torch::Tensor& grad_output, + const torch::Tensor& topk_scores, + const torch::Tensor& expert_counts, + const torch::Tensor& token_src_metadata, + const torch::Tensor& w13_weight, + const torch::Tensor& w13_scale, + const torch::Tensor& w2_weight, + const torch::Tensor& w2_scale +) { + auto [grad_x, grad_scores] = persistent_backward_activation( + buffer, buffer_ptrs, rank, context_tokens_per_rank, + x, grad_output, topk_scores, expert_counts, token_src_metadata, + w13_weight, w13_scale, w2_weight, w2_scale); + + const int64_t local_experts = + kPersistentExperts / static_cast(buffer_ptrs.size()); + const auto options = x.options().dtype(torch::kBFloat16); + auto grad_w1 = torch::empty( + {local_experts, kPersistentIntermediate, kPersistentHidden}, + options); + auto grad_w2 = torch::empty( + {local_experts, kPersistentHidden, kPersistentIntermediate}, + options); + auto grad_w3 = torch::empty( + {local_experts, kPersistentIntermediate, kPersistentHidden}, + options); + if (x.size(0) == 0) { + grad_w1.zero_(); + grad_w2.zero_(); + grad_w3.zero_(); + return {grad_x, grad_scores, grad_w1, grad_w2, grad_w3}; + } + + const uint32_t num_sms = get_persistent_sm_count(buffer); + + const PersistentWorkspaceLayout layout( + buffer.data_ptr(), static_cast(buffer_ptrs.size()), + static_cast(context_tokens_per_rank)); + if (num_sms == kPersistentLocalSMs) { + if (buffer_ptrs.size() == 2) { + launch_persistent_wgrad<2, kPersistentLocalSMs, true>( + grad_w2, grad_w2, buffer, layout, expert_counts); + launch_persistent_wgrad<2, kPersistentLocalSMs, false>( + grad_w1, grad_w3, buffer, layout, expert_counts); + } else { + launch_persistent_wgrad<16, kPersistentLocalSMs, true>( + grad_w2, grad_w2, buffer, layout, expert_counts); + launch_persistent_wgrad<16, kPersistentLocalSMs, false>( + grad_w1, grad_w3, buffer, layout, expert_counts); + } + } else { + if (buffer_ptrs.size() == 2) { + launch_persistent_wgrad<2, kPersistentProductionSMs, true>( + grad_w2, grad_w2, buffer, layout, expert_counts); + launch_persistent_wgrad<2, kPersistentProductionSMs, false>( + grad_w1, grad_w3, buffer, layout, expert_counts); + } else { + launch_persistent_wgrad<16, kPersistentProductionSMs, true>( + grad_w2, grad_w2, buffer, layout, expert_counts); + launch_persistent_wgrad<16, kPersistentProductionSMs, false>( + grad_w1, grad_w3, buffer, layout, expert_counts); + } + } + return {grad_x, grad_scores, grad_w1, grad_w2, grad_w3}; +} + +std::tuple quantize_bf16(const torch::Tensor& input) { + check_bf16_matrix(input, "input"); + c10::cuda::CUDAGuard guard(input.device()); + const auto rows = input.size(0); + const auto columns = input.size(1); + auto output = torch::empty(input.sizes(), input.options().dtype(torch::kFloat8_e4m3fn)); + auto scales = torch::empty({rows, columns / kBlockK}, input.options().dtype(torch::kFloat32)); + if (rows != 0) { + const auto stream = at::cuda::getCurrentCUDAStream(input.get_device()); + sm103_quantize_bf16_e4m3_group128_kernel<<>>( + reinterpret_cast(input.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(output.data_ptr()), + scales.data_ptr(), rows, columns + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return {output, scales}; +} + +torch::Tensor dequantize_fp8( + const torch::Tensor& input, + const torch::Tensor& scales +) { + check_fp8_matrix_and_scales(input, scales, "input"); + c10::cuda::CUDAGuard guard(input.device()); + const auto rows = input.size(0); + const auto columns = input.size(1); + auto output = torch::empty(input.sizes(), input.options().dtype(torch::kBFloat16)); + if (input.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (input.numel() + threads - 1) / threads; + const auto stream = at::cuda::getCurrentCUDAStream(input.get_device()); + sm103_dequantize_e4m3_group128_kernel<<>>( + reinterpret_cast(input.data_ptr()), + scales.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), rows, columns + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +std::tuple swiglu_quantize(const torch::Tensor& preactivation) { + check_bf16_matrix(preactivation, "preactivation"); + TORCH_CHECK(preactivation.size(1) % (2 * kBlockK) == 0, "preactivation width must be 2 * H with H divisible by 128"); + c10::cuda::CUDAGuard guard(preactivation.device()); + const auto rows = preactivation.size(0); + const auto hidden = preactivation.size(1) / 2; + auto output = torch::empty({rows, hidden}, preactivation.options().dtype(torch::kFloat8_e4m3fn)); + auto scales = torch::empty({rows, hidden / kBlockK}, preactivation.options().dtype(torch::kFloat32)); + if (rows != 0) { + const auto stream = at::cuda::getCurrentCUDAStream(preactivation.get_device()); + sm103_swiglu_quantize_group128_kernel<<>>( + reinterpret_cast(preactivation.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(output.data_ptr()), + scales.data_ptr(), rows, hidden + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return {output, scales}; +} + +torch::Tensor swiglu_backward( + const torch::Tensor& grad_output, + const torch::Tensor& preactivation +) { + check_bf16_matrix(grad_output, "grad_output"); + check_sm103_device(preactivation); + DG_CHECK_CONTIGUOUS(preactivation); + TORCH_CHECK(preactivation.scalar_type() == torch::kBFloat16, "preactivation must be bfloat16"); + TORCH_CHECK(preactivation.dim() == 2, "preactivation must be rank 2"); + TORCH_CHECK(preactivation.size(0) == grad_output.size(0), "row count mismatch"); + TORCH_CHECK(preactivation.size(1) == grad_output.size(1) * 2, "preactivation width mismatch"); + TORCH_CHECK(preactivation.device() == grad_output.device(), "device mismatch"); + c10::cuda::CUDAGuard guard(grad_output.device()); + auto grad_preactivation = torch::empty_like(preactivation); + if (grad_output.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (grad_output.numel() + threads - 1) / threads; + const auto stream = at::cuda::getCurrentCUDAStream(grad_output.get_device()); + sm103_swiglu_backward_kernel<<>>( + reinterpret_cast(grad_output.data_ptr()), + reinterpret_cast(preactivation.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(grad_preactivation.data_ptr()), + grad_output.size(0), grad_output.size(1) + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return grad_preactivation; +} + +torch::Tensor swiglu_backward_canonical( + const torch::Tensor& grad_output, + const torch::Tensor& preactivation +) { + check_bf16_matrix(grad_output, "grad_output"); + check_sm103_device(preactivation); + DG_CHECK_CONTIGUOUS(preactivation); + TORCH_CHECK(preactivation.scalar_type() == torch::kBFloat16, + "preactivation must be bfloat16"); + TORCH_CHECK(preactivation.dim() == 2, "preactivation must be rank 2"); + TORCH_CHECK(preactivation.size(0) == grad_output.size(0), "row count mismatch"); + TORCH_CHECK(preactivation.size(1) == grad_output.size(1) * 2, + "preactivation width mismatch"); + TORCH_CHECK(preactivation.device() == grad_output.device(), "device mismatch"); + c10::cuda::CUDAGuard guard(grad_output.device()); + auto grad_preactivation = torch::empty_like(preactivation); + if (grad_output.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (grad_output.numel() + threads - 1) / threads; + const auto stream = at::cuda::getCurrentCUDAStream(grad_output.get_device()); + sm103_swiglu_backward_canonical_kernel<<>>( + reinterpret_cast(grad_output.data_ptr()), + reinterpret_cast(preactivation.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(grad_preactivation.data_ptr()), + grad_output.size(0), grad_output.size(1) + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return grad_preactivation; +} + +std::tuple route_scale_quantize( + const torch::Tensor& grad_output, + const torch::Tensor& route_scores, + const torch::Tensor& route_order +) { + check_bf16_matrix(grad_output, "grad_output"); + check_sm103_device(route_scores); + DG_CHECK_CONTIGUOUS(route_scores); + DG_CHECK_CUDA(route_order); + DG_CHECK_CONTIGUOUS(route_order); + TORCH_CHECK(route_scores.dim() == 2 && route_scores.scalar_type() == torch::kFloat32, + "route_scores must be contiguous rank-2 float32"); + TORCH_CHECK(route_order.dim() == 1 && route_order.scalar_type() == torch::kInt64, + "route_order must be contiguous rank-1 int64"); + TORCH_CHECK(route_scores.size(0) == grad_output.size(0), "token count mismatch"); + TORCH_CHECK(route_scores.device() == grad_output.device() && route_order.device() == grad_output.device(), + "all tensors must share a device"); + c10::cuda::CUDAGuard guard(grad_output.device()); + const auto routes = route_order.numel(); + const auto hidden = grad_output.size(1); + auto output = torch::empty({routes, hidden}, grad_output.options().dtype(torch::kFloat8_e4m3fn)); + auto scales = torch::empty({routes, hidden / kBlockK}, grad_output.options().dtype(torch::kFloat32)); + if (routes != 0) { + const auto stream = at::cuda::getCurrentCUDAStream(grad_output.get_device()); + sm103_route_scale_quantize_group128_kernel<<>>( + reinterpret_cast(grad_output.data_ptr()), + route_scores.data_ptr(), route_order.data_ptr(), + reinterpret_cast<__nv_fp8_e4m3*>(output.data_ptr()), scales.data_ptr(), + routes, hidden, route_scores.size(1) + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return {output, scales}; +} + +torch::Tensor post_down_combine( + const torch::Tensor& route_output, + const torch::Tensor& route_scores +) { + check_bf16_matrix(route_output, "route_output"); + check_sm103_device(route_scores); + DG_CHECK_CONTIGUOUS(route_scores); + TORCH_CHECK(route_scores.dim() == 2 && route_scores.scalar_type() == torch::kFloat32, + "route_scores must be contiguous rank-2 float32"); + TORCH_CHECK(route_output.size(0) == route_scores.numel(), "route count mismatch"); + TORCH_CHECK(route_scores.device() == route_output.device(), "device mismatch"); + c10::cuda::CUDAGuard guard(route_output.device()); + auto output = torch::empty( + {route_scores.size(0), route_output.size(1)}, + route_output.options().dtype(torch::kBFloat16) + ); + if (output.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (output.numel() + threads - 1) / threads; + const auto stream = at::cuda::getCurrentCUDAStream(route_output.get_device()); + sm103_post_down_combine_kernel<<>>( + reinterpret_cast(route_output.data_ptr()), + route_scores.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + route_scores.size(0), route_output.size(1), route_scores.size(1) + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +torch::Tensor route_sum( + const torch::Tensor& route_grad, + int64_t num_tokens, + int64_t topk +) { + check_bf16_matrix(route_grad, "route_grad"); + TORCH_CHECK(num_tokens >= 0 && topk > 0, "invalid token/top-k dimensions"); + TORCH_CHECK(route_grad.size(0) == num_tokens * topk, "route count mismatch"); + c10::cuda::CUDAGuard guard(route_grad.device()); + auto output = torch::empty({num_tokens, route_grad.size(1)}, route_grad.options()); + if (output.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (output.numel() + threads - 1) / threads; + const auto stream = at::cuda::getCurrentCUDAStream(route_grad.get_device()); + sm103_route_sum_kernel<<>>( + reinterpret_cast(route_grad.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + num_tokens, route_grad.size(1), topk + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +torch::Tensor post_down_score_grad( + const torch::Tensor& route_output, + const torch::Tensor& grad_output, + int64_t topk +) { + check_bf16_matrix(route_output, "route_output"); + check_bf16_matrix(grad_output, "grad_output"); + TORCH_CHECK(topk > 0, "topk must be positive"); + TORCH_CHECK(route_output.size(0) == grad_output.size(0) * topk, "route count mismatch"); + TORCH_CHECK(route_output.size(1) == grad_output.size(1), "hidden dimension mismatch"); + TORCH_CHECK(route_output.device() == grad_output.device(), "device mismatch"); + c10::cuda::CUDAGuard guard(route_output.device()); + auto output = torch::empty({grad_output.size(0), topk}, grad_output.options().dtype(torch::kFloat32)); + if (route_output.size(0) != 0) { + constexpr int threads = 256; + const auto stream = at::cuda::getCurrentCUDAStream(route_output.get_device()); + sm103_post_down_score_grad_kernel<<>>( + reinterpret_cast(route_output.data_ptr()), + reinterpret_cast(grad_output.data_ptr()), + output.data_ptr(), route_output.size(0), route_output.size(1), topk + ); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +void check_expanded_routes(const torch::Tensor& routes, const torch::Device& device) { + check_sm103_device(routes); + TORCH_CHECK(routes.dim() == 2, "expanded routes must be rank 2"); + TORCH_CHECK(routes.scalar_type() == torch::kInt32 || routes.scalar_type() == torch::kInt64, + "expanded routes must be int32 or int64"); + TORCH_CHECK(routes.device() == device, "expanded routes must share the payload device"); + TORCH_CHECK(routes.stride(1) > 0 && routes.stride(0) > 0, + "expanded routes must have positive strides"); +} + +torch::Tensor expanded_post_down_scale( + const torch::Tensor& route_output, + const torch::Tensor& route_scores +) { + check_bf16_matrix(route_output, "route_output"); + check_sm103_device(route_scores); + DG_CHECK_CONTIGUOUS(route_scores); + TORCH_CHECK(route_scores.scalar_type() == torch::kFloat32 && route_scores.dim() == 1, + "expanded route scores must be contiguous rank-1 float32"); + TORCH_CHECK(route_scores.size(0) == route_output.size(0), + "expanded route score row count mismatch"); + TORCH_CHECK(route_scores.device() == route_output.device(), "device mismatch"); + auto output = torch::empty_like(route_output); + if (output.numel() != 0) { + constexpr int threads = 256; + const auto blocks = (output.numel() + threads - 1) / threads; + c10::cuda::CUDAGuard guard(route_output.device()); + const auto stream = at::cuda::getCurrentCUDAStream(route_output.get_device()); + sm103_expanded_post_down_scale_kernel<<>>( + reinterpret_cast(route_output.data_ptr()), + route_scores.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + route_output.size(0), + route_output.size(1)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +torch::Tensor expand_compact_routes( + const torch::Tensor& compact, + const torch::Tensor& routes, + const int64_t expanded_rows +) { + check_bf16_matrix(compact, "compact"); + check_expanded_routes(routes, compact.device()); + TORCH_CHECK(expanded_rows >= 0, "expanded row count must be non-negative"); + TORCH_CHECK(routes.size(0) == compact.size(0), + "route metadata and compact payload row counts differ"); + auto output = torch::empty( + {expanded_rows, compact.size(1)}, compact.options().dtype(torch::kBFloat16)); + if (routes.numel() != 0) { + constexpr int threads = 256; + c10::cuda::CUDAGuard guard(compact.device()); + const auto stream = at::cuda::getCurrentCUDAStream(compact.get_device()); + const auto blocks = routes.numel(); + if (routes.scalar_type() == torch::kInt32) { + sm103_expand_compact_routes_kernel<<>>( + reinterpret_cast(compact.data_ptr()), + routes.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + compact.size(0), routes.size(1), compact.size(1), expanded_rows, + routes.stride(0), routes.stride(1)); + } else { + sm103_expand_compact_routes_kernel<<>>( + reinterpret_cast(compact.data_ptr()), + routes.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + compact.size(0), routes.size(1), compact.size(1), expanded_rows, + routes.stride(0), routes.stride(1)); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +torch::Tensor collapse_expanded_routes( + const torch::Tensor& expanded, + const torch::Tensor& routes +) { + check_bf16_matrix(expanded, "expanded"); + check_expanded_routes(routes, expanded.device()); + auto output = torch::empty( + {routes.size(0), expanded.size(1)}, expanded.options().dtype(torch::kBFloat16)); + if (output.numel() != 0) { + constexpr int threads = 256; + c10::cuda::CUDAGuard guard(expanded.device()); + const auto stream = at::cuda::getCurrentCUDAStream(expanded.get_device()); + const auto blocks = routes.size(0); + if (routes.scalar_type() == torch::kInt32) { + sm103_collapse_expanded_routes_kernel<<>>( + reinterpret_cast(expanded.data_ptr()), + routes.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + routes.size(0), routes.size(1), expanded.size(1), expanded.size(0), + routes.stride(0), routes.stride(1)); + } else { + sm103_collapse_expanded_routes_kernel<<>>( + reinterpret_cast(expanded.data_ptr()), + routes.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), + routes.size(0), routes.size(1), expanded.size(1), expanded.size(0), + routes.stride(0), routes.stride(1)); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +torch::Tensor expanded_post_down_score_grad( + const torch::Tensor& route_output, + const torch::Tensor& grad_output, + const torch::Tensor& routes +) { + check_bf16_matrix(route_output, "route_output"); + check_bf16_matrix(grad_output, "grad_output"); + check_expanded_routes(routes, route_output.device()); + TORCH_CHECK(route_output.sizes() == grad_output.sizes(), + "expanded route output and gradient shapes differ"); + TORCH_CHECK(route_output.device() == grad_output.device(), "device mismatch"); + auto output = torch::empty( + routes.sizes(), route_output.options().dtype(torch::kFloat32)); + if (routes.numel() != 0) { + constexpr int threads = 256; + c10::cuda::CUDAGuard guard(route_output.device()); + const auto stream = at::cuda::getCurrentCUDAStream(route_output.get_device()); + const auto blocks = routes.numel(); + if (routes.scalar_type() == torch::kInt32) { + sm103_expanded_post_down_score_grad_kernel<<>>( + reinterpret_cast(route_output.data_ptr()), + reinterpret_cast(grad_output.data_ptr()), + routes.data_ptr(), output.data_ptr(), + routes.size(0), routes.size(1), route_output.size(1), route_output.size(0), + routes.stride(0), routes.stride(1)); + } else { + sm103_expanded_post_down_score_grad_kernel<<>>( + reinterpret_cast(route_output.data_ptr()), + reinterpret_cast(grad_output.data_ptr()), + routes.data_ptr(), output.data_ptr(), + routes.size(0), routes.size(1), route_output.size(1), route_output.size(0), + routes.stride(0), routes.stride(1)); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return output; +} + +std::tuple expanded_route_scale_quantize( + const torch::Tensor& grad_output, + const torch::Tensor& route_scores +) { + check_bf16_matrix(grad_output, "grad_output"); + check_sm103_device(route_scores); + DG_CHECK_CONTIGUOUS(route_scores); + TORCH_CHECK(route_scores.scalar_type() == torch::kFloat32 && route_scores.dim() == 1, + "expanded route scores must be contiguous rank-1 float32"); + TORCH_CHECK(route_scores.size(0) == grad_output.size(0), + "expanded route score row count mismatch"); + TORCH_CHECK(route_scores.device() == grad_output.device(), "device mismatch"); + const auto rows = grad_output.size(0); + const auto hidden = grad_output.size(1); + auto output = torch::empty( + grad_output.sizes(), grad_output.options().dtype(torch::kFloat8_e4m3fn)); + auto scales = torch::empty( + {rows, hidden / kBlockK}, grad_output.options().dtype(torch::kFloat32)); + if (rows != 0) { + c10::cuda::CUDAGuard guard(grad_output.device()); + const auto stream = at::cuda::getCurrentCUDAStream(grad_output.get_device()); + sm103_expanded_route_scale_quantize_group128_kernel<<< + rows * (hidden / kBlockK), kBlockK, 0, stream>>>( + reinterpret_cast(grad_output.data_ptr()), + route_scores.data_ptr(), + reinterpret_cast<__nv_fp8_e4m3*>(output.data_ptr()), + scales.data_ptr(), rows, hidden); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return {output, scales}; +} + +pybind11::dict capabilities() { + namespace py = pybind11; + py::dict result; + result["name"] = "fp8_block128_mega_moe"; + result["architecture"] = "sm103"; + result["compute_capability"] = py::make_tuple(kRequiredMajor, kRequiredMinor); + result["activation_dtype"] = "float8_e4m3fn"; + result["weight_dtype"] = "float8_e4m3fn"; + result["scale_dtype"] = "float32"; + result["activation_group_k"] = kBlockK; + result["weight_block_m"] = kBlockK; + result["weight_block_k"] = kBlockK; + result["route_score_placement"] = "post_down"; + result["execution"] = "persistent_2cta_ring_l1_l2"; + result["mma"] = "sm103_mxf8f6f4_block_scale_e4m3_e4m3"; + result["scale_values"] = "fp32_power_of_two"; + result["tile"] = py::make_tuple( + kPersistentBlockM, kPersistentBlockN, kPersistentBlockK); + result["workspace_capacity"] = "private_context_over_cp_once"; + result["fallback"] = py::none(); + result["native_symbols"] = py::make_tuple( + "sm103_fp8_block128_persistent_workspace_info", + "sm103_fp8_block128_prepare_persistent_inputs", + "sm103_fp8_block128_persistent_forward", + "sm103_fp8_block128_persistent_backward", + "sm103_fp8_block128_quantize", + "sm103_fp8_block128_dequantize", + "sm103_fp8_block128_grouped_gemm_nt", + "sm103_fp8_block128_grouped_gemm_nn", + "sm103_fp8_block128_grouped_w13_gemm_nt_canonical", + "sm103_fp8_block128_grouped_gemm_nt_expanded", + "sm103_fp8_block128_grouped_gemm_nn_expanded", + "sm103_fp8_block128_grouped_w13_gemm_nt_canonical_expanded", + "sm103_fp8_block128_swiglu_quantize", + "sm103_fp8_block128_swiglu_backward", + "sm103_fp8_block128_swiglu_backward_canonical", + "sm103_fp8_block128_route_scale_quantize", + "sm103_fp8_block128_post_down_combine", + "sm103_fp8_block128_post_down_score_grad", + "sm103_fp8_block128_route_sum", + "sm103_fp8_block128_expanded_post_down_scale", + "sm103_fp8_block128_expand_compact_routes", + "sm103_fp8_block128_collapse_expanded_routes", + "sm103_fp8_block128_expanded_post_down_score_grad", + "sm103_fp8_block128_expanded_route_scale_quantize", + "sm103_fp8_block128_grouped_bf16_wgrad", + "sm103_fp8_block128_grouped_bf16_wgrad_expanded" + ); + return result; +} + +} // namespace + +void register_apis(pybind11::module_& m) { + m.def("get_sm103_fp8_block128_capabilities", &capabilities); + m.def("sm103_fp8_block128_persistent_workspace_info", + &persistent_workspace_info, + pybind11::arg("num_ranks"), + pybind11::arg("context_tokens_per_rank")); + m.def("sm103_fp8_block128_prepare_persistent_inputs", + &prepare_persistent_inputs, + pybind11::arg("buffer"), + pybind11::arg("input"), + pybind11::arg("topk_ids"), + pybind11::arg("topk_scores"), + pybind11::arg("num_ranks"), + pybind11::arg("context_tokens_per_rank")); + m.def("sm103_fp8_block128_persistent_forward", + &persistent_forward, + pybind11::arg("buffer"), + pybind11::arg("buffer_ptrs"), + pybind11::arg("rank"), + pybind11::arg("context_tokens_per_rank"), + pybind11::arg("num_tokens"), + pybind11::arg("w13_weight"), + pybind11::arg("w13_scale"), + pybind11::arg("w2_weight"), + pybind11::arg("w2_scale")); + m.def("sm103_fp8_block128_persistent_backward", + &persistent_backward, + pybind11::arg("buffer"), + pybind11::arg("buffer_ptrs"), + pybind11::arg("rank"), + pybind11::arg("context_tokens_per_rank"), + pybind11::arg("x"), + pybind11::arg("grad_output"), + pybind11::arg("topk_scores"), + pybind11::arg("expert_counts"), + pybind11::arg("token_src_metadata"), + pybind11::arg("w13_weight"), + pybind11::arg("w13_scale"), + pybind11::arg("w2_weight"), + pybind11::arg("w2_scale")); + m.def("sm103_fp8_block128_quantize", &quantize_bf16, pybind11::arg("input")); + m.def("sm103_fp8_block128_dequantize", &dequantize_fp8, + pybind11::arg("input"), pybind11::arg("scales")); + m.def("sm103_fp8_block128_grouped_gemm_nt", &grouped_fp8_block128_gemm_nt, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("weights"), pybind11::arg("weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_grouped_gemm_nn", &grouped_fp8_block128_gemm_nn, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("weights"), pybind11::arg("weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_grouped_w13_gemm_nt_canonical", + &grouped_fp8_block128_w13_gemm_nt_canonical, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("canonical_weights"), pybind11::arg("canonical_weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_grouped_gemm_nt_expanded", + &grouped_fp8_block128_gemm_nt_expanded, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("weights"), pybind11::arg("weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_grouped_gemm_nn_expanded", + &grouped_fp8_block128_gemm_nn_expanded, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("weights"), pybind11::arg("weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_grouped_w13_gemm_nt_canonical_expanded", + &grouped_fp8_block128_w13_gemm_nt_canonical_expanded, + pybind11::arg("activations"), pybind11::arg("activation_scales"), + pybind11::arg("canonical_weights"), pybind11::arg("canonical_weight_scales"), + pybind11::arg("group_counts")); + m.def("sm103_fp8_block128_swiglu_quantize", &swiglu_quantize, + pybind11::arg("preactivation")); + m.def("sm103_fp8_block128_swiglu_backward", &swiglu_backward, + pybind11::arg("grad_output"), pybind11::arg("preactivation")); + m.def("sm103_fp8_block128_swiglu_backward_canonical", &swiglu_backward_canonical, + pybind11::arg("grad_output"), pybind11::arg("preactivation")); + m.def("sm103_fp8_block128_route_scale_quantize", &route_scale_quantize, + pybind11::arg("grad_output"), pybind11::arg("route_scores"), pybind11::arg("route_order")); + m.def("sm103_fp8_block128_post_down_combine", &post_down_combine, + pybind11::arg("route_output"), pybind11::arg("route_scores")); + m.def("sm103_fp8_block128_post_down_score_grad", &post_down_score_grad, + pybind11::arg("route_output"), pybind11::arg("grad_output"), pybind11::arg("topk")); + m.def("sm103_fp8_block128_route_sum", &route_sum, + pybind11::arg("route_grad"), pybind11::arg("num_tokens"), pybind11::arg("topk")); + m.def("sm103_fp8_block128_expanded_post_down_scale", &expanded_post_down_scale, + pybind11::arg("route_output"), pybind11::arg("route_scores")); + m.def("sm103_fp8_block128_expand_compact_routes", &expand_compact_routes, + pybind11::arg("compact"), pybind11::arg("routes"), pybind11::arg("expanded_rows")); + m.def("sm103_fp8_block128_collapse_expanded_routes", &collapse_expanded_routes, + pybind11::arg("expanded"), pybind11::arg("routes")); + m.def("sm103_fp8_block128_expanded_post_down_score_grad", + &expanded_post_down_score_grad, + pybind11::arg("route_output"), pybind11::arg("grad_output"), pybind11::arg("routes")); + m.def("sm103_fp8_block128_expanded_route_scale_quantize", + &expanded_route_scale_quantize, + pybind11::arg("grad_output"), pybind11::arg("route_scores")); + m.def("sm103_fp8_block128_grouped_bf16_wgrad", + &grouped_bf16_wgrad, + pybind11::arg("left"), pybind11::arg("right"), + pybind11::arg("padded_group_counts")); + m.def("sm103_fp8_block128_grouped_bf16_wgrad_expanded", + &grouped_bf16_wgrad_expanded, + pybind11::arg("left"), pybind11::arg("right"), pybind11::arg("psum")); +} + +} // namespace deep_gemm::sm103_fp8_block128 diff --git a/deep_gemm/__init__.py b/deep_gemm/__init__.py index 4e9c924e66..4f37f1a700 100644 --- a/deep_gemm/__init__.py +++ b/deep_gemm/__init__.py @@ -14,6 +14,10 @@ # Configs from . import _C + +# Full, build-time source identity. Unlike the local version suffix this is +# never shortened and remains available after the source checkout is absent. +__git_commit__ = _C.__git_commit__ from ._C import ( set_num_sms, get_num_sms, @@ -87,6 +91,9 @@ transform_weights_for_mega_moe, fp8_fp4_mega_moe, bf16_mega_moe, + fp8_block128_mega_moe, + get_fp8_block128_mega_moe_capabilities, + transform_glm_w13_for_fp8_block128_mega_moe, ) # Some utils diff --git a/deep_gemm/include/deep_gemm/common/types.cuh b/deep_gemm/include/deep_gemm/common/types.cuh index d9b05355ea..2531fcb245 100644 --- a/deep_gemm/include/deep_gemm/common/types.cuh +++ b/deep_gemm/include/deep_gemm/common/types.cuh @@ -56,4 +56,15 @@ enum class ActivationType { GeGLU = 1, }; +enum class RouteWeightMode { + PreDown = 0, + PostDown = 1, +}; + +enum class CombineOrderMode { + FixedTopK = 0, + DeepEP = 1, + DeepEPV1 = 2, +}; + } // namespace deep_gemm diff --git a/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh b/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh index cfad2e92cd..db65af3b1f 100644 --- a/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh @@ -1,6 +1,7 @@ #pragma once #include +#include #include #include @@ -19,24 +20,36 @@ namespace deep_gemm { +// Store one reduction value into the same shared-memory address on the peer +// CTA. The FP8-block128 L1 epilogue uses this only to finish the activation +// amax reduction across the two 64-feature halves already owned by the +// upstream 2-CTA task. +CUTLASS_DEVICE void store_cluster_float2( + float2* ptr, const uint32_t& cta_rank, const float2& value) { + const uint32_t remote_addr = cute::set_block_rank( + cute::cast_smem_ptr_to_uint(ptr), cta_rank); + asm volatile( + "st.shared::cluster.v2.f32 [%0], {%1, %2};\n" + :: "r"(remote_addr), "f"(value.x), "f"(value.y) + : "memory"); +} + 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 STORE_BLOCK_M, uint32_t SF_BLOCK_M, uint32_t SF_BLOCK_N, - uint32_t kNumRingTokens, - uint32_t kNumSFRingTokens, uint32_t kNumStages, uint32_t kNumBytesPerPull, uint32_t kNumDispatchThreads, uint32_t kNumNonEpilogueThreads, uint32_t kNumEpilogueThreads, uint32_t kNumSMs, uint32_t kNumRanks, - float kActivationClamp, + uint32_t kActivationClampBits, bool kFastMath, ActivationType kActivationType, + bool kFP8Block128Weights = false, uint32_t L1_SHAPE_N = kIntermediateHidden * 2, uint32_t L1_SHAPE_K = kHidden, uint32_t L2_SHAPE_N = kHidden, @@ -47,23 +60,37 @@ template < uint32_t kNumEpilogueWarpgroups = kNumEpilogueWarps / 4, uint32_t kNumThreads = kNumDispatchThreads + kNumNonEpilogueThreads + kNumEpilogueThreads, uint32_t kNumTokensPerWarp = 32 / kNumTopk, - uint32_t kNumExpertsPerRank = kNumExperts / kNumRanks, - uint32_t kNumRingBlocks = kNumRingTokens / BLOCK_M + uint32_t kNumExpertsPerRank = kNumExperts / kNumRanks > CUTLASS_GLOBAL __launch_bounds__(kNumThreads, 1) void sm100_fp8_fp4_mega_moe_impl(void* y, int* cumulative_local_expert_recv_stats, + layout::TokenSrcMetadata* saved_token_src_metadata, const uint32_t num_tokens, + const uint32_t num_max_tokens_per_rank, + const uint32_t num_ring_tokens, + const uint32_t num_sf_ring_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 __grid_constant__ cute::TmaDescriptor tensor_map_l1_weights_sf, const __grid_constant__ cute::TmaDescriptor tensor_map_l1_output, + const __grid_constant__ cute::TmaDescriptor tensor_map_l1_saved_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 __grid_constant__ cute::TmaDescriptor tensor_map_l2_weights_sf) { + const __grid_constant__ cute::TmaDescriptor tensor_map_l2_weights_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_l2_saved_output, + const float* l1_block128_scales = nullptr, + const float* l2_block128_scales = nullptr, + cutlass::float_e4m3_t* saved_l1_acts = nullptr, + uint32_t* saved_l1_acts_sf = nullptr, + float* saved_route_weights = nullptr, + nv_bfloat16* saved_l1_preact = nullptr, + cutlass::float_e4m3_t* saved_l2_acts = nullptr, + uint32_t* saved_l2_acts_sf = nullptr, + nv_bfloat16* saved_down_unweighted = nullptr) { #if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 1000)) or defined(__CLION_IDE__) using Barrier = cutlass::arch::ClusterTransactionBarrier; using Allocator = cute::TMEM::Allocator2Sm; @@ -73,6 +100,15 @@ sm100_fp8_fp4_mega_moe_impl(void* y, DG_STATIC_ASSERT(kNumNonEpilogueThreads == 128, "Invalid number of MMA non-epilogue threads"); DG_STATIC_ASSERT(kNumEpilogueThreads % 128 == 0, "Invalid number of MMA epilogue and combine threads"); DG_STATIC_ASSERT(kNumExperts % kNumRanks == 0, "Invalid number of experts or ranks"); + DG_DEVICE_ASSERT(num_max_tokens_per_rank > 0); + DG_DEVICE_ASSERT(num_tokens <= num_max_tokens_per_rank); + DG_DEVICE_ASSERT(num_ring_tokens > 0 and num_ring_tokens % BLOCK_M == 0); + DG_DEVICE_ASSERT(num_sf_ring_tokens > 0 and num_sf_ring_tokens % 4 == 0); + if constexpr (kFP8Block128Weights) { + DG_STATIC_ASSERT(BLOCK_M == 192 and BLOCK_N == 128 and BLOCK_K == 128, + "GLM FP8-block128 MegaMoE uses one fixed large-M tile"); + DG_DEVICE_ASSERT(l1_block128_scales != nullptr and l2_block128_scales != nullptr); + } // Thread indices const bool is_leader_cta = cute::block_rank_in_cluster() == 0; @@ -88,38 +124,41 @@ sm100_fp8_fp4_mega_moe_impl(void* y, cute::prefetch_tma_descriptor(&tensor_map_l1_weights); cute::prefetch_tma_descriptor(&tensor_map_l1_weights_sf); cute::prefetch_tma_descriptor(&tensor_map_l1_output); + cute::prefetch_tma_descriptor(&tensor_map_l1_saved_output); cute::prefetch_tma_descriptor(&tensor_map_l2_acts); cute::prefetch_tma_descriptor(&tensor_map_l2_acts_sf); cute::prefetch_tma_descriptor(&tensor_map_l2_weights); cute::prefetch_tma_descriptor(&tensor_map_l2_weights_sf); + cute::prefetch_tma_descriptor(&tensor_map_l2_saved_output); } // Workspaces const auto workspace = layout::Workspace( - sym_buffer.get_base_ptr(), kNumRanks, kNumExperts, kNumMaxTokensPerRank, kNumTopk, kNumRingTokens); + sym_buffer.get_base_ptr(), kNumRanks, kNumExperts, num_max_tokens_per_rank, kNumTopk, num_ring_tokens); + const uint32_t num_ring_blocks = num_ring_tokens / BLOCK_M; // Token and buffer layouts - 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); - constexpr auto fp8_sf_layout = layout::Data(kHidden / 32); - constexpr auto fp8_intermediate_sf_layout = layout::Data(kIntermediateHidden / 32); - 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); + const auto fp8_token_layout = layout::Data(kHidden); + const auto bf16_token_layout = layout::Data(kHidden * sizeof(nv_bfloat16)); + const auto fp8_intermediate_token_layout = layout::Data(kIntermediateHidden); + const auto fp8_sf_layout = layout::Data(kHidden / 32); + const auto fp8_intermediate_sf_layout = layout::Data(kIntermediateHidden / 32); + const auto input_topk_idx_layout = layout::Data(kNumTopk * sizeof(int64_t), false); + const auto input_topk_weights_layout = layout::Data(kNumTopk * sizeof(float), false); + const auto l1_topk_weights_layout = layout::Data(sizeof(float), false); // Registered inputs const auto input_token_buffer = layout::Buffer( - fp8_token_layout, 1, kNumMaxTokensPerRank, + fp8_token_layout, 1, num_max_tokens_per_rank, workspace.get_end_ptr()); const auto input_sf_buffer = layout::Buffer( - fp8_sf_layout, 1, kNumMaxTokensPerRank, + 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, kNumMaxTokensPerRank, + 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, kNumMaxTokensPerRank, + input_topk_weights_layout, 1, num_max_tokens_per_rank, input_topk_idx_buffer.get_end_ptr()); // SF and its buffer configs @@ -137,35 +176,38 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // L1 inputs const auto l1_token_buffer = layout::Buffer( - fp8_token_layout, 1, kNumRingTokens, + fp8_token_layout, 1, num_ring_tokens, input_topk_weights_buffer.get_end_ptr()); const auto l1_sf_buffer = layout::Buffer( - fp8_sf_layout, 1, kNumSFRingTokens, + fp8_sf_layout, 1, num_sf_ring_tokens, l1_token_buffer.get_end_ptr()); const auto l1_topk_weights_buffer = layout::Buffer( - l1_topk_weights_layout, 1, kNumRingTokens, + l1_topk_weights_layout, 1, num_ring_tokens, l1_sf_buffer.get_end_ptr()); // L2 inputs const auto l2_token_buffer = layout::Buffer( - fp8_intermediate_token_layout, 1, kNumRingTokens, + fp8_intermediate_token_layout, 1, num_ring_tokens, l1_topk_weights_buffer.get_end_ptr() ); const auto l2_sf_buffer = layout::Buffer( - fp8_intermediate_sf_layout, 1, kNumSFRingTokens, + fp8_intermediate_sf_layout, 1, num_sf_ring_tokens, l2_token_buffer.get_end_ptr() ); // Combine inputs const auto combine_token_buffer = layout::Buffer( - bf16_token_layout, kNumTopk, kNumMaxTokensPerRank, + bf16_token_layout, kNumTopk, num_max_tokens_per_rank, l2_sf_buffer.get_end_ptr() ); // Data types // NOTES: activations are FP8 (e4m3), weights are FP4 (e2m1) using a_dtype_t = cutlass::float_e4m3_t; - using b_dtype_t = cutlass::detail::float_e2m1_unpacksmem_t; + using b_dtype_t = std::conditional_t< + kFP8Block128Weights, + cutlass::float_e4m3_t, + cutlass::detail::float_e2m1_unpacksmem_t>; // MMA configs // NOTES: always swap A/B, 2-CTA MMA, and matrices are K-major @@ -210,19 +252,35 @@ sm100_fp8_fp4_mega_moe_impl(void* y, uint32_t smem_sfa[kNumStages][SF_BLOCK_M * (BLOCK_K / 128)]; uint32_t smem_sfb[kNumStages][SF_BLOCK_N * (BLOCK_K / 128)]; float2 amax_reduction[kNumEpilogueWarps][AMAX_REDUCTION_WARP_BUFFER_SIZE]; + // GLM quantizes one contiguous 128-value activation group. Each + // physical CTA owns 64 values, so its peer writes the other half's + // reduced amax here before both CTAs derive one shared exponent. + float2 l1_peer_amax + [kFP8Block128Weights ? kNumEpilogueWarpgroups : 1] + [kFP8Block128Weights ? AMAX_REDUCTION_WARP_BUFFER_SIZE : 1]; + // Canonical GLM W13 is loaded as contiguous [up; gate] 64-row + // planes. Corresponding accumulator rows therefore land in warp + // pairs (0, 2) and (1, 3). Exchange only the BF16 half each partner + // consumes; this preserves the fused epilogue without repacking W13. + uint2 l1_pair_exchange + [kFP8Block128Weights ? kNumEpilogueWarpgroups : 1] + [kFP8Block128Weights ? 4 : 1] + [kFP8Block128Weights ? 32 : 1]; Barrier dispatch_barriers[kNumDispatchWarps]; Barrier full_barriers[kNumStages]; Barrier empty_barriers[kNumStages]; Barrier tmem_full_barriers[kNumEpilogueStages]; Barrier tmem_empty_barriers[kNumEpilogueStages]; Barrier combine_barriers[kNumEpilogueWarps * 2]; + Barrier l1_scale_barriers + [kFP8Block128Weights ? kNumEpilogueWarpgroups : 1]; uint32_t tmem_ptr_in_smem; }; constexpr uint32_t kNumReusableSmemBytes = offsetof(SharedStorage, dispatch_barriers); SharedStorage &shared_storage = *reinterpret_cast(smem_buffer); // Send buffers - constexpr auto pull_layout = layout::Data(kNumBytesPerPull); + const auto pull_layout = layout::Data(kNumBytesPerPull); const auto smem_send_buffers = layout::Buffer( pull_layout, kNumDispatchWarps, 1, static_cast(shared_storage.dispatch_send_buffer)); @@ -274,6 +332,11 @@ sm100_fp8_fp4_mega_moe_impl(void* y, #pragma unroll for (uint32_t i = 0; i < kNumEpilogueWarps * 2; ++ i) shared_storage.combine_barriers[i].init(1); + if constexpr (kFP8Block128Weights) { + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueWarpgroups; ++ i) + shared_storage.l1_scale_barriers[i].init(1); + } } cutlass::arch::fence_barrier_init(); } else if (warp_idx == 3) { @@ -528,15 +591,18 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Wait for ring buffer slot to be available (previous consumer must have finished all N blocks) constexpr uint32_t kNumL1BlockNs = L1_SHAPE_N / BLOCK_N; - const auto l1_empty_count_target = (pool_block_idx / kNumRingBlocks) * kNumL1BlockNs; + const auto l1_empty_count_target = (pool_block_idx / num_ring_blocks) * kNumL1BlockNs; if (l1_empty_count_target > 0) { - const auto empty_ptr = workspace.get_l1_empty_count_ptr(pool_block_idx % kNumRingBlocks); + const auto empty_ptr = workspace.get_l1_empty_count_ptr(pool_block_idx % num_ring_blocks); while (ptx::ld_acq(empty_ptr) < l1_empty_count_target); } const auto src_base_ptr = sym_buffer.map( input_token_buffer.get_data_buffer(src_token_idx).get_base_ptr(), current_rank_in_expert_idx); - const auto dst_base_ptr = l1_token_buffer.get_data_buffer(pool_token_idx % kNumRingTokens).get_base_ptr(); + const auto dst_base_ptr = l1_token_buffer.get_data_buffer(pool_token_idx % num_ring_tokens).get_base_ptr(); + const auto saved_dst_base_ptr = saved_l1_acts != nullptr + ? saved_l1_acts + static_cast(pool_token_idx) * kHidden + : nullptr; const auto issue_and_wait_pull_store = [&](const uint32_t& i) { ptx::mbarrier_wait_and_flip_phase(pull_mbarrier, pull_mbarrier_phase); ptx::tma_store_1d( @@ -544,6 +610,13 @@ sm100_fp8_fp4_mega_moe_impl(void* y, pull_buffer.get_base_ptr(), kNumBytesPerPull ); cute::tma_store_arrive(); + if (saved_dst_base_ptr != nullptr) { + ptx::tma_store_1d( + math::advance_ptr( + saved_dst_base_ptr, i * kNumBytesPerPull), + pull_buffer.get_base_ptr(), kNumBytesPerPull); + cute::tma_store_arrive(); + } ptx::tma_store_wait<0>(); }; if (cute::elect_one_sync()) { @@ -567,15 +640,24 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 ring_block_idx = pool_block_idx % kNumRingBlocks; + const uint32_t ring_block_idx = pool_block_idx % num_ring_blocks; const uint32_t token_idx_in_block = token_idx_in_expert % BLOCK_M; const auto sf_ring_token_idx = ring_block_idx * SF_BLOCK_M + transform_sf_token_idx(token_idx_in_block); #pragma unroll for (uint32_t i = 0; i < math::constexpr_ceil_div(kNumSFUint32, 32u); ++ i) { const uint32_t j = i * 32 + lane_idx; - if (j < kNumSFUint32) - local_sf_ptr[j * kNumSFRingTokens + sf_ring_token_idx] = remote_sf_ptr[j]; + if (j < kNumSFUint32) { + const uint32_t packed_sf = remote_sf_ptr[j]; + local_sf_ptr[j * num_sf_ring_tokens + sf_ring_token_idx] = + packed_sf; + if (saved_l1_acts_sf != nullptr) { + saved_l1_acts_sf[ + static_cast(pool_token_idx) * + kNumSFUint32 + + j] = packed_sf; + } + } } __syncwarp(); @@ -585,17 +667,23 @@ sm100_fp8_fp4_mega_moe_impl(void* y, const auto weight = *sym_buffer.map( input_topk_weights_buffer.get_base_ptr() + src_token_topk_idx, current_rank_in_expert_idx); - *l1_topk_weights_buffer.get_data_buffer(pool_token_idx % kNumRingTokens).template get_base_ptr() = weight; + *l1_topk_weights_buffer.get_data_buffer(pool_token_idx % num_ring_tokens).template get_base_ptr() = weight; + if constexpr (kFP8Block128Weights) + *workspace.get_route_weight_ptr(pool_token_idx) = weight; + if (saved_route_weights != nullptr) + saved_route_weights[pool_token_idx] = weight; // Write source metadata for combine write-back (logical pool token) - *workspace.get_token_src_metadata_ptr(pool_token_idx) = + *(saved_token_src_metadata != nullptr + ? saved_token_src_metadata + pool_token_idx + : workspace.get_token_src_metadata_ptr(pool_token_idx)) = {current_rank_in_expert_idx, src_token_idx, src_topk_idx}; // Complete last chunk's store issue_and_wait_pull_store(kNumChunks - 1); const bool is_last_token = (token_idx == expert_end_idx - 1); ptx::red_add_rel( - workspace.get_l1_full_count_ptr(pool_block_idx % kNumRingBlocks), + workspace.get_l1_full_count_ptr(pool_block_idx % num_ring_blocks), is_last_token ? BLOCK_M - (token_idx_in_expert % BLOCK_M) : 1u ); } @@ -643,10 +731,10 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Clean L1 and L2 full stuffs and ring buffer counts for (uint32_t j = thread_idx; j < num_recv_m_blocks; j += kNumDispatchThreads) { - *workspace.get_l1_full_count_ptr((expert_pool_block_offset + j) % kNumRingBlocks) = 0; - *workspace.get_l1_empty_count_ptr((expert_pool_block_offset + j) % kNumRingBlocks) = 0; - *workspace.get_l2_full_count_ptr((expert_pool_block_offset + j) % kNumRingBlocks) = 0; - *workspace.get_l2_empty_count_ptr((expert_pool_block_offset + j) % kNumRingBlocks) = 0; + *workspace.get_l1_full_count_ptr((expert_pool_block_offset + j) % num_ring_blocks) = 0; + *workspace.get_l1_empty_count_ptr((expert_pool_block_offset + j) % num_ring_blocks) = 0; + *workspace.get_l2_full_count_ptr((expert_pool_block_offset + j) % num_ring_blocks) = 0; + *workspace.get_l2_empty_count_ptr((expert_pool_block_offset + j) % num_ring_blocks) = 0; } __syncwarp(); } @@ -679,16 +767,16 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Compute pool block offset for this expert const uint32_t pool_block_idx = scheduler.get_current_pool_block_offset() + m_block_idx; - const uint32_t ring_block_idx = pool_block_idx % kNumRingBlocks; + const uint32_t ring_block_idx = pool_block_idx % num_ring_blocks; // Wait the entire token arrival for linear 1 if (block_phase == sched::BlockPhase::Linear1) { const auto ptr = workspace.get_l1_full_count_ptr(ring_block_idx); - const auto num_expected_tokens = BLOCK_M * (pool_block_idx / kNumRingBlocks + 1); + const auto num_expected_tokens = BLOCK_M * (pool_block_idx / num_ring_blocks + 1); while (ptx::ld_acq(ptr) != num_expected_tokens); } else { const auto ptr = workspace.get_l2_full_count_ptr(ring_block_idx); - const auto num_expected_blocks = (L2_SHAPE_K / BLOCK_N) * 2 * (pool_block_idx / kNumRingBlocks + 1); + const auto num_expected_blocks = (L2_SHAPE_K / BLOCK_N) * 2 * (pool_block_idx / num_ring_blocks + 1); while (ptx::ld_acq(ptr) != num_expected_blocks); } @@ -750,7 +838,91 @@ sm100_fp8_fp4_mega_moe_impl(void* y, uint32_t sfb_k_idx = local_expert_idx * shape_sfb_k + k_block_idx * (BLOCK_K / 128); // TMA copy weights with SF - if (cute::elect_one_sync()) { + if constexpr (kFP8Block128Weights) { + // GLM stores canonical [gate, up] as [2E, H, D]. Load the + // two 64-row canonical planes directly into contiguous + // [up; gate] shared-memory halves. Two TMA transactions + // replace the invalid 16-way 8-row fan-out; only the + // 16-KiB tile is materialized and W13 is never repacked. + if (cute::elect_one_sync()) { + if (block_phase == sched::BlockPhase::Linear1) { + constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; + const uint32_t logical_row = + n_block_idx * kLogicalRowsPerBlock; + const uint32_t up_row = + (local_expert_idx * 2 + 1) * kIntermediateHidden + logical_row; + const uint32_t gate_row = + (local_expert_idx * 2) * kIntermediateHidden + logical_row; + tma::copy( + tensor_map_b_ptr, + &shared_storage.full_barriers[stage_idx], + shared_storage.smem_b[stage_idx], + k_idx, + up_row, + 2); + tma::copy( + tensor_map_b_ptr, + &shared_storage.full_barriers[stage_idx], + shared_storage.smem_b[stage_idx] + + kLogicalRowsPerBlock * BLOCK_K, + k_idx, + gate_row, + 2); + } else { + tma::copy( + tensor_map_b_ptr, + &shared_storage.full_barriers[stage_idx], + shared_storage.smem_b[stage_idx], + k_idx, + local_expert_idx * kHidden + n_block_idx * BLOCK_N, + 2); + } + } + + // Canonical FP32 scales are exact powers of two after the + // owning q/s refresh. Feed their exponent directly to the + // hardware K/32 scale lanes. A K/128 tile therefore + // becomes four identical UE8M0 bytes, with no software + // accumulator scaling and no persistent scale copy. + #pragma unroll + for (uint32_t row = lane_idx; row < BLOCK_N; row += 32) { + float scale; + if (block_phase == sched::BlockPhase::Linear1) { + constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; + const uint32_t logical_row = + n_block_idx * kLogicalRowsPerBlock + + row % kLogicalRowsPerBlock; + const uint32_t canonical_expert = + local_expert_idx * 2 + (row < kLogicalRowsPerBlock ? 1u : 0u); + const uint32_t scale_idx = + (canonical_expert * (kIntermediateHidden / 128) + logical_row / 128) * + (kHidden / 128) + + k_block_idx; + scale = __ldg(l1_block128_scales + scale_idx); + } else { + const uint32_t scale_idx = + (local_expert_idx * (kHidden / 128) + n_block_idx) * + (kIntermediateHidden / 128) + + k_block_idx; + scale = __ldg(l2_block128_scales + scale_idx); + } + const uint32_t exponent = __float_as_uint(scale) >> 23; + shared_storage.smem_sfb[stage_idx][row] = exponent * 0x01010101u; + } + __syncwarp(); + if (cute::elect_one_sync()) { + if (is_leader_cta) + // Both CTAs issue the 2-SM FP8 TMA load and their + // transaction bytes accumulate on CTA 0's + // multicast barrier. Unlike packed FP4, one CTA's + // FP8 contribution already occupies the complete + // shared-memory tile, so account for both CTAs. + shared_storage.full_barriers[stage_idx].arrive_and_expect_tx( + sizeof(SharedStorage::smem_b[0]) * 2); + else + shared_storage.full_barriers[stage_idx].arrive(0u); + } + } else if (cute::elect_one_sync()) { tma::copy( tensor_map_b_ptr, &shared_storage.full_barriers[stage_idx], shared_storage.smem_b[stage_idx], k_idx, n_idx, 2); tma::copy( @@ -918,6 +1090,7 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Persistently schedule over blocks uint32_t current_iter_idx = 0; + uint32_t l1_scale_phase = 0; scheduler.for_each_block([&](const sched::BlockPhase& block_phase, const uint32_t& local_expert_idx, const uint32_t& num_k_blocks, @@ -932,7 +1105,7 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // NOTES: use shuffle here to let NVCC know warp divergence won't happen const uint32_t valid_m = ptx::exchange(scheduler.template get_valid_m(), 0); const uint32_t pool_block_idx = scheduler.get_current_pool_block_offset() + m_block_idx; - const uint32_t ring_block_idx = pool_block_idx % kNumRingBlocks; + const uint32_t ring_block_idx = pool_block_idx % num_ring_blocks; const uint32_t ring_m_idx = ring_block_idx * BLOCK_M; // Ring-buffer offset for reusable data buffers const uint32_t pool_m_idx = pool_block_idx * BLOCK_M; // Full-pool offset for non-ring metadata uint32_t n_idx = n_block_idx * BLOCK_N; @@ -940,12 +1113,13 @@ sm100_fp8_fp4_mega_moe_impl(void* y, if (block_phase == sched::BlockPhase::Linear1) { // Wait L2 block empty const auto l2_empty_ptr = workspace.get_l2_empty_count_ptr(ring_block_idx); - const auto num_expected_blocks = (L2_SHAPE_N / BLOCK_N) * (pool_block_idx / kNumRingBlocks); + const auto num_expected_blocks = (L2_SHAPE_N / BLOCK_N) * (pool_block_idx / num_ring_blocks); while (ptx::ld_acq(l2_empty_ptr) != num_expected_blocks); - // Unified L1 epilogue: gated activation (SwiGLU/GeGLU) in-place using - // granularity 8 interleaved weights. - // With `SM100_TMEM_LOAD_16dp256b1x`, gate/up pairs are: + // Unified L1 epilogue: gated activation (SwiGLU/GeGLU). + // The FP8-block128 path keeps canonical W13 in contiguous + // [up; gate] halves. Accumulator warp pairs exchange the + // BF16 half needed to form each logical feature in place. float stored_cached_weight = 0; #pragma unroll @@ -966,17 +1140,28 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Load weights from global into register cache per 32 tokens DG_STATIC_ASSERT(32 % ATOM_M == 0, "Invalid block size"); - if ((j * ATOM_M) % 32 == 0 and (WG_BLOCK_M % 32 == 0 or j * ATOM_M + lane_idx < WG_BLOCK_M)) { - stored_cached_weight = *l1_topk_weights_buffer - .get_data_buffer(ring_m_idx + epilogue_wg_idx * WG_BLOCK_M + j * ATOM_M + lane_idx) - .template get_base_ptr(); + if constexpr (!kFP8Block128Weights) { + if ((j * ATOM_M) % 32 == 0 and (WG_BLOCK_M % 32 == 0 or j * ATOM_M + lane_idx < WG_BLOCK_M)) { + stored_cached_weight = *l1_topk_weights_buffer + .get_data_buffer(ring_m_idx + epilogue_wg_idx * WG_BLOCK_M + j * ATOM_M + lane_idx) + .template get_base_ptr(); + } } - // Load weights from register cache - const float2 weights = { - ptx::exchange(stored_cached_weight, (j * ATOM_M) % 32 + (lane_idx % 4) * 2 + 0), - ptx::exchange(stored_cached_weight, (j * ATOM_M) % 32 + (lane_idx % 4) * 2 + 1) - }; + // Upstream MegaMoE folds each route score into L1 + // before the low-precision intermediate. GLM's + // canonical POST_DOWN contract instead applies the + // FP32 score after W2 has produced BF16. Retain the + // upstream exchange only for the original FP4 path; + // the FP8-block128 L2 epilogue reads the score for its + // final BF16 row directly. + float2 weights = {1.0f, 1.0f}; + if constexpr (!kFP8Block128Weights) { + weights = { + ptx::exchange(stored_cached_weight, (j * ATOM_M) % 32 + (lane_idx % 4) * 2 + 0), + ptx::exchange(stored_cached_weight, (j * ATOM_M) % 32 + (lane_idx % 4) * 2 + 1) + }; + } // Load from TMEM uint2 raw_values[4]; @@ -993,18 +1178,128 @@ sm100_fp8_fp4_mega_moe_impl(void* y, shared_storage.tmem_empty_barriers[accum_stage_idx].arrive(0u); } - // Apply gated activation: act(gate) * up (SwiGLU or GeGLU) + // Materialize logical gate/up pairs. The upstream FP4 + // layout is already interleaved at granularity 8. The + // canonical FP8 layout is contiguous [up64; gate64], + // so warp pairs (0, 2) and (1, 3) exchange only the + // BF16 half their partner consumes. auto fp32_values = reinterpret_cast(raw_values); + nv_bfloat162 bf16_gate_values[2]; + nv_bfloat162 bf16_up_values[2]; + if constexpr (kFP8Block128Weights) { + const bool owns_up = warp_idx_in_wg < 2; + const uint32_t own_base = owns_up ? 0u : 2u; + const uint32_t outbound_base = owns_up ? 2u : 0u; + const auto outbound_0 = __float22bfloat162_rn( + fp32_values[outbound_base]); + const auto outbound_1 = __float22bfloat162_rn( + fp32_values[outbound_base + 1]); + const uint2 outbound = { + *reinterpret_cast(&outbound_0), + *reinterpret_cast(&outbound_1) + }; + auto outbound_ptr = reinterpret_cast( + &shared_storage.l1_pair_exchange + [epilogue_wg_idx][warp_idx_in_wg][lane_idx]); + ptx::st_shared(outbound_ptr, outbound.x); + ptx::st_shared(outbound_ptr + 1, outbound.y); + ptx::sync_aligned( + 128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + + auto inbound_ptr = reinterpret_cast( + &shared_storage.l1_pair_exchange + [epilogue_wg_idx][warp_idx_in_wg ^ 2u][lane_idx]); + const uint2 inbound = { + ptx::ld_shared(inbound_ptr), + ptx::ld_shared(inbound_ptr + 1) + }; + // No warp may overwrite the single exchange stage + // for the next atom until every partner has read it. + ptx::sync_aligned( + 128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + const auto inbound_bf16 = + reinterpret_cast(&inbound); + + #pragma unroll + for (uint32_t k = 0; k < 2; ++ k) { + const auto own = __float22bfloat162_rn( + fp32_values[own_base + k]); + bf16_gate_values[k] = owns_up ? inbound_bf16[k] : own; + bf16_up_values[k] = owns_up ? own : inbound_bf16[k]; + } + } else { + #pragma unroll + for (uint32_t k = 0; k < 2; ++ k) { + bf16_gate_values[k] = __float22bfloat162_rn( + fp32_values[k * 2]); + bf16_up_values[k] = __float22bfloat162_rn( + fp32_values[k * 2 + 1]); + } + } + + // Apply gated activation: act(gate) * up (SwiGLU or GeGLU) #pragma unroll for (uint32_t k = 0; k < 2; ++ k) { - auto bf16_gate = __float22bfloat162_rn(fp32_values[k * 2 + 0]); - auto bf16_up = __float22bfloat162_rn(fp32_values[k * 2 + 1]); + auto bf16_gate = bf16_gate_values[k]; + auto bf16_up = bf16_up_values[k]; + + if constexpr (kFP8Block128Weights) { + if (saved_l1_preact != nullptr) { + // Preserve upstream's training lifetime: + // save the exact BF16 accumulator rounding + // before clamp/activation. Canonical W13 is + // loaded as [up64; gate64], while the fused + // output chunk order is 0,2,1,3 across the + // four epilogue warps. + const uint32_t logical_chunk = + (warp_idx_in_wg % 2) * 2 + + warp_idx_in_wg / 2; + const uint32_t hidden_col = + n_block_idx * (BLOCK_N / 2) + + logical_chunk * 16 + + (lane_idx / 4) * 2 + k; + const uint32_t row_base = + pool_m_idx + + epilogue_wg_idx * WG_BLOCK_M + + s * STORE_BLOCK_M + + i * ATOM_M + + (lane_idx % 4) * 2; + const uint32_t gate_bits = + *reinterpret_cast( + &bf16_gate); + const uint32_t up_bits = + *reinterpret_cast( + &bf16_up); + #pragma unroll + for (uint32_t r = 0; r < 2; ++r) { + const uint32_t pool_row = row_base + r; + if (pool_row < pool_m_idx + valid_m) { + auto* saved_bits = + reinterpret_cast( + saved_l1_preact + + static_cast( + pool_row) * + (2 * kIntermediateHidden)); + saved_bits[hidden_col] = + static_cast( + gate_bits >> (r * 16)); + saved_bits[ + kIntermediateHidden + hidden_col] = + static_cast( + up_bits >> (r * 16)); + } + } + } + } // Clamp - if constexpr (kActivationClamp != cute::numeric_limits::infinity()) { - bf16_gate = __hmin2(bf16_gate, {kActivationClamp, kActivationClamp}); - bf16_up = __hmax2(bf16_up, {-kActivationClamp, -kActivationClamp}); - bf16_up = __hmin2(bf16_up, {kActivationClamp, kActivationClamp}); + if constexpr (kActivationClampBits != 0x7f800000u) { + const float activation_clamp = __uint_as_float(kActivationClampBits); + const auto clamp = __floats2bfloat162_rn(activation_clamp, activation_clamp); + const auto neg_clamp = __floats2bfloat162_rn(-activation_clamp, -activation_clamp); + bf16_gate = __hmin2(bf16_gate, clamp); + bf16_up = __hmax2(bf16_up, neg_clamp); + bf16_up = __hmin2(bf16_up, clamp); } const auto gate = __bfloat1622float2(bf16_gate); @@ -1037,7 +1332,12 @@ sm100_fp8_fp4_mega_moe_impl(void* y, } else { activated = {gate.x / denom.x, gate.y / denom.y}; } - activation_values[i][k] = __fmul2_rn(__fmul2_rn(activated, up), weights); + const auto unweighted = __fmul2_rn(activated, up); + if constexpr (kFP8Block128Weights) + activation_values[i][k] = unweighted; + else + activation_values[i][k] = + __fmul2_rn(unweighted, weights); } // Amax reduction (thread-level) @@ -1066,14 +1366,92 @@ sm100_fp8_fp4_mega_moe_impl(void* y, ptx::tma_store_wait(); ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + // The upstream task already pairs adjacent N blocks in one + // cluster. For GLM, those CTAs contain adjacent 64-value + // halves of a single block128 activation group. Reduce the + // four local warp fragments, exchange one value per token + // pair with the peer CTA, and derive exactly one UE8M0 + // exponent for all 128 values. This remains inside the + // fused persistent epilogue; no preprocessing/composed EP + // stage is introduced. + if constexpr (kFP8Block128Weights) { + float2 local_amax[kNumAtomsPerStore]; + #pragma unroll + for (uint32_t i = 0; i < kNumAtomsPerStore; ++ i) { + local_amax[i] = {0.0f, 0.0f}; + #pragma unroll + for (uint32_t local_warp = 0; + local_warp < 4; ++ local_warp) { + const float2 value = + shared_storage.amax_reduction + [epilogue_wg_idx * 4 + local_warp] + [i * (ATOM_M / 2) + lane_idx % 4]; + local_amax[i].x = cute::max( + local_amax[i].x, value.x); + local_amax[i].y = cute::max( + local_amax[i].y, value.y); + } + } + + if (warp_idx_in_wg == 0 && lane_idx < 4) { + #pragma unroll + for (uint32_t i = 0; + i < kNumAtomsPerStore; ++ i) { + store_cluster_float2( + &shared_storage.l1_peer_amax + [epilogue_wg_idx] + [i * (ATOM_M / 2) + lane_idx], + cute::block_rank_in_cluster() ^ 1u, + local_amax[i]); + } + } + // Four lanes issue the peer DSM stores, whereas one + // elected lane publishes the remote mbarrier arrival. + // Give every writer cluster-scope release semantics + // before that arrival; a warp execution barrier alone + // does not order another lane's remote shared writes. + if (warp_idx_in_wg == 0) + __threadfence_cluster(); + __syncwarp(); + if (warp_idx_in_wg == 0 && cute::elect_one_sync()) { + shared_storage.l1_scale_barriers + [epilogue_wg_idx].arrive( + cute::block_rank_in_cluster() ^ 1u); + } + shared_storage.l1_scale_barriers + [epilogue_wg_idx].wait(l1_scale_phase); + + #pragma unroll + for (uint32_t i = 0; + i < kNumAtomsPerStore; ++ i) { + const float2 peer_amax = ptx::ld_shared( + &shared_storage.l1_peer_amax + [epilogue_wg_idx] + [i * (ATOM_M / 2) + lane_idx % 4]); + amax_values[i].x = cute::max( + local_amax[i].x, peer_amax.x); + amax_values[i].y = cute::max( + local_amax[i].y, peer_amax.y); + } + l1_scale_phase ^= 1u; + } + // Cast to FP8 E4M3 and store into shared memory #pragma unroll for (uint32_t i = 0; i < kNumAtomsPerStore; ++ i) { // Reduce amax (warp-pair-level) - const float2 wp_amax = - shared_storage.amax_reduction[epilogue_warp_idx ^ 1][i * (ATOM_M / 2) + lane_idx % 4]; - amax_values[i].x = cute::max(amax_values[i].x, wp_amax.x); - amax_values[i].y = cute::max(amax_values[i].y, wp_amax.y); + if constexpr (!kFP8Block128Weights) { + const uint32_t amax_partner = + epilogue_warp_idx ^ 1u; + const float2 wp_amax = + shared_storage.amax_reduction + [amax_partner] + [i * (ATOM_M / 2) + lane_idx % 4]; + amax_values[i].x = cute::max( + amax_values[i].x, wp_amax.x); + amax_values[i].y = cute::max( + amax_values[i].y, wp_amax.y); + } // Calculate SF float2 sf, sf_inv; @@ -1086,7 +1464,12 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // STSM uint32_t row = lane_idx; - uint32_t col = warp_idx_in_wg; + // Contiguous [up; gate] maps logical 16-feature output + // chunks as warp 0, 2, 1, 3. The FP4 path retains its + // original interleaved warp order. + uint32_t col = kFP8Block128Weights + ? (warp_idx_in_wg % 2) * 2 + warp_idx_in_wg / 2 + : warp_idx_in_wg; const auto smem_ptr = reinterpret_cast(shared_storage.smem_d.l1[epilogue_wg_idx][tma_stage_idx]) + i * ATOM_M * L1_OUT_BLOCK_N + row * L1_OUT_BLOCK_N @@ -1097,10 +1480,16 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Store SF to `l2_sf_buffer` as UE8M0 (MN-major layout) // Only one warp per pair writes (both hold the same SF after cross-warp reduce) // Each lane < 4 holds SF for 2 rows (sf.x and sf.y) - if (warp_idx_in_wg % 2 == 0 and lane_idx < 4) { - const uint32_t k_idx = n_block_idx * 2 + warp_idx_in_wg / 2; + const bool writes_sf = kFP8Block128Weights + ? warp_idx_in_wg < 2 + : warp_idx_in_wg % 2 == 0; + if (writes_sf and lane_idx < 4) { + const uint32_t sf_group = kFP8Block128Weights + ? warp_idx_in_wg + : warp_idx_in_wg / 2; + const uint32_t k_idx = n_block_idx * 2 + sf_group; const uint32_t k_uint_idx = k_idx / 4, byte_idx = k_idx % 4; - const uint32_t mn_stride = kNumSFRingTokens * sizeof(uint32_t); + const uint32_t mn_stride = num_sf_ring_tokens * sizeof(uint32_t); const auto sf_base_ptr = l2_sf_buffer.get_base_ptr(); // NOTES: consecutive tokens (t, t + 1) are in the same 32-group, so `sf_idx` differs by 4 // NOTES: originally there was: @@ -1119,6 +1508,43 @@ sm100_fp8_fp4_mega_moe_impl(void* y, (*reinterpret_cast(&sf.x) >> 23); sf_base_ptr[sf_addr + 4 * static_cast(sizeof(uint32_t))] = (*reinterpret_cast(&sf.y) >> 23); + + if constexpr (kFP8Block128Weights) { + if (saved_l2_acts_sf != nullptr && + (n_block_idx & 1u) == 0u && + warp_idx_in_wg == 0) { + const uint32_t token_base_idx = + epilogue_wg_idx * WG_BLOCK_M + + s * STORE_BLOCK_M + + i * ATOM_M; + const uint32_t row0 = + pool_m_idx + token_base_idx + + lane_idx * 2; + const uint32_t row1 = row0 + 1; + const uint32_t scale_block = + n_block_idx / 2; + const uint32_t num_scale_blocks = + kIntermediateHidden / 128; + const uint32_t packed0 = + (*reinterpret_cast( + &sf.x) >> 23) * 0x01010101u; + const uint32_t packed1 = + (*reinterpret_cast( + &sf.y) >> 23) * 0x01010101u; + if (row0 < pool_m_idx + valid_m) { + saved_l2_acts_sf[ + static_cast(row0) * + num_scale_blocks + + scale_block] = packed0; + } + if (row1 < pool_m_idx + valid_m) { + saved_l2_acts_sf[ + static_cast(row1) * + num_scale_blocks + + scale_block] = packed1; + } + } + } } __syncwarp(); } @@ -1134,6 +1560,17 @@ sm100_fp8_fp4_mega_moe_impl(void* y, out_n_idx, ring_m_idx + epilogue_wg_idx * WG_BLOCK_M + s * STORE_BLOCK_M); cute::tma_store_arrive(); + if (saved_l2_acts != nullptr) { + cute::SM90_TMA_STORE_2D::copy( + &tensor_map_l1_saved_output, + shared_storage.smem_d + .l1[epilogue_wg_idx][tma_stage_idx], + out_n_idx, + pool_m_idx + + epilogue_wg_idx * WG_BLOCK_M + + s * STORE_BLOCK_M); + cute::tma_store_arrive(); + } } __syncwarp(); } @@ -1217,6 +1654,42 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // Wait shared memory ready ptx::sync_aligned(128, kEpilogueWGBarrierStartIdx + epilogue_wg_idx); + if constexpr (kFP8Block128Weights) { + if (saved_down_unweighted != nullptr) { + const uint32_t saved_store_row = + pool_m_idx + + epilogue_wg_idx * WG_BLOCK_M + + s * STORE_BLOCK_M; + if (warp_idx_in_wg == 0 && + cute::elect_one_sync()) { + cute::tma_store_fence(); + #pragma unroll + for (uint32_t atom = 0; + atom < + BLOCK_N * sizeof(nv_bfloat16) / + kSwizzleCDMode; + ++atom) { + cute::SM90_TMA_STORE_2D::copy( + &tensor_map_l2_saved_output, + shared_storage.smem_d + .l2[epilogue_wg_idx] + + atom * STORE_BLOCK_M * + (kSwizzleCDMode / + sizeof(nv_bfloat16)), + n_idx + + atom * + (kSwizzleCDMode / + sizeof(nv_bfloat16)), + saved_store_row); + cute::tma_store_arrive(); + } + } + if (warp_idx_in_wg == 0) + cute::tma_store_wait<0>(); + __syncwarp(); + } + } + // Write into remote buffers // Each warp writes 2 rows (lane_idx/16 splits the warp into two halves, one per row) const uint32_t row_in_atom = (warp_idx_in_wg * 2 + lane_idx / 16) % ATOM_M; @@ -1231,7 +1704,9 @@ sm100_fp8_fp4_mega_moe_impl(void* y, if (m_idx_in_block >= valid_m) break; - const auto src_metadata = *workspace.get_token_src_metadata_ptr(pool_m_idx + m_idx_in_block); + const auto src_metadata = *(saved_token_src_metadata != nullptr + ? saved_token_src_metadata + pool_m_idx + m_idx_in_block + : workspace.get_token_src_metadata_ptr(pool_m_idx + m_idx_in_block)); const uint32_t dst_rank_idx = src_metadata.rank_idx; const uint32_t dst_token_idx = src_metadata.token_idx; const uint32_t dst_topk_idx = src_metadata.topk_idx; @@ -1241,7 +1716,37 @@ sm100_fp8_fp4_mega_moe_impl(void* y, (lane_idx % 16 / 8) * STORE_BLOCK_M * kSwizzleCDMode + row_in_store * kSwizzleCDMode + (bank_group_idx ^ row_in_atom) * kNumBankGroupBytes; - const auto packed = ptx::ld_shared(reinterpret_cast(smem_ptr)); + auto packed = ptx::ld_shared(reinterpret_cast(smem_ptr)); + + if constexpr (kFP8Block128Weights) { + // Match GLM's frozen post-down semantics exactly: + // W2's FP32 accumulator was rounded to BF16 when + // written to shared memory above, then converted + // back to FP32 and multiplied by the FP32 route + // score before the final BF16 rounding. Moving + // this multiply into L1 is not equivalent because + // it changes E4M3 requantization before W2. + // Follow upstream persistent POST_DOWN lifetime: + // L1 ring slots may already be reused by a later + // wave, whereas this full-pool entry remains owned + // by the logical route until remote combine. + const float route_weight = + *workspace.get_route_weight_ptr( + pool_m_idx + m_idx_in_block); + auto* packed_bf16 = + reinterpret_cast(&packed); + #pragma unroll + for (uint32_t pair = 0; + pair < sizeof(float4) / sizeof(nv_bfloat162); + ++pair) { + const auto value = + __bfloat1622float2(packed_bf16[pair]); + packed_bf16[pair] = __float22bfloat162_rn( + __fmul2_rn( + value, + {route_weight, route_weight})); + } + } // Write into remote const auto dst_token = combine_token_buffer.get_rank_buffer(dst_topk_idx) diff --git a/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh b/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh new file mode 100644 index 0000000000..e9bf06d969 --- /dev/null +++ b/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh @@ -0,0 +1,7371 @@ +#pragma once + +#include + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace deep_gemm { + +template < + uint32_t kHidden, uint32_t kNumExperts, uint32_t BLOCK_M, + uint32_t kNumSMs, uint32_t kNumThreads, + CombineOrderMode kCombineOrderMode> +__device__ __forceinline__ void +bf16_mega_moe_reduce_post_down_route( + const int* expert_counts, + const cutlass::bfloat16_t* grad_y_unweighted, + const cutlass::bfloat16_t* down_unweighted, + float* grad_route_output, + uint8_t* scratch) { + constexpr uint32_t kTritonRouteBlockH = [] { + uint32_t value = 1; + while (value < kHidden && value < 8192) + value <<= 1; + return value; + }(); + constexpr uint32_t kTritonRouteNumWarps = [] { + uint32_t value = kTritonRouteBlockH / 256; + value = value < 4 ? 4 : value; + return value > 32 ? 32 : value; + }(); + constexpr uint32_t kTritonRouteThreads = + kTritonRouteNumWarps * 32; + constexpr uint32_t kTritonRouteValuesPerThread = + kTritonRouteBlockH / kTritonRouteThreads; + DG_STATIC_ASSERT( + kTritonRouteValuesPerThread == 2 || + kTritonRouteValuesPerThread == 4 || + kTritonRouteValuesPerThread == 8, + "Unsupported Triton route reduction width"); + constexpr uint32_t kRouteInputPow2 = [] { + uint32_t value = 1; + constexpr uint32_t vectorized_columns = kHidden / 4; + while (value < 512 && + (value << 1) <= vectorized_columns) + value <<= 1; + return value; + }(); + + auto* route_lane_sums = reinterpret_cast(scratch); + auto* route_control = reinterpret_cast(scratch); + if constexpr ( + kCombineOrderMode != CombineOrderMode::FixedTopK) { + // A sub-CTA named barrier is not safe while the persistent kernel's + // earlier role-specific register/barrier phases are still live. + // Assign one route to the CTA instead. The first Triton-sized thread + // group keeps exactly the same lane-to-column map and butterfly tree; + // the remaining threads only participate in CTA phase barriers. + auto* route_warp_arrivals = + reinterpret_cast( + scratch + + kNumThreads * sizeof(float)); + if (threadIdx.x == 0) + *route_warp_arrivals = 0; + __syncthreads(); + uint32_t route_pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + for (uint32_t token_idx = blockIdx.x; + token_idx < num_tokens; + token_idx += kNumSMs) { + const uint32_t pool_row = + route_pool_block_offset * BLOCK_M + + token_idx; + const uint32_t route_lane = threadIdx.x; + float grad_route = 0.0f; + if (route_lane < kTritonRouteThreads) { + float grad_y[ + kTritonRouteValuesPerThread]; + float down[ + kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + route_lane + + i * kTritonRouteThreads; + grad_y[i] = + col < kHidden + ? static_cast( + grad_y_unweighted[ + static_cast< + uint64_t>( + pool_row) * + kHidden + + col]) + : 0.0f; + down[i] = + col < kHidden + ? static_cast( + down_unweighted[ + static_cast< + uint64_t>( + pool_row) * + kHidden + + col]) + : 0.0f; + } + if constexpr ( + kTritonRouteValuesPerThread == 2) { + grad_route = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[1], down[1])); + } else if constexpr ( + kTritonRouteValuesPerThread == 4) { + const float even = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float odd = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + grad_route = + __fadd_rn(even, odd); + } else { + const float pair_02 = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float pair_13 = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + const float pair_46 = __fmaf_rn( + grad_y[4], down[4], + __fmul_rn( + grad_y[6], down[6])); + const float pair_57 = __fmaf_rn( + grad_y[5], down[5], + __fmul_rn( + grad_y[7], down[7])); + grad_route = __fadd_rn( + __fadd_rn( + pair_02, pair_46), + __fadd_rn( + pair_13, pair_57)); + } + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, offset)); + } + const uint32_t lane_in_warp = + route_lane & 31; + if (lane_in_warp == 0) { + route_lane_sums[ + route_lane / 32] = + grad_route; + __threadfence_block(); + atomicAdd( + route_warp_arrivals, 1u); + } + } + if (threadIdx.x < 32) { + if (threadIdx.x == 0) { + while (atomicAdd( + route_warp_arrivals, + 0u) != + kTritonRouteNumWarps) { + } + } + __syncwarp(); + grad_route = route_lane_sums[ + threadIdx.x & + (kTritonRouteNumWarps - 1)]; + #pragma unroll + for (uint32_t offset = + kTritonRouteNumWarps / 2; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, offset)); + } + } + if (threadIdx.x == 0) + grad_route_output[pool_row] = + grad_route; + __syncthreads(); + if (threadIdx.x == 0) + *route_warp_arrivals = 0; + __syncthreads(); + } + route_pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } + return; + } + + if (threadIdx.x == 0) { + uint32_t total_route_rows = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + total_route_rows += static_cast( + __ldg(expert_counts + expert_idx)); + } + route_control[0] = total_route_rows; + } + __syncthreads(); + const uint32_t total_route_rows = route_control[0]; + const uint32_t route_output_pow2 = + total_route_rows > 0 + ? 1u << (31 - __clz(total_route_rows)) + : 1u; + constexpr uint32_t kInitialRouteGroupThreads = + cute::min(kRouteInputPow2, 32u); + const uint32_t route_block_height = + cute::min( + route_output_pow2, + 512u / kInitialRouteGroupThreads); + const uint32_t route_group_threads = + kCombineOrderMode != CombineOrderMode::FixedTopK + ? kTritonRouteThreads + : cute::min( + kRouteInputPow2, + 512u / route_block_height); + const uint32_t num_route_groups_per_cta = + kNumThreads / route_group_threads; + const uint32_t route_group_idx = + threadIdx.x / route_group_threads; + const uint32_t route_group_lane_idx = + threadIdx.x & (route_group_threads - 1); + const uint32_t global_route_group = + blockIdx.x * num_route_groups_per_cta + + route_group_idx; + const uint32_t num_route_groups = + kNumSMs * num_route_groups_per_cta; + uint32_t route_pool_block_offset = 0; + + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = static_cast( + __ldg(expert_counts + expert_idx)); + for (uint32_t token_idx = global_route_group; + token_idx < num_tokens; + token_idx += num_route_groups) { + const uint32_t pool_row = + route_pool_block_offset * BLOCK_M + token_idx; + float grad_route = 0.0f; + if constexpr ( + kCombineOrderMode != CombineOrderMode::FixedTopK) { + float grad_y[kTritonRouteValuesPerThread]; + float down[kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t i = 0; + i < kTritonRouteValuesPerThread; ++i) { + const uint32_t col = + route_group_lane_idx + + i * kTritonRouteThreads; + grad_y[i] = + col < kHidden + ? static_cast( + grad_y_unweighted[ + static_cast(pool_row) * + kHidden + + col]) + : 0.0f; + down[i] = + col < kHidden + ? static_cast( + down_unweighted[ + static_cast(pool_row) * + kHidden + + col]) + : 0.0f; + } + + if constexpr (kTritonRouteValuesPerThread == 2) { + grad_route = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn(grad_y[1], down[1])); + } else if constexpr ( + kTritonRouteValuesPerThread == 4) { + const float even = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn(grad_y[2], down[2])); + const float odd = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn(grad_y[3], down[3])); + grad_route = __fadd_rn(even, odd); + } else { + const float pair_02 = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn(grad_y[2], down[2])); + const float pair_13 = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn(grad_y[3], down[3])); + const float pair_46 = __fmaf_rn( + grad_y[4], down[4], + __fmul_rn(grad_y[6], down[6])); + const float pair_57 = __fmaf_rn( + grad_y[5], down[5], + __fmul_rn(grad_y[7], down[7])); + grad_route = __fadd_rn( + __fadd_rn(pair_02, pair_46), + __fadd_rn(pair_13, pair_57)); + } + + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, grad_route, offset)); + } + const uint32_t warp_in_group = + route_group_lane_idx / 32; + const uint32_t lane_in_warp = + route_group_lane_idx & 31; + if (lane_in_warp == 0) { + route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + warp_in_group] = grad_route; + } + ptx::sync_aligned( + kTritonRouteThreads, route_group_idx); + if (warp_in_group == 0) { + grad_route = route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + (lane_in_warp & + (kTritonRouteNumWarps - 1))]; + #pragma unroll + for (uint32_t offset = + kTritonRouteNumWarps / 2; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, offset)); + } + } + } else { + float lane_sums[4] = { + 0.0f, 0.0f, 0.0f, 0.0f}; + for (uint32_t col_base = + route_group_lane_idx * 4; + col_base < kHidden; + col_base += route_group_threads * 4) { + #pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const uint32_t col = col_base + i; + const float grad_y = static_cast( + grad_y_unweighted[ + static_cast(pool_row) * + kHidden + + col]); + const float down = static_cast( + down_unweighted[ + static_cast(pool_row) * + kHidden + + col]); + lane_sums[i] = __fadd_rn( + lane_sums[i], + __fmul_rn(grad_y, down)); + } + } + grad_route = __fadd_rn( + __fadd_rn(lane_sums[0], lane_sums[1]), + lane_sums[2]); + grad_route = + __fadd_rn(grad_route, lane_sums[3]); + route_lane_sums[threadIdx.x] = grad_route; + if (route_group_threads > 32) { + for (uint32_t offset = + route_group_threads / 2; + offset >= 32; offset >>= 1) { + ptx::sync_aligned( + route_group_threads, + route_group_idx); + if (route_group_lane_idx < offset) { + grad_route = __fadd_rn( + grad_route, + route_lane_sums[ + threadIdx.x + offset]); + route_lane_sums[threadIdx.x] = + grad_route; + } + } + } + if (route_group_lane_idx < 32) { + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_down_sync( + 0xffffffff, + grad_route, offset)); + } + } + } + if (route_group_lane_idx == 0) + grad_route_output[pool_row] = grad_route; + if (route_group_threads > 32) { + ptx::sync_aligned( + route_group_threads, route_group_idx); + } else { + __syncwarp(); + } + } + route_pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } +} + +template < + uint32_t kHidden, uint32_t kNumExperts, uint32_t BLOCK_M, + uint32_t kNumSMs, uint32_t kNumRanks, + CombineOrderMode kCombineOrderMode, + bool kDoReverseDispatch = true, + bool kComputeRouteDot = true, + bool kWriteWeighted = true, + bool kWeightedSourceIsRhs = false, + bool kSynchronizeRanks = true, + bool kSynchronizeAfterDispatch = true, + bool kBarrierOnly = false, + bool kXPrepared = false, + uint32_t kRoutePreludeThreads = 256> +CUTLASS_GLOBAL __launch_bounds__(1024, 1) void +sm100_bf16_mega_moe_backward_post_down_prelude( + const int* expert_counts, + const __grid_constant__ layout::Workspace + backward_workspace, + const __grid_constant__ layout::SymBuffer + backward_sym_buffer, + const cutlass::bfloat16_t* backward_grad_y, + const cutlass::bfloat16_t* backward_x, + const float* backward_topk_weights, + float* backward_grad_route, + const layout::TokenSrcMetadata* token_src_metadata, + const uint32_t num_topk, + const uint32_t num_pool_rows, + cutlass::bfloat16_t* grad_y_unweighted_output, + cutlass::bfloat16_t* grad_y_weighted_output, + cutlass::bfloat16_t* x_pool_output, + float* route_weights_output, + const cutlass::bfloat16_t* down_unweighted, + float* grad_route_output) { +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000)) || defined(__CLION_IDE__) + constexpr uint32_t kNumThreads = 1024; + if constexpr (kSynchronizeRanks) { + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, 0, 71>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + } + if constexpr (kBarrierOnly) { + if constexpr (kSynchronizeAfterDispatch) { + // Profiling-only completion barrier. Production folds this into + // the dispatch launch below to avoid an extra kernel launch. + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, 1, 72>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + } + return; + } + constexpr uint32_t kTritonRouteBlockH = [] { + uint32_t value = 1; + while (value < kHidden && value < 8192) + value <<= 1; + return value; + }(); + constexpr uint32_t kTritonRouteNumWarps = [] { + uint32_t value = kTritonRouteBlockH / 256; + value = value < 4 ? 4 : value; + return value > 32 ? 32 : value; + }(); + constexpr uint32_t kTritonRouteThreads = + kTritonRouteNumWarps * 32; + constexpr uint32_t kTritonRouteValuesPerThread = + kTritonRouteBlockH / kTritonRouteThreads; + constexpr bool kVirtualizeRouteLanes = + kRoutePreludeThreads == 128; + DG_STATIC_ASSERT( + kRoutePreludeThreads == 128 || + kRoutePreludeThreads == 256, + "POST_DOWN route prelude requires 128 or 256 physical threads"); + DG_STATIC_ASSERT( + !kVirtualizeRouteLanes || + (kHidden == 2048 && + kCombineOrderMode != + CombineOrderMode::FixedTopK && + kComputeRouteDot), + "128-thread route prelude is only supported for the exact " + "non-fixed H=2048 route-dot path"); + constexpr uint32_t kExactRouteGroupThreads = + kVirtualizeRouteLanes + ? kRoutePreludeThreads + : kTritonRouteThreads; + constexpr uint32_t kRouteVirtualLanes = + kTritonRouteThreads / + kExactRouteGroupThreads; + DG_STATIC_ASSERT( + kRouteVirtualLanes == 1 || + kRouteVirtualLanes == 2, + "Unsupported POST_DOWN route lane virtualization"); + DG_STATIC_ASSERT( + kTritonRouteValuesPerThread == 2 || + kTritonRouteValuesPerThread == 4 || + kTritonRouteValuesPerThread == 8, + "Unsupported Triton route reduction width"); + constexpr uint32_t kRouteInputPow2 = [] { + uint32_t value = 1; + constexpr uint32_t vectorized_columns = + kHidden / 4; + while (value < 512 && + (value << 1) <= vectorized_columns) + value <<= 1; + return value; + }(); + constexpr uint32_t kInitialRouteGroupThreads = + cute::min(kRouteInputPow2, 32u); + + extern __shared__ __align__(1024) uint8_t scratch[]; + auto* route_lane_sums = + reinterpret_cast(scratch); + auto* route_control = + reinterpret_cast(scratch); + if (threadIdx.x == 0) { + uint32_t total_route_rows = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + total_route_rows += static_cast( + __ldg(expert_counts + expert_idx)); + } + route_control[0] = total_route_rows; + } + __syncthreads(); + const uint32_t total_route_rows = route_control[0]; + const uint32_t route_output_pow2 = + total_route_rows > 0 + ? 1u << (31 - __clz(total_route_rows)) + : 1u; + const uint32_t route_block_height = + cute::min( + route_output_pow2, + 512u / kInitialRouteGroupThreads); + // The dispatch-only launch copies vectorized BF16 payloads and does not + // need Triton's exact reduction lane map. Use more route groups per CTA + // so remote reads have enough independent rows to cover NVLink latency. + const uint32_t route_group_threads = + !kComputeRouteDot && !kWriteWeighted + ? cute::min(kRouteInputPow2, 128u) + : kCombineOrderMode != CombineOrderMode::FixedTopK + ? kExactRouteGroupThreads + : cute::min( + kRouteInputPow2, + 512u / route_block_height); + const uint32_t num_route_groups_per_cta = + kNumThreads / route_group_threads; + const uint32_t route_group_idx = + threadIdx.x / route_group_threads; + const uint32_t route_group_lane_idx = + threadIdx.x & (route_group_threads - 1); + const uint32_t global_route_group = + blockIdx.x * num_route_groups_per_cta + + route_group_idx; + const uint32_t num_route_groups = + kNumSMs * num_route_groups_per_cta; + uint32_t route_pool_block_offset = 0; + + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = static_cast( + __ldg(expert_counts + expert_idx)); + for (uint32_t token_idx = global_route_group; + token_idx < num_tokens; + token_idx += num_route_groups) { + const uint32_t pool_row = + route_pool_block_offset * BLOCK_M + + token_idx; + const cutlass::bfloat16_t* remote_grad_y; + const cutlass::bfloat16_t* remote_x; + if constexpr (kDoReverseDispatch) { + const auto metadata = + token_src_metadata[pool_row]; + remote_grad_y = + backward_sym_buffer.map( + backward_grad_y + + static_cast( + metadata.token_idx) * + kHidden, + metadata.rank_idx); + if constexpr (kXPrepared) { + remote_x = + x_pool_output + + static_cast(pool_row) * + kHidden; + } else { + remote_x = + backward_sym_buffer.map( + backward_x + + static_cast( + metadata.token_idx) * + kHidden, + metadata.rank_idx); + } + if (route_group_lane_idx == 0) { + const auto* remote_weight = + backward_sym_buffer.map( + backward_topk_weights + + static_cast( + metadata.token_idx) * + num_topk + + metadata.topk_idx, + metadata.rank_idx); + route_weights_output[pool_row] = + *remote_weight; + } + } else { + remote_grad_y = + grad_y_unweighted_output + + static_cast(pool_row) * + kHidden; + remote_x = + x_pool_output + + static_cast(pool_row) * + kHidden; + } + + float grad_route = 0.0f; + if constexpr (!kComputeRouteDot && !kWriteWeighted) { + constexpr uint32_t kBF16ValuesPerVector = + sizeof(uint4) / + sizeof(cutlass::bfloat16_t); + DG_STATIC_ASSERT( + kHidden % kBF16ValuesPerVector == 0, + "BF16 dispatch requires vector-aligned hidden"); + for (uint32_t col = + route_group_lane_idx * + kBF16ValuesPerVector; + col < kHidden; + col += route_group_threads * + kBF16ValuesPerVector) { + const uint64_t offset = + static_cast( + pool_row) * + kHidden + + col; + reinterpret_cast( + grad_y_unweighted_output)[ + offset / + kBF16ValuesPerVector] = + reinterpret_cast< + const uint4*>( + remote_grad_y)[ + col / + kBF16ValuesPerVector]; + if constexpr (!kXPrepared) { + reinterpret_cast( + x_pool_output)[ + offset / + kBF16ValuesPerVector] = + reinterpret_cast< + const uint4*>( + remote_x)[ + col / + kBF16ValuesPerVector]; + } + } + } else if constexpr ( + !kComputeRouteDot && kWriteWeighted) { + const float route_weight = + route_weights_output[pool_row]; + for (uint32_t col = + route_group_lane_idx; + col < kHidden; + col += route_group_threads) { + const uint64_t offset = + static_cast( + pool_row) * + kHidden + + col; + grad_y_weighted_output[offset] = + cutlass::bfloat16_t( + static_cast( + (kWeightedSourceIsRhs + ? down_unweighted[ + static_cast( + pool_row) * + kHidden + + col] + : remote_grad_y[col])) * + route_weight); + } + } else if constexpr ( + kCombineOrderMode != + CombineOrderMode::FixedTopK) { + if constexpr (kVirtualizeRouteLanes) { + // Each physical lane evaluates logical lanes p and + // p + 128. Their FMA and warp-XOR trees remain separate, + // then the four physical warps publish all eight logical + // Triton warp partials for the unchanged second level. + float weighted_values + [kRouteVirtualLanes] + [kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t virtual_lane = 0; + virtual_lane < kRouteVirtualLanes; + ++virtual_lane) { + const uint32_t logical_route_lane = + route_group_lane_idx + + virtual_lane * + kExactRouteGroupThreads; + float grad_y[ + kTritonRouteValuesPerThread]; + float down[ + kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + logical_route_lane + + i * kTritonRouteThreads; + grad_y[i] = + col < kHidden + ? static_cast( + remote_grad_y[col]) + : 0.0f; + down[i] = + col < kHidden + ? static_cast( + down_unweighted[ + static_cast( + pool_row) * + kHidden + + col]) + : 0.0f; + if constexpr ( + kDoReverseDispatch && + !kXPrepared) { + if (col < kHidden) { + x_pool_output[ + static_cast( + pool_row) * + kHidden + + col] = + remote_x[col]; + } + } + if constexpr (kWriteWeighted) { + weighted_values[virtual_lane][i] = + kWeightedSourceIsRhs + ? down[i] + : grad_y[i]; + } + } + float logical_grad_route; + if constexpr ( + kTritonRouteValuesPerThread == + 2) { + logical_grad_route = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[1], down[1])); + } else if constexpr ( + kTritonRouteValuesPerThread == + 4) { + const float even = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float odd = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + logical_grad_route = + __fadd_rn(even, odd); + } else { + const float pair_02 = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float pair_13 = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + const float pair_46 = __fmaf_rn( + grad_y[4], down[4], + __fmul_rn( + grad_y[6], down[6])); + const float pair_57 = __fmaf_rn( + grad_y[5], down[5], + __fmul_rn( + grad_y[7], down[7])); + logical_grad_route = __fadd_rn( + __fadd_rn( + pair_02, pair_46), + __fadd_rn( + pair_13, pair_57)); + } + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + logical_grad_route = __fadd_rn( + logical_grad_route, + __shfl_xor_sync( + 0xffffffff, + logical_grad_route, + offset)); + } + const uint32_t lane_in_warp = + route_group_lane_idx & 31; + const uint32_t logical_warp = + logical_route_lane / 32; + if (lane_in_warp == 0) { + route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + logical_warp] = + logical_grad_route; + } + } + ptx::sync_aligned( + kExactRouteGroupThreads, + route_group_idx); + const uint32_t physical_warp_in_group = + route_group_lane_idx / 32; + const uint32_t lane_in_warp = + route_group_lane_idx & 31; + if (physical_warp_in_group == 0) { + grad_route = route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + (lane_in_warp & + (kTritonRouteNumWarps - 1))]; + #pragma unroll + for (uint32_t offset = + kTritonRouteNumWarps / 2; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, offset)); + } + } + if constexpr (kWriteWeighted) { + const float route_weight = + route_weights_output[pool_row]; + #pragma unroll + for (uint32_t virtual_lane = 0; + virtual_lane < + kRouteVirtualLanes; + ++virtual_lane) { + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + route_group_lane_idx + + virtual_lane * + kExactRouteGroupThreads + + i * + kTritonRouteThreads; + if (col < kHidden) { + grad_y_weighted_output[ + static_cast( + pool_row) * + kHidden + + col] = + cutlass::bfloat16_t( + weighted_values[ + virtual_lane] + [i] * + route_weight); + } + } + } + } + } else { + float grad_y[ + kTritonRouteValuesPerThread]; + float down[ + kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + route_group_lane_idx + + i * kTritonRouteThreads; + grad_y[i] = + col < kHidden + ? static_cast( + remote_grad_y[col]) + : 0.0f; + down[i] = + col < kHidden + ? static_cast( + down_unweighted[ + static_cast( + pool_row) * + kHidden + + col]) + : 0.0f; + if constexpr ( + kDoReverseDispatch && + !kXPrepared) { + if (col < kHidden) { + x_pool_output[ + static_cast( + pool_row) * + kHidden + + col] = + remote_x[col]; + } + } + } + if constexpr ( + kTritonRouteValuesPerThread == + 2) { + grad_route = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[1], down[1])); + } else if constexpr ( + kTritonRouteValuesPerThread == + 4) { + const float even = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float odd = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + grad_route = + __fadd_rn(even, odd); + } else { + const float pair_02 = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], down[2])); + const float pair_13 = __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], down[3])); + const float pair_46 = __fmaf_rn( + grad_y[4], down[4], + __fmul_rn( + grad_y[6], down[6])); + const float pair_57 = __fmaf_rn( + grad_y[5], down[5], + __fmul_rn( + grad_y[7], down[7])); + grad_route = __fadd_rn( + __fadd_rn( + pair_02, pair_46), + __fadd_rn( + pair_13, pair_57)); + } + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, offset)); + } + const uint32_t warp_in_group = + route_group_lane_idx / 32; + const uint32_t lane_in_warp = + route_group_lane_idx & 31; + if (lane_in_warp == 0) { + route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + warp_in_group] = + grad_route; + } + ptx::sync_aligned( + kTritonRouteThreads, + route_group_idx); + if (warp_in_group == 0) { + grad_route = route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + (lane_in_warp & + (kTritonRouteNumWarps - + 1))]; + #pragma unroll + for (uint32_t offset = + kTritonRouteNumWarps / + 2; + offset > 0; + offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, + offset)); + } + } + if constexpr (kWriteWeighted) { + const float route_weight = + route_weights_output[ + pool_row]; + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + route_group_lane_idx + + i * + kTritonRouteThreads; + if (col < kHidden) { + grad_y_weighted_output[ + static_cast( + pool_row) * + kHidden + + col] = + cutlass::bfloat16_t( + (kWeightedSourceIsRhs + ? down[i] + : grad_y[i]) * + route_weight); + } + } + } + } + } else { + float lane_sums[4] = { + 0.0f, 0.0f, 0.0f, 0.0f}; + const float route_weight = + kWriteWeighted + ? route_weights_output[pool_row] + : 0.0f; + for (uint32_t col_base = + route_group_lane_idx * 4; + col_base < kHidden; + col_base += + route_group_threads * 4) { + #pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const uint32_t col = + col_base + i; + const float grad_y = + static_cast( + remote_grad_y[col]); + const float down = + static_cast( + down_unweighted[ + static_cast( + pool_row) * + kHidden + + col]); + if constexpr (kWriteWeighted) { + grad_y_weighted_output[ + static_cast( + pool_row) * + kHidden + + col] = + cutlass::bfloat16_t( + (kWeightedSourceIsRhs + ? down + : grad_y) * + route_weight); + } + if constexpr (kDoReverseDispatch) { + grad_y_unweighted_output[ + static_cast( + pool_row) * + kHidden + + col] = + cutlass::bfloat16_t( + grad_y); + if constexpr (!kXPrepared) { + x_pool_output[ + static_cast( + pool_row) * + kHidden + + col] = + remote_x[col]; + } + } + lane_sums[i] = __fadd_rn( + lane_sums[i], + __fmul_rn( + grad_y, down)); + } + } + grad_route = __fadd_rn( + __fadd_rn( + lane_sums[0], + lane_sums[1]), + lane_sums[2]); + grad_route = __fadd_rn( + grad_route, lane_sums[3]); + route_lane_sums[threadIdx.x] = + grad_route; + if (route_group_threads > 32) { + for (uint32_t offset = + route_group_threads / 2; + offset >= 32; + offset >>= 1) { + ptx::sync_aligned( + route_group_threads, + route_group_idx); + if (route_group_lane_idx < + offset) { + grad_route = __fadd_rn( + grad_route, + route_lane_sums[ + threadIdx.x + + offset]); + route_lane_sums[ + threadIdx.x] = + grad_route; + } + } + } + if (route_group_lane_idx < 32) { + #pragma unroll + for (uint32_t offset = 16; + offset > 0; offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_down_sync( + 0xffffffff, + grad_route, offset)); + } + } + } + if constexpr (kComputeRouteDot) { + if (route_group_lane_idx == 0) { + grad_route_output[pool_row] = + grad_route; + if (backward_grad_route != nullptr) { + const auto metadata = + token_src_metadata[pool_row]; + auto* remote_grad_route = + backward_sym_buffer.map( + backward_grad_route + + static_cast( + metadata.token_idx) * + num_topk + + metadata.topk_idx, + metadata.rank_idx); + *remote_grad_route = grad_route; + } + } + if (route_group_threads > 32) { + ptx::sync_aligned( + route_group_threads, + route_group_idx); + } else { + __syncwarp(); + } + } + } + route_pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } + + uint32_t padded_pool_blocks = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = static_cast( + __ldg(expert_counts + expert_idx)); + const uint32_t num_blocks = + math::ceil_div(num_tokens, BLOCK_M); + const uint32_t num_padded_tokens = + num_blocks * BLOCK_M; + const uint32_t num_padding_rows = + num_padded_tokens - num_tokens; + for (uint64_t linear = + static_cast(blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + num_padding_rows) * + kHidden; + linear += + static_cast(kNumSMs) * + kNumThreads) { + const uint32_t padding_row = + linear / kHidden; + const uint32_t col = + linear - + static_cast(padding_row) * + kHidden; + const uint32_t pool_row = + padded_pool_blocks * BLOCK_M + + num_tokens + padding_row; + const uint64_t offset = + static_cast(pool_row) * + kHidden + + col; + if constexpr (kDoReverseDispatch) { + grad_y_unweighted_output[offset] = + cutlass::bfloat16_t(0.0f); + if constexpr (!kXPrepared) { + x_pool_output[offset] = + cutlass::bfloat16_t(0.0f); + } + } + if constexpr (kWriteWeighted) { + grad_y_weighted_output[offset] = + cutlass::bfloat16_t(0.0f); + } + } + for (uint32_t padding_row = + blockIdx.x * kNumThreads + + threadIdx.x; + padding_row < num_padding_rows; + padding_row += + kNumSMs * kNumThreads) { + const uint32_t pool_row = + padded_pool_blocks * BLOCK_M + + num_tokens + padding_row; + if constexpr (kDoReverseDispatch) + route_weights_output[pool_row] = 0.0f; + if constexpr (kComputeRouteDot) + grad_route_output[pool_row] = 0.0f; + } + padded_pool_blocks += num_blocks; + } + if constexpr (kSynchronizeAfterDispatch) { + // Every rank may reuse its local symmetric grad-y plane as the + // direct-write grad-x destination as soon as this producer returns. + // Publish completion only after all peers have finished their remote + // reads; an entry-only barrier permits checkpoint/replay rank skew to + // corrupt those in-flight pulls. + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, 1, 72>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + } + // Rows beyond the final padded expert block are capacity only. No + // downstream kernel addresses them, so clearing that high-water tail + // wastes bandwidth and grows with the cached pool margin. +#endif +} + +// Production MegaMoE backward wave. This persistent kernel consumes the +// forward kernel's block-padded expert pool directly and replays gate and up +// together as one W13 FP8xFP4 mainloop before computing the retained dgrads. +template < + uint32_t kHidden, uint32_t kIntermediateHidden, + uint32_t kNumExperts, + uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K, + uint32_t SF_BLOCK_M, uint32_t SF_BLOCK_N, + uint32_t kNumStages, + uint32_t kNumSMs, + uint32_t kNumRanks = 1, + bool kCompileW13Dgrad = true, + bool kBF16Mode = false, + ActivationType kActivationType = ActivationType::SwiGLU, + bool kFastMath = false, + RouteWeightMode kRouteWeightMode = RouteWeightMode::PreDown, + CombineOrderMode kCombineOrderMode = CombineOrderMode::FixedTopK, + bool kInputsPrepared = false, + bool kDispatchInputsPrepared = false, + bool kDirectRemoteGradX = false, + bool kWriteGradXPool = true, + bool kClearWgradPadding = false, + bool kComputeRouteGrad = false, + bool kTraceKernel = false, + bool kVectorizedGradXStore = false, + bool kWideGradXStore = false, + uint32_t kNumNonEpilogueThreads = 128, + uint32_t kNumEpilogueThreads = 128, + uint32_t kNumThreads = + kNumNonEpilogueThreads + kNumEpilogueThreads + + 768> +CUTLASS_GLOBAL __launch_bounds__(kNumThreads, 1) void +sm100_fp8_fp4_mega_moe_backward_wave_impl( + const int* expert_counts, + const __grid_constant__ layout::SymBuffer backward_sym_buffer, + const __grid_constant__ layout::Workspace backward_workspace, + const cutlass::bfloat16_t* backward_grad_y, + const cutlass::bfloat16_t* backward_x, + const float* backward_topk_weights, + float* backward_grad_route, + const layout::TokenSrcMetadata* token_src_metadata, + const uint32_t num_topk, + const uint32_t num_pool_rows, + const uint32_t acts_sf_stride, + const __grid_constant__ cute::TmaDescriptor tensor_map_acts, + const __grid_constant__ cute::TmaDescriptor tensor_map_acts_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_weights, + const __grid_constant__ cute::TmaDescriptor tensor_map_weights_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_output, + const __grid_constant__ cute::TmaDescriptor tensor_map_grad_ye, + const __grid_constant__ cute::TmaDescriptor tensor_map_w2_dequant, + const __grid_constant__ cute::TmaDescriptor tensor_map_w2_weights, + const __grid_constant__ cute::TmaDescriptor tensor_map_w2_scales, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_dequant, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_weights, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_scales, + const __grid_constant__ cute::TmaDescriptor tensor_map_grad_gate_up, + const cutlass::float_e4m3_t* acts_ptr, + const uint32_t* acts_sf_ptr, + const int8_t* w2_weights, + const float* w2_scales, + cutlass::bfloat16_t* w2_dequant_scratch, + const int8_t* w13_weights, + const float* w13_scales, + cutlass::bfloat16_t* w13_dequant_scratch, + const cutlass::bfloat16_t* gate_up_output, + cutlass::bfloat16_t* grad_ye_output, + cutlass::bfloat16_t* grad_y_unweighted_output, + cutlass::bfloat16_t* route_weights, + float* route_weights_fp32, + cutlass::bfloat16_t* grad_h_output, + cutlass::bfloat16_t* grad_gate_up_output, + cutlass::bfloat16_t* h_act_output, + cutlass::bfloat16_t* h_weighted_output, + cutlass::bfloat16_t* x_pool_output, + cutlass::bfloat16_t* grad_x_pool_output, + const cutlass::bfloat16_t* down_unweighted_output, + float* grad_route_output, + uint32_t* weight_tile_states, + const uint32_t launch_epoch, + const float activation_limit, + uint64_t* kernel_trace) { +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000)) || defined(__CLION_IDE__) + using Barrier = cutlass::arch::ClusterTransactionBarrier; + using Allocator = cute::TMEM::Allocator2Sm; + using a_dtype_t = cutlass::float_e4m3_t; + using b_dtype_t = cutlass::detail::float_e2m1_unpacksmem_t; + using cd_dtype_t = cutlass::bfloat16_t; + using dgrad_b_dtype_t = cutlass::bfloat16_t; + + constexpr uint32_t kNumEpilogueStages = 2; + constexpr uint32_t kNumTMAStoreStages = 2; + constexpr uint32_t kNumDispatchThreads = + kNumRanks > 1 ? 128 : 0; + constexpr uint32_t kNumDispatchWarps = + kNumDispatchThreads / 32; + constexpr uint32_t kDispatchWarpStart = + (kNumNonEpilogueThreads + kNumEpilogueThreads) / 32; + constexpr uint32_t kNumDgradEpilogueThreads = + kNumThreads - kNumNonEpilogueThreads; + constexpr uint32_t kGranK = 32; + constexpr uint32_t kNumUTCCPAlignedElems = 128; + constexpr uint32_t kNumBlockNs = (2 * kIntermediateHidden) / BLOCK_N; + constexpr uint32_t kNumDgradBlockNs = kIntermediateHidden / BLOCK_N; + constexpr uint32_t kNumW13DgradBlockNs = kHidden / BLOCK_N; + constexpr uint32_t kNumW13DgradSplits = + kBF16Mode ? 2 : 1; + constexpr uint32_t LAYOUT_AD_M = 128; + constexpr uint32_t UMMA_M = LAYOUT_AD_M * 2; + constexpr uint32_t UMMA_N = BLOCK_M; + constexpr uint32_t UMMA_K = 32; + constexpr uint32_t DGRAD_BLOCK_K = 64; + constexpr uint32_t DGRAD_UMMA_K = 16; + constexpr uint32_t kNumW2WeightTileStates = + kNumExperts * (kHidden / DGRAD_BLOCK_K) * + kNumDgradBlockNs; + constexpr uint32_t LOAD_BLOCK_M = BLOCK_M / 2; + constexpr uint32_t LOAD_BLOCK_N = BLOCK_N; + constexpr uint32_t STORE_BLOCK_M = 16; + constexpr uint32_t STORE_BLOCK_N = BLOCK_N; + constexpr uint32_t kSwizzleAMode = BLOCK_K * sizeof(a_dtype_t); + constexpr uint32_t kSwizzleBMode = BLOCK_K * sizeof(b_dtype_t); + constexpr uint32_t kSwizzleCDMode = 128; + + DG_STATIC_ASSERT(kNumNonEpilogueThreads == 128, "Invalid producer thread count"); + DG_STATIC_ASSERT(kNumEpilogueThreads == 128, "Invalid epilogue thread count"); + DG_STATIC_ASSERT(kNumRanks == 1 || kNumDispatchThreads == 128, + "Invalid backward dispatch thread count"); + DG_STATIC_ASSERT(BLOCK_M % 16 == 0 && BLOCK_N == 128 && BLOCK_K == 128, + "Invalid backward wave tile"); + DG_STATIC_ASSERT(kNumBlockNs % 2 == 0, "Cluster peers must receive adjacent N blocks"); + DG_STATIC_ASSERT(kNumDgradBlockNs % 2 == 0, + "Dgrad cluster peers must receive adjacent N blocks"); + DG_STATIC_ASSERT(kNumW13DgradBlockNs % 2 == 0, + "W13 dgrad cluster peers must receive adjacent N blocks"); + DG_STATIC_ASSERT(SF_BLOCK_M == math::constexpr_align(BLOCK_M, kNumUTCCPAlignedElems), + "Invalid SFA block"); + DG_STATIC_ASSERT(SF_BLOCK_N == BLOCK_N, "Invalid SFB block"); + DG_STATIC_ASSERT(kHidden % BLOCK_K == 0, "Invalid hidden size"); + DG_STATIC_ASSERT(kNumSMs % 2 == 0, "2-CTA clusters require an even SM count"); + + constexpr uint32_t kNumW13WeightTileStates = + kNumExperts * + ((2 * kIntermediateHidden) / DGRAD_BLOCK_K) * + kNumW13DgradBlockNs; + auto* phase_count = + weight_tile_states + kNumW2WeightTileStates + + kNumW13WeightTileStates; + auto* phase_sense = phase_count + 1; + constexpr uint32_t kTraceSiteCount = 22; + constexpr uint32_t kTraceValueCount = 5; + constexpr uint32_t kTraceBeginCycle = 0; + constexpr uint32_t kTraceEndCycle = 1; + constexpr uint32_t kTraceBeginGlobalNs = 2; + constexpr uint32_t kTraceEndGlobalNs = 3; + constexpr uint32_t kTraceSM = 4; + const auto globaltimer = [] { + uint64_t value; + asm volatile( + "mov.u64 %0, %%globaltimer;" : "=l"(value)); + return value; + }; + const auto trace_begin = [&](const uint32_t site) { + if constexpr (kTraceKernel) { + if (threadIdx.x == 0) { + auto* values = + kernel_trace + + (static_cast(site) * kNumSMs + + blockIdx.x) * + kTraceValueCount; + values[kTraceBeginCycle] = clock64(); + values[kTraceBeginGlobalNs] = globaltimer(); + values[kTraceSM] = ptx::get_sm_idx(); + } + } + }; + const auto trace_end = [&](const uint32_t site) { + if constexpr (kTraceKernel) { + if (threadIdx.x == 0) { + auto* values = + kernel_trace + + (static_cast(site) * kNumSMs + + blockIdx.x) * + kTraceValueCount; + values[kTraceEndCycle] = clock64(); + values[kTraceEndGlobalNs] = globaltimer(); + } + } + }; + if constexpr (kTraceKernel) { + DG_STATIC_ASSERT( + kTraceSiteCount == 22, + "Update the host trace-site schema with the kernel"); + trace_begin(0); + } + const auto full_grid_phase_barrier = + [&](const uint32_t trace_site) { + trace_begin(trace_site); + if (threadIdx.x == 0) { + const uint32_t old_sense = + atomicAdd(phase_sense, 0u); + __threadfence(); + const uint32_t ticket = + atomicAdd(phase_count, 1u); + if (ticket == kNumSMs - 1) { + atomicExch(phase_count, 0u); + __threadfence(); + atomicAdd(phase_sense, 1u); + } else { + while (ptx::ld_acq(phase_sense) == + old_sense) { + } + } + } + __syncthreads(); + trace_end(trace_site); + }; + + const bool is_leader_cta = cute::block_rank_in_cluster() == 0; + const uint32_t warp_idx = cutlass::canonical_warp_idx_sync(); + const uint32_t lane_idx = ptx::get_lane_idx(); + + if (warp_idx == 0) { + cute::prefetch_tma_descriptor(&tensor_map_grad_ye); + cute::prefetch_tma_descriptor(&tensor_map_w2_dequant); + cute::prefetch_tma_descriptor(&tensor_map_w13_dequant); + cute::prefetch_tma_descriptor(&tensor_map_grad_gate_up); + if constexpr (!kBF16Mode) { + cute::prefetch_tma_descriptor(&tensor_map_acts); + cute::prefetch_tma_descriptor(&tensor_map_acts_sf); + cute::prefetch_tma_descriptor(&tensor_map_weights); + cute::prefetch_tma_descriptor(&tensor_map_weights_sf); + cute::prefetch_tma_descriptor(&tensor_map_output); + cute::prefetch_tma_descriptor(&tensor_map_w2_weights); + cute::prefetch_tma_descriptor(&tensor_map_w2_scales); + cute::prefetch_tma_descriptor(&tensor_map_w13_weights); + cute::prefetch_tma_descriptor(&tensor_map_w13_scales); + } + } + + constexpr uint32_t SMEM_CD_SIZE_PER_STAGE = + STORE_BLOCK_M * STORE_BLOCK_N * sizeof(cd_dtype_t); + constexpr uint32_t SMEM_CD_SIZE = SMEM_CD_SIZE_PER_STAGE * kNumTMAStoreStages; + 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); + constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = SF_BLOCK_M * sizeof(uint32_t); + constexpr uint32_t SMEM_SFB_SIZE_PER_STAGE = SF_BLOCK_N * sizeof(uint32_t); + constexpr uint32_t SMEM_DISPATCH_SIZE = + kNumDispatchWarps * kHidden * sizeof(cd_dtype_t); + + extern __shared__ __align__(1024) uint8_t smem_buffer[]; + auto* smem_gemm_base = smem_buffer + SMEM_DISPATCH_SIZE; + auto smem_cd = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + smem_gemm_base + i * SMEM_CD_SIZE_PER_STAGE); + }); + auto smem_a = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + smem_gemm_base + SMEM_CD_SIZE + + i * SMEM_A_SIZE_PER_STAGE); + }); + auto smem_b = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + smem_gemm_base + SMEM_CD_SIZE + + kNumStages * SMEM_A_SIZE_PER_STAGE + + i * SMEM_B_SIZE_PER_STAGE); + }); + // The dgrad phase aliases the recompute mainloop storage exactly: + // FP8 A [BLOCK_M/2, 128] == BF16 A [BLOCK_M/2, 64] + // packed-FP4 B [128, 128] == BF16 B [128, 64]. + // W2 is dequantized and transposed directly into the latter by a producer + // warp, so no persistent weight copy or host-side packing is needed. + auto smem_dgrad_a = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + smem_gemm_base + SMEM_CD_SIZE + + i * SMEM_A_SIZE_PER_STAGE); + }); + auto smem_dgrad_b = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + smem_gemm_base + SMEM_CD_SIZE + + kNumStages * SMEM_A_SIZE_PER_STAGE + + i * SMEM_B_SIZE_PER_STAGE); + }); + DG_STATIC_ASSERT( + LOAD_BLOCK_M * DGRAD_BLOCK_K * sizeof(cd_dtype_t) == + SMEM_A_SIZE_PER_STAGE, + "Dgrad A alias size mismatch"); + DG_STATIC_ASSERT( + LOAD_BLOCK_N * DGRAD_BLOCK_K * sizeof(dgrad_b_dtype_t) == + SMEM_B_SIZE_PER_STAGE, + "Dgrad B alias size mismatch"); + auto sf_start_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); + }); + auto smem_sfb = utils::PatternVisitor([=](const uint32_t& i) { + return reinterpret_cast( + sf_start_ptr + kNumStages * SMEM_SFA_SIZE_PER_STAGE + + i * SMEM_SFB_SIZE_PER_STAGE); + }); + + auto barrier_start_ptr = reinterpret_cast(smem_sfb[kNumStages]); + auto full_barriers = utils::PatternVisitor( + [=](const uint32_t& i) { return barrier_start_ptr + i; }); + auto empty_barriers = utils::PatternVisitor( + [=](const uint32_t& i) { return barrier_start_ptr + kNumStages + i; }); + auto tmem_full_barriers = utils::PatternVisitor([=](const uint32_t& i) { + return barrier_start_ptr + 2 * kNumStages + i; + }); + auto tmem_empty_barriers = utils::PatternVisitor([=](const uint32_t& i) { + return barrier_start_ptr + 2 * kNumStages + kNumEpilogueStages + i; + }); + auto dispatch_barriers = utils::PatternVisitor([=](const uint32_t& i) { + return barrier_start_ptr + 2 * kNumStages + + 2 * kNumEpilogueStages + i; + }); + auto tmem_ptr_in_smem = reinterpret_cast( + barrier_start_ptr + 2 * kNumStages + + 2 * kNumEpilogueStages + kNumDispatchWarps); + + constexpr uint32_t kNumAccumTmemCols = UMMA_N * kNumEpilogueStages; + constexpr uint32_t kNumSFATmemCols = SF_BLOCK_M / 32; + constexpr uint32_t kNumSFBTmemCols = SF_BLOCK_N / 32; + constexpr uint32_t kNumTmemCols = + utils::get_num_aligned_tmem_cols< + kNumAccumTmemCols + kNumSFATmemCols + kNumSFBTmemCols>(); + constexpr uint32_t kTmemStartColOfSFA = kNumAccumTmemCols; + constexpr uint32_t kTmemStartColOfSFB = + kNumAccumTmemCols + kNumSFATmemCols; + DG_STATIC_ASSERT(kNumTmemCols <= 512, "Backward recompute exceeds TMEM"); + + if constexpr (!kBF16Mode) { + // Dequantize W2 exactly once per launch into an ephemeral + // [expert, dim, H] BF16 workspace. Keeping the source orientation + // makes both packed-FP4 reads and BF16 writes coalesced; dgrad consumes + // it as an MN-major transposed operand. + constexpr uint32_t kDequantTileK = 256; + constexpr uint32_t kDequantTileN = LOAD_BLOCK_N; + constexpr uint32_t kDequantPairsPerTile = + kDequantTileK * kDequantTileN / 2; + constexpr uint32_t kDequantSFsPerK = + kDequantTileN / 32; + constexpr uint32_t kDequantSFsPerTile = + kDequantTileK * kDequantSFsPerK; + constexpr uint32_t kDequantWeightBytes = + kDequantTileK * (kDequantTileN / 2); + constexpr uint32_t kDequantScaleBytes = + kDequantSFsPerTile * sizeof(float); + constexpr uint32_t kNumDequantKTiles = + kHidden / kDequantTileK; + constexpr uint32_t kNumDequantNTiles = + kIntermediateHidden / kDequantTileN; + constexpr uint32_t kNumDequantTiles = + kNumExperts * kNumDequantKTiles * + kNumDequantNTiles; + auto* dequant_weights = + reinterpret_cast(smem_buffer); + auto* dequant_scales = + reinterpret_cast( + smem_buffer + kDequantWeightBytes); + auto* dequant_scale_half2 = + reinterpret_cast( + smem_buffer + kDequantWeightBytes + + kDequantScaleBytes); + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 0 && cute::elect_one_sync()) { + full_barriers[0]->init(1); + if constexpr (kCompileW13Dgrad) + full_barriers[1]->init(1); + cutlass::arch::fence_barrier_init(); + } + comm::cluster_sync_with_relaxed_arrive(); + uint32_t dequant_phase = 0; + + for (uint32_t tile_idx = blockIdx.x; + tile_idx < kNumDequantTiles; + tile_idx += kNumSMs) { + const uint32_t n_tile_idx = + tile_idx % kNumDequantNTiles; + const uint32_t k_expert_tile_idx = + tile_idx / kNumDequantNTiles; + const uint32_t k_tile_idx = + k_expert_tile_idx % kNumDequantKTiles; + const uint32_t expert_idx = + k_expert_tile_idx / kNumDequantKTiles; + const uint32_t global_k_base = + k_tile_idx * kDequantTileK; + const uint32_t global_n_base = + n_tile_idx * kDequantTileN; + + if (warp_idx == 0 && cute::elect_one_sync()) { + tma::copy< + kDequantTileN / 2, kDequantTileK, 0, + int8_t>( + &tensor_map_w2_weights, + full_barriers[0], dequant_weights, + global_n_base / 2, + expert_idx * kHidden + global_k_base); + tma::copy< + kDequantSFsPerK, kDequantTileK, 0, + float>( + &tensor_map_w2_scales, + full_barriers[0], dequant_scales, + global_n_base / 32, + expert_idx * kHidden + global_k_base); + full_barriers[0]->arrive_and_expect_tx( + kDequantWeightBytes + + kDequantScaleBytes); + } + full_barriers[0]->wait(dequant_phase); + __syncthreads(); + + for (uint32_t scale_idx = threadIdx.x; + scale_idx < kDequantSFsPerTile; + scale_idx += kNumThreads) { + const auto scale_half2 = + __float2half2_rn( + dequant_scales[scale_idx]); + dequant_scale_half2[scale_idx] = + *reinterpret_cast( + &scale_half2); + } + __syncthreads(); + + for (uint32_t pair_idx = threadIdx.x; + pair_idx < kDequantPairsPerTile; + pair_idx += kNumThreads) { + const uint32_t local_k = + pair_idx / (kDequantTileN / 2); + const uint32_t local_n_pair = + pair_idx % (kDequantTileN / 2); + const uint32_t global_k = + global_k_base + local_k; + const uint32_t global_n_pair = + global_n_base / 2 + local_n_pair; + const uint8_t packed = + static_cast( + dequant_weights[ + local_k * + (kDequantTileN / 2) + + local_n_pair]); + uint32_t fp16x2; + asm volatile( + "{\n" + ".reg .b8 fp4;\n" + ".reg .b8 unused1, unused2, unused3;\n" + "mov.b32 {fp4, unused1, unused2, unused3}, %1;\n" + "cvt.rn.f16x2.e2m1x2 %0, fp4;\n" + "}\n" + : "=r"(fp16x2) + : "r"(static_cast(packed))); + auto value_pair = + *reinterpret_cast<__half2*>(&fp16x2); + const uint32_t scale_half2_bits = + dequant_scale_half2[ + local_k * kDequantSFsPerK + + (local_n_pair * 2) / 32]; + const auto scale_half2 = + *reinterpret_cast( + &scale_half2_bits); + value_pair = + __hmul2(value_pair, scale_half2); + const auto value_pair_bf16 = + __float22bfloat162_rn( + __half22float2(value_pair)); + const uint32_t scaled_pair = + *reinterpret_cast( + &value_pair_bf16); + *reinterpret_cast( + w2_dequant_scratch + + (static_cast(expert_idx) * + kHidden + + global_k) * + kIntermediateHidden + + global_n_pair * 2) = + scaled_pair; + } + __syncthreads(); + if (threadIdx.x < + kDequantTileK / DGRAD_BLOCK_K) { + const uint32_t dgrad_k_block_idx = + k_tile_idx * + (kDequantTileK / + DGRAD_BLOCK_K) + + threadIdx.x; + const uint32_t weight_tile_idx = + (expert_idx * + (kHidden / DGRAD_BLOCK_K) + + dgrad_k_block_idx) * + kNumDgradBlockNs + + n_tile_idx; + asm volatile( + "st.release.gpu.global.u32 [%0], %1;" + :: "l"(weight_tile_states + + weight_tile_idx), + "r"(launch_epoch) + : "memory"); + } + __syncthreads(); + dequant_phase ^= 1; + } + + if constexpr (kCompileW13Dgrad) { + constexpr uint32_t kW13DequantTileK = 256; + constexpr uint32_t kW13DequantTileN = LOAD_BLOCK_N; + constexpr uint32_t kW13DequantPairsPerTile = + kW13DequantTileK * kW13DequantTileN / 2; + constexpr uint32_t kW13DequantSFsPerK = + kW13DequantTileN / 32; + constexpr uint32_t kW13DequantSFsPerTile = + kW13DequantTileK * kW13DequantSFsPerK; + constexpr uint32_t kW13DequantWeightBytes = + kW13DequantTileK * (kW13DequantTileN / 2); + constexpr uint32_t kW13DequantScaleBytes = + kW13DequantSFsPerTile * sizeof(float); + constexpr uint32_t kNumW13DequantKTiles = + (2 * kIntermediateHidden) / kW13DequantTileK; + constexpr uint32_t kNumW13DequantNTiles = + kHidden / kW13DequantTileN; + constexpr uint32_t kNumW13DequantTiles = + kNumExperts * kNumW13DequantKTiles * + kNumW13DequantNTiles; + const uint32_t w13_launch_epoch = + launch_epoch ^ 0x80000000u; + uint32_t w13_dequant_phase = 0; + + for (uint32_t tile_idx = blockIdx.x; + tile_idx < kNumW13DequantTiles; + tile_idx += kNumSMs) { + const uint32_t n_tile_idx = + tile_idx % kNumW13DequantNTiles; + const uint32_t k_expert_tile_idx = + tile_idx / kNumW13DequantNTiles; + const uint32_t k_tile_idx = + k_expert_tile_idx % kNumW13DequantKTiles; + const uint32_t expert_idx = + k_expert_tile_idx / kNumW13DequantKTiles; + const uint32_t global_k_base = + k_tile_idx * kW13DequantTileK; + const uint32_t global_n_base = + n_tile_idx * kW13DequantTileN; + + if (warp_idx == 0 && cute::elect_one_sync()) { + tma::copy< + kW13DequantTileN / 2, + kW13DequantTileK, 0, int8_t>( + &tensor_map_w13_weights, + full_barriers[1], + dequant_weights, + global_n_base / 2, + expert_idx * + (2 * kIntermediateHidden) + + global_k_base); + tma::copy< + kW13DequantSFsPerK, + kW13DequantTileK, 0, float>( + &tensor_map_w13_scales, + full_barriers[1], + dequant_scales, + global_n_base / 32, + expert_idx * + (2 * kIntermediateHidden) + + global_k_base); + full_barriers[1]->arrive_and_expect_tx( + kW13DequantWeightBytes + + kW13DequantScaleBytes); + } + full_barriers[1]->wait( + w13_dequant_phase); + __syncthreads(); + + for (uint32_t scale_idx = threadIdx.x; + scale_idx < kW13DequantSFsPerTile; + scale_idx += kNumThreads) { + const auto scale_half2 = + __float2half2_rn( + dequant_scales[scale_idx]); + dequant_scale_half2[scale_idx] = + *reinterpret_cast< + const uint32_t*>( + &scale_half2); + } + __syncthreads(); + + for (uint32_t pair_idx = threadIdx.x; + pair_idx < + kW13DequantPairsPerTile; + pair_idx += kNumThreads) { + const uint32_t local_k = + pair_idx / + (kW13DequantTileN / 2); + const uint32_t local_n_pair = + pair_idx % + (kW13DequantTileN / 2); + const uint32_t global_k = + global_k_base + local_k; + const uint32_t global_n_pair = + global_n_base / 2 + + local_n_pair; + const uint8_t packed = + static_cast( + dequant_weights[ + local_k * + (kW13DequantTileN / 2) + + local_n_pair]); + uint32_t fp16x2; + asm volatile( + "{\n" + ".reg .b8 fp4;\n" + ".reg .b8 unused1, unused2, unused3;\n" + "mov.b32 {fp4, unused1, unused2, unused3}, %1;\n" + "cvt.rn.f16x2.e2m1x2 %0, fp4;\n" + "}\n" + : "=r"(fp16x2) + : "r"( + static_cast( + packed))); + auto value_pair = + *reinterpret_cast<__half2*>( + &fp16x2); + const uint32_t scale_half2_bits = + dequant_scale_half2[ + local_k * + kW13DequantSFsPerK + + (local_n_pair * 2) / 32]; + value_pair = __hmul2( + value_pair, + *reinterpret_cast< + const __half2*>( + &scale_half2_bits)); + const auto value_pair_bf16 = + __float22bfloat162_rn( + __half22float2( + value_pair)); + *reinterpret_cast( + w13_dequant_scratch + + (static_cast( + expert_idx) * + (2 * + kIntermediateHidden) + + global_k) * + kHidden + + global_n_pair * 2) = + *reinterpret_cast< + const uint32_t*>( + &value_pair_bf16); + } + __syncthreads(); + + if (threadIdx.x < + kW13DequantTileK / + DGRAD_BLOCK_K) { + const uint32_t + dgrad_k_block_idx = + k_tile_idx * + (kW13DequantTileK / + DGRAD_BLOCK_K) + + threadIdx.x; + const uint32_t weight_tile_idx = + (expert_idx * + ((2 * + kIntermediateHidden) / + DGRAD_BLOCK_K) + + dgrad_k_block_idx) * + kNumW13DgradBlockNs + + n_tile_idx; + asm volatile( + "st.release.gpu.global.u32 [%0], %1;" + :: "l"(weight_tile_states + + kNumW2WeightTileStates + + weight_tile_idx), + "r"(w13_launch_epoch) + : "memory"); + } + __syncthreads(); + w13_dequant_phase ^= 1; + } + } + } + trace_begin(1); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(1); + if (warp_idx == 0 && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kNumStages; ++i) { + full_barriers[i]->init(4); + empty_barriers[i]->init(1); + } + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { + tmem_full_barriers[i]->init(1); + tmem_empty_barriers[i]->init(2 * kNumEpilogueThreads); + } + #pragma unroll + for (uint32_t i = 0; i < kNumDispatchWarps; ++i) + dispatch_barriers[i]->init(1); + cutlass::arch::fence_barrier_init(); + } else if (warp_idx == 1) { + Allocator().allocate(kNumTmemCols, tmem_ptr_in_smem); + } + trace_begin(2); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(2); + + // Every role walks this deterministic schedule independently. Pool offsets + // are prefixes of ceil(count/BLOCK_M), matching the forward MegaMoE layout. + const auto for_each_block = [&](const auto& func) { + uint32_t next_assigned_block = blockIdx.x; + uint32_t global_block = 0; + uint32_t pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast(__ldg(expert_counts + expert_idx)); + const uint32_t num_m_blocks = math::ceil_div(num_tokens, BLOCK_M); + const uint32_t expert_blocks = num_m_blocks * kNumBlockNs; + const uint32_t expert_end = global_block + expert_blocks; + + while (next_assigned_block < global_block) + next_assigned_block += kNumSMs; + while (next_assigned_block < expert_end) { + const uint32_t local_block = + next_assigned_block - global_block; + const uint32_t m_block_idx = local_block / kNumBlockNs; + const uint32_t n_block_idx = + local_block - m_block_idx * kNumBlockNs; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + func(expert_idx, pool_block_offset, m_block_idx, + n_block_idx, valid_m); + next_assigned_block += kNumSMs; + } + global_block = expert_end; + pool_block_offset += num_m_blocks; + } + }; + + uint32_t stage_idx = 0; + uint32_t phase = 0; + const 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; + }; + + constexpr uint32_t kNumProducerRegisters = 40; + constexpr uint32_t kNumEpilogueRegisters = 208; + + if constexpr (!kBF16Mode) { + if (warp_idx == 0) { + cutlass::arch::warpgroup_reg_dealloc(); + for_each_block([&](const uint32_t&, const uint32_t& pool_block_offset, + const uint32_t& m_block_idx, const uint32_t&, + const uint32_t& valid_m) { + const uint32_t pool_block_idx = + pool_block_offset + m_block_idx; + #pragma unroll + for (uint32_t k_block_idx = 0; + k_block_idx < kHidden / BLOCK_K; + advance_pipeline(k_block_idx)) { + empty_barriers[stage_idx]->wait(phase ^ 1); + uint32_t m_idx = pool_block_idx * BLOCK_M; + if (!is_leader_cta) + m_idx += math::align(valid_m, 16u) / 2; + if (cute::elect_one_sync()) { + tma::copy( + &tensor_map_acts, full_barriers[stage_idx], + smem_a[stage_idx], k_block_idx * BLOCK_K, m_idx, 2); + tma::copy( + &tensor_map_acts_sf, full_barriers[stage_idx], + smem_sfa[stage_idx], + pool_block_idx * SF_BLOCK_M, k_block_idx, 2); + if (is_leader_cta) { + full_barriers[stage_idx]->arrive_and_expect_tx( + SMEM_A_SIZE_PER_STAGE * 2 + + SF_BLOCK_M * sizeof(uint32_t) * 2); + } else { + full_barriers[stage_idx]->arrive(0u); + } + } + __syncwarp(); + } + }); + } else if (warp_idx == 1) { + cutlass::arch::warpgroup_reg_dealloc(); + for_each_block([&](const uint32_t& expert_idx, const uint32_t&, + const uint32_t&, const uint32_t& n_block_idx, + const uint32_t&) { + #pragma unroll + for (uint32_t k_block_idx = 0; + k_block_idx < kHidden / BLOCK_K; + advance_pipeline(k_block_idx)) { + empty_barriers[stage_idx]->wait(phase ^ 1); + if (cute::elect_one_sync()) { + tma::copy( + &tensor_map_weights, full_barriers[stage_idx], + smem_b[stage_idx], k_block_idx * BLOCK_K, + expert_idx * 2 * kIntermediateHidden + + n_block_idx * BLOCK_N, + 2); + tma::copy( + &tensor_map_weights_sf, full_barriers[stage_idx], + smem_sfb[stage_idx], n_block_idx * BLOCK_N, + expert_idx * (kHidden / (kGranK * 4)) + + k_block_idx, + 2); + if (is_leader_cta) { + full_barriers[stage_idx]->arrive_and_expect_tx( + SMEM_B_SIZE_PER_STAGE + + BLOCK_N * sizeof(uint32_t) * 2); + } else { + full_barriers[stage_idx]->arrive(0u); + } + } + __syncwarp(); + } + }); + } else if (warp_idx == 2) { + cutlass::arch::warpgroup_reg_dealloc(); + if (is_leader_cta) { + auto instr_desc = + cute::UMMA::make_instr_desc_block_scaled< + b_dtype_t, a_dtype_t, float, cutlass::float_ue8m0_t, + UMMA_M, UMMA_N, cute::UMMA::Major::K, + cute::UMMA::Major::K>(); + auto sf_desc = mma::sm100::make_sf_desc(nullptr); + auto a_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::K, LOAD_BLOCK_M, BLOCK_K, + kSwizzleAMode>(smem_a[0], 0, 0); + auto b_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::K, LOAD_BLOCK_N, BLOCK_K, + kSwizzleBMode>(smem_b[0], 0, 0); + const uint32_t a_desc_lo = lane_idx < kNumStages + ? a_desc.lo + lane_idx * SMEM_A_SIZE_PER_STAGE / 16 + : 0; + const uint32_t b_desc_lo = lane_idx < kNumStages + ? b_desc.lo + lane_idx * SMEM_B_SIZE_PER_STAGE / 16 + : 0; + uint32_t current_iter = 0; + + for_each_block([&](const uint32_t&, const uint32_t&, + const uint32_t&, const uint32_t&, + const uint32_t& valid_m) { + mma::sm100::update_instr_desc_with_umma_n( + instr_desc, math::align(valid_m, 16u)); + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1; + tmem_empty_barriers[accum_stage]->wait( + accum_phase ^ 1); + ptx::tcgen05_after_thread_sync(); + + #pragma unroll + for (uint32_t k_block_idx = 0; + k_block_idx < kHidden / BLOCK_K; + advance_pipeline(k_block_idx)) { + full_barriers[stage_idx]->wait(phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t a_desc_base = + ptx::exchange(a_desc_lo, stage_idx); + const uint32_t b_desc_base = + ptx::exchange(b_desc_lo, stage_idx); + if (cute::elect_one_sync()) { + using utccp_t = + cute::SM100_UTCCP_4x32dp128bit_2cta; + #pragma unroll + for (uint32_t i = 0; + i < SF_BLOCK_M / + kNumUTCCPAlignedElems; + ++i) { + mma::sm100::replace_smem_desc_addr( + sf_desc, + smem_sfa[stage_idx] + + i * kNumUTCCPAlignedElems); + utccp_t::copy( + sf_desc, + kTmemStartColOfSFA + i * 4); + } + mma::sm100::replace_smem_desc_addr( + sf_desc, smem_sfb[stage_idx]); + utccp_t::copy(sf_desc, kTmemStartColOfSFB); + + #pragma unroll + for (uint32_t k = 0; + k < BLOCK_K / UMMA_K; ++k) { + const auto runtime_instr_desc = + mma::sm100:: + make_runtime_instr_desc_with_sf_id( + instr_desc, k, k); + a_desc.lo = + mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::K, + LOAD_BLOCK_M, kSwizzleAMode, + a_dtype_t>( + a_desc_base, 0, k * UMMA_K); + b_desc.lo = + mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::K, + LOAD_BLOCK_N, kSwizzleBMode, + b_dtype_t>( + b_desc_base, 0, k * UMMA_K); + ptx::SM100_MMA_MXF8F6F4_2x1SM_SS::fma( + b_desc, a_desc, + accum_stage * UMMA_N, + k_block_idx > 0 || k > 0, + runtime_instr_desc, + kTmemStartColOfSFB, + kTmemStartColOfSFA); + } + } + __syncwarp(); + + constexpr uint16_t kCTAMask = 0x3; + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + empty_barriers[stage_idx]), + kCTAMask); + if (k_block_idx == + kHidden / BLOCK_K - 1) { + cutlass::arch:: + umma_arrive_multicast_2x1SM( + reinterpret_cast( + tmem_full_barriers[ + accum_stage]), + kCTAMask); + } + __syncwarp(); + } + }); + if (current_iter > 0) { + const uint32_t last = current_iter - 1; + tmem_empty_barriers[ + last % kNumEpilogueStages] + ->wait((last / kNumEpilogueStages) & 1); + } + } + } else if (warp_idx == 3) { + cutlass::arch::warpgroup_reg_dealloc(); + } else if ( + warp_idx < + (kNumNonEpilogueThreads + + kNumEpilogueThreads) / + 32) { + cutlass::arch::warpgroup_reg_alloc(); + DG_TRAP_ONLY_DEVICE_ASSERT( + ptx::ld_shared(tmem_ptr_in_smem) == 0); + const uint32_t epilogue_warp_idx = warp_idx - 4; + uint32_t current_iter = 0; + uint32_t tma_stage_idx = 0; + + for_each_block([&](const uint32_t&, const uint32_t& pool_block_offset, + const uint32_t& m_block_idx, + const uint32_t& n_block_idx, + const uint32_t& valid_m) { + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1; + tmem_full_barriers[accum_stage]->wait(accum_phase); + ptx::tcgen05_after_thread_sync(); + + epilogue::sm100_store_cd_swap_ab< + BLOCK_M, BLOCK_N, STORE_BLOCK_M, + STORE_BLOCK_N, kSwizzleCDMode, + kNumTMAStoreStages, kNumEpilogueThreads, + GemmType::Normal, false, cd_dtype_t, + epilogue::transform::EpilogueIdentity>( + smem_cd, tma_stage_idx, + accum_stage * UMMA_N, + (pool_block_offset + m_block_idx) * BLOCK_M, + n_block_idx * BLOCK_N, 0, + math::align(valid_m, 16u), + epilogue_warp_idx, lane_idx, + tmem_empty_barriers[accum_stage], + tensor_map_output); + }); + + // The dgrad phase consumes gate/up from global memory. Drain the final + // two TMA-store stages before publishing phase completion. + if (epilogue_warp_idx == 0) + cute::tma_store_wait<0>(); + __syncwarp(); + } + } + if ( + warp_idx >= kDispatchWarpStart && + warp_idx < kDispatchWarpStart + kNumDispatchWarps) { + // The 1024-thread dgrad launch already reserves extra warps. Reuse one + // warpgroup as the third role instead of increasing the launch size. + constexpr uint32_t kNumDispatchRegisters = + kBF16Mode ? 56 : 48; + cutlass::arch::warpgroup_reg_dealloc(); + if constexpr ( + kNumRanks > 1 && + !(kBF16Mode && kDispatchInputsPrepared)) { + const uint32_t dispatch_warp_idx = + warp_idx - kDispatchWarpStart; + const uint32_t dispatch_thread_idx = + dispatch_warp_idx * 32 + lane_idx; + constexpr uint32_t kDispatchGridSyncIndex = 0; + constexpr uint32_t kDispatchDoneGridSyncIndex = 1; + constexpr uint32_t kBeforeBackwardPullBarrierTag = 4; + constexpr uint32_t kDispatchNamedBarrierIdx = 15; + + // All ranks stage their local BF16 grad-y before launch. This + // system-scope barrier publishes those stores before remote TMA. + trace_begin(3); + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumDispatchThreads, + kDispatchGridSyncIndex, + kBeforeBackwardPullBarrierTag>( + backward_workspace, backward_sym_buffer, + blockIdx.x, dispatch_thread_idx, + [=]() { + ptx::sync_aligned( + kNumDispatchThreads, + kDispatchNamedBarrierIdx); + }, + true, true); + trace_end(3); + + auto* pull_buffer = + reinterpret_cast(smem_buffer) + + dispatch_warp_idx * kHidden; + auto* pull_mbarrier = + dispatch_barriers[dispatch_warp_idx]; + uint32_t pull_mbarrier_phase = 0; + uint32_t pool_block_offset = 0; + + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + for (uint32_t token_idx = + blockIdx.x * kNumDispatchWarps + + dispatch_warp_idx; + token_idx < num_tokens; + token_idx += + kNumSMs * kNumDispatchWarps) { + const uint32_t pool_row = + pool_block_offset * BLOCK_M + + token_idx; + const auto metadata = + token_src_metadata[pool_row]; + const auto* remote_grad_y = + backward_sym_buffer.map( + backward_grad_y + + static_cast( + metadata.token_idx) * + kHidden, + metadata.rank_idx); + + if (cute::elect_one_sync()) { + ptx::tma_load_1d( + pull_buffer, remote_grad_y, + pull_mbarrier, + kHidden * sizeof(cd_dtype_t)); + } + __syncwarp(); + + if (cute::elect_one_sync()) { + const auto* remote_weight = + backward_sym_buffer.map( + backward_topk_weights + + static_cast( + metadata.token_idx) * + num_topk + + metadata.topk_idx, + metadata.rank_idx); + if (route_weights_fp32 != nullptr) { + route_weights_fp32[pool_row] = + *remote_weight; + } else { + route_weights[pool_row] = + cd_dtype_t(*remote_weight); + } + + ptx::mbarrier_arrive_and_set_tx( + pull_mbarrier, + kHidden * sizeof(cd_dtype_t)); + ptx::mbarrier_wait_and_flip_phase( + pull_mbarrier, + pull_mbarrier_phase); + ptx::tma_store_1d( + grad_y_unweighted_output + + static_cast( + pool_row) * + kHidden, + pull_buffer, + kHidden * sizeof(cd_dtype_t)); + cute::tma_store_arrive(); + ptx::tma_store_wait<0>(); + } + __syncwarp(); + } + pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } + + // Stronger than the eventual per-expert handshake: every L2 tile + // sees every dispatched row. This barrier runs concurrently with + // recompute and joins only at the phase boundary below. + trace_begin(4); + comm::grid_sync< + kNumSMs, kDispatchDoneGridSyncIndex>( + backward_workspace, blockIdx.x, + dispatch_thread_idx, + [=]() { + ptx::sync_aligned( + kNumDispatchThreads, + kDispatchNamedBarrierIdx); + }); + trace_end(4); + } + } else if (warp_idx >= 12) { + // W13 wgrad needs the exact BF16 value represented by the forward + // FP8+UE8M0 pool. Produce it while the recompute MMA is running, using + // otherwise-idle warps. Padding rows are explicitly zeroed so the + // k-grouped wgrad mainloop can round K up to 64 without reading the + // following expert. + constexpr uint32_t kNumXPoolRegisters = + kBF16Mode ? 56 : 40; + cutlass::arch::warpgroup_reg_dealloc< + kNumXPoolRegisters>(); + constexpr uint32_t kFirstXPoolWarp = 12; + constexpr uint32_t kNumXPoolThreads = + kNumThreads - kFirstXPoolWarp * 32; + const uint32_t x_thread_idx = + (warp_idx - kFirstXPoolWarp) * 32 + lane_idx; + if constexpr (!(kBF16Mode && kDispatchInputsPrepared)) { + uint32_t pool_block_offset = 0; + uint32_t global_pool_block = 0; + #pragma unroll + for (uint32_t expert_idx = 0; expert_idx < kNumExperts; + ++expert_idx) { + const uint32_t num_tokens = + static_cast(__ldg(expert_counts + expert_idx)); + const uint32_t num_blocks = + math::ceil_div(num_tokens, BLOCK_M); + for (uint32_t m_block_idx = 0; m_block_idx < num_blocks; + ++m_block_idx, ++global_pool_block) { + if (global_pool_block % kNumSMs != blockIdx.x) + continue; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + const uint32_t pool_block = + pool_block_offset + m_block_idx; + for (uint32_t linear = x_thread_idx; + linear < BLOCK_M * kHidden; + linear += kNumXPoolThreads) { + const uint32_t row = linear / kHidden; + const uint32_t col = linear - row * kHidden; + const uint32_t pool_row = + pool_block * BLOCK_M + row; + cd_dtype_t value = cd_dtype_t(0.0f); + if (row < valid_m) { + if constexpr (kBF16Mode) { + const auto metadata = + token_src_metadata[pool_row]; + value = *backward_sym_buffer.map( + backward_x + + static_cast( + metadata.token_idx) * + kHidden + + col, + metadata.rank_idx); + if constexpr (kNumRanks == 1) { + grad_y_unweighted_output[ + static_cast(pool_row) * + kHidden + + col] = + *backward_sym_buffer.map( + backward_grad_y + + static_cast( + metadata.token_idx) * + kHidden + + col, + metadata.rank_idx); + if (col == 0) { + const float weight = + *backward_sym_buffer.map( + backward_topk_weights + + static_cast( + metadata.token_idx) * + num_topk + + metadata.topk_idx, + metadata.rank_idx); + route_weights_fp32[pool_row] = + weight; + } + } + } else { + const uint32_t idx = row % BLOCK_M; + const uint32_t sf_token = + pool_block * SF_BLOCK_M + + (idx & ~127u) + + (idx & 31u) * 4 + + ((idx >> 5) & 3u); + const uint32_t sf_group = col / 128; + const uint32_t sf_byte = (col / 32) & 3u; + const uint32_t packed_sf = + acts_sf_ptr[ + sf_group * acts_sf_stride + + sf_token]; + const uint32_t exponent = + (packed_sf >> (sf_byte * 8)) & + 0xffu; + const uint32_t scale_bits = + exponent << 23; + const float scale = + *reinterpret_cast( + &scale_bits); + value = cd_dtype_t( + static_cast( + acts_ptr[ + static_cast( + pool_row) * + kHidden + + col]) * + scale); + } + } + x_pool_output[ + static_cast(pool_row) * kHidden + + col] = value; + } + } + pool_block_offset += num_blocks; + } + } + } else { + constexpr uint32_t kNumIdleRegisters = + kBF16Mode ? 56 : 24; + cutlass::arch::warpgroup_reg_dealloc< + kNumIdleRegisters>(); + } + + { + __syncthreads(); + if constexpr ( + kBF16Mode && kNumRanks == 1 && + !kDispatchInputsPrepared) { + // In single-rank BF16 mode the x-pool warps also stage grad-y. + // Their pool-block assignment is independent of the dgrad tile + // assignment, so a cluster barrier is insufficient before W2 + // dgrad starts consuming the completed expert pool. + constexpr uint32_t kLocalDispatchDoneGridSyncIndex = 1; + trace_begin(5); + comm::grid_sync< + kNumSMs, kLocalDispatchDoneGridSyncIndex>( + backward_workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + trace_end(5); + } + if constexpr ( + (kBF16Mode || + kRouteWeightMode == + RouteWeightMode::PostDown) && + !kDispatchInputsPrepared) { + uint32_t grad_pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + const uint32_t num_padded_tokens = + math::ceil_div(num_tokens, BLOCK_M) * + BLOCK_M; + for (uint64_t linear = + static_cast(blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + num_padded_tokens) * + kHidden; + linear += + static_cast(kNumSMs) * + kNumThreads) { + const uint32_t token_idx = + linear / kHidden; + const uint32_t col = + linear - + static_cast(token_idx) * + kHidden; + const uint32_t pool_row = + grad_pool_block_offset * BLOCK_M + + token_idx; + grad_ye_output[ + static_cast(pool_row) * + kHidden + + col] = + token_idx >= num_tokens + ? cd_dtype_t(0.0f) + : kRouteWeightMode == + RouteWeightMode::PostDown + ? cd_dtype_t( + static_cast( + grad_y_unweighted_output[ + static_cast( + pool_row) * + kHidden + + col]) * + (route_weights_fp32 != nullptr + ? route_weights_fp32[ + pool_row] + : static_cast( + route_weights[ + pool_row]))) + : grad_y_unweighted_output[ + static_cast( + pool_row) * + kHidden + + col]; + } + grad_pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } + if constexpr (kBF16Mode) { + constexpr uint32_t + kW2GradInputGridSyncIndex = 0; + trace_begin(6); + comm::grid_sync< + kNumSMs, + kW2GradInputGridSyncIndex>( + backward_workspace, blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + trace_end(6); + } else { + // Standalone MXFP4 backward has no symmetric Workspace. + // Reuse its launch-epoch grid state for the same publication + // barrier before W2 dgrad consumes weighted grad-y. + full_grid_phase_barrier(6); + } + } + if constexpr (kDirectRemoteGradX) { + if constexpr (kNumRanks > 1) { + // backward_grad_y aliases combine plane zero. All ranks must + // finish remotely pulling it before any W13 dgrad epilogue + // reuses the combine planes for direct grad-x writes. + if constexpr ( + !(kBF16Mode && kDispatchInputsPrepared)) { + constexpr uint32_t + kBeforeDirectGradXGridSyncIndex = 2; + constexpr uint32_t + kBeforeDirectGradXBarrierTag = 7; + trace_begin(7); + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, + kBeforeDirectGradXGridSyncIndex, + kBeforeDirectGradXBarrierTag>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + trace_end(7); + } + + } + + if constexpr ( + kCombineOrderMode == + CombineOrderMode::FixedTopK) { + // FixedTopK consumes every physical slot, including invalid + // routes. Clear all slot planes only after all grad-y pulls + // have completed, then publish the clear before any direct + // remote stores. This also makes repeated and single-rank + // calls independent of stale valid routes. + auto* combine_buffer = + const_cast(backward_grad_y); + const uint64_t num_plane_values = + static_cast(num_topk) * + backward_workspace.num_max_tokens_per_rank * + kHidden; + for (uint64_t linear = + static_cast(blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < num_plane_values; + linear += + static_cast(kNumSMs) * + kNumThreads) { + combine_buffer[linear] = + cd_dtype_t(0.0f); + } + + if constexpr (kNumRanks > 1) { + // Do not let a rank remotely write direct grad-x until + // every destination has finished clearing its local + // slot planes. + constexpr uint32_t + kAfterGradYClearGridSyncIndex = 3; + constexpr uint32_t + kAfterGradYClearBarrierTag = 8; + trace_begin(8); + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, + kAfterGradYClearGridSyncIndex, + kAfterGradYClearBarrierTag>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + trace_end(8); + } + } + } + if constexpr (!kBF16Mode) { + if (warp_idx >= kDispatchWarpStart && + warp_idx < + kDispatchWarpStart + + kNumDispatchWarps) { + // Dispatch used 48 registers; transition down to the common + // dgrad epilogue budget with dealloc, not reg_alloc + // (allocating a lower count is illegal on SM100). + cutlass::arch::warpgroup_reg_dealloc<40>(); + } else if (warp_idx >= kDispatchWarpStart) { + cutlass::arch::warpgroup_reg_alloc<40>(); + } + } + const auto for_each_dgrad_block = [&](const auto& func) { + uint32_t next_assigned_block = blockIdx.x; + uint32_t global_block = 0; + uint32_t pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + const uint32_t num_m_blocks = + math::ceil_div(num_tokens, BLOCK_M); + const uint32_t expert_blocks = + num_m_blocks * kNumDgradBlockNs; + const uint32_t expert_end = + global_block + expert_blocks; + + while (next_assigned_block < global_block) + next_assigned_block += kNumSMs; + while (next_assigned_block < expert_end) { + const uint32_t local_block = + next_assigned_block - global_block; + const uint32_t m_block_idx = + local_block / kNumDgradBlockNs; + const uint32_t n_block_idx = + local_block - + m_block_idx * kNumDgradBlockNs; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * BLOCK_M, + BLOCK_M); + func( + expert_idx, pool_block_offset, + m_block_idx, n_block_idx, valid_m); + next_assigned_block += kNumSMs; + } + global_block = expert_end; + pool_block_offset += num_m_blocks; + } + }; + + trace_begin(9); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(9); + + trace_begin(10); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(10); + + // Reinitialize the drained pipelines in-place. The FP32 accumulator + // columns are phase-aliased; dgrad does not need the SFA/SFB columns. + if (warp_idx == 0 && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kNumStages; ++i) { + // A and transposed-W2 TMA warps in both CTAs. + full_barriers[i]->init(4); + empty_barriers[i]->init(1); + } + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { + tmem_full_barriers[i]->init(1); + tmem_empty_barriers[i]->init( + 2 * kNumDgradEpilogueThreads); + } + cutlass::arch::fence_barrier_init(); + } + trace_begin(11); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(11); + + stage_idx = 0; + phase = 0; + if (warp_idx == 0) { + // BF16 grad_y producer. The 2-SM TMA instruction writes each + // CTA's half-M operand and completes transactions on CTA0's + // cluster barrier. + for_each_dgrad_block( + [&](const uint32_t&, const uint32_t& pool_block_offset, + const uint32_t& m_block_idx, const uint32_t&, + const uint32_t& valid_m) { + const uint32_t pool_block_idx = + pool_block_offset + m_block_idx; + #pragma unroll 1 + for (uint32_t k_block_idx = 0; + k_block_idx < kHidden / DGRAD_BLOCK_K; + advance_pipeline(k_block_idx)) { + empty_barriers[stage_idx]->wait(phase ^ 1); + uint32_t m_idx = pool_block_idx * BLOCK_M; + if (!is_leader_cta) + m_idx += math::align(valid_m, 16u) / 2; + if (cute::elect_one_sync()) { + tma::copy< + DGRAD_BLOCK_K, LOAD_BLOCK_M, + DGRAD_BLOCK_K * sizeof(cd_dtype_t), + cd_dtype_t>( + &tensor_map_grad_ye, + full_barriers[stage_idx], + smem_dgrad_a[stage_idx], + k_block_idx * DGRAD_BLOCK_K, + m_idx, 2); + if (is_leader_cta) { + full_barriers[stage_idx] + ->arrive_and_expect_tx( + SMEM_A_SIZE_PER_STAGE * 2); + } else { + full_barriers[stage_idx]->arrive(0u); + } + } + __syncwarp(); + } + }); + } else if (warp_idx == 1) { + // Load the in-kernel dequantized transposed W2 workspace. It is + // shared by all M tiles for this expert instead of reconverting + // the same packed weights for every token block. + for_each_dgrad_block( + [&](const uint32_t& expert_idx, const uint32_t&, + const uint32_t&, const uint32_t& n_block_idx, + const uint32_t&) { + #pragma unroll 1 + for (uint32_t k_block_idx = 0; + k_block_idx < kHidden / DGRAD_BLOCK_K; + advance_pipeline(k_block_idx)) { + const uint32_t weight_tile_idx = + (expert_idx * + (kHidden / DGRAD_BLOCK_K) + + k_block_idx) * + kNumDgradBlockNs + + n_block_idx; + if constexpr (!kBF16Mode) { + while (ptx::ld_acq( + weight_tile_states + + weight_tile_idx) != + launch_epoch) { + } + } + constexpr bool weight_tile_ready = true; + empty_barriers[stage_idx]->wait(phase ^ 1); + if (weight_tile_ready) { + if (cute::elect_one_sync()) { + tma::copy< + LOAD_BLOCK_N, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * + sizeof( + dgrad_b_dtype_t), + dgrad_b_dtype_t>( + &tensor_map_w2_dequant, + full_barriers[stage_idx], + smem_dgrad_b[stage_idx], + n_block_idx * + BLOCK_N, + expert_idx * + kHidden + + k_block_idx * + DGRAD_BLOCK_K, + 2); + if (is_leader_cta) { + full_barriers[stage_idx] + ->arrive_and_expect_tx( + SMEM_B_SIZE_PER_STAGE * + 2); + } else { + full_barriers[stage_idx] + ->arrive(0u); + } + } + } else { + constexpr uint32_t + kPairsPerTile = + LOAD_BLOCK_N * + DGRAD_BLOCK_K / 2; + auto* smem_b_bytes = + reinterpret_cast( + smem_dgrad_b[stage_idx]); + for (uint32_t pair_idx = lane_idx; + pair_idx < kPairsPerTile; + pair_idx += 32) { + const uint32_t local_k = + pair_idx / + (LOAD_BLOCK_N / 2); + const uint32_t + local_n_pair = + pair_idx % + (LOAD_BLOCK_N / 2); + const uint32_t global_k = + k_block_idx * + DGRAD_BLOCK_K + + local_k; + const uint32_t + global_n_pair = + n_block_idx * + (LOAD_BLOCK_N / + 2) + + local_n_pair; + const uint8_t packed = + static_cast( + __ldg( + w2_weights + + (static_cast< + uint64_t>( + expert_idx) * + kHidden + + global_k) * + (kIntermediateHidden / + 2) + + global_n_pair)); + const float scale = __ldg( + w2_scales + + (static_cast( + expert_idx) * + kHidden + + global_k) * + (kIntermediateHidden / + 32) + + n_block_idx * + (LOAD_BLOCK_N / 32) + + (local_n_pair * 2) / + 32); + uint32_t fp16x2; + asm volatile( + "{\n" + ".reg .b8 fp4;\n" + ".reg .b8 unused1, unused2, unused3;\n" + "mov.b32 {fp4, unused1, unused2, unused3}, %1;\n" + "cvt.rn.f16x2.e2m1x2 %0, fp4;\n" + "}\n" + : "=r"(fp16x2) + : "r"( + static_cast< + uint32_t>( + packed))); + auto value_pair = + *reinterpret_cast< + __half2*>(&fp16x2); + value_pair = __hmul2( + value_pair, + __float2half2_rn( + scale)); + const auto + value_pair_bf16 = + __float22bfloat162_rn( + __half22float2( + value_pair)); + const uint32_t + scaled_pair = + *reinterpret_cast< + const uint32_t*>( + &value_pair_bf16); + #pragma unroll + for (uint32_t i = 0; + i < 2; ++i) { + const uint32_t local_n = + local_n_pair * 2 + i; + const uint32_t row = + local_n & 7; + const uint32_t col_byte = + local_k * + sizeof( + dgrad_b_dtype_t); + const uint32_t + byte_offset = + (local_n >> 3) * + 8 * 128 + + row * 128 + + ((col_byte >> 4) ^ + row) * + 16 + + (col_byte & 15); + *reinterpret_cast< + uint16_t*>( + smem_b_bytes + + byte_offset) = + static_cast< + uint16_t>( + scaled_pair >> + (i * 16)); + } + } + cutlass::arch:: + fence_view_async_shared(); + if (cute::elect_one_sync()) + full_barriers[stage_idx] + ->arrive(0u); + } + __syncwarp(); + } + }); + } else if (warp_idx == 2) { + if (is_leader_cta) { + auto instr_desc = + cute::UMMA::make_instr_desc< + dgrad_b_dtype_t, cd_dtype_t, float, + UMMA_M, UMMA_N, + cute::UMMA::Major::MN, + cute::UMMA::Major::K>(); + auto a_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::K, LOAD_BLOCK_M, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * sizeof(cd_dtype_t)>( + smem_dgrad_a[0], 0, 0); + auto b_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::MN, LOAD_BLOCK_N, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * sizeof(dgrad_b_dtype_t)>( + smem_dgrad_b[0], 0, 0); + const uint32_t a_desc_lo = lane_idx < kNumStages + ? a_desc.lo + + lane_idx * SMEM_A_SIZE_PER_STAGE / 16 + : 0; + const uint32_t b_desc_lo = lane_idx < kNumStages + ? b_desc.lo + + lane_idx * SMEM_B_SIZE_PER_STAGE / 16 + : 0; + uint32_t current_iter = 0; + + for_each_dgrad_block( + [&](const uint32_t&, const uint32_t&, + const uint32_t&, const uint32_t&, + const uint32_t& valid_m) { + mma::sm100::update_instr_desc_with_umma_n( + instr_desc, + math::align(valid_m, 16u)); + const auto runtime_instr_desc = + cute::UMMA::make_runtime_instr_desc( + instr_desc); + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / + kNumEpilogueStages) & + 1; + tmem_empty_barriers[accum_stage]->wait( + accum_phase ^ 1); + ptx::tcgen05_after_thread_sync(); + + #pragma unroll 1 + for (uint32_t k_block_idx = 0; + k_block_idx < + kHidden / DGRAD_BLOCK_K; + advance_pipeline(k_block_idx)) { + full_barriers[stage_idx]->wait(phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t a_desc_base = + ptx::exchange( + a_desc_lo, stage_idx); + const uint32_t b_desc_base = + ptx::exchange( + b_desc_lo, stage_idx); + if (cute::elect_one_sync()) { + #pragma unroll + for (uint32_t k = 0; + k < + DGRAD_BLOCK_K / + DGRAD_UMMA_K; + ++k) { + a_desc.lo = + mma::sm100:: + advance_umma_desc_lo< + cute::UMMA::Major::K, + LOAD_BLOCK_M, + DGRAD_BLOCK_K * + sizeof( + cd_dtype_t), + cd_dtype_t>( + a_desc_base, 0, + k * + DGRAD_UMMA_K); + b_desc.lo = + mma::sm100:: + advance_umma_desc_lo< + cute::UMMA::Major::MN, + LOAD_BLOCK_N, + DGRAD_BLOCK_K * + sizeof( + dgrad_b_dtype_t), + dgrad_b_dtype_t>( + b_desc_base, 0, + k * + DGRAD_UMMA_K); + ptx:: + SM100_MMA_F16BF16_2x1SM_SS:: + fma( + b_desc, a_desc, + accum_stage * + UMMA_N, + k_block_idx > 0 || + k > 0, + runtime_instr_desc); + } + } + __syncwarp(); + constexpr uint16_t kCTAMask = 0x3; + cutlass::arch:: + umma_arrive_multicast_2x1SM( + reinterpret_cast( + empty_barriers[ + stage_idx]), + kCTAMask); + if (k_block_idx == + kHidden / DGRAD_BLOCK_K - 1) { + cutlass::arch:: + umma_arrive_multicast_2x1SM( + reinterpret_cast( + tmem_full_barriers[ + accum_stage]), + kCTAMask); + } + __syncwarp(); + } + }); + if (current_iter > 0) { + const uint32_t last = current_iter - 1; + tmem_empty_barriers[ + last % kNumEpilogueStages] + ->wait( + (last / kNumEpilogueStages) & 1); + } + } + } else if (warp_idx >= 4) { + const uint32_t epilogue_warp_idx = warp_idx - 4; + const uint32_t epilogue_thread_idx = + epilogue_warp_idx * 32 + lane_idx; + uint32_t current_iter = 0; + + for_each_dgrad_block( + [&](const uint32_t&, const uint32_t& pool_block_offset, + const uint32_t& m_block_idx, + const uint32_t& n_block_idx, + const uint32_t& valid_m) { + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / + kNumEpilogueStages) & + 1; + tmem_full_barriers[accum_stage]->wait( + accum_phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t effective_m = + math::align(valid_m, 16u); + + for (uint32_t s = 0; + s < effective_m / STORE_BLOCK_M; ++s) { + cutlass::arch::NamedBarrier::sync( + kNumDgradEpilogueThreads, 0); + if (epilogue_warp_idx < + kNumEpilogueThreads / 32) { + #pragma unroll + for (uint32_t i = 0; + i < STORE_BLOCK_M / 8; ++i) { + const uint32_t tmem_addr = + accum_stage * UMMA_N + + s * STORE_BLOCK_M + i * 8; + uint32_t values[8]; + cute::SM100_TMEM_LOAD_16dp256b1x:: + copy( + tmem_addr, values[0], + values[1], values[2], + values[3]); + cute::SM100_TMEM_LOAD_16dp256b1x:: + copy( + tmem_addr | 0x00100000, + values[4], values[5], + values[6], values[7]); + cutlass::arch:: + fence_view_async_tmem_load(); + + constexpr uint32_t kBankBytes = 16; + const uint32_t outer_atom = + (epilogue_warp_idx / 2) * + STORE_BLOCK_M * 128; + const uint32_t inner_atom = + i * 8 * 128; + const uint32_t row = lane_idx % 8; + const uint32_t col = + (epilogue_warp_idx % 2) * 4 + + lane_idx / 8; + auto* smem_ptr = + reinterpret_cast( + smem_cd[0]) + + outer_atom + inner_atom + + row * (kBankBytes * 8) + + (col ^ row) * kBankBytes; + ptx::SM90_U32x4_STSM_T::copy( + math::cast_into_bf16_and_pack( + values[0], values[1]), + math::cast_into_bf16_and_pack( + values[2], values[3]), + math::cast_into_bf16_and_pack( + values[4], values[5]), + math::cast_into_bf16_and_pack( + values[6], values[7]), + smem_ptr); + } + } + cutlass::arch::NamedBarrier::sync( + kNumDgradEpilogueThreads, 0); + + #pragma unroll + for (uint32_t linear = + epilogue_thread_idx; + linear < + STORE_BLOCK_M * BLOCK_N; + linear += + kNumDgradEpilogueThreads) { + const uint32_t row = + linear / BLOCK_N; + const uint32_t n = + linear - row * BLOCK_N; + const uint32_t local_m = + s * STORE_BLOCK_M + row; + if (local_m >= valid_m) + continue; + + const uint32_t n_atom = n / 64; + const uint32_t n_in_atom = + n - n_atom * 64; + const uint32_t row_in_atom = + row & 7; + const uint32_t smem_byte_offset = + n_atom * + STORE_BLOCK_M * 128 + + (row >> 3) * 8 * 128 + + row_in_atom * 128 + + ((n_in_atom >> 3) ^ + row_in_atom) * + 16 + + (n_in_atom & 7) * + sizeof(cd_dtype_t); + const cd_dtype_t grad_h_w2 = + *reinterpret_cast< + cd_dtype_t*>( + reinterpret_cast< + uint8_t*>( + smem_cd[0]) + + smem_byte_offset); + const uint32_t pool_row = + (pool_block_offset + + m_block_idx) * + BLOCK_M + + local_m; + const uint32_t hidden_col = + n_block_idx * BLOCK_N + n; + const float route_weight = + route_weights_fp32 != nullptr + ? route_weights_fp32[pool_row] + : static_cast( + route_weights[pool_row]); + const cd_dtype_t grad_h_bf16 = + kRouteWeightMode == + RouteWeightMode::PostDown + ? grad_h_w2 + : cd_dtype_t( + static_cast( + grad_h_w2) * + route_weight); + const float grad_h = + static_cast(grad_h_bf16); + // This is the W2 dgrad output before any pre-down + // route multiplication. In post-down mode its GEMM + // input was already weighted BF16 grad-y. + grad_h_output[ + static_cast( + pool_row) * + kIntermediateHidden + + hidden_col] = + grad_h_w2; + const uint32_t chunk = + hidden_col / 8; + const uint32_t in_chunk = + hidden_col & 7; + const uint32_t gate_col = + kBF16Mode + ? hidden_col + : chunk * 16 + in_chunk; + const uint32_t up_col = + kBF16Mode + ? kIntermediateHidden + + hidden_col + : gate_col + 8; + const float gate_unclamped = + static_cast( + gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + gate_col]); + const float up_unclamped = + static_cast( + gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + up_col]); + + const bool has_activation_clamp = + kBF16Mode + ? activation_limit != + cute::numeric_limits< + float>::infinity() + : activation_limit > 0.0f; + const bool gate_in_range = + !has_activation_clamp || + gate_unclamped <= + activation_limit; + const bool up_in_range = + !has_activation_clamp || + (up_unclamped >= + -activation_limit && + up_unclamped <= + activation_limit); + const float gate = + has_activation_clamp + ? cute::min( + gate_unclamped, + activation_limit) + : gate_unclamped; + const float up = + has_activation_clamp + ? cute::min( + cute::max( + up_unclamped, + -activation_limit), + activation_limit) + : up_unclamped; + float z; + float dz_dgate; + if constexpr ( + kActivationType == + ActivationType::GeGLU) { + constexpr float kAlpha = + 1.5957691216057308f; + constexpr float kBeta = 0.044715f; + // Python evaluates 3.0 * beta in FP64 before + // converting the scalar to FP32. Multiplying + // the already-rounded kBeta by 3.0f is one ULP + // lower and changes BF16 ties in GeGLU dgate. + constexpr float kThreeBeta = 0.134145f; + const float gate_sq = + __fmul_rn(gate, gate); + z = __fmul_rn( + __fmul_rn(kAlpha, gate), + __fadd_rn( + 1.0f, + __fmul_rn( + kBeta, gate_sq))); + dz_dgate = __fmul_rn( + kAlpha, + __fadd_rn( + 1.0f, + __fmul_rn( + kThreeBeta, + gate_sq))); + } else { + z = gate; + dz_dgate = 1.0f; + } + const float neg_exp = + !kBF16Mode || kFastMath + ? __expf(-z) + : expf(-z); + const float denom = + __fadd_rn(1.0f, neg_exp); + const float sig = + 1.0f / denom; + cd_dtype_t h_act_bf16; + cd_dtype_t grad_gate_bf16; + cd_dtype_t grad_up_bf16; + if constexpr ( + kBF16Mode && + kActivationType == + ActivationType::SwiGLU) { + if (!has_activation_clamp) { + // Native grouped experts materialize + // BF16 SiLU, BF16 SiLU*up, and BF16 + // grad_h*up before aten::silu_backward. + const cd_dtype_t silu_bf16 = + cd_dtype_t( + gate / denom); + h_act_bf16 = + cd_dtype_t( + __fmul_rn( + static_cast< + float>( + silu_bf16), + up)); + const cd_dtype_t + grad_silu_bf16 = + cd_dtype_t( + __fmul_rn( + grad_h, + up)); + const float + one_minus_sig = + __fsub_rn( + 1.0f, sig); + const float + silu_inner = + __fadd_rn( + 1.0f, + __fmul_rn( + gate, + one_minus_sig)); + const float + grad_silu_sig = + __fmul_rn( + static_cast< + float>( + grad_silu_bf16), + sig); + grad_gate_bf16 = + cd_dtype_t( + __fmul_rn( + grad_silu_sig, + silu_inner)); + grad_up_bf16 = + cd_dtype_t( + __fmul_rn( + grad_h, + static_cast< + float>( + silu_bf16))); + } else { + const float + activated_gate = + __fmul_rn( + gate, sig); + h_act_bf16 = + cd_dtype_t( + __fmul_rn( + activated_gate, + up)); + const float + one_minus_sig = + __fsub_rn( + 1.0f, sig); + const float gate_sig = + __fmul_rn( + gate, sig); + const float + activation_grad = + __fadd_rn( + sig, + __fmul_rn( + __fmul_rn( + gate_sig, + one_minus_sig), + dz_dgate)); + grad_gate_bf16 = + cd_dtype_t( + gate_in_range + ? __fmul_rn( + __fmul_rn( + grad_h, + up), + activation_grad) + : 0.0f); + grad_up_bf16 = + cd_dtype_t( + up_in_range + ? __fmul_rn( + grad_h, + activated_gate) + : 0.0f); + } + } else { + const float activated_gate = + __fmul_rn(gate, sig); + h_act_bf16 = + cd_dtype_t( + __fmul_rn( + activated_gate, up)); + const float one_minus_sig = + __fsub_rn(1.0f, sig); + const float gate_sig = + __fmul_rn(gate, sig); + const float activation_grad = + __fadd_rn( + sig, + __fmul_rn( + __fmul_rn( + gate_sig, + one_minus_sig), + dz_dgate)); + grad_gate_bf16 = + cd_dtype_t( + gate_in_range + ? __fmul_rn( + __fmul_rn( + grad_h, + up), + activation_grad) + : 0.0f); + grad_up_bf16 = + cd_dtype_t( + up_in_range + ? __fmul_rn( + grad_h, + activated_gate) + : 0.0f); + } + h_act_output[ + static_cast( + pool_row) * + kIntermediateHidden + + hidden_col] = + h_act_bf16; + // PRE_DOWN may phase-alias h_act and h_weighted. + // Preserve unweighted h until its route-gradient + // reduction, then overwrite it in a later phase. + if (!( + kBF16Mode && + kRouteWeightMode == + RouteWeightMode::PreDown && + h_act_output == + h_weighted_output)) { + h_weighted_output[ + static_cast( + pool_row) * + kIntermediateHidden + + hidden_col] = + kRouteWeightMode == + RouteWeightMode::PostDown + ? h_act_bf16 + : cd_dtype_t( + static_cast( + h_act_bf16) * + route_weight); + } + grad_gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + hidden_col] = + grad_gate_bf16; + grad_gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + kIntermediateHidden + + hidden_col] = + grad_up_bf16; + } + } + ptx::tcgen05_before_thread_sync(); + tmem_empty_barriers[accum_stage]->arrive(0u); + }); + + } + + __syncthreads(); + if constexpr (kCompileW13Dgrad) { + // W13 dgrad consumes grad_gate_up rows produced by every CTA in + // the preceding L2-dgrad/SwiGLU phase. Cluster synchronization is + // insufficient here: an early cluster can otherwise read rows + // whose owning cluster has not stored them yet. + full_grid_phase_barrier(12); + + if constexpr (kBF16Mode) { + // In phase-ordered mode these outputs may still contain the + // forward gate values or reverse-dispatched grad-y in every + // row that the active activation tiles did not visit. Clear + // per-expert block padding and the unused capacity tail only + // after all active gate reads have completed. + uint32_t padding_pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg( + expert_counts + + expert_idx)); + const uint32_t num_blocks = + math::ceil_div( + num_tokens, BLOCK_M); + const uint32_t num_padded_tokens = + num_blocks * BLOCK_M; + const uint32_t padding_rows = + num_padded_tokens - num_tokens; + for (uint64_t linear = + static_cast( + blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + padding_rows) * + (2 * + kIntermediateHidden); + linear += + static_cast( + kNumSMs) * + kNumThreads) { + const uint32_t row = + linear / + (2 * kIntermediateHidden); + const uint32_t col = + linear - + static_cast(row) * + (2 * + kIntermediateHidden); + const uint32_t pool_row = + padding_pool_block_offset * + BLOCK_M + + num_tokens + row; + grad_gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + col] = + cd_dtype_t(0.0f); + } + for (uint64_t linear = + static_cast( + blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + padding_rows) * + kIntermediateHidden; + linear += + static_cast( + kNumSMs) * + kNumThreads) { + const uint32_t row = + linear / + kIntermediateHidden; + const uint32_t col = + linear - + static_cast(row) * + kIntermediateHidden; + const uint32_t pool_row = + padding_pool_block_offset * + BLOCK_M + + num_tokens + row; + h_weighted_output[ + static_cast( + pool_row) * + kIntermediateHidden + + col] = + cd_dtype_t(0.0f); + } + padding_pool_block_offset += + num_blocks; + } + const uint32_t capacity_tail_start = + padding_pool_block_offset * BLOCK_M; + const uint32_t capacity_tail_rows = + num_pool_rows - capacity_tail_start; + for (uint64_t linear = + static_cast( + blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + capacity_tail_rows) * + (2 * + kIntermediateHidden); + linear += + static_cast( + kNumSMs) * + kNumThreads) { + grad_gate_up_output[ + static_cast( + capacity_tail_start) * + (2 * + kIntermediateHidden) + + linear] = + cd_dtype_t(0.0f); + } + for (uint64_t linear = + static_cast( + blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + capacity_tail_rows) * + kIntermediateHidden; + linear += + static_cast( + kNumSMs) * + kNumThreads) { + h_weighted_output[ + static_cast( + capacity_tail_start) * + kIntermediateHidden + + linear] = + cd_dtype_t(0.0f); + } + full_grid_phase_barrier(13); + } + + if constexpr ( + kComputeRouteGrad && !kInputsPrepared) { + // The activation epilogue spans multiple N-tile CTAs. Reduce + // each route term only after all tiles are visible so the + // router gradient has a fixed FP32 summation order instead of + // depending on cross-CTA atomic arrival order. + constexpr uint32_t kRouteColumns = + kRouteWeightMode == + RouteWeightMode::PostDown + ? kHidden + : kIntermediateHidden; + // FireTitan's POST_DOWN path uses Triton's tl.sum with a + // power-of-two BLOCK_H and BLOCK_H / 256 warps (clamped to + // [4, 32]). For the production hidden sizes this gives 2, 4, + // or 8 elements per thread. Preserve that exact logical + // layout; changing the columns assigned to a lane changes the + // FP32 reduction result. + constexpr uint32_t kTritonRouteBlockH = [] { + uint32_t value = 1; + while (value < kRouteColumns && value < 8192) + value <<= 1; + return value; + }(); + constexpr uint32_t kTritonRouteNumWarps = [] { + uint32_t value = + kTritonRouteBlockH / 256; + value = value < 4 ? 4 : value; + return value > 32 ? 32 : value; + }(); + constexpr uint32_t kTritonRouteThreads = + kTritonRouteNumWarps * 32; + constexpr uint32_t + kTritonRouteValuesPerThread = + kTritonRouteBlockH / + kTritonRouteThreads; + DG_STATIC_ASSERT( + kTritonRouteValuesPerThread == 2 || + kTritonRouteValuesPerThread == 4 || + kTritonRouteValuesPerThread == 8, + "Unsupported Triton route reduction width"); + constexpr uint32_t kRouteInputPow2 = [] { + uint32_t value = 1; + const uint32_t vectorized_columns = + kRouteColumns / 4; + while (value < 512 && + (value << 1) <= + vectorized_columns) + value <<= 1; + return value; + }(); + auto* route_lane_sums = + reinterpret_cast(smem_gemm_base); + auto* route_control = + reinterpret_cast(smem_gemm_base); + if (threadIdx.x == 0) { + uint32_t total_route_rows = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; + ++expert_idx) { + total_route_rows += + static_cast( + __ldg( + expert_counts + + expert_idx)); + } + route_control[0] = + total_route_rows; + } + __syncthreads(); + const uint32_t total_route_rows = + route_control[0]; + const uint32_t route_output_pow2 = + total_route_rows > 0 + ? 1u << (31 - __clz( + total_route_rows)) + : 1u; + constexpr uint32_t + kInitialRouteGroupThreads = + cute::min( + kRouteInputPow2, 32u); + const uint32_t route_block_height = + cute::min( + route_output_pow2, + 512u / + kInitialRouteGroupThreads); + const uint32_t route_group_threads = + cute::min( + kRouteInputPow2, + 512u / route_block_height); + const uint32_t num_route_groups_per_cta = + kNumThreads / route_group_threads; + const uint32_t route_group_idx = + threadIdx.x / route_group_threads; + const uint32_t route_group_lane_idx = + threadIdx.x & + (route_group_threads - 1); + const uint32_t global_route_group = + blockIdx.x * + num_route_groups_per_cta + + route_group_idx; + const uint32_t num_route_groups = + kNumSMs * + num_route_groups_per_cta; + uint32_t route_pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + for (uint32_t token_idx = global_route_group; + token_idx < num_tokens; + token_idx += num_route_groups) { + const uint32_t pool_row = + route_pool_block_offset * BLOCK_M + + token_idx; + float grad_route = 0.0f; + if constexpr (false) { + float grad_y[ + kTritonRouteValuesPerThread]; + float down[ + kTritonRouteValuesPerThread]; + #pragma unroll + for (uint32_t i = 0; + i < + kTritonRouteValuesPerThread; + ++i) { + const uint32_t col = + route_group_lane_idx + + i * kTritonRouteThreads; + grad_y[i] = + col < kHidden + ? static_cast( + grad_y_unweighted_output[ + static_cast( + pool_row) * + kHidden + + col]) + : 0.0f; + down[i] = + col < kHidden + ? static_cast( + down_unweighted_output[ + static_cast( + pool_row) * + kHidden + + col]) + : 0.0f; + } + + if constexpr ( + kTritonRouteValuesPerThread == + 2) { + grad_route = __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[1], down[1])); + } else if constexpr ( + kTritonRouteValuesPerThread == + 4) { + const float even = + __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], + down[2])); + const float odd = + __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], + down[3])); + grad_route = + __fadd_rn(even, odd); + } else { + const float pair_02 = + __fmaf_rn( + grad_y[0], down[0], + __fmul_rn( + grad_y[2], + down[2])); + const float pair_13 = + __fmaf_rn( + grad_y[1], down[1], + __fmul_rn( + grad_y[3], + down[3])); + const float pair_46 = + __fmaf_rn( + grad_y[4], down[4], + __fmul_rn( + grad_y[6], + down[6])); + const float pair_57 = + __fmaf_rn( + grad_y[5], down[5], + __fmul_rn( + grad_y[7], + down[7])); + grad_route = __fadd_rn( + __fadd_rn( + pair_02, pair_46), + __fadd_rn( + pair_13, pair_57)); + } + + // Triton first performs a butterfly reduction + // within each physical warp. + #pragma unroll + for (uint32_t offset = 16; + offset > 0; + offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, + offset)); + } + + const uint32_t warp_in_group = + route_group_lane_idx / 32; + const uint32_t lane_in_warp = + route_group_lane_idx & 31; + if (lane_in_warp == 0) { + route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + warp_in_group] = + grad_route; + } + ptx::sync_aligned( + kTritonRouteThreads, + route_group_idx); + + // Triton loads the power-of-two set of warp + // partials into the first warp and reduces it with + // the same butterfly tree. + if (warp_in_group == 0) { + grad_route = + route_lane_sums[ + route_group_idx * + kTritonRouteNumWarps + + (lane_in_warp & + (kTritonRouteNumWarps - + 1))]; + #pragma unroll + for (uint32_t offset = + kTritonRouteNumWarps / + 2; + offset > 0; + offset >>= 1) { + grad_route = __fadd_rn( + grad_route, + __shfl_xor_sync( + 0xffffffff, + grad_route, + offset)); + } + } + } else { + float lane_sums[4] = { + 0.0f, 0.0f, 0.0f, 0.0f}; + if constexpr ( + kRouteWeightMode == + RouteWeightMode::PostDown) { + for (uint32_t col_base = + route_group_lane_idx * + 4; + col_base < kHidden; + col_base += + route_group_threads * + 4) { + #pragma unroll + for (uint32_t i = 0; i < 4; + ++i) { + const uint32_t col = + col_base + i; + const float grad_y = + static_cast( + grad_y_unweighted_output[ + static_cast< + uint64_t>( + pool_row) * + kHidden + + col]); + const float down = + static_cast( + down_unweighted_output[ + static_cast< + uint64_t>( + pool_row) * + kHidden + + col]); + lane_sums[i] = + __fadd_rn( + lane_sums[i], + __fmul_rn( + grad_y, + down)); + } + } + } else { + for (uint32_t col_base = + route_group_lane_idx * + 4; + col_base < + kIntermediateHidden; + col_base += + route_group_threads * + 4) { + #pragma unroll + for (uint32_t i = 0; i < 4; + ++i) { + const uint32_t col = + col_base + i; + const float grad_h = + static_cast( + grad_h_output[ + static_cast< + uint64_t>( + pool_row) * + kIntermediateHidden + + col]); + const float h_act = + static_cast( + h_act_output[ + static_cast< + uint64_t>( + pool_row) * + kIntermediateHidden + + col]); + lane_sums[i] = + __fadd_rn( + lane_sums[i], + __fmul_rn( + grad_h, + h_act)); + } + } + } + grad_route = __fadd_rn( + __fadd_rn( + lane_sums[0], + lane_sums[1]), + lane_sums[2]); + grad_route = __fadd_rn( + grad_route, lane_sums[3]); + + route_lane_sums[threadIdx.x] = + grad_route; + if (route_group_threads > 32) { + for (uint32_t offset = + route_group_threads / + 2; + offset >= 32; + offset >>= 1) { + ptx::sync_aligned( + route_group_threads, + route_group_idx); + if (route_group_lane_idx < + offset) { + grad_route = + __fadd_rn( + grad_route, + route_lane_sums[ + threadIdx.x + + offset]); + route_lane_sums[ + threadIdx.x] = + grad_route; + } + } + } + if (route_group_lane_idx < 32) { + #pragma unroll + for (uint32_t offset = 16; + offset > 0; + offset >>= 1) { + grad_route = + __fadd_rn( + grad_route, + __shfl_down_sync( + 0xffffffff, + grad_route, + offset)); + } + } + } + if (route_group_lane_idx == 0) { + grad_route_output[pool_row] = + grad_route; + if (backward_grad_route != nullptr) { + const auto metadata = + token_src_metadata[pool_row]; + auto* remote_grad_route = + backward_sym_buffer.map( + backward_grad_route + + static_cast( + metadata.token_idx) * + num_topk + + metadata.topk_idx, + metadata.rank_idx); + *remote_grad_route = grad_route; + } + } + if (route_group_threads > 32) { + ptx::sync_aligned( + route_group_threads, + route_group_idx); + } else { + __syncwarp(); + } + } + route_pool_block_offset += + math::ceil_div(num_tokens, BLOCK_M); + } + } + if constexpr ( + kBF16Mode && + kRouteWeightMode == RouteWeightMode::PreDown && + !kInputsPrepared) { + if (h_act_output == h_weighted_output) { + // Every route reduction must consume unweighted h before + // the shared storage becomes the W2-wgrad input. + full_grid_phase_barrier(14); + uint32_t pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg( + expert_counts + + expert_idx)); + for (uint64_t linear = + static_cast( + blockIdx.x) * + kNumThreads + + threadIdx.x; + linear < + static_cast( + num_tokens) * + kIntermediateHidden; + linear += + static_cast( + kNumSMs) * + kNumThreads) { + const uint32_t token_idx = + linear / + kIntermediateHidden; + const uint32_t col = + linear - + static_cast( + token_idx) * + kIntermediateHidden; + const uint32_t pool_row = + pool_block_offset * BLOCK_M + + token_idx; + const uint64_t offset = + static_cast( + pool_row) * + kIntermediateHidden + + col; + h_weighted_output[offset] = + cd_dtype_t( + static_cast( + h_act_output[offset]) * + route_weights_fp32[ + pool_row]); + } + pool_block_offset += + math::ceil_div( + num_tokens, BLOCK_M); + } + } + } + + // Phase 3: dequantize canonical [W1; W3] once per launch, then + // consume it as the transposed BF16 operand for W13 dgrad. This + // phase starts only after L2 dgrad/SwiGLU has drained both TMEM + // accumulator stages, so the same 512-column allocation is reused. + + const uint32_t w13_launch_epoch = + launch_epoch ^ 0x80000000u; + + const auto for_each_w13_dgrad_block = + [&](const auto& func) { + uint32_t next_assigned_block = + blockIdx.x; + uint32_t global_block = 0; + uint32_t pool_block_offset = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; + ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg( + expert_counts + + expert_idx)); + const uint32_t num_m_blocks = + math::ceil_div( + num_tokens, BLOCK_M); + const uint32_t expert_blocks = + num_m_blocks * + kNumW13DgradBlockNs; + const uint32_t expert_end = + global_block + + expert_blocks; + + while (next_assigned_block < + global_block) + next_assigned_block += + kNumSMs; + while (next_assigned_block < + expert_end) { + const uint32_t local_block = + next_assigned_block - + global_block; + const uint32_t + m_block_idx = + local_block / + kNumW13DgradBlockNs; + const uint32_t + n_block_idx = + local_block - + m_block_idx * + kNumW13DgradBlockNs; + const uint32_t valid_m = + cute::min( + num_tokens - + m_block_idx * + BLOCK_M, + BLOCK_M); + func( + expert_idx, + pool_block_offset, + m_block_idx, + n_block_idx, + valid_m); + next_assigned_block += + kNumSMs; + } + global_block = expert_end; + pool_block_offset += + num_m_blocks; + } + }; + + trace_begin(15); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(15); + if (warp_idx == 0 && + cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; + i < kNumStages; ++i) { + full_barriers[i]->init(4); + empty_barriers[i]->init(1); + } + #pragma unroll + for (uint32_t i = 0; + i < kNumEpilogueStages; ++i) { + tmem_full_barriers[i]->init(1); + tmem_empty_barriers[i]->init( + 2 * + kNumDgradEpilogueThreads); + } + cutlass::arch::fence_barrier_init(); + } + trace_begin(16); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(16); + trace_begin(21); + + stage_idx = 0; + phase = 0; + if (warp_idx == 0) { + for_each_w13_dgrad_block( + [&](const uint32_t&, + const uint32_t& + pool_block_offset, + const uint32_t& + m_block_idx, + const uint32_t&, + const uint32_t& valid_m) { + const uint32_t pool_block_idx = + pool_block_offset + + m_block_idx; + #pragma unroll + for (uint32_t split_idx = 0; + split_idx < + kNumW13DgradSplits; + ++split_idx) { + #pragma unroll 1 + for (uint32_t k_block_idx = 0; + k_block_idx < + (2 * + kIntermediateHidden) / + (DGRAD_BLOCK_K * + kNumW13DgradSplits); + advance_pipeline( + k_block_idx)) { + empty_barriers[stage_idx] + ->wait(phase ^ 1); + uint32_t m_idx = + pool_block_idx * + BLOCK_M; + if (!is_leader_cta) + m_idx += + math::align( + valid_m, 16u) / + 2; + if (cute::elect_one_sync()) { + tma::copy< + DGRAD_BLOCK_K, + LOAD_BLOCK_M, + DGRAD_BLOCK_K * + sizeof( + cd_dtype_t), + cd_dtype_t>( + &tensor_map_grad_gate_up, + full_barriers[ + stage_idx], + smem_dgrad_a[ + stage_idx], + split_idx * + ((2 * + kIntermediateHidden) / + kNumW13DgradSplits) + + k_block_idx * + DGRAD_BLOCK_K, + m_idx, 2); + if (is_leader_cta) { + full_barriers[ + stage_idx] + ->arrive_and_expect_tx( + SMEM_A_SIZE_PER_STAGE * + 2); + } else { + full_barriers[ + stage_idx] + ->arrive(0u); + } + } + __syncwarp(); + } + } + }); + } else if (warp_idx == 1) { + for_each_w13_dgrad_block( + [&](const uint32_t& expert_idx, + const uint32_t&, + const uint32_t&, + const uint32_t& + n_block_idx, + const uint32_t&) { + #pragma unroll + for (uint32_t split_idx = 0; + split_idx < + kNumW13DgradSplits; + ++split_idx) { + #pragma unroll 1 + for (uint32_t k_block_idx = 0; + k_block_idx < + (2 * + kIntermediateHidden) / + (DGRAD_BLOCK_K * + kNumW13DgradSplits); + advance_pipeline( + k_block_idx)) { + const uint32_t + global_k_block_idx = + split_idx * + ((2 * + kIntermediateHidden) / + (DGRAD_BLOCK_K * + kNumW13DgradSplits)) + + k_block_idx; + const uint32_t + weight_tile_idx = + (expert_idx * + ((2 * + kIntermediateHidden) / + DGRAD_BLOCK_K) + + global_k_block_idx) * + kNumW13DgradBlockNs + + n_block_idx; + if constexpr (!kBF16Mode) { + while (ptx::ld_acq( + weight_tile_states + + kNumW2WeightTileStates + + weight_tile_idx) != + w13_launch_epoch) { + } + } + empty_barriers[stage_idx] + ->wait(phase ^ 1); + if (cute::elect_one_sync()) { + tma::copy< + LOAD_BLOCK_N, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * + sizeof( + dgrad_b_dtype_t), + dgrad_b_dtype_t>( + &tensor_map_w13_dequant, + full_barriers[ + stage_idx], + smem_dgrad_b[ + stage_idx], + n_block_idx * + BLOCK_N, + expert_idx * + (2 * + kIntermediateHidden) + + global_k_block_idx * + DGRAD_BLOCK_K, + 2); + if (is_leader_cta) { + full_barriers[ + stage_idx] + ->arrive_and_expect_tx( + SMEM_B_SIZE_PER_STAGE * + 2); + } else { + full_barriers[ + stage_idx] + ->arrive(0u); + } + } + __syncwarp(); + } + } + }); + } else if (warp_idx == 2) { + if (is_leader_cta) { + auto instr_desc = + cute::UMMA::make_instr_desc< + dgrad_b_dtype_t, + cd_dtype_t, float, + UMMA_M, UMMA_N, + cute::UMMA::Major::MN, + cute::UMMA::Major::K>(); + auto a_desc = + mma::sm100::make_umma_desc< + cute::UMMA::Major::K, + LOAD_BLOCK_M, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * + sizeof(cd_dtype_t)>( + smem_dgrad_a[0], 0, 0); + auto b_desc = + mma::sm100::make_umma_desc< + cute::UMMA::Major::MN, + LOAD_BLOCK_N, + DGRAD_BLOCK_K, + DGRAD_BLOCK_K * + sizeof( + dgrad_b_dtype_t)>( + smem_dgrad_b[0], 0, 0); + const uint32_t a_desc_lo = + lane_idx < kNumStages + ? a_desc.lo + + lane_idx * + SMEM_A_SIZE_PER_STAGE / + 16 + : 0; + const uint32_t b_desc_lo = + lane_idx < kNumStages + ? b_desc.lo + + lane_idx * + SMEM_B_SIZE_PER_STAGE / + 16 + : 0; + uint32_t current_iter = 0; + + for_each_w13_dgrad_block( + [&](const uint32_t&, + const uint32_t&, + const uint32_t&, + const uint32_t&, + const uint32_t& + valid_m) { + mma::sm100:: + update_instr_desc_with_umma_n( + instr_desc, + math::align( + valid_m, 16u)); + const auto + runtime_instr_desc = + cute::UMMA:: + make_runtime_instr_desc( + instr_desc); + #pragma unroll + for (uint32_t split_idx = 0; + split_idx < + kNumW13DgradSplits; + ++split_idx) { + const uint32_t accum_stage = + current_iter % + kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / + kNumEpilogueStages) & + 1; + tmem_empty_barriers[ + accum_stage] + ->wait( + accum_phase ^ 1); + ptx::tcgen05_after_thread_sync(); + + #pragma unroll 1 + for (uint32_t + k_block_idx = 0; + k_block_idx < + (2 * + kIntermediateHidden) / + (DGRAD_BLOCK_K * + kNumW13DgradSplits); + advance_pipeline( + k_block_idx)) { + full_barriers[ + stage_idx] + ->wait(phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t + a_desc_base = + ptx::exchange( + a_desc_lo, + stage_idx); + const uint32_t + b_desc_base = + ptx::exchange( + b_desc_lo, + stage_idx); + if (cute::elect_one_sync()) { + #pragma unroll + for (uint32_t k = 0; + k < + DGRAD_BLOCK_K / + DGRAD_UMMA_K; + ++k) { + a_desc.lo = + mma::sm100:: + advance_umma_desc_lo< + cute::UMMA::Major::K, + LOAD_BLOCK_M, + DGRAD_BLOCK_K * + sizeof( + cd_dtype_t), + cd_dtype_t>( + a_desc_base, + 0, + k * + DGRAD_UMMA_K); + b_desc.lo = + mma::sm100:: + advance_umma_desc_lo< + cute::UMMA::Major::MN, + LOAD_BLOCK_N, + DGRAD_BLOCK_K * + sizeof( + dgrad_b_dtype_t), + dgrad_b_dtype_t>( + b_desc_base, + 0, + k * + DGRAD_UMMA_K); + ptx:: + SM100_MMA_F16BF16_2x1SM_SS:: + fma( + b_desc, + a_desc, + accum_stage * + UMMA_N, + k_block_idx > + 0 || + k > 0, + runtime_instr_desc); + } + } + __syncwarp(); + constexpr uint16_t + kCTAMask = 0x3; + cutlass::arch:: + umma_arrive_multicast_2x1SM( + reinterpret_cast< + uint64_t*>( + empty_barriers[ + stage_idx]), + kCTAMask); + if (k_block_idx == + (2 * + kIntermediateHidden) / + (DGRAD_BLOCK_K * + kNumW13DgradSplits) - + 1) { + cutlass::arch:: + umma_arrive_multicast_2x1SM( + reinterpret_cast< + uint64_t*>( + tmem_full_barriers[ + accum_stage]), + kCTAMask); + } + __syncwarp(); + } + } + }); + if (current_iter > 0) { + const uint32_t last = + current_iter - 1; + tmem_empty_barriers[ + last % + kNumEpilogueStages] + ->wait( + (last / + kNumEpilogueStages) & + 1); + } + } + } else if (warp_idx >= 4) { + const uint32_t epilogue_warp_idx = + warp_idx - 4; + const uint32_t epilogue_thread_idx = + epilogue_warp_idx * 32 + + lane_idx; + uint32_t current_iter = 0; + + for_each_w13_dgrad_block( + [&](const uint32_t&, + const uint32_t& + pool_block_offset, + const uint32_t& + m_block_idx, + const uint32_t& + n_block_idx, + const uint32_t& valid_m) { + const uint32_t accum_stage = + current_iter % + kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter / + kNumEpilogueStages) & + 1; + current_iter += + kNumW13DgradSplits; + tmem_full_barriers[accum_stage]->wait( + accum_phase); + if constexpr (kBF16Mode) + tmem_full_barriers[ + accum_stage ^ 1] + ->wait(accum_phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t effective_m = + math::align(valid_m, 16u); + + for (uint32_t s = 0; + s < + effective_m / + STORE_BLOCK_M; + ++s) { + cutlass::arch:: + NamedBarrier::sync( + kNumDgradEpilogueThreads, + 0); + // The four proven loader warps own the 4 KiB + // TMEM-to-shared mapping. Extra dgrad epilogue + // warps participate only in the global scatter. + if (epilogue_warp_idx < + kNumEpilogueThreads / + 32) { + #pragma unroll + for (uint32_t i = 0; + i < + STORE_BLOCK_M / + 8; + ++i) { + const uint32_t tmem_addr = + accum_stage * UMMA_N + + s * STORE_BLOCK_M + + i * 8; + uint32_t w1_values[8]; + uint32_t w3_values[8]; + cute:: + SM100_TMEM_LOAD_16dp256b1x:: + copy( + tmem_addr, + w1_values[0], + w1_values[1], + w1_values[2], + w1_values[3]); + cute:: + SM100_TMEM_LOAD_16dp256b1x:: + copy( + tmem_addr | + 0x00100000, + w1_values[4], + w1_values[5], + w1_values[6], + w1_values[7]); + if constexpr (kBF16Mode) { + const uint32_t + w3_tmem_addr = + (accum_stage ^ + 1) * + UMMA_N + + s * + STORE_BLOCK_M + + i * 8; + cute:: + SM100_TMEM_LOAD_16dp256b1x:: + copy( + w3_tmem_addr, + w3_values[0], + w3_values[1], + w3_values[2], + w3_values[3]); + cute:: + SM100_TMEM_LOAD_16dp256b1x:: + copy( + w3_tmem_addr | + 0x00100000, + w3_values[4], + w3_values[5], + w3_values[6], + w3_values[7]); + } + cutlass::arch:: + fence_view_async_tmem_load(); + + constexpr uint32_t + kBankBytes = 16; + const uint32_t + outer_atom = + (epilogue_warp_idx / + 2) * + STORE_BLOCK_M * + 128; + const uint32_t + inner_atom = + i * 8 * 128; + const uint32_t row = + lane_idx % 8; + const uint32_t col = + (epilogue_warp_idx % + 2) * + 4 + + lane_idx / 8; + auto* smem_ptr = + reinterpret_cast< + uint8_t*>( + smem_cd[0]) + + outer_atom + + inner_atom + + row * + (kBankBytes * + 8) + + (col ^ row) * + kBankBytes; + const auto add_bf16_pair = + [](uint32_t a, + uint32_t b, + uint32_t c, + uint32_t d) { + const uint32_t + w1_packed = + math:: + cast_into_bf16_and_pack( + a, + b); + const uint32_t + w3_packed = + math:: + cast_into_bf16_and_pack( + c, + d); + const auto w1 = + *reinterpret_cast< + const nv_bfloat162*>( + &w1_packed); + const auto w3 = + *reinterpret_cast< + const nv_bfloat162*>( + &w3_packed); + const auto sum = + __hadd2_rn( + w1, w3); + return *reinterpret_cast< + const uint32_t*>( + &sum); + }; + if constexpr (kBF16Mode) { + ptx:: + SM90_U32x4_STSM_T:: + copy( + add_bf16_pair( + w1_values[0], + w1_values[1], + w3_values[0], + w3_values[1]), + add_bf16_pair( + w1_values[2], + w1_values[3], + w3_values[2], + w3_values[3]), + add_bf16_pair( + w1_values[4], + w1_values[5], + w3_values[4], + w3_values[5]), + add_bf16_pair( + w1_values[6], + w1_values[7], + w3_values[6], + w3_values[7]), + smem_ptr); + } else { + ptx:: + SM90_U32x4_STSM_T:: + copy( + math:: + cast_into_bf16_and_pack( + w1_values[0], + w1_values[1]), + math:: + cast_into_bf16_and_pack( + w1_values[2], + w1_values[3]), + math:: + cast_into_bf16_and_pack( + w1_values[4], + w1_values[5]), + math:: + cast_into_bf16_and_pack( + w1_values[6], + w1_values[7]), + smem_ptr); + } + } + } + cutlass::arch:: + NamedBarrier::sync( + kNumDgradEpilogueThreads, + 0); + + if constexpr ( + kWideGradXStore) { + DG_STATIC_ASSERT( + BLOCK_N % 8 == 0, + "Wide grad-x stores require eight-column alignment"); + DG_STATIC_ASSERT( + kHidden % 8 == 0, + "Wide grad-x stores require aligned output rows"); + #pragma unroll + for (uint32_t linear = + epilogue_thread_idx; + linear < + STORE_BLOCK_M * + (BLOCK_N / 8); + linear += + kNumDgradEpilogueThreads) { + const uint32_t row = + linear / + (BLOCK_N / 8); + const uint32_t n = + (linear - + row * + (BLOCK_N / 8)) * + 8; + const uint32_t local_m = + s * STORE_BLOCK_M + + row; + if (local_m >= valid_m) + continue; + const uint32_t n_atom = + n / 64; + const uint32_t + n_in_atom = + n - + n_atom * 64; + const uint32_t + row_in_atom = + row & 7; + const uint32_t + smem_byte_offset = + n_atom * + STORE_BLOCK_M * + 128 + + (row >> 3) * + 8 * 128 + + row_in_atom * 128 + + ((n_in_atom >> 3) ^ + row_in_atom) * + 16; + const auto packed = + *reinterpret_cast< + const uint4*>( + reinterpret_cast< + const uint8_t*>( + smem_cd[0]) + + smem_byte_offset); + const uint32_t pool_row = + (pool_block_offset + + m_block_idx) * + BLOCK_M + + local_m; + const uint32_t out_col = + n_block_idx * + BLOCK_N + + n; + if constexpr ( + kWriteGradXPool) { + *reinterpret_cast< + uint4*>( + grad_x_pool_output + + static_cast< + uint64_t>( + pool_row) * + kHidden + + out_col) = packed; + } + if constexpr ( + kDirectRemoteGradX) { + const auto metadata = + token_src_metadata[ + pool_row]; + auto* combine_buffer = + const_cast< + cd_dtype_t*>( + backward_grad_y); + auto* dst = + combine_buffer + + ((static_cast< + uint64_t>( + metadata + .topk_idx) * + backward_workspace + .num_max_tokens_per_rank + + metadata + .token_idx) * + kHidden + + out_col); + *reinterpret_cast< + uint4*>( + backward_sym_buffer + .map( + dst, + metadata + .rank_idx)) = + packed; + } + } + } else if constexpr ( + kVectorizedGradXStore) { + #pragma unroll + for (uint32_t linear = + epilogue_thread_idx; + linear < + STORE_BLOCK_M * + (BLOCK_N / 2); + linear += + kNumDgradEpilogueThreads) { + const uint32_t row = + linear / + (BLOCK_N / 2); + const uint32_t n = + (linear - + row * + (BLOCK_N / 2)) * + 2; + const uint32_t local_m = + s * STORE_BLOCK_M + + row; + if (local_m >= valid_m) + continue; + const uint32_t + row_in_atom = + row & 7; + const auto load_bf16_bits = + [&](const uint32_t + element_n) { + const uint32_t + n_atom = + element_n / + 64; + const uint32_t + n_in_atom = + element_n - + n_atom * + 64; + const uint32_t + smem_byte_offset = + n_atom * + STORE_BLOCK_M * + 128 + + (row >> 3) * + 8 * + 128 + + row_in_atom * + 128 + + ((n_in_atom >> + 3) ^ + row_in_atom) * + 16 + + (n_in_atom & + 7) * + sizeof( + cd_dtype_t); + return *reinterpret_cast< + uint16_t*>( + reinterpret_cast< + uint8_t*>( + smem_cd[0]) + + smem_byte_offset); + }; + const uint32_t packed = + static_cast( + load_bf16_bits(n)) | + (static_cast( + load_bf16_bits( + n + 1)) + << 16); + const uint32_t pool_row = + (pool_block_offset + + m_block_idx) * + BLOCK_M + + local_m; + const uint32_t out_col = + n_block_idx * + BLOCK_N + + n; + if constexpr ( + kWriteGradXPool) { + *reinterpret_cast< + uint32_t*>( + grad_x_pool_output + + static_cast< + uint64_t>( + pool_row) * + kHidden + + out_col) = packed; + } + if constexpr ( + kDirectRemoteGradX) { + const auto metadata = + token_src_metadata[ + pool_row]; + auto* combine_buffer = + const_cast< + cd_dtype_t*>( + backward_grad_y); + auto* dst = + combine_buffer + + ((static_cast< + uint64_t>( + metadata + .topk_idx) * + backward_workspace + .num_max_tokens_per_rank + + metadata + .token_idx) * + kHidden + + out_col); + *reinterpret_cast< + uint32_t*>( + backward_sym_buffer + .map( + dst, + metadata + .rank_idx)) = + packed; + } + } + } else { + #pragma unroll + for (uint32_t linear = + epilogue_thread_idx; + linear < + STORE_BLOCK_M * + BLOCK_N; + linear += + kNumDgradEpilogueThreads) { + const uint32_t row = + linear / BLOCK_N; + const uint32_t n = + linear - + row * BLOCK_N; + const uint32_t local_m = + s * + STORE_BLOCK_M + + row; + if (local_m >= valid_m) + continue; + const uint32_t n_atom = + n / 64; + const uint32_t + n_in_atom = + n - + n_atom * 64; + const uint32_t + row_in_atom = + row & 7; + const uint32_t + smem_byte_offset = + n_atom * + STORE_BLOCK_M * + 128 + + (row >> 3) * + 8 * 128 + + row_in_atom * + 128 + + ((n_in_atom >> 3) ^ + row_in_atom) * + 16 + + (n_in_atom & 7) * + sizeof( + cd_dtype_t); + const uint32_t pool_row = + (pool_block_offset + + m_block_idx) * + BLOCK_M + + local_m; + const uint32_t out_col = + n_block_idx * + BLOCK_N + + n; + const auto value = + *reinterpret_cast< + cd_dtype_t*>( + reinterpret_cast< + uint8_t*>( + smem_cd[0]) + + smem_byte_offset); + if constexpr ( + kWriteGradXPool) { + grad_x_pool_output[ + static_cast< + uint64_t>( + pool_row) * + kHidden + + out_col] = value; + } + if constexpr ( + kDirectRemoteGradX) { + const auto metadata = + token_src_metadata[ + pool_row]; + auto* combine_buffer = + const_cast< + cd_dtype_t*>( + backward_grad_y); + auto* dst = + combine_buffer + + ((static_cast< + uint64_t>( + metadata + .topk_idx) * + backward_workspace + .num_max_tokens_per_rank + + metadata + .token_idx) * + kHidden + + out_col); + *backward_sym_buffer + .map( + dst, + metadata + .rank_idx) = + value; + } + } + } + } + ptx::tcgen05_before_thread_sync(); + tmem_empty_barriers[accum_stage] + ->arrive(0u); + if constexpr (kBF16Mode) + tmem_empty_barriers[ + accum_stage ^ 1] + ->arrive(0u); + }); + } + } + + trace_end(21); + if constexpr (kNumRanks > 1) { + if constexpr ( + kDirectRemoteGradX || kComputeRouteGrad) { + // Publish every direct grad-x and route-gradient NVLink store + // before any destination rank consumes its source planes. + constexpr uint32_t + kDirectGradXDoneGridSyncIndex = 1; + constexpr uint32_t + kDirectGradXDoneBarrierTag = 9; + if constexpr (kTraceKernel) { + // Decompose the otherwise identical NVLink barrier so the + // trace distinguishes local compute/grid skew from the + // cross-rank signal and its publication grid sync. + trace_begin(17); + comm::grid_sync< + kNumSMs, + kDirectGradXDoneGridSyncIndex>( + backward_workspace, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + trace_end(17); + + if (blockIdx.x == 0) + trace_begin(18); + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, + kDirectGradXDoneGridSyncIndex, + kDirectGradXDoneBarrierTag>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }, + false, false); + if (blockIdx.x == 0) + trace_end(18); + + trace_begin(19); + comm::grid_sync< + kNumSMs, + kDirectGradXDoneGridSyncIndex>( + backward_workspace, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + trace_end(19); + } else { + comm::nvlink_barrier< + kNumRanks, kNumSMs, kNumThreads, + kDirectGradXDoneGridSyncIndex, + kDirectGradXDoneBarrierTag>( + backward_workspace, + backward_sym_buffer, + blockIdx.x, + threadIdx.x, + []() { __syncthreads(); }); + } + } + } + + const auto clear_wgrad_padding_rows = [&]() { + uint32_t pad_pool_block_offset = 0; + uint32_t pad_global_block = 0; + #pragma unroll + for (uint32_t expert_idx = 0; + expert_idx < kNumExperts; ++expert_idx) { + const uint32_t num_tokens = + static_cast( + __ldg(expert_counts + expert_idx)); + const uint32_t num_blocks = + math::ceil_div(num_tokens, BLOCK_M); + if (num_blocks != 0) { + const uint32_t last_valid = + num_tokens - (num_blocks - 1) * BLOCK_M; + const uint32_t pool_block = + pad_pool_block_offset + num_blocks - 1; + if (pad_global_block % kNumSMs == + blockIdx.x) { + for (uint32_t linear = threadIdx.x; + linear < + (BLOCK_M - last_valid) * + (kHidden + + 3 * kIntermediateHidden); + linear += kNumThreads) { + const uint32_t row_delta = + linear / + (kHidden + + 3 * kIntermediateHidden); + const uint32_t col = + linear - + row_delta * + (kHidden + + 3 * + kIntermediateHidden); + const uint32_t pool_row = + pool_block * BLOCK_M + + last_valid + row_delta; + if (col < kHidden) { + grad_ye_output[ + static_cast( + pool_row) * + kHidden + + col] = cd_dtype_t(0.0f); + } else if ( + col < + kHidden + + kIntermediateHidden) { + h_weighted_output[ + static_cast( + pool_row) * + kIntermediateHidden + + col - kHidden] = + cd_dtype_t(0.0f); + } else { + grad_gate_up_output[ + static_cast( + pool_row) * + (2 * + kIntermediateHidden) + + col - kHidden - + kIntermediateHidden] = + cd_dtype_t(0.0f); + } + } + } + ++pad_global_block; + } + pad_pool_block_offset += num_blocks; + } + }; + + // Standalone Kernel B consumes only these three padded operands. Valid + // rows are fully overwritten above; clear only the final partial block + // of each expert instead of memset'ing every active scratch prefix. + if constexpr (kClearWgradPadding) + clear_wgrad_padding_rows(); + + __syncthreads(); + trace_begin(20); + comm::cluster_sync_with_relaxed_arrive(); + trace_end(20); + trace_end(0); + if (warp_idx == 0) + Allocator().free(0, kNumTmemCols); + } +#endif +} + +// -------------------------------------------------------------------------- +// Fixed SM103 E4M3-block128 training reverse used by GLM-5.2. +// +// This is deliberately AOT-only. It keeps the upstream 2-CTA/TMEM structure +// but turns the reverse into one resident expert-wave loop. Compact x and dy +// are power-of-two quantized into symmetric storage by this kernel, each +// expert is pulled into the private ring, and the three hardware-scaled GEMMs +// execute before that ring slot is reused. Only the two compressed operands +// retained for the dedicated wgrad kernels span the full padded route pool. +// -------------------------------------------------------------------------- + +namespace sm103_block128_backward { + +static constexpr uint32_t kHidden = 6144; +static constexpr uint32_t kIntermediate = 2048; +static constexpr uint32_t kGlobalExperts = 256; +static constexpr uint32_t kTopK = 8; +static constexpr uint32_t kBlockM = 192; +static constexpr uint32_t kBlockN = 128; +static constexpr uint32_t kBlockK = 128; +static constexpr uint32_t kSFBlockM = 256; +static constexpr uint32_t kSFBlockN = 128; +static constexpr uint32_t kStages = 6; +static constexpr uint32_t kThreads = 512; +static constexpr uint32_t kNumDispatchWarps = 4; +// Match the upstream persistent dispatch granularity. Reverse pulls BF16 +// grad-y, so four 3-KiB transactions cover one 6,144-value token. Two shared +// stages let the next remote TMA overlap score/dz work on the current chunk. +static constexpr uint32_t kDispatchPullBytes = 3072; +static constexpr uint32_t kDispatchPullStages = 2; +static constexpr uint32_t kDispatchPullElements = + kDispatchPullBytes / sizeof(cutlass::bfloat16_t); +static constexpr uint32_t kDispatchPullChunks = + kHidden * sizeof(cutlass::bfloat16_t) / kDispatchPullBytes; +// The reverse epilogue retains sm100_store_cd_swap_ab's exact 128-thread TMEM +// mapping over the two 64-column BF16 atoms of BLOCK_N, but consumes the tile +// in place: W2 feeds fused SwiGLU backward/quantization and W13 feeds remote +// combine. STORE_BLOCK_M=16 covers every 16-row-aligned UMMA-N extent. +static constexpr uint32_t kStoreBlockM = 16; +static constexpr uint32_t kEpilogueThreads = 128; +static constexpr uint32_t kNumEpilogueStages = 2; +static constexpr uint32_t kNumTMAStoreStages = 2; +static constexpr uint32_t kLoadBlockM = kBlockM / 2; +static constexpr uint32_t kLoadBlockN = kBlockN; +static constexpr uint32_t kUMMAM = 256; +static constexpr uint32_t kUMMAN = kBlockM; +static constexpr uint32_t kUMMAK = 32; +static constexpr uint32_t kSwizzle = 128; +static constexpr uint32_t kUTCCPAlignedElements = 128; +// CUTLASS user named-barrier ID 0 maps to physical BAR 8. The four independent +// reduction groups use BAR 4..7, so neither overlaps BAR 0 (__syncthreads) or +// the TMA epilogue. +static constexpr uint32_t kReductionBarrierBase = 4; +static constexpr uint32_t kNumTmemAccumCols = + kUMMAN * kNumEpilogueStages; +static constexpr uint32_t kNumTmemSFACols = kSFBlockM / 32; +static constexpr uint32_t kNumTmemSFBCols = kSFBlockN / 32; +static constexpr uint32_t kTmemSFAStart = kNumTmemAccumCols; +static constexpr uint32_t kTmemSFBStart = + kNumTmemAccumCols + kNumTmemSFACols; +static constexpr uint32_t kNumTmemCols = + utils::get_num_aligned_tmem_cols< + kNumTmemAccumCols + kNumTmemSFACols + kNumTmemSFBCols>(); + +using fp8_t = cutlass::float_e4m3_t; +using bf16_t = cutlass::bfloat16_t; +using Barrier = cutlass::arch::ClusterTransactionBarrier; + +struct alignas(1024) SharedStorage { + alignas(1024) bf16_t dispatch_pull[kNumDispatchWarps] + [kDispatchPullStages] + [kDispatchPullElements]; + alignas(1024) bf16_t smem_cd[kNumTMAStoreStages] + [kStoreBlockM * kBlockN]; + alignas(1024) fp8_t smem_a[kStages][kLoadBlockM * kBlockK]; + alignas(1024) fp8_t smem_b[kStages][kLoadBlockN * kBlockK]; + alignas(1024) uint32_t smem_sfa[kStages][kSFBlockM]; + alignas(1024) uint32_t smem_sfb[kStages][kSFBlockN]; + alignas(128) float reduce_values[4][8]; + Barrier full_barriers[kStages]; + Barrier empty_barriers[kStages]; + Barrier tmem_full_barriers[kNumEpilogueStages]; + Barrier tmem_empty_barriers[kNumEpilogueStages]; + Barrier dispatch_barriers[kNumDispatchWarps] + [kDispatchPullStages]; + uint32_t tmem_ptr; +}; + +DG_STATIC_ASSERT(kNumTmemCols <= 512, "SM103 backward exceeds TMEM"); +DG_STATIC_ASSERT( + kEpilogueThreads == 128 && kStoreBlockM == 16, + "SM103 swap-AB epilogue requires one warpgroup and 16-row stores"); +DG_STATIC_ASSERT( + kDispatchPullChunks * kDispatchPullBytes == + kHidden * sizeof(bf16_t), + "reverse dispatch pull must cover one complete hidden row"); + +CUTLASS_DEVICE float warp_reduce_max(float value) { + #pragma unroll + for (uint32_t offset = 16; offset > 0; offset >>= 1) + value = cute::max( + value, __shfl_down_sync(0xffffffff, value, offset)); + return value; +} + +CUTLASS_DEVICE float warp_reduce_sum(float value) { + #pragma unroll + for (uint32_t offset = 16; offset > 0; offset >>= 1) + value = __fadd_rn( + value, __shfl_down_sync(0xffffffff, value, offset)); + return value; +} + +template +CUTLASS_DEVICE float reduce_group_128( + float value, SharedStorage& storage, const uint32_t group_idx) { + const uint32_t lane = threadIdx.x & 31; + const uint32_t warp_in_group = (threadIdx.x >> 5) & 3; + value = kMax ? warp_reduce_max(value) : warp_reduce_sum(value); + if (lane == 0) + storage.reduce_values[group_idx][warp_in_group] = value; + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); + if (warp_in_group == 0) { + value = lane < 4 + ? storage.reduce_values[group_idx][lane] + : (kMax ? 0.0f : 0.0f); + value = kMax ? warp_reduce_max(value) : warp_reduce_sum(value); + if (lane == 0) + storage.reduce_values[group_idx][4] = value; + } + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); + return storage.reduce_values[group_idx][4]; +} + +CUTLASS_DEVICE float2 reduce_group_128_max_pair( + float2 value, SharedStorage& storage, const uint32_t group_idx) { + const uint32_t lane = threadIdx.x & 31; + const uint32_t warp_in_group = (threadIdx.x >> 5) & 3; + value.x = warp_reduce_max(value.x); + value.y = warp_reduce_max(value.y); + if (lane == 0) { + storage.reduce_values[group_idx][warp_in_group] = value.x; + storage.reduce_values[group_idx][4 + warp_in_group] = value.y; + } + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); + if (warp_in_group == 0) { + value.x = lane < 4 + ? storage.reduce_values[group_idx][lane] + : 0.0f; + value.y = lane < 4 + ? storage.reduce_values[group_idx][4 + lane] + : 0.0f; + value.x = warp_reduce_max(value.x); + value.y = warp_reduce_max(value.y); + if (lane == 0) { + storage.reduce_values[group_idx][0] = value.x; + storage.reduce_values[group_idx][4] = value.y; + } + } + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); + return { + storage.reduce_values[group_idx][0], + storage.reduce_values[group_idx][4]}; +} + +CUTLASS_DEVICE uint32_t packed_power2_scale( + const float amax, float& scale_inv) { + const float raw = cute::max(amax * (1.0f / 448.0f), 0x1p-127f); + const int exponent = math::fast_log2_ceil(raw); + // UE8M0 code zero represents 2^-127. It is an FP32 subnormal and cannot + // be constructed by merely shifting an IEEE exponent field. + const float scale = exponent == -127 + ? 0x1p-127f + : math::fast_pow2(exponent); + scale_inv = math::fast_pow2(-exponent); + return ((*reinterpret_cast(&scale)) >> 23) * + 0x01010101u; +} + +CUTLASS_DEVICE uint32_t transform_sf_row( + const uint32_t row) { + const uint32_t in_block = row % kBlockM; + return row / kBlockM * kSFBlockM + + (in_block & ~127u) + (in_block & 31u) * 4 + + ((in_block >> 5) & 3u); +} + +template +CUTLASS_DEVICE uint32_t phase_shape_n() { + if constexpr (kPhase == sched::BackwardBlockPhase::RecomputeW13) + return 2 * kIntermediate; + if constexpr (kPhase == sched::BackwardBlockPhase::W2Dgrad) + return kIntermediate; + return kHidden; +} + +template +CUTLASS_DEVICE uint32_t phase_shape_k() { + if constexpr (kPhase == sched::BackwardBlockPhase::W13Dgrad) + return 2 * kIntermediate; + return kHidden; +} + +template +CUTLASS_DEVICE void run_gemm_phase( + SharedStorage& storage, + const uint32_t local_expert_idx, + const uint32_t num_tokens, + const cute::TmaDescriptor& tensor_map_a, + const cute::TmaDescriptor& tensor_map_sfa, + const cute::TmaDescriptor& tensor_map_w13_recompute, + const cute::TmaDescriptor& tensor_map_w2_dgrad, + const cute::TmaDescriptor& tensor_map_w13_dgrad, + const cute::TmaDescriptor& tensor_map_output, + const float* w13_scales, + const float* w2_scales, + const uint32_t sf_ring_tokens) { + constexpr uint32_t shape_n = + kPhase == sched::BackwardBlockPhase::RecomputeW13 + ? 2 * kIntermediate + : kPhase == sched::BackwardBlockPhase::W2Dgrad + ? kIntermediate + : kHidden; + constexpr uint32_t shape_k = + kPhase == sched::BackwardBlockPhase::W13Dgrad + ? 2 * kIntermediate + : kHidden; + constexpr uint32_t num_block_ns = shape_n / kBlockN; + constexpr uint32_t num_block_ks = shape_k / kBlockK; + constexpr cute::UMMA::Major major_b = + kPhase == sched::BackwardBlockPhase::RecomputeW13 + ? cute::UMMA::Major::K + : cute::UMMA::Major::MN; + + const bool leader_cta = cute::block_rank_in_cluster() == 0; + const uint32_t warp_idx = cutlass::canonical_warp_idx_sync(); + const uint32_t lane_idx = ptx::get_lane_idx(); + const uint32_t num_m_blocks = math::ceil_div(num_tokens, kBlockM); + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 4 && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kStages; ++i) { + storage.full_barriers[i].init(4); + storage.empty_barriers[i].init(1); + } + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { + storage.tmem_full_barriers[i].init(1); + storage.tmem_empty_barriers[i].init(2 * kEpilogueThreads); + } + cutlass::arch::fence_barrier_init(); + } + comm::cluster_sync_with_relaxed_arrive(); + + uint32_t stage_idx = 0; + uint32_t phase = 0; + const auto advance_pipeline = [&](uint32_t& k_block_idx) { + ++k_block_idx; + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + }; + + if (warp_idx == 4) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + for (uint32_t block = blockIdx.x; + block < num_m_blocks * num_block_ns; block += kNumSMs) { + const uint32_t m_block_idx = block / num_block_ns; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * 192u, 192u); + #pragma unroll 2 + for (uint32_t k_block_idx = 0; k_block_idx < num_block_ks; + advance_pipeline(k_block_idx)) { + storage.empty_barriers[stage_idx].wait(phase ^ 1); + uint32_t m_idx = m_block_idx * 192u; + if (!leader_cta) + m_idx += math::align(valid_m, 16u) / 2; + if (cute::elect_one_sync()) { + tma::copy( + &tensor_map_a, &storage.full_barriers[stage_idx], + storage.smem_a[stage_idx], + k_block_idx * kBlockK, m_idx, 2); + tma::copy( + &tensor_map_sfa, &storage.full_barriers[stage_idx], + storage.smem_sfa[stage_idx], + m_block_idx * kSFBlockM, + k_block_idx, 2); + if (leader_cta) { + storage.full_barriers[stage_idx].arrive_and_expect_tx( + sizeof(storage.smem_a[0]) * 2 + + sizeof(storage.smem_sfa[0]) * 2); + } else { + storage.full_barriers[stage_idx].arrive(0u); + } + } + __syncwarp(); + } + } + } else if (warp_idx == 5) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + for (uint32_t block = blockIdx.x; + block < num_m_blocks * num_block_ns; block += kNumSMs) { + const uint32_t n_block_idx = block % num_block_ns; + #pragma unroll 2 + for (uint32_t k_block_idx = 0; k_block_idx < num_block_ks; + advance_pipeline(k_block_idx)) { + storage.empty_barriers[stage_idx].wait(phase ^ 1); + if constexpr ( + kPhase == sched::BackwardBlockPhase::RecomputeW13) { + if (cute::elect_one_sync()) { + constexpr uint32_t logical_rows = kBlockN / 2; + const uint32_t logical_row = + n_block_idx * logical_rows; + const uint32_t up_row = + (local_expert_idx * 2 + 1) * kIntermediate + + logical_row; + const uint32_t gate_row = + (local_expert_idx * 2) * kIntermediate + + logical_row; + tma::copy( + &tensor_map_w13_recompute, + &storage.full_barriers[stage_idx], + storage.smem_b[stage_idx], + k_block_idx * kBlockK, up_row, 2); + tma::copy( + &tensor_map_w13_recompute, + &storage.full_barriers[stage_idx], + storage.smem_b[stage_idx] + + logical_rows * kBlockK, + k_block_idx * kBlockK, gate_row, 2); + } + } else { + const auto* map = + kPhase == sched::BackwardBlockPhase::W2Dgrad + ? &tensor_map_w2_dgrad + : &tensor_map_w13_dgrad; + if (cute::elect_one_sync()) { + const uint32_t outer_k = + kPhase == sched::BackwardBlockPhase::W2Dgrad + ? local_expert_idx * kHidden + + k_block_idx * kBlockK + : local_expert_idx * (2 * kIntermediate) + + k_block_idx * kBlockK; + tma::copy( + map, &storage.full_barriers[stage_idx], + storage.smem_b[stage_idx], + n_block_idx * kBlockN, outer_k, 2); + } + } + + #pragma unroll + for (uint32_t row = lane_idx; row < kBlockN; row += 32) { + float scale; + if constexpr ( + kPhase == sched::BackwardBlockPhase::RecomputeW13) { + constexpr uint32_t logical_rows = kBlockN / 2; + const uint32_t logical_row = + n_block_idx * logical_rows + + row % logical_rows; + const uint32_t canonical_expert = + local_expert_idx * 2 + + (row < logical_rows ? 1u : 0u); + scale = __ldg( + w13_scales + + (canonical_expert * (kIntermediate / 128) + + logical_row / 128) * + (kHidden / 128) + + k_block_idx); + } else if constexpr ( + kPhase == sched::BackwardBlockPhase::W2Dgrad) { + scale = __ldg( + w2_scales + + (local_expert_idx * (kHidden / 128) + + k_block_idx) * + (kIntermediate / 128) + + n_block_idx); + } else { + const uint32_t plane = + k_block_idx / (kIntermediate / 128); + const uint32_t row_block = + k_block_idx % (kIntermediate / 128); + scale = __ldg( + w13_scales + + ((local_expert_idx * 2 + plane) * + (kIntermediate / 128) + + row_block) * + (kHidden / 128) + + n_block_idx); + } + const uint32_t exponent = + __float_as_uint(scale) >> 23; + storage.smem_sfb[stage_idx][row] = + exponent * 0x01010101u; + } + __syncwarp(); + if (cute::elect_one_sync()) { + if (leader_cta) { + storage.full_barriers[stage_idx] + .arrive_and_expect_tx( + sizeof(storage.smem_b[0]) * 2); + } else { + storage.full_barriers[stage_idx].arrive(0u); + } + } + __syncwarp(); + } + } + } else if (warp_idx == 6) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + if (leader_cta) { + auto instr_desc = + cute::UMMA::make_instr_desc_block_scaled< + fp8_t, fp8_t, float, cutlass::float_ue8m0_t, + kUMMAM, kUMMAN, major_b, + cute::UMMA::Major::K>(); + auto sf_desc = mma::sm100::make_sf_desc(nullptr); + auto a_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::K, kLoadBlockM, kBlockK, kSwizzle>( + storage.smem_a[0], 0, 0); + auto b_desc = mma::sm100::make_umma_desc< + major_b, kLoadBlockN, kBlockK, kSwizzle>( + storage.smem_b[0], 0, 0); + const uint32_t a_desc_lo = + lane_idx < kStages + ? a_desc.lo + + lane_idx * sizeof(storage.smem_a[0]) / 16 + : 0u; + const uint32_t b_desc_lo = + lane_idx < kStages + ? b_desc.lo + + lane_idx * sizeof(storage.smem_b[0]) / 16 + : 0u; + uint32_t current_iter = 0; + for (uint32_t block = blockIdx.x; + block < num_m_blocks * num_block_ns; + block += kNumSMs) { + const uint32_t m_block_idx = block / num_block_ns; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * 192u, 192u); + mma::sm100::update_instr_desc_with_umma_n( + instr_desc, math::align(valid_m, 16u)); + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1; + storage.tmem_empty_barriers[accum_stage].wait( + accum_phase ^ 1); + ptx::tcgen05_after_thread_sync(); + + #pragma unroll 2 + for (uint32_t k_block_idx = 0; + k_block_idx < num_block_ks; + advance_pipeline(k_block_idx)) { + storage.full_barriers[stage_idx].wait(phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t a_base = + ptx::exchange(a_desc_lo, stage_idx); + const uint32_t b_base = + ptx::exchange(b_desc_lo, stage_idx); + if (cute::elect_one_sync()) { + using utccp_t = + cute::SM100_UTCCP_4x32dp128bit_2cta; + #pragma unroll + for (uint32_t i = 0; + i < kSFBlockM / kUTCCPAlignedElements; ++i) { + mma::sm100::replace_smem_desc_addr( + sf_desc, + storage.smem_sfa[stage_idx] + + i * kUTCCPAlignedElements); + utccp_t::copy( + sf_desc, kTmemSFAStart + i * 4); + } + #pragma unroll + for (uint32_t i = 0; + i < kSFBlockN / kUTCCPAlignedElements; ++i) { + mma::sm100::replace_smem_desc_addr( + sf_desc, + storage.smem_sfb[stage_idx] + + i * kUTCCPAlignedElements); + utccp_t::copy( + sf_desc, kTmemSFBStart + i * 4); + } + #pragma unroll + for (uint32_t k = 0; k < kBlockK / kUMMAK; ++k) { + const auto runtime_desc = + mma::sm100::make_runtime_instr_desc_with_sf_id( + instr_desc, k, k); + a_desc.lo = mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::K, kLoadBlockM, + kSwizzle, fp8_t>(a_base, 0, k * kUMMAK); + b_desc.lo = mma::sm100::advance_umma_desc_lo< + major_b, kLoadBlockN, kSwizzle, fp8_t>( + b_base, 0, k * kUMMAK); + ptx::SM100_MMA_MXF8F6F4_2x1SM_SS::fma( + b_desc, a_desc, + accum_stage * kUMMAN, + k_block_idx > 0 || k > 0, + runtime_desc, + 392u, 384u); + } + } + __syncwarp(); + constexpr uint16_t cta_mask = 3; + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.empty_barriers[stage_idx]), + cta_mask); + if (k_block_idx == num_block_ks - 1) { + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.tmem_full_barriers[accum_stage]), + cta_mask); + } + __syncwarp(); + } + } + if (current_iter > 0) { + const uint32_t last = current_iter - 1; + storage.tmem_empty_barriers[ + last % kNumEpilogueStages] + .wait((last / kNumEpilogueStages) & 1); + } + } + } else if (warp_idx >= 8 && warp_idx < 12) { + cutlass::arch::warpgroup_reg_alloc<208>(); + const uint32_t epilogue_warp_idx = warp_idx - 8; + uint32_t current_iter = 0; + uint32_t tma_stage_idx = 0; + auto smem_cd = utils::PatternVisitor([&](const uint32_t& i) { + return storage.smem_cd[i]; + }); + for (uint32_t block = blockIdx.x; + block < num_m_blocks * num_block_ns; block += kNumSMs) { + const uint32_t m_block_idx = block / num_block_ns; + const uint32_t n_block_idx = block % num_block_ns; + const uint32_t valid_m = cute::min( + num_tokens - m_block_idx * 192u, 192u); + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1; + storage.tmem_full_barriers[accum_stage].wait(accum_phase); + ptx::tcgen05_after_thread_sync(); + epilogue::sm100_store_cd_swap_ab< + kBlockM, kBlockN, kStoreBlockM, kBlockN, + kSwizzle, kNumTMAStoreStages, kEpilogueThreads, + GemmType::Normal, false, bf16_t, + epilogue::transform::EpilogueIdentity>( + smem_cd, tma_stage_idx, + accum_stage * kUMMAN, + m_block_idx * kBlockM, + n_block_idx * kBlockN, 0, + math::align(valid_m, 16u), + epilogue_warp_idx, lane_idx, + &storage.tmem_empty_barriers[accum_stage], + tensor_map_output); + } + if (epilogue_warp_idx == 0) + cute::tma_store_wait<0>(); + __syncwarp(); + } else { + cutlass::arch::warpgroup_reg_dealloc<40>(); + } + + // Match the upstream 2-CTA GEMM role join. The epilogue has drained its + // final TMA stage before its warpgroup reaches this cluster boundary. + comm::cluster_sync_with_relaxed_arrive(); +} + +template +CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void +sm103_fp8_block128_mega_moe_backward_impl( + const int* expert_counts, + const layout::TokenSrcMetadata* token_src_metadata, + const uint32_t num_tokens, + const uint32_t capacity, + const uint32_t ring_tokens, + const uint32_t sf_ring_tokens, + const uint32_t max_pool_tokens, + const __grid_constant__ layout::SymBuffer sym_buffer, + const __grid_constant__ layout::Workspace workspace, + const bf16_t* compact_x, + const bf16_t* compact_grad_y, + const float* compact_scores, + fp8_t* symmetric_x, + uint32_t* symmetric_x_sf, + fp8_t* symmetric_grad_y, + uint32_t* symmetric_grad_y_sf, + float* symmetric_scores, + float* symmetric_grad_scores, + bf16_t* combine_slots, + fp8_t* ring_x, + uint32_t* ring_x_sf, + fp8_t* ring_grad_y, + uint32_t* ring_grad_y_sf, + float* ring_scores, + fp8_t* ring_h, + uint32_t* ring_h_sf, + fp8_t* ring_grad_preact, + uint32_t* ring_grad_preact_sf, + bf16_t* ring_bf16, + float* ring_dscore, + fp8_t* full_x, + uint32_t* full_x_sf, + fp8_t* full_grad_y, + uint32_t* full_grad_y_sf, + float* full_scores, + fp8_t* full_h, + uint32_t* full_h_sf, + fp8_t* full_grad_preact, + uint32_t* full_grad_preact_sf, + bf16_t* grad_x, + float* grad_scores, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_x, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_x_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_y, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_y_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_h, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_h_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_preact, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_preact_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_recompute, + const __grid_constant__ cute::TmaDescriptor tensor_map_w2_dgrad, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_dgrad, + const __grid_constant__ cute::TmaDescriptor tensor_map_gate_up, + const __grid_constant__ cute::TmaDescriptor tensor_map_grad_h, + const __grid_constant__ cute::TmaDescriptor tensor_map_grad_x, + const float* w13_scales, + const float* w2_scales) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + constexpr uint32_t kLocalExperts = kGlobalExperts / kNumRanks; + const uint32_t global_thread = + blockIdx.x * kThreads + threadIdx.x; + const uint32_t global_stride = kNumSMs * kThreads; + const uint32_t group_idx = threadIdx.x / 128; + const uint32_t group_lane = threadIdx.x % 128; + const uint32_t group_global = blockIdx.x * 4 + group_idx; + const uint32_t group_stride = kNumSMs * 4; + const uint32_t hidden_blocks = kHidden / 128; + extern __shared__ __align__(1024) uint8_t smem_buffer[]; + SharedStorage& storage = + *reinterpret_cast(smem_buffer); + + // Install this call's ordinary PyTorch inputs into the symmetric compact + // planes. Four 128-thread groups quantize independent rows/blocks; the + // resulting FP32 ABI is represented by four repeated UE8M0 bytes. + for (uint64_t work = group_global; + work < static_cast(num_tokens) * hidden_blocks; + work += group_stride) { + const uint32_t row = work / hidden_blocks; + const uint32_t block = work - + static_cast(row) * hidden_blocks; + const uint32_t col = block * 128 + group_lane; + const uint64_t index = + static_cast(row) * kHidden + col; + const float x = static_cast(compact_x[index]); + const float dy = static_cast(compact_grad_y[index]); + const float x_amax = reduce_group_128( + cute::abs(x), storage, group_idx); + const float dy_amax = reduce_group_128( + cute::abs(dy), storage, group_idx); + float x_inv, dy_inv; + const uint32_t x_scale = packed_power2_scale(x_amax, x_inv); + const uint32_t dy_scale = packed_power2_scale(dy_amax, dy_inv); + if (group_lane == 0) { + symmetric_x_sf[row * hidden_blocks + block] = x_scale; + symmetric_grad_y_sf[row * hidden_blocks + block] = dy_scale; + } + symmetric_x[index] = fp8_t(x * x_inv); + symmetric_grad_y[index] = fp8_t(dy * dy_inv); + if (block == 0 && group_lane < kTopK) { + const uint64_t route = + static_cast(row) * kTopK + group_lane; + symmetric_scores[route] = compact_scores[route]; + } + } + __syncthreads(); + + // Publish compact q/s to peers before any expert wave performs remote + // gathers. Grid/NVLink counters are reusable across calls and layers. + comm::nvlink_barrier( + workspace, sym_buffer, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + if (cutlass::canonical_warp_idx_sync() == 7 && + cute::elect_one_sync()) { + cutlass::arch::fence_barrier_init(); + } + comm::cluster_sync_with_relaxed_arrive(); + if (cutlass::canonical_warp_idx_sync() == 7) + cute::TMEM::Allocator2Sm().allocate(kNumTmemCols, &storage.tmem_ptr); + comm::cluster_sync_with_relaxed_arrive(); + + uint32_t pool_block_offset = 0; + #pragma unroll 1 + for (uint32_t expert = 0; expert < kLocalExperts; ++expert) { + const uint32_t count = + static_cast(__ldg(expert_counts + expert)); + const uint32_t pool_row_offset = pool_block_offset * kBlockM; + DG_DEVICE_ASSERT(count <= ring_tokens); + DG_DEVICE_ASSERT( + pool_row_offset + count <= max_pool_tokens); + + // Gather both compact operands and their exact route score into the + // local expert ring. FP8 vectors stay FP8; no dequantized x/dy pool is + // constructed. + constexpr uint32_t vecs_per_row = kHidden / sizeof(uint4); + for (uint64_t linear = global_thread; + linear < static_cast(count) * vecs_per_row; + linear += global_stride) { + const uint32_t row = linear / vecs_per_row; + const uint32_t vec = linear - + static_cast(row) * vecs_per_row; + const auto metadata = + token_src_metadata[pool_row_offset + row]; + const auto* remote_x = sym_buffer.map( + reinterpret_cast(symmetric_x) + + static_cast(metadata.token_idx) * vecs_per_row + + vec, + metadata.rank_idx); + const auto* remote_dy = sym_buffer.map( + reinterpret_cast(symmetric_grad_y) + + static_cast(metadata.token_idx) * vecs_per_row + + vec, + metadata.rank_idx); + const uint4 x_value = *remote_x; + const uint4 grad_y_value = *remote_dy; + reinterpret_cast(ring_x)[ + static_cast(row) * vecs_per_row + vec] = x_value; + reinterpret_cast(ring_grad_y)[ + static_cast(row) * vecs_per_row + vec] = + grad_y_value; + const uint64_t full_vec = + static_cast(pool_row_offset + row) * + vecs_per_row + + vec; + reinterpret_cast(full_x)[full_vec] = x_value; + reinterpret_cast(full_grad_y)[full_vec] = grad_y_value; + } + for (uint64_t linear = global_thread; + linear < static_cast(count) * hidden_blocks; + linear += global_stride) { + const uint32_t row = linear / hidden_blocks; + const uint32_t block = linear - + static_cast(row) * hidden_blocks; + const auto metadata = + token_src_metadata[pool_row_offset + row]; + const uint64_t remote_index = + static_cast(metadata.token_idx) * hidden_blocks + + block; + const uint32_t sf_row = transform_sf_row(row); + const uint32_t x_scale = *sym_buffer.map( + symmetric_x_sf + remote_index, metadata.rank_idx); + const uint32_t grad_y_scale = *sym_buffer.map( + symmetric_grad_y_sf + remote_index, metadata.rank_idx); + ring_x_sf[block * sf_ring_tokens + sf_row] = x_scale; + ring_grad_y_sf[block * sf_ring_tokens + sf_row] = grad_y_scale; + const uint64_t full_scale = + static_cast(pool_row_offset + row) * + hidden_blocks + + block; + full_x_sf[full_scale] = x_scale; + full_grad_y_sf[full_scale] = grad_y_scale; + } + for (uint32_t row = global_thread; row < count; + row += global_stride) { + const auto metadata = + token_src_metadata[pool_row_offset + row]; + const float score = *sym_buffer.map( + symmetric_scores + + static_cast(metadata.token_idx) * kTopK + + metadata.topk_idx, + metadata.rank_idx); + ring_scores[row] = score; + full_scores[pool_row_offset + row] = score; + ring_dscore[row] = 0.0f; + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + if (count != 0) { + run_gemm_phase< + sched::BackwardBlockPhase::RecomputeW13, kNumSMs>( + storage, expert, count, + tensor_map_ring_x, tensor_map_ring_x_sf, + tensor_map_w13_recompute, tensor_map_w2_dgrad, + tensor_map_w13_dgrad, tensor_map_gate_up, + w13_scales, w2_scales, sf_ring_tokens); + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + // Contiguous [up64; gate64] W13 output -> logical h, quantized per + // 128 columns with the same power-of-two recipe as compact x. + constexpr uint32_t h_blocks = kIntermediate / 128; + for (uint64_t work = group_global; + work < static_cast(count) * h_blocks; + work += group_stride) { + const uint32_t row = work / h_blocks; + const uint32_t block = work - + static_cast(row) * h_blocks; + const uint32_t h_col = block * 128 + group_lane; + const uint32_t w13_block = h_col / 64; + const uint32_t in_block = h_col % 64; + const uint32_t physical_up = + w13_block * 128 + in_block; + const uint32_t physical_gate = physical_up + 64; + const float up = static_cast( + ring_bf16[static_cast(row) * kHidden + + physical_up]); + const float gate = static_cast( + ring_bf16[static_cast(row) * kHidden + + physical_gate]); + // Match the upstream MegaMoE default and the retained SM103 + // control's fast-math SwiGLU. Full libdevice expf over every + // routed feature is both a precision mismatch and a dominant + // large-M serialization cost. + const float sigmoid = math::fast_rcp(1.0f + __expf(-gate)); + const float h = up * gate * sigmoid; + const float amax = reduce_group_128( + cute::abs(h), storage, group_idx); + float inv; + const uint32_t packed = packed_power2_scale(amax, inv); + const uint64_t ring_index = + static_cast(row) * kIntermediate + h_col; + const uint64_t full_index = + static_cast(pool_row_offset + row) * + kIntermediate + + h_col; + ring_h[ring_index] = fp8_t(h * inv); + full_h[full_index] = fp8_t(h * inv); + if (group_lane == 0) { + const uint32_t sf_row = transform_sf_row(row); + ring_h_sf[block * sf_ring_tokens + sf_row] = packed; + full_h_sf[ + static_cast(pool_row_offset + row) * + h_blocks + + block] = packed; + } + } + __syncthreads(); + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + if (count != 0) { + run_gemm_phase< + sched::BackwardBlockPhase::W2Dgrad, kNumSMs>( + storage, expert, count, + tensor_map_ring_grad_y, + tensor_map_ring_grad_y_sf, + tensor_map_w13_recompute, tensor_map_w2_dgrad, + tensor_map_w13_dgrad, tensor_map_grad_h, + w13_scales, w2_scales, sf_ring_tokens); + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + // Exact per-route dscore = dot(W2^T dy, h) in a deterministic + // 128-thread tree. This avoids retaining the 6144-wide down output. + for (uint32_t row_work = group_global; row_work < count; + row_work += group_stride) { + float partial = 0.0f; + #pragma unroll + for (uint32_t i = 0; i < kIntermediate / 128; ++i) { + const uint32_t col = group_lane + i * 128; + const uint32_t packed = + ring_h_sf[(col / 128) * sf_ring_tokens + + transform_sf_row(row_work)]; + const uint32_t exponent = packed & 0xffu; + const float scale = + __uint_as_float(exponent << 23); + const float h = static_cast( + ring_h[static_cast(row_work) * + kIntermediate + + col]) * + scale; + const float dh = static_cast( + ring_bf16[static_cast(row_work) * kHidden + + 2 * kIntermediate + + col]); + partial = __fmaf_rn(dh, h, partial); + } + partial = reduce_group_128( + partial, storage, group_idx); + if (group_lane == 0) + ring_dscore[row_work] = partial; + } + __syncthreads(); + + constexpr uint32_t grad_blocks = (2 * kIntermediate) / 128; + for (uint64_t work = group_global; + work < static_cast(count) * grad_blocks; + work += group_stride) { + const uint32_t row = work / grad_blocks; + const uint32_t grad_block = work - + static_cast(row) * grad_blocks; + const bool gate_plane = grad_block < h_blocks; + const uint32_t block = + gate_plane ? grad_block : grad_block - h_blocks; + const uint32_t col = block * 128 + group_lane; + const uint32_t w13_block = col / 64; + const uint32_t in_block = col % 64; + const uint32_t physical_up = + w13_block * 128 + in_block; + const uint32_t physical_gate = physical_up + 64; + const float up = static_cast( + ring_bf16[static_cast(row) * kHidden + + physical_up]); + const float gate = static_cast( + ring_bf16[static_cast(row) * kHidden + + physical_gate]); + const float dy_h = static_cast( + ring_bf16[ + static_cast(row) * kHidden + + 2 * kIntermediate + + col]) * + ring_scores[row]; + const float sigmoid = math::fast_rcp(1.0f + __expf(-gate)); + const float grad_value = gate_plane + ? dy_h * up * sigmoid * + (1.0f + gate * (1.0f - sigmoid)) + : dy_h * gate * sigmoid; + const float amax = reduce_group_128( + cute::abs(grad_value), storage, group_idx); + float inv; + const uint32_t packed = packed_power2_scale(amax, inv); + const uint32_t logical_col = grad_block * 128 + group_lane; + const uint64_t ring_index = + static_cast(row) * + (2 * kIntermediate) + + logical_col; + const uint64_t full_index = + static_cast(pool_row_offset + row) * + (2 * kIntermediate) + + logical_col; + ring_grad_preact[ring_index] = fp8_t(grad_value * inv); + full_grad_preact[full_index] = fp8_t(grad_value * inv); + if (group_lane == 0) { + const uint32_t sf_row = transform_sf_row(row); + ring_grad_preact_sf[ + grad_block * sf_ring_tokens + sf_row] = packed; + full_grad_preact_sf[ + static_cast(pool_row_offset + row) * + grad_blocks + + grad_block] = packed; + } + } + __syncthreads(); + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + if (count != 0) { + run_gemm_phase< + sched::BackwardBlockPhase::W13Dgrad, kNumSMs>( + storage, expert, count, + tensor_map_ring_grad_preact, + tensor_map_ring_grad_preact_sf, + tensor_map_w13_recompute, tensor_map_w2_dgrad, + tensor_map_w13_dgrad, tensor_map_grad_x, + w13_scales, w2_scales, sf_ring_tokens); + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + // Publish each route directly into its immutable source top-k slot. + for (uint64_t linear = global_thread; + linear < static_cast(count) * vecs_per_row; + linear += global_stride) { + const uint32_t row = linear / vecs_per_row; + const uint32_t vec = linear - + static_cast(row) * vecs_per_row; + const auto metadata = + token_src_metadata[pool_row_offset + row]; + auto* remote_dst = sym_buffer.map( + reinterpret_cast(combine_slots) + + (static_cast(metadata.topk_idx) * capacity + + metadata.token_idx) * + vecs_per_row + + vec, + metadata.rank_idx); + *remote_dst = reinterpret_cast(ring_bf16)[ + static_cast(row) * vecs_per_row + vec]; + } + for (uint32_t row = global_thread; row < count; + row += global_stride) { + const auto metadata = + token_src_metadata[pool_row_offset + row]; + *sym_buffer.map( + symmetric_grad_scores + + static_cast(metadata.token_idx) * kTopK + + metadata.topk_idx, + metadata.rank_idx) = ring_dscore[row]; + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + pool_block_offset += math::ceil_div(count, kBlockM); + } + + comm::nvlink_barrier( + workspace, sym_buffer, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + for (uint64_t linear = global_thread; + linear < static_cast(num_tokens) * kHidden; + linear += global_stride) { + const uint32_t token = linear / kHidden; + const uint32_t col = linear - + static_cast(token) * kHidden; + float value = 0.0f; + #pragma unroll + for (uint32_t slot = 0; slot < kTopK; ++slot) { + value = __fadd_rn( + value, + static_cast( + combine_slots[ + (static_cast(slot) * capacity + token) * + kHidden + + col])); + } + grad_x[linear] = bf16_t(value); + } + for (uint64_t route = global_thread; + route < static_cast(num_tokens) * kTopK; + route += global_stride) { + grad_scores[route] = symmetric_grad_scores[route]; + } + __syncthreads(); + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + if (cutlass::canonical_warp_idx_sync() == 7) + cute::TMEM::Allocator2Sm().free(0, kNumTmemCols); +#endif +} + +// Reference-shaped GLM reverse. The role layout, two-CTA MMA pipeline, +// wave scheduler, ring counters, and remote publication mirror the upstream +// forward kernel. Only the operands and fused epilogue differ: reverse +// dispatch forms score-scaled dz, W2 dgrad emits quantized dpreact, and W13 +// dgrad publishes dX. Forward-saved BF16 preactivation and unweighted down +// output remove the recompute/standalone score phases entirely. +template +CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void +sm103_fp8_block128_mega_moe_backward_persistent_impl( + const int* expert_counts, + const layout::TokenSrcMetadata* token_src_metadata, + const uint32_t num_tokens, + const uint32_t capacity, + const uint32_t ring_tokens, + const uint32_t sf_ring_tokens, + const uint32_t max_pool_tokens, + const __grid_constant__ layout::SymBuffer sym_buffer, + const __grid_constant__ layout::Workspace workspace, + const bf16_t* compact_grad_y, + bf16_t* symmetric_grad_y, + bf16_t* symmetric_grad_x, + fp8_t* ring_grad_y, + uint32_t* ring_grad_y_sf, + fp8_t* ring_grad_preact, + uint32_t* ring_grad_preact_sf, + fp8_t* full_grad_y, + uint32_t* full_grad_y_sf, + const float* full_scores, + const bf16_t* saved_l1_preact, + const bf16_t* saved_down_unweighted, + fp8_t* full_grad_preact, + uint32_t* full_grad_preact_sf, + bf16_t* grad_x, + float* symmetric_grad_scores, + float* grad_scores, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_y, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_y_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_preact, + const __grid_constant__ cute::TmaDescriptor tensor_map_ring_grad_preact_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_w2_dgrad, + const __grid_constant__ cute::TmaDescriptor tensor_map_w13_dgrad, + const float* w13_scales, + const float* w2_scales) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + constexpr uint32_t kLocalExperts = kGlobalExperts / kNumRanks; + constexpr uint32_t kW2BlockNs = kIntermediate / kBlockN; + constexpr uint32_t kW13BlockNs = kHidden / kBlockN; + constexpr uint32_t kHiddenScaleBlocks = kHidden / 128; + constexpr uint32_t kGradPreactScaleBlocks = + (2 * kIntermediate) / 128; + constexpr uint32_t kDispatchThreads = kNumDispatchWarps * 32; + constexpr uint32_t kEpilogueBarrier = 9; + using Scheduler = sched::MegaMoEBackwardScheduler< + kBlockM, kBlockN, kBlockK, + kHidden, kIntermediate, + kLocalExperts, 1, kNumSMs>; + + const uint32_t warp_idx = cutlass::canonical_warp_idx_sync(); + const uint32_t lane_idx = ptx::get_lane_idx(); + const uint32_t global_thread = + blockIdx.x * kThreads + threadIdx.x; + const uint32_t global_stride = kNumSMs * kThreads; + const bool leader_cta = cute::block_rank_in_cluster() == 0; + const uint32_t num_ring_blocks = ring_tokens / kBlockM; + extern __shared__ __align__(1024) uint8_t smem_buffer[]; + SharedStorage& storage = + *reinterpret_cast(smem_buffer); + + if (warp_idx == 0) { + cute::prefetch_tma_descriptor(&tensor_map_ring_grad_y); + cute::prefetch_tma_descriptor(&tensor_map_ring_grad_y_sf); + cute::prefetch_tma_descriptor(&tensor_map_ring_grad_preact); + cute::prefetch_tma_descriptor(&tensor_map_ring_grad_preact_sf); + cute::prefetch_tma_descriptor(&tensor_map_w2_dgrad); + cute::prefetch_tma_descriptor(&tensor_map_w13_dgrad); + } + + // Match upstream's separate input and combine transport lifetimes. The + // immutable grad-y plane may be pulled throughout dispatch while W13 + // concurrently publishes dX into the independent combine plane. This is + // what permits ring generations to wrap without a global dispatch latch. + for (uint64_t linear = global_thread; + linear < static_cast(num_tokens) * kHidden; + linear += global_stride) { + symmetric_grad_y[linear] = compact_grad_y[linear]; + } + comm::nvlink_barrier( + workspace, sym_buffer, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 7 && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kStages; ++i) { + storage.full_barriers[i].init(4); + storage.empty_barriers[i].init(1); + } + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { + storage.tmem_full_barriers[i].init(1); + storage.tmem_empty_barriers[i].init( + 2 * kEpilogueThreads); + } + #pragma unroll + for (uint32_t warp = 0; warp < kNumDispatchWarps; ++warp) { + #pragma unroll + for (uint32_t stage = 0; + stage < kDispatchPullStages; ++stage) { + storage.dispatch_barriers[warp][stage].init(1); + } + } + cutlass::arch::fence_barrier_init(); + } + if (warp_idx == 7) + cute::TMEM::Allocator2Sm().allocate( + kNumTmemCols, &storage.tmem_ptr); + comm::cluster_sync_with_relaxed_arrive(); + + uint32_t stage_idx = 0; + uint32_t pipeline_phase = 0; + const auto advance_pipeline = [&]() { + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + pipeline_phase ^= stage_idx == 0; + }; + + if (warp_idx < kNumDispatchWarps) { + cutlass::arch::warpgroup_reg_dealloc<48>(); + const uint32_t dispatch_warp = warp_idx; + const uint32_t global_warp = + blockIdx.x * kNumDispatchWarps + dispatch_warp; + constexpr uint32_t kGlobalWarps = + kNumSMs * kNumDispatchWarps; + uint32_t pull_phase[kDispatchPullStages] = {}; + uint32_t pool_block_offset = 0; + + #pragma unroll 1 + for (uint32_t expert = 0; expert < kLocalExperts; ++expert) { + const uint32_t count = static_cast( + __ldg(expert_counts + expert)); + for (uint32_t row = global_warp; row < count; + row += kGlobalWarps) { + const uint32_t pool_row = + pool_block_offset * kBlockM + row; + DG_DEVICE_ASSERT(pool_row < max_pool_tokens); + const uint32_t pool_block = pool_row / kBlockM; + const uint32_t ring_block = + pool_block % num_ring_blocks; + const uint32_t ring_row = + ring_block * kBlockM + row % kBlockM; + const uint32_t empty_target = + (pool_block / num_ring_blocks) * kW2BlockNs; + if (empty_target != 0) { + while (ptx::ld_acq( + workspace.get_l1_empty_count_ptr( + ring_block)) < empty_target) { + } + } + + const auto metadata = token_src_metadata[pool_row]; + const auto* remote_grad_y = sym_buffer.map( + symmetric_grad_y + + static_cast(metadata.token_idx) * + kHidden, + metadata.rank_idx); + const float score = full_scores[pool_row]; + float dscore = 0.0f; + + // Follow the upstream persistent dispatch transport: one + // elected lane issues a chunked remote TMA into a private + // warp pull buffer while all lanes transform the preceding + // chunk. Only the precision-side work differs here: BF16 + // dy is score-scaled, reduced for dscore, and quantized to + // the existing E4M3/UE8M0 block128 ring representation. + const auto issue_pull = [&](const uint32_t chunk) { + const uint32_t pull_bytes = kDispatchPullBytes; + const uint32_t pull_stage = + chunk % kDispatchPullStages; + auto* pull_barrier = + &storage.dispatch_barriers[dispatch_warp] + [pull_stage]; + ptx::tma_load_1d( + storage.dispatch_pull[dispatch_warp] + [pull_stage], + remote_grad_y + + static_cast(chunk) * + kDispatchPullElements, + pull_barrier, pull_bytes); + ptx::mbarrier_arrive_and_set_tx( + pull_barrier, pull_bytes); + }; + const auto wait_pull = [&](const uint32_t chunk) { + const uint32_t pull_stage = + chunk % kDispatchPullStages; + ptx::mbarrier_wait_and_flip_phase( + &storage.dispatch_barriers[dispatch_warp] + [pull_stage], + pull_phase[pull_stage]); + }; + + if (cute::elect_one_sync()) { + issue_pull(0); + wait_pull(0); + } + __syncwarp(); + + constexpr uint32_t kScaleBlocksPerPull = + kDispatchPullElements / 128; + DG_STATIC_ASSERT( + kScaleBlocksPerPull * kDispatchPullChunks == + kHiddenScaleBlocks, + "reverse pull chunks must preserve block128 groups"); + #pragma unroll 1 + for (uint32_t chunk = 0; + chunk < kDispatchPullChunks; ++chunk) { + const uint32_t pull_stage = + chunk % kDispatchPullStages; + if (chunk + 1 < kDispatchPullChunks && + cute::elect_one_sync()) { + issue_pull(chunk + 1); + } + + #pragma unroll 1 + for (uint32_t chunk_block = 0; + chunk_block < kScaleBlocksPerPull; + ++chunk_block) { + const uint32_t block = + chunk * kScaleBlocksPerPull + + chunk_block; + float dz_values[4]; + float local_amax = 0.0f; + #pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const uint32_t col_in_chunk = + chunk_block * 128 + lane_idx * 4 + i; + const uint32_t col = + chunk * kDispatchPullElements + + col_in_chunk; + const float dy = static_cast( + storage.dispatch_pull[dispatch_warp] + [pull_stage] + [col_in_chunk]); + const float down = static_cast( + saved_down_unweighted[ + static_cast(pool_row) * + kHidden + + col]); + dscore = __fmaf_rn(dy, down, dscore); + dz_values[i] = static_cast( + bf16_t(__fmul_rn(dy, score))); + local_amax = cute::max( + local_amax, + cute::abs(dz_values[i])); + } + local_amax = warp_reduce_max(local_amax); + local_amax = __shfl_sync( + 0xffffffff, local_amax, 0); + float scale_inv; + const uint32_t packed_scale = + packed_power2_scale( + local_amax, scale_inv); + #pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const uint32_t col = + block * 128 + lane_idx * 4 + i; + const fp8_t value( + dz_values[i] * scale_inv); + ring_grad_y[ + static_cast(ring_row) * + kHidden + + col] = value; + full_grad_y[ + static_cast(pool_row) * + kHidden + + col] = value; + } + if (lane_idx == 0) { + ring_grad_y_sf[ + block * sf_ring_tokens + + transform_sf_row(ring_row)] = + packed_scale; + full_grad_y_sf[ + static_cast(pool_row) * + kHiddenScaleBlocks + + block] = packed_scale; + } + } + + if (chunk + 1 < kDispatchPullChunks && + cute::elect_one_sync()) { + wait_pull(chunk + 1); + } + __syncwarp(); + } + + dscore = warp_reduce_sum(dscore); + __syncwarp(); + if (lane_idx == 0) { + *sym_buffer.map( + symmetric_grad_scores + + static_cast( + metadata.token_idx) * kTopK + + metadata.topk_idx, + metadata.rank_idx) = dscore; + __threadfence(); + const bool is_last = row + 1 == count; + ptx::red_add_rel( + workspace.get_l1_full_count_ptr(ring_block), + is_last + ? kBlockM - (row % kBlockM) + : 1u); + } + __syncwarp(); + } + pool_block_offset += math::ceil_div(count, kBlockM); + } + + constexpr uint32_t kDispatchNamedBarrier = 12; + comm::nvlink_barrier< + kNumRanks, kNumSMs, kDispatchThreads, 1, 92>( + workspace, sym_buffer, blockIdx.x, + dispatch_warp * 32 + lane_idx, + []() { + ptx::sync_aligned( + kDispatchThreads, kDispatchNamedBarrier); + }); + } else if (warp_idx == 4) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + Scheduler scheduler(expert_counts); + scheduler.for_each_block( + [&](const sched::BackwardBlockPhase block_phase, + const uint32_t&, const uint32_t num_k_blocks, + const uint32_t m_block_idx, + const uint32_t&) { + const uint32_t pool_block = + scheduler.get_current_pool_block_offset() + + m_block_idx; + const uint32_t ring_block = + pool_block % num_ring_blocks; + const uint32_t generation = + pool_block / num_ring_blocks; + const uint32_t full_target = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? kBlockM * (generation + 1) + : (2 * kW2BlockNs) * (generation + 1); + auto* full_ptr = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? workspace.get_l1_full_count_ptr(ring_block) + : workspace.get_l2_full_count_ptr(ring_block); + while (ptx::ld_acq(full_ptr) != full_target) { + } + const auto* map_a = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? &tensor_map_ring_grad_y + : &tensor_map_ring_grad_preact; + const auto* map_sfa = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? &tensor_map_ring_grad_y_sf + : &tensor_map_ring_grad_preact_sf; + const uint32_t ring_m = ring_block * kBlockM; + const uint32_t sf_m = ring_block * kSFBlockM; + const uint32_t valid_m = + scheduler.template get_valid_m(); + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block, advance_pipeline()) { + storage.empty_barriers[stage_idx].wait( + pipeline_phase ^ 1u); + uint32_t m_idx = ring_m; + if (!leader_cta) + m_idx += math::align(valid_m, 16u) / 2; + if (cute::elect_one_sync()) { + tma::copy< + kBlockK, kLoadBlockM, kSwizzle, fp8_t>( + map_a, + &storage.full_barriers[stage_idx], + storage.smem_a[stage_idx], + k_block * kBlockK, m_idx, 2); + tma::copy( + map_sfa, + &storage.full_barriers[stage_idx], + storage.smem_sfa[stage_idx], + sf_m, k_block, 2); + if (leader_cta) { + storage.full_barriers[stage_idx] + .arrive_and_expect_tx( + sizeof(storage.smem_a[0]) * 2 + + sizeof(storage.smem_sfa[0]) * 2); + } else { + storage.full_barriers[stage_idx].arrive(0u); + } + } + __syncwarp(); + } + }); + } else if (warp_idx == 5) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + Scheduler scheduler(expert_counts); + scheduler.for_each_block( + [&](const sched::BackwardBlockPhase block_phase, + const uint32_t expert, const uint32_t num_k_blocks, + const uint32_t&, const uint32_t n_block) { + const auto* map_b = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? &tensor_map_w2_dgrad + : &tensor_map_w13_dgrad; + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block, advance_pipeline()) { + storage.empty_barriers[stage_idx].wait( + pipeline_phase ^ 1u); + if (cute::elect_one_sync()) { + const uint32_t outer_k = + block_phase == + sched::BackwardBlockPhase::W2Dgrad + ? expert * kHidden + + k_block * kBlockK + : expert * (2 * kIntermediate) + + k_block * kBlockK; + tma::copy< + kBlockN, kBlockK, kSwizzle, fp8_t>( + map_b, + &storage.full_barriers[stage_idx], + storage.smem_b[stage_idx], + n_block * kBlockN, outer_k, 2); + } + #pragma unroll + for (uint32_t row = lane_idx; + row < kBlockN; row += 32) { + float scale; + if (block_phase == + sched::BackwardBlockPhase::W2Dgrad) { + scale = __ldg( + w2_scales + + (expert * (kHidden / 128) + + k_block) * + (kIntermediate / 128) + + n_block); + } else { + const uint32_t plane = + k_block / (kIntermediate / 128); + const uint32_t row_block = + k_block % (kIntermediate / 128); + scale = __ldg( + w13_scales + + ((expert * 2 + plane) * + (kIntermediate / 128) + + row_block) * + (kHidden / 128) + + n_block); + } + storage.smem_sfb[stage_idx][row] = + (__float_as_uint(scale) >> 23) * + 0x01010101u; + } + __syncwarp(); + if (cute::elect_one_sync()) { + if (leader_cta) { + storage.full_barriers[stage_idx] + .arrive_and_expect_tx( + sizeof(storage.smem_b[0]) * 2); + } else { + storage.full_barriers[stage_idx].arrive(0u); + } + } + __syncwarp(); + } + }); + } else if (warp_idx == 6) { + cutlass::arch::warpgroup_reg_dealloc<40>(); + if (leader_cta) { + auto instr_desc = + cute::UMMA::make_instr_desc_block_scaled< + fp8_t, fp8_t, float, + cutlass::float_ue8m0_t, + kUMMAM, kUMMAN, + cute::UMMA::Major::MN, + cute::UMMA::Major::K>(); + auto sf_desc = mma::sm100::make_sf_desc(nullptr); + auto a_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::K, kLoadBlockM, + kBlockK, kSwizzle>(storage.smem_a[0], 0, 0); + auto b_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::MN, kLoadBlockN, + kBlockK, kSwizzle>(storage.smem_b[0], 0, 0); + const uint32_t a_desc_lo = lane_idx < kStages + ? a_desc.lo + + lane_idx * sizeof(storage.smem_a[0]) / 16 + : 0u; + const uint32_t b_desc_lo = lane_idx < kStages + ? b_desc.lo + + lane_idx * sizeof(storage.smem_b[0]) / 16 + : 0u; + uint32_t current_iter = 0; + Scheduler scheduler(expert_counts); + scheduler.for_each_block( + [&](const sched::BackwardBlockPhase, + const uint32_t&, const uint32_t num_k_blocks, + const uint32_t&, const uint32_t&) { + mma::sm100::update_instr_desc_with_umma_n( + instr_desc, + scheduler.template get_valid_m()); + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1u; + storage.tmem_empty_barriers[accum_stage].wait( + accum_phase ^ 1u); + ptx::tcgen05_after_thread_sync(); + for (uint32_t k_block = 0; + k_block < num_k_blocks; + ++k_block, advance_pipeline()) { + storage.full_barriers[stage_idx].wait( + pipeline_phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t a_base = + ptx::exchange(a_desc_lo, stage_idx); + const uint32_t b_base = + ptx::exchange(b_desc_lo, stage_idx); + if (cute::elect_one_sync()) { + using utccp_t = + cute::SM100_UTCCP_4x32dp128bit_2cta; + #pragma unroll + for (uint32_t i = 0; + i < kSFBlockM / + kUTCCPAlignedElements; + ++i) { + mma::sm100::replace_smem_desc_addr( + sf_desc, + storage.smem_sfa[stage_idx] + + i * kUTCCPAlignedElements); + utccp_t::copy( + sf_desc, kTmemSFAStart + i * 4); + } + mma::sm100::replace_smem_desc_addr( + sf_desc, + storage.smem_sfb[stage_idx]); + // Keep these device-lambda constants literal. NVCC + // otherwise ODR-uses the namespace-scope constexpr + // and looks for a device symbol during template + // instantiation. + utccp_t::copy(sf_desc, 392u); + #pragma unroll + for (uint32_t k = 0; + k < kBlockK / kUMMAK; ++k) { + const auto runtime_desc = + mma::sm100:: + make_runtime_instr_desc_with_sf_id( + instr_desc, k, k); + a_desc.lo = + mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::K, + kLoadBlockM, kSwizzle, fp8_t>( + a_base, 0, k * kUMMAK); + b_desc.lo = + mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::MN, + kLoadBlockN, kSwizzle, fp8_t>( + b_base, 0, k * kUMMAK); + ptx::SM100_MMA_MXF8F6F4_2x1SM_SS::fma( + b_desc, a_desc, + accum_stage * kUMMAN, + k_block > 0 || k > 0, + runtime_desc, 392u, 384u); + } + } + __syncwarp(); + constexpr uint16_t kCTAMask = 3; + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.empty_barriers[stage_idx]), + kCTAMask); + if (k_block + 1 == num_k_blocks) { + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.tmem_full_barriers[ + accum_stage]), + kCTAMask); + } + __syncwarp(); + } + }); + if (current_iter != 0) { + const uint32_t last = current_iter - 1; + storage.tmem_empty_barriers[ + last % kNumEpilogueStages] + .wait((last / kNumEpilogueStages) & 1u); + } + } + } else if (warp_idx >= 8 && warp_idx < 12) { + cutlass::arch::warpgroup_reg_alloc<208>(); + const uint32_t epilogue_warp = warp_idx - 8; + const uint32_t epilogue_thread = + epilogue_warp * 32 + lane_idx; + uint32_t current_iter = 0; + Scheduler scheduler(expert_counts); + scheduler.for_each_block( + [&](const sched::BackwardBlockPhase block_phase, + const uint32_t&, const uint32_t, + const uint32_t m_block, const uint32_t n_block) { + const uint32_t accum_stage = + current_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (current_iter++ / kNumEpilogueStages) & 1u; + storage.tmem_full_barriers[accum_stage].wait( + accum_phase); + ptx::tcgen05_after_thread_sync(); + + const uint32_t pool_block = + scheduler.get_current_pool_block_offset() + + m_block; + const uint32_t ring_block = + pool_block % num_ring_blocks; + const uint32_t generation = + pool_block / num_ring_blocks; + const uint32_t ring_m = ring_block * kBlockM; + const uint32_t pool_m = pool_block * kBlockM; + const uint32_t valid_m = + scheduler.template get_valid_m(); + if (block_phase == + sched::BackwardBlockPhase::W2Dgrad) { + const uint32_t empty_target = + kW13BlockNs * generation; + while (ptx::ld_acq( + workspace.get_l2_empty_count_ptr( + ring_block)) != empty_target) { + } + } + const uint32_t aligned_m = math::align(valid_m, 16u); + const uint32_t num_stores = aligned_m / kStoreBlockM; + constexpr uint32_t kRowsPerLoad = 8; + constexpr uint32_t kLoadsPerStore = + kStoreBlockM / kRowsPerLoad; + DG_STATIC_ASSERT( + kLoadsPerStore == 2, + "reverse fused epilogue assumes two 8-row TMEM atoms"); + if (block_phase == + sched::BackwardBlockPhase::W2Dgrad) { + // Consume W2 dgrad directly from TMEM. Each epilogue + // thread owns one of the 128 logical hidden columns and + // eight token rows per TMEM load. Round the accumulator + // to BF16 in registers, exactly as the generic store + // epilogue did, then fuse SwiGLU backward and E4M3/UE8M0 + // quantization. There is no BF16 global-ring store/reload + // boundary between W2 dgrad and W13 dgrad. + const uint32_t hidden_col = + n_block * kBlockN + epilogue_thread; + for (uint32_t store = 0; store < num_stores; + ++store) { + #pragma unroll + for (uint32_t atom = 0; + atom < kLoadsPerStore; ++atom) { + const uint32_t row_base = + store * kStoreBlockM + + atom * kRowsPerLoad; + const uint32_t tmem_addr = + accum_stage * kUMMAN + row_base; + uint32_t values[kRowsPerLoad]; + cute::SM100_TMEM_LOAD_16dp256b1x::copy( + tmem_addr, + values[0], values[1], values[2], values[3]); + cute::SM100_TMEM_LOAD_16dp256b1x::copy( + tmem_addr | 0x00100000, + values[4], values[5], values[6], values[7]); + cutlass::arch::fence_view_async_tmem_load(); + if (store + 1 == num_stores && + atom + 1 == kLoadsPerStore) { + ptx::tcgen05_before_thread_sync(); + storage.tmem_empty_barriers[accum_stage] + .arrive(0u); + } + + uint32_t packed_dh[kRowsPerLoad / 2]; + #pragma unroll + for (uint32_t pair = 0; + pair < kRowsPerLoad / 2; ++pair) { + packed_dh[pair] = + math::cast_into_bf16_and_pack( + values[pair * 2], + values[pair * 2 + 1]); + } + #pragma unroll + for (uint32_t row_in_atom = 0; + row_in_atom < kRowsPerLoad; + ++row_in_atom) { + const uint32_t row = + row_base + row_in_atom; + const auto dh_pair = + __bfloat1622float2( + *reinterpret_cast( + &packed_dh[row_in_atom / 2])); + const float dh = (row_in_atom & 1u) + ? dh_pair.y : dh_pair.x; + const uint32_t ring_row = ring_m + row; + const uint32_t pool_row = pool_m + row; + const bool valid = row < valid_m; + const float gate = valid + ? static_cast( + saved_l1_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + hidden_col]) + : 0.0f; + const float up = valid + ? static_cast( + saved_l1_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + kIntermediate + hidden_col]) + : 0.0f; + const float sigmoid = + math::fast_rcp(1.0f + __expf(-gate)); + const float dgate = valid + ? dh * up * sigmoid * + (1.0f + gate * (1.0f - sigmoid)) + : 0.0f; + const float dup = valid + ? dh * gate * sigmoid : 0.0f; + const float2 amax = + reduce_group_128_max_pair( + {cute::abs(dgate), cute::abs(dup)}, + storage, 2); + float gate_inv, up_inv; + const uint32_t gate_scale = + packed_power2_scale( + amax.x, gate_inv); + const uint32_t up_scale = + packed_power2_scale( + amax.y, up_inv); + if (valid) { + const uint32_t gate_col = hidden_col; + const uint32_t up_col = + kIntermediate + hidden_col; + const fp8_t gate_value( + dgate * gate_inv); + const fp8_t up_value(dup * up_inv); + ring_grad_preact[ + static_cast(ring_row) * + (2 * kIntermediate) + + gate_col] = gate_value; + ring_grad_preact[ + static_cast(ring_row) * + (2 * kIntermediate) + + up_col] = up_value; + full_grad_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + gate_col] = gate_value; + full_grad_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + up_col] = up_value; + if (epilogue_thread == 0) { + const uint32_t sf_row = + transform_sf_row(ring_row); + ring_grad_preact_sf[ + n_block * sf_ring_tokens + + sf_row] = gate_scale; + ring_grad_preact_sf[ + (kW2BlockNs + n_block) * + sf_ring_tokens + + sf_row] = up_scale; + full_grad_preact_sf[ + static_cast(pool_row) * + kGradPreactScaleBlocks + + n_block] = gate_scale; + full_grad_preact_sf[ + static_cast(pool_row) * + kGradPreactScaleBlocks + + kW2BlockNs + n_block] = up_scale; + } + } + } + } + } + ptx::sync_aligned( + 128u, kEpilogueBarrier); + if (epilogue_warp == 0 && + cute::elect_one_sync()) { + __threadfence(); + ptx::red_add_rel( + workspace.get_l2_full_count_ptr( + ring_block), + 2u); + ptx::red_add( + workspace.get_l1_empty_count_ptr( + ring_block), + 1u); + } + __syncwarp(); + } else { + // Mirror upstream's L2 epilogue: round W13 dgrad to BF16 + // while moving TMEM into one shared store tile, then publish + // that tile directly to the owning rank's combine plane. + // The reusable global BF16 ring is deliberately absent. + constexpr uint32_t kBankGroupBytes = 16; + constexpr uint32_t kVecElements = + sizeof(uint4) / sizeof(bf16_t); + constexpr uint32_t kVecsPerNBlock = + kBlockN / kVecElements; + constexpr uint32_t kRowsPerWarp = + kStoreBlockM / 8; + for (uint32_t store = 0; store < num_stores; + ++store) { + #pragma unroll + for (uint32_t atom = 0; + atom < kLoadsPerStore; ++atom) { + const uint32_t row_base = + store * kStoreBlockM + + atom * kRowsPerLoad; + const uint32_t tmem_addr = + accum_stage * kUMMAN + row_base; + uint32_t values[kRowsPerLoad]; + cute::SM100_TMEM_LOAD_16dp256b1x::copy( + tmem_addr, + values[0], values[1], values[2], values[3]); + cute::SM100_TMEM_LOAD_16dp256b1x::copy( + tmem_addr | 0x00100000, + values[4], values[5], values[6], values[7]); + cutlass::arch::fence_view_async_tmem_load(); + const uint32_t outer_atom_offset = + (epilogue_warp / 2) * + kStoreBlockM * kSwizzle; + const uint32_t inner_atom_offset = + atom * kRowsPerLoad * kSwizzle; + auto* smem_base = + reinterpret_cast( + storage.smem_cd[0]) + + outer_atom_offset + inner_atom_offset; + const uint32_t smem_row = lane_idx % 8; + const uint32_t smem_col = + (epilogue_warp % 2) * 4 + lane_idx / 8; + auto* smem_ptr = smem_base + + smem_row * (kBankGroupBytes * 8) + + (smem_col ^ smem_row) * kBankGroupBytes; + ptx::SM90_U32x4_STSM_T::copy( + math::cast_into_bf16_and_pack( + values[0], values[1]), + math::cast_into_bf16_and_pack( + values[2], values[3]), + math::cast_into_bf16_and_pack( + values[4], values[5]), + math::cast_into_bf16_and_pack( + values[6], values[7]), + smem_ptr); + if (store + 1 == num_stores && + atom + 1 == kLoadsPerStore) { + ptx::tcgen05_before_thread_sync(); + storage.tmem_empty_barriers[accum_stage] + .arrive(0u); + } + } + ptx::sync_aligned(128u, kEpilogueBarrier); + + const uint32_t row_in_atom = + (epilogue_warp * 2 + lane_idx / 16) % 8; + const uint32_t bank_group = lane_idx % 8; + #pragma unroll + for (uint32_t j = 0; j < kRowsPerWarp; ++j) { + const uint32_t row_in_store = + j * 8 + epilogue_warp * 2 + lane_idx / 16; + const uint32_t row = + store * kStoreBlockM + row_in_store; + if (row < valid_m) { + const uint32_t pool_row = pool_m + row; + const auto metadata = + token_src_metadata[pool_row]; + const uint32_t vec = lane_idx % 16; + const uint32_t source_vec = + n_block * kVecsPerNBlock + vec; + const auto* smem_ptr = + reinterpret_cast( + storage.smem_cd[0]) + + (lane_idx % 16 / 8) * + kStoreBlockM * kSwizzle + + row_in_store * kSwizzle + + (bank_group ^ row_in_atom) * + kBankGroupBytes; + const uint4 value = ptx::ld_shared( + reinterpret_cast( + smem_ptr)); + auto* remote = sym_buffer.map( + reinterpret_cast( + symmetric_grad_x + + (static_cast( + metadata.topk_idx) * + capacity + + metadata.token_idx) * + kHidden) + + source_vec, + metadata.rank_idx); + *remote = value; + } + } + ptx::sync_aligned(128u, kEpilogueBarrier); + } + if (epilogue_warp == 0 && + cute::elect_one_sync()) { + __threadfence_system(); + ptx::red_add_rel( + workspace.get_l2_empty_count_ptr( + ring_block), + 1u); + } + __syncwarp(); + } + }); + } else { + cutlass::arch::warpgroup_reg_dealloc<40>(); + } + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 7) + cute::TMEM::Allocator2Sm().free(0, kNumTmemCols); + + comm::nvlink_barrier( + workspace, sym_buffer, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + for (uint64_t linear = global_thread; + linear < static_cast(num_tokens) * kHidden; + linear += global_stride) { + const uint32_t token = linear / kHidden; + const uint32_t col = linear - + static_cast(token) * kHidden; + float value = 0.0f; + #pragma unroll + for (uint32_t slot = 0; slot < kTopK; ++slot) { + value = __fadd_rn( + value, + static_cast( + symmetric_grad_x[ + (static_cast(slot) * capacity + + token) * + kHidden + + col])); + } + grad_x[linear] = bf16_t(value); + } + for (uint64_t route = global_thread; + route < static_cast(num_tokens) * kTopK; + route += global_stride) { + grad_scores[route] = symmetric_grad_scores[route]; + } + + for (uint32_t block = global_thread; + block < num_ring_blocks; block += global_stride) { + *workspace.get_l1_full_count_ptr(block) = 0u; + *workspace.get_l1_empty_count_ptr(block) = 0u; + *workspace.get_l2_full_count_ptr(block) = 0u; + *workspace.get_l2_empty_count_ptr(block) = 0u; + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); +#endif +} + +} // namespace sm103_block128_backward + +} // namespace deep_gemm diff --git a/deep_gemm/include/deep_gemm/impls/sm103_fp8_block128_mega_moe_wgrad.cuh b/deep_gemm/include/deep_gemm/impls/sm103_fp8_block128_mega_moe_wgrad.cuh new file mode 100644 index 0000000000..445b18680a --- /dev/null +++ b/deep_gemm/include/deep_gemm/impls/sm103_fp8_block128_mega_moe_wgrad.cuh @@ -0,0 +1,517 @@ +#pragma once +#pragma clang diagnostic push +#pragma clang diagnostic ignored "-Wunknown-attributes" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace deep_gemm::sm103_block128_wgrad { + +// The two CTAs form the same logical 256x256x64 production BF16 tile used by +// DeepGEMM's grouped large-M kernels. Its two load warps read canonical E4M3 +// route tiles and FP32 power-of-two scales directly, materialize the exact BF16 +// operands in the native UMMA swizzle, and publish the stage to the MMA warp. +// No full-route BF16 backing or grid-wide conversion phase exists. +static constexpr uint32_t kHidden = 6144; +static constexpr uint32_t kIntermediate = 2048; +static constexpr uint32_t kGlobalExperts = 256; +static constexpr uint32_t kRouteBlockM = 192; +static constexpr uint32_t kBlockM = 128; +static constexpr uint32_t kBlockN = 256; +static constexpr uint32_t kBlockK = 64; +static constexpr uint32_t kLoadBlockN = kBlockN / 2; +static constexpr uint32_t kStages = 3; +static constexpr uint32_t kLoadAWarp = 0; +static constexpr uint32_t kMMAWarp = 1; +static constexpr uint32_t kLoadBWarp = 2; +static constexpr uint32_t kControlWarp = 3; +static constexpr uint32_t kEpilogueFirstWarp = 4; +static constexpr uint32_t kEpilogueThreads = 128; +static constexpr uint32_t kThreads = + kEpilogueFirstWarp * 32 + kEpilogueThreads; +static constexpr uint32_t kNumEpilogueStages = 2; +static constexpr uint32_t kNumTMAStoreStages = 2; +static constexpr uint32_t kStoreBlockM = 128; +static constexpr uint32_t kStoreBlockN = 64; +static constexpr uint32_t kUMMAM = 256; +static constexpr uint32_t kUMMAN = 256; +static constexpr uint32_t kUMMAK = 16; +static constexpr uint32_t kSwizzle = 128; +static constexpr uint32_t kNumTmemAccumCols = + kNumEpilogueStages * kUMMAN; +static constexpr uint32_t kNumTmemCols = + utils::get_num_aligned_tmem_cols(); + +using fp8_t = cutlass::float_e4m3_t; +using bf16_t = cutlass::bfloat16_t; +using Barrier = cutlass::arch::ClusterTransactionBarrier; + +struct alignas(1024) SharedStorage { + alignas(1024) bf16_t smem_cd[kNumTMAStoreStages] + [kStoreBlockM * kStoreBlockN]; + alignas(1024) bf16_t smem_a[kStages][kBlockK * kBlockM]; + alignas(1024) bf16_t smem_b[kStages][kBlockK * kLoadBlockN]; + Barrier load_full_barriers[kStages]; + Barrier load_empty_barriers[kStages]; + Barrier tmem_full_barriers[kNumEpilogueStages]; + Barrier tmem_empty_barriers[kNumEpilogueStages]; + uint32_t tmem_ptr; +}; + +DG_STATIC_ASSERT(kThreads == 256, "SM103 wgrad role layout changed"); +DG_STATIC_ASSERT(kLoadBlockN == kBlockM, + "wgrad A/B prologues must share one tile shape"); +DG_STATIC_ASSERT(kNumTmemCols == 512, "SM103 wgrad TMEM layout changed"); + +CUTLASS_DEVICE uint16_t fold_power2_scale_into_bf16( + const uint16_t half_bits, + const uint32_t scale_exponent) { + const uint16_t sign = half_bits & 0x8000u; + const uint32_t half_exponent = (half_bits >> 10) & 0x1fu; + if (half_exponent == 0u) + return sign; + if (half_exponent == 0x1fu) + return sign | 0x7fc0u; + + const uint32_t mantissa = (half_bits >> 3) & 0x7fu; + const int32_t bf16_exponent = + static_cast(half_exponent) + + static_cast(scale_exponent) - 15; + if (bf16_exponent >= 0xff) + return sign | 0x7f80u; + if (bf16_exponent > 0) + return sign | + static_cast(bf16_exponent << 7) | + static_cast(mantissa); + + // The E4M3 significand has only four bits, so folding a power-of-two scale + // into BF16 is exact except when the result reaches BF16's subnormal range. + // Reproduce round-to-nearest-even there without materializing FP32. + const uint32_t significand = 0x80u | mantissa; + const uint32_t shift = static_cast(1 - bf16_exponent); + if (shift > 8u) + return sign; + const uint32_t truncated = significand >> shift; + const uint32_t remainder = + significand & ((1u << shift) - 1u); + const uint32_t halfway = 1u << (shift - 1u); + const uint32_t rounded = truncated + + (remainder > halfway || + (remainder == halfway && (truncated & 1u))); + return sign | static_cast(rounded); +} + +CUTLASS_DEVICE uint32_t convert_fp8x2_power2_to_bf16x2( + const uint16_t fp8x2, + const uint32_t scale_exponent) { + uint32_t half2; + asm("cvt.rn.f16x2.e4m3x2 %0, %1;\n" + : "=r"(half2) : "h"(fp8x2)); + const uint32_t lo = fold_power2_scale_into_bf16( + static_cast(half2), scale_exponent); + const uint32_t hi = fold_power2_scale_into_bf16( + static_cast(half2 >> 16), scale_exponent); + return lo | (hi << 16); +} + +CUTLASS_DEVICE uint32_t get_bf16_mn_swizzled_pair_offset( + const uint32_t mn, const uint32_t k) { + const uint32_t row = mn & 7u; + const uint32_t col_byte = k * sizeof(bf16_t); + return (mn >> 3) * 8u * kSwizzle + + row * kSwizzle + + ((col_byte >> 4) ^ row) * 16u + + (col_byte & 15u); +} + +template +CUTLASS_DEVICE void load_dequant_bf16_tile( + const fp8_t* source, + const uint32_t* scales, + const uint32_t route_base, + const uint32_t feature_block, + const uint32_t valid_k, + bf16_t* destination, + const uint32_t lane_idx) { + constexpr uint32_t kScaleBlocksPerRow = kShape / 128; + DG_STATIC_ASSERT(kShape % 128 == 0, + "wgrad dequant shape must be block128 aligned"); + DG_STATIC_ASSERT(kBlockM == 128 && kLoadBlockN == 128, + "wgrad tiled dequant requires one block128 feature group"); + + auto* destination_bytes = reinterpret_cast(destination); + #pragma unroll 1 + for (uint32_t local_k = 0; local_k < kBlockK; ++local_k) { + const bool valid = local_k < valid_k; + uint32_t scale_exponent = 0; + if (lane_idx == 0 && valid) { + scale_exponent = __ldg( + scales + + static_cast(route_base + local_k) * + kScaleBlocksPerRow + + feature_block) & 0xffu; + } + scale_exponent = __shfl_sync( + 0xffffffffu, scale_exponent, 0); + + #pragma unroll + for (uint32_t pair_group = 0; pair_group < 2; ++pair_group) { + const uint32_t pair = lane_idx + pair_group * 32u; + const uint32_t local_mn = pair * 2u; + uint32_t bf16x2 = 0; + if (valid) { + const uint16_t fp8x2 = + *reinterpret_cast( + source + + static_cast(route_base + local_k) * + kShape + + feature_block * 128u + local_mn); + bf16x2 = convert_fp8x2_power2_to_bf16x2( + fp8x2, scale_exponent); + } + #pragma unroll + for (uint32_t value = 0; value < 2; ++value) { + *reinterpret_cast( + destination_bytes + + get_bf16_mn_swizzled_pair_offset( + local_mn + value, local_k)) = + static_cast( + bf16x2 >> (value * 16u)); + } + } + } +} + +// This is the production grouped-GEMM L2 swizzle specialized to the fixed GLM +// shapes. Eight adjacent M blocks sweep all N blocks before moving to the next +// M group. The group size and every expert's tile count are even, so adjacent +// physical CTAs always remain the two M halves of one 2-CTA output tile. +template +struct WgradTileScheduler { + static constexpr uint32_t kLocalExperts = + kGlobalExperts / kNumRanks; + static constexpr uint32_t kShapeM = + kW2 ? kHidden : 2 * kIntermediate; + static constexpr uint32_t kShapeN = + kW2 ? kIntermediate : kHidden; + static constexpr uint32_t kNumMBlocks = kShapeM / kBlockM; + static constexpr uint32_t kNumNBlocks = kShapeN / kBlockN; + static constexpr uint32_t kMBlocksPerL2Group = 8; + static constexpr uint32_t kTilesPerL2Group = + kMBlocksPerL2Group * kNumNBlocks; + static constexpr uint32_t kTilesPerExpert = + kNumMBlocks * kNumNBlocks; + static constexpr uint32_t kTotalTiles = + kLocalExperts * kTilesPerExpert; + + const int* expert_counts; + uint32_t linear_tile = blockIdx.x; + uint32_t cached_expert = 0; + uint32_t cached_pool_row = 0; + + CUTLASS_DEVICE explicit WgradTileScheduler(const int* counts) + : expert_counts(counts) { + DG_STATIC_ASSERT(kNumMBlocks % kMBlocksPerL2Group == 0, + "fixed GLM M shape must fit the L2 swizzle"); + DG_STATIC_ASSERT(kTilesPerL2Group % 2 == 0, + "L2 groups must preserve 2-CTA pairing"); + DG_STATIC_ASSERT(kNumSMs % 2 == 0, + "2-CTA wgrad requires an even SM count"); + } + + CUTLASS_DEVICE bool get_next( + uint32_t& expert, + uint32_t& count, + uint32_t& pool_row, + uint32_t& m_block, + uint32_t& n_block) { + if (linear_tile >= kTotalTiles) + return false; + expert = linear_tile / kTilesPerExpert; + const uint32_t expert_tile = + linear_tile - expert * kTilesPerExpert; + while (cached_expert < expert) { + const uint32_t previous_count = static_cast( + __ldg(expert_counts + cached_expert)); + cached_pool_row += + math::ceil_div(previous_count, kRouteBlockM) * + kRouteBlockM; + ++cached_expert; + } + count = static_cast(__ldg(expert_counts + expert)); + pool_row = cached_pool_row; + const uint32_t l2_group = expert_tile / kTilesPerL2Group; + const uint32_t tile_in_group = + expert_tile - l2_group * kTilesPerL2Group; + m_block = l2_group * kMBlocksPerL2Group + + tile_in_group % kMBlocksPerL2Group; + n_block = tile_in_group / kMBlocksPerL2Group; + linear_tile += kNumSMs; + return true; + } +}; + +template +CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void +sm103_fp8_block128_mega_moe_wgrad_impl( + const int* expert_counts, + const uint32_t max_pool_tokens, + const fp8_t* full_a, + const fp8_t* full_b, + const uint32_t* full_a_sf, + const uint32_t* full_b_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_output_0, + const __grid_constant__ cute::TmaDescriptor tensor_map_output_1) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == 1030 + using Scheduler = WgradTileScheduler; + constexpr uint32_t kShapeM = Scheduler::kShapeM; + constexpr uint32_t kShapeN = Scheduler::kShapeN; + + extern __shared__ __align__(1024) uint8_t smem_buffer[]; + SharedStorage& storage = *reinterpret_cast(smem_buffer); + const uint32_t warp_idx = cutlass::canonical_warp_idx_sync(); + const uint32_t lane_idx = ptx::get_lane_idx(); + const bool leader_cta = cute::block_rank_in_cluster() == 0; + const uint32_t cta_rank = cute::block_rank_in_cluster(); + + if (warp_idx == kLoadAWarp) { + cute::prefetch_tma_descriptor(&tensor_map_output_0); + cute::prefetch_tma_descriptor(&tensor_map_output_1); + } + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == kControlWarp && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kStages; ++i) { + // One A producer and one B producer in each CTA publish directly + // to CTA 0, matching upstream's four-arrival 2-CTA load stage. + storage.load_full_barriers[i].init(4); + storage.load_empty_barriers[i].init(1); + } + #pragma unroll + for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { + storage.tmem_full_barriers[i].init(1); + storage.tmem_empty_barriers[i].init( + 2 * kEpilogueThreads); + } + cutlass::arch::fence_barrier_init(); + } + __syncwarp(); + if (warp_idx == kControlWarp) + cute::TMEM::Allocator2Sm().allocate( + kNumTmemCols, &storage.tmem_ptr); + comm::cluster_sync_with_relaxed_arrive(); + + auto instr_desc = cute::UMMA::make_instr_desc< + bf16_t, bf16_t, float, + kUMMAM, kUMMAN, + cute::UMMA::Major::MN, + cute::UMMA::Major::MN>(); + auto a_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::MN, kBlockM, kBlockK, kSwizzle>( + storage.smem_a[0], 0, 0); + auto b_desc = mma::sm100::make_umma_desc< + cute::UMMA::Major::MN, kLoadBlockN, kBlockK, kSwizzle>( + storage.smem_b[0], 0, 0); + const uint32_t a_desc_lo = lane_idx < kStages + ? a_desc.lo + lane_idx * sizeof(storage.smem_a[0]) / 16 + : 0u; + const uint32_t b_desc_lo = lane_idx < kStages + ? b_desc.lo + lane_idx * sizeof(storage.smem_b[0]) / 16 + : 0u; + const auto runtime_instr_desc = + cute::UMMA::make_runtime_instr_desc(instr_desc); + auto smem_cd = utils::PatternVisitor([&](const uint32_t& i) { + return storage.smem_cd[i]; + }); + + uint32_t stage_idx = 0; + uint32_t phase = 0; + auto advance_pipeline = [&]() { + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + }; + + if (warp_idx == kLoadAWarp || warp_idx == kLoadBWarp) { + // Preserve upstream's two producer-warps-per-CTA pipeline. Each + // producer loads and dequantizes one canonical FP8 operand directly + // into its native BF16 UMMA tile; no route-pool backing is formed. + const bool load_a = warp_idx == kLoadAWarp; + Scheduler scheduler(expert_counts); + uint32_t expert, count, pool_row, m_block, n_block; + while (scheduler.get_next( + expert, count, pool_row, m_block, n_block)) { + DG_DEVICE_ASSERT(pool_row + count <= max_pool_tokens); + const uint32_t num_k_blocks = + cute::max(1u, math::ceil_div(count, kBlockK)); + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block) { + storage.load_empty_barriers[stage_idx].wait( + phase ^ 1u); + const uint32_t valid_k = count == 0 + ? 0u + : cute::min( + 64u, + count - k_block * kBlockK); + if (load_a) { + load_dequant_bf16_tile( + full_a, full_a_sf, + pool_row + k_block * kBlockK, + m_block, + valid_k, + storage.smem_a[stage_idx], + lane_idx); + } else { + const uint32_t b_feature_block = + n_block * (kBlockN / 128u) + cta_rank; + load_dequant_bf16_tile( + full_b, full_b_sf, + pool_row + k_block * kBlockK, + b_feature_block, + valid_k, + storage.smem_b[stage_idx], + lane_idx); + } + cutlass::arch::fence_view_async_shared(); + __syncwarp(); + if (cute::elect_one_sync()) + storage.load_full_barriers[stage_idx].arrive(0u); + advance_pipeline(); + } + } + } else if (warp_idx == kMMAWarp && leader_cta) { + Scheduler scheduler(expert_counts); + uint32_t expert, count, pool_row, m_block, n_block; + uint32_t output_iter = 0; + while (scheduler.get_next( + expert, count, pool_row, m_block, n_block)) { + const uint32_t accum_stage = + output_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (output_iter / kNumEpilogueStages) & 1u; + ++output_iter; + storage.tmem_empty_barriers[accum_stage].wait( + accum_phase ^ 1u); + ptx::tcgen05_after_thread_sync(); + + const uint32_t num_k_blocks = + cute::max(1u, math::ceil_div(count, kBlockK)); + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block) { + storage.load_full_barriers[stage_idx].wait(phase); + ptx::tcgen05_after_thread_sync(); + const uint32_t a_base = + ptx::exchange(a_desc_lo, stage_idx); + const uint32_t b_base = + ptx::exchange(b_desc_lo, stage_idx); + if (cute::elect_one_sync()) { + #pragma unroll + for (uint32_t k = 0; k < kBlockK / kUMMAK; ++k) { + a_desc.lo = mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::MN, kBlockM, + kSwizzle, bf16_t>( + a_base, 0, k * kUMMAK); + b_desc.lo = mma::sm100::advance_umma_desc_lo< + cute::UMMA::Major::MN, kLoadBlockN, + kSwizzle, bf16_t>( + b_base, 0, k * kUMMAK); + ptx::SM100_MMA_F16BF16_2x1SM_SS::fma( + a_desc, b_desc, + accum_stage * kUMMAN, + k_block > 0 || k > 0, + runtime_instr_desc); + } + } + __syncwarp(); + constexpr uint16_t kCTAMask = 3; + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.load_empty_barriers[stage_idx]), + kCTAMask); + if (k_block == num_k_blocks - 1) { + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.tmem_full_barriers[accum_stage]), + kCTAMask); + } + __syncwarp(); + advance_pipeline(); + } + } + } else if ( + warp_idx >= kEpilogueFirstWarp && + warp_idx < kEpilogueFirstWarp + kEpilogueThreads / 32) { + Scheduler scheduler(expert_counts); + const uint32_t epilogue_warp_idx = + warp_idx - kEpilogueFirstWarp; + uint32_t expert, count, pool_row, m_block, n_block; + uint32_t output_iter = 0; + uint32_t tma_stage_idx = 0; + DG_TRAP_ONLY_DEVICE_ASSERT( + ptx::ld_shared(&storage.tmem_ptr) == 0); + while (scheduler.get_next( + expert, count, pool_row, m_block, n_block)) { + const uint32_t accum_stage = + output_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (output_iter / kNumEpilogueStages) & 1u; + ++output_iter; + storage.tmem_full_barriers[accum_stage].wait(accum_phase); + ptx::tcgen05_after_thread_sync(); + + const cute::TmaDescriptor* output_map = + &tensor_map_output_0; + uint32_t output_m = + expert * (kW2 ? kHidden : kIntermediate); + if constexpr (kW2) { + output_m += m_block * kBlockM; + } else { + const uint32_t plane = + m_block / (kIntermediate / kBlockM); + const uint32_t plane_m_block = + m_block % (kIntermediate / kBlockM); + output_map = plane == 0 + ? &tensor_map_output_0 + : &tensor_map_output_1; + output_m += plane_m_block * kBlockM; + } + epilogue::sm100_store_cd< + kBlockM, kBlockN, + kStoreBlockM, kStoreBlockN, + kSwizzle, kNumTMAStoreStages, kEpilogueThreads, + GemmType::Normal, false, bf16_t, + epilogue::transform::EpilogueIdentity>( + smem_cd, tma_stage_idx, + accum_stage * kUMMAN, + output_m, n_block * kBlockN, 0, + epilogue_warp_idx, lane_idx, + &storage.tmem_empty_barriers[accum_stage], + *output_map); + } + if (epilogue_warp_idx == 0) + cute::tma_store_wait<0>(); + __syncwarp(); + } + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == kControlWarp) + cute::TMEM::Allocator2Sm().free(0, kNumTmemCols); +#else + if (blockIdx.x == 0 && threadIdx.x == 0) + DG_DEVICE_ASSERT(false && "SM103 MegaMoE wgrad has no fallback"); +#endif +} + +} // namespace deep_gemm::sm103_block128_wgrad + +#pragma clang diagnostic pop diff --git a/deep_gemm/include/deep_gemm/impls/smxx_layout.cuh b/deep_gemm/include/deep_gemm/impls/smxx_layout.cuh index bf9495e9bf..63a3ad0832 100644 --- a/deep_gemm/include/deep_gemm/impls/smxx_layout.cuh +++ b/deep_gemm/include/deep_gemm/impls/smxx_layout.cuh @@ -52,7 +52,8 @@ CUTLASS_GLOBAL void transpose_fp32(const float* sf, float* out, const uint32_t m // NOTES: the two kernels below always pack the K dimension template + uint32_t kNumPsumGroups = 1, bool kUsePsumLayout = false, + bool kSFColumnMajor = false> CUTLASS_GLOBAL void transpose_and_pack_fp32_into_ue8m0(float* sf, uint32_t* out, const uint32_t mn, const uint32_t* grouped_layout, const uint32_t m_alignment) { extern __shared__ uint32_t smem_buffer[]; @@ -94,23 +95,37 @@ CUTLASS_GLOBAL void transpose_and_pack_fp32_into_ue8m0(float* sf, uint32_t* out, }; // Shift into the group - sf = sf + static_cast(blockIdx.y) * mn * SF_K; + sf = sf + static_cast(blockIdx.y) * + (kSFColumnMajor ? tma_aligned_mn * SF_K : mn * SF_K); out = out + static_cast(blockIdx.y) * tma_aligned_mn * kNumPackedSFK; // Load FP32 SFs DG_STATIC_ASSERT(BLOCK_MN % 4 == 0, "Invalid block size"); - const auto local_sf = reinterpret_cast(sf + static_cast(blockIdx.x) * (BLOCK_MN * SF_K)); const auto num_values = in_block_mn * SF_K; - const auto num_uint4 = num_values / 4; - #pragma unroll - for (uint32_t i = threadIdx.x; i < num_uint4; i += kNumThreads) { - const auto& [x, y, z, w] = reinterpret_cast(local_sf)[i]; - ptx::st_shared(reinterpret_cast(sf_smem_buffer) + i, x, y, z, w); - } + if constexpr (kSFColumnMajor) { + const auto sf_bits = reinterpret_cast(sf); + const auto global_mn_start = blockIdx.x * BLOCK_MN; + #pragma unroll + for (uint32_t i = threadIdx.x; i < num_values; i += kNumThreads) { + const auto local_mn_idx = i / SF_K, sf_k_idx = i % SF_K; + ptx::st_shared( + sf_smem_buffer + i, + sf_bits[sf_k_idx * tma_aligned_mn + global_mn_start + local_mn_idx]); + } + } else { + const auto local_sf = reinterpret_cast( + sf + static_cast(blockIdx.x) * (BLOCK_MN * SF_K)); + const auto num_uint4 = num_values / 4; + #pragma unroll + for (uint32_t i = threadIdx.x; i < num_uint4; i += kNumThreads) { + const auto& [x, y, z, w] = reinterpret_cast(local_sf)[i]; + ptx::st_shared(reinterpret_cast(sf_smem_buffer) + i, x, y, z, w); + } - // Fill unaligned values as well - if (const auto unaligned_idx = num_uint4 * 4 + threadIdx.x; unaligned_idx < num_values) - ptx::st_shared(sf_smem_buffer + unaligned_idx, local_sf[unaligned_idx]); + // Fill unaligned values as well + if (const auto unaligned_idx = num_uint4 * 4 + threadIdx.x; unaligned_idx < num_values) + ptx::st_shared(sf_smem_buffer + unaligned_idx, local_sf[unaligned_idx]); + } __syncthreads(); // Pack into UE8M0 and store diff --git a/deep_gemm/include/deep_gemm/layout/mega_moe.cuh b/deep_gemm/include/deep_gemm/layout/mega_moe.cuh index 93694befbd..626bbac0e3 100644 --- a/deep_gemm/include/deep_gemm/layout/mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/layout/mega_moe.cuh @@ -102,6 +102,11 @@ struct Workspace { // Combine push source indices (full) num_bytes += num_max_pool_tokens * sizeof(TokenSrcMetadata); + // POST_DOWN route weights (full). The reusable L1 ring is released + // before L2 remote combine, so route weights needed by that epilogue + // must follow the logical pool-token lifetime instead of the ring. + num_bytes += num_max_pool_tokens * sizeof(float); + // Align to TMA descriptor requirements num_bytes = math::align(num_bytes, 16); return num_bytes; @@ -192,6 +197,15 @@ struct Workspace { const auto base = reinterpret_cast(get_src_token_topk_idx_ptr(num_experts_per_rank)); return base + pool_token_idx; } + + // Full-pool POST_DOWN state, matching upstream persistent MegaMoE. This + // is private symmetric-workspace storage derived from context/CP capacity. + CUTLASS_DEVICE + float* get_route_weight_ptr(const uint32_t& pool_token_idx = 0) const { + const auto base = reinterpret_cast( + get_token_src_metadata_ptr(num_max_pool_tokens)); + return base + pool_token_idx; + } }; struct Data { diff --git a/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh index 05021ec89f..52a5f6de5a 100644 --- a/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh @@ -16,6 +16,201 @@ enum class BlockPhase { Linear2 = 2 }; +// Fixed training reverse for the GLM large-M specialization. It is the same +// two-phase wave state machine as MegaMoEScheduler, with immutable expert +// counts saved by forward instead of live dispatch counters. Reverse L1 is W2 +// dgrad; its fused epilogue publishes quantized dpreact into the L2 ring, and +// reverse L2 is W13 dgrad. Ring readiness counters carry the dependency, so +// CTAs never serialize phases behind a grid-wide barrier. +enum class BackwardBlockPhase { + None = 0, + W2Dgrad = 1, + W13Dgrad = 2, + // Legacy helper selector only. The persistent scheduler never emits a + // recompute phase; forward-saved BF16 preactivation is authoritative. + RecomputeW13 = 3, +}; + +template < + uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K, + uint32_t kHidden, uint32_t kIntermediateHidden, + uint32_t kNumExpertsPerRank, uint32_t kNumExpertsPerWave, + uint32_t kNumSMs, + uint32_t kNumExpertsPerLane = + math::constexpr_ceil_div(kNumExpertsPerRank, 32u)> +struct MegaMoEBackwardScheduler { + static constexpr uint32_t kW2DgradBlockNs = + kIntermediateHidden / BLOCK_N; + static constexpr uint32_t kW13DgradBlockNs = + kHidden / BLOCK_N; + static constexpr uint32_t kW2DgradBlockKs = kHidden / BLOCK_K; + static constexpr uint32_t kW13DgradBlockKs = + (2 * kIntermediateHidden) / BLOCK_K; + + DG_STATIC_ASSERT(kNumSMs % 2 == 0, + "Backward 2-CTA scheduler requires an even SM count"); + DG_STATIC_ASSERT(kW2DgradBlockNs % 2 == 0 && + kW13DgradBlockNs % 2 == 0, + "Every backward phase must assign adjacent N blocks to a cluster"); + DG_STATIC_ASSERT(kNumExpertsPerWave > 0 && + kNumExpertsPerWave <= kNumExpertsPerRank, + "Invalid backward wave size"); + + const int* expert_counts; + BackwardBlockPhase phase = BackwardBlockPhase::W2Dgrad; + uint32_t current_local_expert_idx = 0; + uint32_t current_pool_block_offset = 0; + uint32_t current_num_tokens = 0; + uint32_t block_idx = 0; + uint32_t m_block_idx = 0; + uint32_t n_block_idx = 0; + uint32_t stored_num_tokens_per_expert[kNumExpertsPerLane] = {}; + + CUTLASS_DEVICE explicit MegaMoEBackwardScheduler(const int* counts) + : expert_counts(counts), block_idx(blockIdx.x) { + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++i) { + const uint32_t expert_idx = i * 32 + ptx::get_lane_idx(); + stored_num_tokens_per_expert[i] = + expert_idx < kNumExpertsPerRank + ? static_cast(__ldg(expert_counts + expert_idx)) + : 0u; + } + __syncwarp(); + set_expert_idx(0); + } + + CUTLASS_DEVICE uint32_t get_wave_expert_end_idx() const { + return cute::min( + math::align(current_local_expert_idx + 1, + kNumExpertsPerWave), + kNumExpertsPerRank); + } + + CUTLASS_DEVICE uint32_t get_num_tokens( + const uint32_t expert_idx) const { + uint32_t value = 0; + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++i) { + if (expert_idx == i * 32 + ptx::get_lane_idx()) + value = stored_num_tokens_per_expert[i]; + } + return ptx::exchange(value, expert_idx % 32); + } + + CUTLASS_DEVICE uint32_t get_pool_block_offset( + const uint32_t expert_idx) const { + uint32_t blocks = 0; + #pragma unroll + for (uint32_t i = 0; i < kNumExpertsPerLane; ++i) { + if (i * 32 + ptx::get_lane_idx() < expert_idx) { + blocks += math::ceil_div( + stored_num_tokens_per_expert[i], BLOCK_M); + } + } + return __reduce_add_sync(0xffffffff, blocks); + } + + 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); + } + + CUTLASS_DEVICE void advance_expert_idx() { + current_pool_block_offset += get_current_num_m_blocks(); + ++current_local_expert_idx; + if (current_local_expert_idx < kNumExpertsPerRank) + current_num_tokens = get_num_tokens(current_local_expert_idx); + } + + 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 value = cute::min( + current_num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + return kDoUMMAAligned ? math::align(value, 16u) : value; + } + + CUTLASS_DEVICE bool fetch_next_w2_block() { + const uint32_t wave_end = get_wave_expert_end_idx(); + while (current_local_expert_idx < wave_end) { + const uint32_t num_m_blocks = get_current_num_m_blocks(); + m_block_idx = block_idx / kW2DgradBlockNs; + if (m_block_idx < num_m_blocks) + return true; + block_idx -= num_m_blocks * kW2DgradBlockNs; + advance_expert_idx(); + } + return false; + } + + CUTLASS_DEVICE bool fetch_next_w13_block() { + const uint32_t wave_end = get_wave_expert_end_idx(); + while (current_local_expert_idx < wave_end) { + const uint32_t num_m_blocks = get_current_num_m_blocks(); + if (block_idx < num_m_blocks * kW13DgradBlockNs) { + m_block_idx = block_idx / kW13DgradBlockNs; + return true; + } + block_idx -= num_m_blocks * kW13DgradBlockNs; + advance_expert_idx(); + } + return false; + } + + CUTLASS_DEVICE cute::tuple + get_next_block() { + while (current_local_expert_idx < kNumExpertsPerRank) { + if (phase == BackwardBlockPhase::W2Dgrad) { + if (fetch_next_w2_block()) { + n_block_idx = + block_idx - m_block_idx * kW2DgradBlockNs; + block_idx += kNumSMs; + return {phase, current_local_expert_idx, + m_block_idx, n_block_idx}; + } + phase = BackwardBlockPhase::W13Dgrad; + set_expert_idx(math::align( + current_local_expert_idx - 1, + kNumExpertsPerWave)); + } else if (fetch_next_w13_block()) { + n_block_idx = + block_idx - m_block_idx * kW13DgradBlockNs; + block_idx += kNumSMs; + return {phase, current_local_expert_idx, + m_block_idx, n_block_idx}; + } else { + phase = BackwardBlockPhase::W2Dgrad; + } + } + return {BackwardBlockPhase::None, 0, 0, 0}; + } + + template + CUTLASS_DEVICE void for_each_block(Func&& func) { + while (true) { + CUTE_TIE_DECL(get_next_block(), current_phase, expert_idx, + current_m_block_idx, current_n_block_idx); + if (current_phase == BackwardBlockPhase::None) + break; + const uint32_t num_k_blocks = + current_phase == BackwardBlockPhase::W2Dgrad + ? kW2DgradBlockKs + : kW13DgradBlockKs; + func(current_phase, expert_idx, num_k_blocks, + current_m_block_idx, current_n_block_idx); + } + } +}; + template dict[str, Any]: + """Return a non-launching exact capability manifest for preflight.""" + native = dict(_C.get_sm103_fp8_block128_capabilities()) + missing = [name for name in REQUIRED_NATIVE_SYMBOLS if not hasattr(_C, name)] + try: + import torch.distributed._symmetric_memory as symm_mem + + if not callable(getattr(symm_mem, "empty", None)) or not callable( + getattr(symm_mem, "rendezvous", None) + ): + missing.append("torch.distributed._symmetric_memory") + except (AttributeError, ImportError): + missing.append("torch.distributed._symmetric_memory") + native.update( + { + "native_symbols": REQUIRED_NATIVE_SYMBOLS, + "python_symbols": REQUIRED_PYTHON_SYMBOLS, + "forward": not missing, + "backward": not missing, + "distributed_transport": "persistent_symmetric_ring", + "transport_layout": "ring_l1_l2_wave", + "transport_scale_layout": "ue8m0_power2_group128", + "transport_deterministic": True, + "combine_reductions": 1, + "wgrad_backend": "two_persistent_fused_dequant_bf16", + "supported_ep": _SUPPORTED_EP, + "missing_symbols": tuple(missing), + } + ) + return native + + +def transform_glm_w13_for_fp8_block128_mega_moe( + canonical_weight: torch.Tensor, + canonical_scale: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Validate and return canonical ``[gate, up]`` storage without a copy. + + The persistent scheduler issues independent TMA offsets for gate and up + and presents ``[up; gate]`` only as a logical MMA/output order. Returning + the original tensors is intentional and is part of the no-repack ABI. + """ + if canonical_weight.ndim != 3 or canonical_scale.ndim != 3: + raise ValueError("canonical W13 weight and scale must both be rank 3") + if canonical_weight.shape[0] % 2: + raise ValueError("canonical W13 must contain gate/up pairs") + experts = canonical_weight.shape[0] // 2 + hidden, model_dim = canonical_weight.shape[1:] + expected_scale = (experts * 2, hidden // _BLOCK, model_dim // _BLOCK) + if hidden % _BLOCK or model_dim % _BLOCK: + raise ValueError("canonical W13 dimensions must be divisible by 128") + if tuple(canonical_scale.shape) != expected_scale: + raise ValueError( + f"canonical W13 scale shape must be {expected_scale}, " + f"got {tuple(canonical_scale.shape)}" + ) + if not canonical_weight.is_contiguous() or not canonical_scale.is_contiguous(): + raise ValueError("canonical W13 q/s must be contiguous") + return canonical_weight, canonical_scale + + +@dataclass(frozen=True) +class _GroupState: + group: Any + rank: int + world_size: int + + +def _resolve_group(group: Any) -> _GroupState: + if not dist.is_available() or not dist.is_initialized(): + raise RuntimeError("persistent MegaMoE requires initialized torch.distributed") + resolved = dist.group.WORLD if group is None else group + state = _GroupState( + group=resolved, + rank=dist.get_rank(resolved), + world_size=dist.get_world_size(resolved), + ) + if state.world_size not in _SUPPORTED_EP: + raise RuntimeError( + f"persistent GLM MegaMoE supports EP{_SUPPORTED_EP}, " + f"got EP{state.world_size}" + ) + return state + + +def _check_tensor( + tensor: torch.Tensor, + *, + name: str, + ndim: int, + dtype: torch.dtype, + device: torch.device, +) -> None: + if tensor.ndim != ndim: + raise ValueError(f"{name} must be rank {ndim}, got rank {tensor.ndim}") + if tensor.dtype != dtype: + raise TypeError(f"{name} must have dtype {dtype}, got {tensor.dtype}") + if tensor.device != device: + raise ValueError(f"{name} must be on {device}, got {tensor.device}") + if not tensor.is_contiguous(): + raise ValueError(f"{name} must be contiguous") + + +def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: + to_local = getattr(tensor, "to_local", None) + return to_local() if callable(to_local) else tensor + + +def _validate_master_tensor( + tensor: torch.Tensor, + *, + name: str, + local_shape: tuple[int, int, int], + global_shape: tuple[int, int, int], + device: torch.device, + master_gradient_wrapper: Any, +) -> None: + local = _local_tensor(tensor) + _check_tensor(local, name=name, ndim=3, dtype=torch.bfloat16, device=device) + if tuple(local.shape) == local_shape: + return + is_distributed = callable(getattr(tensor, "to_local", None)) + if not is_distributed or master_gradient_wrapper is None: + raise ValueError( + f"{name} must have resident shape {local_shape}, got {tuple(local.shape)}" + ) + if tuple(tensor.shape) != global_shape: + raise ValueError( + f"{name} distributed logical shape must be {global_shape}, " + f"got {tuple(tensor.shape)}" + ) + if local.numel() == 0: + raise ValueError(f"{name} distributed local shard must be resident") + + +def _validate_inputs( + x: torch.Tensor, + topk_ids: torch.Tensor, + topk_scores: torch.Tensor, + w13_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_scale: torch.Tensor, + w1_master: torch.Tensor, + w2_master: torch.Tensor, + w3_master: torch.Tensor, + group_state: _GroupState, + master_gradient_wrapper: Any, +) -> None: + if not x.is_cuda: + raise ValueError("FP8-block128 MegaMoE requires CUDA") + capability = torch.cuda.get_device_capability(x.device) + if capability != (10, 3): + raise RuntimeError( + f"FP8-block128 MegaMoE is SM103-only; runtime capability is " + f"{capability}. No fallback is available." + ) + if master_gradient_wrapper is not None and not callable(master_gradient_wrapper): + raise TypeError("master_gradient_wrapper must be callable or None") + device = x.device + _check_tensor(x, name="x", ndim=2, dtype=torch.bfloat16, device=device) + _check_tensor(topk_ids, name="topk_ids", ndim=2, dtype=torch.int64, device=device) + _check_tensor( + topk_scores, + name="topk_scores", + ndim=2, + dtype=torch.float32, + device=device, + ) + _check_tensor( + w13_weight, + name="w13_weight", + ndim=3, + dtype=torch.float8_e4m3fn, + device=device, + ) + _check_tensor( + w13_scale, + name="w13_scale", + ndim=3, + dtype=torch.float32, + device=device, + ) + _check_tensor( + w2_weight, + name="w2_weight", + ndim=3, + dtype=torch.float8_e4m3fn, + device=device, + ) + _check_tensor( + w2_scale, + name="w2_scale", + ndim=3, + dtype=torch.float32, + device=device, + ) + tokens = x.shape[0] + if tuple(x.shape[1:]) != (_MODEL_DIM,): + raise ValueError(f"persistent GLM MegaMoE requires token width {_MODEL_DIM}") + if tuple(topk_ids.shape) != (tokens, _TOPK) or topk_scores.shape != topk_ids.shape: + raise ValueError(f"top-k IDs/scores must have shape [tokens, {_TOPK}]") + + local_experts = _GLOBAL_EXPERTS // group_state.world_size + expected_w13 = (2 * local_experts, _INTERMEDIATE, _MODEL_DIM) + expected_w13_scale = (2 * local_experts, _INTERMEDIATE // 128, _MODEL_DIM // 128) + expected_w2 = (local_experts, _MODEL_DIM, _INTERMEDIATE) + expected_w2_scale = (local_experts, _MODEL_DIM // 128, _INTERMEDIATE // 128) + if tuple(w13_weight.shape) != expected_w13: + raise ValueError(f"canonical W13 must have shape {expected_w13}") + if tuple(w13_scale.shape) != expected_w13_scale: + raise ValueError(f"canonical W13 scales must have shape {expected_w13_scale}") + if tuple(w2_weight.shape) != expected_w2: + raise ValueError(f"W2 must have shape {expected_w2}") + if tuple(w2_scale.shape) != expected_w2_scale: + raise ValueError(f"W2 scales must have shape {expected_w2_scale}") + + local_w13 = (local_experts, _INTERMEDIATE, _MODEL_DIM) + global_w13 = (_GLOBAL_EXPERTS, _INTERMEDIATE, _MODEL_DIM) + local_w2 = (local_experts, _MODEL_DIM, _INTERMEDIATE) + global_w2 = (_GLOBAL_EXPERTS, _MODEL_DIM, _INTERMEDIATE) + _validate_master_tensor( + w1_master, + name="w1_master", + local_shape=local_w13, + global_shape=global_w13, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) + _validate_master_tensor( + w2_master, + name="w2_master", + local_shape=local_w2, + global_shape=global_w2, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) + _validate_master_tensor( + w3_master, + name="w3_master", + local_shape=local_w13, + global_shape=global_w13, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) + if topk_ids.numel(): + torch._assert_async( + ((topk_ids >= 0) & (topk_ids < _GLOBAL_EXPERTS)).all(), + f"top-k IDs must lie in [0, {_GLOBAL_EXPERTS})", + ) + + +@dataclass(frozen=True) +class _PersistentBufferState: + buffer: torch.Tensor + handle: Any + buffer_ptrs: tuple[int, ...] + rank: int + context_tokens_per_rank: int + workspace_info: dict[str, Any] + + +_persistent_context_tokens_per_rank: dict[int, int] = {} +_persistent_buffers: dict[tuple[int, int], _PersistentBufferState] = {} + + +def _configure_fp8_block128_mega_moe_transport( + group: Any, + *, + context_tokens_per_rank: int, +) -> None: + """Install the model's existing ``seq_len / CP`` envelope once. + + This internal setup hook is not an operation flag. FireTitan derives the + value from its model context and CP degree; callers cannot tune route-pool + capacity, and the hot path never resizes or runs a sizing collective. + """ + if ( + isinstance(context_tokens_per_rank, bool) + or not isinstance(context_tokens_per_rank, int) + or context_tokens_per_rank < 1 + ): + raise ValueError("context_tokens_per_rank must be a positive integer") + group_state = _resolve_group(group) + key = id(group_state.group) + existing = _persistent_context_tokens_per_rank.get(key) + if existing is not None and existing != context_tokens_per_rank: + raise RuntimeError( + "MegaMoE context/CP envelope changed after setup: " + f"configured={context_tokens_per_rank}, existing={existing}" + ) + if existing is None and any(buffer_key[0] == key for buffer_key in _persistent_buffers): + raise RuntimeError("MegaMoE context must be installed before arena allocation") + _persistent_context_tokens_per_rank.setdefault(key, context_tokens_per_rank) + + +def _get_persistent_buffer( + group_state: _GroupState, + *, + device: torch.device, + tokens: int, +) -> _PersistentBufferState: + context_tokens = _persistent_context_tokens_per_rank.get(id(group_state.group)) + if context_tokens is None: + raise RuntimeError("MegaMoE was not configured from the owning model context") + if tokens > context_tokens: + raise RuntimeError( + "SM103 MegaMoE input exceeds the derived seq_len/CP envelope: " + f"actual={tokens}, envelope={context_tokens}" + ) + device_index = device.index if device.index is not None else torch.cuda.current_device() + key = (id(group_state.group), device_index) + cached = _persistent_buffers.get(key) + if cached is not None: + return cached + + import torch.distributed._symmetric_memory as symm_mem + + info = dict( + _C.sm103_fp8_block128_persistent_workspace_info( + group_state.world_size, context_tokens + ) + ) + buffer = symm_mem.empty( + int(info["num_bytes"]), dtype=torch.int8, device=device + ) + handle = symm_mem.rendezvous(buffer, group=group_state.group) + buffer.zero_() + # One setup synchronization publishes the fixed arena before any layer can + # use remote pointers. There is no per-call capacity synchronization. + dist.barrier(group=group_state.group, device_ids=[device_index]) + torch.cuda.synchronize(device) + pointers = tuple(int(pointer) for pointer in handle.buffer_ptrs) + if len(pointers) != group_state.world_size: + raise RuntimeError("symmetric-memory rendezvous returned an invalid rank map") + state = _PersistentBufferState( + buffer=buffer, + handle=handle, + buffer_ptrs=pointers, + rank=group_state.rank, + context_tokens_per_rank=context_tokens, + workspace_info=info, + ) + _persistent_buffers[key] = state + return state + + +class _FP8Block128MegaMoEPersistent(torch.autograd.Function): + @staticmethod + def forward( + ctx: Any, + x: torch.Tensor, + topk_ids: torch.Tensor, + topk_scores: torch.Tensor, + w13_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_scale: torch.Tensor, + w1_master: torch.Tensor, + w2_master: torch.Tensor, + w3_master: torch.Tensor, + group: Any, + master_gradient_wrapper: Any, + ) -> torch.Tensor: + group_state = _resolve_group(group) + _validate_inputs( + x, + topk_ids, + topk_scores, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + w1_master, + w2_master, + w3_master, + group_state, + master_gradient_wrapper, + ) + state = _get_persistent_buffer( + group_state, device=x.device, tokens=x.shape[0] + ) + with torch.autograd.profiler.record_function( + "sm103_fp8_block128_megamoe_persistent_forward" + ): + _C.sm103_fp8_block128_prepare_persistent_inputs( + state.buffer, + x, + topk_ids, + topk_scores, + group_state.world_size, + state.context_tokens_per_rank, + ) + output, expert_counts, token_src_metadata = ( + _C.sm103_fp8_block128_persistent_forward( + state.buffer, + list(state.buffer_ptrs), + state.rank, + state.context_tokens_per_rank, + x.shape[0], + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) + ) + ctx.buffer_state = state + ctx.master_gradient_wrapper = master_gradient_wrapper + ctx.save_for_backward( + x, + topk_scores, + expert_counts, + token_src_metadata, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) + return output + + @staticmethod + def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: + ( + x, + topk_scores, + expert_counts, + token_src_metadata, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) = ctx.saved_tensors + state = ctx.buffer_state + grad_output = grad_output.contiguous() + if grad_output.dtype != torch.bfloat16: + grad_output = grad_output.to(torch.bfloat16) + with torch.autograd.profiler.record_function( + "sm103_fp8_block128_megamoe_persistent_backward" + ): + grad_x, grad_scores, grad_w1, grad_w2, grad_w3 = ( + _C.sm103_fp8_block128_persistent_backward( + state.buffer, + list(state.buffer_ptrs), + state.rank, + state.context_tokens_per_rank, + x, + grad_output, + topk_scores, + expert_counts, + token_src_metadata, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) + ) + wrapper = ctx.master_gradient_wrapper + if wrapper is not None: + grad_w1, grad_w2, grad_w3 = wrapper(grad_w1, grad_w2, grad_w3) + return ( + grad_x, + None, + grad_scores, + None, + None, + None, + None, + grad_w1, + grad_w2, + grad_w3, + None, + None, + ) + + +def fp8_block128_mega_moe( + x: torch.Tensor, + topk_ids: torch.Tensor, + topk_scores: torch.Tensor, + w13_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_scale: torch.Tensor, + w1_master: torch.Tensor, + w2_master: torch.Tensor, + w3_master: torch.Tensor, + group: Any = None, + master_gradient_wrapper: Any = None, +) -> torch.Tensor: + """Run GLM's complete routed branch in the fixed SM103 pipeline. + + Workspace capacity comes only from the model's preinstalled + ``seq_len / CP`` envelope. The operation exposes no route-pool bound, + performs no arena growth, and consumes canonical gate/up expert pairs + directly with separate TMA offsets. + """ + group_state = _resolve_group(group) + return _FP8Block128MegaMoEPersistent.apply( + x, + topk_ids, + topk_scores, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + w1_master, + w2_master, + w3_master, + group_state.group, + master_gradient_wrapper, + ) diff --git a/setup.py b/setup.py index c4d74ae929..6f48ca4782 100644 --- a/setup.py +++ b/setup.py @@ -14,7 +14,7 @@ from setuptools.command.build_py import build_py from packaging.version import parse from pathlib import Path -from torch.utils.cpp_extension import CUDAExtension, CUDA_HOME +from torch.utils.cpp_extension import BuildExtension, CUDAExtension, CUDA_HOME from wheel.bdist_wheel import bdist_wheel as _bdist_wheel from scripts.generate_pyi import generate_pyi_file @@ -24,21 +24,54 @@ DG_USE_LOCAL_VERSION = int(os.getenv('DG_USE_LOCAL_VERSION', '1')) == 1 DG_JIT_USE_RUNTIME_API = int(os.environ.get('DG_JIT_USE_RUNTIME_API', '0')) == 1 + +def get_source_git_commit() -> str: + """Return the exact source OID embedded in the native extension. + + Release builders without a ``.git`` directory must provide the same value + explicitly. A shortened or symbolic revision is deliberately rejected: + consumers use this value as a fail-closed ABI/provenance boundary. + """ + revision = os.environ.get('DEEP_GEMM_GIT_COMMIT') + if revision is None: + revision = subprocess.check_output( + ['git', '-C', os.path.dirname(os.path.realpath(__file__)), 'rev-parse', 'HEAD'], + text=True, + ).strip() + if re.fullmatch(r'[0-9a-f]{40}', revision) is None: + raise RuntimeError( + 'DEEP_GEMM_GIT_COMMIT must be the full 40-character lowercase source OID; ' + f'got {revision!r}' + ) + return revision + + +SOURCE_GIT_COMMIT = get_source_git_commit() +GIT_COMMIT_DEFINE = f'-DDEEP_GEMM_GIT_COMMIT_TOKEN={SOURCE_GIT_COMMIT}' + # Compiler flags cxx_flags = ['-std=c++17', '-O3', '-fPIC', '-Wno-psabi', '-Wno-deprecated-declarations', - f'-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}'] + f'-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}', + GIT_COMMIT_DEFINE] +nvcc_flags = [ + '-std=c++17', '-O3', '--expt-relaxed-constexpr', '--expt-extended-lambda', + '--generate-code=arch=compute_103a,code=sm_103a', + f'-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}', + GIT_COMMIT_DEFINE, +] if DG_JIT_USE_RUNTIME_API: cxx_flags.append('-DDG_JIT_USE_RUNTIME_API') # Sources current_dir = os.path.dirname(os.path.realpath(__file__)) -sources = ['csrc/python_api.cpp'] +sources = ['csrc/python_api.cpp', 'csrc/sm103_fp8_block128.cu'] build_include_dirs = [ f'{CUDA_HOME}/include', f'{CUDA_HOME}/include/cccl', - 'deep_gemm/include', - 'third-party/cutlass/include', - 'third-party/fmt/include', + os.path.join(current_dir, 'deep_gemm/include'), + os.path.join(current_dir, 'third-party/cutlass/include'), + os.path.join(current_dir, 'third-party/cutlass/tools/util/include'), + os.path.join(current_dir, 'third-party/fmt/include'), ] build_libraries = ['cudart', 'nvrtc'] build_library_dirs = [f'{CUDA_HOME}/lib64'] @@ -108,7 +141,7 @@ def get_ext_modules(): include_dirs=build_include_dirs, libraries=build_libraries, library_dirs=build_library_dirs, - extra_compile_args=cxx_flags)] + extra_compile_args={'cxx': cxx_flags, 'nvcc': nvcc_flags})] class CustomBuildPy(build_py): @@ -209,6 +242,7 @@ def run(self): zip_safe=False, cmdclass={ 'build_py': CustomBuildPy, + 'build_ext': BuildExtension, 'bdist_wheel': CachedWheelsCommand, }, ) diff --git a/tests/benchmark_fp8_block128_mega_moe.py b/tests/benchmark_fp8_block128_mega_moe.py new file mode 100644 index 0000000000..23e092c4fb --- /dev/null +++ b/tests/benchmark_fp8_block128_mega_moe.py @@ -0,0 +1,250 @@ +"""Structured SM103 performance smoke for FP8-block128 MegaMoE. + +The defaults are intentionally small enough for routine companion validation. +Use GLM-5.2 dimensions explicitly for integration evidence, for example:: + + torchrun --standalone --nproc-per-node 2 \ + tests/benchmark_fp8_block128_mega_moe.py \ + --tokens 15625 --experts 256 --model-dim 6144 --hidden 2048 --topk 8 + +``--experts`` is the global expert count; every rank owns an equal contiguous +shard. The process(es) must be launched on otherwise idle SM103 GPUs. Rank zero +prints one JSON object with every rank's metrics so callers can archive the +exact dimensions and straggler behavior. +""" + +from __future__ import annotations + +import argparse +import json +import os +import statistics +import time +from typing import Callable + +import torch +import torch.distributed as dist + +import deep_gemm + + +def _blockwise_quantize(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + groups, rows, columns = weight.shape + blocks = ( + weight.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + ) + scales = blocks.abs().amax(dim=(-1, -2)) / 448.0 + scales = torch.where(scales == 0, torch.ones_like(scales), scales) + quantized = ( + (blocks / scales[..., None, None]) + .to(torch.float8_e4m3fn) + .permute(0, 1, 3, 2, 4) + .reshape_as(weight) + .contiguous() + ) + return quantized, scales.contiguous() + + +def _elapsed_ms(operation: Callable[[], None], iterations: int) -> list[float]: + measurements: list[float] = [] + for _ in range(iterations): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + operation() + end.record() + end.synchronize() + measurements.append(float(start.elapsed_time(end))) + return measurements + + +def _summary(values: list[float]) -> dict[str, float]: + ordered = sorted(values) + p95_index = max(0, min(len(ordered) - 1, int(0.95 * len(ordered) + 0.999999) - 1)) + return { + "minimum_ms": min(values), + "median_ms": statistics.median(values), + "p95_ms": ordered[p95_index], + "maximum_ms": max(values), + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--tokens", type=int, default=1024) + parser.add_argument("--experts", type=int, default=8) + parser.add_argument("--topk", type=int, default=8) + parser.add_argument("--model-dim", type=int, default=1024) + parser.add_argument("--hidden", type=int, default=512) + parser.add_argument("--warmup", type=int, default=2) + parser.add_argument("--iterations", type=int, default=5) + args = parser.parse_args() + + world_size = int(os.environ.get("WORLD_SIZE", "1")) + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + rank = int(os.environ.get("RANK", "0")) + if world_size > 1: + torch.cuda.set_device(local_rank) + dist.init_process_group("nccl") + if not torch.cuda.is_available(): + raise RuntimeError("the performance benchmark requires CUDA") + capability = torch.cuda.get_device_capability() + if capability != (10, 3): + raise RuntimeError( + f"the performance benchmark is SM103-only; runtime capability is {capability}" + ) + for name in ("model_dim", "hidden"): + if getattr(args, name) <= 0 or getattr(args, name) % 128: + raise ValueError( + f"--{name.replace('_', '-')} must be a positive multiple of 128" + ) + if args.tokens <= 0 or args.experts <= 0 or args.topk <= 0: + raise ValueError("tokens, experts, and topk must be positive") + if args.topk > args.experts: + raise ValueError("topk cannot exceed experts") + if args.experts % world_size: + raise ValueError("global experts must be divisible by world size") + if args.warmup < 0 or args.iterations <= 0: + raise ValueError("warmup must be non-negative and iterations must be positive") + + torch.manual_seed(20260721 + rank) + device = torch.device("cuda", local_rank) + local_experts = args.experts // world_size + x = ( + torch.randn(args.tokens, args.model_dim, device=device, dtype=torch.bfloat16) + * 0.02 + ).requires_grad_() + w1_master = ( + torch.randn( + local_experts, + args.hidden, + args.model_dim, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + w3_master = ( + torch.randn( + local_experts, + args.hidden, + args.model_dim, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + w2_master = ( + torch.randn( + local_experts, + args.model_dim, + args.hidden, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + canonical_w13 = torch.stack( + (w1_master.detach(), w3_master.detach()), dim=1 + ).flatten(0, 1) + w13_q, w13_s = _blockwise_quantize(canonical_w13) + w2_q, w2_s = _blockwise_quantize(w2_master.detach()) + routes = torch.arange(args.tokens * args.topk, device=device, dtype=torch.int64) + topk_ids = torch.remainder(routes * 17 + 3, args.experts).view( + args.tokens, args.topk + ) + raw_scores = torch.sigmoid( + torch.randn(args.tokens, args.topk, device=device, dtype=torch.float32) + ) + scores = ( + (raw_scores / raw_scores.sum(dim=-1, keepdim=True) * 2.5) + .detach() + .requires_grad_() + ) + upstream = torch.randn_like(x) + + latest_output: torch.Tensor | None = None + + def forward() -> None: + nonlocal latest_output + latest_output = deep_gemm.fp8_block128_mega_moe( + x, + topk_ids, + scores, + w13_q, + w13_s, + w2_q, + w2_s, + w1_master, + w2_master, + w3_master, + group=dist.group.WORLD if world_size > 1 else None, + ) + + def forward_backward() -> None: + forward() + assert latest_output is not None + latest_output.backward(upstream) + x.grad = None + scores.grad = None + w1_master.grad = None + w2_master.grad = None + w3_master.grad = None + + for _ in range(args.warmup): + forward_backward() + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + started = time.monotonic() + forward_times = _elapsed_ms(forward, args.iterations) + # Do not charge destruction of the final forward-only autograd graph to + # the first complete forward/backward sample. + latest_output = None + torch.cuda.synchronize() + forward_backward_times = _elapsed_ms(forward_backward, args.iterations) + wall_seconds = time.monotonic() - started + + rank_result = { + "rank": rank, + "forward": _summary(forward_times), + "forward_backward": _summary(forward_backward_times), + "peak_allocated_bytes": torch.cuda.max_memory_allocated(), + "peak_reserved_bytes": torch.cuda.max_memory_reserved(), + "wall_seconds": wall_seconds, + } + rank_results: list[dict[str, object] | None] | None = None + if world_size > 1: + rank_results = [None] * world_size if rank == 0 else None + dist.gather_object(rank_result, rank_results, dst=0) + else: + rank_results = [rank_result] + result = { + "schema_version": 1, + "backend": "fp8_block128_mega_moe", + "architecture": "sm103", + "device": torch.cuda.get_device_name(), + "compute_capability": list(capability), + "deep_gemm_git_commit": deep_gemm.__git_commit__, + "dimensions": { + "tokens": args.tokens, + "experts": args.experts, + "local_experts": local_experts, + "topk": args.topk, + "model_dim": args.model_dim, + "hidden": args.hidden, + }, + "warmup": args.warmup, + "iterations": args.iterations, + "world_size": world_size, + "ranks": rank_results, + } + if rank == 0: + print(json.dumps(result, sort_keys=True)) + if world_size > 1: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/tests/test_fp8_block128_capabilities.py b/tests/test_fp8_block128_capabilities.py new file mode 100644 index 0000000000..774819ed70 --- /dev/null +++ b/tests/test_fp8_block128_capabilities.py @@ -0,0 +1,38 @@ +"""Non-launching provenance and capability checks for the SM103 backend.""" + +from __future__ import annotations + +import re + +import deep_gemm + + +def test_fp8_block128_capability_manifest_is_exact_and_fail_closed() -> None: + assert re.fullmatch(r"[0-9a-f]{40}", deep_gemm.__git_commit__) + + capabilities = deep_gemm.get_fp8_block128_mega_moe_capabilities() + assert capabilities["name"] == "fp8_block128_mega_moe" + assert capabilities["architecture"] == "sm103" + assert capabilities["compute_capability"] == (10, 3) + assert capabilities["activation_dtype"] == "float8_e4m3fn" + assert capabilities["weight_dtype"] == "float8_e4m3fn" + assert capabilities["scale_dtype"] == "float32" + assert capabilities["activation_group_k"] == 128 + assert capabilities["weight_block_m"] == 128 + assert capabilities["weight_block_k"] == 128 + assert capabilities["route_score_placement"] == "post_down" + assert capabilities["distributed_transport"] == "deep_ep.ElasticBuffer.expanded" + assert capabilities["transport_layout"] == "expert_aligned_128" + assert capabilities["transport_scale_layout"] == "row_major_fp32_group128" + assert capabilities["transport_deterministic"] is True + assert capabilities["combine_reductions"] == 1 + assert capabilities["wgrad_backend"] == "sm103_companion_grouped_bf16" + assert capabilities["forward"] is True + assert capabilities["backward"] is True + assert capabilities["missing_symbols"] == () + assert capabilities["fallback"] is None + + for symbol in capabilities["native_symbols"]: + assert hasattr(deep_gemm._C, symbol) + for symbol in capabilities["python_symbols"]: + assert hasattr(deep_gemm, symbol) diff --git a/tests/test_fp8_block128_mega_moe.py b/tests/test_fp8_block128_mega_moe.py new file mode 100644 index 0000000000..04c4d84250 --- /dev/null +++ b/tests/test_fp8_block128_mega_moe.py @@ -0,0 +1,358 @@ +import pytest +import torch + +import deep_gemm +from deep_gemm.mega.fp8_block128 import ( + _bf16_grouped_wgrad_expanded, + _validate_master_tensor, +) + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 3), + reason="SM103-only MegaMoE test", +) + + +def _blockwise_weight_quantize( + weight: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + groups, rows, columns = weight.shape + blocks = ( + weight.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + ) + scales = blocks.abs().amax(dim=(-1, -2)) / 448.0 + scales = torch.where(scales == 0, torch.ones_like(scales), scales) + quantized = ( + (blocks / scales[..., None, None]) + .to(torch.float8_e4m3fn) + .permute(0, 1, 3, 2, 4) + .reshape_as(weight) + .contiguous() + ) + return quantized, scales.contiguous() + + +def _blockwise_weight_dequantize( + quantized: torch.Tensor, scales: torch.Tensor +) -> torch.Tensor: + groups, rows, columns = quantized.shape + return ( + quantized.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + .mul(scales[..., None, None]) + .permute(0, 1, 3, 2, 4) + .reshape(groups, rows, columns) + ) + + +def _make_case(tokens: int, *, all_to_one: bool = False) -> dict[str, torch.Tensor]: + torch.manual_seed(7000 + tokens + int(all_to_one)) + experts, model_dim, hidden, topk = 4, 256, 128, 2 + x = ( + torch.randn(tokens, model_dim, device="cuda", dtype=torch.bfloat16) * 0.1 + ).requires_grad_() + w1_master = ( + torch.randn(experts, hidden, model_dim, device="cuda", dtype=torch.bfloat16) + * 0.05 + ).requires_grad_() + w3_master = ( + torch.randn(experts, hidden, model_dim, device="cuda", dtype=torch.bfloat16) + * 0.05 + ).requires_grad_() + w2_master = ( + torch.randn(experts, model_dim, hidden, device="cuda", dtype=torch.bfloat16) + * 0.05 + ).requires_grad_() + canonical_w13 = torch.stack( + (w1_master.detach(), w3_master.detach()), dim=1 + ).flatten(0, 1) + w13_q, w13_s = _blockwise_weight_quantize(canonical_w13) + w2_q, w2_s = _blockwise_weight_quantize(w2_master.detach()) + if all_to_one: + topk_ids = torch.zeros(tokens, topk, device="cuda", dtype=torch.int64) + else: + route = torch.arange(tokens * topk, device="cuda", dtype=torch.int64) + topk_ids = torch.remainder(route * 3 + 1, experts).view(tokens, topk) + scores = torch.sigmoid( + torch.randn(tokens, topk, device="cuda", dtype=torch.float32) + ) + scores = (scores / scores.sum(dim=-1, keepdim=True) * 2.5).detach().requires_grad_() + return { + "x": x, + "ids": topk_ids.contiguous(), + "scores": scores.contiguous(), + "w13_q": w13_q, + "w13_s": w13_s, + "w2_q": w2_q, + "w2_s": w2_s, + "w1_master": w1_master, + "w2_master": w2_master, + "w3_master": w3_master, + } + + +def _forward_reference( + case: dict[str, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + x_q, x_s = deep_gemm._C.sm103_fp8_block128_quantize(case["x"].detach()) + x_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize(x_q, x_s).float() + canonical_w13 = _blockwise_weight_dequantize(case["w13_q"], case["w13_s"]) + experts = case["w2_q"].shape[0] + hidden = case["w2_q"].shape[2] + model_dim = case["w2_q"].shape[1] + canonical_pairs = canonical_w13.view(experts, 2, hidden, model_dim) + w13 = torch.stack((canonical_pairs[:, 1], canonical_pairs[:, 0]), dim=1).reshape( + experts, hidden * 2, model_dim + ) + w2 = _blockwise_weight_dequantize(case["w2_q"], case["w2_s"]) + flat_ids = case["ids"].flatten() + route_x = x_dequantized.repeat_interleave(case["ids"].shape[1], dim=0) + preactivation = ( + torch.bmm(w13.index_select(0, flat_ids), route_x.unsqueeze(-1)) + .squeeze(-1) + .to(torch.bfloat16) + ) + hidden_q, hidden_s = deep_gemm._C.sm103_fp8_block128_swiglu_quantize(preactivation) + hidden = deep_gemm._C.sm103_fp8_block128_dequantize(hidden_q, hidden_s).float() + route_output = ( + torch.bmm(w2.index_select(0, flat_ids), hidden.unsqueeze(-1)) + .squeeze(-1) + .to(torch.bfloat16) + ) + output = deep_gemm._C.sm103_fp8_block128_post_down_combine( + route_output, case["scores"] + ) + return output, route_output + + +def _ste_reference_backward( + case: dict[str, torch.Tensor], upstream: torch.Tensor +) -> dict[str, torch.Tensor]: + x = case["x"].detach().clone().requires_grad_() + scores = case["scores"].detach().clone().requires_grad_() + w1_master = case["w1_master"].detach().clone().requires_grad_() + w2_master = case["w2_master"].detach().clone().requires_grad_() + w3_master = case["w3_master"].detach().clone().requires_grad_() + experts, model_dim, hidden = w2_master.shape + + x_q, x_s = deep_gemm._C.sm103_fp8_block128_quantize(x.detach()) + x_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize(x_q, x_s).float() + x_effective = x.float() + (x_dequantized - x.float()).detach() + w13_active_master = torch.stack((w3_master, w1_master), dim=1).reshape( + experts, hidden * 2, model_dim + ) + canonical_w13_dequantized = _blockwise_weight_dequantize( + case["w13_q"], case["w13_s"] + ).view(experts, 2, hidden, model_dim) + w13_effective = ( + w13_active_master.float() + + ( + torch.stack( + (canonical_w13_dequantized[:, 1], canonical_w13_dequantized[:, 0]), + dim=1, + ).reshape(experts, hidden * 2, model_dim) + - w13_active_master.float() + ).detach() + ) + w2_dequantized = _blockwise_weight_dequantize(case["w2_q"], case["w2_s"]) + w2_effective = w2_master.float() + (w2_dequantized - w2_master.float()).detach() + + flat_ids = case["ids"].flatten() + route_x = x_effective.repeat_interleave(case["ids"].shape[1], dim=0) + preactivation = ( + torch.bmm(w13_effective.index_select(0, flat_ids), route_x.unsqueeze(-1)) + .squeeze(-1) + .to(torch.bfloat16) + ) + up, gate = preactivation.float().chunk(2, dim=-1) + hidden_raw = torch.nn.functional.silu(gate) * up + hidden_q, hidden_s = deep_gemm._C.sm103_fp8_block128_swiglu_quantize( + preactivation.detach() + ) + hidden_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize( + hidden_q, hidden_s + ).float() + hidden_effective = hidden_raw + (hidden_dequantized - hidden_raw).detach() + route_output = ( + torch.bmm( + w2_effective.index_select(0, flat_ids), hidden_effective.unsqueeze(-1) + ) + .squeeze(-1) + .to(torch.bfloat16) + ) + tokens, topk = case["ids"].shape + output_float = torch.zeros(tokens, model_dim, device="cuda", dtype=torch.float32) + route_output_view = route_output.view(tokens, topk, model_dim) + for route in range(topk): + output_float = ( + output_float + route_output_view[:, route].float() * scores[:, route, None] + ) + output = output_float.to(torch.bfloat16) + output.backward(upstream) + return { + "output": output.detach(), + "x_grad": x.grad.detach(), + "score_grad": scores.grad.detach(), + "w1_grad": w1_master.grad.detach(), + "w2_grad": w2_master.grad.detach(), + "w3_grad": w3_master.grad.detach(), + } + + +def _normalized_difference(actual: torch.Tensor, expected: torch.Tensor) -> float: + actual_f = actual.float().flatten() + expected_f = expected.float().flatten() + return float( + 1 + - 2 + * (actual_f @ expected_f) + / (actual_f.square().sum() + expected_f.square().sum() + 1e-12) + ) + + +class _DistributedMasterFixture: + def __init__( + self, + local: torch.Tensor, + logical_shape: tuple[int, int, int], + ) -> None: + self._local = local + self.shape = logical_shape + + def to_local(self) -> torch.Tensor: + return self._local + + +def test_distributed_master_validation_accepts_resident_efsdp_shard_with_wrapper() -> ( + None +): + local = torch.empty(1, 128, 256, device="cuda", dtype=torch.bfloat16) + master = _DistributedMasterFixture(local, (4, 128, 256)) + + _validate_master_tensor( + master, + name="w1_master", + local_shape=(2, 128, 256), + global_shape=(4, 128, 256), + device=local.device, + master_gradient_wrapper=lambda *_grads: None, + ) + + +def test_distributed_master_validation_requires_gradient_wrapper() -> None: + local = torch.empty(1, 128, 256, device="cuda", dtype=torch.bfloat16) + master = _DistributedMasterFixture(local, (4, 128, 256)) + + with pytest.raises(ValueError, match="resident shape"): + _validate_master_tensor( + master, + name="w1_master", + local_shape=(2, 128, 256), + global_shape=(4, 128, 256), + device=local.device, + master_gradient_wrapper=None, + ) + + +def test_expanded_bf16_wgrad_uses_only_psum_valid_rows() -> None: + torch.manual_seed(20260722) + psum = torch.tensor([3, 133], dtype=torch.int32, device="cuda") + left = torch.randn(256, 128, dtype=torch.bfloat16, device="cuda") + right = torch.randn(256, 256, dtype=torch.bfloat16, device="cuda") + left[3:128].fill_(float("nan")) + left[133:].fill_(float("nan")) + right[3:128].fill_(float("nan")) + right[133:].fill_(float("nan")) + + actual = _bf16_grouped_wgrad_expanded(left, right, psum) + expected = torch.stack( + ( + left[:3].transpose(0, 1) @ right[:3], + left[128:133].transpose(0, 1) @ right[128:133], + ) + ) + assert torch.isfinite(actual).all() + torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2) + + +@pytest.mark.parametrize( + ("tokens", "all_to_one"), + [(1, True), (63, False), (64, True), (65, False)], +) +def test_single_rank_forward_matches_reference_across_padding_boundaries( + tokens: int, all_to_one: bool +) -> None: + case = _make_case(tokens, all_to_one=all_to_one) + actual = deep_gemm.fp8_block128_mega_moe( + case["x"], + case["ids"], + case["scores"], + case["w13_q"], + case["w13_s"], + case["w2_q"], + case["w2_s"], + case["w1_master"], + case["w2_master"], + case["w3_master"], + ) + expected, _ = _forward_reference(case) + torch.testing.assert_close(actual.float(), expected.float(), rtol=0.08, atol=0.08) + + +def test_glm_w13_transform_is_gate_up_to_up_gate() -> None: + gate = torch.full((128, 256), 3, device="cuda", dtype=torch.float8_e4m3fn) + up = torch.full((128, 256), 7, device="cuda", dtype=torch.float8_e4m3fn) + gate_second = torch.full_like(gate, 6) + up_second = torch.full_like(up, 14) + canonical = torch.stack((gate, up, gate_second, up_second)) + scales = torch.arange(8, device="cuda", dtype=torch.float32).view(4, 1, 2) + active, active_scales = deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( + canonical, scales + ) + assert torch.equal(active[0, :128], up) + assert torch.equal(active[0, 128:], gate) + assert torch.equal(active[1, :128], up_second) + assert torch.equal(active[1, 128:], gate_second) + assert torch.equal(active_scales[0, 0], scales[1, 0]) + assert torch.equal(active_scales[0, 1], scales[0, 0]) + + +def test_single_rank_backward_returns_input_score_and_canonical_master_grads() -> None: + case = _make_case(9, all_to_one=False) + upstream = torch.randn(9, 256, device="cuda", dtype=torch.bfloat16) + expected = _ste_reference_backward(case, upstream) + output = deep_gemm.fp8_block128_mega_moe( + case["x"], + case["ids"], + case["scores"], + case["w13_q"], + case["w13_s"], + case["w2_q"], + case["w2_s"], + case["w1_master"], + case["w2_master"], + case["w3_master"], + ) + output.backward(upstream) + + torch.testing.assert_close( + output.float(), expected["output"].float(), rtol=0.08, atol=0.08 + ) + torch.testing.assert_close( + case["scores"].grad, expected["score_grad"], rtol=3e-4, atol=3e-3 + ) + assert case["x"].grad.shape == case["x"].shape + assert case["x"].grad.dtype == torch.bfloat16 + assert case["w1_master"].grad.shape == case["w1_master"].shape + assert case["w1_master"].grad.dtype == torch.bfloat16 + assert case["w2_master"].grad.shape == case["w2_master"].shape + assert case["w2_master"].grad.dtype == torch.bfloat16 + assert case["w3_master"].grad.shape == case["w3_master"].shape + assert case["w3_master"].grad.dtype == torch.bfloat16 + assert _normalized_difference(case["x"].grad, expected["x_grad"]) < 0.12 + assert _normalized_difference(case["w1_master"].grad, expected["w1_grad"]) < 0.15 + assert _normalized_difference(case["w2_master"].grad, expected["w2_grad"]) < 0.12 + assert _normalized_difference(case["w3_master"].grad, expected["w3_grad"]) < 0.15 diff --git a/tests/test_fp8_block128_mega_moe_distributed.py b/tests/test_fp8_block128_mega_moe_distributed.py new file mode 100644 index 0000000000..b71e5b64eb --- /dev/null +++ b/tests/test_fp8_block128_mega_moe_distributed.py @@ -0,0 +1,570 @@ +"""Two-rank distributed parity/adversarial coverage for FP8-block128 MegaMoE.""" + +from __future__ import annotations + +import os +import socket + +import pytest +import torch +import torch.multiprocessing as mp +from torch.utils.checkpoint import checkpoint + + +def _free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return int(sock.getsockname()[1]) + + +def _weight_quantize(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + groups, rows, columns = weight.shape + blocks = ( + weight.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + ) + scales = blocks.abs().amax(dim=(-1, -2)) / 448.0 + scales = torch.where(scales == 0, torch.ones_like(scales), scales) + quantized = ( + (blocks / scales[..., None, None]) + .to(torch.float8_e4m3fn) + .permute(0, 1, 3, 2, 4) + .reshape_as(weight) + .contiguous() + ) + return quantized, scales.contiguous() + + +def _weight_dequantize(quantized: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + groups, rows, columns = quantized.shape + return ( + quantized.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + .mul(scales[..., None, None]) + .permute(0, 1, 3, 2, 4) + .reshape(groups, rows, columns) + ) + + +def _rank_inputs(rank: int, device: torch.device) -> tuple[torch.Tensor, ...]: + tokens, topk = 5 + rank * 2, 2 + generator = torch.Generator(device=device).manual_seed(9000 + rank) + x = ( + torch.randn( + tokens, 256, generator=generator, device=device, dtype=torch.bfloat16 + ) + * 0.1 + ).requires_grad_() + token_index = torch.arange(tokens, device=device, dtype=torch.int64) + # Every rank sends remotely, expert 2 receives skewed traffic, and expert 3 + # is globally empty. The unequal token counts stress asymmetric splits. + first = torch.remainder(token_index + rank, 2) + second = torch.full_like(first, 2) + ids = torch.stack((first, second), dim=-1).contiguous() + raw_scores = torch.sigmoid( + torch.randn( + tokens, topk, generator=generator, device=device, dtype=torch.float32 + ) + ) + scores = ( + (raw_scores / raw_scores.sum(dim=-1, keepdim=True) * 2.5) + .detach() + .requires_grad_() + ) + upstream = torch.randn( + tokens, 256, generator=generator, device=device, dtype=torch.bfloat16 + ) + return x, ids, scores.contiguous(), upstream + + +def _reference( + deep_gemm, + x: torch.Tensor, + ids: torch.Tensor, + scores: torch.Tensor, + full_w13_q: torch.Tensor, + full_w13_s: torch.Tensor, + full_w2_q: torch.Tensor, + full_w2_s: torch.Tensor, + upstream: torch.Tensor, +) -> tuple[torch.Tensor, ...]: + x_ref = x.detach().clone() + score_ref = scores.detach().clone() + x_q, x_s = deep_gemm._C.sm103_fp8_block128_quantize(x_ref) + x_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize(x_q, x_s).float() + w13 = _weight_dequantize(full_w13_q, full_w13_s) + w2 = _weight_dequantize(full_w2_q, full_w2_s) + flat_ids = ids.flatten() + route_x = x_dequantized.repeat_interleave(ids.shape[1], dim=0) + preactivation = ( + torch.bmm(w13.index_select(0, flat_ids), route_x.unsqueeze(-1)) + .squeeze(-1) + .to(torch.bfloat16) + ) + up, gate = preactivation.float().chunk(2, dim=-1) + hidden_q, hidden_s = deep_gemm._C.sm103_fp8_block128_swiglu_quantize( + preactivation.detach() + ) + hidden_dequantized_bf16 = deep_gemm._C.sm103_fp8_block128_dequantize( + hidden_q, hidden_s + ) + route_output = ( + torch.bmm( + w2.index_select(0, flat_ids), hidden_dequantized_bf16.float().unsqueeze(-1) + ) + .squeeze(-1) + .to(torch.bfloat16) + ) + tokens, topk = ids.shape + output_float = torch.zeros(tokens, 256, device=x.device, dtype=torch.float32) + route_view = route_output.view(tokens, topk, 256) + for route in range(topk): + output_float = ( + output_float + route_view[:, route].float() * score_ref[:, route, None] + ) + output = output_float.to(torch.bfloat16) + + upstream_routes = upstream.repeat_interleave(topk, dim=0).float() + grad_score = (route_output.float() * upstream_routes).sum(dim=-1).view(tokens, topk) + grad_route_float = upstream_routes * score_ref.flatten()[:, None] + grad_route_blocks = grad_route_float.view( + grad_route_float.shape[0], grad_route_float.shape[1] // 128, 128 + ) + grad_route_scales = grad_route_blocks.abs().amax(dim=-1) / 448.0 + grad_route_scales = torch.where( + grad_route_scales == 0, + torch.ones_like(grad_route_scales), + grad_route_scales, + ) + grad_route_q = (grad_route_blocks / grad_route_scales[..., None]).to( + torch.float8_e4m3fn + ) + grad_route = ( + (grad_route_q.float() * grad_route_scales[..., None]) + .reshape_as(grad_route_float) + .to(torch.bfloat16) + ) + grad_hidden = ( + torch.bmm( + w2.index_select(0, flat_ids).transpose(1, 2), + grad_route.float().unsqueeze(-1), + ) + .squeeze(-1) + .to(torch.bfloat16) + ) + sigmoid_gate = torch.sigmoid(gate) + grad_gate = ( + grad_hidden.float() * up * sigmoid_gate * (1.0 + gate * (1.0 - sigmoid_gate)) + ).to(torch.bfloat16) + grad_up = (grad_hidden.float() * gate * sigmoid_gate).to(torch.bfloat16) + grad_preactivation = torch.cat((grad_gate, grad_up), dim=-1) + grad_pre_q, grad_pre_s = deep_gemm._C.sm103_fp8_block128_quantize( + grad_preactivation + ) + grad_pre_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize( + grad_pre_q, grad_pre_s + ).float() + canonical_w13 = torch.cat( + (w13[:, w13.shape[1] // 2 :], w13[:, : w13.shape[1] // 2]), dim=1 + ) + grad_input_routes = ( + torch.bmm( + canonical_w13.index_select(0, flat_ids).transpose(1, 2), + grad_pre_dequantized.unsqueeze(-1), + ) + .squeeze(-1) + .to(torch.bfloat16) + ) + grad_input_float = torch.zeros_like(x_ref, dtype=torch.float32) + grad_input_view = grad_input_routes.view(tokens, topk, x_ref.shape[1]) + for route in range(topk): + grad_input_float += grad_input_view[:, route].float() + + experts, hidden, model_dim = full_w2_q.shape[0], full_w2_q.shape[2], x.shape[1] + grad_w1 = torch.zeros( + experts, hidden, model_dim, device=x.device, dtype=torch.float32 + ) + grad_w2 = torch.zeros_like(w2, dtype=torch.float32) + grad_w3 = torch.zeros_like(grad_w1) + route_x_bf16 = route_x.to(torch.bfloat16) + for expert in range(experts): + mask = flat_ids == expert + if not mask.any(): + continue + grad_w1[expert] = ( + grad_gate[mask].float().transpose(0, 1) @ route_x_bf16[mask].float() + ) + grad_w2[expert] = ( + grad_route[mask].float().transpose(0, 1) + @ hidden_dequantized_bf16[mask].float() + ) + grad_w3[expert] = ( + grad_up[mask].float().transpose(0, 1) @ route_x_bf16[mask].float() + ) + return ( + output, + grad_input_float.to(torch.bfloat16), + grad_score, + grad_w1, + grad_w2, + grad_w3, + ) + + +def _normalized_difference(actual: torch.Tensor, expected: torch.Tensor) -> float: + actual_f, expected_f = actual.float().flatten(), expected.float().flatten() + return float( + 1 + - 2 + * (actual_f @ expected_f) + / (actual_f.square().sum() + expected_f.square().sum() + 1e-12) + ) + + +def _worker(rank: int, world_size: int, port: int) -> None: + os.environ["DG_JIT_CACHE_DIR"] = f"/tmp/deepgemm-sm103-dist-rank{rank}" + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + import torch.distributed as dist + + dist.init_process_group( + "nccl", + init_method=f"tcp://127.0.0.1:{port}", + rank=rank, + world_size=world_size, + ) + try: + import deep_gemm + from deep_gemm.mega.fp8_block128 import ( + _configure_fp8_block128_mega_moe_transport, + ) + + assert torch.cuda.get_device_capability(device) == (10, 3) + _configure_fp8_block128_mega_moe_transport( + dist.group.WORLD, + context_tokens_per_rank=64, + ) + experts, local_experts, model_dim, hidden = 4, 2, 256, 128 + generator = torch.Generator(device=device).manual_seed(8800) + full_w13_master = ( + torch.randn( + experts * 2, + hidden, + model_dim, + generator=generator, + device=device, + dtype=torch.bfloat16, + ) + * 0.05 + ) + full_w2_master = ( + torch.randn( + experts, + model_dim, + hidden, + generator=generator, + device=device, + dtype=torch.bfloat16, + ) + * 0.05 + ) + full_w13_q_canonical, full_w13_s_canonical = _weight_quantize(full_w13_master) + full_w13_q_active, full_w13_s_active = ( + deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( + full_w13_q_canonical, full_w13_s_canonical + ) + ) + full_w2_q, full_w2_s = _weight_quantize(full_w2_master) + expert_start = rank * local_experts + expert_end = expert_start + local_experts + local_w13_q = full_w13_q_canonical[ + expert_start * 2 : expert_end * 2 + ].contiguous() + local_w13_s = full_w13_s_canonical[ + expert_start * 2 : expert_end * 2 + ].contiguous() + local_w2_q = full_w2_q[expert_start:expert_end].contiguous() + local_w2_s = full_w2_s[expert_start:expert_end].contiguous() + local_w1_master = ( + full_w13_master[expert_start * 2 : expert_end * 2 : 2] + .clone() + .detach() + .requires_grad_() + ) + local_w3_master = ( + full_w13_master[expert_start * 2 + 1 : expert_end * 2 : 2] + .clone() + .detach() + .requires_grad_() + ) + local_w2_master = ( + full_w2_master[expert_start:expert_end].clone().detach().requires_grad_() + ) + x, ids, scores, upstream = _rank_inputs(rank, device) + # Reproduce the production lifecycle: a small warmup constructs the + # first arena, whose fixed target envelope must admit the real batch. + with torch.no_grad(): + warmup_output = deep_gemm.fp8_block128_mega_moe( + x[:1].contiguous(), + ids[:1].contiguous(), + scores[:1].contiguous(), + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w1_master, + local_w2_master, + local_w3_master, + group=dist.group.WORLD, + ) + assert warmup_output.shape == (1, model_dim) + ( + expected_output, + expected_x_grad, + expected_score_grad, + expected_w1_grad, + expected_w2_grad, + expected_w3_grad, + ) = _reference( + deep_gemm, + x, + ids, + scores, + full_w13_q_active, + full_w13_s_active, + full_w2_q, + full_w2_s, + upstream, + ) + # Every sharded expert receives routes from every source rank. Sum the + # FP32 reference partials before comparing with the owning rank's one + # grouped BF16-master wgrad result. + for expected_weight_grad in ( + expected_w1_grad, + expected_w2_grad, + expected_w3_grad, + ): + dist.all_reduce(expected_weight_grad, group=dist.group.WORLD) + expected_w1_grad = expected_w1_grad.to(torch.bfloat16) + expected_w2_grad = expected_w2_grad.to(torch.bfloat16) + expected_w3_grad = expected_w3_grad.to(torch.bfloat16) + + output = deep_gemm.fp8_block128_mega_moe( + x, + ids, + scores, + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w1_master, + local_w2_master, + local_w3_master, + group=dist.group.WORLD, + ) + output.backward(upstream) + torch.testing.assert_close( + output.float(), expected_output.float(), rtol=0.08, atol=0.08 + ) + torch.testing.assert_close( + scores.grad, expected_score_grad, rtol=3e-4, atol=3e-3 + ) + assert _normalized_difference(x.grad, expected_x_grad) < 0.12 + weight_differences = { + "w1": _normalized_difference( + local_w1_master.grad, + expected_w1_grad[expert_start:expert_end], + ), + "w2": _normalized_difference( + local_w2_master.grad, + expected_w2_grad[expert_start:expert_end], + ), + "w3": _normalized_difference( + local_w3_master.grad, + expected_w3_grad[expert_start:expert_end], + ), + } + assert weight_differences["w1"] < 0.15, weight_differences + assert weight_differences["w2"] < 0.12, weight_differences + assert weight_differences["w3"] < 0.15, weight_differences + for gradient in ( + x.grad, + local_w1_master.grad, + local_w2_master.grad, + local_w3_master.grad, + ): + assert gradient is not None + assert gradient.dtype == torch.bfloat16 + assert torch.isfinite(gradient).all() + if rank == 1: + # Global expert 3 has no routes on either source rank. + assert torch.count_nonzero(local_w1_master.grad[1]) == 0 + assert torch.count_nonzero(local_w3_master.grad[1]) == 0 + assert torch.count_nonzero(local_w2_master.grad[1]) == 0 + + direct_output = output.detach() + direct_x_grad = x.grad.detach().clone() + direct_score_grad = scores.grad.detach().clone() + direct_weight_grads = tuple( + value.grad.detach().clone() + for value in (local_w1_master, local_w2_master, local_w3_master) + ) + for value in (local_w1_master, local_w2_master, local_w3_master): + value.grad = None + checkpoint_x = x.detach().clone().requires_grad_() + checkpoint_scores = scores.detach().clone().requires_grad_() + + def checkpointed_megamoe( + checkpoint_input: torch.Tensor, + checkpoint_route_scores: torch.Tensor, + ) -> torch.Tensor: + return deep_gemm.fp8_block128_mega_moe( + checkpoint_input, + ids, + checkpoint_route_scores, + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w1_master, + local_w2_master, + local_w3_master, + group=dist.group.WORLD, + ) + + checkpoint_output = checkpoint( + checkpointed_megamoe, + checkpoint_x, + checkpoint_scores, + use_reentrant=False, + ) + checkpoint_output.backward(upstream) + torch.testing.assert_close(checkpoint_output, direct_output, rtol=0, atol=0) + torch.testing.assert_close(checkpoint_x.grad, direct_x_grad, rtol=0, atol=0) + torch.testing.assert_close( + checkpoint_scores.grad, direct_score_grad, rtol=0, atol=0 + ) + for checkpoint_grad, direct_grad in zip( + (local_w1_master.grad, local_w2_master.grad, local_w3_master.grad), + direct_weight_grads, + ): + torch.testing.assert_close(checkpoint_grad, direct_grad, rtol=0, atol=0) + + # Exercise the production lifetime where multiple layer forwards reuse + # one ElasticBuffer before their backwards consume distinct handles. + # The second routing pattern activates expert 3 instead of expert 2 so + # stale/overwritten handle metadata cannot accidentally pass. + second_x = (x.detach() * 0.75).contiguous().requires_grad_() + second_ids = ids.detach().clone() + second_ids[:, 0] = 1 - second_ids[:, 0] + second_ids[:, 1] = 3 + second_scores = scores.detach().flip(1).contiguous().requires_grad_() + second_upstream = upstream.flip(0).contiguous() + ( + second_expected_output, + second_expected_x_grad, + second_expected_score_grad, + second_expected_w1_grad, + second_expected_w2_grad, + second_expected_w3_grad, + ) = _reference( + deep_gemm, + second_x, + second_ids, + second_scores, + full_w13_q_active, + full_w13_s_active, + full_w2_q, + full_w2_s, + second_upstream, + ) + for second_expected_weight_grad in ( + second_expected_w1_grad, + second_expected_w2_grad, + second_expected_w3_grad, + ): + dist.all_reduce(second_expected_weight_grad, group=dist.group.WORLD) + + for value in (local_w1_master, local_w2_master, local_w3_master): + value.grad = None + first_x = x.detach().clone().requires_grad_() + first_scores = scores.detach().clone().requires_grad_() + first_output = deep_gemm.fp8_block128_mega_moe( + first_x, + ids, + first_scores, + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w1_master, + local_w2_master, + local_w3_master, + group=dist.group.WORLD, + ) + second_output = deep_gemm.fp8_block128_mega_moe( + second_x, + second_ids, + second_scores, + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w1_master, + local_w2_master, + local_w3_master, + group=dist.group.WORLD, + ) + torch.autograd.backward( + (first_output, second_output), (upstream, second_upstream) + ) + torch.testing.assert_close(first_output, direct_output, rtol=0, atol=0) + torch.testing.assert_close( + second_output.float(), + second_expected_output.float(), + rtol=0.08, + atol=0.08, + ) + torch.testing.assert_close(first_x.grad, direct_x_grad, rtol=0, atol=0) + torch.testing.assert_close(first_scores.grad, direct_score_grad, rtol=0, atol=0) + assert _normalized_difference(second_x.grad, second_expected_x_grad) < 0.12 + torch.testing.assert_close( + second_scores.grad, + second_expected_score_grad, + rtol=3e-4, + atol=3e-3, + ) + expected_combined_weight_grads = ( + expected_w1_grad.float() + + second_expected_w1_grad.to(torch.bfloat16).float(), + expected_w2_grad.float() + + second_expected_w2_grad.to(torch.bfloat16).float(), + expected_w3_grad.float() + + second_expected_w3_grad.to(torch.bfloat16).float(), + ) + combined_weight_differences = { + name: _normalized_difference( + actual.grad, + expected[expert_start:expert_end], + ) + for name, actual, expected in zip( + ("w1", "w2", "w3"), + (local_w1_master, local_w2_master, local_w3_master), + expected_combined_weight_grads, + ) + } + assert combined_weight_differences["w1"] < 0.15, combined_weight_differences + assert combined_weight_differences["w2"] < 0.12, combined_weight_differences + assert combined_weight_differences["w3"] < 0.15, combined_weight_differences + dist.barrier() + finally: + dist.destroy_process_group() + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires two SM103 GPUs") +def test_two_rank_uneven_cross_rank_forward_backward() -> None: + if any(torch.cuda.get_device_capability(index) != (10, 3) for index in range(2)): + pytest.skip("requires two SM103 GPUs") + mp.spawn(_worker, args=(2, _free_port()), nprocs=2, join=True) diff --git a/tests/test_sm103_fp8_block128_primitives.py b/tests/test_sm103_fp8_block128_primitives.py new file mode 100644 index 0000000000..ebe959e4c9 --- /dev/null +++ b/tests/test_sm103_fp8_block128_primitives.py @@ -0,0 +1,396 @@ +import re + +import pytest +import torch + +import deep_gemm + + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 3), + reason="SM103-only primitive test", +) + + +def _reference_quantize(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + rows, columns = x.shape + grouped = x.float().view(rows, columns // 128, 128) + scales = grouped.abs().amax(dim=-1) / 448.0 + scales = torch.where(scales == 0, torch.ones_like(scales), scales) + quantized = (grouped / scales.unsqueeze(-1)).to(torch.float8_e4m3fn).view_as(x) + return quantized, scales + + +def _assert_quantized_close( + actual_quantized: torch.Tensor, + actual_scales: torch.Tensor, + expected_quantized: torch.Tensor, + expected_scales: torch.Tensor, +) -> None: + torch.testing.assert_close(actual_scales, expected_scales, rtol=2e-6, atol=1e-8) + if actual_quantized.numel() == 0: + return + # CUDA's native FP8 constructor and PyTorch's vectorized cast can choose + # adjacent E4M3 bins at exact midpoints. Bound that difference explicitly + # and compare the represented (dequantized) values rather than requiring + # implementation-specific byte identity. + mismatch_ratio = ( + actual_quantized.view(torch.uint8) != expected_quantized.view(torch.uint8) + ).float().mean() + assert mismatch_ratio <= 0.03 + columns = actual_quantized.size(1) + actual = ( + actual_quantized.float().view(-1, columns // 128, 128) + * actual_scales.unsqueeze(-1) + ).flatten(1) + expected = ( + expected_quantized.float().view(-1, columns // 128, 128) + * expected_scales.unsqueeze(-1) + ).flatten(1) + torch.testing.assert_close(actual, expected, rtol=0.13, atol=0.02) + + +def _blockwise_weight_quantize(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + groups, rows, columns = weight.shape + blocks = ( + weight.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + ) + scales = blocks.abs().amax(dim=(-1, -2)) / 448.0 + scales = torch.where(scales == 0, torch.ones_like(scales), scales) + quantized = ( + (blocks / scales[..., None, None]) + .to(torch.float8_e4m3fn) + .permute(0, 1, 3, 2, 4) + .reshape_as(weight) + .contiguous() + ) + return quantized, scales.contiguous() + + +def _blockwise_weight_dequantize( + quantized: torch.Tensor, scales: torch.Tensor +) -> torch.Tensor: + groups, rows, columns = quantized.shape + return ( + quantized.float() + .view(groups, rows // 128, 128, columns // 128, 128) + .permute(0, 1, 3, 2, 4) + .mul(scales[..., None, None]) + .permute(0, 1, 3, 2, 4) + .reshape(groups, rows, columns) + ) + + +def _ceil_to_ue8m0(scales: torch.Tensor) -> torch.Tensor: + bits = scales.abs().float().view(torch.int32) + exponent = ((bits >> 23) & 0xFF) + (bits & 0x7FFFFF).bool().to(torch.int32) + return (exponent.clamp(1, 254) << 23).view(torch.float32) + + +def _ue8m0_per_token_quantize(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + rows, columns = x.shape + grouped = x.float().view(rows, columns // 128, 128) + scales = _ceil_to_ue8m0(grouped.abs().amax(dim=-1).clamp_min(1e-4) / 448.0) + quantized = (grouped / scales.unsqueeze(-1)).to(torch.float8_e4m3fn).view_as(x) + return quantized.contiguous(), scales.contiguous() + + +def _ue8m0_blockwise_weight_quantize( + weight: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + groups, rows, columns = weight.shape + blocks = weight.float().view(groups, rows // 128, 128, columns // 128, 128).permute(0, 1, 3, 2, 4) + scales = _ceil_to_ue8m0(blocks.abs().amax(dim=(-1, -2)).clamp_min(1e-4) / 448.0) + quantized = ( + (blocks / scales[..., None, None]) + .to(torch.float8_e4m3fn) + .permute(0, 1, 3, 2, 4) + .reshape_as(weight) + .contiguous() + ) + return quantized, scales.contiguous() + + +def test_build_provenance_and_capabilities_are_fail_closed() -> None: + assert re.fullmatch(r"[0-9a-f]{40}", deep_gemm.__git_commit__) + capabilities = deep_gemm._C.get_sm103_fp8_block128_capabilities() + assert capabilities["architecture"] == "sm103" + assert capabilities["compute_capability"] == (10, 3) + assert capabilities["activation_dtype"] == "float8_e4m3fn" + assert capabilities["weight_dtype"] == "float8_e4m3fn" + assert capabilities["scale_dtype"] == "float32" + assert capabilities["activation_group_k"] == 128 + assert capabilities["weight_block_m"] == 128 + assert capabilities["weight_block_k"] == 128 + assert capabilities["route_score_placement"] == "post_down" + assert capabilities["fallback"] is None + + +@pytest.mark.parametrize("rows", [0, 1, 5, 128, 129]) +def test_quantize_and_dequantize_match_reference(rows: int) -> None: + torch.manual_seed(1000 + rows) + x = torch.randn(rows, 256, device="cuda", dtype=torch.bfloat16) + quantized, scales = deep_gemm._C.sm103_fp8_block128_quantize(x) + expected_quantized, expected_scales = _reference_quantize(x) + _assert_quantized_close(quantized, scales, expected_quantized, expected_scales) + + dequantized = deep_gemm._C.sm103_fp8_block128_dequantize(quantized, scales) + expected_dequantized = ( + quantized.float().view(rows, 2, 128) * scales.unsqueeze(-1) + ).view_as(x).to(torch.bfloat16) + assert torch.equal(dequantized, expected_dequantized) + + +@pytest.mark.parametrize("rows", [0, 3, 129]) +def test_swiglu_forward_and_backward_match_reference(rows: int) -> None: + torch.manual_seed(2000 + rows) + preactivation = torch.randn(rows, 512, device="cuda", dtype=torch.bfloat16) + quantized, scales = deep_gemm._C.sm103_fp8_block128_swiglu_quantize(preactivation) + up, gate = preactivation.float().chunk(2, dim=-1) + expected_quantized, expected_scales = _reference_quantize(torch.nn.functional.silu(gate) * up) + _assert_quantized_close(quantized, scales, expected_quantized, expected_scales) + + grad_output = torch.randn(rows, 256, device="cuda", dtype=torch.bfloat16) + actual_grad = deep_gemm._C.sm103_fp8_block128_swiglu_backward(grad_output, preactivation) + preactivation_ref = preactivation.float().detach().requires_grad_(True) + up_ref, gate_ref = preactivation_ref.chunk(2, dim=-1) + (torch.nn.functional.silu(gate_ref) * up_ref).backward(grad_output.float()) + torch.testing.assert_close( + actual_grad.float(), preactivation_ref.grad.to(torch.bfloat16).float(), rtol=3e-2, atol=2e-2 + ) + canonical_grad = deep_gemm._C.sm103_fp8_block128_swiglu_backward_canonical( + grad_output, preactivation + ) + active_up_grad, active_gate_grad = preactivation_ref.grad.chunk(2, dim=-1) + expected_canonical = torch.cat((active_gate_grad, active_up_grad), dim=-1) + torch.testing.assert_close( + canonical_grad.float(), + expected_canonical.to(torch.bfloat16).float(), + rtol=3e-2, + atol=2e-2, + ) + + +def test_post_down_combine_score_grad_and_route_sum() -> None: + torch.manual_seed(3000) + tokens, topk, hidden = 7, 8, 256 + route_output = torch.randn(tokens * topk, hidden, device="cuda", dtype=torch.bfloat16) + route_scores = torch.randn(tokens, topk, device="cuda", dtype=torch.float32) + grad_output = torch.randn(tokens, hidden, device="cuda", dtype=torch.bfloat16) + + actual = deep_gemm._C.sm103_fp8_block128_post_down_combine(route_output, route_scores) + expected = torch.zeros(tokens, hidden, device="cuda", dtype=torch.float32) + for route in range(topk): + expected.add_(route_output.view(tokens, topk, hidden)[:, route].float() * route_scores[:, route, None]) + assert torch.equal(actual, expected.to(torch.bfloat16)) + + actual_score_grad = deep_gemm._C.sm103_fp8_block128_post_down_score_grad( + route_output, grad_output, topk + ) + expected_score_grad = ( + route_output.view(tokens, topk, hidden).float() * grad_output[:, None].float() + ).sum(dim=-1) + torch.testing.assert_close(actual_score_grad, expected_score_grad, rtol=2e-5, atol=2e-4) + + actual_sum = deep_gemm._C.sm103_fp8_block128_route_sum(route_output, tokens, topk) + expected_sum = torch.zeros(tokens, hidden, device="cuda", dtype=torch.float32) + for route in range(topk): + expected_sum.add_(route_output.view(tokens, topk, hidden)[:, route].float()) + assert torch.equal(actual_sum, expected_sum.to(torch.bfloat16)) + + +def test_route_scale_quantize_uses_flat_route_order() -> None: + torch.manual_seed(4000) + tokens, topk, hidden = 5, 3, 256 + grad_output = torch.randn(tokens, hidden, device="cuda", dtype=torch.bfloat16) + route_scores = torch.randn(tokens, topk, device="cuda", dtype=torch.float32) + route_order = torch.randperm(tokens * topk, device="cuda", dtype=torch.int64) + + quantized, scales = deep_gemm._C.sm103_fp8_block128_route_scale_quantize( + grad_output, route_scores, route_order + ) + flat = route_order.cpu().tolist() + expanded = torch.stack( + [grad_output[route // topk].float() * route_scores.view(-1)[route] for route in flat] + ) + expected_quantized, expected_scales = _reference_quantize(expanded) + _assert_quantized_close(quantized, scales, expected_quantized, expected_scales) + + +@pytest.mark.parametrize("weight_is_k_by_n", [False, True]) +def test_grouped_fp8_block128_gemm_matches_dequantized_reference( + weight_is_k_by_n: bool, +) -> None: + torch.manual_seed(5000 + int(weight_is_k_by_n)) + counts = [4, 8, 0, 12] + rows, groups, k, n = sum(counts), len(counts), 256, 256 + activations_bf16 = ( + torch.randn(rows, k, device="cuda", dtype=torch.bfloat16) * 0.1 + ) + activations, activation_scales = deep_gemm._C.sm103_fp8_block128_quantize( + activations_bf16 + ) + physical_shape = (groups, k, n) if weight_is_k_by_n else (groups, n, k) + weights_bf16 = ( + torch.randn(*physical_shape, device="cuda", dtype=torch.bfloat16) * 0.1 + ) + weights, weight_scales = _blockwise_weight_quantize(weights_bf16) + + symbol = ( + deep_gemm._C.sm103_fp8_block128_grouped_gemm_nn + if weight_is_k_by_n + else deep_gemm._C.sm103_fp8_block128_grouped_gemm_nt + ) + actual = symbol( + activations, activation_scales, weights, weight_scales, counts + ) + + activations_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize( + activations, activation_scales + ).float() + weights_dequantized = _blockwise_weight_dequantize(weights, weight_scales) + expected_parts = [] + offset = 0 + for expert, count in enumerate(counts): + if count: + expert_weight = weights_dequantized[expert] + expected_parts.append( + activations_dequantized[offset : offset + count] + @ (expert_weight if weight_is_k_by_n else expert_weight.t()) + ) + offset += count + expected = torch.cat(expected_parts).to(torch.bfloat16) + torch.testing.assert_close(actual.float(), expected.float(), rtol=0.08, atol=0.08) + + +def test_grouped_fp8_block128_gemm_covers_zero_routes_and_rejects_unpadded_groups() -> None: + activations = torch.empty(0, 256, device="cuda", dtype=torch.float8_e4m3fn) + activation_scales = torch.empty(0, 2, device="cuda", dtype=torch.float32) + weights_bf16 = torch.randn(3, 256, 256, device="cuda", dtype=torch.bfloat16) + weights, weight_scales = _blockwise_weight_quantize(weights_bf16) + output = deep_gemm._C.sm103_fp8_block128_grouped_gemm_nt( + activations, activation_scales, weights, weight_scales, [0, 0, 0] + ) + assert output.shape == (0, 256) + assert output.dtype == torch.bfloat16 + + one_row = torch.zeros(1, 256, device="cuda", dtype=torch.bfloat16) + one_q, one_s = deep_gemm._C.sm103_fp8_block128_quantize(one_row) + with pytest.raises(RuntimeError, match="multiple of four"): + deep_gemm._C.sm103_fp8_block128_grouped_gemm_nt( + one_q, one_s, weights, weight_scales, [1, 0, 0] + ) + + +def test_canonical_w13_gemm_presents_up_then_gate_without_weight_copy() -> None: + torch.manual_seed(5500) + counts = [4, 0, 8] + experts, hidden, model_dim = len(counts), 128, 256 + activations_bf16 = torch.randn( + sum(counts), model_dim, device="cuda", dtype=torch.bfloat16 + ) + activations, activation_scales = deep_gemm._C.sm103_fp8_block128_quantize( + activations_bf16 + ) + canonical_bf16 = torch.randn( + experts * 2, hidden, model_dim, device="cuda", dtype=torch.bfloat16 + ) + canonical, canonical_scales = _blockwise_weight_quantize(canonical_bf16) + actual = deep_gemm._C.sm103_fp8_block128_grouped_w13_gemm_nt_canonical( + activations, + activation_scales, + canonical, + canonical_scales, + counts, + ) + + x_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize( + activations, activation_scales + ).float() + w_dequantized = _blockwise_weight_dequantize(canonical, canonical_scales) + expected_parts = [] + offset = 0 + for expert, count in enumerate(counts): + if count: + gate = x_dequantized[offset : offset + count] @ w_dequantized[2 * expert].t() + up = x_dequantized[offset : offset + count] @ w_dequantized[2 * expert + 1].t() + expected_parts.append(torch.cat((up, gate), dim=-1)) + offset += count + expected = torch.cat(expected_parts).to(torch.bfloat16) + torch.testing.assert_close(actual.float(), expected.float(), rtol=0.08, atol=0.08) + + +def test_legacy_psum_grouped_gemm_accepts_deepep_column_major_scales() -> None: + """Preserve the expanded DeepEP API consumed by the existing routed stack.""" + torch.manual_seed(5750) + alignment = 128 + deep_gemm.set_mk_alignment_for_contiguous_layout(alignment) + groups, rows, n, k = 2, 2 * alignment, 256, 256 + psum = torch.tensor([5, alignment + 5], device="cuda", dtype=torch.int32) + + activations_bf16 = torch.zeros(rows, k, device="cuda", dtype=torch.bfloat16) + activations_bf16[:5].normal_() + activations_bf16[alignment : alignment + 5].normal_() + activations, contiguous_scales = _ue8m0_per_token_quantize(activations_bf16) + column_major_scales = torch.empty_strided( + contiguous_scales.shape, + (1, rows), + dtype=contiguous_scales.dtype, + device=contiguous_scales.device, + ) + column_major_scales.copy_(contiguous_scales) + assert not column_major_scales.is_contiguous() + assert column_major_scales.stride() == (1, rows) + + weights_bf16 = torch.randn(groups, n, k, device="cuda", dtype=torch.bfloat16) + weights, weight_scales = _ue8m0_blockwise_weight_quantize(weights_bf16) + expected = torch.empty(rows, n, device="cuda", dtype=torch.bfloat16) + actual = torch.empty_like(expected) + deep_gemm.m_grouped_fp8_gemm_nt_contiguous( + (activations, contiguous_scales), + (weights, weight_scales), + expected, + psum, + use_psum_layout=True, + ) + deep_gemm.m_grouped_fp8_gemm_nt_contiguous( + (activations, column_major_scales), + (weights, weight_scales), + actual, + psum, + use_psum_layout=True, + ) + assert torch.equal(actual, expected) + + +def test_k_grouped_bf16_wgrad_supports_no_accumulator_and_empty_expert() -> None: + torch.manual_seed(6000) + counts = [128, 0, 128] + total, m, n = sum(counts), 256, 256 + left = torch.randn(total, m, device="cuda", dtype=torch.bfloat16) + right = torch.randn(total, n, device="cuda", dtype=torch.bfloat16) + output = torch.zeros(len(counts), m, n, device="cuda", dtype=torch.bfloat16) + grouped_layout = torch.tensor( + [128, 128, 256], device="cuda", dtype=torch.int32 + ) + + deep_gemm.k_grouped_bf16_gemm_tn_contiguous( + left, + right, + output, + None, + grouped_layout, + None, + use_psum_layout=True, + ) + torch.testing.assert_close( + output[0].float(), (left[:128].t() @ right[:128]).to(torch.bfloat16).float(), + rtol=2e-2, atol=0.2, + ) + assert torch.count_nonzero(output[1]) == 0 + torch.testing.assert_close( + output[2].float(), (left[128:].t() @ right[128:]).to(torch.bfloat16).float(), + rtol=2e-2, atol=0.2, + )