Skip to content
KernelIndex
Search⌘K

submission 379258

rt11 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 3741 lines, June 9 Researcher Reciprocity License v1.0.

v_ai_8_2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-379258?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
15.6µs
#80 of 161
2026-01-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b12810d9dab237a026c4f47c282d561b50f27810b39c6b56c50c9a84357bb969
license declaredunknown
license concludedunknown
authorsrt11
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

cluster__cluster_dims__(CLUSTER_M, 1, 1)
fp4Provenance: copied from `nvfp4/dual_gemm/submissions/0101/v2/v_ai_7.py`.
fused-epilogue- tcgen05.ld to read accumulators and fuse SiLU+mul in epilogue
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
shared-memory__device__ inline void tma_gmem2smem_multicast(
stages = 5constexpr int NUM_STAGES = 5;
tcgen05- tcgen05.cp to move block scale factors to TMEM
tile-k = 256constexpr int BLOCK_K = 256;
tile-m = 128constexpr int BLOCK_M = 128;
tile-n = 64constexpr int BLOCK_N = 64;
tma- TMA (cuTensorMap + cp.async.bulk.tensor) to stage FP4 tiles
vector-width = half2reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});

Kernel source

v_ai_8_2.py3741 lines
#!POPCORN leaderboard modal_nvfp4_dual_gemm
#!POPCORN gpu B200

"""
Provenance: copied from `nvfp4/dual_gemm/submissions/0101/v2/v_ai_7.py`.
Change: inline fully standalone `cpp_src_cfg1`/`cuda_src_cfg1` ... `cpp_src_cfg4`/`cuda_src_cfg4` literals (no templates / no `.replace()` generation), and keep each CUDA source single-shape in `run_one_cfg` (no `else if (m == ...)` / `else if (k == ...)` dispatch).

Fused NVFP4 dual GEMM for B200 (SM100a), implemented as a raw C++/CUDA kernel.

Computation:
  C = silu(A @ B1) * (A @ B2)

This version intentionally avoids the Python CuTe DSL and follows the same low-level
approach as `nvfp4/gemm/top/ranked/r001/submission.py`:
- TMA (cuTensorMap + cp.async.bulk.tensor) to stage FP4 tiles
- tcgen05.cp to move block scale factors to TMEM
- tcgen05.mma mxf4nvf4 block-scaled MMA
- tcgen05.ld to read accumulators and fuse SiLU+mul in epilogue

v_z is a per-shape hybrid:
- m=256: reuse v_i's cluster-multicast (CLUSTER_M=2) path to speed up small-M shapes
- m=512: reuse v_t's merged-B MMA path (MMA_N=256) to reduce MMA instruction count
"""

from __future__ import annotations

import hashlib
import os

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

# Compile for SM100a (B200).
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0a")

# Per-leaderboard-config sources (standalone literals).
# Each cuda_src contains a single-shape run_one_cfg (no runtime m/k dispatch).

# cfg1: m=256 n=4096 k=7168
cpp_src_cfg1 = r"""
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    // cfg: cfg1 (m256_n4096_k7168)
    int nvfp4_dual_gemm_fused_run(
        uint64_t a_ptr,
        uint64_t b1_ptr,
        uint64_t b2_ptr,
        uint64_t sfa_ptr,
        uint64_t sfb1_ptr,
        uint64_t sfb2_ptr,
        uint64_t out_ptr,
        int m, int n, int k, int l);

    std::tuple<int, int> nvfp4_dual_gemm_last_error();
    """

