From a0cba68c6e1701023ce93e3eaef5ad08d9b5d23f Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 21 Jul 2026 19:33:00 +0800 Subject: [PATCH 01/29] Add SM103 FP8 block128 MegaMoE training backend --- README.md | 20 +- csrc/apis/gemm.hpp | 2 +- csrc/apis/sm103_fp8_block128.hpp | 9 + csrc/python_api.cpp | 9 + csrc/sm103_fp8_block128.cu | 967 ++++++++++++++++++ deep_gemm/__init__.py | 7 + deep_gemm/mega/__init__.py | 7 + deep_gemm/mega/fp8_block128.py | 684 +++++++++++++ setup.py | 48 +- tests/benchmark_fp8_block128_mega_moe.py | 204 ++++ tests/test_fp8_block128_capabilities.py | 32 + tests/test_fp8_block128_mega_moe.py | 245 +++++ .../test_fp8_block128_mega_moe_distributed.py | 245 +++++ tests/test_sm103_fp8_block128_primitives.py | 274 +++++ 14 files changed, 2744 insertions(+), 9 deletions(-) create mode 100644 csrc/apis/sm103_fp8_block128.hpp create mode 100644 csrc/sm103_fp8_block128.cu create mode 100644 deep_gemm/mega/fp8_block128.py create mode 100644 tests/benchmark_fp8_block128_mega_moe.py create mode 100644 tests/test_fp8_block128_capabilities.py create mode 100644 tests/test_fp8_block128_mega_moe.py create mode 100644 tests/test_fp8_block128_mega_moe_distributed.py create mode 100644 tests/test_sm103_fp8_block128_primitives.py diff --git a/README.md b/README.md index 6ef705ffce..93518bfed1 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,24 @@ 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. 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. + +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/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..0aa3794b2d --- /dev/null +++ b/csrc/sm103_fp8_block128.cu @@ -0,0 +1,967 @@ +#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 + +// 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 KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm103 final + : KernelSchedule1Sm, 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; + +#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 scale = amax == 0.0f ? 1.0f : amax / kE4M3Max; + if (threadIdx.x == 0) { + scales[row * num_blocks_k + block_k] = scale; + } + output[offset] = __nv_fp8_e4m3(value / scale); +#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 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_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 +} + +__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 +} + +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::PtrArrayTmaWarpSpecialized1Sm>::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::KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm103>::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; +}; + +template +torch::Tensor copy_metadata_to_device( + const std::vector& host, + const torch::TensorOptions& options, + cudaStream_t stream +) { + const int64_t num_bytes = static_cast(host.size() * sizeof(T)); + auto storage = torch::empty({std::max(num_bytes, 1)}, options.dtype(torch::kUInt8)); + if (num_bytes != 0) { + C10_CUDA_CHECK(cudaMemcpyAsync( + storage.data_ptr(), host.data(), num_bytes, cudaMemcpyHostToDevice, stream + )); + } + return storage; +} + +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 +) { + 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 total_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(count == 0 || count % 4 == 0, + "active group counts must be padded to a multiple of four"); + total_rows += count; + active_groups += count != 0; + } + TORCH_CHECK(total_rows == activations.size(0), "sum(group_counts) must equal activation rows"); + + auto output = torch::empty({total_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 += m; + } + + const auto metadata_options = activations.options().dtype(torch::kUInt8); + auto problems_device = copy_metadata_to_device(problems, metadata_options, stream); + auto ptr_a_device = copy_metadata_to_device(ptr_a, metadata_options, stream); + auto ptr_b_device = copy_metadata_to_device(ptr_b, metadata_options, stream); + auto ptr_c_device = copy_metadata_to_device(ptr_c, metadata_options, stream); + auto ptr_d_device = copy_metadata_to_device(ptr_d, metadata_options, stream); + auto ptr_sfa_device = copy_metadata_to_device(ptr_sfa, metadata_options, stream); + auto ptr_sfb_device = copy_metadata_to_device(ptr_sfb, metadata_options, stream); + auto stride_a_device = copy_metadata_to_device(stride_a, metadata_options, stream); + auto stride_b_device = copy_metadata_to_device(stride_b, metadata_options, stream); + auto stride_c_device = copy_metadata_to_device(stride_c, metadata_options, stream); + auto stride_d_device = copy_metadata_to_device(stride_d, metadata_options, stream); + auto layout_sfa_device = copy_metadata_to_device(layout_sfa, metadata_options, stream); + auto layout_sfb_device = copy_metadata_to_device(layout_sfb, metadata_options, stream); + + 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(problems_device.data_ptr()), + problems.data()}, + {reinterpret_cast(ptr_a_device.data_ptr()), + reinterpret_cast(stride_a_device.data_ptr()), + reinterpret_cast(ptr_b_device.data_ptr()), + reinterpret_cast(stride_b_device.data_ptr()), + reinterpret_cast(ptr_sfa_device.data_ptr()), + reinterpret_cast(layout_sfa_device.data_ptr()), + reinterpret_cast(ptr_sfb_device.data_ptr()), + reinterpret_cast(layout_sfb_device.data_ptr())}, + {{}, + reinterpret_cast(ptr_c_device.data_ptr()), + reinterpret_cast(stride_c_device.data_ptr()), + reinterpret_cast(ptr_d_device.data_ptr()), + reinterpret_cast(stride_d_device.data_ptr())}, + 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; +} + +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); +} + +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); +} + +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; +} + +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; +} + +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["fallback"] = py::none(); + result["native_symbols"] = py::make_tuple( + "sm103_fp8_block128_quantize", + "sm103_fp8_block128_dequantize", + "sm103_fp8_block128_grouped_gemm_nt", + "sm103_fp8_block128_grouped_gemm_nn", + "sm103_fp8_block128_swiglu_quantize", + "sm103_fp8_block128_swiglu_backward", + "sm103_fp8_block128_route_scale_quantize", + "sm103_fp8_block128_post_down_combine", + "sm103_fp8_block128_post_down_score_grad", + "sm103_fp8_block128_route_sum" + ); + return result; +} + +} // namespace + +void register_apis(pybind11::module_& m) { + m.def("get_sm103_fp8_block128_capabilities", &capabilities); + 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_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_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")); +} + +} // 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/mega/__init__.py b/deep_gemm/mega/__init__.py index bf6ea94230..4ba460c255 100644 --- a/deep_gemm/mega/__init__.py +++ b/deep_gemm/mega/__init__.py @@ -13,6 +13,13 @@ print(f'Failed to load mega kernels, please check your PyTorch version: {exception}') from .. import _C +from .fp8_block128 import ( + REQUIRED_NATIVE_SYMBOLS as FP8_BLOCK128_MEGAMOE_REQUIRED_NATIVE_SYMBOLS, + REQUIRED_PYTHON_SYMBOLS as FP8_BLOCK128_MEGAMOE_REQUIRED_PYTHON_SYMBOLS, + fp8_block128_mega_moe, + get_fp8_block128_mega_moe_capabilities, + transform_glm_w13_for_fp8_block128_mega_moe, +) class SymmBuffer: diff --git a/deep_gemm/mega/fp8_block128.py b/deep_gemm/mega/fp8_block128.py new file mode 100644 index 0000000000..6a3633bb3a --- /dev/null +++ b/deep_gemm/mega/fp8_block128.py @@ -0,0 +1,684 @@ +"""SM103-only distributed FP8-block128 MegaMoE training path. + +The public operation owns routed-token quantization, transport, expert compute, +POST_DOWN combine, and the complete routed backward. Weight tensors retain +GLM's E4M3 + FP32 128x128 block-scale contract; no BF16 weight dequantization, +MXFP4 transcode, or architecture fallback exists in this path. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from itertools import accumulate +from typing import Any, Sequence + +import torch +import torch.distributed as dist + +from .. import _C + + +_BLOCK = 128 +_PAD_ROWS = 128 + +REQUIRED_NATIVE_SYMBOLS = ( + "get_sm103_fp8_block128_capabilities", + "sm103_fp8_block128_quantize", + "sm103_fp8_block128_dequantize", + "sm103_fp8_block128_grouped_gemm_nt", + "sm103_fp8_block128_grouped_gemm_nn", + "sm103_fp8_block128_swiglu_quantize", + "sm103_fp8_block128_swiglu_backward", + "sm103_fp8_block128_route_scale_quantize", + "sm103_fp8_block128_post_down_combine", + "sm103_fp8_block128_post_down_score_grad", + "sm103_fp8_block128_route_sum", + "k_grouped_bf16_gemm_tn_contiguous", +) + +REQUIRED_PYTHON_SYMBOLS = ( + "fp8_block128_mega_moe", + "transform_glm_w13_for_fp8_block128_mega_moe", + "get_fp8_block128_mega_moe_capabilities", +) + + +def get_fp8_block128_mega_moe_capabilities() -> 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)] + native.update( + { + "native_symbols": REQUIRED_NATIVE_SYMBOLS, + "python_symbols": REQUIRED_PYTHON_SYMBOLS, + "forward": not missing, + "backward": not missing, + "distributed_transport": "torch.distributed.all_to_all_single", + "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]: + """Convert canonical ``[gate, up]`` GLM storage to ``[up; gate]``. + + ``canonical_weight`` is ``[2E, H, D]`` and ``canonical_scale`` is + ``[2E, H/128, D/128]``. Returned tensors are contiguous + ``[E, 2H, D]`` and ``[E, 2H/128, D/128]`` respectively. + """ + 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:] + if hidden % _BLOCK or model_dim % _BLOCK: + raise ValueError("W13 dimensions must be divisible by 128") + expected_scale_shape = (experts * 2, hidden // _BLOCK, model_dim // _BLOCK) + if tuple(canonical_scale.shape) != expected_scale_shape: + raise ValueError( + f"canonical W13 scale shape must be {expected_scale_shape}, got {tuple(canonical_scale.shape)}" + ) + # Canonical pair index 0 is gate and 1 is up. The fused preactivation ABI + # requires up first, followed by gate. + pair_order = torch.tensor([1, 0], dtype=torch.int64, device=canonical_weight.device) + active_weight = ( + canonical_weight.view(experts, 2, hidden, model_dim) + .index_select(1, pair_order) + .reshape(experts, hidden * 2, model_dim) + .contiguous() + ) + active_scale = ( + canonical_scale.view( + experts, 2, hidden // _BLOCK, model_dim // _BLOCK + ) + .index_select(1, pair_order) + .reshape(experts, hidden * 2 // _BLOCK, model_dim // _BLOCK) + .contiguous() + ) + return active_weight, active_scale + + +def _active_w13_grad_to_canonical(active_grad: torch.Tensor) -> torch.Tensor: + experts, doubled_hidden, model_dim = active_grad.shape + hidden = doubled_hidden // 2 + up, gate = active_grad.view(experts, 2, hidden, model_dim).unbind(dim=1) + return torch.stack((gate, up), dim=1).reshape(experts * 2, hidden, model_dim).contiguous() + + +@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(): + if group is not None: + raise RuntimeError("a process group was provided before torch.distributed initialization") + return _GroupState(group=None, rank=0, world_size=1) + return _GroupState( + group=group, + rank=dist.get_rank(group), + world_size=dist.get_world_size(group), + ) + + +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 _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, + w13_master: torch.Tensor, + w2_master: torch.Tensor, + group_state: _GroupState, +) -> tuple[int, int, int, int, int]: + if not x.is_cuda: + raise ValueError("FP8-block128 MegaMoE requires CUDA") + if torch.cuda.get_device_capability(x.device) != (10, 3): + capability = torch.cuda.get_device_capability(x.device) + raise RuntimeError( + f"FP8-block128 MegaMoE is SM103-only; runtime capability is {capability}. " + "No fallback is available." + ) + 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) + _check_tensor(w13_master, name="w13_master", ndim=3, dtype=torch.bfloat16, device=device) + _check_tensor(w2_master, name="w2_master", ndim=3, dtype=torch.bfloat16, device=device) + + tokens, model_dim = x.shape + if model_dim % _BLOCK: + raise ValueError("model dimension must be divisible by 128") + if topk_ids.shape != topk_scores.shape or topk_ids.shape[0] != tokens: + raise ValueError("top-k IDs/scores must have identical [tokens, top_k] shape") + topk = topk_ids.shape[1] + if topk <= 0: + raise ValueError("top_k must be positive") + + local_experts, doubled_hidden, w13_k = w13_weight.shape + if local_experts <= 0 or doubled_hidden % (2 * _BLOCK) or w13_k != model_dim: + raise ValueError("W13 must have shape [local_experts, 2H, D] with D/H divisible by 128") + hidden = doubled_hidden // 2 + if tuple(w13_scale.shape) != ( + local_experts, + doubled_hidden // _BLOCK, + model_dim // _BLOCK, + ): + raise ValueError("W13 scale shape does not match 128x128 weight blocks") + if tuple(w2_weight.shape) != (local_experts, model_dim, hidden): + raise ValueError("W2 must have shape [local_experts, D, H]") + if tuple(w2_scale.shape) != ( + local_experts, + model_dim // _BLOCK, + hidden // _BLOCK, + ): + raise ValueError("W2 scale shape does not match 128x128 weight blocks") + if tuple(w13_master.shape) != (local_experts * 2, hidden, model_dim): + raise ValueError("canonical BF16 W13 master must have shape [2E, H, D]") + if tuple(w2_master.shape) != tuple(w2_weight.shape): + raise ValueError("BF16 W2 master shape must match W2") + + global_experts = local_experts * group_state.world_size + if topk_ids.numel(): + minimum, maximum = torch.aminmax(topk_ids) + if minimum.item() < 0 or maximum.item() >= global_experts: + raise ValueError( + f"top-k IDs must lie in [0, {global_experts}); got [{minimum.item()}, {maximum.item()}]" + ) + return tokens, model_dim, hidden, local_experts, topk + + +def _exchange_counts( + send_counts: Sequence[int], group_state: _GroupState, device: torch.device +) -> list[int]: + if group_state.world_size == 1: + return list(send_counts) + send = torch.tensor(send_counts, device=device, dtype=torch.int64) + receive = torch.empty_like(send) + dist.all_to_all_single(receive, send, group=group_state.group) + return [int(value) for value in receive.cpu().tolist()] + + +def _all_to_all_rows( + tensor: torch.Tensor, + send_counts: Sequence[int], + receive_counts: Sequence[int], + group_state: _GroupState, +) -> torch.Tensor: + if group_state.world_size == 1: + return tensor + output = torch.empty( + (sum(receive_counts), *tensor.shape[1:]), + dtype=tensor.dtype, + device=tensor.device, + ) + source = tensor.view(torch.uint8) if tensor.dtype == torch.float8_e4m3fn else tensor + destination = output.view(torch.uint8) if output.dtype == torch.float8_e4m3fn else output + dist.all_to_all_single( + destination, + source, + output_split_sizes=list(receive_counts), + input_split_sizes=list(send_counts), + group=group_state.group, + ) + return output + + +def _inverse_permutation(order: torch.Tensor) -> torch.Tensor: + inverse = torch.empty_like(order) + inverse.scatter_(0, order, torch.arange(order.numel(), device=order.device)) + return inverse + + +def _padding_state( + counts: Sequence[int], device: torch.device +) -> tuple[list[int], torch.Tensor]: + padded_counts = [((count + _PAD_ROWS - 1) // _PAD_ROWS) * _PAD_ROWS if count else 0 for count in counts] + total_actual = sum(counts) + if total_actual == 0: + return padded_counts, torch.empty(0, dtype=torch.int64, device=device) + count_tensor = torch.tensor(counts, dtype=torch.int64, device=device) + padded_tensor = torch.tensor(padded_counts, dtype=torch.int64, device=device) + group_ids = torch.repeat_interleave( + torch.arange(len(counts), dtype=torch.int64, device=device), count_tensor + ) + padding_before = torch.cumsum(padded_tensor - count_tensor, dim=0) - ( + padded_tensor - count_tensor + ) + actual_to_padded = torch.arange(total_actual, dtype=torch.int64, device=device) + actual_to_padded.add_(padding_before.index_select(0, group_ids)) + return padded_counts, actual_to_padded + + +def _pad_rows( + tensor: torch.Tensor, + actual_to_padded: torch.Tensor, + padded_rows: int, + *, + fill_value: float, +) -> torch.Tensor: + output = torch.full( + (padded_rows, *tensor.shape[1:]), + fill_value, + dtype=tensor.dtype, + device=tensor.device, + ) + if tensor.shape[0]: + if tensor.dtype == torch.float8_e4m3fn: + output.view(torch.uint8).index_copy_( + 0, actual_to_padded, tensor.view(torch.uint8) + ) + else: + output.index_copy_(0, actual_to_padded, tensor) + return output + + +def _unpad_rows(tensor: torch.Tensor, actual_to_padded: torch.Tensor) -> torch.Tensor: + if tensor.dtype == torch.float8_e4m3fn: + output = tensor.view(torch.uint8).index_select(0, actual_to_padded) + return output.view(torch.float8_e4m3fn) + return tensor.index_select(0, actual_to_padded) + + +def _bf16_grouped_wgrad( + left: torch.Tensor, + right: torch.Tensor, + padded_counts: Sequence[int], +) -> torch.Tensor: + output = torch.zeros( + (len(padded_counts), left.shape[1], right.shape[1]), + dtype=torch.bfloat16, + device=left.device, + ) + if left.shape[0] == 0: + return output + grouped_layout = torch.tensor( + list(accumulate(padded_counts)), dtype=torch.int32, device=left.device + ) + _C.k_grouped_bf16_gemm_tn_contiguous( + left.contiguous(), + right.contiguous(), + output, + None, + grouped_layout, + None, + "mn", + True, + ) + return output + + +class _FP8Block128MegaMoE(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, + w13_master: torch.Tensor, + w2_master: torch.Tensor, + group: Any, + ) -> torch.Tensor: + group_state = _resolve_group(group) + tokens, model_dim, hidden, local_experts, topk = _validate_inputs( + x, + topk_ids, + topk_scores, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + w13_master, + w2_master, + group_state, + ) + + with torch.autograd.profiler.record_function( + "sm103_fp8_block128_megamoe_forward" + ): + flat_ids = topk_ids.flatten() + send_order = torch.argsort(flat_ids, stable=True) + sorted_ids = flat_ids.index_select(0, send_order) + destinations = torch.div(sorted_ids, local_experts, rounding_mode="floor") + send_counts = [ + int(value) + for value in torch.bincount( + destinations, minlength=group_state.world_size + ) + .cpu() + .tolist() + ] + receive_counts = _exchange_counts(send_counts, group_state, x.device) + + token_quantized, token_scales = _C.sm103_fp8_block128_quantize(x) + sorted_tokens = torch.div(send_order, topk, rounding_mode="floor") + send_activations = token_quantized.index_select(0, sorted_tokens) + send_activation_scales = token_scales.index_select(0, sorted_tokens) + receive_activations = _all_to_all_rows( + send_activations, send_counts, receive_counts, group_state + ) + receive_activation_scales = _all_to_all_rows( + send_activation_scales, send_counts, receive_counts, group_state + ) + receive_ids = _all_to_all_rows( + sorted_ids, send_counts, receive_counts, group_state + ) + + local_ids = torch.remainder(receive_ids, local_experts) + group_order = torch.argsort(local_ids, stable=True) + ungroup_order = _inverse_permutation(group_order) + grouped_activations = receive_activations.index_select(0, group_order) + grouped_activation_scales = receive_activation_scales.index_select( + 0, group_order + ) + grouped_local_ids = local_ids.index_select(0, group_order) + actual_counts = [ + int(value) + for value in torch.bincount( + grouped_local_ids, minlength=local_experts + ) + .cpu() + .tolist() + ] + padded_counts, actual_to_padded = _padding_state( + actual_counts, x.device + ) + padded_rows = sum(padded_counts) + padded_activations = _pad_rows( + grouped_activations, + actual_to_padded, + padded_rows, + fill_value=0, + ) + padded_activation_scales = _pad_rows( + grouped_activation_scales, + actual_to_padded, + padded_rows, + fill_value=1, + ) + + preactivation = _C.sm103_fp8_block128_grouped_gemm_nt( + padded_activations, + padded_activation_scales, + w13_weight, + w13_scale, + padded_counts, + ) + hidden_quantized, hidden_scales = ( + _C.sm103_fp8_block128_swiglu_quantize(preactivation) + ) + routed_output_padded = _C.sm103_fp8_block128_grouped_gemm_nt( + hidden_quantized, + hidden_scales, + w2_weight, + w2_scale, + padded_counts, + ) + routed_output_grouped = _unpad_rows( + routed_output_padded, actual_to_padded + ) + routed_output_receive_order = routed_output_grouped.index_select( + 0, ungroup_order + ) + routed_output_send_order = _all_to_all_rows( + routed_output_receive_order, + receive_counts, + send_counts, + group_state, + ) + routed_output = torch.empty( + (tokens * topk, model_dim), + dtype=torch.bfloat16, + device=x.device, + ) + if routed_output.shape[0]: + routed_output.index_copy_( + 0, send_order, routed_output_send_order + ) + output = _C.sm103_fp8_block128_post_down_combine( + routed_output, topk_scores + ) + + ctx.group_state = group_state + ctx.send_counts = send_counts + ctx.receive_counts = receive_counts + ctx.padded_counts = padded_counts + ctx.tokens = tokens + ctx.model_dim = model_dim + ctx.hidden = hidden + ctx.local_experts = local_experts + ctx.topk = topk + ctx.save_for_backward( + topk_scores, + send_order, + group_order, + ungroup_order, + actual_to_padded, + padded_activations, + padded_activation_scales, + preactivation, + hidden_quantized, + hidden_scales, + routed_output, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) + return output + + @staticmethod + def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: + ( + topk_scores, + send_order, + group_order, + ungroup_order, + actual_to_padded, + padded_activations, + padded_activation_scales, + preactivation, + hidden_quantized, + hidden_scales, + routed_output, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) = ctx.saved_tensors + 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_backward" + ): + grad_scores = _C.sm103_fp8_block128_post_down_score_grad( + routed_output, grad_output, ctx.topk + ) + grad_route_quantized_send, grad_route_scales_send = ( + _C.sm103_fp8_block128_route_scale_quantize( + grad_output, topk_scores, send_order + ) + ) + grad_route_quantized_receive = _all_to_all_rows( + grad_route_quantized_send, + ctx.send_counts, + ctx.receive_counts, + ctx.group_state, + ) + grad_route_scales_receive = _all_to_all_rows( + grad_route_scales_send, + ctx.send_counts, + ctx.receive_counts, + ctx.group_state, + ) + grad_route_quantized_grouped = grad_route_quantized_receive.index_select( + 0, group_order + ) + grad_route_scales_grouped = grad_route_scales_receive.index_select( + 0, group_order + ) + padded_rows = sum(ctx.padded_counts) + grad_route_quantized = _pad_rows( + grad_route_quantized_grouped, + actual_to_padded, + padded_rows, + fill_value=0, + ) + grad_route_scales = _pad_rows( + grad_route_scales_grouped, + actual_to_padded, + padded_rows, + fill_value=1, + ) + + grad_hidden = _C.sm103_fp8_block128_grouped_gemm_nn( + grad_route_quantized, + grad_route_scales, + w2_weight, + w2_scale, + ctx.padded_counts, + ) + grad_preactivation = _C.sm103_fp8_block128_swiglu_backward( + grad_hidden, preactivation + ) + grad_preactivation_quantized, grad_preactivation_scales = ( + _C.sm103_fp8_block128_quantize(grad_preactivation) + ) + grad_input_padded = _C.sm103_fp8_block128_grouped_gemm_nn( + grad_preactivation_quantized, + grad_preactivation_scales, + w13_weight, + w13_scale, + ctx.padded_counts, + ) + + grad_route_dequantized = _C.sm103_fp8_block128_dequantize( + grad_route_quantized, grad_route_scales + ) + hidden_dequantized = _C.sm103_fp8_block128_dequantize( + hidden_quantized, hidden_scales + ) + grad_w2 = _bf16_grouped_wgrad( + grad_route_dequantized, + hidden_dequantized, + ctx.padded_counts, + ) + input_dequantized = _C.sm103_fp8_block128_dequantize( + padded_activations, padded_activation_scales + ) + grad_w13_active = _bf16_grouped_wgrad( + grad_preactivation, + input_dequantized, + ctx.padded_counts, + ) + grad_w13 = _active_w13_grad_to_canonical(grad_w13_active) + + grad_input_grouped = _unpad_rows( + grad_input_padded, actual_to_padded + ) + grad_input_receive_order = grad_input_grouped.index_select( + 0, ungroup_order + ) + grad_input_send_order = _all_to_all_rows( + grad_input_receive_order, + ctx.receive_counts, + ctx.send_counts, + ctx.group_state, + ) + grad_input_routes = torch.empty( + (ctx.tokens * ctx.topk, ctx.model_dim), + dtype=torch.bfloat16, + device=grad_output.device, + ) + if grad_input_routes.shape[0]: + grad_input_routes.index_copy_( + 0, send_order, grad_input_send_order + ) + grad_input = _C.sm103_fp8_block128_route_sum( + grad_input_routes, ctx.tokens, ctx.topk + ) + + return ( + grad_input, + None, + grad_scores, + None, + None, + None, + None, + grad_w13, + grad_w2, + 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, + w13_master: torch.Tensor, + w2_master: torch.Tensor, + group: Any = None, +) -> torch.Tensor: + """Run the complete SM103 FP8-block128 routed branch. + + W13 quantized tensors use active ``[up; gate]`` ordering while the BF16 + W13 master remains canonical interleaved ``[gate, up]`` storage. Route + scores are applied only after W2 and their gradients are accumulated in + FP32. ``group`` is the expert-parallel process group; no other token + transport may wrap this operation. + """ + return _FP8Block128MegaMoE.apply( + x, + topk_ids, + topk_scores, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + w13_master, + w2_master, + group, + ) 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..d3870ef822 --- /dev/null +++ b/tests/benchmark_fp8_block128_mega_moe.py @@ -0,0 +1,204 @@ +"""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:: + + python tests/benchmark_fp8_block128_mega_moe.py \ + --tokens 15625 --experts 16 --model-dim 6144 --hidden 2048 --topk 8 + +The process must be launched on an otherwise idle SM103 GPU. The benchmark +prints one JSON object so callers can archive the exact dimensions and metrics. +""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +from typing import Callable + +import torch + +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() + + 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.warmup < 0 or args.iterations <= 0: + raise ValueError("warmup must be non-negative and iterations must be positive") + + torch.manual_seed(20260721) + device = torch.device("cuda") + x = ( + torch.randn(args.tokens, args.model_dim, device=device, dtype=torch.bfloat16) + * 0.02 + ).requires_grad_() + canonical_w13 = ( + torch.randn( + args.experts * 2, + args.hidden, + args.model_dim, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + w2_master = ( + torch.randn( + args.experts, + args.model_dim, + args.hidden, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + canonical_w13_q, canonical_w13_s = _blockwise_quantize(canonical_w13.detach()) + w13_q, w13_s = deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( + canonical_w13_q, canonical_w13_s + ) + 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, + canonical_w13, + w2_master, + ) + + def forward_backward() -> None: + forward() + assert latest_output is not None + latest_output.backward(upstream) + x.grad = None + scores.grad = None + canonical_w13.grad = None + w2_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 + + 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, + "topk": args.topk, + "model_dim": args.model_dim, + "hidden": args.hidden, + }, + "warmup": args.warmup, + "iterations": args.iterations, + "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, + } + print(json.dumps(result, sort_keys=True)) + + +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..7a0e92a918 --- /dev/null +++ b/tests/test_fp8_block128_capabilities.py @@ -0,0 +1,32 @@ +"""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["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..229a89719f --- /dev/null +++ b/tests/test_fp8_block128_mega_moe.py @@ -0,0 +1,245 @@ +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 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_() + canonical_w13 = ( + torch.randn(experts * 2, 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_q, canonical_w13_s = _blockwise_weight_quantize(canonical_w13.detach()) + w13_q, w13_s = deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( + canonical_w13_q, canonical_w13_s + ) + 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, + "w13_master": canonical_w13, + "w2_master": w2_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() + w13 = _blockwise_weight_dequantize(case["w13_q"], case["w13_s"]) + 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_() + canonical_w13 = case["w13_master"].detach().clone().requires_grad_() + w2_master = case["w2_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_dequantized = _blockwise_weight_dequantize( + case["w13_q"], case["w13_s"] + ) + canonical_pairs = canonical_w13.view(experts, 2, hidden, model_dim) + w13_active_master = torch.stack( + (canonical_pairs[:, 1], canonical_pairs[:, 0]), dim=1 + ).reshape(experts, hidden * 2, model_dim) + w13_effective = w13_active_master.float() + ( + w13_dequantized - 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(), + "w13_grad": canonical_w13.grad.detach(), + "w2_grad": w2_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) + ) + + +@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["w13_master"], + case["w2_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["w13_master"], + case["w2_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["w13_master"].grad.shape == case["w13_master"].shape + assert case["w13_master"].grad.dtype == torch.bfloat16 + assert case["w2_master"].grad.shape == case["w2_master"].shape + assert case["w2_master"].grad.dtype == torch.bfloat16 + assert _normalized_difference(case["x"].grad, expected["x_grad"]) < 0.12 + assert _normalized_difference(case["w13_master"].grad, expected["w13_grad"]) < 0.15 + assert _normalized_difference(case["w2_master"].grad, expected["w2_grad"]) < 0.12 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..9870aeef5c --- /dev/null +++ b/tests/test_fp8_block128_mega_moe_distributed.py @@ -0,0 +1,245 @@ +"""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 + + +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, experts = 5 + rank * 2, 2, 4 + 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, torch.Tensor, torch.Tensor]: + x_ref = x.detach().clone().requires_grad_() + score_ref = scores.detach().clone().requires_grad_() + x_q, x_s = deep_gemm._C.sm103_fp8_block128_quantize(x_ref.detach()) + x_dequantized = deep_gemm._C.sm103_fp8_block128_dequantize(x_q, x_s).float() + x_effective = x_ref.float() + (x_dequantized - x_ref.float()).detach() + 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_effective.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_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.index_select(0, flat_ids), hidden_effective.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) + output.backward(upstream) + return output.detach(), x_ref.grad.detach(), score_ref.grad.detach() + + +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 + + assert torch.cuda.get_device_capability(device) == (10, 3) + 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, full_w13_s = ( + 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[expert_start:expert_end].contiguous() + local_w13_s = full_w13_s[expert_start:expert_end].contiguous() + local_w2_q = full_w2_q[expert_start:expert_end].contiguous() + local_w2_s = full_w2_s[expert_start:expert_end].contiguous() + local_w13_master = ( + full_w13_master[expert_start * 2 : expert_end * 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) + expected_output, expected_x_grad, expected_score_grad = _reference( + deep_gemm, + x, + ids, + scores, + full_w13_q, + full_w13_s, + full_w2_q, + full_w2_s, + upstream, + ) + + output = deep_gemm.fp8_block128_mega_moe( + x, + ids, + scores, + local_w13_q, + local_w13_s, + local_w2_q, + local_w2_s, + local_w13_master, + local_w2_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 + for gradient in (x.grad, local_w13_master.grad, local_w2_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_w13_master.grad[2:]) == 0 + assert torch.count_nonzero(local_w2_master.grad[1]) == 0 + 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..1812de7da4 --- /dev/null +++ b/tests/test_sm103_fp8_block128_primitives.py @@ -0,0 +1,274 @@ +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 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 + ) + + +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_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, + ) From 6b307df547fe474cc45e3d3c6d5a60738eb4c330 Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 21 Jul 2026 20:03:13 +0800 Subject: [PATCH 02/29] Avoid packed GLM master and active-weight copies --- README.md | 11 +- csrc/sm103_fp8_block128.cu | 283 ++++++++++++++++++ deep_gemm/mega/fp8_block128.py | 115 ++++--- tests/benchmark_fp8_block128_mega_moe.py | 28 +- tests/test_fp8_block128_mega_moe.py | 63 ++-- .../test_fp8_block128_mega_moe_distributed.py | 37 ++- tests/test_sm103_fp8_block128_primitives.py | 49 +++ 7 files changed, 508 insertions(+), 78 deletions(-) diff --git a/README.md b/README.md index 93518bfed1..2b4e49e01a 100644 --- a/README.md +++ b/README.md @@ -144,10 +144,13 @@ For the full example with multi-process setup and benchmarking, please refer to `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. 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. +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. The backend requires CUDA compute capability exactly 10.3. It has no SM100, SM90, generic, or compatibility fallback. Use diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 0aa3794b2d..c15c3070fd 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -269,6 +269,38 @@ __global__ void sm103_swiglu_backward_kernel( #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, @@ -677,6 +709,217 @@ torch::Tensor grouped_fp8_block128_gemm_impl( 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 +) { + // 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 total_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(count == 0 || count % 4 == 0, + "active group counts must be padded to a multiple of four"); + total_rows += count; + active_groups += count != 0; + } + TORCH_CHECK(total_rows == activations.size(0), + "sum(group_counts) must equal activation rows"); + + auto output = torch::empty( + {total_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 += m; + } + + const auto metadata_options = activations.options().dtype(torch::kUInt8); + auto problems_device = copy_metadata_to_device(problems, metadata_options, stream); + auto ptr_a_device = copy_metadata_to_device(ptr_a, metadata_options, stream); + auto ptr_b_device = copy_metadata_to_device(ptr_b, metadata_options, stream); + auto ptr_c_device = copy_metadata_to_device(ptr_c, metadata_options, stream); + auto ptr_d_device = copy_metadata_to_device(ptr_d, metadata_options, stream); + auto ptr_sfa_device = copy_metadata_to_device(ptr_sfa, metadata_options, stream); + auto ptr_sfb_device = copy_metadata_to_device(ptr_sfb, metadata_options, stream); + auto stride_a_device = copy_metadata_to_device(stride_a, metadata_options, stream); + auto stride_b_device = copy_metadata_to_device(stride_b, metadata_options, stream); + auto stride_c_device = copy_metadata_to_device(stride_c, metadata_options, stream); + auto stride_d_device = copy_metadata_to_device(stride_d, metadata_options, stream); + auto layout_sfa_device = copy_metadata_to_device(layout_sfa, metadata_options, stream); + auto layout_sfb_device = copy_metadata_to_device(layout_sfb, metadata_options, stream); + + 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(problems_device.data_ptr()), + problems.data()}, + {reinterpret_cast(ptr_a_device.data_ptr()), + reinterpret_cast(stride_a_device.data_ptr()), + reinterpret_cast(ptr_b_device.data_ptr()), + reinterpret_cast(stride_b_device.data_ptr()), + reinterpret_cast(ptr_sfa_device.data_ptr()), + reinterpret_cast(layout_sfa_device.data_ptr()), + reinterpret_cast(ptr_sfb_device.data_ptr()), + reinterpret_cast(layout_sfb_device.data_ptr())}, + {{}, + reinterpret_cast(ptr_c_device.data_ptr()), + reinterpret_cast(stride_c_device.data_ptr()), + reinterpret_cast(ptr_d_device.data_ptr()), + reinterpret_cast(stride_d_device.data_ptr())}, + 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, @@ -790,6 +1033,37 @@ torch::Tensor swiglu_backward( 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, @@ -925,8 +1199,10 @@ pybind11::dict capabilities() { "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_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", @@ -950,10 +1226,17 @@ void register_apis(pybind11::module_& m) { 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_impl, + 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, diff --git a/deep_gemm/mega/fp8_block128.py b/deep_gemm/mega/fp8_block128.py index 6a3633bb3a..e4472ff577 100644 --- a/deep_gemm/mega/fp8_block128.py +++ b/deep_gemm/mega/fp8_block128.py @@ -27,8 +27,10 @@ "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_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", @@ -103,13 +105,6 @@ def transform_glm_w13_for_fp8_block128_mega_moe( return active_weight, active_scale -def _active_w13_grad_to_canonical(active_grad: torch.Tensor) -> torch.Tensor: - experts, doubled_hidden, model_dim = active_grad.shape - hidden = doubled_hidden // 2 - up, gate = active_grad.view(experts, 2, hidden, model_dim).unbind(dim=1) - return torch.stack((gate, up), dim=1).reshape(experts * 2, hidden, model_dim).contiguous() - - @dataclass(frozen=True) class _GroupState: group: Any @@ -147,6 +142,17 @@ def _check_tensor( raise ValueError(f"{name} must be contiguous") +def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: + """Return a plain local tensor without taking ownership from FireTitan. + + Expert masters may be DTensors. Their placement and gradient reduction + remain a FireTitan concern; this companion only validates their resident + local slice and returns local BF16 gradients through the optional wrapper. + """ + to_local = getattr(tensor, "to_local", None) + return to_local() if callable(to_local) else tensor + + def _validate_inputs( x: torch.Tensor, topk_ids: torch.Tensor, @@ -155,8 +161,9 @@ def _validate_inputs( w13_scale: torch.Tensor, w2_weight: torch.Tensor, w2_scale: torch.Tensor, - w13_master: torch.Tensor, + w1_master: torch.Tensor, w2_master: torch.Tensor, + w3_master: torch.Tensor, group_state: _GroupState, ) -> tuple[int, int, int, int, int]: if not x.is_cuda: @@ -175,8 +182,12 @@ def _validate_inputs( _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) - _check_tensor(w13_master, name="w13_master", ndim=3, dtype=torch.bfloat16, device=device) - _check_tensor(w2_master, name="w2_master", ndim=3, dtype=torch.bfloat16, device=device) + w1_master_local = _local_tensor(w1_master) + w2_master_local = _local_tensor(w2_master) + w3_master_local = _local_tensor(w3_master) + _check_tensor(w1_master_local, name="w1_master", ndim=3, dtype=torch.bfloat16, device=device) + _check_tensor(w2_master_local, name="w2_master", ndim=3, dtype=torch.bfloat16, device=device) + _check_tensor(w3_master_local, name="w3_master", ndim=3, dtype=torch.bfloat16, device=device) tokens, model_dim = x.shape if model_dim % _BLOCK: @@ -187,16 +198,18 @@ def _validate_inputs( if topk <= 0: raise ValueError("top_k must be positive") - local_experts, doubled_hidden, w13_k = w13_weight.shape - if local_experts <= 0 or doubled_hidden % (2 * _BLOCK) or w13_k != model_dim: - raise ValueError("W13 must have shape [local_experts, 2H, D] with D/H divisible by 128") - hidden = doubled_hidden // 2 + if w13_weight.shape[0] % 2: + raise ValueError("canonical W13 must contain gate/up pairs") + local_experts = w13_weight.shape[0] // 2 + hidden, w13_k = w13_weight.shape[1:] + if local_experts <= 0 or hidden % _BLOCK or w13_k != model_dim: + raise ValueError("canonical W13 must have shape [2E, H, D] with D/H divisible by 128") if tuple(w13_scale.shape) != ( - local_experts, - doubled_hidden // _BLOCK, + local_experts * 2, + hidden // _BLOCK, model_dim // _BLOCK, ): - raise ValueError("W13 scale shape does not match 128x128 weight blocks") + raise ValueError("canonical W13 scale shape does not match 128x128 weight blocks") if tuple(w2_weight.shape) != (local_experts, model_dim, hidden): raise ValueError("W2 must have shape [local_experts, D, H]") if tuple(w2_scale.shape) != ( @@ -205,9 +218,12 @@ def _validate_inputs( hidden // _BLOCK, ): raise ValueError("W2 scale shape does not match 128x128 weight blocks") - if tuple(w13_master.shape) != (local_experts * 2, hidden, model_dim): - raise ValueError("canonical BF16 W13 master must have shape [2E, H, D]") - if tuple(w2_master.shape) != tuple(w2_weight.shape): + master_w13_shape = (local_experts, hidden, model_dim) + if tuple(w1_master_local.shape) != master_w13_shape: + raise ValueError("BF16 gate master must have shape [E, H, D]") + if tuple(w3_master_local.shape) != master_w13_shape: + raise ValueError("BF16 up master must have shape [E, H, D]") + if tuple(w2_master_local.shape) != tuple(w2_weight.shape): raise ValueError("BF16 W2 master shape must match W2") global_experts = local_experts * group_state.world_size @@ -351,9 +367,11 @@ def forward( w13_scale: torch.Tensor, w2_weight: torch.Tensor, w2_scale: torch.Tensor, - w13_master: 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) tokens, model_dim, hidden, local_experts, topk = _validate_inputs( @@ -364,10 +382,13 @@ def forward( w13_scale, w2_weight, w2_scale, - w13_master, + w1_master, w2_master, + w3_master, group_state, ) + if master_gradient_wrapper is not None and not callable(master_gradient_wrapper): + raise TypeError("master_gradient_wrapper must be callable or None") with torch.autograd.profiler.record_function( "sm103_fp8_block128_megamoe_forward" @@ -433,7 +454,7 @@ def forward( fill_value=1, ) - preactivation = _C.sm103_fp8_block128_grouped_gemm_nt( + preactivation = _C.sm103_fp8_block128_grouped_w13_gemm_nt_canonical( padded_activations, padded_activation_scales, w13_weight, @@ -484,6 +505,7 @@ def forward( ctx.hidden = hidden ctx.local_experts = local_experts ctx.topk = topk + ctx.master_gradient_wrapper = master_gradient_wrapper ctx.save_for_backward( topk_scores, send_order, @@ -576,7 +598,7 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: w2_scale, ctx.padded_counts, ) - grad_preactivation = _C.sm103_fp8_block128_swiglu_backward( + grad_preactivation = _C.sm103_fp8_block128_swiglu_backward_canonical( grad_hidden, preactivation ) grad_preactivation_quantized, grad_preactivation_scales = ( @@ -585,8 +607,14 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: grad_input_padded = _C.sm103_fp8_block128_grouped_gemm_nn( grad_preactivation_quantized, grad_preactivation_scales, - w13_weight, - w13_scale, + w13_weight.view( + ctx.local_experts, ctx.hidden * 2, ctx.model_dim + ), + w13_scale.view( + ctx.local_experts, + ctx.hidden * 2 // _BLOCK, + ctx.model_dim // _BLOCK, + ), ctx.padded_counts, ) @@ -604,12 +632,13 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: input_dequantized = _C.sm103_fp8_block128_dequantize( padded_activations, padded_activation_scales ) - grad_w13_active = _bf16_grouped_wgrad( + grad_w13_canonical = _bf16_grouped_wgrad( grad_preactivation, input_dequantized, ctx.padded_counts, ) - grad_w13 = _active_w13_grad_to_canonical(grad_w13_active) + grad_w1 = grad_w13_canonical[:, : ctx.hidden].contiguous() + grad_w3 = grad_w13_canonical[:, ctx.hidden :].contiguous() grad_input_grouped = _unpad_rows( grad_input_padded, actual_to_padded @@ -636,6 +665,10 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: grad_input_routes, ctx.tokens, ctx.topk ) + 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_input, None, @@ -644,8 +677,10 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: None, None, None, - grad_w13, + grad_w1, grad_w2, + grad_w3, + None, None, ) @@ -658,17 +693,23 @@ def fp8_block128_mega_moe( w13_scale: torch.Tensor, w2_weight: torch.Tensor, w2_scale: torch.Tensor, - w13_master: 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 the complete SM103 FP8-block128 routed branch. - W13 quantized tensors use active ``[up; gate]`` ordering while the BF16 - W13 master remains canonical interleaved ``[gate, up]`` storage. Route - scores are applied only after W2 and their gradients are accumulated in - FP32. ``group`` is the expert-parallel process group; no other token - transport may wrap this operation. + Quantized W13 tensors retain FireTitan's canonical interleaved + ``[gate, up]`` storage. The native grouped W13 launch presents its + preactivation as ``[up; gate]`` without copying that storage. BF16 gate, + down, and up masters are passed separately, matching their checkpoint + FQNs without a packed-master allocation. Route scores are applied only + after W2 and their gradients are accumulated in FP32. ``group`` is the + expert-parallel process group; no other token transport may wrap this + operation. ``master_gradient_wrapper`` lets the owning framework restore + DTensor placement/reduction metadata to the three local BF16 gradients. """ return _FP8Block128MegaMoE.apply( x, @@ -678,7 +719,9 @@ def fp8_block128_mega_moe( w13_scale, w2_weight, w2_scale, - w13_master, + w1_master, w2_master, + w3_master, group, + master_gradient_wrapper, ) diff --git a/tests/benchmark_fp8_block128_mega_moe.py b/tests/benchmark_fp8_block128_mega_moe.py index d3870ef822..dc75b74881 100644 --- a/tests/benchmark_fp8_block128_mega_moe.py +++ b/tests/benchmark_fp8_block128_mega_moe.py @@ -100,9 +100,19 @@ def main() -> None: torch.randn(args.tokens, args.model_dim, device=device, dtype=torch.bfloat16) * 0.02 ).requires_grad_() - canonical_w13 = ( + w1_master = ( torch.randn( - args.experts * 2, + args.experts, + args.hidden, + args.model_dim, + device=device, + dtype=torch.bfloat16, + ) + * 0.02 + ).requires_grad_() + w3_master = ( + torch.randn( + args.experts, args.hidden, args.model_dim, device=device, @@ -120,9 +130,11 @@ def main() -> None: ) * 0.02 ).requires_grad_() - canonical_w13_q, canonical_w13_s = _blockwise_quantize(canonical_w13.detach()) - w13_q, w13_s = deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( - canonical_w13_q, canonical_w13_s + 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) @@ -149,8 +161,9 @@ def forward() -> None: w13_s, w2_q, w2_s, - canonical_w13, + w1_master, w2_master, + w3_master, ) def forward_backward() -> None: @@ -159,8 +172,9 @@ def forward_backward() -> None: latest_output.backward(upstream) x.grad = None scores.grad = None - canonical_w13.grad = None + w1_master.grad = None w2_master.grad = None + w3_master.grad = None for _ in range(args.warmup): forward_backward() diff --git a/tests/test_fp8_block128_mega_moe.py b/tests/test_fp8_block128_mega_moe.py index 229a89719f..a3b00093e8 100644 --- a/tests/test_fp8_block128_mega_moe.py +++ b/tests/test_fp8_block128_mega_moe.py @@ -47,15 +47,18 @@ def _make_case(tokens: int, *, all_to_one: bool = False) -> dict[str, torch.Tens 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_() - canonical_w13 = ( - torch.randn(experts * 2, hidden, model_dim, device="cuda", dtype=torch.bfloat16) * 0.05 + 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_q, canonical_w13_s = _blockwise_weight_quantize(canonical_w13.detach()) - w13_q, w13_s = deep_gemm.transform_glm_w13_for_fp8_block128_mega_moe( - canonical_w13_q, canonical_w13_s + 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: @@ -73,15 +76,23 @@ def _make_case(tokens: int, *, all_to_one: bool = False) -> dict[str, torch.Tens "w13_s": w13_s, "w2_q": w2_q, "w2_s": w2_s, - "w13_master": canonical_w13, + "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() - w13 = _blockwise_weight_dequantize(case["w13_q"], case["w13_s"]) + 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) @@ -104,22 +115,26 @@ def _ste_reference_backward( ) -> dict[str, torch.Tensor]: x = case["x"].detach().clone().requires_grad_() scores = case["scores"].detach().clone().requires_grad_() - canonical_w13 = case["w13_master"].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_dequantized = _blockwise_weight_dequantize( - case["w13_q"], case["w13_s"] + w13_active_master = torch.stack((w3_master, w1_master), dim=1).reshape( + experts, hidden * 2, model_dim ) - canonical_pairs = canonical_w13.view(experts, 2, hidden, model_dim) - w13_active_master = torch.stack( - (canonical_pairs[:, 1], canonical_pairs[:, 0]), 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() + ( - w13_dequantized - 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() + ( @@ -154,8 +169,9 @@ def _ste_reference_backward( "output": output.detach(), "x_grad": x.grad.detach(), "score_grad": scores.grad.detach(), - "w13_grad": canonical_w13.grad.detach(), + "w1_grad": w1_master.grad.detach(), "w2_grad": w2_master.grad.detach(), + "w3_grad": w3_master.grad.detach(), } @@ -186,8 +202,9 @@ def test_single_rank_forward_matches_reference_across_padding_boundaries( case["w13_s"], case["w2_q"], case["w2_s"], - case["w13_master"], + 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) @@ -223,8 +240,9 @@ def test_single_rank_backward_returns_input_score_and_canonical_master_grads() - case["w13_s"], case["w2_q"], case["w2_s"], - case["w13_master"], + case["w1_master"], case["w2_master"], + case["w3_master"], ) output.backward(upstream) @@ -236,10 +254,13 @@ def test_single_rank_backward_returns_input_score_and_canonical_master_grads() - ) assert case["x"].grad.shape == case["x"].shape assert case["x"].grad.dtype == torch.bfloat16 - assert case["w13_master"].grad.shape == case["w13_master"].shape - assert case["w13_master"].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["w13_master"].grad, expected["w13_grad"]) < 0.15 + 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 index 9870aeef5c..7165872353 100644 --- a/tests/test_fp8_block128_mega_moe_distributed.py +++ b/tests/test_fp8_block128_mega_moe_distributed.py @@ -168,7 +168,7 @@ def _worker(rank: int, world_size: int, port: int) -> None: full_w13_q_canonical, full_w13_s_canonical = _weight_quantize( full_w13_master ) - full_w13_q, full_w13_s = ( + 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 ) @@ -176,12 +176,22 @@ def _worker(rank: int, world_size: int, port: int) -> None: 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[expert_start:expert_end].contiguous() - local_w13_s = full_w13_s[expert_start:expert_end].contiguous() + 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_w13_master = ( - full_w13_master[expert_start * 2 : expert_end * 2] + 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_() @@ -198,8 +208,8 @@ def _worker(rank: int, world_size: int, port: int) -> None: x, ids, scores, - full_w13_q, - full_w13_s, + full_w13_q_active, + full_w13_s_active, full_w2_q, full_w2_s, upstream, @@ -213,8 +223,9 @@ def _worker(rank: int, world_size: int, port: int) -> None: local_w13_s, local_w2_q, local_w2_s, - local_w13_master, + local_w1_master, local_w2_master, + local_w3_master, group=dist.group.WORLD, ) output.backward(upstream) @@ -225,13 +236,19 @@ def _worker(rank: int, world_size: int, port: int) -> None: scores.grad, expected_score_grad, rtol=3e-4, atol=3e-3 ) assert _normalized_difference(x.grad, expected_x_grad) < 0.12 - for gradient in (x.grad, local_w13_master.grad, local_w2_master.grad): + 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_w13_master.grad[2:]) == 0 + 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 dist.barrier() finally: diff --git a/tests/test_sm103_fp8_block128_primitives.py b/tests/test_sm103_fp8_block128_primitives.py index 1812de7da4..49189344c7 100644 --- a/tests/test_sm103_fp8_block128_primitives.py +++ b/tests/test_sm103_fp8_block128_primitives.py @@ -130,6 +130,17 @@ def test_swiglu_forward_and_backward_match_reference(rows: int) -> None: 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: @@ -243,6 +254,44 @@ def test_grouped_fp8_block128_gemm_covers_zero_routes_and_rejects_unpadded_group ) +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_k_grouped_bf16_wgrad_supports_no_accumulator_and_empty_expert() -> None: torch.manual_seed(6000) counts = [128, 0, 128] From a1c4baf8337993faf9d0114375e7733ebe3a130d Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 21 Jul 2026 20:54:29 +0800 Subject: [PATCH 03/29] Support sharded MegaMoE master anchors --- deep_gemm/mega/fp8_block128.py | 82 +++++++++++++++++++++++------ tests/test_fp8_block128_mega_moe.py | 43 +++++++++++++++ 2 files changed, 108 insertions(+), 17 deletions(-) diff --git a/deep_gemm/mega/fp8_block128.py b/deep_gemm/mega/fp8_block128.py index e4472ff577..f438303b17 100644 --- a/deep_gemm/mega/fp8_block128.py +++ b/deep_gemm/mega/fp8_block128.py @@ -153,6 +153,39 @@ def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: 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: + """Validate either a resident EP-local master or an at-rest DTensor shard. + + The masters are autograd anchors only; MegaMoE never reads their values. + FireTitan may therefore keep them eFSDP-sharded while the active q/s + tensors are all-gathered. In that case the framework-provided gradient + wrapper maps the full EP-local wgrad back to the master's DTensor layout. + """ + 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}, 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, @@ -165,6 +198,7 @@ def _validate_inputs( w2_master: torch.Tensor, w3_master: torch.Tensor, group_state: _GroupState, + master_gradient_wrapper: Any, ) -> tuple[int, int, int, int, int]: if not x.is_cuda: raise ValueError("FP8-block128 MegaMoE requires CUDA") @@ -182,13 +216,6 @@ def _validate_inputs( _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) - w1_master_local = _local_tensor(w1_master) - w2_master_local = _local_tensor(w2_master) - w3_master_local = _local_tensor(w3_master) - _check_tensor(w1_master_local, name="w1_master", ndim=3, dtype=torch.bfloat16, device=device) - _check_tensor(w2_master_local, name="w2_master", ndim=3, dtype=torch.bfloat16, device=device) - _check_tensor(w3_master_local, name="w3_master", ndim=3, dtype=torch.bfloat16, device=device) - tokens, model_dim = x.shape if model_dim % _BLOCK: raise ValueError("model dimension must be divisible by 128") @@ -218,15 +245,35 @@ def _validate_inputs( hidden // _BLOCK, ): raise ValueError("W2 scale shape does not match 128x128 weight blocks") - master_w13_shape = (local_experts, hidden, model_dim) - if tuple(w1_master_local.shape) != master_w13_shape: - raise ValueError("BF16 gate master must have shape [E, H, D]") - if tuple(w3_master_local.shape) != master_w13_shape: - raise ValueError("BF16 up master must have shape [E, H, D]") - if tuple(w2_master_local.shape) != tuple(w2_weight.shape): - raise ValueError("BF16 W2 master shape must match W2") - global_experts = local_experts * group_state.world_size + local_w13_shape = (local_experts, hidden, model_dim) + global_w13_shape = (global_experts, hidden, model_dim) + local_w2_shape = tuple(w2_weight.shape) + global_w2_shape = (global_experts, model_dim, hidden) + _validate_master_tensor( + w1_master, + name="w1_master", + local_shape=local_w13_shape, + global_shape=global_w13_shape, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) + _validate_master_tensor( + w2_master, + name="w2_master", + local_shape=local_w2_shape, + global_shape=global_w2_shape, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) + _validate_master_tensor( + w3_master, + name="w3_master", + local_shape=local_w13_shape, + global_shape=global_w13_shape, + device=device, + master_gradient_wrapper=master_gradient_wrapper, + ) if topk_ids.numel(): minimum, maximum = torch.aminmax(topk_ids) if minimum.item() < 0 or maximum.item() >= global_experts: @@ -374,6 +421,8 @@ def forward( master_gradient_wrapper: Any, ) -> torch.Tensor: group_state = _resolve_group(group) + if master_gradient_wrapper is not None and not callable(master_gradient_wrapper): + raise TypeError("master_gradient_wrapper must be callable or None") tokens, model_dim, hidden, local_experts, topk = _validate_inputs( x, topk_ids, @@ -386,9 +435,8 @@ def forward( w2_master, w3_master, group_state, + master_gradient_wrapper, ) - if master_gradient_wrapper is not None and not callable(master_gradient_wrapper): - raise TypeError("master_gradient_wrapper must be callable or None") with torch.autograd.profiler.record_function( "sm103_fp8_block128_megamoe_forward" diff --git a/tests/test_fp8_block128_mega_moe.py b/tests/test_fp8_block128_mega_moe.py index a3b00093e8..af76c98224 100644 --- a/tests/test_fp8_block128_mega_moe.py +++ b/tests/test_fp8_block128_mega_moe.py @@ -2,6 +2,7 @@ import torch import deep_gemm +from deep_gemm.mega.fp8_block128 import _validate_master_tensor pytestmark = pytest.mark.skipif( @@ -186,6 +187,48 @@ def _normalized_difference(actual: torch.Tensor, expected: torch.Tensor) -> floa ) +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, + ) + + @pytest.mark.parametrize( ("tokens", "all_to_one"), [(1, True), (63, False), (64, True), (65, False)], From 7a099751d8f648a5b7dd4ca7eded83cb4ec45e24 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 01:21:17 +0800 Subject: [PATCH 04/29] Preserve expanded DeepEP scale layout --- csrc/jit_kernels/impls/smxx_layout.hpp | 19 +++-- .../include/deep_gemm/impls/smxx_layout.cuh | 39 +++++++--- tests/test_sm103_fp8_block128_primitives.py | 73 +++++++++++++++++++ 3 files changed, 113 insertions(+), 18 deletions(-) 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/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/tests/test_sm103_fp8_block128_primitives.py b/tests/test_sm103_fp8_block128_primitives.py index 49189344c7..ebe959e4c9 100644 --- a/tests/test_sm103_fp8_block128_primitives.py +++ b/tests/test_sm103_fp8_block128_primitives.py @@ -83,6 +83,36 @@ def _blockwise_weight_dequantize( ) +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() @@ -292,6 +322,49 @@ def test_canonical_w13_gemm_presents_up_then_gate_without_weight_copy() -> None: 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] From 1f8b00146b7fcd962a5fc36b3a254b7affd9ab44 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 05:35:05 +0800 Subject: [PATCH 05/29] Add device-driven SM103 expanded MegaMoE path --- README.md | 11 + csrc/sm103_fp8_block128.cu | 845 ++++++++++++++++-- deep_gemm/mega/fp8_block128.py | 687 ++++++++++++-- tests/benchmark_fp8_block128_mega_moe.py | 76 +- tests/test_fp8_block128_capabilities.py | 6 + tests/test_fp8_block128_mega_moe.py | 123 ++- .../test_fp8_block128_mega_moe_distributed.py | 370 +++++++- 7 files changed, 1868 insertions(+), 250 deletions(-) diff --git a/README.md b/README.md index 2b4e49e01a..0d65932888 100644 --- a/README.md +++ b/README.md @@ -152,6 +152,17 @@ 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 diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index c15c3070fd..2a81959752 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -1,6 +1,7 @@ #include "apis/sm103_fp8_block128.hpp" #include +#include #include #include #include @@ -20,6 +21,8 @@ #include #include +#include +#include #include #include #include @@ -89,6 +92,10 @@ constexpr int kRequiredMinor = 3; constexpr int kBlockK = 128; constexpr float kE4M3Max = 448.0f; +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) \ @@ -428,6 +435,250 @@ __global__ void sm103_post_down_score_grad_kernel( #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>; @@ -507,20 +758,156 @@ struct SM103GroupedBlockwiseGemm { using StrideD = typename GemmKernel::InternalStrideD; }; -template -torch::Tensor copy_metadata_to_device( - const std::vector& host, - const torch::TensorOptions& options, - cudaStream_t stream +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 ) { - const int64_t num_bytes = static_cast(host.size() * sizeof(T)); - auto storage = torch::empty({std::max(num_bytes, 1)}, options.dtype(torch::kUInt8)); - if (num_bytes != 0) { - C10_CUDA_CHECK(cudaMemcpyAsync( - storage.data_ptr(), host.data(), num_bytes, cudaMemcpyHostToDevice, stream - )); - } - return storage; + 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 @@ -529,7 +916,8 @@ torch::Tensor grouped_fp8_block128_gemm_impl( const torch::Tensor& activation_scales, const torch::Tensor& weights, const torch::Tensor& weight_scales, - const std::vector& group_counts + const std::vector& group_counts, + const bool expanded_layout ) { using Config = SM103GroupedBlockwiseGemm; using Problem = typename Config::ProblemShape::UnderlyingProblemShape; @@ -566,18 +954,19 @@ torch::Tensor grouped_fp8_block128_gemm_impl( "weight scales must have shape [G, N/128, K/128]"); } - int64_t total_rows = 0; + 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(count == 0 || count % 4 == 0, - "active group counts must be padded to a multiple of four"); - total_rows += count; + 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(total_rows == activations.size(0), "sum(group_counts) must equal activation rows"); + TORCH_CHECK(storage_rows == activations.size(0), + "group counts do not match the activation storage rows"); - auto output = torch::empty({total_rows, n}, activations.options().dtype(torch::kBFloat16)); + auto output = torch::empty({storage_rows, n}, activations.options().dtype(torch::kBFloat16)); if (active_groups == 0) { return output; } @@ -644,23 +1033,26 @@ torch::Tensor grouped_fp8_block128_gemm_impl( 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 += m; + row_offset += expanded_layout ? align_rows(m) : m; } const auto metadata_options = activations.options().dtype(torch::kUInt8); - auto problems_device = copy_metadata_to_device(problems, metadata_options, stream); - auto ptr_a_device = copy_metadata_to_device(ptr_a, metadata_options, stream); - auto ptr_b_device = copy_metadata_to_device(ptr_b, metadata_options, stream); - auto ptr_c_device = copy_metadata_to_device(ptr_c, metadata_options, stream); - auto ptr_d_device = copy_metadata_to_device(ptr_d, metadata_options, stream); - auto ptr_sfa_device = copy_metadata_to_device(ptr_sfa, metadata_options, stream); - auto ptr_sfb_device = copy_metadata_to_device(ptr_sfb, metadata_options, stream); - auto stride_a_device = copy_metadata_to_device(stride_a, metadata_options, stream); - auto stride_b_device = copy_metadata_to_device(stride_b, metadata_options, stream); - auto stride_c_device = copy_metadata_to_device(stride_c, metadata_options, stream); - auto stride_d_device = copy_metadata_to_device(stride_d, metadata_options, stream); - auto layout_sfa_device = copy_metadata_to_device(layout_sfa, metadata_options, stream); - auto layout_sfb_device = copy_metadata_to_device(layout_sfb, metadata_options, stream); + 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(); @@ -670,21 +1062,21 @@ torch::Tensor grouped_fp8_block128_gemm_impl( typename Config::Gemm::Arguments arguments{ cutlass::gemm::GemmUniversalMode::kGrouped, {static_cast(active_groups), - reinterpret_cast(problems_device.data_ptr()), + reinterpret_cast(metadata_base + problems_offset), problems.data()}, - {reinterpret_cast(ptr_a_device.data_ptr()), - reinterpret_cast(stride_a_device.data_ptr()), - reinterpret_cast(ptr_b_device.data_ptr()), - reinterpret_cast(stride_b_device.data_ptr()), - reinterpret_cast(ptr_sfa_device.data_ptr()), - reinterpret_cast(layout_sfa_device.data_ptr()), - reinterpret_cast(ptr_sfb_device.data_ptr()), - reinterpret_cast(layout_sfb_device.data_ptr())}, + {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(ptr_c_device.data_ptr()), - reinterpret_cast(stride_c_device.data_ptr()), - reinterpret_cast(ptr_d_device.data_ptr()), - reinterpret_cast(stride_d_device.data_ptr())}, + 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; @@ -715,7 +1107,8 @@ torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_impl( const torch::Tensor& activation_scales, const torch::Tensor& canonical_weights, const torch::Tensor& canonical_weight_scales, - const std::vector& group_counts + 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 @@ -760,20 +1153,20 @@ torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_impl( torch::IntArrayRef({groups * 2, hidden / kBlockK, k / kBlockK}), "canonical W13 scales must have shape [2E, H/128, D/128]"); - int64_t total_rows = 0; + 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(count == 0 || count % 4 == 0, - "active group counts must be padded to a multiple of four"); - total_rows += count; + 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(total_rows == activations.size(0), - "sum(group_counts) must equal activation rows"); + TORCH_CHECK(storage_rows == activations.size(0), + "group counts do not match the activation storage rows"); auto output = torch::empty( - {total_rows, doubled_hidden}, activations.options().dtype(torch::kBFloat16)); + {storage_rows, doubled_hidden}, activations.options().dtype(torch::kBFloat16)); if (active_groups == 0) { return output; } @@ -855,23 +1248,26 @@ torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_impl( 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 += m; + row_offset += expanded_layout ? align_rows(m) : m; } const auto metadata_options = activations.options().dtype(torch::kUInt8); - auto problems_device = copy_metadata_to_device(problems, metadata_options, stream); - auto ptr_a_device = copy_metadata_to_device(ptr_a, metadata_options, stream); - auto ptr_b_device = copy_metadata_to_device(ptr_b, metadata_options, stream); - auto ptr_c_device = copy_metadata_to_device(ptr_c, metadata_options, stream); - auto ptr_d_device = copy_metadata_to_device(ptr_d, metadata_options, stream); - auto ptr_sfa_device = copy_metadata_to_device(ptr_sfa, metadata_options, stream); - auto ptr_sfb_device = copy_metadata_to_device(ptr_sfb, metadata_options, stream); - auto stride_a_device = copy_metadata_to_device(stride_a, metadata_options, stream); - auto stride_b_device = copy_metadata_to_device(stride_b, metadata_options, stream); - auto stride_c_device = copy_metadata_to_device(stride_c, metadata_options, stream); - auto stride_d_device = copy_metadata_to_device(stride_d, metadata_options, stream); - auto layout_sfa_device = copy_metadata_to_device(layout_sfa, metadata_options, stream); - auto layout_sfb_device = copy_metadata_to_device(layout_sfb, metadata_options, stream); + 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(); @@ -881,21 +1277,21 @@ torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_impl( typename Config::Gemm::Arguments arguments{ cutlass::gemm::GemmUniversalMode::kGrouped, {static_cast(active_problems), - reinterpret_cast(problems_device.data_ptr()), + reinterpret_cast(metadata_base + problems_offset), problems.data()}, - {reinterpret_cast(ptr_a_device.data_ptr()), - reinterpret_cast(stride_a_device.data_ptr()), - reinterpret_cast(ptr_b_device.data_ptr()), - reinterpret_cast(stride_b_device.data_ptr()), - reinterpret_cast(ptr_sfa_device.data_ptr()), - reinterpret_cast(layout_sfa_device.data_ptr()), - reinterpret_cast(ptr_sfb_device.data_ptr()), - reinterpret_cast(layout_sfb_device.data_ptr())}, + {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(ptr_c_device.data_ptr()), - reinterpret_cast(stride_c_device.data_ptr()), - reinterpret_cast(ptr_d_device.data_ptr()), - reinterpret_cast(stride_d_device.data_ptr())}, + 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; @@ -928,7 +1324,7 @@ torch::Tensor grouped_fp8_block128_gemm_nt( const std::vector& group_counts ) { return grouped_fp8_block128_gemm_impl( - activations, activation_scales, weights, weight_scales, group_counts); + activations, activation_scales, weights, weight_scales, group_counts, false); } torch::Tensor grouped_fp8_block128_gemm_nn( @@ -939,7 +1335,61 @@ torch::Tensor grouped_fp8_block128_gemm_nn( const std::vector& group_counts ) { return grouped_fp8_block128_gemm_impl( - activations, activation_scales, weights, weight_scales, group_counts); + 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); } std::tuple quantize_bf16(const torch::Tensor& input) { @@ -1180,6 +1630,185 @@ torch::Tensor post_down_score_grad( 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; @@ -1200,13 +1829,23 @@ pybind11::dict capabilities() { "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_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; } @@ -1227,7 +1866,22 @@ void register_apis(pybind11::module_& m) { 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_impl, + &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")); @@ -1245,6 +1899,25 @@ void register_apis(pybind11::module_& m) { 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/mega/fp8_block128.py b/deep_gemm/mega/fp8_block128.py index f438303b17..0dc42edf3a 100644 --- a/deep_gemm/mega/fp8_block128.py +++ b/deep_gemm/mega/fp8_block128.py @@ -8,8 +8,8 @@ from __future__ import annotations +from copy import copy from dataclasses import dataclass -from itertools import accumulate from typing import Any, Sequence import torch @@ -17,7 +17,6 @@ from .. import _C - _BLOCK = 128 _PAD_ROWS = 128 @@ -28,6 +27,9 @@ "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", @@ -35,7 +37,13 @@ "sm103_fp8_block128_post_down_combine", "sm103_fp8_block128_post_down_score_grad", "sm103_fp8_block128_route_sum", - "k_grouped_bf16_gemm_tn_contiguous", + "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", ) REQUIRED_PYTHON_SYMBOLS = ( @@ -49,13 +57,26 @@ def get_fp8_block128_mega_moe_capabilities() -> 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: + from deep_ep import ElasticBuffer + + has_elastic_buffer = callable(ElasticBuffer) + except (AttributeError, ImportError): + has_elastic_buffer = False + if not has_elastic_buffer: + missing.append("deep_ep.ElasticBuffer") native.update( { "native_symbols": REQUIRED_NATIVE_SYMBOLS, "python_symbols": REQUIRED_PYTHON_SYMBOLS, "forward": not missing, "backward": not missing, - "distributed_transport": "torch.distributed.all_to_all_single", + "distributed_transport": "deep_ep.ElasticBuffer.expanded", + "transport_layout": "expert_aligned_128", + "transport_scale_layout": "row_major_fp32_group128", + "transport_deterministic": True, + "combine_reductions": 1, + "wgrad_backend": "sm103_companion_grouped_bf16", "missing_symbols": tuple(missing), } ) @@ -95,9 +116,7 @@ def transform_glm_w13_for_fp8_block128_mega_moe( .contiguous() ) active_scale = ( - canonical_scale.view( - experts, 2, hidden // _BLOCK, model_dim // _BLOCK - ) + canonical_scale.view(experts, 2, hidden // _BLOCK, model_dim // _BLOCK) .index_select(1, pair_order) .reshape(experts, hidden * 2 // _BLOCK, model_dim // _BLOCK) .contiguous() @@ -115,12 +134,15 @@ class _GroupState: def _resolve_group(group: Any) -> _GroupState: if not dist.is_available() or not dist.is_initialized(): if group is not None: - raise RuntimeError("a process group was provided before torch.distributed initialization") + raise RuntimeError( + "a process group was provided before torch.distributed initialization" + ) return _GroupState(group=None, rank=0, world_size=1) + resolved_group = dist.group.WORLD if group is None else group return _GroupState( - group=group, - rank=dist.get_rank(group), - world_size=dist.get_world_size(group), + group=resolved_group, + rank=dist.get_rank(resolved_group), + world_size=dist.get_world_size(resolved_group), ) @@ -211,10 +233,18 @@ def _validate_inputs( 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( + 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, model_dim = x.shape if model_dim % _BLOCK: @@ -230,13 +260,17 @@ def _validate_inputs( local_experts = w13_weight.shape[0] // 2 hidden, w13_k = w13_weight.shape[1:] if local_experts <= 0 or hidden % _BLOCK or w13_k != model_dim: - raise ValueError("canonical W13 must have shape [2E, H, D] with D/H divisible by 128") + raise ValueError( + "canonical W13 must have shape [2E, H, D] with D/H divisible by 128" + ) if tuple(w13_scale.shape) != ( local_experts * 2, hidden // _BLOCK, model_dim // _BLOCK, ): - raise ValueError("canonical W13 scale shape does not match 128x128 weight blocks") + raise ValueError( + "canonical W13 scale shape does not match 128x128 weight blocks" + ) if tuple(w2_weight.shape) != (local_experts, model_dim, hidden): raise ValueError("W2 must have shape [local_experts, D, H]") if tuple(w2_scale.shape) != ( @@ -275,14 +309,192 @@ def _validate_inputs( master_gradient_wrapper=master_gradient_wrapper, ) if topk_ids.numel(): - minimum, maximum = torch.aminmax(topk_ids) - if minimum.item() < 0 or maximum.item() >= global_experts: - raise ValueError( - f"top-k IDs must lie in [0, {global_experts}); got [{minimum.item()}, {maximum.item()}]" - ) + # Keep the hot path free of a device-to-host scalar synchronization. + # The assertion is enqueued on the current CUDA stream and fails the + # operation rather than admitting an out-of-range route. + torch._assert_async( + ((topk_ids >= 0) & (topk_ids < global_experts)).all(), + f"top-k IDs must lie in [0, {global_experts})", + ) return tokens, model_dim, hidden, local_experts, topk +@dataclass(frozen=True) +class _DeepEPBufferState: + buffer: Any + capacity: int + num_sms: int + num_qps: int + + +_deepep_buffers: dict[tuple[int, int, int, int, int], _DeepEPBufferState] = {} +_deepep_context_tokens_per_rank: dict[int, int] = {} + + +def _configure_fp8_block128_mega_moe_transport( + group: Any, + *, + context_tokens_per_rank: int, +) -> None: + """Register the owning model's existing context/CP envelope once.""" + 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) + if group_state.world_size <= 1: + raise RuntimeError("expanded DeepEP transport requires a multi-rank group") + key = id(group_state.group) + existing = _deepep_context_tokens_per_rank.get(key) + if existing is not None and existing != context_tokens_per_rank: + raise RuntimeError( + "MegaMoE context envelope changed after transport configuration: " + f"configured={context_tokens_per_rank}, existing={existing}" + ) + if existing is None and any(buffer_key[0] == key for buffer_key in _deepep_buffers): + raise RuntimeError("MegaMoE transport cannot be configured after arena construction") + if existing is None: + # DeepEP reads the native NCCL communicator while calculating its arena + # size. Materialize that communicator once during model setup; otherwise + # a first-use EP group can expose an uninitialized handle to + # ncclTeamWorld. This is not a per-layer sizing collective. + if dist.get_backend(group_state.group) == "nccl": + dist.barrier( + group=group_state.group, + device_ids=[torch.cuda.current_device()], + ) + _deepep_context_tokens_per_rank[key] = context_tokens_per_rank + + +def _get_deepep_buffer( + group_state: _GroupState, + *, + device: torch.device, + tokens: int, + model_dim: int, + topk: int, + global_experts: int, +) -> _DeepEPBufferState: + """Create one runtime-context-sized DeepEP arena per EP/model shape. + + FireTitan registers its already-resolved context/CP envelope when it + installs the EP group. Warmup shape therefore never affects sizing, while + changing the existing context length or CP configuration only requires a + normal trainer restart, not an image rebuild. No hot-path sizing collective + or caller-visible operation argument exists. + """ + if group_state.world_size <= 1: + raise RuntimeError("expanded DeepEP requires a multi-rank process group") + capacity = _deepep_context_tokens_per_rank.get(id(group_state.group)) + if capacity is None: + raise RuntimeError( + "MegaMoE transport was not configured from the owning model context" + ) + if tokens > capacity: + raise RuntimeError( + "SM103 MegaMoE input exceeds the resolved context/CP envelope: " + f"actual={tokens}, capacity={capacity}" + ) + + device_index = ( + device.index if device.index is not None else torch.cuda.current_device() + ) + key = (id(group_state.group), device_index, model_dim, topk, global_experts) + cached = _deepep_buffers.get(key) + if cached is not None: + return cached + + try: + from deep_ep import ElasticBuffer + except (AttributeError, ImportError) as exc: + raise RuntimeError( + "MegaMoE requires deep_ep.ElasticBuffer; no transport fallback exists" + ) from exc + + buffer_kwargs = dict( + num_max_tokens_per_rank=capacity, + hidden=model_dim, + num_topk=topk, + allow_hybrid_mode=True, + allow_multiple_reduction=False, + ) + fp8_bytes = ElasticBuffer.get_buffer_size_hint( + group_state.group, use_fp8_dispatch=True, **buffer_kwargs + ) + bf16_bytes = ElasticBuffer.get_buffer_size_hint( + group_state.group, use_fp8_dispatch=False, **buffer_kwargs + ) + buffer = ElasticBuffer( + group_state.group, + num_bytes=max(fp8_bytes, bf16_bytes), + use_fp8_dispatch=True, + deterministic=True, + prefer_overlap_with_compute=True, + **buffer_kwargs, + ) + num_sms = int(buffer.get_theoretical_num_sms(global_experts, topk)) + device_sms = torch.cuda.get_device_properties(device).multi_processor_count + if not 1 <= num_sms <= device_sms: + raise RuntimeError( + "DeepEP returned an invalid SM103 launch width: " + f"selected={num_sms}, device_sms={device_sms}" + ) + num_qps = int(buffer.get_theoretical_num_qps(num_sms)) + if num_qps < 1: + raise RuntimeError(f"DeepEP returned an invalid QP count: {num_qps}") + state = _DeepEPBufferState( + buffer=buffer, + capacity=capacity, + num_sms=num_sms, + num_qps=num_qps, + ) + _deepep_buffers[key] = state + return state + + +def _host_expanded_storage_counts(handle: Any, local_experts: int) -> tuple[int, ...]: + counts = getattr(handle, "num_recv_tokens_per_expert_list", None) + if counts is None: + raise RuntimeError( + "DeepEP expanded dispatch did not return host storage counts" + ) + if isinstance(counts, torch.Tensor): + if counts.device.type != "cpu": + raise RuntimeError("DeepEP host expert counts unexpectedly reside on CUDA") + counts = counts.tolist() + result = tuple(int(value) for value in counts) + if ( + len(result) != local_experts + or any(value < 0 for value in result) + or any(value and value % _PAD_ROWS for value in result) + ): + raise RuntimeError( + "DeepEP expanded storage counts violate the aligned local expert contract: " + f"expected={local_experts}, actual={result}" + ) + return result + + +def _expanded_routes(handle: Any) -> torch.Tensor: + metadata = getattr(handle, "recv_src_metadata", None) + if ( + not isinstance(metadata, torch.Tensor) + or metadata.ndim != 2 + or metadata.shape[1] < 3 + ): + shape = tuple(metadata.shape) if isinstance(metadata, torch.Tensor) else None + raise RuntimeError(f"invalid DeepEP expanded route metadata: {shape}") + return metadata[:, 2:] + + +def _shadow_compact_handle(handle: Any) -> Any: + shadow = copy(handle) + shadow.do_expand = False + return shadow + + def _exchange_counts( send_counts: Sequence[int], group_state: _GroupState, device: torch.device ) -> list[int]: @@ -308,7 +520,9 @@ def _all_to_all_rows( device=tensor.device, ) source = tensor.view(torch.uint8) if tensor.dtype == torch.float8_e4m3fn else tensor - destination = output.view(torch.uint8) if output.dtype == torch.float8_e4m3fn else output + destination = ( + output.view(torch.uint8) if output.dtype == torch.float8_e4m3fn else output + ) dist.all_to_all_single( destination, source, @@ -328,7 +542,10 @@ def _inverse_permutation(order: torch.Tensor) -> torch.Tensor: def _padding_state( counts: Sequence[int], device: torch.device ) -> tuple[list[int], torch.Tensor]: - padded_counts = [((count + _PAD_ROWS - 1) // _PAD_ROWS) * _PAD_ROWS if count else 0 for count in counts] + padded_counts = [ + ((count + _PAD_ROWS - 1) // _PAD_ROWS) * _PAD_ROWS if count else 0 + for count in counts + ] total_actual = sum(counts) if total_actual == 0: return padded_counts, torch.empty(0, dtype=torch.int64, device=device) @@ -380,30 +597,26 @@ def _bf16_grouped_wgrad( right: torch.Tensor, padded_counts: Sequence[int], ) -> torch.Tensor: - output = torch.zeros( - (len(padded_counts), left.shape[1], right.shape[1]), - dtype=torch.bfloat16, - device=left.device, - ) - if left.shape[0] == 0: - return output - grouped_layout = torch.tensor( - list(accumulate(padded_counts)), dtype=torch.int32, device=left.device + return _C.sm103_fp8_block128_grouped_bf16_wgrad( + left.contiguous(), + right.contiguous(), + list(padded_counts), ) - _C.k_grouped_bf16_gemm_tn_contiguous( + + +def _bf16_grouped_wgrad_expanded( + left: torch.Tensor, + right: torch.Tensor, + psum: torch.Tensor, +) -> torch.Tensor: + return _C.sm103_fp8_block128_grouped_bf16_wgrad_expanded( left.contiguous(), right.contiguous(), - output, - None, - grouped_layout, - None, - "mn", - True, + psum, ) - return output -class _FP8Block128MegaMoE(torch.autograd.Function): +class _FP8Block128MegaMoELocal(torch.autograd.Function): @staticmethod def forward( ctx: Any, @@ -421,7 +634,11 @@ def forward( master_gradient_wrapper: Any, ) -> torch.Tensor: group_state = _resolve_group(group) - if master_gradient_wrapper is not None and not callable(master_gradient_wrapper): + if group_state.world_size != 1: + raise RuntimeError("the packed local MegaMoE path is single-rank only") + if master_gradient_wrapper is not None and not callable( + master_gradient_wrapper + ): raise TypeError("master_gradient_wrapper must be callable or None") tokens, model_dim, hidden, local_experts, topk = _validate_inputs( x, @@ -479,15 +696,11 @@ def forward( grouped_local_ids = local_ids.index_select(0, group_order) actual_counts = [ int(value) - for value in torch.bincount( - grouped_local_ids, minlength=local_experts - ) + for value in torch.bincount(grouped_local_ids, minlength=local_experts) .cpu() .tolist() ] - padded_counts, actual_to_padded = _padding_state( - actual_counts, x.device - ) + padded_counts, actual_to_padded = _padding_state(actual_counts, x.device) padded_rows = sum(padded_counts) padded_activations = _pad_rows( grouped_activations, @@ -509,8 +722,8 @@ def forward( w13_scale, padded_counts, ) - hidden_quantized, hidden_scales = ( - _C.sm103_fp8_block128_swiglu_quantize(preactivation) + hidden_quantized, hidden_scales = _C.sm103_fp8_block128_swiglu_quantize( + preactivation ) routed_output_padded = _C.sm103_fp8_block128_grouped_gemm_nt( hidden_quantized, @@ -519,9 +732,7 @@ def forward( w2_scale, padded_counts, ) - routed_output_grouped = _unpad_rows( - routed_output_padded, actual_to_padded - ) + routed_output_grouped = _unpad_rows(routed_output_padded, actual_to_padded) routed_output_receive_order = routed_output_grouped.index_select( 0, ungroup_order ) @@ -537,12 +748,8 @@ def forward( device=x.device, ) if routed_output.shape[0]: - routed_output.index_copy_( - 0, send_order, routed_output_send_order - ) - output = _C.sm103_fp8_block128_post_down_combine( - routed_output, topk_scores - ) + routed_output.index_copy_(0, send_order, routed_output_send_order) + output = _C.sm103_fp8_block128_post_down_combine(routed_output, topk_scores) ctx.group_state = group_state ctx.send_counts = send_counts @@ -655,9 +862,7 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: grad_input_padded = _C.sm103_fp8_block128_grouped_gemm_nn( grad_preactivation_quantized, grad_preactivation_scales, - w13_weight.view( - ctx.local_experts, ctx.hidden * 2, ctx.model_dim - ), + w13_weight.view(ctx.local_experts, ctx.hidden * 2, ctx.model_dim), w13_scale.view( ctx.local_experts, ctx.hidden * 2 // _BLOCK, @@ -688,12 +893,8 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: grad_w1 = grad_w13_canonical[:, : ctx.hidden].contiguous() grad_w3 = grad_w13_canonical[:, ctx.hidden :].contiguous() - grad_input_grouped = _unpad_rows( - grad_input_padded, actual_to_padded - ) - grad_input_receive_order = grad_input_grouped.index_select( - 0, ungroup_order - ) + grad_input_grouped = _unpad_rows(grad_input_padded, actual_to_padded) + grad_input_receive_order = grad_input_grouped.index_select(0, ungroup_order) grad_input_send_order = _all_to_all_rows( grad_input_receive_order, ctx.receive_counts, @@ -706,9 +907,7 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: device=grad_output.device, ) if grad_input_routes.shape[0]: - grad_input_routes.index_copy_( - 0, send_order, grad_input_send_order - ) + grad_input_routes.index_copy_(0, send_order, grad_input_send_order) grad_input = _C.sm103_fp8_block128_route_sum( grad_input_routes, ctx.tokens, ctx.topk ) @@ -733,6 +932,336 @@ def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: ) +class _FP8Block128MegaMoEDeepEP(torch.autograd.Function): + """Invocation-owned SM103 compute over DeepEP's expanded route layout.""" + + @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) + if group_state.world_size <= 1: + raise RuntimeError( + "the expanded DeepEP MegaMoE path requires multiple ranks" + ) + if master_gradient_wrapper is not None and not callable( + master_gradient_wrapper + ): + raise TypeError("master_gradient_wrapper must be callable or None") + tokens, model_dim, hidden, local_experts, topk = _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, + ) + global_experts = local_experts * group_state.world_size + buffer_state = _get_deepep_buffer( + group_state, + device=x.device, + tokens=tokens, + model_dim=model_dim, + topk=topk, + global_experts=global_experts, + ) + + with torch.autograd.profiler.record_function( + "sm103_fp8_block128_megamoe_forward" + ): + token_quantized, token_scales = _C.sm103_fp8_block128_quantize(x) + payload, recv_ids, routed_scores, handle, event = ( + buffer_state.buffer.dispatch( + (token_quantized, token_scales), + topk_idx=topk_ids, + topk_weights=topk_scores, + num_experts=global_experts, + num_max_tokens_per_rank=buffer_state.capacity, + expert_alignment=_PAD_ROWS, + num_sms=buffer_state.num_sms, + num_qps=buffer_state.num_qps, + async_with_compute_stream=True, + allocate_on_comm_stream=False, + do_cpu_sync=True, + do_expand=True, + use_tma_aligned_col_major_sf=False, + ) + ) + event.current_stream_wait() + if recv_ids is not None or routed_scores is None: + raise RuntimeError( + "DeepEP did not return an expanded scored FP8 payload" + ) + if not isinstance(payload, tuple) or len(payload) != 2: + raise RuntimeError("DeepEP expanded dispatch did not return FP8 q/s") + receive_activations, receive_activation_scales = payload + if ( + receive_activations.dtype != torch.float8_e4m3fn + or receive_activation_scales.dtype != torch.float32 + or not receive_activations.is_contiguous() + or not receive_activation_scales.is_contiguous() + ): + raise RuntimeError( + "DeepEP expanded payload must retain contiguous E4M3 q and row-major FP32 scales" + ) + expanded_group_rows = _host_expanded_storage_counts(handle, local_experts) + expanded_rows = sum(expanded_group_rows) + if receive_activations.shape != (expanded_rows, model_dim): + raise RuntimeError( + "DeepEP expanded activation shape disagrees with expert counts: " + f"payload={tuple(receive_activations.shape)}, rows={expanded_rows}" + ) + if receive_activation_scales.shape != ( + expanded_rows, + model_dim // _BLOCK, + ): + raise RuntimeError("DeepEP expanded activation-scale shape mismatch") + routed_scores = routed_scores.contiguous().view(-1) + if ( + routed_scores.dtype != torch.float32 + or routed_scores.numel() != expanded_rows + ): + raise RuntimeError("DeepEP expanded route-score shape/dtype mismatch") + routes = _expanded_routes(handle) + psum = handle.psum_num_recv_tokens_per_expert + if ( + routes.device != x.device + or routes.shape[1] != topk + or psum.device != x.device + or psum.dtype != torch.int32 + or not psum.is_contiguous() + or psum.numel() != local_experts + ): + raise RuntimeError("DeepEP expanded metadata violates the MegaMoE ABI") + + preactivation = ( + _C.sm103_fp8_block128_grouped_w13_gemm_nt_canonical_expanded( + receive_activations, + receive_activation_scales, + w13_weight, + w13_scale, + expanded_group_rows, + ) + ) + hidden_quantized, hidden_scales = _C.sm103_fp8_block128_swiglu_quantize( + preactivation + ) + routed_output = _C.sm103_fp8_block128_grouped_gemm_nt_expanded( + hidden_quantized, + hidden_scales, + w2_weight, + w2_scale, + expanded_group_rows, + ) + scaled_output = _C.sm103_fp8_block128_expanded_post_down_scale( + routed_output, routed_scores + ) + output, combined_scores, event = buffer_state.buffer.combine( + scaled_output, + handle=handle, + num_sms=buffer_state.num_sms, + num_qps=buffer_state.num_qps, + async_with_compute_stream=True, + allocate_on_comm_stream=False, + ) + event.current_stream_wait() + if combined_scores is not None or output.shape != x.shape: + raise RuntimeError( + "DeepEP expanded combine violated the output contract" + ) + + ctx.group_state = group_state + ctx.buffer_state = buffer_state + ctx.handle = handle + ctx.expanded_group_rows = expanded_group_rows + ctx.tokens = tokens + ctx.model_dim = model_dim + ctx.hidden = hidden + ctx.local_experts = local_experts + ctx.topk = topk + ctx.expanded_rows = expanded_rows + ctx.master_gradient_wrapper = master_gradient_wrapper + ctx.save_for_backward( + routed_scores, + routes, + psum, + receive_activations, + receive_activation_scales, + preactivation, + hidden_quantized, + hidden_scales, + routed_output, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) + return output + + @staticmethod + def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: + ( + routed_scores, + routes, + psum, + receive_activations, + receive_activation_scales, + preactivation, + hidden_quantized, + hidden_scales, + routed_output, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + ) = ctx.saved_tensors + 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_backward" + ): + ( + grad_output_compact, + _recv_ids, + _recv_scores, + _reverse_handle, + event, + ) = ctx.buffer_state.buffer.dispatch( + grad_output, + handle=ctx.handle, + num_sms=ctx.buffer_state.num_sms, + num_qps=ctx.buffer_state.num_qps, + async_with_compute_stream=True, + allocate_on_comm_stream=False, + do_cpu_sync=False, + do_expand=False, + ) + event.current_stream_wait() + if not isinstance(grad_output_compact, torch.Tensor): + raise RuntimeError( + "DeepEP reverse dispatch returned a non-tensor payload" + ) + if grad_output_compact.shape != (routes.shape[0], ctx.model_dim): + raise RuntimeError("DeepEP reverse dispatch compact shape mismatch") + grad_output_expanded = _C.sm103_fp8_block128_expand_compact_routes( + grad_output_compact.contiguous(), routes, ctx.expanded_rows + ) + grad_scores_compact = _C.sm103_fp8_block128_expanded_post_down_score_grad( + routed_output, grad_output_expanded, routes + ) + grad_route_quantized, grad_route_scales = ( + _C.sm103_fp8_block128_expanded_route_scale_quantize( + grad_output_expanded, routed_scores + ) + ) + + grad_hidden = _C.sm103_fp8_block128_grouped_gemm_nn_expanded( + grad_route_quantized, + grad_route_scales, + w2_weight, + w2_scale, + ctx.expanded_group_rows, + ) + grad_preactivation = _C.sm103_fp8_block128_swiglu_backward_canonical( + grad_hidden, preactivation + ) + grad_preactivation_quantized, grad_preactivation_scales = ( + _C.sm103_fp8_block128_quantize(grad_preactivation) + ) + grad_input_expanded = _C.sm103_fp8_block128_grouped_gemm_nn_expanded( + grad_preactivation_quantized, + grad_preactivation_scales, + w13_weight.view(ctx.local_experts, ctx.hidden * 2, ctx.model_dim), + w13_scale.view( + ctx.local_experts, + ctx.hidden * 2 // _BLOCK, + ctx.model_dim // _BLOCK, + ), + ctx.expanded_group_rows, + ) + + grad_route_dequantized = _C.sm103_fp8_block128_dequantize( + grad_route_quantized, grad_route_scales + ) + hidden_dequantized = _C.sm103_fp8_block128_dequantize( + hidden_quantized, hidden_scales + ) + grad_w2 = _bf16_grouped_wgrad_expanded( + grad_route_dequantized, + hidden_dequantized, + psum, + ) + del grad_route_dequantized, hidden_dequantized + input_dequantized = _C.sm103_fp8_block128_dequantize( + receive_activations, receive_activation_scales + ) + grad_w13_canonical = _bf16_grouped_wgrad_expanded( + grad_preactivation, + input_dequantized, + psum, + ) + grad_w1 = grad_w13_canonical[:, : ctx.hidden].contiguous() + grad_w3 = grad_w13_canonical[:, ctx.hidden :].contiguous() + + grad_input_compact = _C.sm103_fp8_block128_collapse_expanded_routes( + grad_input_expanded, routes + ) + grad_input, grad_scores, event = ctx.buffer_state.buffer.combine( + grad_input_compact, + handle=_shadow_compact_handle(ctx.handle), + topk_weights=grad_scores_compact, + num_sms=ctx.buffer_state.num_sms, + num_qps=ctx.buffer_state.num_qps, + async_with_compute_stream=True, + allocate_on_comm_stream=False, + ) + event.current_stream_wait() + if grad_scores is None: + raise RuntimeError( + "DeepEP backward combine omitted route-score gradients" + ) + grad_scores = grad_scores.to(torch.float32) + + 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_input, + 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, @@ -758,8 +1287,18 @@ def fp8_block128_mega_moe( expert-parallel process group; no other token transport may wrap this operation. ``master_gradient_wrapper`` lets the owning framework restore DTensor placement/reduction metadata to the three local BF16 gradients. + Multi-rank execution owns one automatically sized expanded + ``deep_ep.ElasticBuffer`` transport path internally. The single-rank + specialization exists only for focused kernel validation; neither path + admits another GPU architecture or backend. """ - return _FP8Block128MegaMoE.apply( + group_state = _resolve_group(group) + operation = ( + _FP8Block128MegaMoELocal + if group_state.world_size == 1 + else _FP8Block128MegaMoEDeepEP + ) + return operation.apply( x, topk_ids, topk_scores, @@ -770,6 +1309,6 @@ def fp8_block128_mega_moe( w1_master, w2_master, w3_master, - group, + group_state.group, master_gradient_wrapper, ) diff --git a/tests/benchmark_fp8_block128_mega_moe.py b/tests/benchmark_fp8_block128_mega_moe.py index dc75b74881..23e092c4fb 100644 --- a/tests/benchmark_fp8_block128_mega_moe.py +++ b/tests/benchmark_fp8_block128_mega_moe.py @@ -3,22 +3,27 @@ The defaults are intentionally small enough for routine companion validation. Use GLM-5.2 dimensions explicitly for integration evidence, for example:: - python tests/benchmark_fp8_block128_mega_moe.py \ - --tokens 15625 --experts 16 --model-dim 6144 --hidden 2048 --topk 8 - -The process must be launched on an otherwise idle SM103 GPU. The benchmark -prints one JSON object so callers can archive the exact dimensions and metrics. + 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 @@ -77,6 +82,12 @@ def main() -> None: 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() @@ -86,23 +97,28 @@ def main() -> None: ) 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") + 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) - device = torch.device("cuda") + 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( - args.experts, + local_experts, args.hidden, args.model_dim, device=device, @@ -112,7 +128,7 @@ def main() -> None: ).requires_grad_() w3_master = ( torch.randn( - args.experts, + local_experts, args.hidden, args.model_dim, device=device, @@ -122,7 +138,7 @@ def main() -> None: ).requires_grad_() w2_master = ( torch.randn( - args.experts, + local_experts, args.model_dim, args.hidden, device=device, @@ -133,9 +149,7 @@ def main() -> None: canonical_w13 = torch.stack( (w1_master.detach(), w3_master.detach()), dim=1 ).flatten(0, 1) - w13_q, w13_s = _blockwise_quantize( - canonical_w13 - ) + 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( @@ -145,8 +159,10 @@ def main() -> None: 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_() + (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 @@ -164,6 +180,7 @@ def forward() -> None: w1_master, w2_master, w3_master, + group=dist.group.WORLD if world_size > 1 else None, ) def forward_backward() -> None: @@ -189,6 +206,20 @@ def forward_backward() -> None: 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", @@ -199,19 +230,20 @@ def forward_backward() -> None: "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, - "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, + "world_size": world_size, + "ranks": rank_results, } - print(json.dumps(result, sort_keys=True)) + if rank == 0: + print(json.dumps(result, sort_keys=True)) + if world_size > 1: + dist.destroy_process_group() if __name__ == "__main__": diff --git a/tests/test_fp8_block128_capabilities.py b/tests/test_fp8_block128_capabilities.py index 7a0e92a918..774819ed70 100644 --- a/tests/test_fp8_block128_capabilities.py +++ b/tests/test_fp8_block128_capabilities.py @@ -21,6 +21,12 @@ def test_fp8_block128_capability_manifest_is_exact_and_fail_closed() -> None: 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"] == () diff --git a/tests/test_fp8_block128_mega_moe.py b/tests/test_fp8_block128_mega_moe.py index af76c98224..04c4d84250 100644 --- a/tests/test_fp8_block128_mega_moe.py +++ b/tests/test_fp8_block128_mega_moe.py @@ -2,8 +2,10 @@ import torch import deep_gemm -from deep_gemm.mega.fp8_block128 import _validate_master_tensor - +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), @@ -11,7 +13,9 @@ ) -def _blockwise_weight_quantize(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: +def _blockwise_weight_quantize( + weight: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: groups, rows, columns = weight.shape blocks = ( weight.float() @@ -47,27 +51,34 @@ def _blockwise_weight_dequantize( 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_() + 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 + 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 + 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 + 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 - ) + 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 = 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, @@ -83,7 +94,9 @@ def _make_case(tokens: int, *, all_to_one: bool = False) -> dict[str, torch.Tens } -def _forward_reference(case: dict[str, torch.Tensor]) -> tuple[torch.Tensor, torch.Tensor]: +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"]) @@ -97,14 +110,18 @@ def _forward_reference(case: dict[str, torch.Tensor]) -> tuple[torch.Tensor, tor 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) + 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) + 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"] ) @@ -130,23 +147,26 @@ def _ste_reference_backward( 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() + 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() + 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) + 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( @@ -156,14 +176,20 @@ def _ste_reference_backward( 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) + 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_float = ( + output_float + route_output_view[:, route].float() * scores[:, route, None] + ) output = output_float.to(torch.bfloat16) output.backward(upstream) return { @@ -200,7 +226,9 @@ def to_local(self) -> torch.Tensor: return self._local -def test_distributed_master_validation_accepts_resident_efsdp_shard_with_wrapper() -> None: +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)) @@ -229,6 +257,27 @@ def test_distributed_master_validation_requires_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)], diff --git a/tests/test_fp8_block128_mega_moe_distributed.py b/tests/test_fp8_block128_mega_moe_distributed.py index 7165872353..b71e5b64eb 100644 --- a/tests/test_fp8_block128_mega_moe_distributed.py +++ b/tests/test_fp8_block128_mega_moe_distributed.py @@ -8,6 +8,7 @@ import pytest import torch import torch.multiprocessing as mp +from torch.utils.checkpoint import checkpoint def _free_port() -> int: @@ -48,10 +49,12 @@ def _weight_dequantize(quantized: torch.Tensor, scales: torch.Tensor) -> torch.T def _rank_inputs(rank: int, device: torch.device) -> tuple[torch.Tensor, ...]: - tokens, topk, experts = 5 + rank * 2, 2, 4 + 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) + 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) @@ -61,9 +64,15 @@ def _rank_inputs(rank: int, device: torch.device) -> tuple[torch.Tensor, ...]: 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) + 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_() ) - 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 ) @@ -80,39 +89,128 @@ def _reference( full_w2_q: torch.Tensor, full_w2_s: torch.Tensor, upstream: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - x_ref = x.detach().clone().requires_grad_() - score_ref = scores.detach().clone().requires_grad_() - x_q, x_s = deep_gemm._C.sm103_fp8_block128_quantize(x_ref.detach()) +) -> 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() - x_effective = x_ref.float() + (x_dequantized - x_ref.float()).detach() 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_effective.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) + 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_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_dequantized_bf16 = 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.index_select(0, flat_ids), hidden_effective.unsqueeze(-1) - ).squeeze(-1).to(torch.bfloat16) + ) + 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_float = ( + output_float + route_view[:, route].float() * score_ref[:, route, None] + ) output = output_float.to(torch.bfloat16) - output.backward(upstream) - return output.detach(), x_ref.grad.detach(), score_ref.grad.detach() + + 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: @@ -139,8 +237,15 @@ def _worker(rank: int, world_size: int, port: int) -> None: ) 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 = ( @@ -165,9 +270,7 @@ def _worker(rank: int, world_size: int, port: int) -> None: ) * 0.05 ) - full_w13_q_canonical, full_w13_s_canonical = _weight_quantize( - full_w13_master - ) + 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 @@ -197,13 +300,34 @@ def _worker(rank: int, world_size: int, port: int) -> None: .requires_grad_() ) local_w2_master = ( - full_w2_master[expert_start:expert_end] - .clone() - .detach() - .requires_grad_() + full_w2_master[expert_start:expert_end].clone().detach().requires_grad_() ) x, ids, scores, upstream = _rank_inputs(rank, device) - expected_output, expected_x_grad, expected_score_grad = _reference( + # 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, @@ -214,6 +338,18 @@ def _worker(rank: int, world_size: int, port: int) -> None: 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, @@ -236,6 +372,23 @@ def _worker(rank: int, world_size: int, port: int) -> None: 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, @@ -250,6 +403,161 @@ def _worker(rank: int, world_size: int, port: int) -> None: 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() From 9c7bb4c3accfa290ef9e41f482f26bf73b920c43 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 08:03:07 +0800 Subject: [PATCH 06/29] Use 2-SM grouped blockwise MegaMoE GEMMs --- csrc/sm103_fp8_block128.cu | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 2a81959752..4ed392d059 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -35,8 +35,8 @@ // translation unit itself contains only sm_103a code and every host entry point // checks compute capability 10.3 exactly. namespace cutlass::gemm { -struct KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm103 final - : KernelSchedule1Sm, KernelScheduleSm100PtrArrayBlockwise {}; +struct KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm103 final + : KernelSchedule2Sm, KernelScheduleSm100PtrArrayBlockwise {}; } // namespace cutlass::gemm namespace cutlass::gemm::collective { @@ -699,8 +699,8 @@ struct SM103GroupedBlockwiseGemm { static constexpr int AlignmentB = 16; static constexpr int AlignmentC = 8; static constexpr int AlignmentD = 8; - using MmaTileShape = cute::Shape; - using ClusterShape = cute::Shape; + 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 @@ -728,7 +728,7 @@ struct SM103GroupedBlockwiseGemm { ElementD, LayoutD*, AlignmentD, - cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm>::CollectiveOp; + cutlass::epilogue::PtrArrayTmaWarpSpecialized2Sm>::CollectiveOp; using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< cutlass::arch::Sm103, @@ -744,7 +744,7 @@ struct SM103GroupedBlockwiseGemm { ClusterShape, cutlass::gemm::collective::StageCountAutoCarveout< static_cast(sizeof(typename CollectiveEpilogue::SharedStorage))>, - cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm103>::CollectiveOp; + cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm103>::CollectiveOp; using GemmKernel = cutlass::gemm::kernel::GemmUniversal< ProblemShape, From a5cf385bc4de0d9a6e0027794a78e09cb70a4085 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 10:16:32 +0800 Subject: [PATCH 07/29] feat: add persistent SM103 FP8 MegaMoE training pipeline --- .../impls/sm100_fp8_fp4_mega_moe.hpp | 30 +- csrc/sm103_fp8_block128.cu | 1179 ++- deep_gemm/include/deep_gemm/common/types.cuh | 11 + .../impls/sm100_fp8_fp4_mega_moe.cuh | 224 +- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 6320 +++++++++++++++++ .../sm103_fp8_block128_mega_moe_wgrad.cuh | 513 ++ .../include/deep_gemm/scheduler/mega_moe.cuh | 126 + deep_gemm/mega/fp8_block128.py | 1221 +--- 8 files changed, 8568 insertions(+), 1056 deletions(-) create mode 100644 deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh create mode 100644 deep_gemm/include/deep_gemm/impls/sm103_fp8_block128_mega_moe_wgrad.cuh 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..0d29ef806b 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,7 +113,11 @@ 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, @@ -118,7 +127,9 @@ static void __instantiate_kernel() {{ 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, + nullptr, + nullptr )); } }; @@ -221,6 +232,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/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 4ed392d059..0816c4c5d3 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -19,6 +19,15 @@ #include #include +#include +#include +#include +#include +#include + +#include "utils/system.hpp" +#include "jit_kernels/impls/runtime_utils.hpp" + #include #include #include @@ -92,6 +101,210 @@ 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 kPersistentSMs = 152; +constexpr uint32_t kPersistentSmemBytes = 212260; +constexpr uint32_t kWorkspaceAlignment = + deep_gemm::layout::kLCMCandidateBlockM; + +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_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; + + 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_h( + deep_gemm::layout::Data(kPersistentIntermediate), + 1, + workspace.num_max_pool_tokens, + backward_ring_dscore.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()) {} + + int64_t num_bytes() const { + return reinterpret_cast( + backward_full_grad_preact_scales.get_end_ptr()) - + reinterpret_cast(workspace.base); + } +}; + constexpr int64_t align_rows(const int64_t rows) { return (rows + kBlockK - 1) / kBlockK * kBlockK; } @@ -189,7 +402,8 @@ __global__ void sm103_quantize_bf16_e4m3_group128_kernel( const float value = __bfloat162float(input[offset]); __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; + 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; } @@ -197,6 +411,49 @@ __global__ void sm103_quantize_bf16_e4m3_group128_kernel( #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, @@ -241,7 +498,8 @@ __global__ void sm103_swiglu_quantize_group128_kernel( const float value = up * gate * sigmoid_gate; __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; + 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; } @@ -1392,6 +1650,875 @@ torch::Tensor grouped_fp8_block128_w13_gemm_nt_canonical_expanded( 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["num_sms"] = kPersistentSMs; + 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); + + 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); + + 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 8-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, + 8, + 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_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); + + 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, + kPersistentSMs, + kNumRanks, + 0x7f800000u, + false, + 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, + kPersistentSMs, + kNumRanks, + 0x7f800000u, + false, + 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(kPersistentSMs, 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_l2_acts, + tensor_map_l2_acts_sf, + tensor_map_l2_weights, + tensor_map_l2_acts_sf, + w13_scale.data_ptr(), + w2_scale.data_ptr())); + 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()); + cudaDeviceProp properties{}; + C10_CUDA_CHECK(cudaGetDeviceProperties( + &properties, buffer.get_device())); + TORCH_CHECK(properties.multiProcessorCount == kPersistentSMs, + "persistent GLM MegaMoE requires the 152-SM SM103 target, got ", + properties.multiProcessorCount); + 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 (buffer_ptrs.size() == 2) { + launch_persistent_forward<2>( + 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>( + 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); + const auto bf16_options = torch::TensorOptions() + .dtype(torch::kBFloat16) + .device(device); + + auto ring_x = torch::from_blob( + layout.l1_tokens.base, + {layout.ring_tokens, kPersistentHidden}, fp8_options); + auto ring_x_sf = torch::from_blob( + layout.l1_scales.base, + {layout.sf_ring_tokens, kPersistentHidden / 128}, + {1, static_cast(layout.sf_ring_tokens)}, + int_options); + 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_h = torch::from_blob( + layout.l2_tokens.base, + {layout.ring_tokens, kPersistentIntermediate}, fp8_options); + auto ring_h_sf = torch::from_blob( + layout.l2_scales.base, + {layout.sf_ring_tokens, kPersistentIntermediate / 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); + + auto gate_up = torch::from_blob( + layout.backward_ring_bf16.base, + {layout.ring_tokens, 2 * kPersistentIntermediate}, + {static_cast(kPersistentHidden), 1}, + bf16_options); + auto grad_h = torch::from_blob( + layout.backward_ring_bf16.get_base_ptr<__nv_bfloat16>() + + 2 * kPersistentIntermediate, + {layout.ring_tokens, kPersistentIntermediate}, + {static_cast(kPersistentHidden), 1}, + bf16_options); + auto ring_grad_x = torch::from_blob( + layout.backward_ring_bf16.base, + {layout.ring_tokens, kPersistentHidden}, + {static_cast(kPersistentHidden), 1}, + bf16_options); + + const auto tensor_map_ring_x = deep_gemm::make_tma_2d_desc( + ring_x, kPersistentHidden, layout.ring_tokens, + kPersistentBlockK, kPersistentBlockM / 2, + kPersistentHidden, 128); + const auto tensor_map_ring_x_sf = deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, ring_x_sf, + layout.sf_ring_tokens, kPersistentHidden, + kPersistentSFBlockM, 32, 1, 0, 0, false, 1); + 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_h = deep_gemm::make_tma_2d_desc( + ring_h, kPersistentIntermediate, layout.ring_tokens, + kPersistentBlockK, kPersistentBlockM / 2, + kPersistentIntermediate, 128); + const auto tensor_map_ring_h_sf = deep_gemm::make_tma_sf_desc( + cute::UMMA::Major::MN, ring_h_sf, + layout.sf_ring_tokens, kPersistentIntermediate, + 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_w13_recompute = deep_gemm::make_tma_2d_desc( + w13_weight, kPersistentHidden, + static_cast(w13_weight.size(0) * w13_weight.size(1)), + kPersistentBlockK, 8, + static_cast(w13_weight.stride(-2)), 128); + 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); + const auto tensor_map_gate_up = deep_gemm::make_tma_2d_desc( + gate_up, 2 * kPersistentIntermediate, layout.ring_tokens, + kPersistentBlockN, kPersistentStoreBlockM, + kPersistentHidden, 128); + const auto tensor_map_grad_h = deep_gemm::make_tma_2d_desc( + grad_h, kPersistentIntermediate, layout.ring_tokens, + kPersistentBlockN, kPersistentStoreBlockM, + kPersistentHidden, 128); + const auto tensor_map_grad_x = deep_gemm::make_tma_2d_desc( + ring_grad_x, kPersistentHidden, layout.ring_tokens, + kPersistentBlockN, kPersistentStoreBlockM, + kPersistentHidden, 128); + + using Kernel = decltype( + &deep_gemm::sm103_block128_backward:: + sm103_fp8_block128_mega_moe_backward_impl); + Kernel kernel = + &deep_gemm::sm103_block128_backward:: + sm103_fp8_block128_mega_moe_backward_impl; + 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(kPersistentSMs, 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(x.data_ptr()), + reinterpret_cast( + grad_output.data_ptr()), + topk_scores.data_ptr(), + layout.input_tokens.get_base_ptr(), + layout.input_scales.get_base_ptr(), + layout.backward_grad_y_tokens + .get_base_ptr(), + layout.backward_grad_y_scales.get_base_ptr(), + layout.input_topk_scores.get_base_ptr(), + layout.backward_grad_scores.get_base_ptr(), + layout.combine_tokens.get_base_ptr(), + layout.l1_tokens.get_base_ptr(), + layout.l1_scales.get_base_ptr(), + layout.backward_ring_grad_y + .get_base_ptr(), + layout.backward_ring_grad_y_scales.get_base_ptr(), + layout.l1_scores.get_base_ptr(), + layout.l2_tokens.get_base_ptr(), + layout.l2_scales.get_base_ptr(), + layout.backward_ring_grad_preact + .get_base_ptr(), + layout.backward_ring_grad_preact_scales.get_base_ptr(), + layout.backward_ring_bf16.get_base_ptr(), + layout.backward_ring_dscore.get_base_ptr(), + layout.backward_full_h.get_base_ptr(), + layout.backward_full_h_scales.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()), + grad_scores.data_ptr(), + tensor_map_ring_x, tensor_map_ring_x_sf, + tensor_map_ring_grad_y, tensor_map_ring_grad_y_sf, + tensor_map_ring_h, tensor_map_ring_h_sf, + 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_gate_up, + tensor_map_grad_h, tensor_map_grad_x, + 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 std::vector& buffer_ptrs, + const int64_t rank, + const PersistentWorkspaceLayout& layout, + const torch::Tensor& expert_counts, + const torch::Tensor& token_src_metadata +) { + constexpr int64_t output_rows = + kW2 ? kPersistentHidden : kPersistentIntermediate; + constexpr int64_t output_columns = + kW2 ? kPersistentIntermediate : kPersistentHidden; + const int64_t local_experts = kPersistentExperts / kNumRanks; + 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); + Kernel kernel = + &deep_gemm::sm103_block128_wgrad:: + sm103_fp8_block128_mega_moe_wgrad_impl; + 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(kPersistentSMs, 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; + const auto sym_buffer = deep_gemm::layout::SymBuffer( + buffer_ptrs, static_cast(rank)); + + auto* ring_operand = kW2 + ? layout.backward_ring_grad_y + .get_base_ptr() + : layout.l1_tokens.get_base_ptr(); + auto* ring_operand_sf = kW2 + ? layout.backward_ring_grad_y_scales.get_base_ptr() + : layout.l1_scales.get_base_ptr(); + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, kernel, + expert_counts.data_ptr(), + reinterpret_cast( + token_src_metadata.data_ptr()), + layout.sf_ring_tokens, + sym_buffer, layout.workspace, + layout.input_tokens.get_base_ptr(), + layout.input_scales.get_base_ptr(), + layout.backward_grad_y_tokens + .get_base_ptr(), + layout.backward_grad_y_scales.get_base_ptr(), + layout.input_topk_scores.get_base_ptr(), + ring_operand, ring_operand_sf, + layout.l1_scores.get_base_ptr(), + layout.backward_full_h.get_base_ptr(), + layout.backward_full_h_scales.get_base_ptr(), + layout.backward_full_grad_preact + .get_base_ptr(), + layout.backward_full_grad_preact_scales.get_base_ptr(), + 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()); + cudaDeviceProp properties{}; + C10_CUDA_CHECK(cudaGetDeviceProperties( + &properties, buffer.get_device())); + TORCH_CHECK(properties.multiProcessorCount == kPersistentSMs, + "persistent GLM MegaMoE requires the 152-SM SM103 target, got ", + properties.multiProcessorCount); + 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 (buffer_ptrs.size() == 2) { + launch_persistent_backward_activation<2>( + 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>( + 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 PersistentWorkspaceLayout layout( + buffer.data_ptr(), static_cast(buffer_ptrs.size()), + static_cast(context_tokens_per_rank)); + if (buffer_ptrs.size() == 2) { + launch_persistent_wgrad<2, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<2, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } else { + launch_persistent_wgrad<16, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<16, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } + 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()); @@ -1822,8 +2949,18 @@ pybind11::dict capabilities() { 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", @@ -1854,6 +2991,44 @@ pybind11::dict capabilities() { 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")); 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..7746de93f7 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 @@ -20,23 +21,21 @@ namespace deep_gemm { 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,13 +46,16 @@ 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, @@ -63,7 +65,9 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 float* l1_block128_scales = nullptr, + const float* l2_block128_scales = 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 +77,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; @@ -96,30 +109,31 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // 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 +151,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 @@ -222,7 +239,7 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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)); @@ -528,15 +545,15 @@ 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 issue_and_wait_pull_store = [&](const uint32_t& i) { ptx::mbarrier_wait_and_flip_phase(pull_mbarrier, pull_mbarrier_phase); ptx::tma_store_1d( @@ -567,7 +584,7 @@ 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); @@ -575,7 +592,7 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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]; + local_sf_ptr[j * num_sf_ring_tokens + sf_ring_token_idx] = remote_sf_ptr[j]; } __syncwarp(); @@ -585,17 +602,19 @@ 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; // 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 +662,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 +698,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 +769,94 @@ 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]. The L1 + // epilogue consumes 8-row [up, gate] pairs, so issue TMA + // loads from the two canonical expert planes directly into + // the logical interleave in shared memory. Only the 16-KiB + // shared tile is materialized; the multi-GiB weight tensor + // is never copied or repacked. + if (cute::elect_one_sync()) { + if (block_phase == sched::BlockPhase::Linear1) { + constexpr uint32_t kPairGranularity = 8; + constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; + #pragma unroll + for (uint32_t group = 0; group < kLogicalRowsPerBlock / kPairGranularity; ++group) { + const uint32_t logical_row = + n_block_idx * kLogicalRowsPerBlock + group * kPairGranularity; + 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] + + (group * 2) * kPairGranularity * BLOCK_K, + k_idx, + up_row, + 2); + tma::copy( + tensor_map_b_ptr, + &shared_storage.full_barriers[stage_idx], + shared_storage.smem_b[stage_idx] + + (group * 2 + 1) * kPairGranularity * 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 kPairGranularity = 8; + constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; + const uint32_t segment = row / kPairGranularity; + const uint32_t logical_row = + n_block_idx * kLogicalRowsPerBlock + + (segment / 2) * kPairGranularity; + const uint32_t canonical_expert = + local_expert_idx * 2 + ((segment & 1u) ? 0u : 1u); + 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) + shared_storage.full_barriers[stage_idx].arrive_and_expect_tx( + sizeof(SharedStorage::smem_b[0])); + 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( @@ -932,7 +1038,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,7 +1046,7 @@ 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 @@ -997,14 +1103,22 @@ sm100_fp8_fp4_mega_moe_impl(void* y, auto fp32_values = reinterpret_cast(raw_values); #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]); + // The upstream transformed FP4 tensor is + // [gate, up]. Canonical GLM is loaded logically as + // [up, gate] from its two expert planes. + auto bf16_gate = __float22bfloat162_rn( + fp32_values[k * 2 + (kFP8Block128Weights ? 1 : 0)]); + auto bf16_up = __float22bfloat162_rn( + fp32_values[k * 2 + (kFP8Block128Weights ? 0 : 1)]); // 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); @@ -1100,7 +1214,7 @@ sm100_fp8_fp4_mega_moe_impl(void* 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 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: @@ -1231,7 +1345,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; 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..5ea44b105c --- /dev/null +++ b/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh @@ -0,0 +1,6320 @@ +#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 kNumSMs = 152; +static constexpr uint32_t kThreads = 512; +static constexpr uint32_t kStoreBlockM = 32; +static constexpr uint32_t kEpilogueThreads = 256; +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 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 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]; + uint32_t tmem_ptr; +}; + +DG_STATIC_ASSERT(kNumTmemCols <= 512, "SM103 backward exceeds TMEM"); + +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, 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, group_idx); + return 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 gran = 8; + constexpr uint32_t logical_rows = kBlockN / 2; + #pragma unroll + for (uint32_t group = 0; + group < logical_rows / gran; ++group) { + const uint32_t logical_row = + n_block_idx * logical_rows + group * gran; + 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] + + (group * 2) * gran * kBlockK, + k_block_idx * kBlockK, up_row, 2); + tma::copy( + &tensor_map_w13_recompute, + &storage.full_barriers[stage_idx], + storage.smem_b[stage_idx] + + (group * 2 + 1) * gran * 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 gran = 8; + constexpr uint32_t logical_rows = kBlockN / 2; + const uint32_t segment = row / gran; + const uint32_t logical_row = + n_block_idx * logical_rows + + (segment / 2) * gran; + const uint32_t canonical_expert = + local_expert_idx * 2 + + ((segment & 1u) ? 0u : 1u); + 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])); + } 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()) { + auto* sfa = storage.smem_sfa[stage_idx]; + mma::sm100::replace_smem_desc_addr(sf_desc, sfa); + cute::SM100_UTCCP_4x32dp128bit_2cta::copy( + sf_desc, 384u); + mma::sm100::replace_smem_desc_addr( + sf_desc, storage.smem_sfb[stage_idx]); + cute::SM100_UTCCP_4x32dp128bit_2cta::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< + 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) { + 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>(); + } + + __syncthreads(); + 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_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); + reinterpret_cast(ring_x)[ + static_cast(row) * vecs_per_row + vec] = + *remote_x; + reinterpret_cast(ring_grad_y)[ + static_cast(row) * vecs_per_row + vec] = + *remote_dy; + } + 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); + ring_x_sf[block * sf_ring_tokens + sf_row] = + *sym_buffer.map(symmetric_x_sf + remote_index, + metadata.rank_idx); + ring_grad_y_sf[block * sf_ring_tokens + sf_row] = + *sym_buffer.map(symmetric_grad_y_sf + remote_index, + metadata.rank_idx); + } + for (uint32_t row = global_thread; row < count; + row += global_stride) { + const auto metadata = + token_src_metadata[pool_row_offset + row]; + ring_scores[row] = *sym_buffer.map( + symmetric_scores + + static_cast(metadata.token_idx) * kTopK + + metadata.topk_idx, + metadata.rank_idx); + ring_dscore[row] = 0.0f; + } + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + if (count != 0) { + run_gemm_phase( + 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(); }); + + // [up8,gate8] physical 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 / 8) * 16 + in_block % 8; + const uint32_t physical_gate = physical_up + 8; + 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 sigmoid = 1.0f / (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( + 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 / 8) * 16 + in_block % 8; + 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_up + 8]); + const float dy_h = static_cast( + ring_bf16[ + static_cast(row) * kHidden + + 2 * kIntermediate + + col]) * + ring_scores[row]; + const float sigmoid = 1.0f / (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( + 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 +} + +} // 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..b5ce86ed82 --- /dev/null +++ b/deep_gemm/include/deep_gemm/impls/sm103_fp8_block128_mega_moe_wgrad.cuh @@ -0,0 +1,513 @@ +#pragma once +#pragma clang diagnostic push +#pragma clang diagnostic ignored "-Wunknown-attributes" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace deep_gemm::sm103_block128_wgrad { + +// GLM-5.2 has enough routes per expert that a single fixed 2-CTA family is +// preferable to a runtime configuration matrix. K is the route dimension; +// BF16 UMMA consumes it in 64-row atoms after FP8 dequantization in the load +// prologue. No full BF16 route pool is materialized. +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 kRouteBlockM = 192; +static constexpr uint32_t kBlockM = 128; +static constexpr uint32_t kBlockN = 128; +static constexpr uint32_t kBlockK = 64; +static constexpr uint32_t kLoadBlockN = kBlockN / 2; +static constexpr uint32_t kStages = 6; +static constexpr uint32_t kNumSMs = 152; +static constexpr uint32_t kThreads = 256; +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 kEpilogueThreads = 128; +static constexpr uint32_t kUMMAM = 256; +static constexpr uint32_t kUMMAN = 128; +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][kBlockM * kBlockK]; + alignas(1024) bf16_t smem_b[kStages][kLoadBlockN * kBlockK]; + Barrier full_barriers[kStages]; + Barrier empty_barriers[kStages]; + Barrier tmem_full_barriers[kNumEpilogueStages]; + Barrier tmem_empty_barriers[kNumEpilogueStages]; + uint32_t tmem_ptr; +}; + +DG_STATIC_ASSERT(kNumTmemCols <= 512, "SM103 wgrad exceeds TMEM"); + +CUTLASS_DEVICE uint32_t transform_sf_row(const uint32_t row) { + const uint32_t in_block = row % kRouteBlockM; + return row / kRouteBlockM * 256u + + (in_block & ~127u) + (in_block & 31u) * 4u + + ((in_block >> 5) & 3u); +} + +CUTLASS_DEVICE float unpack_power2_scale(const uint32_t packed) { + const uint32_t exponent = packed & 0xffu; + // UE8M0 code zero denotes 2^-127. That value is an FP32 subnormal, so + // constructing it by shifting an IEEE exponent field would incorrectly + // produce zero. + return exponent == 0u ? 0x1p-127f : __uint_as_float(exponent << 23); +} + +template +CUTLASS_DEVICE void store_mn_swizzle128( + bf16_t* base, + const uint32_t row, + const uint32_t k, + const bf16_t value) { + DG_STATIC_ASSERT(kRows == 64 || kRows == 128, "invalid BF16 SMEM rows"); + const uint32_t row_in_atom = row & 7u; + const uint32_t col_byte = k * sizeof(bf16_t); + const uint32_t byte_offset = + (row >> 3) * 8u * kSwizzle + row_in_atom * kSwizzle + + ((col_byte >> 4) ^ row_in_atom) * 16u + (col_byte & 15u); + *reinterpret_cast( + reinterpret_cast(base) + byte_offset) = value; +} + +template +CUTLASS_DEVICE void gather_compact_operand( + const uint32_t count, + const uint32_t pool_row_offset, + const uint32_t sf_ring_tokens, + const layout::TokenSrcMetadata* token_src_metadata, + const layout::SymBuffer& sym_buffer, + const fp8_t* symmetric_x, + const uint32_t* symmetric_x_sf, + const fp8_t* symmetric_grad_y, + const uint32_t* symmetric_grad_y_sf, + const float* symmetric_scores, + fp8_t* ring_operand, + uint32_t* ring_operand_sf, + float* ring_scores) { + constexpr uint32_t kHiddenBlocks = kHidden / 128; + constexpr uint32_t kVecsPerRow = kHidden / sizeof(uint4); + const uint32_t global_thread = blockIdx.x * kThreads + threadIdx.x; + const uint32_t global_stride = kNumSMs * kThreads; + const fp8_t* source = kW2 ? symmetric_grad_y : symmetric_x; + const uint32_t* source_sf = + kW2 ? symmetric_grad_y_sf : symmetric_x_sf; + + for (uint64_t linear = global_thread; + linear < static_cast(count) * kVecsPerRow; + linear += global_stride) { + const uint32_t row = linear / kVecsPerRow; + const uint32_t vec = linear - + static_cast(row) * kVecsPerRow; + const auto metadata = token_src_metadata[pool_row_offset + row]; + const auto* remote = sym_buffer.map( + reinterpret_cast(source) + + static_cast(metadata.token_idx) * kVecsPerRow + + vec, + metadata.rank_idx); + reinterpret_cast(ring_operand)[ + static_cast(row) * kVecsPerRow + vec] = *remote; + } + for (uint64_t linear = global_thread; + linear < static_cast(count) * kHiddenBlocks; + linear += global_stride) { + const uint32_t row = linear / kHiddenBlocks; + const uint32_t block = linear - + static_cast(row) * kHiddenBlocks; + const auto metadata = token_src_metadata[pool_row_offset + row]; + const uint64_t remote_index = + static_cast(metadata.token_idx) * kHiddenBlocks + block; + ring_operand_sf[ + block * sf_ring_tokens + transform_sf_row(row)] = + *sym_buffer.map(source_sf + remote_index, metadata.rank_idx); + } + if constexpr (kW2) { + for (uint32_t row = global_thread; row < count; + row += global_stride) { + const auto metadata = token_src_metadata[pool_row_offset + row]; + ring_scores[row] = *sym_buffer.map( + symmetric_scores + + static_cast(metadata.token_idx) * kTopK + + metadata.topk_idx, + metadata.rank_idx); + } + } +} + +template +CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void +sm103_fp8_block128_mega_moe_wgrad_impl( + const int* expert_counts, + const layout::TokenSrcMetadata* token_src_metadata, + const uint32_t sf_ring_tokens, + const __grid_constant__ layout::SymBuffer sym_buffer, + const __grid_constant__ layout::Workspace workspace, + const fp8_t* symmetric_x, + const uint32_t* symmetric_x_sf, + const fp8_t* symmetric_grad_y, + const uint32_t* symmetric_grad_y_sf, + const float* symmetric_scores, + fp8_t* ring_operand, + uint32_t* ring_operand_sf, + float* ring_scores, + const fp8_t* full_h, + const uint32_t* full_h_sf, + const fp8_t* full_grad_preact, + const uint32_t* full_grad_preact_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 + constexpr uint32_t kLocalExperts = kGlobalExperts / kNumRanks; + constexpr uint32_t kShapeM = kW2 ? kHidden : 2 * kIntermediate; + constexpr uint32_t kShapeN = kW2 ? kIntermediate : kHidden; + constexpr uint32_t kNumMBlocks = kShapeM / kBlockM; + constexpr uint32_t kNumNBlocks = kShapeN / kBlockN; + constexpr uint32_t kNumTilesPerExpert = kNumMBlocks * kNumNBlocks; + constexpr uint32_t kFullABlocks = + (kW2 ? kHidden : 2 * kIntermediate) / 128; + constexpr uint32_t kFullBBlocks = + (kW2 ? kIntermediate : kHidden) / 128; + + 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; + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 3 && cute::elect_one_sync()) { + #pragma unroll + for (uint32_t i = 0; i < kStages; ++i) { + // A and B producer warps in both CTAs remotely arrive at CTA 0. + 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(); + } + __syncwarp(); + if (warp_idx == 3) + cute::TMEM::Allocator2Sm().allocate(kNumTmemCols, &storage.tmem_ptr); + comm::cluster_sync_with_relaxed_arrive(); + + if (warp_idx == 3) { + cute::prefetch_tma_descriptor(&tensor_map_output_0); + cute::prefetch_tma_descriptor(&tensor_map_output_1); + } + + 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; + uint32_t output_iter = 0; + uint32_t tma_stage_idx = 0; + 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 * kRouteBlockM; + DG_DEVICE_ASSERT(count <= workspace.num_ring_tokens); + DG_DEVICE_ASSERT( + pool_row_offset + count <= workspace.num_max_pool_tokens); + + // Each dedicated wgrad kernel transports its one compact operand once + // per expert. All output tiles then reuse the local FP8 ring. + gather_compact_operand( + count, pool_row_offset, sf_ring_tokens, + token_src_metadata, sym_buffer, + symmetric_x, symmetric_x_sf, + symmetric_grad_y, symmetric_grad_y_sf, + symmetric_scores, + ring_operand, ring_operand_sf, ring_scores); + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + + const uint32_t num_k_blocks = + cute::max(1u, math::ceil_div(count, kBlockK)); + + if (warp_idx == 0) { + // A: (output-M, routes). W2 reads remotely transported dy and + // applies score before BF16 rounding; W13 reads local dpreact. + for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; + tile += kNumSMs) { + const uint32_t m_block = tile % kNumMBlocks; + #pragma unroll 1 + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block) { + storage.empty_barriers[stage_idx].wait(phase ^ 1); + for (uint32_t linear = lane_idx; + linear < kBlockM * kBlockK; linear += 32) { + const uint32_t row = linear / kBlockK; + const uint32_t k = linear - row * kBlockK; + const uint32_t route = k_block * kBlockK + k; + const uint32_t m = m_block * kBlockM + row; + float value = 0.0f; + if (route < count) { + if constexpr (kW2) { + const uint32_t packed = ring_operand_sf[ + (m / 128) * sf_ring_tokens + + transform_sf_row(route)]; + value = static_cast( + ring_operand[ + static_cast(route) * + kHidden + + m]) * + unpack_power2_scale(packed) * + ring_scores[route]; + } else { + const uint64_t full_row = + pool_row_offset + route; + const uint32_t packed = + full_grad_preact_sf[ + full_row * kFullABlocks + m / 128]; + value = static_cast( + full_grad_preact[ + full_row * kShapeM + m]) * + unpack_power2_scale(packed); + } + } + store_mn_swizzle128( + storage.smem_a[stage_idx], row, k, + bf16_t(value)); + } + cutlass::arch::fence_view_async_shared(); + if (cute::elect_one_sync()) + storage.full_barriers[stage_idx].arrive(0u); + __syncwarp(); + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + } + } + } else if (warp_idx == 1) { + // B: (output-N, routes). W2 reads local h; W13 reads the compact + // x operand transported into the ring above. + for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; + tile += kNumSMs) { + const uint32_t n_block = tile / kNumMBlocks; + const uint32_t cta_n_base = + n_block * kBlockN + + cute::block_rank_in_cluster() * kLoadBlockN; + #pragma unroll 1 + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block) { + storage.empty_barriers[stage_idx].wait(phase ^ 1); + for (uint32_t linear = lane_idx; + linear < kLoadBlockN * kBlockK; linear += 32) { + const uint32_t row = linear / kBlockK; + const uint32_t k = linear - row * kBlockK; + const uint32_t route = k_block * kBlockK + k; + const uint32_t n = cta_n_base + row; + float value = 0.0f; + if (route < count) { + if constexpr (kW2) { + const uint64_t full_row = + pool_row_offset + route; + const uint32_t packed = full_h_sf[ + full_row * kFullBBlocks + n / 128]; + value = static_cast( + full_h[ + full_row * kShapeN + n]) * + unpack_power2_scale(packed); + } else { + const uint32_t packed = ring_operand_sf[ + (n / 128) * sf_ring_tokens + + transform_sf_row(route)]; + value = static_cast( + ring_operand[ + static_cast(route) * + kHidden + + n]) * + unpack_power2_scale(packed); + } + } + store_mn_swizzle128( + storage.smem_b[stage_idx], row, k, + bf16_t(value)); + } + cutlass::arch::fence_view_async_shared(); + if (cute::elect_one_sync()) + storage.full_barriers[stage_idx].arrive(0u); + __syncwarp(); + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + } + } + } else if (warp_idx == 2 && leader_cta) { + for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; + tile += kNumSMs, ++output_iter) { + const uint32_t accum_stage = + output_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (output_iter / kNumEpilogueStages) & 1u; + storage.tmem_empty_barriers[accum_stage].wait( + accum_phase ^ 1u); + ptx::tcgen05_after_thread_sync(); + + #pragma unroll 1 + for (uint32_t k_block = 0; k_block < num_k_blocks; + ++k_block) { + 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()) { + #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.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(); + stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; + phase ^= stage_idx == 0; + } + } + } else if (warp_idx >= 4) { + const uint32_t epilogue_warp_idx = warp_idx - 4; + DG_TRAP_ONLY_DEVICE_ASSERT( + ptx::ld_shared(&storage.tmem_ptr) == 0); + for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; + tile += kNumSMs, ++output_iter) { + const uint32_t m_block = tile % kNumMBlocks; + const uint32_t n_block = tile / kNumMBlocks; + const uint32_t accum_stage = + output_iter % kNumEpilogueStages; + const uint32_t accum_phase = + (output_iter / kNumEpilogueStages) & 1u; + 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(); + } + + // Do not overwrite the expert ring or start a new output wave until + // every CTA has completed all stores for this expert. + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); + pool_block_offset += math::ceil_div(count, kRouteBlockM); + } + + comm::cluster_sync_with_relaxed_arrive(); + if (warp_idx == 3) + 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/scheduler/mega_moe.cuh b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh index 05021ec89f..0b1de58b7c 100644 --- a/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh @@ -16,6 +16,132 @@ enum class BlockPhase { Linear2 = 2 }; +// Fixed training reverse for the GLM large-M specialization. Unlike the +// forward scheduler, the expert counts are immutable outputs from the saved +// forward call, so no Workspace polling is needed. One expert constitutes a +// wave: its W13 recompute, W2 dgrad, and W13 dgrad traverse the same ring slot +// before the next expert may reuse it. +enum class BackwardBlockPhase { + None = 0, + RecomputeW13 = 1, + W2Dgrad = 2, + W13Dgrad = 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 kNumSMs> +struct MegaMoEBackwardScheduler { + static constexpr uint32_t kRecomputeBlockNs = + (2 * kIntermediateHidden) / BLOCK_N; + static constexpr uint32_t kW2DgradBlockNs = + kIntermediateHidden / BLOCK_N; + static constexpr uint32_t kW13DgradBlockNs = + kHidden / BLOCK_N; + static constexpr uint32_t kRecomputeBlockKs = kHidden / BLOCK_K; + 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(kRecomputeBlockNs % 2 == 0 && + kW2DgradBlockNs % 2 == 0 && + kW13DgradBlockNs % 2 == 0, + "Every backward phase must assign adjacent N blocks to a cluster"); + + const int* expert_counts; + BackwardBlockPhase phase = BackwardBlockPhase::RecomputeW13; + uint32_t local_expert_idx = 0; + uint32_t pool_block_offset = 0; + uint32_t num_tokens = 0; + uint32_t block_idx = 0; + uint32_t m_block_idx = 0; + uint32_t n_block_idx = 0; + + CUTLASS_DEVICE explicit MegaMoEBackwardScheduler(const int* counts) + : expert_counts(counts), block_idx(blockIdx.x) { + num_tokens = static_cast(__ldg(expert_counts)); + } + + CUTLASS_DEVICE uint32_t get_num_block_ns() const { + return phase == BackwardBlockPhase::RecomputeW13 + ? kRecomputeBlockNs + : phase == BackwardBlockPhase::W2Dgrad + ? kW2DgradBlockNs + : kW13DgradBlockNs; + } + + CUTLASS_DEVICE uint32_t get_num_block_ks() const { + return phase == BackwardBlockPhase::RecomputeW13 + ? kRecomputeBlockKs + : phase == BackwardBlockPhase::W2Dgrad + ? kW2DgradBlockKs + : kW13DgradBlockKs; + } + + CUTLASS_DEVICE uint32_t get_current_pool_block_offset() const { + return pool_block_offset; + } + + CUTLASS_DEVICE uint32_t get_current_num_m_blocks() const { + return math::ceil_div(num_tokens, BLOCK_M); + } + + template + CUTLASS_DEVICE uint32_t get_valid_m() const { + const auto value = cute::min( + num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + return kDoUMMAAligned ? math::align(value, 16u) : value; + } + + CUTLASS_DEVICE cute::tuple + get_next_block() { + while (local_expert_idx < kNumExpertsPerRank) { + const uint32_t block_ns = get_num_block_ns(); + const uint32_t phase_blocks = get_current_num_m_blocks() * block_ns; + if (block_idx < phase_blocks) { + m_block_idx = block_idx / block_ns; + n_block_idx = block_idx - m_block_idx * block_ns; + block_idx += kNumSMs; + return {phase, local_expert_idx, m_block_idx, n_block_idx}; + } + + // Every role restarts from its physical CTA index for the next + // phase. Readiness counters, not a grid-wide host launch, carry + // the producer/consumer dependency between the three phases. + block_idx = blockIdx.x; + if (phase == BackwardBlockPhase::RecomputeW13) { + phase = BackwardBlockPhase::W2Dgrad; + } else if (phase == BackwardBlockPhase::W2Dgrad) { + phase = BackwardBlockPhase::W13Dgrad; + } else { + pool_block_offset += get_current_num_m_blocks(); + ++local_expert_idx; + if (local_expert_idx >= kNumExpertsPerRank) + break; + num_tokens = static_cast( + __ldg(expert_counts + local_expert_idx)); + phase = BackwardBlockPhase::RecomputeW13; + } + } + 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; + func(current_phase, expert_idx, get_num_block_ks(), + current_m_block_idx, current_n_block_idx); + } + } +}; + template dict[str, Any]: - """Return a non-launching, exact capability manifest for preflight.""" + """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: - from deep_ep import ElasticBuffer + import torch.distributed._symmetric_memory as symm_mem - has_elastic_buffer = callable(ElasticBuffer) + 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): - has_elastic_buffer = False - if not has_elastic_buffer: - missing.append("deep_ep.ElasticBuffer") + 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": "deep_ep.ElasticBuffer.expanded", - "transport_layout": "expert_aligned_128", - "transport_scale_layout": "row_major_fp32_group128", + "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": "sm103_companion_grouped_bf16", + "wgrad_backend": "two_persistent_fused_dequant_bf16", + "supported_ep": _SUPPORTED_EP, "missing_symbols": tuple(missing), } ) @@ -87,11 +76,11 @@ def transform_glm_w13_for_fp8_block128_mega_moe( canonical_weight: torch.Tensor, canonical_scale: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: - """Convert canonical ``[gate, up]`` GLM storage to ``[up; gate]``. + """Validate and return canonical ``[gate, up]`` storage without a copy. - ``canonical_weight`` is ``[2E, H, D]`` and ``canonical_scale`` is - ``[2E, H/128, D/128]``. Returned tensors are contiguous - ``[E, 2H, D]`` and ``[E, 2H/128, D/128]`` respectively. + 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") @@ -99,29 +88,17 @@ def transform_glm_w13_for_fp8_block128_mega_moe( 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("W13 dimensions must be divisible by 128") - expected_scale_shape = (experts * 2, hidden // _BLOCK, model_dim // _BLOCK) - if tuple(canonical_scale.shape) != expected_scale_shape: + 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_shape}, got {tuple(canonical_scale.shape)}" + f"canonical W13 scale shape must be {expected_scale}, " + f"got {tuple(canonical_scale.shape)}" ) - # Canonical pair index 0 is gate and 1 is up. The fused preactivation ABI - # requires up first, followed by gate. - pair_order = torch.tensor([1, 0], dtype=torch.int64, device=canonical_weight.device) - active_weight = ( - canonical_weight.view(experts, 2, hidden, model_dim) - .index_select(1, pair_order) - .reshape(experts, hidden * 2, model_dim) - .contiguous() - ) - active_scale = ( - canonical_scale.view(experts, 2, hidden // _BLOCK, model_dim // _BLOCK) - .index_select(1, pair_order) - .reshape(experts, hidden * 2 // _BLOCK, model_dim // _BLOCK) - .contiguous() - ) - return active_weight, active_scale + 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) @@ -133,17 +110,19 @@ class _GroupState: def _resolve_group(group: Any) -> _GroupState: if not dist.is_available() or not dist.is_initialized(): - if group is not None: - raise RuntimeError( - "a process group was provided before torch.distributed initialization" - ) - return _GroupState(group=None, rank=0, world_size=1) - resolved_group = dist.group.WORLD if group is None else group - return _GroupState( - group=resolved_group, - rank=dist.get_rank(resolved_group), - world_size=dist.get_world_size(resolved_group), + 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( @@ -165,12 +144,6 @@ def _check_tensor( def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: - """Return a plain local tensor without taking ownership from FireTitan. - - Expert masters may be DTensors. Their placement and gradient reduction - remain a FireTitan concern; this companion only validates their resident - local slice and returns local BF16 gradients through the optional wrapper. - """ to_local = getattr(tensor, "to_local", None) return to_local() if callable(to_local) else tensor @@ -184,13 +157,6 @@ def _validate_master_tensor( device: torch.device, master_gradient_wrapper: Any, ) -> None: - """Validate either a resident EP-local master or an at-rest DTensor shard. - - The masters are autograd anchors only; MegaMoE never reads their values. - FireTitan may therefore keep them eFSDP-sharded while the active q/s - tensors are all-gathered. In that case the framework-provided gradient - wrapper maps the full EP-local wgrad back to the master's DTensor layout. - """ local = _local_tensor(tensor) _check_tensor(local, name=name, ndim=3, dtype=torch.bfloat16, device=device) if tuple(local.shape) == local_shape: @@ -202,7 +168,8 @@ def _validate_master_tensor( ) if tuple(tensor.shape) != global_shape: raise ValueError( - f"{name} distributed logical shape must be {global_shape}, got {tuple(tensor.shape)}" + 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") @@ -221,114 +188,122 @@ def _validate_inputs( w3_master: torch.Tensor, group_state: _GroupState, master_gradient_wrapper: Any, -) -> tuple[int, int, int, int, int]: +) -> None: if not x.is_cuda: raise ValueError("FP8-block128 MegaMoE requires CUDA") - if torch.cuda.get_device_capability(x.device) != (10, 3): - capability = torch.cuda.get_device_capability(x.device) + capability = torch.cuda.get_device_capability(x.device) + if capability != (10, 3): raise RuntimeError( - f"FP8-block128 MegaMoE is SM103-only; runtime capability is {capability}. " - "No fallback is available." + 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 + 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 + 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 + 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 + 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, model_dim = x.shape - if model_dim % _BLOCK: - raise ValueError("model dimension must be divisible by 128") - if topk_ids.shape != topk_scores.shape or topk_ids.shape[0] != tokens: - raise ValueError("top-k IDs/scores must have identical [tokens, top_k] shape") - topk = topk_ids.shape[1] - if topk <= 0: - raise ValueError("top_k must be positive") - - if w13_weight.shape[0] % 2: - raise ValueError("canonical W13 must contain gate/up pairs") - local_experts = w13_weight.shape[0] // 2 - hidden, w13_k = w13_weight.shape[1:] - if local_experts <= 0 or hidden % _BLOCK or w13_k != model_dim: - raise ValueError( - "canonical W13 must have shape [2E, H, D] with D/H divisible by 128" - ) - if tuple(w13_scale.shape) != ( - local_experts * 2, - hidden // _BLOCK, - model_dim // _BLOCK, - ): - raise ValueError( - "canonical W13 scale shape does not match 128x128 weight blocks" - ) - if tuple(w2_weight.shape) != (local_experts, model_dim, hidden): - raise ValueError("W2 must have shape [local_experts, D, H]") - if tuple(w2_scale.shape) != ( - local_experts, - model_dim // _BLOCK, - hidden // _BLOCK, - ): - raise ValueError("W2 scale shape does not match 128x128 weight blocks") - global_experts = local_experts * group_state.world_size - local_w13_shape = (local_experts, hidden, model_dim) - global_w13_shape = (global_experts, hidden, model_dim) - local_w2_shape = tuple(w2_weight.shape) - global_w2_shape = (global_experts, model_dim, hidden) + _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_shape, - global_shape=global_w13_shape, + 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_shape, - global_shape=global_w2_shape, + 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_shape, - global_shape=global_w13_shape, + local_shape=local_w13, + global_shape=global_w13, device=device, master_gradient_wrapper=master_gradient_wrapper, ) if topk_ids.numel(): - # Keep the hot path free of a device-to-host scalar synchronization. - # The assertion is enqueued on the current CUDA stream and fails the - # operation rather than admitting an out-of-range route. torch._assert_async( - ((topk_ids >= 0) & (topk_ids < global_experts)).all(), - f"top-k IDs must lie in [0, {global_experts})", + ((topk_ids >= 0) & (topk_ids < _GLOBAL_EXPERTS)).all(), + f"top-k IDs must lie in [0, {_GLOBAL_EXPERTS})", ) - return tokens, model_dim, hidden, local_experts, topk @dataclass(frozen=True) -class _DeepEPBufferState: - buffer: Any - capacity: int - num_sms: int - num_qps: int +class _PersistentBufferState: + buffer: torch.Tensor + handle: Any + buffer_ptrs: tuple[int, ...] + rank: int + context_tokens_per_rank: int + workspace_info: dict[str, Any] -_deepep_buffers: dict[tuple[int, int, int, int, int], _DeepEPBufferState] = {} -_deepep_context_tokens_per_rank: dict[int, int] = {} +_persistent_context_tokens_per_rank: dict[int, int] = {} +_persistent_buffers: dict[tuple[int, int], _PersistentBufferState] = {} def _configure_fp8_block128_mega_moe_transport( @@ -336,7 +311,12 @@ def _configure_fp8_block128_mega_moe_transport( *, context_tokens_per_rank: int, ) -> None: - """Register the owning model's existing context/CP envelope once.""" + """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) @@ -344,279 +324,70 @@ def _configure_fp8_block128_mega_moe_transport( ): raise ValueError("context_tokens_per_rank must be a positive integer") group_state = _resolve_group(group) - if group_state.world_size <= 1: - raise RuntimeError("expanded DeepEP transport requires a multi-rank group") key = id(group_state.group) - existing = _deepep_context_tokens_per_rank.get(key) + existing = _persistent_context_tokens_per_rank.get(key) if existing is not None and existing != context_tokens_per_rank: raise RuntimeError( - "MegaMoE context envelope changed after transport configuration: " + "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 _deepep_buffers): - raise RuntimeError("MegaMoE transport cannot be configured after arena construction") - if existing is None: - # DeepEP reads the native NCCL communicator while calculating its arena - # size. Materialize that communicator once during model setup; otherwise - # a first-use EP group can expose an uninitialized handle to - # ncclTeamWorld. This is not a per-layer sizing collective. - if dist.get_backend(group_state.group) == "nccl": - dist.barrier( - group=group_state.group, - device_ids=[torch.cuda.current_device()], - ) - _deepep_context_tokens_per_rank[key] = context_tokens_per_rank + 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_deepep_buffer( +def _get_persistent_buffer( group_state: _GroupState, *, device: torch.device, tokens: int, - model_dim: int, - topk: int, - global_experts: int, -) -> _DeepEPBufferState: - """Create one runtime-context-sized DeepEP arena per EP/model shape. - - FireTitan registers its already-resolved context/CP envelope when it - installs the EP group. Warmup shape therefore never affects sizing, while - changing the existing context length or CP configuration only requires a - normal trainer restart, not an image rebuild. No hot-path sizing collective - or caller-visible operation argument exists. - """ - if group_state.world_size <= 1: - raise RuntimeError("expanded DeepEP requires a multi-rank process group") - capacity = _deepep_context_tokens_per_rank.get(id(group_state.group)) - if capacity is None: - raise RuntimeError( - "MegaMoE transport was not configured from the owning model context" - ) - if tokens > capacity: +) -> _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 resolved context/CP envelope: " - f"actual={tokens}, capacity={capacity}" + "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, model_dim, topk, global_experts) - cached = _deepep_buffers.get(key) + 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 - try: - from deep_ep import ElasticBuffer - except (AttributeError, ImportError) as exc: - raise RuntimeError( - "MegaMoE requires deep_ep.ElasticBuffer; no transport fallback exists" - ) from exc + import torch.distributed._symmetric_memory as symm_mem - buffer_kwargs = dict( - num_max_tokens_per_rank=capacity, - hidden=model_dim, - num_topk=topk, - allow_hybrid_mode=True, - allow_multiple_reduction=False, - ) - fp8_bytes = ElasticBuffer.get_buffer_size_hint( - group_state.group, use_fp8_dispatch=True, **buffer_kwargs - ) - bf16_bytes = ElasticBuffer.get_buffer_size_hint( - group_state.group, use_fp8_dispatch=False, **buffer_kwargs + info = dict( + _C.sm103_fp8_block128_persistent_workspace_info( + group_state.world_size, context_tokens + ) ) - buffer = ElasticBuffer( - group_state.group, - num_bytes=max(fp8_bytes, bf16_bytes), - use_fp8_dispatch=True, - deterministic=True, - prefer_overlap_with_compute=True, - **buffer_kwargs, + buffer = symm_mem.empty( + int(info["num_bytes"]), dtype=torch.int8, device=device ) - num_sms = int(buffer.get_theoretical_num_sms(global_experts, topk)) - device_sms = torch.cuda.get_device_properties(device).multi_processor_count - if not 1 <= num_sms <= device_sms: - raise RuntimeError( - "DeepEP returned an invalid SM103 launch width: " - f"selected={num_sms}, device_sms={device_sms}" - ) - num_qps = int(buffer.get_theoretical_num_qps(num_sms)) - if num_qps < 1: - raise RuntimeError(f"DeepEP returned an invalid QP count: {num_qps}") - state = _DeepEPBufferState( + 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, - capacity=capacity, - num_sms=num_sms, - num_qps=num_qps, + handle=handle, + buffer_ptrs=pointers, + rank=group_state.rank, + context_tokens_per_rank=context_tokens, + workspace_info=info, ) - _deepep_buffers[key] = state + _persistent_buffers[key] = state return state -def _host_expanded_storage_counts(handle: Any, local_experts: int) -> tuple[int, ...]: - counts = getattr(handle, "num_recv_tokens_per_expert_list", None) - if counts is None: - raise RuntimeError( - "DeepEP expanded dispatch did not return host storage counts" - ) - if isinstance(counts, torch.Tensor): - if counts.device.type != "cpu": - raise RuntimeError("DeepEP host expert counts unexpectedly reside on CUDA") - counts = counts.tolist() - result = tuple(int(value) for value in counts) - if ( - len(result) != local_experts - or any(value < 0 for value in result) - or any(value and value % _PAD_ROWS for value in result) - ): - raise RuntimeError( - "DeepEP expanded storage counts violate the aligned local expert contract: " - f"expected={local_experts}, actual={result}" - ) - return result - - -def _expanded_routes(handle: Any) -> torch.Tensor: - metadata = getattr(handle, "recv_src_metadata", None) - if ( - not isinstance(metadata, torch.Tensor) - or metadata.ndim != 2 - or metadata.shape[1] < 3 - ): - shape = tuple(metadata.shape) if isinstance(metadata, torch.Tensor) else None - raise RuntimeError(f"invalid DeepEP expanded route metadata: {shape}") - return metadata[:, 2:] - - -def _shadow_compact_handle(handle: Any) -> Any: - shadow = copy(handle) - shadow.do_expand = False - return shadow - - -def _exchange_counts( - send_counts: Sequence[int], group_state: _GroupState, device: torch.device -) -> list[int]: - if group_state.world_size == 1: - return list(send_counts) - send = torch.tensor(send_counts, device=device, dtype=torch.int64) - receive = torch.empty_like(send) - dist.all_to_all_single(receive, send, group=group_state.group) - return [int(value) for value in receive.cpu().tolist()] - - -def _all_to_all_rows( - tensor: torch.Tensor, - send_counts: Sequence[int], - receive_counts: Sequence[int], - group_state: _GroupState, -) -> torch.Tensor: - if group_state.world_size == 1: - return tensor - output = torch.empty( - (sum(receive_counts), *tensor.shape[1:]), - dtype=tensor.dtype, - device=tensor.device, - ) - source = tensor.view(torch.uint8) if tensor.dtype == torch.float8_e4m3fn else tensor - destination = ( - output.view(torch.uint8) if output.dtype == torch.float8_e4m3fn else output - ) - dist.all_to_all_single( - destination, - source, - output_split_sizes=list(receive_counts), - input_split_sizes=list(send_counts), - group=group_state.group, - ) - return output - - -def _inverse_permutation(order: torch.Tensor) -> torch.Tensor: - inverse = torch.empty_like(order) - inverse.scatter_(0, order, torch.arange(order.numel(), device=order.device)) - return inverse - - -def _padding_state( - counts: Sequence[int], device: torch.device -) -> tuple[list[int], torch.Tensor]: - padded_counts = [ - ((count + _PAD_ROWS - 1) // _PAD_ROWS) * _PAD_ROWS if count else 0 - for count in counts - ] - total_actual = sum(counts) - if total_actual == 0: - return padded_counts, torch.empty(0, dtype=torch.int64, device=device) - count_tensor = torch.tensor(counts, dtype=torch.int64, device=device) - padded_tensor = torch.tensor(padded_counts, dtype=torch.int64, device=device) - group_ids = torch.repeat_interleave( - torch.arange(len(counts), dtype=torch.int64, device=device), count_tensor - ) - padding_before = torch.cumsum(padded_tensor - count_tensor, dim=0) - ( - padded_tensor - count_tensor - ) - actual_to_padded = torch.arange(total_actual, dtype=torch.int64, device=device) - actual_to_padded.add_(padding_before.index_select(0, group_ids)) - return padded_counts, actual_to_padded - - -def _pad_rows( - tensor: torch.Tensor, - actual_to_padded: torch.Tensor, - padded_rows: int, - *, - fill_value: float, -) -> torch.Tensor: - output = torch.full( - (padded_rows, *tensor.shape[1:]), - fill_value, - dtype=tensor.dtype, - device=tensor.device, - ) - if tensor.shape[0]: - if tensor.dtype == torch.float8_e4m3fn: - output.view(torch.uint8).index_copy_( - 0, actual_to_padded, tensor.view(torch.uint8) - ) - else: - output.index_copy_(0, actual_to_padded, tensor) - return output - - -def _unpad_rows(tensor: torch.Tensor, actual_to_padded: torch.Tensor) -> torch.Tensor: - if tensor.dtype == torch.float8_e4m3fn: - output = tensor.view(torch.uint8).index_select(0, actual_to_padded) - return output.view(torch.float8_e4m3fn) - return tensor.index_select(0, actual_to_padded) - - -def _bf16_grouped_wgrad( - left: torch.Tensor, - right: torch.Tensor, - padded_counts: Sequence[int], -) -> torch.Tensor: - return _C.sm103_fp8_block128_grouped_bf16_wgrad( - left.contiguous(), - right.contiguous(), - list(padded_counts), - ) - - -def _bf16_grouped_wgrad_expanded( - left: torch.Tensor, - right: torch.Tensor, - psum: torch.Tensor, -) -> torch.Tensor: - return _C.sm103_fp8_block128_grouped_bf16_wgrad_expanded( - left.contiguous(), - right.contiguous(), - psum, - ) - - -class _FP8Block128MegaMoELocal(torch.autograd.Function): +class _FP8Block128MegaMoEPersistent(torch.autograd.Function): @staticmethod def forward( ctx: Any, @@ -634,13 +405,7 @@ def forward( master_gradient_wrapper: Any, ) -> torch.Tensor: group_state = _resolve_group(group) - if group_state.world_size != 1: - raise RuntimeError("the packed local MegaMoE path is single-rank only") - if master_gradient_wrapper is not None and not callable( - master_gradient_wrapper - ): - raise TypeError("master_gradient_wrapper must be callable or None") - tokens, model_dim, hidden, local_experts, topk = _validate_inputs( + _validate_inputs( x, topk_ids, topk_scores, @@ -654,125 +419,40 @@ def forward( 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_forward" + "sm103_fp8_block128_megamoe_persistent_forward" ): - flat_ids = topk_ids.flatten() - send_order = torch.argsort(flat_ids, stable=True) - sorted_ids = flat_ids.index_select(0, send_order) - destinations = torch.div(sorted_ids, local_experts, rounding_mode="floor") - send_counts = [ - int(value) - for value in torch.bincount( - destinations, minlength=group_state.world_size + _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, ) - .cpu() - .tolist() - ] - receive_counts = _exchange_counts(send_counts, group_state, x.device) - - token_quantized, token_scales = _C.sm103_fp8_block128_quantize(x) - sorted_tokens = torch.div(send_order, topk, rounding_mode="floor") - send_activations = token_quantized.index_select(0, sorted_tokens) - send_activation_scales = token_scales.index_select(0, sorted_tokens) - receive_activations = _all_to_all_rows( - send_activations, send_counts, receive_counts, group_state - ) - receive_activation_scales = _all_to_all_rows( - send_activation_scales, send_counts, receive_counts, group_state - ) - receive_ids = _all_to_all_rows( - sorted_ids, send_counts, receive_counts, group_state - ) - - local_ids = torch.remainder(receive_ids, local_experts) - group_order = torch.argsort(local_ids, stable=True) - ungroup_order = _inverse_permutation(group_order) - grouped_activations = receive_activations.index_select(0, group_order) - grouped_activation_scales = receive_activation_scales.index_select( - 0, group_order - ) - grouped_local_ids = local_ids.index_select(0, group_order) - actual_counts = [ - int(value) - for value in torch.bincount(grouped_local_ids, minlength=local_experts) - .cpu() - .tolist() - ] - padded_counts, actual_to_padded = _padding_state(actual_counts, x.device) - padded_rows = sum(padded_counts) - padded_activations = _pad_rows( - grouped_activations, - actual_to_padded, - padded_rows, - fill_value=0, - ) - padded_activation_scales = _pad_rows( - grouped_activation_scales, - actual_to_padded, - padded_rows, - fill_value=1, - ) - - preactivation = _C.sm103_fp8_block128_grouped_w13_gemm_nt_canonical( - padded_activations, - padded_activation_scales, - w13_weight, - w13_scale, - padded_counts, - ) - hidden_quantized, hidden_scales = _C.sm103_fp8_block128_swiglu_quantize( - preactivation - ) - routed_output_padded = _C.sm103_fp8_block128_grouped_gemm_nt( - hidden_quantized, - hidden_scales, - w2_weight, - w2_scale, - padded_counts, - ) - routed_output_grouped = _unpad_rows(routed_output_padded, actual_to_padded) - routed_output_receive_order = routed_output_grouped.index_select( - 0, ungroup_order ) - routed_output_send_order = _all_to_all_rows( - routed_output_receive_order, - receive_counts, - send_counts, - group_state, - ) - routed_output = torch.empty( - (tokens * topk, model_dim), - dtype=torch.bfloat16, - device=x.device, - ) - if routed_output.shape[0]: - routed_output.index_copy_(0, send_order, routed_output_send_order) - output = _C.sm103_fp8_block128_post_down_combine(routed_output, topk_scores) - - ctx.group_state = group_state - ctx.send_counts = send_counts - ctx.receive_counts = receive_counts - ctx.padded_counts = padded_counts - ctx.tokens = tokens - ctx.model_dim = model_dim - ctx.hidden = hidden - ctx.local_experts = local_experts - ctx.topk = topk + ctx.buffer_state = state ctx.master_gradient_wrapper = master_gradient_wrapper ctx.save_for_backward( + x, topk_scores, - send_order, - group_order, - ungroup_order, - actual_to_padded, - padded_activations, - padded_activation_scales, - preactivation, - hidden_quantized, - hidden_scales, - routed_output, + expert_counts, + token_src_metadata, w13_weight, w13_scale, w2_weight, @@ -783,471 +463,44 @@ def forward( @staticmethod def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: ( + x, topk_scores, - send_order, - group_order, - ungroup_order, - actual_to_padded, - padded_activations, - padded_activation_scales, - preactivation, - hidden_quantized, - hidden_scales, - routed_output, + 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_backward" - ): - grad_scores = _C.sm103_fp8_block128_post_down_score_grad( - routed_output, grad_output, ctx.topk - ) - grad_route_quantized_send, grad_route_scales_send = ( - _C.sm103_fp8_block128_route_scale_quantize( - grad_output, topk_scores, send_order - ) - ) - grad_route_quantized_receive = _all_to_all_rows( - grad_route_quantized_send, - ctx.send_counts, - ctx.receive_counts, - ctx.group_state, - ) - grad_route_scales_receive = _all_to_all_rows( - grad_route_scales_send, - ctx.send_counts, - ctx.receive_counts, - ctx.group_state, - ) - grad_route_quantized_grouped = grad_route_quantized_receive.index_select( - 0, group_order - ) - grad_route_scales_grouped = grad_route_scales_receive.index_select( - 0, group_order - ) - padded_rows = sum(ctx.padded_counts) - grad_route_quantized = _pad_rows( - grad_route_quantized_grouped, - actual_to_padded, - padded_rows, - fill_value=0, - ) - grad_route_scales = _pad_rows( - grad_route_scales_grouped, - actual_to_padded, - padded_rows, - fill_value=1, - ) - - grad_hidden = _C.sm103_fp8_block128_grouped_gemm_nn( - grad_route_quantized, - grad_route_scales, - w2_weight, - w2_scale, - ctx.padded_counts, - ) - grad_preactivation = _C.sm103_fp8_block128_swiglu_backward_canonical( - grad_hidden, preactivation - ) - grad_preactivation_quantized, grad_preactivation_scales = ( - _C.sm103_fp8_block128_quantize(grad_preactivation) - ) - grad_input_padded = _C.sm103_fp8_block128_grouped_gemm_nn( - grad_preactivation_quantized, - grad_preactivation_scales, - w13_weight.view(ctx.local_experts, ctx.hidden * 2, ctx.model_dim), - w13_scale.view( - ctx.local_experts, - ctx.hidden * 2 // _BLOCK, - ctx.model_dim // _BLOCK, - ), - ctx.padded_counts, - ) - - grad_route_dequantized = _C.sm103_fp8_block128_dequantize( - grad_route_quantized, grad_route_scales - ) - hidden_dequantized = _C.sm103_fp8_block128_dequantize( - hidden_quantized, hidden_scales - ) - grad_w2 = _bf16_grouped_wgrad( - grad_route_dequantized, - hidden_dequantized, - ctx.padded_counts, - ) - input_dequantized = _C.sm103_fp8_block128_dequantize( - padded_activations, padded_activation_scales - ) - grad_w13_canonical = _bf16_grouped_wgrad( - grad_preactivation, - input_dequantized, - ctx.padded_counts, - ) - grad_w1 = grad_w13_canonical[:, : ctx.hidden].contiguous() - grad_w3 = grad_w13_canonical[:, ctx.hidden :].contiguous() - - grad_input_grouped = _unpad_rows(grad_input_padded, actual_to_padded) - grad_input_receive_order = grad_input_grouped.index_select(0, ungroup_order) - grad_input_send_order = _all_to_all_rows( - grad_input_receive_order, - ctx.receive_counts, - ctx.send_counts, - ctx.group_state, - ) - grad_input_routes = torch.empty( - (ctx.tokens * ctx.topk, ctx.model_dim), - dtype=torch.bfloat16, - device=grad_output.device, - ) - if grad_input_routes.shape[0]: - grad_input_routes.index_copy_(0, send_order, grad_input_send_order) - grad_input = _C.sm103_fp8_block128_route_sum( - grad_input_routes, ctx.tokens, ctx.topk - ) - - 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_input, - None, - grad_scores, - None, - None, - None, - None, - grad_w1, - grad_w2, - grad_w3, - None, - None, - ) - - -class _FP8Block128MegaMoEDeepEP(torch.autograd.Function): - """Invocation-owned SM103 compute over DeepEP's expanded route layout.""" - - @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) - if group_state.world_size <= 1: - raise RuntimeError( - "the expanded DeepEP MegaMoE path requires multiple ranks" - ) - if master_gradient_wrapper is not None and not callable( - master_gradient_wrapper - ): - raise TypeError("master_gradient_wrapper must be callable or None") - tokens, model_dim, hidden, local_experts, topk = _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, - ) - global_experts = local_experts * group_state.world_size - buffer_state = _get_deepep_buffer( - group_state, - device=x.device, - tokens=tokens, - model_dim=model_dim, - topk=topk, - global_experts=global_experts, - ) - with torch.autograd.profiler.record_function( - "sm103_fp8_block128_megamoe_forward" + "sm103_fp8_block128_megamoe_persistent_backward" ): - token_quantized, token_scales = _C.sm103_fp8_block128_quantize(x) - payload, recv_ids, routed_scores, handle, event = ( - buffer_state.buffer.dispatch( - (token_quantized, token_scales), - topk_idx=topk_ids, - topk_weights=topk_scores, - num_experts=global_experts, - num_max_tokens_per_rank=buffer_state.capacity, - expert_alignment=_PAD_ROWS, - num_sms=buffer_state.num_sms, - num_qps=buffer_state.num_qps, - async_with_compute_stream=True, - allocate_on_comm_stream=False, - do_cpu_sync=True, - do_expand=True, - use_tma_aligned_col_major_sf=False, - ) - ) - event.current_stream_wait() - if recv_ids is not None or routed_scores is None: - raise RuntimeError( - "DeepEP did not return an expanded scored FP8 payload" - ) - if not isinstance(payload, tuple) or len(payload) != 2: - raise RuntimeError("DeepEP expanded dispatch did not return FP8 q/s") - receive_activations, receive_activation_scales = payload - if ( - receive_activations.dtype != torch.float8_e4m3fn - or receive_activation_scales.dtype != torch.float32 - or not receive_activations.is_contiguous() - or not receive_activation_scales.is_contiguous() - ): - raise RuntimeError( - "DeepEP expanded payload must retain contiguous E4M3 q and row-major FP32 scales" - ) - expanded_group_rows = _host_expanded_storage_counts(handle, local_experts) - expanded_rows = sum(expanded_group_rows) - if receive_activations.shape != (expanded_rows, model_dim): - raise RuntimeError( - "DeepEP expanded activation shape disagrees with expert counts: " - f"payload={tuple(receive_activations.shape)}, rows={expanded_rows}" - ) - if receive_activation_scales.shape != ( - expanded_rows, - model_dim // _BLOCK, - ): - raise RuntimeError("DeepEP expanded activation-scale shape mismatch") - routed_scores = routed_scores.contiguous().view(-1) - if ( - routed_scores.dtype != torch.float32 - or routed_scores.numel() != expanded_rows - ): - raise RuntimeError("DeepEP expanded route-score shape/dtype mismatch") - routes = _expanded_routes(handle) - psum = handle.psum_num_recv_tokens_per_expert - if ( - routes.device != x.device - or routes.shape[1] != topk - or psum.device != x.device - or psum.dtype != torch.int32 - or not psum.is_contiguous() - or psum.numel() != local_experts - ): - raise RuntimeError("DeepEP expanded metadata violates the MegaMoE ABI") - - preactivation = ( - _C.sm103_fp8_block128_grouped_w13_gemm_nt_canonical_expanded( - receive_activations, - receive_activation_scales, + 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, - expanded_group_rows, + w2_weight, + w2_scale, ) ) - hidden_quantized, hidden_scales = _C.sm103_fp8_block128_swiglu_quantize( - preactivation - ) - routed_output = _C.sm103_fp8_block128_grouped_gemm_nt_expanded( - hidden_quantized, - hidden_scales, - w2_weight, - w2_scale, - expanded_group_rows, - ) - scaled_output = _C.sm103_fp8_block128_expanded_post_down_scale( - routed_output, routed_scores - ) - output, combined_scores, event = buffer_state.buffer.combine( - scaled_output, - handle=handle, - num_sms=buffer_state.num_sms, - num_qps=buffer_state.num_qps, - async_with_compute_stream=True, - allocate_on_comm_stream=False, - ) - event.current_stream_wait() - if combined_scores is not None or output.shape != x.shape: - raise RuntimeError( - "DeepEP expanded combine violated the output contract" - ) - - ctx.group_state = group_state - ctx.buffer_state = buffer_state - ctx.handle = handle - ctx.expanded_group_rows = expanded_group_rows - ctx.tokens = tokens - ctx.model_dim = model_dim - ctx.hidden = hidden - ctx.local_experts = local_experts - ctx.topk = topk - ctx.expanded_rows = expanded_rows - ctx.master_gradient_wrapper = master_gradient_wrapper - ctx.save_for_backward( - routed_scores, - routes, - psum, - receive_activations, - receive_activation_scales, - preactivation, - hidden_quantized, - hidden_scales, - routed_output, - w13_weight, - w13_scale, - w2_weight, - w2_scale, - ) - return output - - @staticmethod - def backward(ctx: Any, grad_output: torch.Tensor) -> tuple[Any, ...]: - ( - routed_scores, - routes, - psum, - receive_activations, - receive_activation_scales, - preactivation, - hidden_quantized, - hidden_scales, - routed_output, - w13_weight, - w13_scale, - w2_weight, - w2_scale, - ) = ctx.saved_tensors - 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_backward" - ): - ( - grad_output_compact, - _recv_ids, - _recv_scores, - _reverse_handle, - event, - ) = ctx.buffer_state.buffer.dispatch( - grad_output, - handle=ctx.handle, - num_sms=ctx.buffer_state.num_sms, - num_qps=ctx.buffer_state.num_qps, - async_with_compute_stream=True, - allocate_on_comm_stream=False, - do_cpu_sync=False, - do_expand=False, - ) - event.current_stream_wait() - if not isinstance(grad_output_compact, torch.Tensor): - raise RuntimeError( - "DeepEP reverse dispatch returned a non-tensor payload" - ) - if grad_output_compact.shape != (routes.shape[0], ctx.model_dim): - raise RuntimeError("DeepEP reverse dispatch compact shape mismatch") - grad_output_expanded = _C.sm103_fp8_block128_expand_compact_routes( - grad_output_compact.contiguous(), routes, ctx.expanded_rows - ) - grad_scores_compact = _C.sm103_fp8_block128_expanded_post_down_score_grad( - routed_output, grad_output_expanded, routes - ) - grad_route_quantized, grad_route_scales = ( - _C.sm103_fp8_block128_expanded_route_scale_quantize( - grad_output_expanded, routed_scores - ) - ) - - grad_hidden = _C.sm103_fp8_block128_grouped_gemm_nn_expanded( - grad_route_quantized, - grad_route_scales, - w2_weight, - w2_scale, - ctx.expanded_group_rows, - ) - grad_preactivation = _C.sm103_fp8_block128_swiglu_backward_canonical( - grad_hidden, preactivation - ) - grad_preactivation_quantized, grad_preactivation_scales = ( - _C.sm103_fp8_block128_quantize(grad_preactivation) - ) - grad_input_expanded = _C.sm103_fp8_block128_grouped_gemm_nn_expanded( - grad_preactivation_quantized, - grad_preactivation_scales, - w13_weight.view(ctx.local_experts, ctx.hidden * 2, ctx.model_dim), - w13_scale.view( - ctx.local_experts, - ctx.hidden * 2 // _BLOCK, - ctx.model_dim // _BLOCK, - ), - ctx.expanded_group_rows, - ) - - grad_route_dequantized = _C.sm103_fp8_block128_dequantize( - grad_route_quantized, grad_route_scales - ) - hidden_dequantized = _C.sm103_fp8_block128_dequantize( - hidden_quantized, hidden_scales - ) - grad_w2 = _bf16_grouped_wgrad_expanded( - grad_route_dequantized, - hidden_dequantized, - psum, - ) - del grad_route_dequantized, hidden_dequantized - input_dequantized = _C.sm103_fp8_block128_dequantize( - receive_activations, receive_activation_scales - ) - grad_w13_canonical = _bf16_grouped_wgrad_expanded( - grad_preactivation, - input_dequantized, - psum, - ) - grad_w1 = grad_w13_canonical[:, : ctx.hidden].contiguous() - grad_w3 = grad_w13_canonical[:, ctx.hidden :].contiguous() - - grad_input_compact = _C.sm103_fp8_block128_collapse_expanded_routes( - grad_input_expanded, routes - ) - grad_input, grad_scores, event = ctx.buffer_state.buffer.combine( - grad_input_compact, - handle=_shadow_compact_handle(ctx.handle), - topk_weights=grad_scores_compact, - num_sms=ctx.buffer_state.num_sms, - num_qps=ctx.buffer_state.num_qps, - async_with_compute_stream=True, - allocate_on_comm_stream=False, - ) - event.current_stream_wait() - if grad_scores is None: - raise RuntimeError( - "DeepEP backward combine omitted route-score gradients" - ) - grad_scores = grad_scores.to(torch.float32) - 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_input, + grad_x, None, grad_scores, None, @@ -1276,29 +529,15 @@ def fp8_block128_mega_moe( group: Any = None, master_gradient_wrapper: Any = None, ) -> torch.Tensor: - """Run the complete SM103 FP8-block128 routed branch. + """Run GLM's complete routed branch in the fixed SM103 pipeline. - Quantized W13 tensors retain FireTitan's canonical interleaved - ``[gate, up]`` storage. The native grouped W13 launch presents its - preactivation as ``[up; gate]`` without copying that storage. BF16 gate, - down, and up masters are passed separately, matching their checkpoint - FQNs without a packed-master allocation. Route scores are applied only - after W2 and their gradients are accumulated in FP32. ``group`` is the - expert-parallel process group; no other token transport may wrap this - operation. ``master_gradient_wrapper`` lets the owning framework restore - DTensor placement/reduction metadata to the three local BF16 gradients. - Multi-rank execution owns one automatically sized expanded - ``deep_ep.ElasticBuffer`` transport path internally. The single-rank - specialization exists only for focused kernel validation; neither path - admits another GPU architecture or backend. + 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) - operation = ( - _FP8Block128MegaMoELocal - if group_state.world_size == 1 - else _FP8Block128MegaMoEDeepEP - ) - return operation.apply( + return _FP8Block128MegaMoEPersistent.apply( x, topk_ids, topk_scores, From 7126d245399611c821be532f81937a71b535e0df Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 10:43:31 +0800 Subject: [PATCH 08/29] fix: select persistent SM103 topology automatically --- csrc/sm103_fp8_block128.cu | 176 ++++++++++++------ .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 14 +- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 7 +- 3 files changed, 128 insertions(+), 69 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 0816c4c5d3..5ad5477059 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -119,11 +119,26 @@ constexpr uint32_t kPersistentEpilogueThreads = 256; constexpr uint32_t kPersistentThreads = kPersistentDispatchThreads + kPersistentNonEpilogueThreads + kPersistentEpilogueThreads; -constexpr uint32_t kPersistentSMs = 152; +constexpr uint32_t kPersistentLocalSMs = 148; +constexpr uint32_t kPersistentProductionSMs = 152; constexpr uint32_t kPersistentSmemBytes = 212260; 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; @@ -1671,7 +1686,8 @@ pybind11::dict persistent_workspace_info( result["block_m"] = kPersistentBlockM; result["block_n"] = kPersistentBlockN; result["block_k"] = kPersistentBlockK; - result["num_sms"] = kPersistentSMs; + result["supported_num_sms"] = pybind11::make_tuple( + kPersistentLocalSMs, kPersistentProductionSMs); return result; } @@ -1732,7 +1748,7 @@ void prepare_persistent_inputs( C10_CUDA_KERNEL_LAUNCH_CHECK(); } -template +template void launch_persistent_forward( const torch::Tensor& output, const torch::Tensor& expert_counts, @@ -1858,7 +1874,7 @@ void launch_persistent_forward( kPersistentDispatchThreads, kPersistentNonEpilogueThreads, kPersistentEpilogueThreads, - kPersistentSMs, + kNumSMs, kNumRanks, 0x7f800000u, false, @@ -1881,7 +1897,7 @@ void launch_persistent_forward( kPersistentDispatchThreads, kPersistentNonEpilogueThreads, kPersistentEpilogueThreads, - kPersistentSMs, + kNumSMs, kNumRanks, 0x7f800000u, false, @@ -1896,7 +1912,7 @@ void launch_persistent_forward( attribute.id = cudaLaunchAttributeClusterDimension; attribute.val.clusterDim = {2, 1, 1}; cudaLaunchConfig_t config{}; - config.gridDim = dim3(kPersistentSMs, 1, 1); + config.gridDim = dim3(kNumSMs, 1, 1); config.blockDim = dim3(kPersistentThreads, 1, 1); config.dynamicSmemBytes = kPersistentSmemBytes; config.stream = at::cuda::getCurrentCUDAStream(buffer.get_device()); @@ -1944,12 +1960,7 @@ std::tuple persistent_forward( ) { check_sm103_device(buffer); c10::cuda::CUDAGuard guard(buffer.device()); - cudaDeviceProp properties{}; - C10_CUDA_CHECK(cudaGetDeviceProperties( - &properties, buffer.get_device())); - TORCH_CHECK(properties.multiProcessorCount == kPersistentSMs, - "persistent GLM MegaMoE requires the 152-SM SM103 target, got ", - properties.multiProcessorCount); + 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()), @@ -2008,21 +2019,35 @@ std::tuple persistent_forward( if (num_tokens == 0) return {output, expert_counts, token_src_metadata}; - if (buffer_ptrs.size() == 2) { - launch_persistent_forward<2>( - output, expert_counts, token_src_metadata, - buffer, buffer_ptrs, rank, layout, - w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + 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 { - launch_persistent_forward<16>( - output, expert_counts, token_src_metadata, - buffer, buffer_ptrs, rank, layout, - w13_weight, w13_scale, w2_weight, w2_scale, num_tokens); + 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 +template void launch_persistent_backward_activation( const torch::Tensor& grad_x, const torch::Tensor& grad_scores, @@ -2169,10 +2194,10 @@ void launch_persistent_backward_activation( using Kernel = decltype( &deep_gemm::sm103_block128_backward:: - sm103_fp8_block128_mega_moe_backward_impl); + sm103_fp8_block128_mega_moe_backward_impl); Kernel kernel = &deep_gemm::sm103_block128_backward:: - sm103_fp8_block128_mega_moe_backward_impl; + sm103_fp8_block128_mega_moe_backward_impl; constexpr uint32_t smem_bytes = sizeof( deep_gemm::sm103_block128_backward::SharedStorage); C10_CUDA_CHECK(cudaFuncSetAttribute( @@ -2181,7 +2206,7 @@ void launch_persistent_backward_activation( attribute.id = cudaLaunchAttributeClusterDimension; attribute.val.clusterDim = {2, 1, 1}; cudaLaunchConfig_t config{}; - config.gridDim = dim3(kPersistentSMs, 1, 1); + config.gridDim = dim3(kNumSMs, 1, 1); config.blockDim = dim3( deep_gemm::sm103_block128_backward::kThreads, 1, 1); config.dynamicSmemBytes = smem_bytes; @@ -2243,7 +2268,7 @@ void launch_persistent_backward_activation( C10_CUDA_KERNEL_LAUNCH_CHECK(); } -template +template void launch_persistent_wgrad( const torch::Tensor& output_0, const torch::Tensor& output_1, @@ -2280,10 +2305,12 @@ void launch_persistent_wgrad( using Kernel = decltype( &deep_gemm::sm103_block128_wgrad:: - sm103_fp8_block128_mega_moe_wgrad_impl); + sm103_fp8_block128_mega_moe_wgrad_impl< + kW2, kNumRanks, kNumSMs>); Kernel kernel = &deep_gemm::sm103_block128_wgrad:: - sm103_fp8_block128_mega_moe_wgrad_impl; + 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( @@ -2292,7 +2319,7 @@ void launch_persistent_wgrad( attribute.id = cudaLaunchAttributeClusterDimension; attribute.val.clusterDim = {2, 1, 1}; cudaLaunchConfig_t config{}; - config.gridDim = dim3(kPersistentSMs, 1, 1); + config.gridDim = dim3(kNumSMs, 1, 1); config.blockDim = dim3( deep_gemm::sm103_block128_wgrad::kThreads, 1, 1); config.dynamicSmemBytes = smem_bytes; @@ -2362,12 +2389,7 @@ persistent_backward_activation( std::numeric_limits::max(), "persistent backward context/CP envelope must be positive uint32"); c10::cuda::CUDAGuard guard(buffer.device()); - cudaDeviceProp properties{}; - C10_CUDA_CHECK(cudaGetDeviceProperties( - &properties, buffer.get_device())); - TORCH_CHECK(properties.multiProcessorCount == kPersistentSMs, - "persistent GLM MegaMoE requires the 152-SM SM103 target, got ", - properties.multiProcessorCount); + 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() && @@ -2441,18 +2463,34 @@ persistent_backward_activation( if (x.size(0) == 0) { return {grad_x, grad_scores}; } - if (buffer_ptrs.size() == 2) { - launch_persistent_backward_activation<2>( - 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); + 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 { - launch_persistent_backward_activation<16>( - 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); + 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}; } @@ -2498,23 +2536,43 @@ persistent_backward( 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 (buffer_ptrs.size() == 2) { - launch_persistent_wgrad<2, true>( - grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); - launch_persistent_wgrad<2, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + if (num_sms == kPersistentLocalSMs) { + if (buffer_ptrs.size() == 2) { + launch_persistent_wgrad<2, kPersistentLocalSMs, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<2, kPersistentLocalSMs, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } else { + launch_persistent_wgrad<16, kPersistentLocalSMs, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<16, kPersistentLocalSMs, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } } else { - launch_persistent_wgrad<16, true>( - grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); - launch_persistent_wgrad<16, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + if (buffer_ptrs.size() == 2) { + launch_persistent_wgrad<2, kPersistentProductionSMs, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<2, kPersistentProductionSMs, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } else { + launch_persistent_wgrad<16, kPersistentProductionSMs, true>( + grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + launch_persistent_wgrad<16, kPersistentProductionSMs, false>( + grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, + expert_counts, token_src_metadata); + } } return {grad_x, grad_scores, grad_w1, grad_w2, grad_w3}; } 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 index 5ea44b105c..286b255b54 100644 --- 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 @@ -5392,7 +5392,6 @@ 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 kNumSMs = 152; static constexpr uint32_t kThreads = 512; static constexpr uint32_t kStoreBlockM = 32; static constexpr uint32_t kEpilogueThreads = 256; @@ -5511,7 +5510,7 @@ CUTLASS_DEVICE uint32_t phase_shape_k() { return kHidden; } -template +template CUTLASS_DEVICE void run_gemm_phase( SharedStorage& storage, const uint32_t local_expert_idx, @@ -5870,7 +5869,7 @@ CUTLASS_DEVICE void run_gemm_phase( comm::cluster_sync_with_relaxed_arrive(); } -template +template CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void sm103_fp8_block128_mega_moe_backward_impl( const int* expert_counts, @@ -6063,7 +6062,8 @@ sm103_fp8_block128_mega_moe_backward_impl( []() { __syncthreads(); }); if (count != 0) { - run_gemm_phase( + 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, @@ -6124,7 +6124,8 @@ sm103_fp8_block128_mega_moe_backward_impl( []() { __syncthreads(); }); if (count != 0) { - run_gemm_phase( + run_gemm_phase< + sched::BackwardBlockPhase::W2Dgrad, kNumSMs>( storage, expert, count, tensor_map_ring_grad_y, tensor_map_ring_grad_y_sf, @@ -6231,7 +6232,8 @@ sm103_fp8_block128_mega_moe_backward_impl( []() { __syncthreads(); }); if (count != 0) { - run_gemm_phase( + run_gemm_phase< + sched::BackwardBlockPhase::W13Dgrad, kNumSMs>( storage, expert, count, tensor_map_ring_grad_preact, tensor_map_ring_grad_preact_sf, 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 index b5ce86ed82..f7f3b99acd 100644 --- 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 @@ -32,7 +32,6 @@ static constexpr uint32_t kBlockN = 128; static constexpr uint32_t kBlockK = 64; static constexpr uint32_t kLoadBlockN = kBlockN / 2; static constexpr uint32_t kStages = 6; -static constexpr uint32_t kNumSMs = 152; static constexpr uint32_t kThreads = 256; static constexpr uint32_t kNumEpilogueStages = 2; static constexpr uint32_t kNumTMAStoreStages = 2; @@ -97,7 +96,7 @@ CUTLASS_DEVICE void store_mn_swizzle128( reinterpret_cast(base) + byte_offset) = value; } -template +template CUTLASS_DEVICE void gather_compact_operand( const uint32_t count, const uint32_t pool_row_offset, @@ -161,7 +160,7 @@ CUTLASS_DEVICE void gather_compact_operand( } } -template +template CUTLASS_GLOBAL __launch_bounds__(kThreads, 1) void sm103_fp8_block128_mega_moe_wgrad_impl( const int* expert_counts, @@ -267,7 +266,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( // Each dedicated wgrad kernel transports its one compact operand once // per expert. All output tiles then reuse the local FP8 ring. - gather_compact_operand( + gather_compact_operand( count, pool_row_offset, sf_ring_tokens, token_src_metadata, sym_buffer, symmetric_x, symmetric_x_sf, From c84bbed0e4d48aa21a7681a9dd3a46a2c1724694 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 11:28:11 +0800 Subject: [PATCH 09/29] fix: load canonical W13 in two persistent TMA planes --- csrc/sm103_fp8_block128.cu | 4 +- .../impls/sm100_fp8_fp4_mega_moe.cuh | 169 ++++++++++++------ 2 files changed, 120 insertions(+), 53 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 5ad5477059..59c3b72de0 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -1811,13 +1811,13 @@ void launch_persistent_forward( false, 1); // Canonical [2E,H,D] is addressed as a 2-D [2E*H,D] plane. The - // persistent kernel issues separate 8-row TMA offsets for up and gate. + // 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, - 8, + kPersistentBlockN / 2, static_cast(w13_weight.stride(-2)), 128); const auto tensor_map_l1_output = deep_gemm::make_tma_2d_desc( 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 7746de93f7..0cfe47f56c 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 @@ -227,6 +227,14 @@ 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]; + // 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]; @@ -770,41 +778,35 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // TMA copy weights with SF if constexpr (kFP8Block128Weights) { - // GLM stores canonical [gate, up] as [2E, H, D]. The L1 - // epilogue consumes 8-row [up, gate] pairs, so issue TMA - // loads from the two canonical expert planes directly into - // the logical interleave in shared memory. Only the 16-KiB - // shared tile is materialized; the multi-GiB weight tensor - // is never copied or repacked. + // 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 kPairGranularity = 8; constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; - #pragma unroll - for (uint32_t group = 0; group < kLogicalRowsPerBlock / kPairGranularity; ++group) { - const uint32_t logical_row = - n_block_idx * kLogicalRowsPerBlock + group * kPairGranularity; - 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] + - (group * 2) * kPairGranularity * BLOCK_K, - k_idx, - up_row, - 2); - tma::copy( - tensor_map_b_ptr, - &shared_storage.full_barriers[stage_idx], - shared_storage.smem_b[stage_idx] + - (group * 2 + 1) * kPairGranularity * BLOCK_K, - k_idx, - gate_row, - 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, @@ -825,14 +827,12 @@ sm100_fp8_fp4_mega_moe_impl(void* y, for (uint32_t row = lane_idx; row < BLOCK_N; row += 32) { float scale; if (block_phase == sched::BlockPhase::Linear1) { - constexpr uint32_t kPairGranularity = 8; constexpr uint32_t kLogicalRowsPerBlock = BLOCK_N / 2; - const uint32_t segment = row / kPairGranularity; const uint32_t logical_row = n_block_idx * kLogicalRowsPerBlock + - (segment / 2) * kPairGranularity; + row % kLogicalRowsPerBlock; const uint32_t canonical_expert = - local_expert_idx * 2 + ((segment & 1u) ? 0u : 1u); + local_expert_idx * 2 + (row < kLogicalRowsPerBlock ? 1u : 0u); const uint32_t scale_idx = (canonical_expert * (kIntermediateHidden / 128) + logical_row / 128) * (kHidden / 128) + @@ -1049,9 +1049,10 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 @@ -1099,17 +1100,70 @@ 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) { - // The upstream transformed FP4 tensor is - // [gate, up]. Canonical GLM is loaded logically as - // [up, gate] from its two expert planes. - auto bf16_gate = __float22bfloat162_rn( - fp32_values[k * 2 + (kFP8Block128Weights ? 1 : 0)]); - auto bf16_up = __float22bfloat162_rn( - fp32_values[k * 2 + (kFP8Block128Weights ? 0 : 1)]); + auto bf16_gate = bf16_gate_values[k]; + auto bf16_up = bf16_up_values[k]; // Clamp if constexpr (kActivationClampBits != 0x7f800000u) { @@ -1184,8 +1238,10 @@ sm100_fp8_fp4_mega_moe_impl(void* y, #pragma unroll for (uint32_t i = 0; i < kNumAtomsPerStore; ++ i) { // Reduce amax (warp-pair-level) + const uint32_t amax_partner = epilogue_warp_idx ^ + (kFP8Block128Weights ? 2u : 1u); const float2 wp_amax = - shared_storage.amax_reduction[epilogue_warp_idx ^ 1][i * (ATOM_M / 2) + lane_idx % 4]; + 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); @@ -1200,7 +1256,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 @@ -1211,8 +1272,14 @@ 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 = num_sf_ring_tokens * sizeof(uint32_t); const auto sf_base_ptr = l2_sf_buffer.get_base_ptr(); From 9c63f41e82809b94a49415a92b6d977112fbc1e1 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 11:49:28 +0800 Subject: [PATCH 10/29] fix: allocate canonical W13 exchange storage --- csrc/sm103_fp8_block128.cu | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 59c3b72de0..bc929e4e0b 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -121,7 +121,16 @@ constexpr uint32_t kPersistentThreads = kPersistentEpilogueThreads; constexpr uint32_t kPersistentLocalSMs = 148; constexpr uint32_t kPersistentProductionSMs = 152; -constexpr uint32_t kPersistentSmemBytes = 212260; +// 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: 2 warpgroups * 4 warps * 32 lanes * sizeof(uint2). +// 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 kPersistentSmemBytes = + kPersistentUpstreamSmemBytes + kPersistentW13PairExchangeBytes; constexpr uint32_t kWorkspaceAlignment = deep_gemm::layout::kLCMCandidateBlockM; From 383c6ae4d7e2e36e16ebaf50c389a3009fee1589 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 12:08:49 +0800 Subject: [PATCH 11/29] fix: account for both persistent FP8 TMA peers --- .../include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) 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 0cfe47f56c..916e1b5eff 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 @@ -851,8 +851,13 @@ sm100_fp8_fp4_mega_moe_impl(void* y, __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])); + sizeof(SharedStorage::smem_b[0]) * 2); else shared_storage.full_barriers[stage_idx].arrive(0u); } From c49a0170945b19442ef5afc9b1662604c083b290 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 12:49:27 +0800 Subject: [PATCH 12/29] fix: load canonical W13 directly in persistent backward --- csrc/sm103_fp8_block128.cu | 2 +- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 60 ++++++++----------- 2 files changed, 27 insertions(+), 35 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index bc929e4e0b..6dc60bfd6b 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -2174,7 +2174,7 @@ void launch_persistent_backward_activation( const auto tensor_map_w13_recompute = deep_gemm::make_tma_2d_desc( w13_weight, kPersistentHidden, static_cast(w13_weight.size(0) * w13_weight.size(1)), - kPersistentBlockK, 8, + kPersistentBlockK, kPersistentBlockN / 2, static_cast(w13_weight.stride(-2)), 128); const auto tensor_map_w2_dgrad = deep_gemm::make_tma_b_desc( cute::UMMA::Major::MN, w2_weight, 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 index 286b255b54..2964d8e243 100644 --- 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 @@ -5617,32 +5617,26 @@ CUTLASS_DEVICE void run_gemm_phase( if constexpr ( kPhase == sched::BackwardBlockPhase::RecomputeW13) { if (cute::elect_one_sync()) { - constexpr uint32_t gran = 8; constexpr uint32_t logical_rows = kBlockN / 2; - #pragma unroll - for (uint32_t group = 0; - group < logical_rows / gran; ++group) { - const uint32_t logical_row = - n_block_idx * logical_rows + group * gran; - 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] + - (group * 2) * gran * kBlockK, - k_block_idx * kBlockK, up_row, 2); - tma::copy( - &tensor_map_w13_recompute, - &storage.full_barriers[stage_idx], - storage.smem_b[stage_idx] + - (group * 2 + 1) * gran * kBlockK, - k_block_idx * kBlockK, gate_row, 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 = @@ -5668,15 +5662,13 @@ CUTLASS_DEVICE void run_gemm_phase( float scale; if constexpr ( kPhase == sched::BackwardBlockPhase::RecomputeW13) { - constexpr uint32_t gran = 8; constexpr uint32_t logical_rows = kBlockN / 2; - const uint32_t segment = row / gran; const uint32_t logical_row = n_block_idx * logical_rows + - (segment / 2) * gran; + row % logical_rows; const uint32_t canonical_expert = local_expert_idx * 2 + - ((segment & 1u) ? 0u : 1u); + (row < logical_rows ? 1u : 0u); scale = __ldg( w13_scales + (canonical_expert * (kIntermediate / 128) + @@ -5714,7 +5706,7 @@ CUTLASS_DEVICE void run_gemm_phase( if (leader_cta) { storage.full_barriers[stage_idx] .arrive_and_expect_tx( - sizeof(storage.smem_b[0])); + sizeof(storage.smem_b[0]) * 2); } else { storage.full_barriers[stage_idx].arrive(0u); } @@ -6074,8 +6066,8 @@ sm103_fp8_block128_mega_moe_backward_impl( workspace, blockIdx.x, threadIdx.x, []() { __syncthreads(); }); - // [up8,gate8] physical W13 output -> logical h, quantized per 128 - // columns with the same power-of-two recipe as compact x. + // 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; @@ -6087,8 +6079,8 @@ sm103_fp8_block128_mega_moe_backward_impl( 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 / 8) * 16 + in_block % 8; - const uint32_t physical_gate = physical_up + 8; + 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]); From 338843f7c81afb450f2e335285cddcd4cce60fe5 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 13:18:38 +0800 Subject: [PATCH 13/29] fix: isolate persistent backward synchronization --- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 47 ++++++++++++++----- 1 file changed, 34 insertions(+), 13 deletions(-) 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 index 2964d8e243..0dd30c8b00 100644 --- 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 @@ -5403,6 +5403,12 @@ 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; +// BAR 0 is owned by __syncthreads and BAR 8 by the 256-thread TMA epilogue. +// Four independent reduction groups therefore use BAR 4..7. Reusing BAR 0 +// for a 128-thread group immediately before a 512-thread __syncthreads is a +// barrier-count race when another group reaches the block barrier first. +static constexpr uint32_t kReductionBarrierBase = 4; static constexpr uint32_t kNumTmemAccumCols = kUMMAN * kNumEpilogueStages; static constexpr uint32_t kNumTmemSFACols = kSFBlockM / 32; @@ -5459,7 +5465,7 @@ CUTLASS_DEVICE float reduce_group_128( 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, group_idx); + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); if (warp_in_group == 0) { value = lane < 4 ? storage.reduce_values[group_idx][lane] @@ -5468,7 +5474,7 @@ CUTLASS_DEVICE float reduce_group_128( if (lane == 0) storage.reduce_values[group_idx][4] = value; } - ptx::sync_aligned(128, group_idx); + ptx::sync_aligned(128, kReductionBarrierBase + group_idx); return storage.reduce_values[group_idx][4]; } @@ -5767,14 +5773,28 @@ CUTLASS_DEVICE void run_gemm_phase( const uint32_t b_base = ptx::exchange(b_desc_lo, stage_idx); if (cute::elect_one_sync()) { - auto* sfa = storage.smem_sfa[stage_idx]; - mma::sm100::replace_smem_desc_addr(sf_desc, sfa); - cute::SM100_UTCCP_4x32dp128bit_2cta::copy( - sf_desc, 384u); - mma::sm100::replace_smem_desc_addr( - sf_desc, storage.smem_sfb[stage_idx]); - cute::SM100_UTCCP_4x32dp128bit_2cta::copy( - sf_desc, 392u); + 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 = @@ -5791,7 +5811,7 @@ CUTLASS_DEVICE void run_gemm_phase( accum_stage * kUMMAN, k_block_idx > 0 || k > 0, runtime_desc, - 392u, 384u); + kTmemSFBStart, kTmemSFAStart); } } __syncwarp(); @@ -6175,13 +6195,14 @@ sm103_fp8_block128_mega_moe_backward_impl( const uint32_t w13_block = col / 64; const uint32_t in_block = col % 64; const uint32_t physical_up = - w13_block * 128 + (in_block / 8) * 16 + in_block % 8; + 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_up + 8]); + physical_gate]); const float dy_h = static_cast( ring_bf16[ static_cast(row) * kHidden + From ae383e3599a333750ff92f413d7e4c203c7e2cc1 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 13:21:39 +0800 Subject: [PATCH 14/29] fix: keep TMEM scale bases immediate --- .../include/deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 index 0dd30c8b00..382c671010 100644 --- 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 @@ -5811,7 +5811,7 @@ CUTLASS_DEVICE void run_gemm_phase( accum_stage * kUMMAN, k_block_idx > 0 || k > 0, runtime_desc, - kTmemSFBStart, kTmemSFAStart); + 392u, 384u); } } __syncwarp(); From 22c14e1052ff561e9de7d96bf7bfc7d097fb86fe Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 13:48:49 +0800 Subject: [PATCH 15/29] fix: separate reverse epilogue barrier --- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) 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 index 382c671010..2b82d2d60d 100644 --- 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 @@ -5404,10 +5404,10 @@ static constexpr uint32_t kUMMAN = kBlockM; static constexpr uint32_t kUMMAK = 32; static constexpr uint32_t kSwizzle = 128; static constexpr uint32_t kUTCCPAlignedElements = 128; -// BAR 0 is owned by __syncthreads and BAR 8 by the 256-thread TMA epilogue. -// Four independent reduction groups therefore use BAR 4..7. Reusing BAR 0 -// for a 128-thread group immediately before a 512-thread __syncthreads is a -// barrier-count race when another group reaches the block barrier first. +// BAR 0 is used both by __syncthreads and by the 256-thread TMA epilogue's +// named barrier. Four independent reduction groups therefore use BAR 4..7, +// and no block-wide barrier may run concurrently with a GEMM epilogue. +// Reusing BAR 0 with a different arrival count is a barrier-count race. static constexpr uint32_t kReductionBarrierBase = 4; static constexpr uint32_t kNumTmemAccumCols = kUMMAN * kNumEpilogueStages; @@ -5877,7 +5877,13 @@ CUTLASS_DEVICE void run_gemm_phase( cutlass::arch::warpgroup_reg_dealloc<40>(); } - __syncthreads(); + // The TMA epilogue uses named barrier 0 with a 256-thread arrival count. + // A block-wide __syncthreads here would reuse that physical barrier while + // the epilogue can still be draining, so non-epilogue warps could corrupt + // its count and strand the active CTA. The upstream 2-CTA GEMM terminates + // this role split with cluster synchronization only; the caller's next + // grid barrier supplies the subsequent block-wide synchronization after + // every thread has crossed this cluster boundary. comm::cluster_sync_with_relaxed_arrive(); } From b1319f4a45f942354741c63948259ea64d09b41f Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 14:33:42 +0800 Subject: [PATCH 16/29] fix: honor reverse swap-ab epilogue contract --- csrc/sm103_fp8_block128.cu | 7 +++-- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 29 ++++++++++--------- 2 files changed, 19 insertions(+), 17 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 6dc60bfd6b..7cd1938bc6 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -109,6 +109,7 @@ constexpr uint32_t kPersistentBlockM = 192; constexpr uint32_t kPersistentBlockN = 128; constexpr uint32_t kPersistentBlockK = 128; constexpr uint32_t kPersistentStoreBlockM = 32; +constexpr uint32_t kPersistentReverseStoreBlockM = 16; constexpr uint32_t kPersistentSFBlockM = 256; constexpr uint32_t kPersistentSFBlockN = 128; constexpr uint32_t kPersistentStages = 6; @@ -2190,15 +2191,15 @@ void launch_persistent_backward_activation( static_cast(w13_weight.size(0) / 2), 128); const auto tensor_map_gate_up = deep_gemm::make_tma_2d_desc( gate_up, 2 * kPersistentIntermediate, layout.ring_tokens, - kPersistentBlockN, kPersistentStoreBlockM, + kPersistentBlockN, kPersistentReverseStoreBlockM, kPersistentHidden, 128); const auto tensor_map_grad_h = deep_gemm::make_tma_2d_desc( grad_h, kPersistentIntermediate, layout.ring_tokens, - kPersistentBlockN, kPersistentStoreBlockM, + kPersistentBlockN, kPersistentReverseStoreBlockM, kPersistentHidden, 128); const auto tensor_map_grad_x = deep_gemm::make_tma_2d_desc( ring_grad_x, kPersistentHidden, layout.ring_tokens, - kPersistentBlockN, kPersistentStoreBlockM, + kPersistentBlockN, kPersistentReverseStoreBlockM, kPersistentHidden, 128); using Kernel = decltype( 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 index 2b82d2d60d..6d804083c4 100644 --- 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 @@ -5393,8 +5393,12 @@ 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 kStoreBlockM = 32; -static constexpr uint32_t kEpilogueThreads = 256; +// sm100_store_cd_swap_ab maps one 128-thread warpgroup over the two 64-column +// BF16 atoms of BLOCK_N. STORE_BLOCK_M=16 is also required so every valid +// UMMA-N extent (which is 16-row aligned) emits at least one store and returns +// the TMEM-empty arrival. +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; @@ -5404,10 +5408,9 @@ static constexpr uint32_t kUMMAN = kBlockM; static constexpr uint32_t kUMMAK = 32; static constexpr uint32_t kSwizzle = 128; static constexpr uint32_t kUTCCPAlignedElements = 128; -// BAR 0 is used both by __syncthreads and by the 256-thread TMA epilogue's -// named barrier. Four independent reduction groups therefore use BAR 4..7, -// and no block-wide barrier may run concurrently with a GEMM epilogue. -// Reusing BAR 0 with a different arrival count is a barrier-count race. +// 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; @@ -5440,6 +5443,9 @@ struct alignas(1024) SharedStorage { }; 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"); CUTLASS_DEVICE float warp_reduce_max(float value) { #pragma unroll @@ -5836,7 +5842,7 @@ CUTLASS_DEVICE void run_gemm_phase( .wait((last / kNumEpilogueStages) & 1); } } - } else if (warp_idx >= 8) { + } 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; @@ -5877,13 +5883,8 @@ CUTLASS_DEVICE void run_gemm_phase( cutlass::arch::warpgroup_reg_dealloc<40>(); } - // The TMA epilogue uses named barrier 0 with a 256-thread arrival count. - // A block-wide __syncthreads here would reuse that physical barrier while - // the epilogue can still be draining, so non-epilogue warps could corrupt - // its count and strand the active CTA. The upstream 2-CTA GEMM terminates - // this role split with cluster synchronization only; the caller's next - // grid barrier supplies the subsequent block-wide synchronization after - // every thread has crossed this cluster boundary. + // 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(); } From 1218b60a51cd1d85ae108b6eae924632a99d0e8c Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 15:40:27 +0800 Subject: [PATCH 17/29] feat: make MegaMoE wgrad globally persistent --- csrc/sm103_fp8_block128.cu | 93 +-- .../impls/sm100_fp8_fp4_mega_moe.cuh | 65 +- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 40 +- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 640 +++++++++--------- 4 files changed, 440 insertions(+), 398 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 7cd1938bc6..1e1d266d0b 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -179,6 +179,11 @@ struct PersistentWorkspaceLayout { 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; @@ -301,11 +306,36 @@ struct PersistentWorkspaceLayout { 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_ring_dscore.get_end_ptr()), + backward_full_scores.get_end_ptr()), backward_full_h_scales( deep_gemm::layout::Data(kPersistentIntermediate / 32), 1, @@ -2260,6 +2290,11 @@ void launch_persistent_backward_activation( layout.backward_ring_grad_preact_scales.get_base_ptr(), layout.backward_ring_bf16.get_base_ptr(), layout.backward_ring_dscore.get_base_ptr(), + layout.backward_full_x.get_base_ptr(), + layout.backward_full_x_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.backward_full_h.get_base_ptr(), layout.backward_full_h_scales.get_base_ptr(), layout.backward_full_grad_preact @@ -2283,11 +2318,8 @@ void launch_persistent_wgrad( const torch::Tensor& output_0, const torch::Tensor& output_1, const torch::Tensor& buffer, - const std::vector& buffer_ptrs, - const int64_t rank, const PersistentWorkspaceLayout& layout, - const torch::Tensor& expert_counts, - const torch::Tensor& token_src_metadata + const torch::Tensor& expert_counts ) { constexpr int64_t output_rows = kW2 ? kPersistentHidden : kPersistentIntermediate; @@ -2336,31 +2368,16 @@ void launch_persistent_wgrad( 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)); - - auto* ring_operand = kW2 - ? layout.backward_ring_grad_y - .get_base_ptr() - : layout.l1_tokens.get_base_ptr(); - auto* ring_operand_sf = kW2 - ? layout.backward_ring_grad_y_scales.get_base_ptr() - : layout.l1_scales.get_base_ptr(); C10_CUDA_CHECK(cudaLaunchKernelEx( &config, kernel, expert_counts.data_ptr(), - reinterpret_cast( - token_src_metadata.data_ptr()), - layout.sf_ring_tokens, - sym_buffer, layout.workspace, - layout.input_tokens.get_base_ptr(), - layout.input_scales.get_base_ptr(), - layout.backward_grad_y_tokens + layout.workspace.num_max_pool_tokens, + layout.backward_full_x.get_base_ptr(), + layout.backward_full_x_scales.get_base_ptr(), + layout.backward_full_grad_y .get_base_ptr(), - layout.backward_grad_y_scales.get_base_ptr(), - layout.input_topk_scores.get_base_ptr(), - ring_operand, ring_operand_sf, - layout.l1_scores.get_base_ptr(), + layout.backward_full_grad_y_scales.get_base_ptr(), + layout.backward_full_scores.get_base_ptr(), layout.backward_full_h.get_base_ptr(), layout.backward_full_h_scales.get_base_ptr(), layout.backward_full_grad_preact @@ -2554,34 +2571,26 @@ persistent_backward( if (num_sms == kPersistentLocalSMs) { if (buffer_ptrs.size() == 2) { launch_persistent_wgrad<2, kPersistentLocalSMs, true>( - grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w2, grad_w2, buffer, layout, expert_counts); launch_persistent_wgrad<2, kPersistentLocalSMs, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w1, grad_w3, buffer, layout, expert_counts); } else { launch_persistent_wgrad<16, kPersistentLocalSMs, true>( - grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w2, grad_w2, buffer, layout, expert_counts); launch_persistent_wgrad<16, kPersistentLocalSMs, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + 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, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w2, grad_w2, buffer, layout, expert_counts); launch_persistent_wgrad<2, kPersistentProductionSMs, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w1, grad_w3, buffer, layout, expert_counts); } else { launch_persistent_wgrad<16, kPersistentProductionSMs, true>( - grad_w2, grad_w2, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w2, grad_w2, buffer, layout, expert_counts); launch_persistent_wgrad<16, kPersistentProductionSMs, false>( - grad_w1, grad_w3, buffer, buffer_ptrs, rank, layout, - expert_counts, token_src_metadata); + grad_w1, grad_w3, buffer, layout, expert_counts); } } return {grad_x, grad_scores, grad_w1, grad_w2, grad_w3}; 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 916e1b5eff..615633079a 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 @@ -1078,17 +1078,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]; @@ -1210,7 +1221,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) @@ -1429,7 +1445,34 @@ 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; convert that BF16 + // value back to FP32, multiply by the exact FP32 + // route score, and round once more to BF16 before + // remote combine. Moving this multiply into L1 is + // not equivalent because it changes the E4M3 + // requantization scale and rounding before W2. + const float route_weight = *l1_topk_weights_buffer + .get_data_buffer(ring_m_idx + m_idx_in_block) + .template get_base_ptr(); + 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 index 6d804083c4..9d53318bcd 100644 --- 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 @@ -5921,6 +5921,11 @@ sm103_fp8_block128_mega_moe_backward_impl( 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, @@ -6039,12 +6044,19 @@ sm103_fp8_block128_mega_moe_backward_impl( 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] = - *remote_x; + static_cast(row) * vecs_per_row + vec] = x_value; reinterpret_cast(ring_grad_y)[ static_cast(row) * vecs_per_row + vec] = - *remote_dy; + 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; @@ -6058,22 +6070,30 @@ sm103_fp8_block128_mega_moe_backward_impl( static_cast(metadata.token_idx) * hidden_blocks + block; const uint32_t sf_row = transform_sf_row(row); - ring_x_sf[block * sf_ring_tokens + sf_row] = - *sym_buffer.map(symmetric_x_sf + remote_index, - metadata.rank_idx); - ring_grad_y_sf[block * sf_ring_tokens + sf_row] = - *sym_buffer.map(symmetric_grad_y_sf + remote_index, - metadata.rank_idx); + 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]; - ring_scores[row] = *sym_buffer.map( + 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( 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 index f7f3b99acd..ed7c758f31 100644 --- 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 @@ -10,34 +10,37 @@ #include #include #include -#include -#include #include #include #include namespace deep_gemm::sm103_block128_wgrad { -// GLM-5.2 has enough routes per expert that a single fixed 2-CTA family is -// preferable to a runtime configuration matrix. K is the route dimension; -// BF16 UMMA consumes it in 64-row atoms after FP8 dequantization in the load -// prologue. No full BF16 route pool is materialized. +// GLM-5.2 has enough routes per expert that one fixed 2-CTA family covers both +// target topologies. K is the route dimension. FP8 operands stay resident in +// the private expert-padded pool; the producer groups dequantize directly into +// pipelined BF16 shared-memory tiles consumed by native BF16 UMMA. 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 kRouteBlockM = 192; static constexpr uint32_t kBlockM = 128; static constexpr uint32_t kBlockN = 128; static constexpr uint32_t kBlockK = 64; static constexpr uint32_t kLoadBlockN = kBlockN / 2; static constexpr uint32_t kStages = 6; -static constexpr uint32_t kThreads = 256; +static constexpr uint32_t kAProducerThreads = 128; +static constexpr uint32_t kBProducerThreads = 64; +static constexpr uint32_t kMMAWarp = 6; +static constexpr uint32_t kControlWarp = 7; +static constexpr uint32_t kEpilogueFirstWarp = 8; +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 kEpilogueThreads = 128; static constexpr uint32_t kUMMAM = 256; static constexpr uint32_t kUMMAN = 128; static constexpr uint32_t kUMMAK = 16; @@ -63,20 +66,13 @@ struct alignas(1024) SharedStorage { uint32_t tmem_ptr; }; +DG_STATIC_ASSERT(kThreads == 384, "SM103 wgrad role layout changed"); DG_STATIC_ASSERT(kNumTmemCols <= 512, "SM103 wgrad exceeds TMEM"); -CUTLASS_DEVICE uint32_t transform_sf_row(const uint32_t row) { - const uint32_t in_block = row % kRouteBlockM; - return row / kRouteBlockM * 256u + - (in_block & ~127u) + (in_block & 31u) * 4u + - ((in_block >> 5) & 3u); -} - CUTLASS_DEVICE float unpack_power2_scale(const uint32_t packed) { const uint32_t exponent = packed & 0xffu; - // UE8M0 code zero denotes 2^-127. That value is an FP32 subnormal, so - // constructing it by shifting an IEEE exponent field would incorrectly - // produce zero. + // UE8M0 code zero denotes 2^-127. Construct it explicitly because + // shifting zero into an IEEE exponent field would produce zero. return exponent == 0u ? 0x1p-127f : __uint_as_float(exponent << 23); } @@ -86,7 +82,8 @@ CUTLASS_DEVICE void store_mn_swizzle128( const uint32_t row, const uint32_t k, const bf16_t value) { - DG_STATIC_ASSERT(kRows == 64 || kRows == 128, "invalid BF16 SMEM rows"); + DG_STATIC_ASSERT(kRows == 64 || kRows == 128, + "invalid BF16 SMEM rows"); const uint32_t row_in_atom = row & 7u; const uint32_t col_byte = k * sizeof(bf16_t); const uint32_t byte_offset = @@ -96,86 +93,72 @@ CUTLASS_DEVICE void store_mn_swizzle128( reinterpret_cast(base) + byte_offset) = value; } +// Each physical CTA starts at its own output tile and advances by the fixed SM +// count. Because every expert matrix has an even number of M tiles, adjacent +// CTAs in a 2-CTA cluster always address adjacent M halves of the same expert/N +// tile. Pool offsets advance only when the monotonic tile stream crosses an +// expert boundary; there is no global queue, host metadata, or per-expert sync. template -CUTLASS_DEVICE void gather_compact_operand( - const uint32_t count, - const uint32_t pool_row_offset, - const uint32_t sf_ring_tokens, - const layout::TokenSrcMetadata* token_src_metadata, - const layout::SymBuffer& sym_buffer, - const fp8_t* symmetric_x, - const uint32_t* symmetric_x_sf, - const fp8_t* symmetric_grad_y, - const uint32_t* symmetric_grad_y_sf, - const float* symmetric_scores, - fp8_t* ring_operand, - uint32_t* ring_operand_sf, - float* ring_scores) { - constexpr uint32_t kHiddenBlocks = kHidden / 128; - constexpr uint32_t kVecsPerRow = kHidden / sizeof(uint4); - const uint32_t global_thread = blockIdx.x * kThreads + threadIdx.x; - const uint32_t global_stride = kNumSMs * kThreads; - const fp8_t* source = kW2 ? symmetric_grad_y : symmetric_x; - const uint32_t* source_sf = - kW2 ? symmetric_grad_y_sf : symmetric_x_sf; +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 kTilesPerExpert = + kNumMBlocks * kNumNBlocks; + static constexpr uint32_t kTotalTiles = + kLocalExperts * kTilesPerExpert; - for (uint64_t linear = global_thread; - linear < static_cast(count) * kVecsPerRow; - linear += global_stride) { - const uint32_t row = linear / kVecsPerRow; - const uint32_t vec = linear - - static_cast(row) * kVecsPerRow; - const auto metadata = token_src_metadata[pool_row_offset + row]; - const auto* remote = sym_buffer.map( - reinterpret_cast(source) + - static_cast(metadata.token_idx) * kVecsPerRow + - vec, - metadata.rank_idx); - reinterpret_cast(ring_operand)[ - static_cast(row) * kVecsPerRow + vec] = *remote; - } - for (uint64_t linear = global_thread; - linear < static_cast(count) * kHiddenBlocks; - linear += global_stride) { - const uint32_t row = linear / kHiddenBlocks; - const uint32_t block = linear - - static_cast(row) * kHiddenBlocks; - const auto metadata = token_src_metadata[pool_row_offset + row]; - const uint64_t remote_index = - static_cast(metadata.token_idx) * kHiddenBlocks + block; - ring_operand_sf[ - block * sf_ring_tokens + transform_sf_row(row)] = - *sym_buffer.map(source_sf + remote_index, metadata.rank_idx); - } - if constexpr (kW2) { - for (uint32_t row = global_thread; row < count; - row += global_stride) { - const auto metadata = token_src_metadata[pool_row_offset + row]; - ring_scores[row] = *sym_buffer.map( - symmetric_scores + - static_cast(metadata.token_idx) * kTopK + - metadata.topk_idx, - metadata.rank_idx); + 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) {} + + 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; + m_block = expert_tile % kNumMBlocks; + n_block = expert_tile / kNumMBlocks; + 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 layout::TokenSrcMetadata* token_src_metadata, - const uint32_t sf_ring_tokens, - const __grid_constant__ layout::SymBuffer sym_buffer, - const __grid_constant__ layout::Workspace workspace, - const fp8_t* symmetric_x, - const uint32_t* symmetric_x_sf, - const fp8_t* symmetric_grad_y, - const uint32_t* symmetric_grad_y_sf, - const float* symmetric_scores, - fp8_t* ring_operand, - uint32_t* ring_operand_sf, - float* ring_scores, + const uint32_t max_pool_tokens, + const fp8_t* full_x, + const uint32_t* full_x_sf, + const fp8_t* full_grad_y, + const uint32_t* full_grad_y_sf, + const float* full_scores, const fp8_t* full_h, const uint32_t* full_h_sf, const fp8_t* full_grad_preact, @@ -183,16 +166,15 @@ sm103_fp8_block128_mega_moe_wgrad_impl( 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 - constexpr uint32_t kLocalExperts = kGlobalExperts / kNumRanks; - constexpr uint32_t kShapeM = kW2 ? kHidden : 2 * kIntermediate; - constexpr uint32_t kShapeN = kW2 ? kIntermediate : kHidden; - constexpr uint32_t kNumMBlocks = kShapeM / kBlockM; - constexpr uint32_t kNumNBlocks = kShapeN / kBlockN; - constexpr uint32_t kNumTilesPerExpert = kNumMBlocks * kNumNBlocks; - constexpr uint32_t kFullABlocks = - (kW2 ? kHidden : 2 * kIntermediate) / 128; - constexpr uint32_t kFullBBlocks = - (kW2 ? kIntermediate : kHidden) / 128; + using Scheduler = WgradTileScheduler; + constexpr uint32_t kShapeM = Scheduler::kShapeM; + constexpr uint32_t kShapeN = Scheduler::kShapeN; + constexpr uint32_t kFullABlocks = kShapeM / 128; + constexpr uint32_t kFullBBlocks = kShapeN / 128; + DG_STATIC_ASSERT(Scheduler::kNumMBlocks % 2 == 0, + "2-CTA wgrad requires paired M tiles"); + DG_STATIC_ASSERT(kNumSMs % 2 == 0, + "2-CTA wgrad requires an even SM count"); extern __shared__ __align__(1024) uint8_t smem_buffer[]; SharedStorage& storage = *reinterpret_cast(smem_buffer); @@ -201,26 +183,28 @@ sm103_fp8_block128_mega_moe_wgrad_impl( const bool leader_cta = cute::block_rank_in_cluster() == 0; comm::cluster_sync_with_relaxed_arrive(); - if (warp_idx == 3 && cute::elect_one_sync()) { + if (warp_idx == kControlWarp && cute::elect_one_sync()) { #pragma unroll for (uint32_t i = 0; i < kStages; ++i) { - // A and B producer warps in both CTAs remotely arrive at CTA 0. + // A and B producer groups in both CTAs complete each stage. 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); + storage.tmem_empty_barriers[i].init( + 2 * kEpilogueThreads); } cutlass::arch::fence_barrier_init(); } __syncwarp(); - if (warp_idx == 3) - cute::TMEM::Allocator2Sm().allocate(kNumTmemCols, &storage.tmem_ptr); + if (warp_idx == kControlWarp) + cute::TMEM::Allocator2Sm().allocate( + kNumTmemCols, &storage.tmem_ptr); comm::cluster_sync_with_relaxed_arrive(); - if (warp_idx == 3) { + if (warp_idx == kControlWarp) { cute::prefetch_tma_descriptor(&tensor_map_output_0); cute::prefetch_tma_descriptor(&tensor_map_output_1); } @@ -250,256 +234,242 @@ sm103_fp8_block128_mega_moe_wgrad_impl( uint32_t stage_idx = 0; uint32_t phase = 0; - uint32_t output_iter = 0; - uint32_t tma_stage_idx = 0; - 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 * kRouteBlockM; - DG_DEVICE_ASSERT(count <= workspace.num_ring_tokens); - DG_DEVICE_ASSERT( - pool_row_offset + count <= workspace.num_max_pool_tokens); - - // Each dedicated wgrad kernel transports its one compact operand once - // per expert. All output tiles then reuse the local FP8 ring. - gather_compact_operand( - count, pool_row_offset, sf_ring_tokens, - token_src_metadata, sym_buffer, - symmetric_x, symmetric_x_sf, - symmetric_grad_y, symmetric_grad_y_sf, - symmetric_scores, - ring_operand, ring_operand_sf, ring_scores); - comm::grid_sync( - workspace, blockIdx.x, threadIdx.x, - []() { __syncthreads(); }); - const uint32_t num_k_blocks = - cute::max(1u, math::ceil_div(count, kBlockK)); - - if (warp_idx == 0) { - // A: (output-M, routes). W2 reads remotely transported dy and - // applies score before BF16 rounding; W13 reads local dpreact. - for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; - tile += kNumSMs) { - const uint32_t m_block = tile % kNumMBlocks; - #pragma unroll 1 - for (uint32_t k_block = 0; k_block < num_k_blocks; - ++k_block) { - storage.empty_barriers[stage_idx].wait(phase ^ 1); - for (uint32_t linear = lane_idx; - linear < kBlockM * kBlockK; linear += 32) { - const uint32_t row = linear / kBlockK; - const uint32_t k = linear - row * kBlockK; - const uint32_t route = k_block * kBlockK + k; - const uint32_t m = m_block * kBlockM + row; - float value = 0.0f; - if (route < count) { - if constexpr (kW2) { - const uint32_t packed = ring_operand_sf[ - (m / 128) * sf_ring_tokens + - transform_sf_row(route)]; - value = static_cast( - ring_operand[ - static_cast(route) * - kHidden + - m]) * - unpack_power2_scale(packed) * - ring_scores[route]; - } else { - const uint64_t full_row = - pool_row_offset + route; - const uint32_t packed = - full_grad_preact_sf[ - full_row * kFullABlocks + m / 128]; - value = static_cast( - full_grad_preact[ - full_row * kShapeM + m]) * - unpack_power2_scale(packed); - } + if (threadIdx.x < kAProducerThreads) { + // A is [output-M, routes]. One 128-thread group consumes one + // contiguous feature block for each route, then transposes into the + // UMMA MN-major shared-memory tile. W2 applies the exact FP32 score + // before the required BF16 rounding. + 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 feature = + m_block * kBlockM + threadIdx.x; + 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.empty_barriers[stage_idx].wait(phase ^ 1u); + #pragma unroll + for (uint32_t k = 0; k < kBlockK; ++k) { + const uint32_t route = k_block * kBlockK + k; + float value = 0.0f; + if (route < count) { + const uint64_t full_row = pool_row + route; + if constexpr (kW2) { + const uint32_t packed = full_grad_y_sf[ + full_row * kFullABlocks + + feature / 128]; + value = + static_cast(full_grad_y[ + full_row * kShapeM + feature]) * + unpack_power2_scale(packed) * + full_scores[full_row]; + } else { + const uint32_t packed = + full_grad_preact_sf[ + full_row * kFullABlocks + + feature / 128]; + value = + static_cast(full_grad_preact[ + full_row * kShapeM + feature]) * + unpack_power2_scale(packed); } - store_mn_swizzle128( - storage.smem_a[stage_idx], row, k, - bf16_t(value)); } - cutlass::arch::fence_view_async_shared(); - if (cute::elect_one_sync()) - storage.full_barriers[stage_idx].arrive(0u); - __syncwarp(); - stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; - phase ^= stage_idx == 0; + store_mn_swizzle128( + storage.smem_a[stage_idx], + threadIdx.x, k, bf16_t(value)); } + ptx::sync_aligned(128, 1); + cutlass::arch::fence_view_async_shared(); + if (threadIdx.x == 0) + storage.full_barriers[stage_idx].arrive(0u); + stage_idx = stage_idx == kStages - 1 + ? 0 + : stage_idx + 1; + phase ^= stage_idx == 0; } - } else if (warp_idx == 1) { - // B: (output-N, routes). W2 reads local h; W13 reads the compact - // x operand transported into the ring above. - for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; - tile += kNumSMs) { - const uint32_t n_block = tile / kNumMBlocks; - const uint32_t cta_n_base = - n_block * kBlockN + - cute::block_rank_in_cluster() * kLoadBlockN; - #pragma unroll 1 - for (uint32_t k_block = 0; k_block < num_k_blocks; - ++k_block) { - storage.empty_barriers[stage_idx].wait(phase ^ 1); - for (uint32_t linear = lane_idx; - linear < kLoadBlockN * kBlockK; linear += 32) { - const uint32_t row = linear / kBlockK; - const uint32_t k = linear - row * kBlockK; - const uint32_t route = k_block * kBlockK + k; - const uint32_t n = cta_n_base + row; - float value = 0.0f; - if (route < count) { - if constexpr (kW2) { - const uint64_t full_row = - pool_row_offset + route; - const uint32_t packed = full_h_sf[ - full_row * kFullBBlocks + n / 128]; - value = static_cast( - full_h[ - full_row * kShapeN + n]) * - unpack_power2_scale(packed); - } else { - const uint32_t packed = ring_operand_sf[ - (n / 128) * sf_ring_tokens + - transform_sf_row(route)]; - value = static_cast( - ring_operand[ - static_cast(route) * - kHidden + - n]) * - unpack_power2_scale(packed); - } + } + } else if ( + threadIdx.x >= kAProducerThreads && + threadIdx.x < kAProducerThreads + kBProducerThreads) { + // B is [output-N, routes]. The two CTAs load adjacent 64-feature + // halves, so all 128 N features for the cluster are coalesced. + Scheduler scheduler(expert_counts); + const uint32_t producer_lane = + threadIdx.x - kAProducerThreads; + 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 feature = + n_block * kBlockN + + cute::block_rank_in_cluster() * kLoadBlockN + + producer_lane; + 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.empty_barriers[stage_idx].wait(phase ^ 1u); + #pragma unroll + for (uint32_t k = 0; k < kBlockK; ++k) { + const uint32_t route = k_block * kBlockK + k; + float value = 0.0f; + if (route < count) { + const uint64_t full_row = pool_row + route; + if constexpr (kW2) { + const uint32_t packed = full_h_sf[ + full_row * kFullBBlocks + + feature / 128]; + value = + static_cast(full_h[ + full_row * kShapeN + feature]) * + unpack_power2_scale(packed); + } else { + const uint32_t packed = full_x_sf[ + full_row * kFullBBlocks + + feature / 128]; + value = + static_cast(full_x[ + full_row * kShapeN + feature]) * + unpack_power2_scale(packed); } - store_mn_swizzle128( - storage.smem_b[stage_idx], row, k, - bf16_t(value)); } - cutlass::arch::fence_view_async_shared(); - if (cute::elect_one_sync()) - storage.full_barriers[stage_idx].arrive(0u); - __syncwarp(); - stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; - phase ^= stage_idx == 0; + store_mn_swizzle128( + storage.smem_b[stage_idx], + producer_lane, k, bf16_t(value)); } + ptx::sync_aligned(64, 2); + cutlass::arch::fence_view_async_shared(); + if (producer_lane == 0) + storage.full_barriers[stage_idx].arrive(0u); + stage_idx = stage_idx == kStages - 1 + ? 0 + : stage_idx + 1; + phase ^= stage_idx == 0; } - } else if (warp_idx == 2 && leader_cta) { - for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; - tile += kNumSMs, ++output_iter) { - const uint32_t accum_stage = - output_iter % kNumEpilogueStages; - const uint32_t accum_phase = - (output_iter / kNumEpilogueStages) & 1u; - storage.tmem_empty_barriers[accum_stage].wait( - accum_phase ^ 1u); - ptx::tcgen05_after_thread_sync(); + } + } 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(); - #pragma unroll 1 - for (uint32_t k_block = 0; k_block < num_k_blocks; - ++k_block) { - 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()) { - #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); - } + 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.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; + } + __syncwarp(); + constexpr uint16_t kCTAMask = 3; + cutlass::arch::umma_arrive_multicast_2x1SM( + reinterpret_cast( + &storage.empty_barriers[stage_idx]), + kCTAMask); + if (k_block == num_k_blocks - 1) { cutlass::arch::umma_arrive_multicast_2x1SM( reinterpret_cast( - &storage.empty_barriers[stage_idx]), + &storage.tmem_full_barriers[accum_stage]), kCTAMask); - if (k_block == num_k_blocks - 1) { - cutlass::arch::umma_arrive_multicast_2x1SM( - reinterpret_cast( - &storage.tmem_full_barriers[accum_stage]), - kCTAMask); - } - __syncwarp(); - stage_idx = stage_idx == kStages - 1 ? 0 : stage_idx + 1; - phase ^= stage_idx == 0; } + __syncwarp(); + stage_idx = stage_idx == kStages - 1 + ? 0 + : stage_idx + 1; + phase ^= stage_idx == 0; } - } else if (warp_idx >= 4) { - const uint32_t epilogue_warp_idx = warp_idx - 4; - DG_TRAP_ONLY_DEVICE_ASSERT( - ptx::ld_shared(&storage.tmem_ptr) == 0); - for (uint32_t tile = blockIdx.x; tile < kNumTilesPerExpert; - tile += kNumSMs, ++output_iter) { - const uint32_t m_block = tile % kNumMBlocks; - const uint32_t n_block = tile / kNumMBlocks; - const uint32_t accum_stage = - output_iter % kNumEpilogueStages; - const uint32_t accum_phase = - (output_iter / kNumEpilogueStages) & 1u; - storage.tmem_full_barriers[accum_stage].wait(accum_phase); - ptx::tcgen05_after_thread_sync(); + } + } 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); + 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; } - if (epilogue_warp_idx == 0) - cute::tma_store_wait<0>(); - __syncwarp(); + 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); } - - // Do not overwrite the expert ring or start a new output wave until - // every CTA has completed all stores for this expert. - comm::grid_sync( - workspace, blockIdx.x, threadIdx.x, - []() { __syncthreads(); }); - pool_block_offset += math::ceil_div(count, kRouteBlockM); + if (epilogue_warp_idx == 0) + cute::tma_store_wait<0>(); + __syncwarp(); } comm::cluster_sync_with_relaxed_arrive(); - if (warp_idx == 3) + if (warp_idx == kControlWarp) cute::TMEM::Allocator2Sm().free(0, kNumTmemCols); #else if (blockIdx.x == 0 && threadIdx.x == 0) From 06246eb9870dd94d45d53725961089149dd019d1 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 16:30:53 +0800 Subject: [PATCH 18/29] feat: pipeline MegaMoE BF16 wgrad loads --- csrc/sm103_fp8_block128.cu | 54 ++- .../impls/sm100_fp8_fp4_mega_moe.cuh | 12 +- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 4 +- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 361 ++++++++++-------- 4 files changed, 261 insertions(+), 170 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 1e1d266d0b..ac1bc6c1bb 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -2321,11 +2321,52 @@ void launch_persistent_wgrad( const PersistentWorkspaceLayout& layout, const torch::Tensor& expert_counts ) { + constexpr int64_t shape_m = + kW2 ? kPersistentHidden : 2 * kPersistentIntermediate; + constexpr int64_t shape_n = + kW2 ? kPersistentIntermediate : kPersistentHidden; constexpr int64_t output_rows = kW2 ? kPersistentHidden : kPersistentIntermediate; constexpr int64_t output_columns = kW2 ? kPersistentIntermediate : kPersistentHidden; const int64_t local_experts = kPersistentExperts / kNumRanks; + const auto fp8_options = torch::TensorOptions() + .dtype(torch::kFloat8_e4m3fn) + .device(buffer.device()); + 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(); + auto full_a = torch::from_blob( + full_a_base, + {static_cast(layout.workspace.num_max_pool_tokens), shape_m}, + fp8_options); + auto full_b = torch::from_blob( + full_b_base, + {static_cast(layout.workspace.num_max_pool_tokens), shape_n}, + fp8_options); + const auto tensor_map_a = deep_gemm::make_tma_2d_desc( + full_a, + static_cast(shape_m), + static_cast(layout.workspace.num_max_pool_tokens), + deep_gemm::sm103_block128_wgrad::kBlockM, + deep_gemm::sm103_block128_wgrad::kBlockK, + static_cast(shape_m), 0); + const auto tensor_map_b = deep_gemm::make_tma_2d_desc( + full_b, + static_cast(shape_n), + static_cast(layout.workspace.num_max_pool_tokens), + deep_gemm::sm103_block128_wgrad::kLoadBlockN, + deep_gemm::sm103_block128_wgrad::kBlockK, + static_cast(shape_n), 0); const auto output_0_flat = output_0.view( {local_experts * output_rows, output_columns}); const auto output_1_flat = output_1.view( @@ -2372,17 +2413,10 @@ void launch_persistent_wgrad( &config, kernel, expert_counts.data_ptr(), layout.workspace.num_max_pool_tokens, - layout.backward_full_x.get_base_ptr(), - layout.backward_full_x_scales.get_base_ptr(), - layout.backward_full_grad_y - .get_base_ptr(), - layout.backward_full_grad_y_scales.get_base_ptr(), + full_a_sf, + full_b_sf, layout.backward_full_scores.get_base_ptr(), - layout.backward_full_h.get_base_ptr(), - layout.backward_full_h_scales.get_base_ptr(), - layout.backward_full_grad_preact - .get_base_ptr(), - layout.backward_full_grad_preact_scales.get_base_ptr(), + tensor_map_a, tensor_map_b, tensor_map_output_0, tensor_map_output_1)); C10_CUDA_KERNEL_LAUNCH_CHECK(); } 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 615633079a..5b357c67f3 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 @@ -1450,15 +1450,19 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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; convert that BF16 - // value back to FP32, multiply by the exact FP32 - // route score, and round once more to BF16 before + // written to shared memory above. The retained GLM + // path also casts the route score to the output + // dtype before the multiply, so round the score to + // BF16, convert both operands back to FP32 for the + // multiply, and round once more to BF16 before // remote combine. Moving this multiply into L1 is // not equivalent because it changes the E4M3 // requantization scale and rounding before W2. - const float route_weight = *l1_topk_weights_buffer + const float route_weight_fp32 = *l1_topk_weights_buffer .get_data_buffer(ring_m_idx + m_idx_in_block) .template get_base_ptr(); + const float route_weight = __bfloat162float( + __float2bfloat16_rn(route_weight_fp32)); auto* packed_bf16 = reinterpret_cast(&packed); #pragma unroll 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 index 9d53318bcd..2393b74f8a 100644 --- 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 @@ -6230,12 +6230,14 @@ sm103_fp8_block128_mega_moe_backward_impl( const float gate = static_cast( ring_bf16[static_cast(row) * kHidden + physical_gate]); + const float rounded_score = static_cast( + bf16_t(ring_scores[row])); const float dy_h = static_cast( ring_bf16[ static_cast(row) * kHidden + 2 * kIntermediate + col]) * - ring_scores[row]; + rounded_score; const float sigmoid = 1.0f / (1.0f + expf(-gate)); const float grad_value = gate_plane ? dy_h * up * sigmoid * 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 index ed7c758f31..e06a5cb5ef 100644 --- 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 @@ -7,6 +7,7 @@ #include #include #include +#include #include #include #include @@ -16,24 +17,25 @@ namespace deep_gemm::sm103_block128_wgrad { -// GLM-5.2 has enough routes per expert that one fixed 2-CTA family covers both -// target topologies. K is the route dimension. FP8 operands stay resident in -// the private expert-padded pool; the producer groups dequantize directly into -// pipelined BF16 shared-memory tiles consumed by native BF16 UMMA. +// The two CTAs form the same logical 256x256x64 production BF16 tile used by +// DeepGEMM's grouped large-M kernels. Each CTA owns 128 M rows and 128 N +// columns. FP8 route pools are TMA-staged, then converted in the load prologue +// into the exact MN-major BF16 shared-memory layout consumed by native UMMA. 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 = 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 = 6; -static constexpr uint32_t kAProducerThreads = 128; -static constexpr uint32_t kBProducerThreads = 64; -static constexpr uint32_t kMMAWarp = 6; -static constexpr uint32_t kControlWarp = 7; -static constexpr uint32_t kEpilogueFirstWarp = 8; +static constexpr uint32_t kStages = 3; +static constexpr uint32_t kTMAWarp = 0; +static constexpr uint32_t kMMAWarp = 1; +static constexpr uint32_t kConvertFirstWarp = 2; +static constexpr uint32_t kConvertThreads = 128; +static constexpr uint32_t kControlWarp = 6; +static constexpr uint32_t kEpilogueFirstWarp = 7; static constexpr uint32_t kEpilogueThreads = 128; static constexpr uint32_t kThreads = kEpilogueFirstWarp * 32 + kEpilogueThreads; @@ -42,7 +44,7 @@ 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 = 128; +static constexpr uint32_t kUMMAN = 256; static constexpr uint32_t kUMMAK = 16; static constexpr uint32_t kSwizzle = 128; static constexpr uint32_t kNumTmemAccumCols = @@ -57,17 +59,23 @@ 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][kBlockM * kBlockK]; - alignas(1024) bf16_t smem_b[kStages][kLoadBlockN * kBlockK]; - Barrier full_barriers[kStages]; - Barrier empty_barriers[kStages]; + alignas(1024) fp8_t raw_a[kStages][kBlockK * kBlockM]; + alignas(1024) fp8_t raw_b[kStages][kBlockK * kLoadBlockN]; + alignas(1024) bf16_t smem_a[kStages][kBlockK * kBlockM]; + alignas(1024) bf16_t smem_b[kStages][kBlockK * kLoadBlockN]; + Barrier tma_full_barriers[kStages]; + Barrier tma_empty_barriers[kStages]; + Barrier mma_full_barriers[kStages]; + Barrier mma_empty_barriers[kStages]; Barrier tmem_full_barriers[kNumEpilogueStages]; Barrier tmem_empty_barriers[kNumEpilogueStages]; uint32_t tmem_ptr; }; -DG_STATIC_ASSERT(kThreads == 384, "SM103 wgrad role layout changed"); -DG_STATIC_ASSERT(kNumTmemCols <= 512, "SM103 wgrad exceeds TMEM"); +DG_STATIC_ASSERT(kThreads == 352, "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 float unpack_power2_scale(const uint32_t packed) { const uint32_t exponent = packed & 0xffu; @@ -76,28 +84,70 @@ CUTLASS_DEVICE float unpack_power2_scale(const uint32_t packed) { return exponent == 0u ? 0x1p-127f : __uint_as_float(exponent << 23); } -template -CUTLASS_DEVICE void store_mn_swizzle128( +CUTLASS_DEVICE float round_score_to_bf16(const float score) { + return static_cast(bf16_t(score)); +} + +// Address one 16-byte bank group in the TMA swizzle-128 layout for an +// MN-major [inner-MN, outer-K] tile. TMA splits an inner dimension wider than +// 64 BF16 values into consecutive 64-value atoms. +template +CUTLASS_DEVICE uint8_t* get_bf16_mn_bank_group( bf16_t* base, - const uint32_t row, - const uint32_t k, - const bf16_t value) { - DG_STATIC_ASSERT(kRows == 64 || kRows == 128, - "invalid BF16 SMEM rows"); - const uint32_t row_in_atom = row & 7u; - const uint32_t col_byte = k * sizeof(bf16_t); + const uint32_t inner_mn, + const uint32_t outer_k) { + constexpr uint32_t kBankGroupBytes = 16; + constexpr uint32_t kInnerPerAtom = kSwizzle / sizeof(bf16_t); + DG_STATIC_ASSERT(kInnerMN % kInnerPerAtom == 0, + "MN dimension must contain whole swizzle atoms"); + DG_STATIC_ASSERT(kOuterK % 8 == 0, + "K dimension must contain whole swizzle rows"); + const uint32_t atom = inner_mn / kInnerPerAtom; + const uint32_t inner_in_atom = inner_mn % kInnerPerAtom; + const uint32_t row = outer_k & 7u; + const uint32_t inner_byte = inner_in_atom * sizeof(bf16_t); const uint32_t byte_offset = - (row >> 3) * 8u * kSwizzle + row_in_atom * kSwizzle + - ((col_byte >> 4) ^ row_in_atom) * 16u + (col_byte & 15u); - *reinterpret_cast( - reinterpret_cast(base) + byte_offset) = value; + atom * kOuterK * kSwizzle + + (outer_k >> 3) * 8u * kSwizzle + + row * kSwizzle + + ((inner_byte >> 4) ^ row) * kBankGroupBytes + + (inner_byte & (kBankGroupBytes - 1)); + return reinterpret_cast(base) + byte_offset; +} + +template +CUTLASS_DEVICE void convert_and_store_eight( + const fp8_t* source, + bf16_t* destination, + const uint32_t inner_mn, + const uint32_t outer_k, + const float dequant_scale, + const float post_scale, + const bool valid) { + uint4 packed{}; + auto* values = reinterpret_cast(&packed); + #pragma unroll + for (uint32_t i = 0; i < 8; ++i) { + if (valid) { + // "BF16-semantics" means the FP8+power-of-two value first becomes + // the BF16 operand represented by the private pool. W2 then + // applies the BF16-rounded route score and rounds to BF16 again. + const float dequantized = static_cast( + bf16_t(static_cast(source[i]) * dequant_scale)); + values[i] = bf16_t(dequantized * post_scale); + } else { + values[i] = bf16_t(0.0f); + } + } + *reinterpret_cast( + get_bf16_mn_bank_group( + destination, inner_mn, outer_k)) = packed; } -// Each physical CTA starts at its own output tile and advances by the fixed SM -// count. Because every expert matrix has an even number of M tiles, adjacent -// CTAs in a 2-CTA cluster always address adjacent M halves of the same expert/N -// tile. Pool offsets advance only when the monotonic tile stream crosses an -// expert boundary; there is no global queue, host metadata, or per-expert sync. +// 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 = @@ -108,6 +158,9 @@ struct WgradTileScheduler { 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 = @@ -119,7 +172,14 @@ struct WgradTileScheduler { uint32_t cached_pool_row = 0; CUTLASS_DEVICE explicit WgradTileScheduler(const int* counts) - : expert_counts(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, @@ -142,8 +202,12 @@ struct WgradTileScheduler { } count = static_cast(__ldg(expert_counts + expert)); pool_row = cached_pool_row; - m_block = expert_tile % kNumMBlocks; - n_block = expert_tile / kNumMBlocks; + 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; } @@ -154,15 +218,11 @@ 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_x, - const uint32_t* full_x_sf, - const fp8_t* full_grad_y, - const uint32_t* full_grad_y_sf, + const uint32_t* full_a_sf, + const uint32_t* full_b_sf, const float* full_scores, - const fp8_t* full_h, - const uint32_t* full_h_sf, - const fp8_t* full_grad_preact, - const uint32_t* full_grad_preact_sf, + const __grid_constant__ cute::TmaDescriptor tensor_map_a, + const __grid_constant__ cute::TmaDescriptor tensor_map_b, 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 @@ -171,24 +231,33 @@ sm103_fp8_block128_mega_moe_wgrad_impl( constexpr uint32_t kShapeN = Scheduler::kShapeN; constexpr uint32_t kFullABlocks = kShapeM / 128; constexpr uint32_t kFullBBlocks = kShapeN / 128; - DG_STATIC_ASSERT(Scheduler::kNumMBlocks % 2 == 0, - "2-CTA wgrad requires paired M tiles"); - DG_STATIC_ASSERT(kNumSMs % 2 == 0, - "2-CTA wgrad requires an even SM count"); + constexpr uint32_t kRawABytes = kBlockK * kBlockM * sizeof(fp8_t); + constexpr uint32_t kRawBBytes = + kBlockK * kLoadBlockN * sizeof(fp8_t); 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 == kTMAWarp) { + cute::prefetch_tma_descriptor(&tensor_map_a); + cute::prefetch_tma_descriptor(&tensor_map_b); + 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) { - // A and B producer groups in both CTAs complete each stage. - storage.full_barriers[i].init(4); - storage.empty_barriers[i].init(1); + storage.tma_full_barriers[i].init(1); + storage.tma_empty_barriers[i].init(1); + // Both CTAs publish their converted halves to CTA 0. + storage.mma_full_barriers[i].init(2); + storage.mma_empty_barriers[i].init(1); } #pragma unroll for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { @@ -204,11 +273,6 @@ sm103_fp8_block128_mega_moe_wgrad_impl( kNumTmemCols, &storage.tmem_ptr); comm::cluster_sync_with_relaxed_arrive(); - if (warp_idx == kControlWarp) { - cute::prefetch_tma_descriptor(&tensor_map_output_0); - cute::prefetch_tma_descriptor(&tensor_map_output_1); - } - auto instr_desc = cute::UMMA::make_instr_desc< bf16_t, bf16_t, float, kUMMAM, kUMMAN, @@ -234,121 +298,111 @@ sm103_fp8_block128_mega_moe_wgrad_impl( 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 (threadIdx.x < kAProducerThreads) { - // A is [output-M, routes]. One 128-thread group consumes one - // contiguous feature block for each route, then transposes into the - // UMMA MN-major shared-memory tile. W2 applies the exact FP32 score - // before the required BF16 rounding. + if (warp_idx == kTMAWarp && cute::elect_one_sync()) { + // The production load warp issues two rectangular TMA transactions per + // K stage. Raw tiles are row-major [route-K, feature-MN]. 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 feature = - m_block * kBlockM + threadIdx.x; 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.empty_barriers[stage_idx].wait(phase ^ 1u); - #pragma unroll - for (uint32_t k = 0; k < kBlockK; ++k) { - const uint32_t route = k_block * kBlockK + k; - float value = 0.0f; - if (route < count) { - const uint64_t full_row = pool_row + route; - if constexpr (kW2) { - const uint32_t packed = full_grad_y_sf[ - full_row * kFullABlocks + - feature / 128]; - value = - static_cast(full_grad_y[ - full_row * kShapeM + feature]) * - unpack_power2_scale(packed) * - full_scores[full_row]; - } else { - const uint32_t packed = - full_grad_preact_sf[ - full_row * kFullABlocks + - feature / 128]; - value = - static_cast(full_grad_preact[ - full_row * kShapeM + feature]) * - unpack_power2_scale(packed); - } - } - store_mn_swizzle128( - storage.smem_a[stage_idx], - threadIdx.x, k, bf16_t(value)); - } - ptx::sync_aligned(128, 1); - cutlass::arch::fence_view_async_shared(); - if (threadIdx.x == 0) - storage.full_barriers[stage_idx].arrive(0u); - stage_idx = stage_idx == kStages - 1 - ? 0 - : stage_idx + 1; - phase ^= stage_idx == 0; + storage.tma_empty_barriers[stage_idx].wait(phase ^ 1u); + const uint32_t route = pool_row + k_block * kBlockK; + const uint32_t a_feature = m_block * kBlockM; + const uint32_t b_feature = + n_block * kBlockN + cta_rank * kLoadBlockN; + tma::copy( + &tensor_map_a, + &storage.tma_full_barriers[stage_idx], + storage.raw_a[stage_idx], + a_feature, route); + tma::copy( + &tensor_map_b, + &storage.tma_full_barriers[stage_idx], + storage.raw_b[stage_idx], + b_feature, route); + storage.tma_full_barriers[stage_idx] + .arrive_and_expect_tx(kRawABytes + kRawBBytes); + advance_pipeline(); } } } else if ( - threadIdx.x >= kAProducerThreads && - threadIdx.x < kAProducerThreads + kBProducerThreads) { - // B is [output-N, routes]. The two CTAs load adjacent 64-feature - // halves, so all 128 N features for the cluster are coalesced. + warp_idx >= kConvertFirstWarp && + warp_idx < kConvertFirstWarp + kConvertThreads / 32) { + // Two converter threads own each route. Each thread loads four aligned + // FP8x16 vectors per operand, reuses one row scale, and emits eight + // aligned BF16x8 bank groups directly into the UMMA swizzle. Scheduler scheduler(expert_counts); - const uint32_t producer_lane = - threadIdx.x - kAProducerThreads; + const uint32_t convert_thread = + threadIdx.x - kConvertFirstWarp * 32; + const uint32_t route_in_k = convert_thread / 2; + const uint32_t half = convert_thread & 1u; 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 feature = - n_block * kBlockN + - cute::block_rank_in_cluster() * kLoadBlockN + - producer_lane; 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.empty_barriers[stage_idx].wait(phase ^ 1u); + storage.tma_full_barriers[stage_idx].wait(phase); + storage.mma_empty_barriers[stage_idx].wait(phase ^ 1u); + const uint32_t route = k_block * kBlockK + route_in_k; + const bool valid = route < count; + const uint64_t full_row = pool_row + route; + float a_scale = 0.0f; + float b_scale = 0.0f; + float a_post_scale = 1.0f; + if (valid) { + a_scale = unpack_power2_scale(full_a_sf[ + full_row * kFullABlocks + m_block]); + b_scale = unpack_power2_scale(full_b_sf[ + full_row * kFullBBlocks + + n_block * (kBlockN / 128) + cta_rank]); + if constexpr (kW2) + a_post_scale = round_score_to_bf16( + full_scores[full_row]); + } #pragma unroll - for (uint32_t k = 0; k < kBlockK; ++k) { - const uint32_t route = k_block * kBlockK + k; - float value = 0.0f; - if (route < count) { - const uint64_t full_row = pool_row + route; - if constexpr (kW2) { - const uint32_t packed = full_h_sf[ - full_row * kFullBBlocks + - feature / 128]; - value = - static_cast(full_h[ - full_row * kShapeN + feature]) * - unpack_power2_scale(packed); - } else { - const uint32_t packed = full_x_sf[ - full_row * kFullBBlocks + - feature / 128]; - value = - static_cast(full_x[ - full_row * kShapeN + feature]) * - unpack_power2_scale(packed); - } - } - store_mn_swizzle128( - storage.smem_b[stage_idx], - producer_lane, k, bf16_t(value)); + for (uint32_t chunk = 0; chunk < 4; ++chunk) { + const uint32_t inner = (half * 4 + chunk) * 16; + const auto* raw_a = storage.raw_a[stage_idx] + + route_in_k * kBlockM + inner; + const auto* raw_b = storage.raw_b[stage_idx] + + route_in_k * kLoadBlockN + inner; + convert_and_store_eight( + raw_a, storage.smem_a[stage_idx], + inner, route_in_k, + a_scale, a_post_scale, valid); + convert_and_store_eight( + raw_a + 8, storage.smem_a[stage_idx], + inner + 8, route_in_k, + a_scale, a_post_scale, valid); + convert_and_store_eight( + raw_b, storage.smem_b[stage_idx], + inner, route_in_k, + b_scale, 1.0f, valid); + convert_and_store_eight( + raw_b + 8, storage.smem_b[stage_idx], + inner + 8, route_in_k, + b_scale, 1.0f, valid); } - ptx::sync_aligned(64, 2); + ptx::sync_aligned(128, 1); cutlass::arch::fence_view_async_shared(); - if (producer_lane == 0) - storage.full_barriers[stage_idx].arrive(0u); - stage_idx = stage_idx == kStages - 1 - ? 0 - : stage_idx + 1; - phase ^= stage_idx == 0; + if (convert_thread == 0) { + storage.tma_empty_barriers[stage_idx].arrive(); + storage.mma_full_barriers[stage_idx].arrive(0u); + } + advance_pipeline(); } } } else if (warp_idx == kMMAWarp && leader_cta) { @@ -370,7 +424,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( cute::max(1u, math::ceil_div(count, kBlockK)); for (uint32_t k_block = 0; k_block < num_k_blocks; ++k_block) { - storage.full_barriers[stage_idx].wait(phase); + storage.mma_full_barriers[stage_idx].wait(phase); ptx::tcgen05_after_thread_sync(); const uint32_t a_base = ptx::exchange(a_desc_lo, stage_idx); @@ -398,7 +452,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( constexpr uint16_t kCTAMask = 3; cutlass::arch::umma_arrive_multicast_2x1SM( reinterpret_cast( - &storage.empty_barriers[stage_idx]), + &storage.mma_empty_barriers[stage_idx]), kCTAMask); if (k_block == num_k_blocks - 1) { cutlass::arch::umma_arrive_multicast_2x1SM( @@ -407,10 +461,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( kCTAMask); } __syncwarp(); - stage_idx = stage_idx == kStages - 1 - ? 0 - : stage_idx + 1; - phase ^= stage_idx == 0; + advance_pipeline(); } } } else if ( From 4d10bc922b18177826ce8966b4fa73f261307eae Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 17:13:35 +0800 Subject: [PATCH 19/29] perf: pack MegaMoE wgrad dequantization --- .../impls/sm100_fp8_fp4_mega_moe.cuh | 17 +-- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 4 +- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 144 ++++++++++++------ 3 files changed, 102 insertions(+), 63 deletions(-) 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 5b357c67f3..b1960ec760 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 @@ -1450,19 +1450,14 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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. The retained GLM - // path also casts the route score to the output - // dtype before the multiply, so round the score to - // BF16, convert both operands back to FP32 for the - // multiply, and round once more to BF16 before - // remote combine. Moving this multiply into L1 is - // not equivalent because it changes the E4M3 - // requantization scale and rounding before W2. - const float route_weight_fp32 = *l1_topk_weights_buffer + // 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. + const float route_weight = *l1_topk_weights_buffer .get_data_buffer(ring_m_idx + m_idx_in_block) .template get_base_ptr(); - const float route_weight = __bfloat162float( - __float2bfloat16_rn(route_weight_fp32)); auto* packed_bf16 = reinterpret_cast(&packed); #pragma unroll 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 index 2393b74f8a..9d53318bcd 100644 --- 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 @@ -6230,14 +6230,12 @@ sm103_fp8_block128_mega_moe_backward_impl( const float gate = static_cast( ring_bf16[static_cast(row) * kHidden + physical_gate]); - const float rounded_score = static_cast( - bf16_t(ring_scores[row])); const float dy_h = static_cast( ring_bf16[ static_cast(row) * kHidden + 2 * kIntermediate + col]) * - rounded_score; + ring_scores[row]; const float sigmoid = 1.0f / (1.0f + expf(-gate)); const float grad_value = gate_plane ? dy_h * up * sigmoid * 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 index e06a5cb5ef..b404415061 100644 --- 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 @@ -33,9 +33,9 @@ static constexpr uint32_t kStages = 3; static constexpr uint32_t kTMAWarp = 0; static constexpr uint32_t kMMAWarp = 1; static constexpr uint32_t kConvertFirstWarp = 2; -static constexpr uint32_t kConvertThreads = 128; -static constexpr uint32_t kControlWarp = 6; -static constexpr uint32_t kEpilogueFirstWarp = 7; +static constexpr uint32_t kConvertThreads = 256; +static constexpr uint32_t kControlWarp = 10; +static constexpr uint32_t kEpilogueFirstWarp = 11; static constexpr uint32_t kEpilogueThreads = 128; static constexpr uint32_t kThreads = kEpilogueFirstWarp * 32 + kEpilogueThreads; @@ -72,20 +72,60 @@ struct alignas(1024) SharedStorage { uint32_t tmem_ptr; }; -DG_STATIC_ASSERT(kThreads == 352, "SM103 wgrad role layout changed"); +DG_STATIC_ASSERT(kThreads == 480, "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 float unpack_power2_scale(const uint32_t packed) { - const uint32_t exponent = packed & 0xffu; - // UE8M0 code zero denotes 2^-127. Construct it explicitly because - // shifting zero into an IEEE exponent field would produce zero. - return exponent == 0u ? 0x1p-127f : __uint_as_float(exponent << 23); +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 float round_score_to_bf16(const float score) { - return static_cast(bf16_t(score)); +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); } // Address one 16-byte bank group in the TMA swizzle-128 layout for an @@ -115,28 +155,34 @@ CUTLASS_DEVICE uint8_t* get_bf16_mn_bank_group( return reinterpret_cast(base) + byte_offset; } -template +template CUTLASS_DEVICE void convert_and_store_eight( const fp8_t* source, bf16_t* destination, const uint32_t inner_mn, const uint32_t outer_k, - const float dequant_scale, + const uint32_t scale_exponent, const float post_scale, const bool valid) { uint4 packed{}; - auto* values = reinterpret_cast(&packed); - #pragma unroll - for (uint32_t i = 0; i < 8; ++i) { - if (valid) { - // "BF16-semantics" means the FP8+power-of-two value first becomes - // the BF16 operand represented by the private pool. W2 then - // applies the BF16-rounded route score and rounds to BF16 again. - const float dequantized = static_cast( - bf16_t(static_cast(source[i]) * dequant_scale)); - values[i] = bf16_t(dequantized * post_scale); - } else { - values[i] = bf16_t(0.0f); + if (valid) { + const uint2 raw = *reinterpret_cast(source); + const uint32_t raw_words[2] = {raw.x, raw.y}; + auto* output_pairs = reinterpret_cast(&packed); + #pragma unroll + for (uint32_t pair = 0; pair < 4; ++pair) { + const uint16_t fp8x2 = static_cast( + raw_words[pair / 2] >> ((pair & 1u) * 16)); + uint32_t bf16x2 = convert_fp8x2_power2_to_bf16x2( + fp8x2, scale_exponent); + if constexpr (kApplyPostScale) { + const auto dequantized = __bfloat1622float2( + *reinterpret_cast(&bf16x2)); + const auto scaled = __float22bfloat162_rn(__fmul2_rn( + dequantized, {post_scale, post_scale})); + bf16x2 = *reinterpret_cast(&scaled); + } + output_pairs[pair] = bf16x2; } } *reinterpret_cast( @@ -338,14 +384,15 @@ sm103_fp8_block128_mega_moe_wgrad_impl( } else if ( warp_idx >= kConvertFirstWarp && warp_idx < kConvertFirstWarp + kConvertThreads / 32) { - // Two converter threads own each route. Each thread loads four aligned - // FP8x16 vectors per operand, reuses one row scale, and emits eight - // aligned BF16x8 bank groups directly into the UMMA swizzle. + // Four converter threads own each route. Each thread loads two aligned + // FP8x16 vectors per operand, converts packed E4M3x2 values, folds the + // power-of-two exponent directly into BF16, and emits four aligned + // BF16x8 bank groups per operand into the UMMA swizzle. Scheduler scheduler(expert_counts); const uint32_t convert_thread = threadIdx.x - kConvertFirstWarp * 32; - const uint32_t route_in_k = convert_thread / 2; - const uint32_t half = convert_thread & 1u; + const uint32_t route_in_k = convert_thread / 4; + const uint32_t quarter = convert_thread & 3u; uint32_t expert, count, pool_row, m_block, n_block; while (scheduler.get_next( expert, count, pool_row, m_block, n_block)) { @@ -359,44 +406,43 @@ sm103_fp8_block128_mega_moe_wgrad_impl( const uint32_t route = k_block * kBlockK + route_in_k; const bool valid = route < count; const uint64_t full_row = pool_row + route; - float a_scale = 0.0f; - float b_scale = 0.0f; + uint32_t a_scale_exponent = 0u; + uint32_t b_scale_exponent = 0u; float a_post_scale = 1.0f; if (valid) { - a_scale = unpack_power2_scale(full_a_sf[ - full_row * kFullABlocks + m_block]); - b_scale = unpack_power2_scale(full_b_sf[ + a_scale_exponent = full_a_sf[ + full_row * kFullABlocks + m_block] & 0xffu; + b_scale_exponent = full_b_sf[ full_row * kFullBBlocks + - n_block * (kBlockN / 128) + cta_rank]); + n_block * (kBlockN / 128) + cta_rank] & 0xffu; if constexpr (kW2) - a_post_scale = round_score_to_bf16( - full_scores[full_row]); + a_post_scale = full_scores[full_row]; } #pragma unroll - for (uint32_t chunk = 0; chunk < 4; ++chunk) { - const uint32_t inner = (half * 4 + chunk) * 16; + for (uint32_t chunk = 0; chunk < 2; ++chunk) { + const uint32_t inner = (quarter * 2 + chunk) * 16; const auto* raw_a = storage.raw_a[stage_idx] + route_in_k * kBlockM + inner; const auto* raw_b = storage.raw_b[stage_idx] + route_in_k * kLoadBlockN + inner; - convert_and_store_eight( + convert_and_store_eight( raw_a, storage.smem_a[stage_idx], inner, route_in_k, - a_scale, a_post_scale, valid); - convert_and_store_eight( + a_scale_exponent, a_post_scale, valid); + convert_and_store_eight( raw_a + 8, storage.smem_a[stage_idx], inner + 8, route_in_k, - a_scale, a_post_scale, valid); - convert_and_store_eight( + a_scale_exponent, a_post_scale, valid); + convert_and_store_eight( raw_b, storage.smem_b[stage_idx], inner, route_in_k, - b_scale, 1.0f, valid); - convert_and_store_eight( + b_scale_exponent, 1.0f, valid); + convert_and_store_eight( raw_b + 8, storage.smem_b[stage_idx], inner + 8, route_in_k, - b_scale, 1.0f, valid); + b_scale_exponent, 1.0f, valid); } - ptx::sync_aligned(128, 1); + ptx::sync_aligned(256, 1); cutlass::arch::fence_view_async_shared(); if (convert_thread == 0) { storage.tma_empty_barriers[stage_idx].arrive(); From 31344a27dd9d3dc415f3d5bdd15a67ad641038b8 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 18:00:33 +0800 Subject: [PATCH 20/29] perf: align persistent MegaMoE block128 pipeline --- csrc/sm103_fp8_block128.cu | 76 +++-- .../impls/sm100_fp8_fp4_mega_moe.cuh | 109 ++++++- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 276 +++++++++--------- 3 files changed, 294 insertions(+), 167 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index ac1bc6c1bb..db20d87bd5 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -123,15 +123,19 @@ constexpr uint32_t kPersistentThreads = 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: 2 warpgroups * 4 warps * 32 lanes * sizeof(uint2). +// 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; + kPersistentUpstreamSmemBytes + kPersistentW13PairExchangeBytes + + kPersistentBlock128ScaleExchangeBytes; constexpr uint32_t kWorkspaceAlignment = deep_gemm::layout::kLCMCandidateBlockM; @@ -188,6 +192,8 @@ struct PersistentWorkspaceLayout { 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 backward_wgrad_bf16_narrow; + deep_gemm::layout::Buffer backward_wgrad_bf16_wide; PersistentWorkspaceLayout( void* base, @@ -351,11 +357,29 @@ struct PersistentWorkspaceLayout { 2 * kPersistentIntermediate / 32), 1, workspace.num_max_pool_tokens, - backward_full_grad_preact.get_end_ptr()) {} + backward_full_grad_preact.get_end_ptr()), + // The two dedicated wgrad kernels reuse these private operands. W13 + // maps grad_preact -> narrow and x -> wide; W2 maps grad_y -> wide and + // h -> narrow. One extra K tile is a permanent zero source for empty + // experts. Capacity remains a once-derived context/CP consequence. + backward_wgrad_bf16_narrow( + deep_gemm::layout::Data( + 2 * kPersistentIntermediate * sizeof(__nv_bfloat16)), + 1, + workspace.num_max_pool_tokens + + deep_gemm::sm103_block128_wgrad::kBlockK, + backward_full_grad_preact_scales.get_end_ptr()), + backward_wgrad_bf16_wide( + deep_gemm::layout::Data( + kPersistentHidden * sizeof(__nv_bfloat16)), + 1, + workspace.num_max_pool_tokens + + deep_gemm::sm103_block128_wgrad::kBlockK, + backward_wgrad_bf16_narrow.get_end_ptr()) {} int64_t num_bytes() const { return reinterpret_cast( - backward_full_grad_preact_scales.get_end_ptr()) - + backward_wgrad_bf16_wide.get_end_ptr()) - reinterpret_cast(workspace.base); } }; @@ -2330,8 +2354,8 @@ void launch_persistent_wgrad( constexpr int64_t output_columns = kW2 ? kPersistentIntermediate : kPersistentHidden; const int64_t local_experts = kPersistentExperts / kNumRanks; - const auto fp8_options = torch::TensorOptions() - .dtype(torch::kFloat8_e4m3fn) + const auto bf16_options = torch::TensorOptions() + .dtype(torch::kBFloat16) .device(buffer.device()); void* full_a_base = kW2 ? layout.backward_full_grad_y.base @@ -2339,34 +2363,39 @@ void launch_persistent_wgrad( void* full_b_base = kW2 ? layout.backward_full_h.base : layout.backward_full_x.base; + void* cached_a_base = kW2 + ? layout.backward_wgrad_bf16_wide.base + : layout.backward_wgrad_bf16_narrow.base; + void* cached_b_base = kW2 + ? layout.backward_wgrad_bf16_narrow.base + : layout.backward_wgrad_bf16_wide.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(); - auto full_a = torch::from_blob( - full_a_base, - {static_cast(layout.workspace.num_max_pool_tokens), shape_m}, - fp8_options); - auto full_b = torch::from_blob( - full_b_base, - {static_cast(layout.workspace.num_max_pool_tokens), shape_n}, - fp8_options); + const int64_t cached_rows = + static_cast(layout.workspace.num_max_pool_tokens) + + deep_gemm::sm103_block128_wgrad::kBlockK; + auto cached_a = torch::from_blob( + cached_a_base, {cached_rows, shape_m}, bf16_options); + auto cached_b = torch::from_blob( + cached_b_base, {cached_rows, shape_n}, bf16_options); const auto tensor_map_a = deep_gemm::make_tma_2d_desc( - full_a, + cached_a, static_cast(shape_m), - static_cast(layout.workspace.num_max_pool_tokens), + static_cast(cached_rows), deep_gemm::sm103_block128_wgrad::kBlockM, deep_gemm::sm103_block128_wgrad::kBlockK, - static_cast(shape_m), 0); + static_cast(shape_m), 128); const auto tensor_map_b = deep_gemm::make_tma_2d_desc( - full_b, + cached_b, static_cast(shape_n), - static_cast(layout.workspace.num_max_pool_tokens), + static_cast(cached_rows), deep_gemm::sm103_block128_wgrad::kLoadBlockN, deep_gemm::sm103_block128_wgrad::kBlockK, - static_cast(shape_n), 0); + static_cast(shape_n), 128); const auto output_0_flat = output_0.view( {local_experts * output_rows, output_columns}); const auto output_1_flat = output_1.view( @@ -2413,9 +2442,14 @@ void launch_persistent_wgrad( &config, kernel, expert_counts.data_ptr(), layout.workspace.num_max_pool_tokens, + layout.workspace, + static_cast(full_a_base), + static_cast(full_b_base), full_a_sf, full_b_sf, layout.backward_full_scores.get_base_ptr(), + static_cast(cached_a_base), + static_cast(cached_b_base), tensor_map_a, tensor_map_b, tensor_map_output_0, tensor_map_output_1)); C10_CUDA_KERNEL_LAUNCH_CHECK(); 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 b1960ec760..c5640f5299 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 @@ -20,6 +20,20 @@ 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 kHidden, uint32_t kIntermediateHidden, uint32_t kNumExperts, uint32_t kNumTopk, @@ -227,6 +241,12 @@ 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 @@ -241,6 +261,8 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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); @@ -299,6 +321,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) { @@ -1029,6 +1056,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, @@ -1255,16 +1283,85 @@ 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]); + } + } + __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 uint32_t amax_partner = epilogue_warp_idx ^ - (kFP8Block128Weights ? 2u : 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); + 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; 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 index b404415061..31ed379477 100644 --- 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 @@ -18,9 +18,12 @@ 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. Each CTA owns 128 M rows and 128 N -// columns. FP8 route pools are TMA-staged, then converted in the load prologue -// into the exact MN-major BF16 shared-memory layout consumed by native UMMA. +// DeepGEMM's grouped large-M kernels. Each dedicated wgrad kernel first converts +// its FP8 route operands exactly once into private BF16 backing inside this same +// persistent launch, grid-fences internally, and then runs the native BF16 TMA +// / UMMA pipeline. Dequantization is therefore part of the dedicated kernel's +// load prologue rather than repeated for every output tile or composed as a +// separate kernel. static constexpr uint32_t kHidden = 6144; static constexpr uint32_t kIntermediate = 2048; static constexpr uint32_t kGlobalExperts = 256; @@ -32,10 +35,9 @@ static constexpr uint32_t kLoadBlockN = kBlockN / 2; static constexpr uint32_t kStages = 3; static constexpr uint32_t kTMAWarp = 0; static constexpr uint32_t kMMAWarp = 1; -static constexpr uint32_t kConvertFirstWarp = 2; -static constexpr uint32_t kConvertThreads = 256; -static constexpr uint32_t kControlWarp = 10; -static constexpr uint32_t kEpilogueFirstWarp = 11; +static constexpr uint32_t kReadyWarp = 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; @@ -59,20 +61,17 @@ using Barrier = cutlass::arch::ClusterTransactionBarrier; struct alignas(1024) SharedStorage { alignas(1024) bf16_t smem_cd[kNumTMAStoreStages] [kStoreBlockM * kStoreBlockN]; - alignas(1024) fp8_t raw_a[kStages][kBlockK * kBlockM]; - alignas(1024) fp8_t raw_b[kStages][kBlockK * kLoadBlockN]; alignas(1024) bf16_t smem_a[kStages][kBlockK * kBlockM]; alignas(1024) bf16_t smem_b[kStages][kBlockK * kLoadBlockN]; Barrier tma_full_barriers[kStages]; Barrier tma_empty_barriers[kStages]; Barrier mma_full_barriers[kStages]; - Barrier mma_empty_barriers[kStages]; Barrier tmem_full_barriers[kNumEpilogueStages]; Barrier tmem_empty_barriers[kNumEpilogueStages]; uint32_t tmem_ptr; }; -DG_STATIC_ASSERT(kThreads == 480, "SM103 wgrad role layout changed"); +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"); @@ -128,66 +127,95 @@ CUTLASS_DEVICE uint32_t convert_fp8x2_power2_to_bf16x2( return lo | (hi << 16); } -// Address one 16-byte bank group in the TMA swizzle-128 layout for an -// MN-major [inner-MN, outer-K] tile. TMA splits an inner dimension wider than -// 64 BF16 values into consecutive 64-value atoms. -template -CUTLASS_DEVICE uint8_t* get_bf16_mn_bank_group( - bf16_t* base, - const uint32_t inner_mn, - const uint32_t outer_k) { - constexpr uint32_t kBankGroupBytes = 16; - constexpr uint32_t kInnerPerAtom = kSwizzle / sizeof(bf16_t); - DG_STATIC_ASSERT(kInnerMN % kInnerPerAtom == 0, - "MN dimension must contain whole swizzle atoms"); - DG_STATIC_ASSERT(kOuterK % 8 == 0, - "K dimension must contain whole swizzle rows"); - const uint32_t atom = inner_mn / kInnerPerAtom; - const uint32_t inner_in_atom = inner_mn % kInnerPerAtom; - const uint32_t row = outer_k & 7u; - const uint32_t inner_byte = inner_in_atom * sizeof(bf16_t); - const uint32_t byte_offset = - atom * kOuterK * kSwizzle + - (outer_k >> 3) * 8u * kSwizzle + - row * kSwizzle + - ((inner_byte >> 4) ^ row) * kBankGroupBytes + - (inner_byte & (kBankGroupBytes - 1)); - return reinterpret_cast(base) + byte_offset; -} - -template -CUTLASS_DEVICE void convert_and_store_eight( +template < + uint32_t kShape, uint32_t kLocalExperts, + uint32_t kNumSMs, uint32_t kNumThreads, + bool kApplyPostScale> +CUTLASS_DEVICE void dequantize_route_pool_once( + const int* expert_counts, + const uint32_t max_pool_tokens, const fp8_t* source, - bf16_t* destination, - const uint32_t inner_mn, - const uint32_t outer_k, - const uint32_t scale_exponent, - const float post_scale, - const bool valid) { - uint4 packed{}; - if (valid) { - const uint2 raw = *reinterpret_cast(source); - const uint32_t raw_words[2] = {raw.x, raw.y}; - auto* output_pairs = reinterpret_cast(&packed); - #pragma unroll - for (uint32_t pair = 0; pair < 4; ++pair) { - const uint16_t fp8x2 = static_cast( - raw_words[pair / 2] >> ((pair & 1u) * 16)); - uint32_t bf16x2 = convert_fp8x2_power2_to_bf16x2( - fp8x2, scale_exponent); - if constexpr (kApplyPostScale) { - const auto dequantized = __bfloat1622float2( - *reinterpret_cast(&bf16x2)); - const auto scaled = __float22bfloat162_rn(__fmul2_rn( - dequantized, {post_scale, post_scale})); - bf16x2 = *reinterpret_cast(&scaled); + const uint32_t* scales, + const float* scores, + bf16_t* destination) { + constexpr uint32_t kValuesPerVector = 8; + constexpr uint32_t kVectorsPerRow = kShape / kValuesPerVector; + constexpr uint32_t kScaleBlocksPerRow = kShape / 128; + DG_STATIC_ASSERT(kShape % 128 == 0, + "wgrad dequant shape must be block128 aligned"); + + const uint64_t global_thread = + static_cast(blockIdx.x) * kNumThreads + threadIdx.x; + constexpr uint64_t kGridThreads = + static_cast(kNumSMs) * kNumThreads; + uint32_t pool_row = 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 padded_count = + math::ceil_div(count, kRouteBlockM) * kRouteBlockM; + const uint64_t num_vectors = + static_cast(padded_count) * kVectorsPerRow; + for (uint64_t linear = global_thread; + linear < num_vectors; linear += kGridThreads) { + const uint32_t route = static_cast( + linear / kVectorsPerRow); + const uint32_t vector_in_row = static_cast( + linear - static_cast(route) * kVectorsPerRow); + const uint32_t feature = + vector_in_row * kValuesPerVector; + const uint64_t full_row = + static_cast(pool_row) + route; + uint4 packed{}; + if (route < count) { + const uint2 raw = *reinterpret_cast( + source + full_row * kShape + feature); + const uint32_t raw_words[2] = {raw.x, raw.y}; + const uint32_t scale_exponent = scales[ + full_row * kScaleBlocksPerRow + feature / 128] & 0xffu; + const float post_scale = kApplyPostScale + ? scores[full_row] + : 1.0f; + auto* output_pairs = reinterpret_cast(&packed); + #pragma unroll + for (uint32_t pair = 0; pair < 4; ++ pair) { + const uint16_t fp8x2 = static_cast( + raw_words[pair / 2] >> ((pair & 1u) * 16)); + uint32_t bf16x2 = convert_fp8x2_power2_to_bf16x2( + fp8x2, scale_exponent); + if constexpr (kApplyPostScale) { + const auto dequantized = __bfloat1622float2( + *reinterpret_cast(&bf16x2)); + const auto scaled = __float22bfloat162_rn( + __fmul2_rn( + dequantized, + {post_scale, post_scale})); + bf16x2 = + *reinterpret_cast(&scaled); + } + output_pairs[pair] = bf16x2; + } } - output_pairs[pair] = bf16x2; + *reinterpret_cast( + destination + full_row * kShape + feature) = packed; } + pool_row += padded_count; + } + DG_DEVICE_ASSERT(pool_row <= max_pool_tokens); + + // Empty experts consume one permanent all-zero K tile beyond the routed + // pool. The host-side private scratch descriptors include these rows. + constexpr uint64_t kZeroVectors = + static_cast(kBlockK) * kVectorsPerRow; + for (uint64_t linear = global_thread; + linear < kZeroVectors; linear += kGridThreads) { + *reinterpret_cast( + destination + + static_cast(max_pool_tokens) * kShape + + linear * kValuesPerVector) = {}; } - *reinterpret_cast( - get_bf16_mn_bank_group( - destination, inner_mn, outer_k)) = packed; } // This is the production grouped-GEMM L2 swizzle specialized to the fixed GLM @@ -264,9 +292,14 @@ 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 __grid_constant__ layout::Workspace workspace, + const fp8_t* full_a, + const fp8_t* full_b, const uint32_t* full_a_sf, const uint32_t* full_b_sf, const float* full_scores, + bf16_t* cached_a, + bf16_t* cached_b, const __grid_constant__ cute::TmaDescriptor tensor_map_a, const __grid_constant__ cute::TmaDescriptor tensor_map_b, const __grid_constant__ cute::TmaDescriptor tensor_map_output_0, @@ -275,11 +308,22 @@ sm103_fp8_block128_mega_moe_wgrad_impl( using Scheduler = WgradTileScheduler; constexpr uint32_t kShapeM = Scheduler::kShapeM; constexpr uint32_t kShapeN = Scheduler::kShapeN; - constexpr uint32_t kFullABlocks = kShapeM / 128; - constexpr uint32_t kFullBBlocks = kShapeN / 128; - constexpr uint32_t kRawABytes = kBlockK * kBlockM * sizeof(fp8_t); - constexpr uint32_t kRawBBytes = - kBlockK * kLoadBlockN * sizeof(fp8_t); + constexpr uint32_t kLocalExperts = Scheduler::kLocalExperts; + + // One fused prologue per dedicated wgrad launch. This removes conversion + // from the output-tile loop while preserving the exact E4M3 + FP32 + // power-of-two scale and BF16-rounding semantics. + dequantize_route_pool_once< + kShapeM, kLocalExperts, kNumSMs, kThreads, kW2>( + expert_counts, max_pool_tokens, + full_a, full_a_sf, full_scores, cached_a); + dequantize_route_pool_once< + kShapeN, kLocalExperts, kNumSMs, kThreads, false>( + expert_counts, max_pool_tokens, + full_b, full_b_sf, full_scores, cached_b); + comm::grid_sync( + workspace, blockIdx.x, threadIdx.x, + []() { __syncthreads(); }); extern __shared__ __align__(1024) uint8_t smem_buffer[]; SharedStorage& storage = *reinterpret_cast(smem_buffer); @@ -301,9 +345,8 @@ sm103_fp8_block128_mega_moe_wgrad_impl( for (uint32_t i = 0; i < kStages; ++i) { storage.tma_full_barriers[i].init(1); storage.tma_empty_barriers[i].init(1); - // Both CTAs publish their converted halves to CTA 0. + // Both CTAs publish their direct-BF16 TMA completion to CTA 0. storage.mma_full_barriers[i].init(2); - storage.mma_empty_barriers[i].init(1); } #pragma unroll for (uint32_t i = 0; i < kNumEpilogueStages; ++i) { @@ -350,8 +393,8 @@ sm103_fp8_block128_mega_moe_wgrad_impl( }; if (warp_idx == kTMAWarp && cute::elect_one_sync()) { - // The production load warp issues two rectangular TMA transactions per - // K stage. Raw tiles are row-major [route-K, feature-MN]. + // The production load warp now reads the once-dequantized BF16 backing + // directly into the native UMMA swizzle. Scheduler scheduler(expert_counts); uint32_t expert, count, pool_row, m_block, n_block; while (scheduler.get_next( @@ -362,92 +405,45 @@ sm103_fp8_block128_mega_moe_wgrad_impl( for (uint32_t k_block = 0; k_block < num_k_blocks; ++k_block) { storage.tma_empty_barriers[stage_idx].wait(phase ^ 1u); - const uint32_t route = pool_row + k_block * kBlockK; + const uint32_t route = count == 0 + ? max_pool_tokens + : pool_row + k_block * kBlockK; const uint32_t a_feature = m_block * kBlockM; const uint32_t b_feature = n_block * kBlockN + cta_rank * kLoadBlockN; - tma::copy( + tma::copy( &tensor_map_a, &storage.tma_full_barriers[stage_idx], - storage.raw_a[stage_idx], + storage.smem_a[stage_idx], a_feature, route); - tma::copy( + tma::copy( &tensor_map_b, &storage.tma_full_barriers[stage_idx], - storage.raw_b[stage_idx], + storage.smem_b[stage_idx], b_feature, route); storage.tma_full_barriers[stage_idx] - .arrive_and_expect_tx(kRawABytes + kRawBBytes); + .arrive_and_expect_tx( + sizeof(storage.smem_a[0]) + + sizeof(storage.smem_b[0])); advance_pipeline(); } } - } else if ( - warp_idx >= kConvertFirstWarp && - warp_idx < kConvertFirstWarp + kConvertThreads / 32) { - // Four converter threads own each route. Each thread loads two aligned - // FP8x16 vectors per operand, converts packed E4M3x2 values, folds the - // power-of-two exponent directly into BF16, and emits four aligned - // BF16x8 bank groups per operand into the UMMA swizzle. + } else if (warp_idx == kReadyWarp && cute::elect_one_sync()) { + // Each CTA waits for its local direct-BF16 TMAs, then contributes one + // arrival to CTA 0. The leader MMA warp therefore observes both halves + // without multicast-copying different GLM feature tiles over each + // other. Scheduler scheduler(expert_counts); - const uint32_t convert_thread = - threadIdx.x - kConvertFirstWarp * 32; - const uint32_t route_in_k = convert_thread / 4; - const uint32_t quarter = convert_thread & 3u; 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.tma_full_barriers[stage_idx].wait(phase); - storage.mma_empty_barriers[stage_idx].wait(phase ^ 1u); - const uint32_t route = k_block * kBlockK + route_in_k; - const bool valid = route < count; - const uint64_t full_row = pool_row + route; - uint32_t a_scale_exponent = 0u; - uint32_t b_scale_exponent = 0u; - float a_post_scale = 1.0f; - if (valid) { - a_scale_exponent = full_a_sf[ - full_row * kFullABlocks + m_block] & 0xffu; - b_scale_exponent = full_b_sf[ - full_row * kFullBBlocks + - n_block * (kBlockN / 128) + cta_rank] & 0xffu; - if constexpr (kW2) - a_post_scale = full_scores[full_row]; - } - #pragma unroll - for (uint32_t chunk = 0; chunk < 2; ++chunk) { - const uint32_t inner = (quarter * 2 + chunk) * 16; - const auto* raw_a = storage.raw_a[stage_idx] + - route_in_k * kBlockM + inner; - const auto* raw_b = storage.raw_b[stage_idx] + - route_in_k * kLoadBlockN + inner; - convert_and_store_eight( - raw_a, storage.smem_a[stage_idx], - inner, route_in_k, - a_scale_exponent, a_post_scale, valid); - convert_and_store_eight( - raw_a + 8, storage.smem_a[stage_idx], - inner + 8, route_in_k, - a_scale_exponent, a_post_scale, valid); - convert_and_store_eight( - raw_b, storage.smem_b[stage_idx], - inner, route_in_k, - b_scale_exponent, 1.0f, valid); - convert_and_store_eight( - raw_b + 8, storage.smem_b[stage_idx], - inner + 8, route_in_k, - b_scale_exponent, 1.0f, valid); - } - ptx::sync_aligned(256, 1); cutlass::arch::fence_view_async_shared(); - if (convert_thread == 0) { - storage.tma_empty_barriers[stage_idx].arrive(); - storage.mma_full_barriers[stage_idx].arrive(0u); - } + storage.mma_full_barriers[stage_idx].arrive(0u); advance_pipeline(); } } @@ -498,7 +494,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( constexpr uint16_t kCTAMask = 3; cutlass::arch::umma_arrive_multicast_2x1SM( reinterpret_cast( - &storage.mma_empty_barriers[stage_idx]), + &storage.tma_empty_barriers[stage_idx]), kCTAMask); if (k_block == num_k_blocks - 1) { cutlass::arch::umma_arrive_multicast_2x1SM( From 749e5698b2a90838f3ed83bfe0c54de7c139b894 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 18:21:57 +0800 Subject: [PATCH 21/29] fix: publish block128 peer amax at cluster scope --- .../include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh | 7 +++++++ 1 file changed, 7 insertions(+) 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 c5640f5299..4ee500917f 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 @@ -1322,6 +1322,13 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 From 8fc792bb15712a5a85ebc1f8fe5db30e6acd8cad Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 19:28:06 +0800 Subject: [PATCH 22/29] fix: preserve post-down scores across MegaMoE waves --- .../deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh | 12 +++++++++--- deep_gemm/include/deep_gemm/layout/mega_moe.cuh | 14 ++++++++++++++ 2 files changed, 23 insertions(+), 3 deletions(-) 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 4ee500917f..e06f1de420 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 @@ -638,6 +638,8 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 % num_ring_tokens).template get_base_ptr() = weight; + if constexpr (kFP8Block128Weights) + *workspace.get_route_weight_ptr(pool_token_idx) = weight; // Write source metadata for combine write-back (logical pool token) *(saved_token_src_metadata != nullptr @@ -1559,9 +1561,13 @@ sm100_fp8_fp4_mega_moe_impl(void* y, // score before the final BF16 rounding. Moving // this multiply into L1 is not equivalent because // it changes E4M3 requantization before W2. - const float route_weight = *l1_topk_weights_buffer - .get_data_buffer(ring_m_idx + m_idx_in_block) - .template get_base_ptr(); + // 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 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 { From c8e6e83d9a9751458312daf7a00c9b48af8ab8d3 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 20:06:11 +0800 Subject: [PATCH 23/29] perf: align MegaMoE SwiGLU fast math --- csrc/sm103_fp8_block128.cu | 4 ++-- .../deep_gemm/impls/sm100_fp8_fp4_mega_moe_backward.cuh | 8 ++++++-- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index db20d87bd5..396d86fc91 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -1941,7 +1941,7 @@ void launch_persistent_forward( kNumSMs, kNumRanks, 0x7f800000u, - false, + true, deep_gemm::ActivationType::SwiGLU, true>); Kernel kernel = &deep_gemm::sm100_fp8_fp4_mega_moe_impl< @@ -1964,7 +1964,7 @@ void launch_persistent_forward( kNumSMs, kNumRanks, 0x7f800000u, - false, + true, deep_gemm::ActivationType::SwiGLU, true>; 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 index 9d53318bcd..a5625a677e 100644 --- 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 @@ -6134,7 +6134,11 @@ sm103_fp8_block128_mega_moe_backward_impl( const float gate = static_cast( ring_bf16[static_cast(row) * kHidden + physical_gate]); - const float sigmoid = 1.0f / (1.0f + expf(-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); @@ -6236,7 +6240,7 @@ sm103_fp8_block128_mega_moe_backward_impl( 2 * kIntermediate + col]) * ring_scores[row]; - const float sigmoid = 1.0f / (1.0f + expf(-gate)); + 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)) From 445396b399115b0d967678930bca539e5c74de22 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 21:10:13 +0800 Subject: [PATCH 24/29] feat: implement persistent MegaMoE reverse pipeline --- .../impls/sm100_fp8_fp4_mega_moe.hpp | 9 + csrc/sm103_fp8_block128.cu | 106 ++- .../impls/sm100_fp8_fp4_mega_moe.cuh | 171 +++- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 779 ++++++++++++++++++ .../sm103_fp8_block128_mega_moe_wgrad.cuh | 5 +- .../include/deep_gemm/scheduler/mega_moe.cuh | 187 +++-- 6 files changed, 1164 insertions(+), 93 deletions(-) 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 0d29ef806b..b70c52fb0b 100644 --- a/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp +++ b/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp @@ -124,10 +124,19 @@ static void __instantiate_kernel() {{ 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_l1_output, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, nullptr, nullptr )); diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 396d86fc91..b33dcc9deb 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -177,6 +177,7 @@ struct PersistentWorkspaceLayout { 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_dispatch_done; 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; @@ -192,6 +193,8 @@ struct PersistentWorkspaceLayout { 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_wgrad_bf16_narrow; deep_gemm::layout::Buffer backward_wgrad_bf16_wide; @@ -280,11 +283,16 @@ struct PersistentWorkspaceLayout { 1, capacity, backward_grad_y_scales.get_end_ptr()), + backward_dispatch_done( + deep_gemm::layout::Data(sizeof(uint32_t), false), + 1, + 1, + backward_grad_scores.get_end_ptr()), backward_ring_grad_y( deep_gemm::layout::Data(kPersistentHidden), 1, ring_tokens, - backward_grad_scores.get_end_ptr()), + backward_dispatch_done.get_end_ptr()), backward_ring_grad_y_scales( deep_gemm::layout::Data(kPersistentHidden / 32), 1, @@ -358,6 +366,19 @@ struct PersistentWorkspaceLayout { 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()), // The two dedicated wgrad kernels reuse these private operands. W13 // maps grad_preact -> narrow and x -> wide; W2 maps grad_y -> wide and // h -> narrow. One extra K tile is a permanent zero source for empty @@ -368,7 +389,7 @@ struct PersistentWorkspaceLayout { 1, workspace.num_max_pool_tokens + deep_gemm::sm103_block128_wgrad::kBlockK, - backward_full_grad_preact_scales.get_end_ptr()), + saved_down_unweighted.get_end_ptr()), backward_wgrad_bf16_wide( deep_gemm::layout::Data( kPersistentHidden * sizeof(__nv_bfloat16)), @@ -1834,6 +1855,9 @@ void launch_persistent_forward( 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, @@ -1853,6 +1877,16 @@ void launch_persistent_forward( {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, @@ -1892,6 +1926,15 @@ void launch_persistent_forward( 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, @@ -1920,6 +1963,15 @@ void launch_persistent_forward( 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, @@ -2002,12 +2054,23 @@ void launch_persistent_forward( 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())); + 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(); } @@ -2179,8 +2242,7 @@ void launch_persistent_backward_activation( {static_cast(kPersistentHidden), 1}, bf16_options); auto grad_h = torch::from_blob( - layout.backward_ring_bf16.get_base_ptr<__nv_bfloat16>() + - 2 * kPersistentIntermediate, + layout.backward_ring_bf16.base, {layout.ring_tokens, kPersistentIntermediate}, {static_cast(kPersistentHidden), 1}, bf16_options); @@ -2258,10 +2320,12 @@ void launch_persistent_backward_activation( using Kernel = decltype( &deep_gemm::sm103_block128_backward:: - sm103_fp8_block128_mega_moe_backward_impl); + sm103_fp8_block128_mega_moe_backward_persistent_impl< + kNumRanks, kNumSMs>); Kernel kernel = &deep_gemm::sm103_block128_backward:: - sm103_fp8_block128_mega_moe_backward_impl; + 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( @@ -2289,49 +2353,31 @@ void launch_persistent_backward_activation( layout.ring_tokens, layout.sf_ring_tokens, layout.workspace.num_max_pool_tokens, sym_buffer, layout.workspace, - reinterpret_cast(x.data_ptr()), reinterpret_cast( grad_output.data_ptr()), - topk_scores.data_ptr(), - layout.input_tokens.get_base_ptr(), - layout.input_scales.get_base_ptr(), - layout.backward_grad_y_tokens - .get_base_ptr(), - layout.backward_grad_y_scales.get_base_ptr(), - layout.input_topk_scores.get_base_ptr(), - layout.backward_grad_scores.get_base_ptr(), layout.combine_tokens.get_base_ptr(), - layout.l1_tokens.get_base_ptr(), - layout.l1_scales.get_base_ptr(), layout.backward_ring_grad_y .get_base_ptr(), layout.backward_ring_grad_y_scales.get_base_ptr(), - layout.l1_scores.get_base_ptr(), - layout.l2_tokens.get_base_ptr(), - layout.l2_scales.get_base_ptr(), layout.backward_ring_grad_preact .get_base_ptr(), layout.backward_ring_grad_preact_scales.get_base_ptr(), layout.backward_ring_bf16.get_base_ptr(), - layout.backward_ring_dscore.get_base_ptr(), - layout.backward_full_x.get_base_ptr(), - layout.backward_full_x_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.backward_full_h.get_base_ptr(), - layout.backward_full_h_scales.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_x, tensor_map_ring_x_sf, + layout.backward_dispatch_done.get_base_ptr(), tensor_map_ring_grad_y, tensor_map_ring_grad_y_sf, - tensor_map_ring_h, tensor_map_ring_h_sf, 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_gate_up, + tensor_map_w2_dgrad, tensor_map_w13_dgrad, tensor_map_grad_h, tensor_map_grad_x, w13_scale.data_ptr(), w2_scale.data_ptr())); C10_CUDA_KERNEL_LAUNCH_CHECK(); 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 e06f1de420..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 @@ -76,12 +76,21 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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_saved_output, const float* l1_block128_scales = nullptr, - const float* l2_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; @@ -115,10 +124,12 @@ 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 @@ -589,6 +600,9 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 % 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( @@ -596,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()) { @@ -626,8 +647,17 @@ sm100_fp8_fp4_mega_moe_impl(void* y, #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 * num_sf_ring_tokens + 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(); @@ -640,6 +670,8 @@ sm100_fp8_fp4_mega_moe_impl(void* y, *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) *(saved_token_src_metadata != nullptr @@ -1211,6 +1243,55 @@ sm100_fp8_fp4_mega_moe_impl(void* y, 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 (kActivationClampBits != 0x7f800000u) { const float activation_clamp = __uint_as_float(kActivationClampBits); @@ -1427,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(); } @@ -1442,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(); } @@ -1525,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; 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 index a5625a677e..b2ea284b9a 100644 --- 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 @@ -6361,6 +6361,785 @@ sm103_fp8_block128_mega_moe_backward_impl( #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_bf16, + fp8_t* ring_grad_y, + uint32_t* ring_grad_y_sf, + fp8_t* ring_grad_preact, + uint32_t* ring_grad_preact_sf, + bf16_t* ring_bf16, + 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, + uint32_t* dispatch_done, + 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 __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; + 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 kDispatchWarps = 4; + constexpr uint32_t kDispatchThreads = kDispatchWarps * 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); + cute::prefetch_tma_descriptor(&tensor_map_grad_h); + cute::prefetch_tma_descriptor(&tensor_map_grad_x); + } + + // The forward combine plane is dead on entry. Reuse slot zero as the + // symmetric BF16 grad-y source until every reverse dispatch warp has + // completed its remote reads. + for (uint64_t linear = global_thread; + linear < static_cast(num_tokens) * kHidden; + linear += global_stride) { + symmetric_bf16[linear] = compact_grad_y[linear]; + } + if (global_thread == 0) + *dispatch_done = 0u; + 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); + } + 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 < kDispatchWarps) { + cutlass::arch::warpgroup_reg_dealloc<48>(); + const uint32_t dispatch_warp = warp_idx; + const uint32_t global_warp = + blockIdx.x * kDispatchWarps + dispatch_warp; + constexpr uint32_t kGlobalWarps = kNumSMs * kDispatchWarps; + 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_bf16 + + static_cast(metadata.token_idx) * + kHidden, + metadata.rank_idx); + const float score = full_scores[pool_row]; + float dscore = 0.0f; + + #pragma unroll 1 + for (uint32_t block = 0; + block < kHiddenScaleBlocks; ++block) { + float dz_values[4]; + float local_amax = 0.0f; + #pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const uint32_t col = + block * 128 + lane_idx * 4 + i; + const float dy = static_cast( + remote_grad_y[col]); + 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; + } + } + + 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); + }); + if (blockIdx.x == 0 && dispatch_warp == 0 && lane_idx == 0) { + __threadfence(); + atomicExch(dispatch_done, 1u); + } + } 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; + uint32_t tma_stage = 0; + auto smem_cd = utils::PatternVisitor([&](const uint32_t& i) { + return storage.smem_cd[i]; + }); + 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(); + const auto* output_map = + block_phase == sched::BackwardBlockPhase::W2Dgrad + ? &tensor_map_grad_h + : &tensor_map_grad_x; + + 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) { + } + } + + epilogue::sm100_store_cd_swap_ab< + kBlockM, kBlockN, kStoreBlockM, kBlockN, + kSwizzle, kNumTMAStoreStages, + 128u, GemmType::Normal, false, + bf16_t, + epilogue::transform::EpilogueIdentity>( + smem_cd, tma_stage, + accum_stage * kUMMAN, + ring_m, n_block * kBlockN, 0, + math::align(valid_m, 16u), + epilogue_warp, lane_idx, + &storage.tmem_empty_barriers[accum_stage], + *output_map); + if (epilogue_warp == 0) + cute::tma_store_wait<0>(); + ptx::sync_aligned( + 128u, kEpilogueBarrier); + + if (block_phase == + sched::BackwardBlockPhase::W2Dgrad) { + const uint32_t hidden_col = + n_block * kBlockN + epilogue_thread; + #pragma unroll 1 + for (uint32_t row = 0; row < valid_m; ++row) { + const uint32_t ring_row = ring_m + row; + const uint32_t pool_row = pool_m + row; + const float dh = static_cast( + ring_bf16[ + static_cast(ring_row) * + kHidden + + hidden_col]); + const float gate = static_cast( + saved_l1_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + hidden_col]); + const float up = static_cast( + saved_l1_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + kIntermediate + hidden_col]); + const float sigmoid = + math::fast_rcp(1.0f + __expf(-gate)); + const float dgate = + dh * up * sigmoid * + (1.0f + gate * (1.0f - sigmoid)); + const float dup = dh * gate * sigmoid; + const float gate_amax = + reduce_group_128( + cute::abs(dgate), storage, 2); + const float up_amax = + reduce_group_128( + cute::abs(dup), storage, 2); + float gate_inv, up_inv; + const uint32_t gate_scale = + packed_power2_scale( + gate_amax, gate_inv); + const uint32_t up_scale = + packed_power2_scale(up_amax, up_inv); + const uint32_t gate_col = hidden_col; + const uint32_t up_col = + kIntermediate + hidden_col; + ring_grad_preact[ + static_cast(ring_row) * + (2 * kIntermediate) + + gate_col] = fp8_t(dgate * gate_inv); + ring_grad_preact[ + static_cast(ring_row) * + (2 * kIntermediate) + + up_col] = fp8_t(dup * up_inv); + full_grad_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + gate_col] = fp8_t(dgate * gate_inv); + full_grad_preact[ + static_cast(pool_row) * + (2 * kIntermediate) + + up_col] = fp8_t(dup * up_inv); + 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 { + while (ptx::ld_acq(dispatch_done) != 1u) { + } + constexpr uint32_t kVecElements = + sizeof(uint4) / sizeof(bf16_t); + constexpr uint32_t kVecsPerNBlock = + kBlockN / kVecElements; + for (uint32_t linear = epilogue_thread; + linear < valid_m * kVecsPerNBlock; + linear += 128u) { + const uint32_t row = + linear / kVecsPerNBlock; + const uint32_t vec = + linear - row * kVecsPerNBlock; + const uint32_t pool_row = pool_m + row; + const uint32_t ring_row = ring_m + row; + const auto metadata = + token_src_metadata[pool_row]; + const uint32_t source_vec = + n_block * kVecsPerNBlock + vec; + const uint4 value = + reinterpret_cast( + ring_bf16 + + static_cast(ring_row) * + kHidden)[source_vec]; + auto* remote = sym_buffer.map( + reinterpret_cast( + symmetric_bf16 + + (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_bf16[ + (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; + } + if (global_thread == 0) + *dispatch_done = 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 index 31ed379477..d928ad537a 100644 --- 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 @@ -313,8 +313,11 @@ sm103_fp8_block128_mega_moe_wgrad_impl( // One fused prologue per dedicated wgrad launch. This removes conversion // from the output-tile loop while preserving the exact E4M3 + FP32 // power-of-two scale and BF16-rounding semantics. + // Reverse dispatch has already formed and power-of-two quantized dz = + // BF16(dy * route_score) for W2. Dequantize it exactly once here; applying + // the route score again would square the router derivative. dequantize_route_pool_once< - kShapeM, kLocalExperts, kNumSMs, kThreads, kW2>( + kShapeM, kLocalExperts, kNumSMs, kThreads, false>( expert_counts, max_pool_tokens, full_a, full_a_sf, full_scores, cached_a); dequantize_route_pool_once< diff --git a/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh index 0b1de58b7c..52a5f6de5a 100644 --- a/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh +++ b/deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh @@ -16,114 +16,179 @@ enum class BlockPhase { Linear2 = 2 }; -// Fixed training reverse for the GLM large-M specialization. Unlike the -// forward scheduler, the expert counts are immutable outputs from the saved -// forward call, so no Workspace polling is needed. One expert constitutes a -// wave: its W13 recompute, W2 dgrad, and W13 dgrad traverse the same ring slot -// before the next expert may reuse it. +// 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, - RecomputeW13 = 1, - W2Dgrad = 2, - W13Dgrad = 3, + 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 kNumSMs> + uint32_t kNumExpertsPerRank, uint32_t kNumExpertsPerWave, + uint32_t kNumSMs, + uint32_t kNumExpertsPerLane = + math::constexpr_ceil_div(kNumExpertsPerRank, 32u)> struct MegaMoEBackwardScheduler { - static constexpr uint32_t kRecomputeBlockNs = - (2 * kIntermediateHidden) / BLOCK_N; static constexpr uint32_t kW2DgradBlockNs = kIntermediateHidden / BLOCK_N; static constexpr uint32_t kW13DgradBlockNs = kHidden / BLOCK_N; - static constexpr uint32_t kRecomputeBlockKs = kHidden / BLOCK_K; 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(kRecomputeBlockNs % 2 == 0 && - kW2DgradBlockNs % 2 == 0 && + 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::RecomputeW13; - uint32_t local_expert_idx = 0; - uint32_t pool_block_offset = 0; - uint32_t num_tokens = 0; + 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) { - num_tokens = static_cast(__ldg(expert_counts)); + #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_block_ns() const { - return phase == BackwardBlockPhase::RecomputeW13 - ? kRecomputeBlockNs - : phase == BackwardBlockPhase::W2Dgrad - ? kW2DgradBlockNs - : kW13DgradBlockNs; + 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_num_block_ks() const { - return phase == BackwardBlockPhase::RecomputeW13 - ? kRecomputeBlockKs - : phase == BackwardBlockPhase::W2Dgrad - ? kW2DgradBlockKs - : kW13DgradBlockKs; + 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 pool_block_offset; + return current_pool_block_offset; } CUTLASS_DEVICE uint32_t get_current_num_m_blocks() const { - return math::ceil_div(num_tokens, BLOCK_M); + return math::ceil_div(current_num_tokens, BLOCK_M); } template CUTLASS_DEVICE uint32_t get_valid_m() const { const auto value = cute::min( - num_tokens - m_block_idx * BLOCK_M, BLOCK_M); + current_num_tokens - m_block_idx * BLOCK_M, BLOCK_M); return kDoUMMAAligned ? math::align(value, 16u) : value; } - CUTLASS_DEVICE cute::tuple - get_next_block() { - while (local_expert_idx < kNumExpertsPerRank) { - const uint32_t block_ns = get_num_block_ns(); - const uint32_t phase_blocks = get_current_num_m_blocks() * block_ns; - if (block_idx < phase_blocks) { - m_block_idx = block_idx / block_ns; - n_block_idx = block_idx - m_block_idx * block_ns; - block_idx += kNumSMs; - return {phase, local_expert_idx, m_block_idx, n_block_idx}; + 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; + } - // Every role restarts from its physical CTA index for the next - // phase. Readiness counters, not a grid-wide host launch, carry - // the producer/consumer dependency between the three phases. - block_idx = blockIdx.x; - if (phase == BackwardBlockPhase::RecomputeW13) { - phase = BackwardBlockPhase::W2Dgrad; - } else if (phase == BackwardBlockPhase::W2Dgrad) { + 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 { - pool_block_offset += get_current_num_m_blocks(); - ++local_expert_idx; - if (local_expert_idx >= kNumExpertsPerRank) - break; - num_tokens = static_cast( - __ldg(expert_counts + local_expert_idx)); - phase = BackwardBlockPhase::RecomputeW13; + phase = BackwardBlockPhase::W2Dgrad; } } return {BackwardBlockPhase::None, 0, 0, 0}; @@ -136,7 +201,11 @@ struct MegaMoEBackwardScheduler { current_m_block_idx, current_n_block_idx); if (current_phase == BackwardBlockPhase::None) break; - func(current_phase, expert_idx, get_num_block_ks(), + 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); } } From 5aab8cefc2b4664ccfebfc4b93ba695c8e7d3768 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 21:25:38 +0800 Subject: [PATCH 25/29] fix: align persistent reverse TMA workspace --- csrc/sm103_fp8_block128.cu | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index b33dcc9deb..34f2ec41bf 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -284,7 +284,11 @@ struct PersistentWorkspaceLayout { capacity, backward_grad_y_scales.get_end_ptr()), backward_dispatch_done( - deep_gemm::layout::Data(sizeof(uint32_t), false), + // The control word itself is four bytes, but every following + // reverse-pipeline operand is TMA-addressed and therefore needs + // a 16-byte-aligned base. Reserve one complete TMA alignment + // unit here instead of shifting the entire suffix by four bytes. + deep_gemm::layout::Data(16, false), 1, 1, backward_grad_scores.get_end_ptr()), From 09610521fb3e075c6f2bd92f6846c9f1a699d6fb Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 21:43:59 +0800 Subject: [PATCH 26/29] fix: separate persistent reverse transport planes --- csrc/sm103_fp8_block128.cu | 19 +++++------- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 29 +++++++------------ 2 files changed, 17 insertions(+), 31 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index 34f2ec41bf..ba8ae5939e 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -177,7 +177,6 @@ struct PersistentWorkspaceLayout { 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_dispatch_done; 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; @@ -283,20 +282,11 @@ struct PersistentWorkspaceLayout { 1, capacity, backward_grad_y_scales.get_end_ptr()), - backward_dispatch_done( - // The control word itself is four bytes, but every following - // reverse-pipeline operand is TMA-addressed and therefore needs - // a 16-byte-aligned base. Reserve one complete TMA alignment - // unit here instead of shifting the entire suffix by four bytes. - deep_gemm::layout::Data(16, false), - 1, - 1, - backward_grad_scores.get_end_ptr()), backward_ring_grad_y( deep_gemm::layout::Data(kPersistentHidden), 1, ring_tokens, - backward_dispatch_done.get_end_ptr()), + backward_grad_scores.get_end_ptr()), backward_ring_grad_y_scales( deep_gemm::layout::Data(kPersistentHidden / 32), 1, @@ -2359,6 +2349,12 @@ void launch_persistent_backward_activation( sym_buffer, layout.workspace, reinterpret_cast( grad_output.data_ptr()), + // Reverse transport needs the same independent input/output + // lifetimes as upstream forward. Reuse the not-yet-live private + // wgrad-wide backing as the immutable symmetric grad-y plane; the + // normal combine plane remains the remote dX destination. + layout.backward_wgrad_bf16_wide + .get_base_ptr(), layout.combine_tokens.get_base_ptr(), layout.backward_ring_grad_y .get_base_ptr(), @@ -2378,7 +2374,6 @@ void launch_persistent_backward_activation( reinterpret_cast(grad_x.data_ptr()), layout.backward_grad_scores.get_base_ptr(), grad_scores.data_ptr(), - layout.backward_dispatch_done.get_base_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, 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 index b2ea284b9a..011425c6f5 100644 --- 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 @@ -6380,7 +6380,8 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( const __grid_constant__ layout::SymBuffer sym_buffer, const __grid_constant__ layout::Workspace workspace, const bf16_t* compact_grad_y, - bf16_t* symmetric_bf16, + 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, @@ -6396,7 +6397,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( bf16_t* grad_x, float* symmetric_grad_scores, float* grad_scores, - uint32_t* dispatch_done, 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, @@ -6444,16 +6444,15 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( cute::prefetch_tma_descriptor(&tensor_map_grad_x); } - // The forward combine plane is dead on entry. Reuse slot zero as the - // symmetric BF16 grad-y source until every reverse dispatch warp has - // completed its remote reads. + // 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_bf16[linear] = compact_grad_y[linear]; + symmetric_grad_y[linear] = compact_grad_y[linear]; } - if (global_thread == 0) - *dispatch_done = 0u; comm::nvlink_barrier( workspace, sym_buffer, blockIdx.x, threadIdx.x, []() { __syncthreads(); }); @@ -6518,7 +6517,7 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( const auto metadata = token_src_metadata[pool_row]; const auto* remote_grad_y = sym_buffer.map( - symmetric_bf16 + + symmetric_grad_y + static_cast(metadata.token_idx) * kHidden, metadata.rank_idx); @@ -6610,10 +6609,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( ptx::sync_aligned( kDispatchThreads, kDispatchNamedBarrier); }); - if (blockIdx.x == 0 && dispatch_warp == 0 && lane_idx == 0) { - __threadfence(); - atomicExch(dispatch_done, 1u); - } } else if (warp_idx == 4) { cutlass::arch::warpgroup_reg_dealloc<40>(); Scheduler scheduler(expert_counts); @@ -7038,8 +7033,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( } __syncwarp(); } else { - while (ptx::ld_acq(dispatch_done) != 1u) { - } constexpr uint32_t kVecElements = sizeof(uint4) / sizeof(bf16_t); constexpr uint32_t kVecsPerNBlock = @@ -7064,7 +7057,7 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( kHidden)[source_vec]; auto* remote = sym_buffer.map( reinterpret_cast( - symmetric_bf16 + + symmetric_grad_x + (static_cast( metadata.topk_idx) * capacity + @@ -7111,7 +7104,7 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( value = __fadd_rn( value, static_cast( - symmetric_bf16[ + symmetric_grad_x[ (static_cast(slot) * capacity + token) * kHidden + @@ -7132,8 +7125,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( *workspace.get_l2_full_count_ptr(block) = 0u; *workspace.get_l2_empty_count_ptr(block) = 0u; } - if (global_thread == 0) - *dispatch_done = 0u; comm::grid_sync( workspace, blockIdx.x, threadIdx.x, []() { __syncthreads(); }); From cf1e4011663f3919a49166cf9377adff5191ef9a Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 22:20:14 +0800 Subject: [PATCH 27/29] perf: use persistent TMA reverse transport --- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 205 +++++++++++++----- 1 file changed, 153 insertions(+), 52 deletions(-) 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 index 011425c6f5..130f95f61d 100644 --- 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 @@ -5393,6 +5393,16 @@ 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; // sm100_store_cd_swap_ab maps one 128-thread warpgroup over the two 64-column // BF16 atoms of BLOCK_N. STORE_BLOCK_M=16 is also required so every valid // UMMA-N extent (which is 16-row aligned) emits at least one store and returns @@ -5428,6 +5438,9 @@ 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]; @@ -5439,6 +5452,8 @@ struct alignas(1024) SharedStorage { Barrier empty_barriers[kStages]; Barrier tmem_full_barriers[kNumEpilogueStages]; Barrier tmem_empty_barriers[kNumEpilogueStages]; + Barrier dispatch_barriers[kNumDispatchWarps] + [kDispatchPullStages]; uint32_t tmem_ptr; }; @@ -5446,6 +5461,10 @@ 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 @@ -6414,8 +6433,7 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( constexpr uint32_t kHiddenScaleBlocks = kHidden / 128; constexpr uint32_t kGradPreactScaleBlocks = (2 * kIntermediate) / 128; - constexpr uint32_t kDispatchWarps = 4; - constexpr uint32_t kDispatchThreads = kDispatchWarps * 32; + constexpr uint32_t kDispatchThreads = kNumDispatchWarps * 32; constexpr uint32_t kEpilogueBarrier = 9; using Scheduler = sched::MegaMoEBackwardScheduler< kBlockM, kBlockN, kBlockK, @@ -6470,6 +6488,14 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( 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) @@ -6484,12 +6510,14 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( pipeline_phase ^= stage_idx == 0; }; - if (warp_idx < kDispatchWarps) { + if (warp_idx < kNumDispatchWarps) { cutlass::arch::warpgroup_reg_dealloc<48>(); const uint32_t dispatch_warp = warp_idx; const uint32_t global_warp = - blockIdx.x * kDispatchWarps + dispatch_warp; - constexpr uint32_t kGlobalWarps = kNumSMs * kDispatchWarps; + 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 @@ -6524,58 +6552,131 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( 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 block = 0; - block < kHiddenScaleBlocks; ++block) { - float dz_values[4]; - float local_amax = 0.0f; - #pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const uint32_t col = - block * 128 + lane_idx * 4 + i; - const float dy = static_cast( - remote_grad_y[col]); - const float down = static_cast( - saved_down_unweighted[ + 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]); - 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; + 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 (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); From 4e1673b6b6c470b4e3c12bd6cc89e1550b0994a1 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 22:57:26 +0800 Subject: [PATCH 28/29] perf: fuse persistent reverse epilogues --- csrc/sm103_fp8_block128.cu | 74 ---- .../impls/sm100_fp8_fp4_mega_moe_backward.cuh | 418 ++++++++++++------ 2 files changed, 276 insertions(+), 216 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index ba8ae5939e..dc9d17a4eb 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -109,7 +109,6 @@ constexpr uint32_t kPersistentBlockM = 192; constexpr uint32_t kPersistentBlockN = 128; constexpr uint32_t kPersistentBlockK = 128; constexpr uint32_t kPersistentStoreBlockM = 32; -constexpr uint32_t kPersistentReverseStoreBlockM = 16; constexpr uint32_t kPersistentSFBlockM = 256; constexpr uint32_t kPersistentSFBlockN = 128; constexpr uint32_t kPersistentStages = 6; @@ -2193,18 +2192,7 @@ void launch_persistent_backward_activation( const auto int_options = torch::TensorOptions() .dtype(torch::kInt) .device(device); - const auto bf16_options = torch::TensorOptions() - .dtype(torch::kBFloat16) - .device(device); - auto ring_x = torch::from_blob( - layout.l1_tokens.base, - {layout.ring_tokens, kPersistentHidden}, fp8_options); - auto ring_x_sf = torch::from_blob( - layout.l1_scales.base, - {layout.sf_ring_tokens, kPersistentHidden / 128}, - {1, static_cast(layout.sf_ring_tokens)}, - int_options); auto ring_grad_y = torch::from_blob( layout.backward_ring_grad_y.base, {layout.ring_tokens, kPersistentHidden}, fp8_options); @@ -2213,14 +2201,6 @@ void launch_persistent_backward_activation( {layout.sf_ring_tokens, kPersistentHidden / 128}, {1, static_cast(layout.sf_ring_tokens)}, int_options); - auto ring_h = torch::from_blob( - layout.l2_tokens.base, - {layout.ring_tokens, kPersistentIntermediate}, fp8_options); - auto ring_h_sf = torch::from_blob( - layout.l2_scales.base, - {layout.sf_ring_tokens, kPersistentIntermediate / 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); @@ -2229,31 +2209,6 @@ void launch_persistent_backward_activation( {layout.sf_ring_tokens, 2 * kPersistentIntermediate / 128}, {1, static_cast(layout.sf_ring_tokens)}, int_options); - - auto gate_up = torch::from_blob( - layout.backward_ring_bf16.base, - {layout.ring_tokens, 2 * kPersistentIntermediate}, - {static_cast(kPersistentHidden), 1}, - bf16_options); - auto grad_h = torch::from_blob( - layout.backward_ring_bf16.base, - {layout.ring_tokens, kPersistentIntermediate}, - {static_cast(kPersistentHidden), 1}, - bf16_options); - auto ring_grad_x = torch::from_blob( - layout.backward_ring_bf16.base, - {layout.ring_tokens, kPersistentHidden}, - {static_cast(kPersistentHidden), 1}, - bf16_options); - - const auto tensor_map_ring_x = deep_gemm::make_tma_2d_desc( - ring_x, kPersistentHidden, layout.ring_tokens, - kPersistentBlockK, kPersistentBlockM / 2, - kPersistentHidden, 128); - const auto tensor_map_ring_x_sf = deep_gemm::make_tma_sf_desc( - cute::UMMA::Major::MN, ring_x_sf, - layout.sf_ring_tokens, kPersistentHidden, - kPersistentSFBlockM, 32, 1, 0, 0, false, 1); const auto tensor_map_ring_grad_y = deep_gemm::make_tma_2d_desc( ring_grad_y, kPersistentHidden, layout.ring_tokens, kPersistentBlockK, kPersistentBlockM / 2, @@ -2262,14 +2217,6 @@ void launch_persistent_backward_activation( 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_h = deep_gemm::make_tma_2d_desc( - ring_h, kPersistentIntermediate, layout.ring_tokens, - kPersistentBlockK, kPersistentBlockM / 2, - kPersistentIntermediate, 128); - const auto tensor_map_ring_h_sf = deep_gemm::make_tma_sf_desc( - cute::UMMA::Major::MN, ring_h_sf, - layout.sf_ring_tokens, kPersistentIntermediate, - 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, @@ -2281,12 +2228,6 @@ void launch_persistent_backward_activation( 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_w13_recompute = 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_w2_dgrad = deep_gemm::make_tma_b_desc( cute::UMMA::Major::MN, w2_weight, kPersistentIntermediate, kPersistentHidden, @@ -2299,19 +2240,6 @@ void launch_persistent_backward_activation( kPersistentBlockN, kPersistentBlockK, static_cast(w13_weight.stride(-2)), static_cast(w13_weight.size(0) / 2), 128); - const auto tensor_map_gate_up = deep_gemm::make_tma_2d_desc( - gate_up, 2 * kPersistentIntermediate, layout.ring_tokens, - kPersistentBlockN, kPersistentReverseStoreBlockM, - kPersistentHidden, 128); - const auto tensor_map_grad_h = deep_gemm::make_tma_2d_desc( - grad_h, kPersistentIntermediate, layout.ring_tokens, - kPersistentBlockN, kPersistentReverseStoreBlockM, - kPersistentHidden, 128); - const auto tensor_map_grad_x = deep_gemm::make_tma_2d_desc( - ring_grad_x, kPersistentHidden, layout.ring_tokens, - kPersistentBlockN, kPersistentReverseStoreBlockM, - kPersistentHidden, 128); - using Kernel = decltype( &deep_gemm::sm103_block128_backward:: sm103_fp8_block128_mega_moe_backward_persistent_impl< @@ -2362,7 +2290,6 @@ void launch_persistent_backward_activation( layout.backward_ring_grad_preact .get_base_ptr(), layout.backward_ring_grad_preact_scales.get_base_ptr(), - layout.backward_ring_bf16.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(), @@ -2377,7 +2304,6 @@ void launch_persistent_backward_activation( 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, - tensor_map_grad_h, tensor_map_grad_x, w13_scale.data_ptr(), w2_scale.data_ptr())); C10_CUDA_KERNEL_LAUNCH_CHECK(); } 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 index 130f95f61d..e9bf06d969 100644 --- 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 @@ -5403,10 +5403,10 @@ static constexpr uint32_t kDispatchPullElements = kDispatchPullBytes / sizeof(cutlass::bfloat16_t); static constexpr uint32_t kDispatchPullChunks = kHidden * sizeof(cutlass::bfloat16_t) / kDispatchPullBytes; -// sm100_store_cd_swap_ab maps one 128-thread warpgroup over the two 64-column -// BF16 atoms of BLOCK_N. STORE_BLOCK_M=16 is also required so every valid -// UMMA-N extent (which is 16-row aligned) emits at least one store and returns -// the TMEM-empty arrival. +// 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; @@ -5503,6 +5503,37 @@ CUTLASS_DEVICE float reduce_group_128( 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); @@ -6405,7 +6436,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( uint32_t* ring_grad_y_sf, fp8_t* ring_grad_preact, uint32_t* ring_grad_preact_sf, - bf16_t* ring_bf16, fp8_t* full_grad_y, uint32_t* full_grad_y_sf, const float* full_scores, @@ -6422,8 +6452,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( 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 __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 @@ -6458,8 +6486,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( 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); - cute::prefetch_tma_descriptor(&tensor_map_grad_h); - cute::prefetch_tma_descriptor(&tensor_map_grad_x); } // Match upstream's separate input and combine transport lifetimes. The @@ -6977,10 +7003,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( const uint32_t epilogue_thread = epilogue_warp * 32 + lane_idx; uint32_t current_iter = 0; - uint32_t tma_stage = 0; - auto smem_cd = utils::PatternVisitor([&](const uint32_t& i) { - return storage.smem_cd[i]; - }); Scheduler scheduler(expert_counts); scheduler.for_each_block( [&](const sched::BackwardBlockPhase block_phase, @@ -7005,11 +7027,6 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( const uint32_t pool_m = pool_block * kBlockM; const uint32_t valid_m = scheduler.template get_valid_m(); - const auto* output_map = - block_phase == sched::BackwardBlockPhase::W2Dgrad - ? &tensor_map_grad_h - : &tensor_map_grad_x; - if (block_phase == sched::BackwardBlockPhase::W2Dgrad) { const uint32_t empty_target = @@ -7019,103 +7036,151 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( ring_block)) != empty_target) { } } - - epilogue::sm100_store_cd_swap_ab< - kBlockM, kBlockN, kStoreBlockM, kBlockN, - kSwizzle, kNumTMAStoreStages, - 128u, GemmType::Normal, false, - bf16_t, - epilogue::transform::EpilogueIdentity>( - smem_cd, tma_stage, - accum_stage * kUMMAN, - ring_m, n_block * kBlockN, 0, - math::align(valid_m, 16u), - epilogue_warp, lane_idx, - &storage.tmem_empty_barriers[accum_stage], - *output_map); - if (epilogue_warp == 0) - cute::tma_store_wait<0>(); - ptx::sync_aligned( - 128u, kEpilogueBarrier); - + 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; - #pragma unroll 1 - for (uint32_t row = 0; row < valid_m; ++row) { - const uint32_t ring_row = ring_m + row; - const uint32_t pool_row = pool_m + row; - const float dh = static_cast( - ring_bf16[ - static_cast(ring_row) * - kHidden + - hidden_col]); - const float gate = static_cast( - saved_l1_preact[ - static_cast(pool_row) * - (2 * kIntermediate) + - hidden_col]); - const float up = static_cast( - saved_l1_preact[ - static_cast(pool_row) * - (2 * kIntermediate) + - kIntermediate + hidden_col]); - const float sigmoid = - math::fast_rcp(1.0f + __expf(-gate)); - const float dgate = - dh * up * sigmoid * - (1.0f + gate * (1.0f - sigmoid)); - const float dup = dh * gate * sigmoid; - const float gate_amax = - reduce_group_128( - cute::abs(dgate), storage, 2); - const float up_amax = - reduce_group_128( - cute::abs(dup), storage, 2); - float gate_inv, up_inv; - const uint32_t gate_scale = - packed_power2_scale( - gate_amax, gate_inv); - const uint32_t up_scale = - packed_power2_scale(up_amax, up_inv); - const uint32_t gate_col = hidden_col; - const uint32_t up_col = - kIntermediate + hidden_col; - ring_grad_preact[ - static_cast(ring_row) * - (2 * kIntermediate) + - gate_col] = fp8_t(dgate * gate_inv); - ring_grad_preact[ - static_cast(ring_row) * - (2 * kIntermediate) + - up_col] = fp8_t(dup * up_inv); - full_grad_preact[ - static_cast(pool_row) * - (2 * kIntermediate) + - gate_col] = fp8_t(dgate * gate_inv); - full_grad_preact[ - static_cast(pool_row) * - (2 * kIntermediate) + - up_col] = fp8_t(dup * up_inv); - 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; + 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( @@ -7134,42 +7199,111 @@ sm103_fp8_block128_mega_moe_backward_persistent_impl( } __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; - for (uint32_t linear = epilogue_thread; - linear < valid_m * kVecsPerNBlock; - linear += 128u) { - const uint32_t row = - linear / kVecsPerNBlock; - const uint32_t vec = - linear - row * kVecsPerNBlock; - const uint32_t pool_row = pool_m + row; - const uint32_t ring_row = ring_m + row; - const auto metadata = - token_src_metadata[pool_row]; - const uint32_t source_vec = - n_block * kVecsPerNBlock + vec; - const uint4 value = - reinterpret_cast( - ring_bf16 + - static_cast(ring_row) * - kHidden)[source_vec]; - 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; + 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); } - ptx::sync_aligned( - 128u, kEpilogueBarrier); if (epilogue_warp == 0 && cute::elect_one_sync()) { __threadfence_system(); From 5018de906eb36731aea4521729a803d614b4da76 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 22 Jul 2026 23:44:32 +0800 Subject: [PATCH 29/29] perf: fuse wgrad dequant load prologues --- csrc/sm103_fp8_block128.cu | 74 +---- .../sm103_fp8_block128_mega_moe_wgrad.cuh | 271 +++++++----------- 2 files changed, 117 insertions(+), 228 deletions(-) diff --git a/csrc/sm103_fp8_block128.cu b/csrc/sm103_fp8_block128.cu index dc9d17a4eb..3178ba7480 100644 --- a/csrc/sm103_fp8_block128.cu +++ b/csrc/sm103_fp8_block128.cu @@ -193,8 +193,7 @@ struct PersistentWorkspaceLayout { 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_wgrad_bf16_narrow; - deep_gemm::layout::Buffer backward_wgrad_bf16_wide; + deep_gemm::layout::Buffer backward_symmetric_grad_y; PersistentWorkspaceLayout( void* base, @@ -372,28 +371,20 @@ struct PersistentWorkspaceLayout { 1, workspace.num_max_pool_tokens, saved_l1_preact.get_end_ptr()), - // The two dedicated wgrad kernels reuse these private operands. W13 - // maps grad_preact -> narrow and x -> wide; W2 maps grad_y -> wide and - // h -> narrow. One extra K tile is a permanent zero source for empty - // experts. Capacity remains a once-derived context/CP consequence. - backward_wgrad_bf16_narrow( - deep_gemm::layout::Data( - 2 * kPersistentIntermediate * sizeof(__nv_bfloat16)), - 1, - workspace.num_max_pool_tokens + - deep_gemm::sm103_block128_wgrad::kBlockK, - saved_down_unweighted.get_end_ptr()), - backward_wgrad_bf16_wide( + // 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, - workspace.num_max_pool_tokens + - deep_gemm::sm103_block128_wgrad::kBlockK, - backward_wgrad_bf16_narrow.get_end_ptr()) {} + capacity, + saved_down_unweighted.get_end_ptr()) {} int64_t num_bytes() const { return reinterpret_cast( - backward_wgrad_bf16_wide.get_end_ptr()) - + backward_symmetric_grad_y.get_end_ptr()) - reinterpret_cast(workspace.base); } }; @@ -2277,11 +2268,9 @@ void launch_persistent_backward_activation( sym_buffer, layout.workspace, reinterpret_cast( grad_output.data_ptr()), - // Reverse transport needs the same independent input/output - // lifetimes as upstream forward. Reuse the not-yet-live private - // wgrad-wide backing as the immutable symmetric grad-y plane; the - // normal combine plane remains the remote dX destination. - layout.backward_wgrad_bf16_wide + // 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 @@ -2316,57 +2305,23 @@ void launch_persistent_wgrad( const PersistentWorkspaceLayout& layout, const torch::Tensor& expert_counts ) { - constexpr int64_t shape_m = - kW2 ? kPersistentHidden : 2 * kPersistentIntermediate; - constexpr int64_t shape_n = - kW2 ? kPersistentIntermediate : kPersistentHidden; constexpr int64_t output_rows = kW2 ? kPersistentHidden : kPersistentIntermediate; constexpr int64_t output_columns = kW2 ? kPersistentIntermediate : kPersistentHidden; const int64_t local_experts = kPersistentExperts / kNumRanks; - const auto bf16_options = torch::TensorOptions() - .dtype(torch::kBFloat16) - .device(buffer.device()); 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; - void* cached_a_base = kW2 - ? layout.backward_wgrad_bf16_wide.base - : layout.backward_wgrad_bf16_narrow.base; - void* cached_b_base = kW2 - ? layout.backward_wgrad_bf16_narrow.base - : layout.backward_wgrad_bf16_wide.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 int64_t cached_rows = - static_cast(layout.workspace.num_max_pool_tokens) + - deep_gemm::sm103_block128_wgrad::kBlockK; - auto cached_a = torch::from_blob( - cached_a_base, {cached_rows, shape_m}, bf16_options); - auto cached_b = torch::from_blob( - cached_b_base, {cached_rows, shape_n}, bf16_options); - const auto tensor_map_a = deep_gemm::make_tma_2d_desc( - cached_a, - static_cast(shape_m), - static_cast(cached_rows), - deep_gemm::sm103_block128_wgrad::kBlockM, - deep_gemm::sm103_block128_wgrad::kBlockK, - static_cast(shape_m), 128); - const auto tensor_map_b = deep_gemm::make_tma_2d_desc( - cached_b, - static_cast(shape_n), - static_cast(cached_rows), - deep_gemm::sm103_block128_wgrad::kLoadBlockN, - deep_gemm::sm103_block128_wgrad::kBlockK, - static_cast(shape_n), 128); const auto output_0_flat = output_0.view( {local_experts * output_rows, output_columns}); const auto output_1_flat = output_1.view( @@ -2413,15 +2368,10 @@ void launch_persistent_wgrad( &config, kernel, expert_counts.data_ptr(), layout.workspace.num_max_pool_tokens, - layout.workspace, static_cast(full_a_base), static_cast(full_b_base), full_a_sf, full_b_sf, - layout.backward_full_scores.get_base_ptr(), - static_cast(cached_a_base), - static_cast(cached_b_base), - tensor_map_a, tensor_map_b, tensor_map_output_0, tensor_map_output_1)); C10_CUDA_KERNEL_LAUNCH_CHECK(); } 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 index d928ad537a..445b18680a 100644 --- 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 @@ -18,12 +18,10 @@ 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. Each dedicated wgrad kernel first converts -// its FP8 route operands exactly once into private BF16 backing inside this same -// persistent launch, grid-fences internally, and then runs the native BF16 TMA -// / UMMA pipeline. Dequantization is therefore part of the dedicated kernel's -// load prologue rather than repeated for every output tile or composed as a -// separate kernel. +// 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; @@ -33,9 +31,9 @@ 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 kTMAWarp = 0; +static constexpr uint32_t kLoadAWarp = 0; static constexpr uint32_t kMMAWarp = 1; -static constexpr uint32_t kReadyWarp = 2; +static constexpr uint32_t kLoadBWarp = 2; static constexpr uint32_t kControlWarp = 3; static constexpr uint32_t kEpilogueFirstWarp = 4; static constexpr uint32_t kEpilogueThreads = 128; @@ -63,9 +61,8 @@ struct alignas(1024) SharedStorage { [kStoreBlockM * kStoreBlockN]; alignas(1024) bf16_t smem_a[kStages][kBlockK * kBlockM]; alignas(1024) bf16_t smem_b[kStages][kBlockK * kLoadBlockN]; - Barrier tma_full_barriers[kStages]; - Barrier tma_empty_barriers[kStages]; - Barrier mma_full_barriers[kStages]; + Barrier load_full_barriers[kStages]; + Barrier load_empty_barriers[kStages]; Barrier tmem_full_barriers[kNumEpilogueStages]; Barrier tmem_empty_barriers[kNumEpilogueStages]; uint32_t tmem_ptr; @@ -127,94 +124,71 @@ CUTLASS_DEVICE uint32_t convert_fp8x2_power2_to_bf16x2( return lo | (hi << 16); } -template < - uint32_t kShape, uint32_t kLocalExperts, - uint32_t kNumSMs, uint32_t kNumThreads, - bool kApplyPostScale> -CUTLASS_DEVICE void dequantize_route_pool_once( - const int* expert_counts, - const uint32_t max_pool_tokens, +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 float* scores, - bf16_t* destination) { - constexpr uint32_t kValuesPerVector = 8; - constexpr uint32_t kVectorsPerRow = kShape / kValuesPerVector; + 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"); - const uint64_t global_thread = - static_cast(blockIdx.x) * kNumThreads + threadIdx.x; - constexpr uint64_t kGridThreads = - static_cast(kNumSMs) * kNumThreads; - uint32_t pool_row = 0; - + auto* destination_bytes = reinterpret_cast(destination); #pragma unroll 1 - for (uint32_t expert = 0; expert < kLocalExperts; ++ expert) { - const uint32_t count = static_cast( - __ldg(expert_counts + expert)); - const uint32_t padded_count = - math::ceil_div(count, kRouteBlockM) * kRouteBlockM; - const uint64_t num_vectors = - static_cast(padded_count) * kVectorsPerRow; - for (uint64_t linear = global_thread; - linear < num_vectors; linear += kGridThreads) { - const uint32_t route = static_cast( - linear / kVectorsPerRow); - const uint32_t vector_in_row = static_cast( - linear - static_cast(route) * kVectorsPerRow); - const uint32_t feature = - vector_in_row * kValuesPerVector; - const uint64_t full_row = - static_cast(pool_row) + route; - uint4 packed{}; - if (route < count) { - const uint2 raw = *reinterpret_cast( - source + full_row * kShape + feature); - const uint32_t raw_words[2] = {raw.x, raw.y}; - const uint32_t scale_exponent = scales[ - full_row * kScaleBlocksPerRow + feature / 128] & 0xffu; - const float post_scale = kApplyPostScale - ? scores[full_row] - : 1.0f; - auto* output_pairs = reinterpret_cast(&packed); - #pragma unroll - for (uint32_t pair = 0; pair < 4; ++ pair) { - const uint16_t fp8x2 = static_cast( - raw_words[pair / 2] >> ((pair & 1u) * 16)); - uint32_t bf16x2 = convert_fp8x2_power2_to_bf16x2( - fp8x2, scale_exponent); - if constexpr (kApplyPostScale) { - const auto dequantized = __bfloat1622float2( - *reinterpret_cast(&bf16x2)); - const auto scaled = __float22bfloat162_rn( - __fmul2_rn( - dequantized, - {post_scale, post_scale})); - bf16x2 = - *reinterpret_cast(&scaled); - } - output_pairs[pair] = bf16x2; - } - } - *reinterpret_cast( - destination + full_row * kShape + feature) = packed; + 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; } - pool_row += padded_count; - } - DG_DEVICE_ASSERT(pool_row <= max_pool_tokens); + scale_exponent = __shfl_sync( + 0xffffffffu, scale_exponent, 0); - // Empty experts consume one permanent all-zero K tile beyond the routed - // pool. The host-side private scratch descriptors include these rows. - constexpr uint64_t kZeroVectors = - static_cast(kBlockK) * kVectorsPerRow; - for (uint64_t linear = global_thread; - linear < kZeroVectors; linear += kGridThreads) { - *reinterpret_cast( - destination + - static_cast(max_pool_tokens) * kShape + - linear * kValuesPerVector) = {}; + #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)); + } + } } } @@ -292,41 +266,16 @@ 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 __grid_constant__ layout::Workspace workspace, const fp8_t* full_a, const fp8_t* full_b, const uint32_t* full_a_sf, const uint32_t* full_b_sf, - const float* full_scores, - bf16_t* cached_a, - bf16_t* cached_b, - const __grid_constant__ cute::TmaDescriptor tensor_map_a, - const __grid_constant__ cute::TmaDescriptor tensor_map_b, 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; - constexpr uint32_t kLocalExperts = Scheduler::kLocalExperts; - - // One fused prologue per dedicated wgrad launch. This removes conversion - // from the output-tile loop while preserving the exact E4M3 + FP32 - // power-of-two scale and BF16-rounding semantics. - // Reverse dispatch has already formed and power-of-two quantized dz = - // BF16(dy * route_score) for W2. Dequantize it exactly once here; applying - // the route score again would square the router derivative. - dequantize_route_pool_once< - kShapeM, kLocalExperts, kNumSMs, kThreads, false>( - expert_counts, max_pool_tokens, - full_a, full_a_sf, full_scores, cached_a); - dequantize_route_pool_once< - kShapeN, kLocalExperts, kNumSMs, kThreads, false>( - expert_counts, max_pool_tokens, - full_b, full_b_sf, full_scores, cached_b); - comm::grid_sync( - workspace, blockIdx.x, threadIdx.x, - []() { __syncthreads(); }); extern __shared__ __align__(1024) uint8_t smem_buffer[]; SharedStorage& storage = *reinterpret_cast(smem_buffer); @@ -335,9 +284,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( const bool leader_cta = cute::block_rank_in_cluster() == 0; const uint32_t cta_rank = cute::block_rank_in_cluster(); - if (warp_idx == kTMAWarp) { - cute::prefetch_tma_descriptor(&tensor_map_a); - cute::prefetch_tma_descriptor(&tensor_map_b); + if (warp_idx == kLoadAWarp) { cute::prefetch_tma_descriptor(&tensor_map_output_0); cute::prefetch_tma_descriptor(&tensor_map_output_1); } @@ -346,10 +293,10 @@ sm103_fp8_block128_mega_moe_wgrad_impl( if (warp_idx == kControlWarp && cute::elect_one_sync()) { #pragma unroll for (uint32_t i = 0; i < kStages; ++i) { - storage.tma_full_barriers[i].init(1); - storage.tma_empty_barriers[i].init(1); - // Both CTAs publish their direct-BF16 TMA completion to CTA 0. - storage.mma_full_barriers[i].init(2); + // 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) { @@ -395,9 +342,11 @@ sm103_fp8_block128_mega_moe_wgrad_impl( phase ^= stage_idx == 0; }; - if (warp_idx == kTMAWarp && cute::elect_one_sync()) { - // The production load warp now reads the once-dequantized BF16 backing - // directly into the native UMMA swizzle. + 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( @@ -407,46 +356,36 @@ sm103_fp8_block128_mega_moe_wgrad_impl( cute::max(1u, math::ceil_div(count, kBlockK)); for (uint32_t k_block = 0; k_block < num_k_blocks; ++k_block) { - storage.tma_empty_barriers[stage_idx].wait(phase ^ 1u); - const uint32_t route = count == 0 - ? max_pool_tokens - : pool_row + k_block * kBlockK; - const uint32_t a_feature = m_block * kBlockM; - const uint32_t b_feature = - n_block * kBlockN + cta_rank * kLoadBlockN; - tma::copy( - &tensor_map_a, - &storage.tma_full_barriers[stage_idx], - storage.smem_a[stage_idx], - a_feature, route); - tma::copy( - &tensor_map_b, - &storage.tma_full_barriers[stage_idx], - storage.smem_b[stage_idx], - b_feature, route); - storage.tma_full_barriers[stage_idx] - .arrive_and_expect_tx( - sizeof(storage.smem_a[0]) + - sizeof(storage.smem_b[0])); - advance_pipeline(); - } - } - } else if (warp_idx == kReadyWarp && cute::elect_one_sync()) { - // Each CTA waits for its local direct-BF16 TMAs, then contributes one - // arrival to CTA 0. The leader MMA warp therefore observes both halves - // without multicast-copying different GLM feature tiles over each - // other. - 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)) { - 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.tma_full_barriers[stage_idx].wait(phase); + 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(); - storage.mma_full_barriers[stage_idx].arrive(0u); + __syncwarp(); + if (cute::elect_one_sync()) + storage.load_full_barriers[stage_idx].arrive(0u); advance_pipeline(); } } @@ -469,7 +408,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( cute::max(1u, math::ceil_div(count, kBlockK)); for (uint32_t k_block = 0; k_block < num_k_blocks; ++k_block) { - storage.mma_full_barriers[stage_idx].wait(phase); + 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); @@ -497,7 +436,7 @@ sm103_fp8_block128_mega_moe_wgrad_impl( constexpr uint16_t kCTAMask = 3; cutlass::arch::umma_arrive_multicast_2x1SM( reinterpret_cast( - &storage.tma_empty_barriers[stage_idx]), + &storage.load_empty_barriers[stage_idx]), kCTAMask); if (k_block == num_k_blocks - 1) { cutlass::arch::umma_arrive_multicast_2x1SM(