Skip to content
KernelIndex
Search⌘K

submission 215834

Quantizr · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

yuh.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-215834?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.1µs
#100 of 420
2025-12-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:361dec90ebe7f5b11cdf0fc6c78d99ea38a67449cd598568f6121ac2a679f877
license declaredunknown
license concludedunknown
authorsQuantizr
imported2026-08-26

Techniques

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

mbarrier__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
split-kconstexpr int SPLIT_K_C = 1;
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tile-k = 256static_assert(BK == 256, "This hardcoded build uses BK=256 only");
tma"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
vector-width = half2half2 h0 = __floats2half2_rn((x0 * s0) * y0, (x1 * s1) * y1);

Kernel source

yuh.py556 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

CUDA_SRC = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;

constexpr int SPLIT_K_C = 1;

constexpr uint64_t EVICT_NORMAL = 0x1000000000000000ULL;
constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

constexpr int SWZ_DESC_128B = 2;

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

__device__ __forceinline__ uint32_t elect_sync() {
  uint32_t pred = 0;
  asm volatile("{\n\t.reg .pred %px;\n\telect.sync _|%px, %1;\n\t@%px mov.s32 %0, 1;\n\t}"
               : "+r"(pred) : "r"(0xFFFFFFFF));
  return pred;
}

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

__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) {
  uint32_t ticks = 0x989680;
  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__ __forceinline__ 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__ __forceinline__ 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__ __forceinline__ 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__ __forceinline__ void tcgen05_mma_nvfp4(
    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));
}

// ---- tcgen05 ld for N-major epilogue ----
struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };
struct NUM   { static constexpr char x8[]  = ".x8"; static constexpr char x16[] = ".x16"; };

template <const char *SHAPE_S, const char *NUM_S>
__device__ __forceinline__ 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_S), "C"(NUM_S)
  );
}

template <const char *SHAPE_S, const char *NUM_S>
__device__ __forceinline__ void tcgen05_ld_64regs(float *tmp, int row, int col) {
  asm volatile(
      "tcgen05.ld.sync.aligned%65%66.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,%33,%34,%35,%36,%37,%38,%39,%40,%41,%42,%43,%44,%45,%46,%47,"
      "%48,%49,%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,%60,%61,%62,%63 }, [%64];"
      : "=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]),
        "=f"(tmp[32]),"=f"(tmp[33]),"=f"(tmp[34]),"=f"(tmp[35]),
        "=f"(tmp[36]),"=f"(tmp[37]),"=f"(tmp[38]),"=f"(tmp[39]),
        "=f"(tmp[40]),"=f"(tmp[41]),"=f"(tmp[42]),"=f"(tmp[43]),
        "=f"(tmp[44]),"=f"(tmp[45]),"=f"(tmp[46]),"=f"(tmp[47]),
        "=f"(tmp[48]),"=f"(tmp[49]),"=f"(tmp[50]),"=f"(tmp[51]),
        "=f"(tmp[52]),"=f"(tmp[53]),"=f"(tmp[54]),"=f"(tmp[55]),
        "=f"(tmp[56]),"=f"(tmp[57]),"=f"(tmp[58]),"=f"(tmp[59]),
        "=f"(tmp[60]),"=f"(tmp[61]),"=f"(tmp[62]),"=f"(tmp[63])
      : "r"((row << 16) | col), "C"(SHAPE_S), "C"(NUM_S)
  );
}

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

// ---- TMA map encode ----
static inline void check_cu(CUresult err) {
  if (err == CUDA_SUCCESS) return;
  const char* msg = "unknown";
  cuGetErrorString(err, &msg);
  TORCH_CHECK(false, "CUDA Driver error: ", msg);
}