cuda_src_cfg1 = r"""
    #include <cuda.h>
    #include <cudaTypedefs.h>
    #include <cuda_fp16.h>
    #include <cuda_runtime.h>
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    constexpr int WARP_SIZE = 32;
    constexpr int MMA_K = 64;  // 32 bytes of FP4 (packed)

    // Cache policy hints (same as r001 solution).
    constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
    constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000;
    constexpr uint64_t EVICT_LAST   = 0x14F0000000000000;

    namespace {
    static int g_last_cuda_error = int(cudaSuccess);
    }  // namespace

    __device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
    __device__ uint32_t elect_sync() {
      uint32_t pred = 0;
      asm volatile(
        "{\n\t"
        ".reg .pred %%px;\n\t"
        "elect.sync _|%%px, %1;\n\t"
        "@%%px mov.s32 %0, 1;\n\t"
        "}"
        : "+r"(pred)
        : "r"(0xFFFFFFFF)
      );
      return pred;
    }

    __device__ inline void mbarrier_init(int mbar_addr, int count) {
      asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
    }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
    __device__ void mbarrier_wait(int mbar_addr, int phase) {
      uint32_t ticks = 0x989680;  // optional
      asm volatile(
        "{\n\t"
        ".reg .pred P1;\n\t"
        "LAB_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
        "@P1 bra.uni DONE;\n\t"
        "bra.uni LAB_WAIT;\n\t"
        "DONE:\n\t"
        "}"
        :: "r"(mbar_addr), "r"(phase), "r"(ticks)
      );
    }

    __device__ inline void tma_gmem2smem(int dst, const void* src, int size, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
      );
    }

    __device__ inline void tma_gmem2smem_multicast(
      int dst,
      const void* src,
      int size,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4, %5;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem(int dst, const void* tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem_multicast(
      int dst,
      const void* tmap_ptr,
      int x,
      int y,
      int z,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6, %7;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
      asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
    }

    __device__ inline void tcgen05_mma_nvfp4_single(
      int d_tmem,
      uint64_t a_desc,
      uint64_t b_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B_tmem,
      int enable_input_d
    ) {
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
        "[%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
      );
    }

    __device__ inline void tcgen05_mma_nvfp4_dualA(
      int d1_tmem,
      int d2_tmem,
      uint64_t a_desc,
      uint64_t b1_desc,
      uint64_t b2_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B1_tmem,
      int scale_B2_tmem,
      int enable_input_d
    ) {
      // Reuse A across the two MMAs via the TensorCore collector buffer.
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %9, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
        "[%0], %2, %3, %5, [%6], [%7], p;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
        "[%1], %2, %4, %5, [%6], [%8], p;\n\t"
        "}"
        :: "r"(d1_tmem), "r"(d2_tmem),
           "l"(a_desc), "l"(b1_desc), "l"(b2_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B1_tmem), "r"(scale_B2_tmem), "r"(enable_input_d)
      );
    }

    struct SHAPE {
      static constexpr char _16x256b[] = ".16x256b";
    };
    struct NUM {
      static constexpr char x4[] = ".x4";
      static constexpr char x8[] = ".x8";
    };

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_16regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%17%18.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15}, [%16];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_32regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%33%34.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15, "
        " %16, %17, %18, %19, %20, %21, %22, %23, "
        " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
          "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
          "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    __device__ inline void tcgen05_ld_16x256bx4(float* tmp, int row, int col) {
      tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col);
    }

    __device__ inline void tcgen05_ld_16x256bx8(float* tmp, int row, int col) {
      tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
    }

    void check_cu(CUresult err) {
      if (err == CUDA_SUCCESS) return;
      const char* error_msg_ptr = nullptr;
      if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS) {
        error_msg_ptr = "unable to get error string";
      }
      TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
    }

    void init_AB_tmap(
      CUtensorMap* tmap,
      const char* ptr,
      uint64_t global_height, uint64_t global_width,
      uint32_t shared_height, uint32_t shared_width,
      CUtensorMapL2promotion l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE
    ) {
      constexpr uint32_t rank = 3;
      uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
      uint64_t globalStrides[rank-1] = {global_width / 2, 128};  // bytes
      uint32_t boxDim[rank]          = {256, shared_height, shared_width / 256};
      uint32_t elementStrides[rank]  = {1, 1, 1};

      auto err = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        (void*)ptr,
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
        l2_promotion,
        CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
      );
      check_cu(err);
    }

    __device__ __forceinline__ float silu(float x) {
      // silu(x) = x / (1 + exp(-x))
      return x / (1.0f + __expf(-x));
    }

    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_M>
    __global__
    __cluster_dims__(CLUSTER_M, 1, 1)
    __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      static_assert(CLUSTER_M >= 1 && CLUSTER_M <= 16);
      constexpr uint16_t cta_mask = uint16_t((1u << CLUSTER_M) - 1u);
      int cta_rank = 0;
      asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      // BlockIdx linearization: run along M first so a cluster spans the full M-slab for a fixed N-tile.
      const int bid_n = bid / grid_m;
      const int bid_m = bid % grid_m;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      // Dynamic shared memory. Stage layout is per r001, extended for B1/B2 and SFB1/SFB2.
      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      // mbarriers: NUM_STAGES for TMA, NUM_STAGES for MMA, 1 for mainloop.
      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      // TMEM layout:
      //   ACC1: [0 .. BLOCK_N-1]
      //   ACC2: [BLOCK_N .. 2*BLOCK_N-1]
      //   SFA:  [2*BLOCK_N .. 2*BLOCK_N+SF_COLS-1]
      //   SFB1: next SF_COLS
      //   SFB2: next SF_COLS
      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
          // Cluster-wide stage reuse sync (B/B2 are multicast to the whole cluster).
          mbarrier_init(mma_mbar_addr + i * 8, CLUSTER_M);
        }
        mbarrier_init(mainloop_mbar_addr, 1);
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        // Allocate TMEM (address is assumed 0).
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      if constexpr (CLUSTER_M > 1) {
        asm volatile("barrier.cluster.arrive.release.aligned;" ::: "memory");
        asm volatile("barrier.cluster.wait.acquire.aligned;" ::: "memory");
      }
      else {
        __syncthreads();
      }

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        // TMA warp.
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_3d_gmem2smem_multicast(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
              tma_3d_gmem2smem_multicast(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          }

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_gmem2smem_multicast(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cta_mask, cache_B);
              tma_gmem2smem_multicast(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
            tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);
          }

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        // Prologue: fill pipeline.
        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }

        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        // MMA warp.
        constexpr int MMA_N = BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
            tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);
              const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
              const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;  // 4 = 256/MMA_K
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
              const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
              const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4_dualA(
                ACC1_TMEM, ACC2_TMEM,
                a_desc, b1_desc, b2_desc, i_desc,
                scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                enable_input_d
              );
            }
          }

          if constexpr (CLUSTER_M > 1) {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
              :: "r"(mma_mbar_addr + stage_id * 8), "h"(cta_mask)
              : "memory"
            );
          }
          else {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
              :: "r"(mma_mbar_addr + stage_id * 8)
              : "memory"
            );
          }
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        // Epilogue threads: fuse silu(acc1) * acc2 and store fp16.
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        // N-major output with half2 stores. Process in 64-column chunks to cap registers.
        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    // Non-cluster kernel for m==512 path: reuse v_t's merged-B MMA (MMA_N=256 when BLOCK_N==128).
    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
    __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel_merged(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      const int bid_m = bid / grid_n;
      const int bid_n = bid % grid_n;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES * 2 + 1; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
        }
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      __syncthreads();

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
          tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }
        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        constexpr bool MERGED_B = (BLOCK_N == 128);
        constexpr int MMA_N = MERGED_B ? 2 * BLOCK_N : BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            if constexpr (MERGED_B) {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8, sfb1_desc);
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8 + 4, sfb2_desc);
            }
            else {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
              tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
            }
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              if constexpr (MERGED_B) {
                const uint64_t b_desc = make_desc_AB(B1_smem + k1 * MMA_N * 128 + k2 * 32);
                const int scale_B_tmem = SFB1_TMEM + k_sf * 8;
                tcgen05_mma_nvfp4_single(
                  ACC1_TMEM,
                  a_desc, b_desc, i_desc,
                  scale_A_tmem, scale_B_tmem,
                  enable_input_d
                );
              }
              else {
                const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
                const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
                const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;
                tcgen05_mma_nvfp4_dualA(
                  ACC1_TMEM, ACC2_TMEM,
                  a_desc, b1_desc, b2_desc, i_desc,
                  scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                  enable_input_d
                );
              }
            }
          }

          asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
            :: "r"(mma_mbar_addr + stage_id * 8)
            : "memory"
          );
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    static int run_one_cfg(
      int k,
      const void* a_ptr,
      const void* b1_ptr,
      const void* b2_ptr,
      const void* sfa_ptr,
      const void* sfb1_ptr,
      const void* sfb2_ptr,
      void* out_ptr,
      int m, int n
    ) {
      // cfg1: m=256 n=4096 k=7168
      constexpr int EXPECT_M = 256;
      constexpr int EXPECT_N = 4096;
      constexpr int EXPECT_K = 7168;
      if (m != EXPECT_M || n != EXPECT_N || k != EXPECT_K) return -13;

      constexpr int BLOCK_M = 128;
      constexpr int BLOCK_N = 64;
      constexpr int BLOCK_K = 256;
      constexpr int NUM_STAGES = 5;
      constexpr int CLUSTER_M = 2;

      CUtensorMap A_tmap, B1_tmap, B2_tmap;
      init_AB_tmap(&A_tmap,  reinterpret_cast<const char*>(a_ptr),  m, k, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B1_tmap, reinterpret_cast<const char*>(b1_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);
      init_AB_tmap(&B2_tmap, reinterpret_cast<const char*>(b2_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);

      const int grid = (m / BLOCK_M) * (n / BLOCK_N);
      const int tb_size = BLOCK_M + 2 * WARP_SIZE;

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
      const int smem_size = STAGE_SIZE * NUM_STAGES;

      auto this_kernel = dual_kernel<EXPECT_K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, CLUSTER_M>;
      if (smem_size > 48'000) cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      this_kernel<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap,
        reinterpret_cast<const char*>(sfa_ptr),
        reinterpret_cast<const char*>(sfb1_ptr),
        reinterpret_cast<const char*>(sfb2_ptr),
        reinterpret_cast<half*>(out_ptr),
        m, n
      );

      auto err = cudaGetLastError();
      g_last_cuda_error = int(err);
      return (err == cudaSuccess) ? 0 : -20;
    }

    int nvfp4_dual_gemm_fused_run(
      uint64_t a_ptr,
      uint64_t b1_ptr,
      uint64_t b2_ptr,
      uint64_t sfa_ptr,
      uint64_t sfb1_ptr,
      uint64_t sfb2_ptr,
      uint64_t out_ptr,
      int m, int n, int k, int l
    ) {
      (void)l;  // only l=1 fast path in Python
      g_last_cuda_error = int(cudaSuccess);

      return run_one_cfg(
        k,
        reinterpret_cast<const void*>(a_ptr),
        reinterpret_cast<const void*>(b1_ptr),
        reinterpret_cast<const void*>(b2_ptr),
        reinterpret_cast<const void*>(sfa_ptr),
        reinterpret_cast<const void*>(sfb1_ptr),
        reinterpret_cast<const void*>(sfb2_ptr),
        reinterpret_cast<void*>(out_ptr),
        m, n
      );
    }

    std::tuple<int, int> nvfp4_dual_gemm_last_error() {
      return {0, g_last_cuda_error};
    }
    """

# cfg2: m=512 n=4096 k=7168
cpp_src_cfg2 = r"""
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    // cfg: cfg2 (m512_n4096_k7168)
    int nvfp4_dual_gemm_fused_run(
        uint64_t a_ptr,
        uint64_t b1_ptr,
        uint64_t b2_ptr,
        uint64_t sfa_ptr,
        uint64_t sfb1_ptr,
        uint64_t sfb2_ptr,
        uint64_t out_ptr,
        int m, int n, int k, int l);

    std::tuple<int, int> nvfp4_dual_gemm_last_error();
    """

