Skip to content
KernelIndex
Search⌘K

submission 491633

XoTic · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v1b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-491633?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 group GEMMsuite of 4 cases
NVIDIA B200
202.8µs
#104 of 145
2026-02-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5800e94fb0005199ff26a801fff3ac124354373f17d56a051dd76fa1079b03d3
license declaredunknown
license concludedunknown
authorsXoTic
imported2026-08-15

Techniques

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

mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
num-warps = 4constexpr int TMA_NUM_WARPS = 4;
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
stages = 2constexpr int NUM_STAGES = 2; // Double buffering
tcgen05"tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;\\n"
tile-k = 256constexpr int TMA_BLOCK_K = 256;
tile-m = 128constexpr int TMA_BLOCK_M = 128;
tile-n = 128constexpr int TMA_BLOCK_N = 128;
tma"cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "
vector-width = half2reinterpret_cast<half2*>(C_ptr + row0 * prob.Cs0 + col)[0] = __halves2half2(h00, h01);

Kernel source

v1b.py663 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu NVIDIA

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

"""
Optimized tcgen05 + TMA kernel v3:
- Batched kernel launch (all tiles per problem in one launch)
- Double-buffered software pipelining (overlap TMA with MMA)
"""

"""
g: 8; k: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; m: [80, 176, 128, 72, 64, 248, 96, 160]; n: [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096]; seed: 1111
 ⏱ 641 ± 1.2 µs
 ⚡ 625 µs 🐌 720 µs

g: 8; k: [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]; m: [40, 76, 168, 72, 164, 148, 196, 160]; n: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; seed: 1111
 ⏱ 451 ± 0.4 µs
 ⚡ 436 µs 🐌 470 µs

g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
 ⏱ 142 ± 0.4 µs
 ⚡ 138 µs 🐌 169 µs

g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
 ⏱ 80.2 ± 0.19 µs
 ⚡ 77.2 µs 🐌 88.3 µs
"""