template<int BK>
void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t g_h, uint64_t g_w, uint32_t s_h, uint32_t s_w) {
  static_assert(BK == 256, "This hardcoded build uses BK=256 only");
  constexpr int INNER = 256;
  constexpr int INNER_BYTES = INNER / 2;

  uint64_t gDim[3] = {(uint64_t)INNER, g_h, g_w / (uint64_t)INNER};
  uint64_t gStrides[2] = {g_w / 2, (uint64_t)INNER_BYTES};
  uint32_t bDim[3] = {(uint32_t)INNER, s_h, (uint32_t)(s_w / (uint32_t)INNER)};
  uint32_t eStrides[3] = {1, 1, 1};

  check_cu(cuTensorMapEncodeTiled(
      tmap,
      CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
      3,
      (void *)ptr,
      gDim,
      gStrides,
      bDim,
      eStrides,
      CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
      CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
      CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
      CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}

template<int BM>
struct CfgT {
  static constexpr int EPIL_THREADS = BM;
  static constexpr int TB_SIZE = EPIL_THREADS + 2 * WARP_SIZE;
  static constexpr int TMA_WARP = EPIL_THREADS / WARP_SIZE;
  static constexpr int MMA_WARP = TMA_WARP + 1;
};

template<int K_FIXED, int BM, int BN, int BK, int NUM_STAGES>
__global__ __launch_bounds__(CfgT<BM>::TB_SIZE)
void 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
) {
  static_assert(BM == 128);
  static_assert(BN == 64 || BN == 128);
  static_assert(BK == 256);
  static_assert(K_FIXED % BK == 0);

  const int tid = threadIdx.x;
  const int warp_id = tid / WARP_SIZE;
  const int lane_id = tid & 31;

  constexpr int A_size   = BM * BK / 2;
  constexpr int B_size   = BN * BK / 2;
  constexpr int SFA_size = 128 * BK / 16;
  constexpr int SFB_size = 128 * BK / 16;
  constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

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

  __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;

  __shared__ int tmem_base;

  constexpr int ACC1_tmem = 0;
  constexpr int ACC2_tmem = BN;
  constexpr int SFA_tmem  = BN * 2;
  constexpr int SFB1_tmem = SFA_tmem  + 4 * (BK / MMA_K);
  constexpr int SFB2_tmem = SFB1_tmem + 4 * (BK / MMA_K);

  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;");
  }
  if (warp_id == 1) {
    const int tmem_base_smem = static_cast<int>(__cvta_generic_to_shared(&tmem_base));
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                 :: "r"(tmem_base_smem), "r"(BN * 4));
  }
  __syncthreads();

  constexpr int num_iters = K_FIXED / BK;
  constexpr int INNER = 256;

  // Scope the whole tile body to shorten live ranges (key “improvement”)
  auto run_one_tile = [&](int bid) {
    const int grid_n = N / BN;
    const int bid_m = bid / grid_n;
    const int bid_n = bid % grid_n;
    const int off_m = bid_m * BM;
    const int off_n = bid_n * BN;

    // Producer warp
    if (warp_id == CfgT<BM>::TMA_WARP && elect_sync()) {
      uint64_t cache_A = (M > N) ? EVICT_FIRST : EVICT_LAST;
      uint64_t cache_B = (M > N) ? EVICT_LAST  : EVICT_FIRST;

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

        const int A_smem    = smem + stage_id * STAGE_SIZE;
        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 * BK;

        tma_3d_gmem2smem(A_smem,  &A_tmap,  0, off_m, off_k / INNER, mbar_addr, cache_A);
        tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / INNER, mbar_addr, cache_B);
        tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / INNER, mbar_addr, cache_B);

        const int rest_k = K_FIXED / 64;
        const char *SFA_src  = SFA_ptr  + ((off_m / 128) * rest_k + off_k / 64) * 512;
        const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
        const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 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");
      };

      #pragma unroll
      for (int i = 0; i < NUM_STAGES; i++) issue_tma(i, i);

      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);
      }
    }

    // Consumer warp
    else if (warp_id == CfgT<BM>::MMA_WARP && elect_sync()) {
      constexpr int MMA_N = BN;
      constexpr int MMA_M = 128;

      constexpr uint32_t i_desc =
          (1U << 7U) | (1U << 10U) |
          ((uint32_t)MMA_N >> 3U << 17U) |
          ((uint32_t)MMA_M >> 7U << 27U);

      const int taddr = tmem_base;

      constexpr int SWZ_MODE  = SWZ_DESC_128B;
      constexpr int SWZ_BYTES = 128;

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

      constexpr uint64_t AB_step = (32ULL >> 4ULL);
      constexpr uint64_t SF_step = (512ULL >> 4ULL);

      const int scale_A_blk = (bid_m % (128 / BM)) * (BM / 32);
      const int scale_B_blk = (bid_n % (128 / BN)) * (BN / 32);

      constexpr int SEG_K = 256;
      constexpr int SEG_MMA_STEPS = SEG_K / MMA_K; // 4
      constexpr int NUM_SEGS = BK / SEG_K;         // 1

      constexpr int A_SEG_BYTES  = BM * SEG_K / 2;
      constexpr int B_SEG_BYTES  = BN * SEG_K / 2;
      constexpr int SF_SEG_BYTES = 128 * SEG_K / 16;

      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 A_smem    = smem + stage_id * STAGE_SIZE;
        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;

        #pragma unroll
        for (int seg = 0; seg < NUM_SEGS; seg++) {
          const int A_seg_smem    = A_smem    + seg * A_SEG_BYTES;
          const int B1_seg_smem   = B1_smem   + seg * B_SEG_BYTES;
          const int B2_seg_smem   = B2_smem   + seg * B_SEG_BYTES;
          const int SFA_seg_smem  = SFA_smem  + seg * SF_SEG_BYTES;
          const int SFB1_seg_smem = SFB1_smem + seg * SF_SEG_BYTES;
          const int SFB2_seg_smem = SFB2_smem + seg * SF_SEG_BYTES;

          const uint64_t SF_desc0 = make_desc_SF(0);
          const uint64_t SFA_desc_base  = SF_desc0 + ((uint64_t)SFA_seg_smem  >> 4ULL);
          const uint64_t SFB1_desc_base = SF_desc0 + ((uint64_t)SFB1_seg_smem >> 4ULL);
          const uint64_t SFB2_desc_base = SF_desc0 + ((uint64_t)SFB2_seg_smem >> 4ULL);

          const uint64_t A_desc0  = make_desc_AB(A_seg_smem);
          const uint64_t B1_desc0 = make_desc_AB(B1_seg_smem);
          const uint64_t B2_desc0 = make_desc_AB(B2_seg_smem);

          #pragma unroll
          for (int kk = 0; kk < SEG_MMA_STEPS; kk++) {
            const int k_global = seg * SEG_MMA_STEPS + kk;
            tcgen05_cp_nvfp4(taddr + SFA_tmem  + k_global * 4, SFA_desc_base  + (uint64_t)kk * SF_step);
            tcgen05_cp_nvfp4(taddr + SFB1_tmem + k_global * 4, SFB1_desc_base + (uint64_t)kk * SF_step);
            tcgen05_cp_nvfp4(taddr + SFB2_tmem + k_global * 4, SFB2_desc_base + (uint64_t)kk * SF_step);
          }

          #pragma unroll
          for (int kk = 0; kk < SEG_MMA_STEPS; kk++) {
            const int k_global = seg * SEG_MMA_STEPS + kk;
            const uint64_t a_desc  = A_desc0  + (uint64_t)kk * AB_step;
            const uint64_t b1_desc = B1_desc0 + (uint64_t)kk * AB_step;
            const uint64_t b2_desc = B2_desc0 + (uint64_t)kk * AB_step;

            const int scale_A_tmem  = (taddr + SFA_tmem)  + k_global * 4 + scale_A_blk;
            const int scale_B1_tmem = (taddr + SFB1_tmem) + k_global * 4 + scale_B_blk;
            const int scale_B2_tmem = (taddr + SFB2_tmem) + k_global * 4 + scale_B_blk;

            const int enable_input_d = (k_global == 0) ? iter_k : 1;
            tcgen05_mma_nvfp4(taddr + ACC1_tmem, a_desc, b1_desc, i_desc, scale_A_tmem, scale_B1_tmem, enable_input_d);
            tcgen05_mma_nvfp4(taddr + ACC2_tmem, a_desc, b2_desc, i_desc, scale_A_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");
    }

    // Epilogue
    if (tid < BM) {
      mbarrier_wait(mainloop_mbar_addr, 0);
      asm volatile("tcgen05.fence::after_thread_sync;");

      #pragma unroll
      for (int m16 = 0; m16 < 32; m16 += 16) {
        float a_tmp[BN / 2];
        float b_tmp[BN / 2];

        const int tmem_row = warp_id * 32 + m16;

        if constexpr (BN == 64) {
          tcgen05_ld_16x256bx8(a_tmp, tmem_row, ACC1_tmem);
          tcgen05_ld_16x256bx8(b_tmp, tmem_row, ACC2_tmem);
        } else {
          tcgen05_ld_16x256bx16(a_tmp, tmem_row, ACC1_tmem);
          tcgen05_ld_16x256bx16(b_tmp, tmem_row, ACC2_tmem);
        }
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        const int row0 = off_m + warp_id * 32 + m16 + (lane_id / 4);
        const int row1 = row0 + 8;
        const int lane2 = (lane_id % 4) * 2;

        #pragma unroll
        for (int i = 0; i < BN / 8; i++) {
          const int col = off_n + i * 8 + lane2;
          const int j = i * 4;

          float x0 = a_tmp[j + 0], x1 = a_tmp[j + 1];
          float x2 = a_tmp[j + 2], x3 = a_tmp[j + 3];
          float y0 = b_tmp[j + 0], y1 = b_tmp[j + 1];
          float y2 = b_tmp[j + 2], y3 = b_tmp[j + 3];

          float s0 = 1.0f / (1.0f + __expf(-x0));
          float s1 = 1.0f / (1.0f + __expf(-x1));
          float s2 = 1.0f / (1.0f + __expf(-x2));
          float s3 = 1.0f / (1.0f + __expf(-x3));

          half2 h0 = __floats2half2_rn((x0 * s0) * y0, (x1 * s1) * y1);
          half2 h1 = __floats2half2_rn((x2 * s2) * y2, (x3 * s3) * y3);

          *reinterpret_cast<half2*>(C_ptr + row0 * N + col) = h0;
          *reinterpret_cast<half2*>(C_ptr + row1 * N + col) = h1;
        }
      }

      asm volatile("bar.sync 1, %0;" :: "r"(BM) : "memory");
      if (tid == 0) {
        const int taddr = tmem_base;
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
                     :: "r"(taddr), "r"(BN * 4));
      }
    }
  };

  // Non-persistent path (still benefits from scoped tile body)
  run_one_tile((int)blockIdx.y);
}