cuda_src_cfg2 = r"""
    #include <cuda.h>
    #include <cudaTypedefs.h>
    #include <cuda_fp16.h>
    #include <cuda_runtime.h>
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    constexpr int WARP_SIZE = 32;
    constexpr int MMA_K = 64;  // 32 bytes of FP4 (packed)

    // Cache policy hints (same as r001 solution).
    constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
    constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000;
    constexpr uint64_t EVICT_LAST   = 0x14F0000000000000;

    namespace {
    static int g_last_cuda_error = int(cudaSuccess);
    }  // namespace

    __device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
    __device__ uint32_t elect_sync() {
      uint32_t pred = 0;
      asm volatile(
        "{\n\t"
        ".reg .pred %%px;\n\t"
        "elect.sync _|%%px, %1;\n\t"
        "@%%px mov.s32 %0, 1;\n\t"
        "}"
        : "+r"(pred)
        : "r"(0xFFFFFFFF)
      );
      return pred;
    }

    __device__ inline void mbarrier_init(int mbar_addr, int count) {
      asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
    }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
    __device__ void mbarrier_wait(int mbar_addr, int phase) {
      uint32_t ticks = 0x989680;  // optional
      asm volatile(
        "{\n\t"
        ".reg .pred P1;\n\t"
        "LAB_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
        "@P1 bra.uni DONE;\n\t"
        "bra.uni LAB_WAIT;\n\t"
        "DONE:\n\t"
        "}"
        :: "r"(mbar_addr), "r"(phase), "r"(ticks)
      );
    }

    __device__ inline void tma_gmem2smem(int dst, const void* src, int size, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
      );
    }

    __device__ inline void tma_gmem2smem_multicast(
      int dst,
      const void* src,
      int size,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4, %5;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem(int dst, const void* tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem_multicast(
      int dst,
      const void* tmap_ptr,
      int x,
      int y,
      int z,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6, %7;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
      asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
    }

    __device__ inline void tcgen05_mma_nvfp4_single(
      int d_tmem,
      uint64_t a_desc,
      uint64_t b_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B_tmem,
      int enable_input_d
    ) {
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
        "[%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
      );
    }

    __device__ inline void tcgen05_mma_nvfp4_dualA(
      int d1_tmem,
      int d2_tmem,
      uint64_t a_desc,
      uint64_t b1_desc,
      uint64_t b2_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B1_tmem,
      int scale_B2_tmem,
      int enable_input_d
    ) {
      // Reuse A across the two MMAs via the TensorCore collector buffer.
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %9, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
        "[%0], %2, %3, %5, [%6], [%7], p;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
        "[%1], %2, %4, %5, [%6], [%8], p;\n\t"
        "}"
        :: "r"(d1_tmem), "r"(d2_tmem),
           "l"(a_desc), "l"(b1_desc), "l"(b2_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B1_tmem), "r"(scale_B2_tmem), "r"(enable_input_d)
      );
    }

    struct SHAPE {
      static constexpr char _16x256b[] = ".16x256b";
    };
    struct NUM {
      static constexpr char x4[] = ".x4";
      static constexpr char x8[] = ".x8";
    };

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_16regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%17%18.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15}, [%16];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_32regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%33%34.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15, "
        " %16, %17, %18, %19, %20, %21, %22, %23, "
        " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
          "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
          "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    __device__ inline void tcgen05_ld_16x256bx4(float* tmp, int row, int col) {
      tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col);
    }

    __device__ inline void tcgen05_ld_16x256bx8(float* tmp, int row, int col) {
      tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
    }

    void check_cu(CUresult err) {
      if (err == CUDA_SUCCESS) return;
      const char* error_msg_ptr = nullptr;
      if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS) {
        error_msg_ptr = "unable to get error string";
      }
      TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
    }

    void init_AB_tmap(
      CUtensorMap* tmap,
      const char* ptr,
      uint64_t global_height, uint64_t global_width,
      uint32_t shared_height, uint32_t shared_width,
      CUtensorMapL2promotion l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE
    ) {
      constexpr uint32_t rank = 3;
      uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
      uint64_t globalStrides[rank-1] = {global_width / 2, 128};  // bytes
      uint32_t boxDim[rank]          = {256, shared_height, shared_width / 256};
      uint32_t elementStrides[rank]  = {1, 1, 1};

      auto err = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        (void*)ptr,
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
        l2_promotion,
        CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
      );
      check_cu(err);
    }

    __device__ __forceinline__ float silu(float x) {
      // silu(x) = x / (1 + exp(-x))
      return x / (1.0f + __expf(-x));
    }

    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_M>
    __global__
    __cluster_dims__(CLUSTER_M, 1, 1)
    __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      static_assert(CLUSTER_M >= 1 && CLUSTER_M <= 16);
      constexpr uint16_t cta_mask = uint16_t((1u << CLUSTER_M) - 1u);
      int cta_rank = 0;
      asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      // BlockIdx linearization: run along M first so a cluster spans the full M-slab for a fixed N-tile.
      const int bid_n = bid / grid_m;
      const int bid_m = bid % grid_m;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      // Dynamic shared memory. Stage layout is per r001, extended for B1/B2 and SFB1/SFB2.
      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      // mbarriers: NUM_STAGES for TMA, NUM_STAGES for MMA, 1 for mainloop.
      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      // TMEM layout:
      //   ACC1: [0 .. BLOCK_N-1]
      //   ACC2: [BLOCK_N .. 2*BLOCK_N-1]
      //   SFA:  [2*BLOCK_N .. 2*BLOCK_N+SF_COLS-1]
      //   SFB1: next SF_COLS
      //   SFB2: next SF_COLS
      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
          // Cluster-wide stage reuse sync (B/B2 are multicast to the whole cluster).
          mbarrier_init(mma_mbar_addr + i * 8, CLUSTER_M);
        }
        mbarrier_init(mainloop_mbar_addr, 1);
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        // Allocate TMEM (address is assumed 0).
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      if constexpr (CLUSTER_M > 1) {
        asm volatile("barrier.cluster.arrive.release.aligned;" ::: "memory");
        asm volatile("barrier.cluster.wait.acquire.aligned;" ::: "memory");
      }
      else {
        __syncthreads();
      }

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        // TMA warp.
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_3d_gmem2smem_multicast(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
              tma_3d_gmem2smem_multicast(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          }

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_gmem2smem_multicast(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cta_mask, cache_B);
              tma_gmem2smem_multicast(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
            tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);
          }

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        // Prologue: fill pipeline.
        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }

        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        // MMA warp.
        constexpr int MMA_N = BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
            tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);
              const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
              const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;  // 4 = 256/MMA_K
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
              const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
              const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4_dualA(
                ACC1_TMEM, ACC2_TMEM,
                a_desc, b1_desc, b2_desc, i_desc,
                scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                enable_input_d
              );
            }
          }

          if constexpr (CLUSTER_M > 1) {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
              :: "r"(mma_mbar_addr + stage_id * 8), "h"(cta_mask)
              : "memory"
            );
          }
          else {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
              :: "r"(mma_mbar_addr + stage_id * 8)
              : "memory"
            );
          }
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        // Epilogue threads: fuse silu(acc1) * acc2 and store fp16.
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        // N-major output with half2 stores. Process in 64-column chunks to cap registers.
        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    // Non-cluster kernel for m==512 path: reuse v_t's merged-B MMA (MMA_N=256 when BLOCK_N==128).
    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
    __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel_merged(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      const int bid_m = bid / grid_n;
      const int bid_n = bid % grid_n;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES * 2 + 1; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
        }
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      __syncthreads();

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
          tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }
        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        constexpr bool MERGED_B = (BLOCK_N == 128);
        constexpr int MMA_N = MERGED_B ? 2 * BLOCK_N : BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            if constexpr (MERGED_B) {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8, sfb1_desc);
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8 + 4, sfb2_desc);
            }
            else {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
              tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
            }
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              if constexpr (MERGED_B) {
                const uint64_t b_desc = make_desc_AB(B1_smem + k1 * MMA_N * 128 + k2 * 32);
                const int scale_B_tmem = SFB1_TMEM + k_sf * 8;
                tcgen05_mma_nvfp4_single(
                  ACC1_TMEM,
                  a_desc, b_desc, i_desc,
                  scale_A_tmem, scale_B_tmem,
                  enable_input_d
                );
              }
              else {
                const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
                const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
                const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;
                tcgen05_mma_nvfp4_dualA(
                  ACC1_TMEM, ACC2_TMEM,
                  a_desc, b1_desc, b2_desc, i_desc,
                  scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                  enable_input_d
                );
              }
            }
          }

          asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
            :: "r"(mma_mbar_addr + stage_id * 8)
            : "memory"
          );
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    static int run_one_cfg(
      int k,
      const void* a_ptr,
      const void* b1_ptr,
      const void* b2_ptr,
      const void* sfa_ptr,
      const void* sfb1_ptr,
      const void* sfb2_ptr,
      void* out_ptr,
      int m, int n
    ) {
      // cfg2: m=512 n=4096 k=7168
      constexpr int EXPECT_M = 512;
      constexpr int EXPECT_N = 4096;
      constexpr int EXPECT_K = 7168;
      if (m != EXPECT_M || n != EXPECT_N || k != EXPECT_K) return -13;

      constexpr int BLOCK_M = 128;
      constexpr int BLOCK_N = 128;
      constexpr int BLOCK_K = 256;
      constexpr int NUM_STAGES = 4;

      CUtensorMap A_tmap, B1_tmap, B2_tmap;
      init_AB_tmap(&A_tmap,  reinterpret_cast<const char*>(a_ptr),  m, k, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B1_tmap, reinterpret_cast<const char*>(b1_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);
      init_AB_tmap(&B2_tmap, reinterpret_cast<const char*>(b2_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);

      const int grid = (m / BLOCK_M) * (n / BLOCK_N);
      const int tb_size = BLOCK_M + 2 * WARP_SIZE;

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
      const int smem_size = STAGE_SIZE * NUM_STAGES;

      auto this_kernel = dual_kernel_merged<EXPECT_K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
      if (smem_size > 48'000) cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      this_kernel<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap,
        reinterpret_cast<const char*>(sfa_ptr),
        reinterpret_cast<const char*>(sfb1_ptr),
        reinterpret_cast<const char*>(sfb2_ptr),
        reinterpret_cast<half*>(out_ptr),
        m, n
      );

      auto err = cudaGetLastError();
      g_last_cuda_error = int(err);
      return (err == cudaSuccess) ? 0 : -20;
    }

    int nvfp4_dual_gemm_fused_run(
      uint64_t a_ptr,
      uint64_t b1_ptr,
      uint64_t b2_ptr,
      uint64_t sfa_ptr,
      uint64_t sfb1_ptr,
      uint64_t sfb2_ptr,
      uint64_t out_ptr,
      int m, int n, int k, int l
    ) {
      (void)l;  // only l=1 fast path in Python
      g_last_cuda_error = int(cudaSuccess);

      return run_one_cfg(
        k,
        reinterpret_cast<const void*>(a_ptr),
        reinterpret_cast<const void*>(b1_ptr),
        reinterpret_cast<const void*>(b2_ptr),
        reinterpret_cast<const void*>(sfa_ptr),
        reinterpret_cast<const void*>(sfb1_ptr),
        reinterpret_cast<const void*>(sfb2_ptr),
        reinterpret_cast<void*>(out_ptr),
        m, n
      );
    }

    std::tuple<int, int> nvfp4_dual_gemm_last_error() {
      return {0, g_last_cuda_error};
    }
    """