CUDA_SRC = """
#include <vector>
#include <cstdint>
#include <cstdio>
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>

static inline int ceil_div(int a, int b) { return (a + b - 1) / b; }

#define CUDA_CHECK(expr)                                                       \\
  do {                                                                         \\
    cudaError_t _err = (expr);                                                 \\
    TORCH_CHECK(_err == cudaSuccess, "CUDA error: ", cudaGetErrorString(_err)); \\
  } while (0)

constexpr int WARP_SIZE = 32;

// Cache hints
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;

// Work item for batched kernel
struct WorkItem {
  int tile_m;
  int tile_n;
};

// Problem descriptor
struct ProblemDesc {
  int M, N, K;
  int64_t Cs0, Cs1, Cs2;
};

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

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

__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size) {
  asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" 
              :: "r"(mbar_addr), "r"(size) : "memory");
}

__device__ 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 LAB_WAIT;\\n\\t"
    "}"
    :: "r"(mbar_addr), "r"(phase), "r"(ticks)
  );
}

// 3D TMA load
template <int CTA_GROUP = 1>
__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::%7.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), "n"(CTA_GROUP)
    : "memory"
  );
}

// 1D TMA load
template <int CTA_GROUP = 1>
__device__ inline void tma_1d_gmem2smem(int dst, const void *tmap_ptr, int x, 
                                         int mbar_addr, uint64_t cache_policy) {
  asm volatile(
    "cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%5.L2::cache_hint "
    "[%0], [%1, {%2}], [%3], %4;"
    :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
    : "memory"
  );
}

template <int CTA_GROUP = 1>
__device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
  asm volatile(
    "tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;\\n"
    :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP)
  );
}

template <int CTA_GROUP = 1>
__device__ __forceinline__ void tcgen05_commit(int mbar_addr) {
  asm volatile(
    "tcgen05.commit.cta_group::%1.mbarrier::arrive::one.shared::cluster.b64 [%0];\\n"
    :: "r"(mbar_addr), "n"(CTA_GROUP) : "memory"
  );
}

template <int CTA_GROUP = 1>
__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::%7.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),
       "n"(CTA_GROUP)
  );
}

// TMEM load helper
struct SHAPE {
  static constexpr char _16x256b[] = ".16x256b";
};
struct NUM {
  static constexpr char x16[] = ".x16";
};

template <const char *SHAPE, const char *NUM>
__device__ inline
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), "C"(NUM));
}

__device__ inline void tcgen05_ld_16x256b_x16(float *tmp, int row, int col) {
  tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col);
}

__device__ __forceinline__ void tcgen05_dealloc_cols_cta1(uint32_t tmem, int count) {
  asm volatile(
    "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\\n"
    :: "r"(tmem), "r"(count)
    : "memory"
  );
}

// Kernel Configuration
constexpr int TMA_BLOCK_M = 128;
constexpr int TMA_BLOCK_N = 128;
constexpr int TMA_BLOCK_K = 256;
constexpr int NUM_STAGES = 2;  // Double buffering
constexpr int TMA_NUM_WARPS = 4;
constexpr int TMEM_COLS = TMA_BLOCK_N * 2;

constexpr int TMA_A_SMEM_BYTES = TMA_BLOCK_M * (TMA_BLOCK_K / 2);   // 16KB
constexpr int TMA_B_SMEM_BYTES = TMA_BLOCK_N * (TMA_BLOCK_K / 2);   // 16KB
constexpr int TMA_SFA_SMEM_BYTES = TMA_BLOCK_M * (TMA_BLOCK_K / 16); // 2KB
constexpr int TMA_SFB_SMEM_BYTES = TMA_BLOCK_N * (TMA_BLOCK_K / 16); // 2KB
constexpr int STAGE_SIZE = TMA_A_SMEM_BYTES + TMA_B_SMEM_BYTES + TMA_SFA_SMEM_BYTES + TMA_SFB_SMEM_BYTES;
constexpr int SMEM_SIZE = STAGE_SIZE * NUM_STAGES + 64;  // 2 stages + mbarriers

// Helper to issue TMA loads for a given stage
__device__ __forceinline__ void issue_tma_loads(
    int A_smem, int B_smem, int SFA_smem, int SFB_smem, int mbar_addr,
    const CUtensorMap* A_tmap, const CUtensorMap* B_tmap,
    const CUtensorMap* SFA_tmap, const CUtensorMap* SFB_tmap,
    int m_offset, int n_offset, int k_iter,
    int m_tile_idx, int n_tile_idx, int sf_bytes_per_m_tile, int sf_k_per_iter
) {
    const int off_k = k_iter * TMA_BLOCK_K;
    tma_3d_gmem2smem<1>(A_smem, A_tmap, 0, m_offset, off_k / 256, mbar_addr, EVICT_NORMAL);
    tma_3d_gmem2smem<1>(B_smem, B_tmap, 0, n_offset, off_k / 256, mbar_addr, EVICT_FIRST);
    
    const int off_sfa = m_tile_idx * sf_bytes_per_m_tile + k_iter * sf_k_per_iter;
    const int off_sfb = n_tile_idx * sf_bytes_per_m_tile + k_iter * sf_k_per_iter;
    tma_1d_gmem2smem<1>(SFA_smem, SFA_tmap, off_sfa / 8, mbar_addr, EVICT_NORMAL);
    tma_1d_gmem2smem<1>(SFB_smem, SFB_tmap, off_sfb / 8, mbar_addr, EVICT_FIRST);
    mbarrier_arrive_expect_tx(mbar_addr, STAGE_SIZE);
}

// Helper to issue MMA for a given stage
__device__ __forceinline__ void issue_mma(
    int A_smem, int B_smem, int SFA_smem, int SFB_smem,
    int d_tmem, int sfa_tmem, int sfb_tmem, uint32_t idesc,
    int mma_mbar, int k_iter, bool first_k_iter
) {
    constexpr uint64_t SF_desc = (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
    const uint64_t SFA_desc = SF_desc | ((uint64_t)SFA_smem >> 4ULL);
    const uint64_t SFB_desc = SF_desc | ((uint64_t)SFB_smem >> 4ULL);

    constexpr uint64_t AB_desc = (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
    const uint64_t A_desc = AB_desc | ((uint64_t)A_smem >> 4ULL);
    const uint64_t B_desc = AB_desc | ((uint64_t)B_smem >> 4ULL);

    #pragma unroll
    for (int k = 0; k < (TMA_BLOCK_K / 64); k++) {
        tcgen05_cp_nvfp4<1>(sfa_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
        tcgen05_cp_nvfp4<1>(sfb_tmem + k * 4, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
    }

    #pragma unroll
    for (int k2 = 0; k2 < (TMA_BLOCK_K / 64); k2++) {
        const uint64_t a_desc = A_desc + (uint64_t)k2 * (32ULL >> 4ULL);
        const uint64_t b_desc = B_desc + (uint64_t)k2 * (32ULL >> 4ULL);
        const int enable_input_d = (first_k_iter && k2 == 0) ? 0 : 1;
        tcgen05_mma_nvfp4<1>(d_tmem, a_desc, b_desc, idesc, sfa_tmem + k2 * 4, sfb_tmem + k2 * 4, enable_input_d);
    }

    tcgen05_commit<1>(mma_mbar);
}

// Pipelined kernel with double buffering
__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)
void grouped_gemm_tcgen_tma_v3(
  const __grid_constant__ CUtensorMap A_tmap,
  const __grid_constant__ CUtensorMap B_tmap,
  const __grid_constant__ CUtensorMap SFA_tmap,
  const __grid_constant__ CUtensorMap SFB_tmap,
  half* __restrict__ C_ptr,
  const ProblemDesc prob,
  const WorkItem* __restrict__ work_items
) {
  const int block_idx = blockIdx.x;
  const WorkItem& work = work_items[block_idx];
  
  const int m_offset = work.tile_m * TMA_BLOCK_M;
  const int n_offset = work.tile_n * TMA_BLOCK_N;
  
  const int M = prob.M;
  const int N = prob.N;
  const int K = prob.K;
  
  const int tid = threadIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  // Shared memory layout with double buffering
  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
  
  // Stage 0 buffers
  const int A_smem_0 = smem;
  const int B_smem_0 = A_smem_0 + TMA_A_SMEM_BYTES;
  const int SFA_smem_0 = B_smem_0 + TMA_B_SMEM_BYTES;
  const int SFB_smem_0 = SFA_smem_0 + TMA_SFA_SMEM_BYTES;
  
  // Stage 1 buffers
  const int A_smem_1 = SFB_smem_0 + TMA_SFB_SMEM_BYTES;
  const int B_smem_1 = A_smem_1 + TMA_A_SMEM_BYTES;
  const int SFA_smem_1 = B_smem_1 + TMA_B_SMEM_BYTES;
  const int SFB_smem_1 = SFA_smem_1 + TMA_SFA_SMEM_BYTES;
  
  // Mbarriers (4 total: TMA0, TMA1, MMA0, MMA1)
  const int mbar_base = SFB_smem_1 + TMA_SFB_SMEM_BYTES;
  const int tma_mbar[2] = {mbar_base, mbar_base + 8};
  const int mma_mbar[2] = {mbar_base + 16, mbar_base + 24};
  
  // Stage buffer arrays
  const int A_smem[2] = {A_smem_0, A_smem_1};
  const int B_smem[2] = {B_smem_0, B_smem_1};
  const int SFA_smem[2] = {SFA_smem_0, SFA_smem_1};
  const int SFB_smem[2] = {SFB_smem_0, SFB_smem_1};

  // Initialize mbarriers
  if (tid == 0) {
    mbarrier_init(tma_mbar[0], 1);
    mbarrier_init(tma_mbar[1], 1);
    mbarrier_init(mma_mbar[0], 1);
    mbarrier_init(mma_mbar[1], 1);
    asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
  }
  __syncthreads();

  // Allocate TMEM
  if (warp_id == 0) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_COLS));
  }
  __syncthreads();

  const int tmem_base = 0;
  const int d_tmem = tmem_base;
  const int sfa_tmem = tmem_base + TMA_BLOCK_N;
  const int sfb_tmem = sfa_tmem + 4 * (TMA_BLOCK_K / 64);

  constexpr uint32_t idesc = (1U << 7U) | (1U << 10U) | ((uint32_t)TMA_BLOCK_N >> 3U << 17U) | ((uint32_t)TMA_BLOCK_M >> 7U << 27U);

  const int num_k_iters = (K + TMA_BLOCK_K - 1) / TMA_BLOCK_K;
  
  // Precompute scale factor offsets
  const int m_tile_idx = m_offset / TMA_BLOCK_M;
  const int n_tile_idx = n_offset / TMA_BLOCK_N;
  const int sf_bytes_per_m_tile = TMA_BLOCK_M * (K / 16);
  const int sf_k_per_iter = TMA_SFA_SMEM_BYTES;
  
  // PROLOGUE: Prime the pipeline - load first stage
  if (tid == 0) {
    issue_tma_loads(A_smem[0], B_smem[0], SFA_smem[0], SFB_smem[0], tma_mbar[0],
                    &A_tmap, &B_tmap, &SFA_tmap, &SFB_tmap,
                    m_offset, n_offset, 0,
                    m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);
  }
  __syncthreads();
  
  // MAIN LOOP: Pipelined execution
  for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {
    const int stage = k_iter % NUM_STAGES;
    const int next_stage = (k_iter + 1) % NUM_STAGES;
    
    // Wait for TMA to complete for current stage
    if (tid == 0) {
      mbarrier_wait(tma_mbar[stage], k_iter / NUM_STAGES);
    }
    __syncthreads();
    
    // Issue TMA for next iteration (overlapped with MMA below)
    if (tid == 0 && (k_iter + 1) < num_k_iters) {
      // Wait for previous MMA to finish with next_stage buffer
      if (k_iter >= 1) {
        mbarrier_wait(mma_mbar[next_stage], (k_iter - 1) / NUM_STAGES);
      }
      
      issue_tma_loads(A_smem[next_stage], B_smem[next_stage], SFA_smem[next_stage], SFB_smem[next_stage], 
                      tma_mbar[next_stage],
                      &A_tmap, &B_tmap, &SFA_tmap, &SFB_tmap,
                      m_offset, n_offset, k_iter + 1,
                      m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);
    }
    
    // Issue MMA for current stage
    if (tid == 0) {
      issue_mma(A_smem[stage], B_smem[stage], SFA_smem[stage], SFB_smem[stage],
                d_tmem, sfa_tmem, sfb_tmem, idesc, mma_mbar[stage], k_iter, k_iter == 0);
    }
    __syncthreads();
  }

  // Wait for last MMA
  if (num_k_iters > 0 && tid == 0) {
    const int last_stage = (num_k_iters - 1) % NUM_STAGES;
    mbarrier_wait(mma_mbar[last_stage], (num_k_iters - 1) / NUM_STAGES);
  }
  __syncthreads();
  
  asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

  // Epilogue
  if (tid < TMA_BLOCK_M) {
    for (int m = 0; m < 32 / 16; m++) {
      float tmp[TMA_BLOCK_N / 2];
      tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, 0);
      asm volatile("tcgen05.wait::ld.sync.aligned;\\n");

      for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
        const int row0 = m_offset + warp_id * 32 + m * 16 + lane_id / 4;
        const int row1 = row0 + 8;
        const int col = n_offset + i * 8 + (lane_id % 4) * 2;

        const int idx = i * 4;
        half h00 = __float2half_rn(tmp[idx + 0]);
        half h01 = __float2half_rn(tmp[idx + 1]);
        half h10 = __float2half_rn(tmp[idx + 2]);
        half h11 = __float2half_rn(tmp[idx + 3]);

        if (row0 < M && col < N) {
          if (prob.Cs1 == 1) {
            if (col + 1 < N) {
              reinterpret_cast<half2*>(C_ptr + row0 * prob.Cs0 + col)[0] = __halves2half2(h00, h01);
            } else {
              C_ptr[row0 * prob.Cs0 + col] = h00;
            }
          } else {
            C_ptr[row0 * prob.Cs0 + col * prob.Cs1] = h00;
            if (col + 1 < N) C_ptr[row0 * prob.Cs0 + (col + 1) * prob.Cs1] = h01;
          }
        }

        if (row1 < M && col < N) {
          if (prob.Cs1 == 1) {
            if (col + 1 < N) {
              reinterpret_cast<half2*>(C_ptr + row1 * prob.Cs0 + col)[0] = __halves2half2(h10, h11);
            } else {
              C_ptr[row1 * prob.Cs0 + col] = h10;
            }
          } else {
            C_ptr[row1 * prob.Cs0 + col * prob.Cs1] = h10;
            if (col + 1 < N) C_ptr[row1 * prob.Cs0 + (col + 1) * prob.Cs1] = h11;
          }
        }
      }
    }
  }

  __syncthreads();
  if (warp_id == 0) {
    tcgen05_dealloc_cols_cta1(tmem_base, TMEM_COLS);
  }
}

// Tensor Map Initialization
void init_AB_tmap_u4(
  CUtensorMap *tmap,
  const void *ptr,
  uint64_t global_height, uint64_t global_width,
  uint32_t shared_height, uint32_t shared_width
) {
  TORCH_CHECK(ptr != nullptr, "init_AB_tmap_u4: ptr is null");
  TORCH_CHECK(((uintptr_t)ptr % 16) == 0, "ptr must be 16-byte aligned");
  TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");
  
  constexpr uint32_t rank = 3;
  uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
  uint64_t globalStrides[rank-1] = {global_width / 2, 128};
  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,
    CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
    CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
  );
  TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed for AB");
}

void init_SF_tmap_reordered(CUtensorMap *tmap, const void *ptr, uint64_t global_size_bytes, uint32_t shared_size_bytes) {
  TORCH_CHECK(ptr != nullptr && ((uintptr_t)ptr % 16) == 0, "SF ptr alignment");
  TORCH_CHECK(global_size_bytes > 0 && global_size_bytes >= shared_size_bytes, "SF size");
  
  constexpr uint32_t rank = 1;
  uint64_t globalDim[rank]       = {global_size_bytes / 8};
  uint64_t globalStrides[rank-1] = {};
  uint32_t boxDim[rank]          = {shared_size_bytes / 8};
  uint32_t elementStrides[rank]  = {1};

  auto err = cuTensorMapEncodeTiled(
    tmap,
    CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_INT64,
    rank, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides,
    CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
    CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE,
    CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
    CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
  );
  TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed for SF");
}

// Host Entry Point
std::vector<at::Tensor> group_gemm(
    std::vector<at::Tensor> A_list,
    std::vector<at::Tensor> B_list,
    std::vector<at::Tensor> C_list,
    std::vector<at::Tensor> sfa_list,
    std::vector<at::Tensor> sfb_list,
    at::Tensor sizes_cpu
) {
    int64_t G = A_list.size();
    TORCH_CHECK(B_list.size() == G && C_list.size() == G, "A/B/C list sizes must match");
    TORCH_CHECK(sfa_list.size() == G && sfb_list.size() == G, "sfa/sfb list sizes must match");
    TORCH_CHECK(sizes_cpu.device().is_cpu() && sizes_cpu.scalar_type() == at::kLong, "sizes must be CPU int64");

    TORCH_CHECK(A_list[0].is_cuda(), "A must be CUDA");
    auto dev = A_list[0].device();
    c10::cuda::CUDAGuard device_guard(dev);

    auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();
    
    CUDA_CHECK(cudaFuncSetAttribute(
        grouped_gemm_tcgen_tma_v3,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        SMEM_SIZE
    ));

    for (int64_t prob_idx = 0; prob_idx < G; prob_idx++) {
        at::Tensor A = A_list[prob_idx];
        at::Tensor B = B_list[prob_idx];
        at::Tensor C = C_list[prob_idx];
        at::Tensor sfa = sfa_list[prob_idx];
        at::Tensor sfb = sfb_list[prob_idx];

        int64_t M = sizes_accessor[prob_idx][0];
        int64_t N = sizes_accessor[prob_idx][1];
        int64_t K = sizes_accessor[prob_idx][2];
        
        if (A.stride(1) != 1 || B.stride(1) != 1) continue;
        TORCH_CHECK((K % TMA_BLOCK_K) == 0, "K must be multiple of ", TMA_BLOCK_K);

        ProblemDesc prob{};
        prob.M = M; prob.N = N; prob.K = K;
        prob.Cs0 = C.stride(0); prob.Cs1 = C.stride(1); prob.Cs2 = C.stride(2);

        int num_tiles_m = ceil_div((int)M, TMA_BLOCK_M);
        int num_tiles_n = ceil_div((int)N, TMA_BLOCK_N);
        int num_tiles = num_tiles_m * num_tiles_n;
        
        std::vector<WorkItem> work_items;
        work_items.reserve(num_tiles);
        for (int tm = 0; tm < num_tiles_m; tm++) {
            for (int tn = 0; tn < num_tiles_n; tn++) {
                work_items.push_back({tm, tn});
            }
        }

        WorkItem* d_work_items;
        CUDA_CHECK(cudaMalloc(&d_work_items, work_items.size() * sizeof(WorkItem)));
        CUDA_CHECK(cudaMemcpy(d_work_items, work_items.data(), work_items.size() * sizeof(WorkItem), cudaMemcpyHostToDevice));

        CUtensorMap A_tmap, B_tmap, SFA_tmap, SFB_tmap;
        init_AB_tmap_u4(&A_tmap, A.data_ptr(), A.size(0), K, TMA_BLOCK_M, TMA_BLOCK_K);
        init_AB_tmap_u4(&B_tmap, B.data_ptr(), B.size(0), K, TMA_BLOCK_N, TMA_BLOCK_K);
        init_SF_tmap_reordered(&SFA_tmap, sfa.data_ptr(), sfa.numel() * sfa.element_size(), TMA_SFA_SMEM_BYTES);
        init_SF_tmap_reordered(&SFB_tmap, sfb.data_ptr(), sfb.numel() * sfb.element_size(), TMA_SFB_SMEM_BYTES);

        grouped_gemm_tcgen_tma_v3<<<num_tiles, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE>>>(
            A_tmap, B_tmap, SFA_tmap, SFB_tmap,
            (half*)C.data_ptr(), prob, d_work_items
        );
        
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaFree(d_work_items));
    }
    
    CUDA_CHECK(cudaGetLastError());
    return C_list;
}

TORCH_LIBRARY(my_module, m) {
    m.def("group_gemm(Tensor[] a, Tensor[] b, Tensor[] c, Tensor[] sfa, Tensor[] sfb, Tensor sizes) -> Tensor[]");
    m.impl("group_gemm", &group_gemm);
}
"""