template<int K_FIXED, int BM, int BN, int BK, int NUM_STAGES>
void launch_one(
  const at::Tensor& A,
  const at::Tensor& B1,
  const at::Tensor& B2,
  const at::Tensor& SFAp,
  const at::Tensor& SFB1p,
  const at::Tensor& SFB2p,
  at::Tensor& C
) {
  const int M = (int)A.size(0);
  const int N = (int)B1.size(0);

  CUtensorMap A_tmap, B1_tmap, B2_tmap;
  init_AB_tmap<BK>(&A_tmap,  (const char*)A.data_ptr(),  (uint64_t)M, (uint64_t)K_FIXED, (uint32_t)BM, (uint32_t)BK);
  init_AB_tmap<BK>(&B1_tmap, (const char*)B1.data_ptr(), (uint64_t)N, (uint64_t)K_FIXED, (uint32_t)BN, (uint32_t)BK);
  init_AB_tmap<BK>(&B2_tmap, (const char*)B2.data_ptr(), (uint64_t)N, (uint64_t)K_FIXED, (uint32_t)BN, (uint32_t)BK);

  const int tiles = (M / BM) * (N / BN);
  dim3 grid(1, tiles);

  int A_sz      = (BM * BK / 2);
  int B_sz      = (BN * BK / 2);
  int SFA_sz    = (128 * (BK / 16));
  int SFB_sz    = (128 * (BK / 16));
  int stage_sz  = A_sz + 2 * B_sz + SFA_sz + 2 * SFB_sz;
  int smem_size = stage_sz * NUM_STAGES;

  auto k = kernel<K_FIXED, BM, BN, BK, NUM_STAGES>;
  if (smem_size > 48000) cudaFuncSetAttribute(k, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

  constexpr int TB = CfgT<BM>::TB_SIZE;
  k<<<grid, TB, smem_size>>>(
      A_tmap, B1_tmap, B2_tmap,
      (const char*)SFAp.data_ptr(),
      (const char*)SFB1p.data_ptr(),
      (const char*)SFB2p.data_ptr(),
      (half*)C.data_ptr(),
      M, N);
}

// Hardcoded dispatch in C++ (cfg15/cfg19 equivalents)
at::Tensor run(
  const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2,
  const at::Tensor& SFAp, const at::Tensor& SFB1p, const at::Tensor& SFB2p,
  at::Tensor& C
) {
  const int M = (int)C.size(0);
  const int N = (int)C.size(1);
  const int K = (int)(A.size(1) * 2);

  // Best from your sweep:
  // cfg15: BM128 BN64  BK256 NS5
  // cfg19: BM128 BN128 BK256 NS4
  if (K == 7168 && M == 256 && N == 4096) { launch_one<7168, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 4096 && M == 256 && N == 3072) { launch_one<4096, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 7168 && M == 512 && N == 4096) { launch_one<7168, 128, 128, 256, 4>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 7168 && M == 512 && N == 3072) { launch_one<7168, 128, 128, 256, 4>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }

  if (K == 7168) { launch_one<7168, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 2304) { launch_one<2304, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 2048) { launch_one<2048, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 1536) { launch_one<1536, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 512) { launch_one<512, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  if (K == 256) { launch_one<256, 128,  64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
  
  TORCH_CHECK(false, "Unsupported shape M=", M, " N=", N, " K=", K);
}

TORCH_LIBRARY(dual_fixed_scoped, m) {
  m.def("run(Tensor A, Tensor B1, Tensor B2, Tensor SFAp, Tensor SFB1p, Tensor SFB2p, Tensor(a!) C) -> Tensor");
  m.impl("run", &run);
}
"""

load_inline(
    name="dual_fixed_scoped",
    cpp_sources="",
    cuda_sources=CUDA_SRC,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
        "-std=c++17",
    ],
    extra_ldflags=["-lcuda"],
)

_run = torch.ops.dual_fixed_scoped.run

def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_perm, sfb1_perm, sfb2_perm, c = data
    _run(a, b1, b2, sfa_perm, sfb1_perm, sfb2_perm, c)
    return c
scrolls · 556 lines total

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

Best evidence level for this revision: reported

JSON