# cfg3: m=256 n=3072 k=4096
cpp_src_cfg3 = r"""
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    // cfg: cfg3 (m256_n3072_k4096)
    int nvfp4_dual_gemm_fused_run(
        uint64_t a_ptr,
        uint64_t b1_ptr,
        uint64_t b2_ptr,
        uint64_t sfa_ptr,
        uint64_t sfb1_ptr,
        uint64_t sfb2_ptr,
        uint64_t out_ptr,
        int m, int n, int k, int l);

    std::tuple<int, int> nvfp4_dual_gemm_last_error();
    """

cuda_src_cfg3 = r"""
    #include <cuda.h>
    #include <cudaTypedefs.h>
    #include <cuda_fp16.h>
    #include <cuda_runtime.h>
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    constexpr int WARP_SIZE = 32;
    constexpr int MMA_K = 64;  // 32 bytes of FP4 (packed)

    // Cache policy hints (same as r001 solution).
    constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
    constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000;
    constexpr uint64_t EVICT_LAST   = 0x14F0000000000000;

    namespace {
    static int g_last_cuda_error = int(cudaSuccess);
    }  // namespace

    __device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
    __device__ uint32_t elect_sync() {
      uint32_t pred = 0;
      asm volatile(
        "{\n\t"
        ".reg .pred %%px;\n\t"
        "elect.sync _|%%px, %1;\n\t"
        "@%%px mov.s32 %0, 1;\n\t"
        "}"
        : "+r"(pred)
        : "r"(0xFFFFFFFF)
      );
      return pred;
    }

    __device__ inline void mbarrier_init(int mbar_addr, int count) {
      asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
    }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
    __device__ void mbarrier_wait(int mbar_addr, int phase) {
      uint32_t ticks = 0x989680;  // optional
      asm volatile(
        "{\n\t"
        ".reg .pred P1;\n\t"
        "LAB_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
        "@P1 bra.uni DONE;\n\t"
        "bra.uni LAB_WAIT;\n\t"
        "DONE:\n\t"
        "}"
        :: "r"(mbar_addr), "r"(phase), "r"(ticks)
      );
    }

    __device__ inline void tma_gmem2smem(int dst, const void* src, int size, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
      );
    }

    __device__ inline void tma_gmem2smem_multicast(
      int dst,
      const void* src,
      int size,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4, %5;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem(int dst, const void* tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem_multicast(
      int dst,
      const void* tmap_ptr,
      int x,
      int y,
      int z,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6, %7;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
      asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
    }

    __device__ inline void tcgen05_mma_nvfp4_single(
      int d_tmem,
      uint64_t a_desc,
      uint64_t b_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B_tmem,
      int enable_input_d
    ) {
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
        "[%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
      );
    }

    __device__ inline void tcgen05_mma_nvfp4_dualA(
      int d1_tmem,
      int d2_tmem,
      uint64_t a_desc,
      uint64_t b1_desc,
      uint64_t b2_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B1_tmem,
      int scale_B2_tmem,
      int enable_input_d
    ) {
      // Reuse A across the two MMAs via the TensorCore collector buffer.
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %9, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
        "[%0], %2, %3, %5, [%6], [%7], p;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
        "[%1], %2, %4, %5, [%6], [%8], p;\n\t"
        "}"
        :: "r"(d1_tmem), "r"(d2_tmem),
           "l"(a_desc), "l"(b1_desc), "l"(b2_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B1_tmem), "r"(scale_B2_tmem), "r"(enable_input_d)
      );
    }

    struct SHAPE {
      static constexpr char _16x256b[] = ".16x256b";
    };
    struct NUM {
      static constexpr char x4[] = ".x4";
      static constexpr char x8[] = ".x8";
    };

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_16regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%17%18.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15}, [%16];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_32regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%33%34.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15, "
        " %16, %17, %18, %19, %20, %21, %22, %23, "
        " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
          "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
          "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    __device__ inline void tcgen05_ld_16x256bx4(float* tmp, int row, int col) {
      tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col);
    }

    __device__ inline void tcgen05_ld_16x256bx8(float* tmp, int row, int col) {
      tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
    }

    void check_cu(CUresult err) {
      if (err == CUDA_SUCCESS) return;
      const char* error_msg_ptr = nullptr;
      if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS) {
        error_msg_ptr = "unable to get error string";
      }
      TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
    }

    void init_AB_tmap(
      CUtensorMap* tmap,
      const char* ptr,
      uint64_t global_height, uint64_t global_width,
      uint32_t shared_height, uint32_t shared_width,
      CUtensorMapL2promotion l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE
    ) {
      constexpr uint32_t rank = 3;
      uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
      uint64_t globalStrides[rank-1] = {global_width / 2, 128};  // bytes
      uint32_t boxDim[rank]          = {256, shared_height, shared_width / 256};
      uint32_t elementStrides[rank]  = {1, 1, 1};

      auto err = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        (void*)ptr,
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
        l2_promotion,
        CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
      );
      check_cu(err);
    }

    __device__ __forceinline__ float silu(float x) {
      // silu(x) = x / (1 + exp(-x))
      return x / (1.0f + __expf(-x));
    }

    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_M>
    __global__
    __cluster_dims__(CLUSTER_M, 1, 1)
    __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      static_assert(CLUSTER_M >= 1 && CLUSTER_M <= 16);
      constexpr uint16_t cta_mask = uint16_t((1u << CLUSTER_M) - 1u);
      int cta_rank = 0;
      asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      // BlockIdx linearization: run along M first so a cluster spans the full M-slab for a fixed N-tile.
      const int bid_n = bid / grid_m;
      const int bid_m = bid % grid_m;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      // Dynamic shared memory. Stage layout is per r001, extended for B1/B2 and SFB1/SFB2.
      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      // mbarriers: NUM_STAGES for TMA, NUM_STAGES for MMA, 1 for mainloop.
      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      // TMEM layout:
      //   ACC1: [0 .. BLOCK_N-1]
      //   ACC2: [BLOCK_N .. 2*BLOCK_N-1]
      //   SFA:  [2*BLOCK_N .. 2*BLOCK_N+SF_COLS-1]
      //   SFB1: next SF_COLS
      //   SFB2: next SF_COLS
      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
          // Cluster-wide stage reuse sync (B/B2 are multicast to the whole cluster).
          mbarrier_init(mma_mbar_addr + i * 8, CLUSTER_M);
        }
        mbarrier_init(mainloop_mbar_addr, 1);
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        // Allocate TMEM (address is assumed 0).
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      if constexpr (CLUSTER_M > 1) {
        asm volatile("barrier.cluster.arrive.release.aligned;" ::: "memory");
        asm volatile("barrier.cluster.wait.acquire.aligned;" ::: "memory");
      }
      else {
        __syncthreads();
      }

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        // TMA warp.
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_3d_gmem2smem_multicast(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
              tma_3d_gmem2smem_multicast(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          }

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_gmem2smem_multicast(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cta_mask, cache_B);
              tma_gmem2smem_multicast(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
            tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);
          }

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        // Prologue: fill pipeline.
        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }

        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        // MMA warp.
        constexpr int MMA_N = BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
            tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);
              const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
              const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;  // 4 = 256/MMA_K
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
              const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
              const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4_dualA(
                ACC1_TMEM, ACC2_TMEM,
                a_desc, b1_desc, b2_desc, i_desc,
                scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                enable_input_d
              );
            }
          }

          if constexpr (CLUSTER_M > 1) {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
              :: "r"(mma_mbar_addr + stage_id * 8), "h"(cta_mask)
              : "memory"
            );
          }
          else {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
              :: "r"(mma_mbar_addr + stage_id * 8)
              : "memory"
            );
          }
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        // Epilogue threads: fuse silu(acc1) * acc2 and store fp16.
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        // N-major output with half2 stores. Process in 64-column chunks to cap registers.
        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    // Non-cluster kernel for m==512 path: reuse v_t's merged-B MMA (MMA_N=256 when BLOCK_N==128).
    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
    __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel_merged(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      const int bid_m = bid / grid_n;
      const int bid_n = bid % grid_n;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES * 2 + 1; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
        }
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      __syncthreads();

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
          tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }
        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        constexpr bool MERGED_B = (BLOCK_N == 128);
        constexpr int MMA_N = MERGED_B ? 2 * BLOCK_N : BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            if constexpr (MERGED_B) {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8, sfb1_desc);
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8 + 4, sfb2_desc);
            }
            else {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
              tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
            }
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              if constexpr (MERGED_B) {
                const uint64_t b_desc = make_desc_AB(B1_smem + k1 * MMA_N * 128 + k2 * 32);
                const int scale_B_tmem = SFB1_TMEM + k_sf * 8;
                tcgen05_mma_nvfp4_single(
                  ACC1_TMEM,
                  a_desc, b_desc, i_desc,
                  scale_A_tmem, scale_B_tmem,
                  enable_input_d
                );
              }
              else {
                const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
                const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
                const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;
                tcgen05_mma_nvfp4_dualA(
                  ACC1_TMEM, ACC2_TMEM,
                  a_desc, b1_desc, b2_desc, i_desc,
                  scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                  enable_input_d
                );
              }
            }
          }

          asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
            :: "r"(mma_mbar_addr + stage_id * 8)
            : "memory"
          );
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    static int run_one_cfg(
      int k,
      const void* a_ptr,
      const void* b1_ptr,
      const void* b2_ptr,
      const void* sfa_ptr,
      const void* sfb1_ptr,
      const void* sfb2_ptr,
      void* out_ptr,
      int m, int n
    ) {
      // cfg3: m=256 n=3072 k=4096
      constexpr int EXPECT_M = 256;
      constexpr int EXPECT_N = 3072;
      constexpr int EXPECT_K = 4096;
      if (m != EXPECT_M || n != EXPECT_N || k != EXPECT_K) return -13;

      constexpr int BLOCK_M = 128;
      constexpr int BLOCK_N = 64;
      constexpr int BLOCK_K = 256;
      constexpr int NUM_STAGES = 5;
      constexpr int CLUSTER_M = 2;

      CUtensorMap A_tmap, B1_tmap, B2_tmap;
      init_AB_tmap(&A_tmap,  reinterpret_cast<const char*>(a_ptr),  m, k, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B1_tmap, reinterpret_cast<const char*>(b1_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);
      init_AB_tmap(&B2_tmap, reinterpret_cast<const char*>(b2_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);

      const int grid = (m / BLOCK_M) * (n / BLOCK_N);
      const int tb_size = BLOCK_M + 2 * WARP_SIZE;

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
      const int smem_size = STAGE_SIZE * NUM_STAGES;

      auto this_kernel = dual_kernel<EXPECT_K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, CLUSTER_M>;
      if (smem_size > 48'000) cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      this_kernel<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap,
        reinterpret_cast<const char*>(sfa_ptr),
        reinterpret_cast<const char*>(sfb1_ptr),
        reinterpret_cast<const char*>(sfb2_ptr),
        reinterpret_cast<half*>(out_ptr),
        m, n
      );

      auto err = cudaGetLastError();
      g_last_cuda_error = int(err);
      return (err == cudaSuccess) ? 0 : -20;
    }

    int nvfp4_dual_gemm_fused_run(
      uint64_t a_ptr,
      uint64_t b1_ptr,
      uint64_t b2_ptr,
      uint64_t sfa_ptr,
      uint64_t sfb1_ptr,
      uint64_t sfb2_ptr,
      uint64_t out_ptr,
      int m, int n, int k, int l
    ) {
      (void)l;  // only l=1 fast path in Python
      g_last_cuda_error = int(cudaSuccess);

      return run_one_cfg(
        k,
        reinterpret_cast<const void*>(a_ptr),
        reinterpret_cast<const void*>(b1_ptr),
        reinterpret_cast<const void*>(b2_ptr),
        reinterpret_cast<const void*>(sfa_ptr),
        reinterpret_cast<const void*>(sfb1_ptr),
        reinterpret_cast<const void*>(sfb2_ptr),
        reinterpret_cast<void*>(out_ptr),
        m, n
      );
    }

    std::tuple<int, int> nvfp4_dual_gemm_last_error() {
      return {0, g_last_cuda_error};
    }
    """