load_inline(
    "group_gemm",
    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",
    ],
    extra_ldflags=["-lcuda"],
)
group_gemm = torch.ops.my_module.group_gemm

def custom_kernel(data: input_t) -> output_t:
    abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data

    def create_aligned_tensor(shape, dtype, device, align_bytes=128):
        numel = 1
        for s in shape:
            numel *= s
        for attempt in range(10):
            tensor_uint8 = torch.zeros(numel, dtype=torch.uint8, device=device)
            if tensor_uint8.data_ptr() % align_bytes == 0:
                return tensor_uint8.view(dtype).view(shape)
        raise RuntimeError(f"Failed to allocate 128-byte aligned tensor")

    A_list = []
    B_list = []
    C_list = [t[2] for t in abc_tensors]
    
    sfa_list = []
    sfb_list = []
    for idx, ((sfa_reord, sfb_reord), (m, n, k, l)) in enumerate(zip(sfasfb_reordered_tensors, problem_sizes)):
        sfa_perm = sfa_reord.permute(2, 4, 0, 1, 3, 5).contiguous()
        sfb_perm = sfb_reord.permute(2, 4, 0, 1, 3, 5).contiguous()
        sfa_list.append(sfa_perm)
        sfb_list.append(sfb_perm)

    for (a, b, c), _ in zip(abc_tensors, problem_sizes):
        is_aligned_a = (a.data_ptr() % 128) == 0
        needs_pad_a = (a.size(0) % 128 != 0)
        
        if needs_pad_a or not is_aligned_a:
            pad_m = 128 - (a.size(0) % 128) if needs_pad_a else 0
            new_shape = list(a.shape)
            new_shape[0] += pad_m
            new_a = create_aligned_tensor(new_shape, a.dtype, a.device, align_bytes=128)
            new_a[:a.size(0), :, :] = a
            A_list.append(new_a)
        else:
            A_list.append(a)

        is_aligned_b = (b.data_ptr() % 128) == 0
        needs_pad_b = (b.size(0) % 128 != 0)

        if needs_pad_b or not is_aligned_b:
            pad_n = 128 - (b.size(0) % 128) if needs_pad_b else 0
            new_shape = list(b.shape)
            new_shape[0] += pad_n
            new_b = create_aligned_tensor(new_shape, b.dtype, b.device, align_bytes=128)
            new_b[:b.size(0), :, :] = b
            B_list.append(new_b)
        else:
            B_list.append(b)
    
    sizes_cpu = torch.tensor(problem_sizes, dtype=torch.int64, device='cpu')
    group_gemm(A_list, B_list, C_list, sfa_list, sfb_list, sizes_cpu)

    return C_list
scrolls · 663 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