# cfg4: m=512 n=3072 k=7168
cpp_src_cfg4 = r"""
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    // cfg: cfg4 (m512_n3072_k7168)
    int nvfp4_dual_gemm_fused_run(
        uint64_t a_ptr,
        uint64_t b1_ptr,
        uint64_t b2_ptr,
        uint64_t sfa_ptr,
        uint64_t sfb1_ptr,
        uint64_t sfb2_ptr,
        uint64_t out_ptr,
        int m, int n, int k, int l);

    std::tuple<int, int> nvfp4_dual_gemm_last_error();
    """

cuda_src_cfg4 = r"""
    #include <cuda.h>
    #include <cudaTypedefs.h>
    #include <cuda_fp16.h>
    #include <cuda_runtime.h>
    #include <cstdint>
    #include <tuple>
    #include <torch/extension.h>

    constexpr int WARP_SIZE = 32;
    constexpr int MMA_K = 64;  // 32 bytes of FP4 (packed)

    // Cache policy hints (same as r001 solution).
    constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
    constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000;
    constexpr uint64_t EVICT_LAST   = 0x14F0000000000000;

    namespace {
    static int g_last_cuda_error = int(cudaSuccess);
    }  // namespace

    __device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
    __device__ uint32_t elect_sync() {
      uint32_t pred = 0;
      asm volatile(
        "{\n\t"
        ".reg .pred %%px;\n\t"
        "elect.sync _|%%px, %1;\n\t"
        "@%%px mov.s32 %0, 1;\n\t"
        "}"
        : "+r"(pred)
        : "r"(0xFFFFFFFF)
      );
      return pred;
    }

    __device__ inline void mbarrier_init(int mbar_addr, int count) {
      asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
    }

    // https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
    __device__ void mbarrier_wait(int mbar_addr, int phase) {
      uint32_t ticks = 0x989680;  // optional
      asm volatile(
        "{\n\t"
        ".reg .pred P1;\n\t"
        "LAB_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
        "@P1 bra.uni DONE;\n\t"
        "bra.uni LAB_WAIT;\n\t"
        "DONE:\n\t"
        "}"
        :: "r"(mbar_addr), "r"(phase), "r"(ticks)
      );
    }

    __device__ inline void tma_gmem2smem(int dst, const void* src, int size, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
      );
    }

    __device__ inline void tma_gmem2smem_multicast(
      int dst,
      const void* src,
      int size,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.L2::cache_hint "
        "[%0], [%1], %2, [%3], %4, %5;"
        :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem(int dst, const void* tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tma_3d_gmem2smem_multicast(
      int dst,
      const void* tmap_ptr,
      int x,
      int y,
      int z,
      int mbar_addr,
      uint16_t cta_mask,
      uint64_t cache_policy
    ) {
      asm volatile(
        "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %6, %7;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy)
        : "memory"
      );
    }

    __device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
      asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
    }

    __device__ inline void tcgen05_mma_nvfp4_single(
      int d_tmem,
      uint64_t a_desc,
      uint64_t b_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B_tmem,
      int enable_input_d
    ) {
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
        "[%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
      );
    }

    __device__ inline void tcgen05_mma_nvfp4_dualA(
      int d1_tmem,
      int d2_tmem,
      uint64_t a_desc,
      uint64_t b1_desc,
      uint64_t b2_desc,
      uint32_t i_desc,
      int scale_A_tmem,
      int scale_B1_tmem,
      int scale_B2_tmem,
      int enable_input_d
    ) {
      // Reuse A across the two MMAs via the TensorCore collector buffer.
      asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %9, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
        "[%0], %2, %3, %5, [%6], [%7], p;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
        "[%1], %2, %4, %5, [%6], [%8], p;\n\t"
        "}"
        :: "r"(d1_tmem), "r"(d2_tmem),
           "l"(a_desc), "l"(b1_desc), "l"(b2_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B1_tmem), "r"(scale_B2_tmem), "r"(enable_input_d)
      );
    }

    struct SHAPE {
      static constexpr char _16x256b[] = ".16x256b";
    };
    struct NUM {
      static constexpr char x4[] = ".x4";
      static constexpr char x8[] = ".x8";
    };

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_16regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%17%18.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15}, [%16];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    template <const char* SHAPE, const char* NUM>
    __device__ inline void tcgen05_ld_32regs(float* tmp, int row, int col) {
      asm volatile(
        "tcgen05.ld.sync.aligned%33%34.b32 "
        "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
        "  %8,  %9, %10, %11, %12, %13, %14, %15, "
        " %16, %17, %18, %19, %20, %21, %22, %23, "
        " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
        : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
          "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
          "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
          "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM)
      );
    }

    __device__ inline void tcgen05_ld_16x256bx4(float* tmp, int row, int col) {
      tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col);
    }

    __device__ inline void tcgen05_ld_16x256bx8(float* tmp, int row, int col) {
      tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
    }

    void check_cu(CUresult err) {
      if (err == CUDA_SUCCESS) return;
      const char* error_msg_ptr = nullptr;
      if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS) {
        error_msg_ptr = "unable to get error string";
      }
      TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
    }

    void init_AB_tmap(
      CUtensorMap* tmap,
      const char* ptr,
      uint64_t global_height, uint64_t global_width,
      uint32_t shared_height, uint32_t shared_width,
      CUtensorMapL2promotion l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE
    ) {
      constexpr uint32_t rank = 3;
      uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
      uint64_t globalStrides[rank-1] = {global_width / 2, 128};  // bytes
      uint32_t boxDim[rank]          = {256, shared_height, shared_width / 256};
      uint32_t elementStrides[rank]  = {1, 1, 1};

      auto err = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        (void*)ptr,
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
        l2_promotion,
        CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
      );
      check_cu(err);
    }

    __device__ __forceinline__ float silu(float x) {
      // silu(x) = x / (1 + exp(-x))
      return x / (1.0f + __expf(-x));
    }

    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_M>
    __global__
    __cluster_dims__(CLUSTER_M, 1, 1)
    __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      static_assert(CLUSTER_M >= 1 && CLUSTER_M <= 16);
      constexpr uint16_t cta_mask = uint16_t((1u << CLUSTER_M) - 1u);
      int cta_rank = 0;
      asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      // BlockIdx linearization: run along M first so a cluster spans the full M-slab for a fixed N-tile.
      const int bid_n = bid / grid_m;
      const int bid_m = bid % grid_m;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      // Dynamic shared memory. Stage layout is per r001, extended for B1/B2 and SFB1/SFB2.
      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      // mbarriers: NUM_STAGES for TMA, NUM_STAGES for MMA, 1 for mainloop.
      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      // TMEM layout:
      //   ACC1: [0 .. BLOCK_N-1]
      //   ACC2: [BLOCK_N .. 2*BLOCK_N-1]
      //   SFA:  [2*BLOCK_N .. 2*BLOCK_N+SF_COLS-1]
      //   SFB1: next SF_COLS
      //   SFB2: next SF_COLS
      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
          // Cluster-wide stage reuse sync (B/B2 are multicast to the whole cluster).
          mbarrier_init(mma_mbar_addr + i * 8, CLUSTER_M);
        }
        mbarrier_init(mainloop_mbar_addr, 1);
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        // Allocate TMEM (address is assumed 0).
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      if constexpr (CLUSTER_M > 1) {
        asm volatile("barrier.cluster.arrive.release.aligned;" ::: "memory");
        asm volatile("barrier.cluster.wait.acquire.aligned;" ::: "memory");
      }
      else {
        __syncthreads();
      }

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        // TMA warp.
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_3d_gmem2smem_multicast(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
              tma_3d_gmem2smem_multicast(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          }

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          if constexpr (CLUSTER_M > 1) {
            if (cta_rank == 0) {
              tma_gmem2smem_multicast(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cta_mask, cache_B);
              tma_gmem2smem_multicast(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cta_mask, cache_B);
            }
          }
          else {
            tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
            tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);
          }

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        // Prologue: fill pipeline.
        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }

        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        // MMA warp.
        constexpr int MMA_N = BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
            tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);
              const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
              const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;  // 4 = 256/MMA_K
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
              const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
              const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4_dualA(
                ACC1_TMEM, ACC2_TMEM,
                a_desc, b1_desc, b2_desc, i_desc,
                scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                enable_input_d
              );
            }
          }

          if constexpr (CLUSTER_M > 1) {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
              :: "r"(mma_mbar_addr + stage_id * 8), "h"(cta_mask)
              : "memory"
            );
          }
          else {
            asm volatile(
              "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
              :: "r"(mma_mbar_addr + stage_id * 8)
              : "memory"
            );
          }
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        // Epilogue threads: fuse silu(acc1) * acc2 and store fp16.
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        // N-major output with half2 stores. Process in 64-column chunks to cap registers.
        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    // Non-cluster kernel for m==512 path: reuse v_t's merged-B MMA (MMA_N=256 when BLOCK_N==128).
    template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
    __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
    void dual_kernel_merged(
      const __grid_constant__ CUtensorMap A_tmap,
      const __grid_constant__ CUtensorMap B1_tmap,
      const __grid_constant__ CUtensorMap B2_tmap,
      const char* SFA_ptr,
      const char* SFB1_ptr,
      const char* SFB2_ptr,
      half* C_ptr,
      int M, int N
    ) {
      const int tid = threadIdx.x;
      const int bid = blockIdx.x;
      const int lane_id = tid % WARP_SIZE;
      const int warp_id = tid / WARP_SIZE;

      const int grid_m = M / BLOCK_M;
      const int grid_n = N / BLOCK_N;
      const int bid_m = bid / grid_n;
      const int bid_n = bid % grid_n;
      const int off_m = bid_m * BLOCK_M;
      const int off_n = bid_n * BLOCK_N;

      constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

      extern __shared__ __align__(1024) char smem_ptr[];
      const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

      #pragma nv_diag_suppress static_var_with_dynamic_init
      __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
      const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
      const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
      const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

      constexpr int SF_COLS = 4 * (BLOCK_K / MMA_K);
      constexpr int ACC1_TMEM = 0;
      constexpr int ACC2_TMEM = BLOCK_N;
      constexpr int SFA_TMEM  = 2 * BLOCK_N;
      constexpr int SFB1_TMEM = SFA_TMEM + SF_COLS;
      constexpr int SFB2_TMEM = SFB1_TMEM + SF_COLS;
      constexpr int TMEM_COLS = 512;

      if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES * 2 + 1; i++) {
          mbarrier_init(tma_mbar_addr + i * 8, 1);
        }
        asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
      }
      else if (warp_id == 1) {
        asm volatile(
          "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
          :: "r"(smem), "r"(TMEM_COLS)
        );
      }
      __syncthreads();

      constexpr int num_iters = K / BLOCK_K;

      if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        const uint64_t cache_A = EVICT_FIRST;
        const uint64_t cache_B = EVICT_FIRST;

        auto issue_tma = [&](int iter_k, int stage_id) {
          const int mbar_addr = tma_mbar_addr + stage_id * 8;

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const int off_k = iter_k * BLOCK_K;
          tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / 256, mbar_addr, cache_A);
          tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
          tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

          const int rest_k = K / 16 / 4;
          const char* SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
          const char* SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;

          tma_gmem2smem(SFA_smem,  SFA_src,  SFA_size, mbar_addr, cache_A);
          tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
          tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);

          asm volatile(
            "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
            :: "r"(mbar_addr), "r"(STAGE_SIZE)
            : "memory"
          );
        };

        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++) {
          issue_tma(iter_k, iter_k);
        }
        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
          issue_tma(iter_k, stage_id);
        }
      }
      else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        constexpr bool MERGED_B = (BLOCK_N == 128);
        constexpr int MMA_N = MERGED_B ? 2 * BLOCK_N : BLOCK_N;
        constexpr int MMA_M = 128;
        constexpr uint32_t i_desc =
            (1U << 7U)   // atype=E2M1
          | (1U << 10U)  // btype=E2M1
          | ((uint32_t)MMA_N >> 3U << 17U)
          | ((uint32_t)MMA_M >> 7U << 27U);

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          const int stage_id = iter_k % NUM_STAGES;
          const int tma_phase = (iter_k / NUM_STAGES) % 2;
          mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

          const int stage_base = smem + stage_id * STAGE_SIZE;
          const int A_smem   = stage_base;
          const int B1_smem  = A_smem + A_size;
          const int B2_smem  = B1_smem + B_size;
          const int SFA_smem = B2_smem + B_size;
          const int SFB1_smem = SFA_smem + SFA_size;
          const int SFB2_smem = SFB1_smem + SFB_size;

          const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem  >> 4ULL);
          const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
          const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            const uint64_t sfa_desc  = SFA_desc  + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
            const uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
            tcgen05_cp_nvfp4(SFA_TMEM  + k * 4, sfa_desc);
            if constexpr (MERGED_B) {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8, sfb1_desc);
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 8 + 4, sfb2_desc);
            }
            else {
              tcgen05_cp_nvfp4(SFB1_TMEM + k * 4, sfb1_desc);
              tcgen05_cp_nvfp4(SFB2_TMEM + k * 4, sfb2_desc);
            }
          }

          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
              const uint64_t a_desc  = make_desc_AB(A_smem  + k1 * BLOCK_M * 128 + k2 * 32);

              const int k_sf = k1 * 4 + k2;
              const int scale_A_tmem = SFA_TMEM + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              if constexpr (MERGED_B) {
                const uint64_t b_desc = make_desc_AB(B1_smem + k1 * MMA_N * 128 + k2 * 32);
                const int scale_B_tmem = SFB1_TMEM + k_sf * 8;
                tcgen05_mma_nvfp4_single(
                  ACC1_TMEM,
                  a_desc, b_desc, i_desc,
                  scale_A_tmem, scale_B_tmem,
                  enable_input_d
                );
              }
              else {
                const uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
                const int scale_B_tmem = (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
                const int scale_B1_tmem = SFB1_TMEM + k_sf * 4 + scale_B_tmem;
                const int scale_B2_tmem = SFB2_TMEM + k_sf * 4 + scale_B_tmem;
                tcgen05_mma_nvfp4_dualA(
                  ACC1_TMEM, ACC2_TMEM,
                  a_desc, b1_desc, b2_desc, i_desc,
                  scale_A_tmem, scale_B1_tmem, scale_B2_tmem,
                  enable_input_d
                );
              }
            }
          }

          asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
            :: "r"(mma_mbar_addr + stage_id * 8)
            : "memory"
          );
        }

        asm volatile(
          "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
          :: "r"(mainloop_mbar_addr)
          : "memory"
        );
      }
      else if (tid < BLOCK_M) {
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

        constexpr int CHUNK_N = 32;
        constexpr int ITERS_N = BLOCK_N / CHUNK_N;
        #pragma unroll
        for (int n_chunk = 0; n_chunk < ITERS_N; n_chunk++) {
          const int col_base = n_chunk * CHUNK_N;

          #pragma unroll
          for (int m16 = 0; m16 < 2; m16++) {
            float acc1[16];
            float acc2[16];
            tcgen05_ld_16x256bx4(acc1, warp_id * 32 + m16 * 16, ACC1_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
            tcgen05_ld_16x256bx4(acc2, warp_id * 32 + m16 * 16, ACC2_TMEM + col_base);
            asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");

            #pragma unroll
            for (int i = 0; i < CHUNK_N / 8; i++) {
              const int row = off_m + warp_id * 32 + m16 * 16 + lane_id / 4;
              const int col = off_n + col_base + i * 8 + (lane_id % 4) * 2;

              const float o0 = silu(acc1[i * 4 + 0]) * acc2[i * 4 + 0];
              const float o1 = silu(acc1[i * 4 + 1]) * acc2[i * 4 + 1];
              const float o2 = silu(acc1[i * 4 + 2]) * acc2[i * 4 + 2];
              const float o3 = silu(acc1[i * 4 + 3]) * acc2[i * 4 + 3];

              reinterpret_cast<half2*>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({o0, o1});
              reinterpret_cast<half2*>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({o2, o3});
            }
          }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0) {
          asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
        }
      }
    }

    static int run_one_cfg(
      int k,
      const void* a_ptr,
      const void* b1_ptr,
      const void* b2_ptr,
      const void* sfa_ptr,
      const void* sfb1_ptr,
      const void* sfb2_ptr,
      void* out_ptr,
      int m, int n
    ) {
      // cfg4: m=512 n=3072 k=7168
      constexpr int EXPECT_M = 512;
      constexpr int EXPECT_N = 3072;
      constexpr int EXPECT_K = 7168;
      if (m != EXPECT_M || n != EXPECT_N || k != EXPECT_K) return -13;

      constexpr int BLOCK_M = 128;
      constexpr int BLOCK_N = 128;
      constexpr int BLOCK_K = 256;
      constexpr int NUM_STAGES = 4;

      CUtensorMap A_tmap, B1_tmap, B2_tmap;
      init_AB_tmap(&A_tmap,  reinterpret_cast<const char*>(a_ptr),  m, k, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B1_tmap, reinterpret_cast<const char*>(b1_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);
      init_AB_tmap(&B2_tmap, reinterpret_cast<const char*>(b2_ptr), n, k, BLOCK_N, BLOCK_K, CU_TENSOR_MAP_L2_PROMOTION_L2_64B);

      const int grid = (m / BLOCK_M) * (n / BLOCK_N);
      const int tb_size = BLOCK_M + 2 * WARP_SIZE;

      constexpr int A_size   = BLOCK_M * BLOCK_K / 2;
      constexpr int B_size   = BLOCK_N * BLOCK_K / 2;
      constexpr int SFA_size = 128 * (BLOCK_K / 16);
      constexpr int SFB_size = 128 * (BLOCK_K / 16);
      constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
      const int smem_size = STAGE_SIZE * NUM_STAGES;

      auto this_kernel = dual_kernel_merged<EXPECT_K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
      if (smem_size > 48'000) cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      this_kernel<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap,
        reinterpret_cast<const char*>(sfa_ptr),
        reinterpret_cast<const char*>(sfb1_ptr),
        reinterpret_cast<const char*>(sfb2_ptr),
        reinterpret_cast<half*>(out_ptr),
        m, n
      );

      auto err = cudaGetLastError();
      g_last_cuda_error = int(err);
      return (err == cudaSuccess) ? 0 : -20;
    }

    int nvfp4_dual_gemm_fused_run(
      uint64_t a_ptr,
      uint64_t b1_ptr,
      uint64_t b2_ptr,
      uint64_t sfa_ptr,
      uint64_t sfb1_ptr,
      uint64_t sfb2_ptr,
      uint64_t out_ptr,
      int m, int n, int k, int l
    ) {
      (void)l;  // only l=1 fast path in Python
      g_last_cuda_error = int(cudaSuccess);

      return run_one_cfg(
        k,
        reinterpret_cast<const void*>(a_ptr),
        reinterpret_cast<const void*>(b1_ptr),
        reinterpret_cast<const void*>(b2_ptr),
        reinterpret_cast<const void*>(sfa_ptr),
        reinterpret_cast<const void*>(sfb1_ptr),
        reinterpret_cast<const void*>(sfb2_ptr),
        reinterpret_cast<void*>(out_ptr),
        m, n
      );
    }

    std::tuple<int, int> nvfp4_dual_gemm_last_error() {
      return {0, g_last_cuda_error};
    }
    """

_cfg_srcs = [
    ("cfg1", cpp_src_cfg1, cuda_src_cfg1),
    ("cfg2", cpp_src_cfg2, cuda_src_cfg2),
    ("cfg3", cpp_src_cfg3, cuda_src_cfg3),
    ("cfg4", cpp_src_cfg4, cuda_src_cfg4),
]

_mod_by_digest = {}
_mod_by_cfg = {}

_extra_cflags = ["-O3"]
_extra_cuda_cflags = [
    "-O3",
    "--use_fast_math",
    "--expt-relaxed-constexpr",
    "--expt-extended-lambda",
    "-std=c++17",
    "-w",
    "-gencode=arch=compute_100a,code=sm_100a",
    "--ptxas-options=--gpu-name=sm_100a",
]
_verbose = bool(int(os.getenv("NVFP4_EXT_VERBOSE", "0")))

for cfg_name, cpp_src, cuda_src in _cfg_srcs:
    digest = hashlib.md5((cpp_src + cuda_src).encode("utf-8")).hexdigest()[:10]
    mod = _mod_by_digest.get(digest)
    if mod is None:
        name = f"nvfp4_dual_gemm_fused_{digest}"
        mod = load_inline(
            name=name,
            cpp_sources=cpp_src,
            cuda_sources=cuda_src,
            functions=[
                "nvfp4_dual_gemm_fused_run",
                "nvfp4_dual_gemm_last_error",
            ],
            extra_cflags=_extra_cflags,
            extra_cuda_cflags=_extra_cuda_cflags,
            extra_ldflags=["-lcuda"],
            verbose=_verbose,
        )
        _mod_by_digest[digest] = mod
    _mod_by_cfg[cfg_name] = mod

mod_cfg1 = _mod_by_cfg["cfg1"]
mod_cfg2 = _mod_by_cfg["cfg2"]
mod_cfg3 = _mod_by_cfg["cfg3"]
mod_cfg4 = _mod_by_cfg["cfg4"]


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data

    m, k_half, l = a.shape
    n, _, _ = b1.shape
    k = k_half * 2

    # Fast path: only the 4 leaderboard benchmark shapes (l=1).
    if l != 1:
        from reference import ref_kernel

        return ref_kernel(data)

    if m == 256 and n == 4096 and k == 7168:
        mod = mod_cfg1
    elif m == 512 and n == 4096 and k == 7168:
        mod = mod_cfg2
    elif m == 256 and n == 3072 and k == 4096:
        mod = mod_cfg3
    elif m == 512 and n == 3072 and k == 7168:
        mod = mod_cfg4
    else:
        from reference import ref_kernel

        return ref_kernel(data)

    rc = int(
        mod.nvfp4_dual_gemm_fused_run(
            a.data_ptr(),
            b1.data_ptr(),
            b2.data_ptr(),
            sfa_permuted.data_ptr(),
            sfb1_permuted.data_ptr(),
            sfb2_permuted.data_ptr(),
            c.data_ptr(),
            m,
            n,
            k,
            l,
        )
    )
    if rc != 0:
        _, cuda_error = mod.nvfp4_dual_gemm_last_error()
        raise RuntimeError(
            f"nvfp4_dual_gemm_fused_run failed: rc={rc} cuda_error={int(cuda_error)}"
        )

    return c

# ---- modal_harness (leaderboard) ----
# cmd: python -m modal_harness.cli nvfp4 submit --leaderboard nvfp4_dual_gemm --mode leaderboard --output "nvfp4\dual_gemm\submissions\0101\v2\v_ai_8.log" "nvfp4\dual_gemm\submissions\0101\v2\v_ai_8.py"
# check: pass
# benchmark.geomean: 15186.93 ns (~15.187 us)
# [0] m=256 n=4096 k=7168 mean: 15044.80 ns (~15.045 us)
# [1] m=512 n=4096 k=7168 mean: 18511.36 ns (~18.511 us)
# [2] m=256 n=3072 k=4096 mean: 10481.39 ns (~10.481 us)
# [3] m=512 n=3072 k=7168 mean: 18223.68 ns (~18.224 us)
scrolls · 3741 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 318682.

- #!POPCORN leaderboard nvfp4_dual_gemm
- #!POPCORN gpu NVIDIA
+ #!POPCORN leaderboard modal_nvfp4_dual_gemm
+ #!POPCORN gpu B200
"""
Provenance: copied from `nvfp4/dual_gemm/submissions/0101/v2/v_ai_7.py`.

Best evidence level for this revision: reported

JSON