Skip to content
KernelIndex
Search⌘K

submission 492630

XoTic · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-492630?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
52.0µs
#58 of 145
2026-02-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:094580a5a3eaf2a3337ed35d48d568300b04dfdca3a0d900e6fe99f94e25367b
license declaredunknown
license concludedunknown
authorsXoTic
imported2026-08-15

Techniques

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

fused-epiloguetemplate <bool LOW_M_EPILOGUE>
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
num-warps = 4constexpr int TMA_NUM_WARPS = 4;
persistent-kerneltemplate <int NUM_STAGES, bool LOW_M_EPILOGUE, bool PERSISTENT>
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
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;
tmaCUtensorMap A_tmap;
vector-width = half2half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);

Kernel source

v3.py929 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

"""
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
 ⏱ 169 ± 0.1 µs
 ⚡ 169 µs 🐌 170 µ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
 ⏱ 159 ± 0.1 µs
 ⚡ 159 µs 🐌 159 µs

g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
 ⏱ 45.6 ± 0.04 µs
 ⚡ 45.5 µs 🐌 45.7 µs

g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
 ⏱ 20.0 ± 0.02 µs
 ⚡ 19.9 µs 🐌 20.0 µs
"""

CUDA_SRC = """
#include <vector>
#include <unordered_map>
#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 persistent kernel
struct WorkItem {
  int problem_idx;
  int tile_m;
  int tile_n;
};

// Global Problem Info stored in Global Memory
struct __align__(128) ProblemInfo {
  CUtensorMap A_tmap;
  CUtensorMap B_tmap;
  const char* SFA_ptr;  // points to underlying contiguous storage in [l, mn/128, (k/16)/4, 32, 4, 4]
  const char* SFB_ptr;
  half* C_ptr;
  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__ inline void mbarrier_inval(int mbar_addr) {
  asm volatile("mbarrier.inval.shared::cta.b64 [%0];" :: "r"(mbar_addr) : "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"
  );
}

__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)
    : "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"
  );
}

// SMEM Layout:
// [ProblemInfo] (aligned to 128)
// [Stage 0]
// [Stage 1]
// [Stage 2]
// [Mbarriers] (6 mbarriers: 3 for TMA complete, 3 for MMA commit)

constexpr int TMA_BLOCK_M = 128;
constexpr int TMA_BLOCK_N = 128;
constexpr int TMA_BLOCK_K = 256;
constexpr int NUM_STAGES_MAIN = 3;  // Triple buffering
constexpr int NUM_STAGES_LOW_M = 2; // Lower SMEM footprint for small-M tiles
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; // ~36KB

constexpr int MBAR_BYTES_MAIN = ((2 * NUM_STAGES_MAIN * 8 + 63) & ~63);
constexpr int MBAR_BYTES_LOW_M = ((2 * NUM_STAGES_LOW_M * 8 + 63) & ~63);
constexpr int SMEM_SIZE_MAIN = STAGE_SIZE * NUM_STAGES_MAIN + MBAR_BYTES_MAIN;
constexpr int SMEM_SIZE_LOW_M = STAGE_SIZE * NUM_STAGES_LOW_M + MBAR_BYTES_LOW_M;
constexpr int LOW_M_THRESHOLD = 96;

template <bool LOW_M_EPILOGUE>
__device__ __forceinline__ void epilogue_store(
    const ProblemInfo& prob,
    int m_offset,
    int n_offset,
    int tid,
    int warp_id,
    int lane_id
) {
  if (tid >= TMA_BLOCK_M) return;

  const int M = prob.M;
  const int N = prob.N;
  half* C_ptr = prob.C_ptr;
  const int64_t Cs0 = prob.Cs0;
  const int64_t Cs1 = prob.Cs1;

  const bool full_n = (n_offset + TMA_BLOCK_N <= N);
  const bool full_m = (m_offset + TMA_BLOCK_M <= M);
  const bool full_tile = full_n && full_m;
  const bool contiguous = (Cs1 == 1);

  const int warp_row_base = m_offset + warp_id * 32;
  if (LOW_M_EPILOGUE && warp_row_base >= M) return;

  int m_iters = 2;
  if (LOW_M_EPILOGUE) {
    const int remaining = M - warp_row_base;
    m_iters = (remaining <= 16) ? 1 : 2;
  }

  const int lane_row = lane_id >> 2;
  const int lane_col = (lane_id & 3) * 2;

  for (int m = 0; m < m_iters; ++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");

    const int row0 = warp_row_base + m * 16 + lane_row;
    const int row1 = row0 + 8;

    if (contiguous) {
      if (full_tile) {
        half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);
        half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset);
        #pragma unroll
        for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
          const int idx = i * 4;
          const int col = i * 8 + lane_col;
          const int h2_idx = col >> 1;
          row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));
          row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));
        }
        continue;
      }

      const bool row0_in = row0 < M;
      const bool row1_in = row1 < M;
      if (full_n) {
        half2* row0_ptr = row0_in ? reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset) : nullptr;
        half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset) : nullptr;
        #pragma unroll
        for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
          const int idx = i * 4;
          const int col = i * 8 + lane_col;
          const int h2_idx = col >> 1;
          if (row0_in) {
            row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));
          }
          if (row1_in) {
            row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));
          }
        }
      } else {
        #pragma unroll
        for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
          const int idx = i * 4;
          const int col = n_offset + i * 8 + lane_col;
          if (col < N) {
            const half h00 = __float2half_rn(tmp[idx + 0]);
            const half h01 = __float2half_rn(tmp[idx + 1]);
            const half h10 = __float2half_rn(tmp[idx + 2]);
            const half h11 = __float2half_rn(tmp[idx + 3]);
            if (row0_in) {
              if (col + 1 < N) {
                reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = __halves2half2(h00, h01);
              } else {
                C_ptr[row0 * Cs0 + col] = h00;
              }
            }
            if (row1_in) {
              if (col + 1 < N) {
                reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = __halves2half2(h10, h11);
              } else {
                C_ptr[row1 * Cs0 + col] = h10;
              }
            }
          }
        }
      }
    } else {
      const bool row0_in = row0 < M;
      const bool row1_in = row1 < M;
      if (full_n) {
        #pragma unroll
        for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
          const int idx = i * 4;
          const int col = n_offset + i * 8 + lane_col;
          const half h00 = __float2half_rn(tmp[idx + 0]);
          const half h01 = __float2half_rn(tmp[idx + 1]);
          const half h10 = __float2half_rn(tmp[idx + 2]);
          const half h11 = __float2half_rn(tmp[idx + 3]);
          if (row0_in) {
            C_ptr[row0 * Cs0 + col * Cs1] = h00;
            C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
          }
          if (row1_in) {
            C_ptr[row1 * Cs0 + col * Cs1] = h10;
            C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
          }
        }
      } else {
        #pragma unroll
        for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
          const int idx = i * 4;
          const int col = n_offset + i * 8 + lane_col;
          if (col < N) {
            const half h00 = __float2half_rn(tmp[idx + 0]);
            const half h01 = __float2half_rn(tmp[idx + 1]);
            const half h10 = __float2half_rn(tmp[idx + 2]);
            const half h11 = __float2half_rn(tmp[idx + 3]);
            if (row0_in) {
              C_ptr[row0 * Cs0 + col * Cs1] = h00;
              if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
            }
            if (row1_in) {
              C_ptr[row1 * Cs0 + col * Cs1] = h10;
              if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
            }
          }
        }
      }
    }
  }
}

// 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 ProblemInfo* prob,
    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, &prob->A_tmap, 0, m_offset, off_k / 256, mbar_addr, EVICT_NORMAL);
    tma_3d_gmem2smem<1>(B_smem, &prob->B_tmap, 0, n_offset, off_k / 256, mbar_addr, EVICT_FIRST);
    
    // Scale factors live in the underlying contiguous storage order:
    //   [l=1, mn/128, (k/16)/4, 32, 4, 4], with each (32,4,4) tile = 512 bytes.
    // For BLOCK_K=256 we need 4 consecutive 512B tiles per k_iter (total 2048B).
    const int rest_k = prob->K / 64;  // (K/16)/4
    const int k_blk = off_k / 64;     // (off_k/16)/4
    const char* SFA_src = prob->SFA_ptr + (int64_t)(m_tile_idx * rest_k + k_blk) * 512;
    const char* SFB_src = prob->SFB_ptr + (int64_t)(n_tile_idx * rest_k + k_blk) * 512;
    tma_gmem2smem(SFA_smem, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, EVICT_NORMAL);
    tma_gmem2smem(SFB_smem, SFB_src, TMA_SFB_SMEM_BYTES, mbar_addr, EVICT_FIRST);
    mbarrier_arrive_expect_tx(mbar_addr, STAGE_SIZE);
}

// Single Kernel (template for main/low-M variants, with optional persistence)
template <int NUM_STAGES, bool LOW_M_EPILOGUE, bool PERSISTENT>
__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)
void grouped_gemm_tcgen_tma_v3_persistent(
  const ProblemInfo* __restrict__ global_probs,
  const WorkItem* __restrict__ work_items,
  int num_items,
  int* __restrict__ work_counter  // Only used when PERSISTENT=true
) {
  const int tid = threadIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  // Shared Memory Setup (constant addresses)
  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

  // Pipeline Buffers
  const int stages_base = smem_base;
  int A_smem[NUM_STAGES];
  #pragma unroll
  for (int i = 0; i < NUM_STAGES; ++i) {
    A_smem[i] = stages_base + STAGE_SIZE * i;
  }
  const int B_off = TMA_A_SMEM_BYTES;
  const int SFA_off = B_off + TMA_B_SMEM_BYTES;
  const int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;
  
  // Mbarriers (constant addresses)
  const int mbar_base = stages_base + STAGE_SIZE * NUM_STAGES;
  int tma_mbar[NUM_STAGES];
  int mma_mbar[NUM_STAGES];
  #pragma unroll
  for (int i = 0; i < NUM_STAGES; ++i) {
    tma_mbar[i] = mbar_base + i * 8;
    mma_mbar[i] = mbar_base + (NUM_STAGES + i) * 8;
  }

  // TMEM addresses (constant)
  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);

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

  // Shared work index (only used in persistent mode)
  __shared__ int shared_work_idx;

  // ================================================================
  // WORK DISPATCH: Persistent loop vs single-tile
  // ================================================================
  int work_idx;
  if constexpr (PERSISTENT) {
    // Fetch work atomically
    if (tid == 0) {
      shared_work_idx = atomicAdd(work_counter, 1);
    }
    __syncthreads();
    work_idx = shared_work_idx;
  } else {
    // Non-persistent: each CTA handles one tile via blockIdx
    work_idx = blockIdx.x;
  }

  // Main processing loop (runs once for non-persistent, loops for persistent)
  while (work_idx < num_items) {
    const WorkItem& work = work_items[work_idx];
    const ProblemInfo& prob = global_probs[work.problem_idx];

    const int m_offset = work.tile_m * TMA_BLOCK_M;
    const int n_offset = work.tile_n * TMA_BLOCK_N;
    const int K = prob.K;
    const int num_k_iters = (K + TMA_BLOCK_K - 1) / TMA_BLOCK_K;

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

    // Initialize or reinit mbarriers
    // Note: mbarrier_inval is not strictly needed if we ensure all threads synced and state is clean via init
    
    if (tid == 0) {
      for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(tma_mbar[i], 1);
      for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(mma_mbar[i], 1);
      asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
    }
    __syncthreads();

  // ----------------------------------------------------------------
  // PRODUCER WARP (Warp 0): Issues TMA
  // ----------------------------------------------------------------
  if (warp_id == 0) {
      if (lane_id == 0) {
          for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {
              const int stage = k_iter % NUM_STAGES;
              // Wait for buffer to be free (consumed by MMA)
              // Ideally, mma_mbar tracks "buffer consumed".
              // Init count 1.
              // Logic:
              // k=0: buffer is free (init state). 
              // But we need to sync with consumer?
              // Let's assume mma_mbar is signaled when MMA is done using the buffer.
              // Initial state: buffers are free. mma_mbar should NOT block for first use.
              // But mbarrier logic: wait() blocks until phase flips.
              // We need to manage phases carefully.
              
              // Correct logic:
              // TMA thread waits for buffer to be available.
              // For k < NUM_STAGES, buffers are initially available.
              // For k >= NUM_STAGES, wait for previous usage to complete.
              
              if (k_iter >= NUM_STAGES) {
                  // Wait for stage to be released by MMA
                  // Corresponding k was k_iter - NUM_STAGES
                  mbarrier_wait(mma_mbar[stage], (k_iter - NUM_STAGES) / NUM_STAGES); // Wait for phase flip?
                  // Or just use the same phase logic as coupled.
                  // Coupled used: mbarrier_wait(mma_mbar[next_stage], ...)
              }

              // Issue TMA
              const int stage_base = A_smem[stage];
              issue_tma_loads(stage_base, stage_base + B_off, stage_base + SFA_off, stage_base + SFB_off,
                              tma_mbar[stage], &prob, m_offset, n_offset, k_iter,
                              m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);
          }
      }
  }

  // ----------------------------------------------------------------
  // CONSUMER WARP (Warp 1): Issues MMA
  // ----------------------------------------------------------------
  else if (warp_id == 1) {
      if (lane_id == 0) {
          for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {
              const int stage = k_iter % NUM_STAGES;
              
              // Wait for data ready (TMA complete)
              mbarrier_wait(tma_mbar[stage], k_iter / NUM_STAGES);
              
              const int stage_base = A_smem[stage];
              
              // Issue MMA
              constexpr uint64_t SF_desc = (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
              const uint64_t SFA_desc = SF_desc | ((uint64_t)(stage_base + SFA_off) >> 4ULL);
              const uint64_t SFB_desc = SF_desc | ((uint64_t)(stage_base + SFB_off) >> 4ULL);

              constexpr uint64_t AB_desc = (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
              const uint64_t A_desc = AB_desc | ((uint64_t)stage_base >> 4ULL);
              const uint64_t B_desc = AB_desc | ((uint64_t)(stage_base + B_off) >> 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 = (k_iter == 0 && 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);
              }
              
              // Commit MMA -> Signals mma_mbar[stage]
              // When commit reaches mma_mbar, it allows TMA producer to reuse buffer for next phase
              tcgen05_commit<1>(mma_mbar[stage]);
          }
          
          // Final Wait
          // Wait for last commit to ensure all instructions retired? 
          // Actually, we must ensure all work is done before exiting.
          // Wait for the last commit to finish
              if (num_k_iters > 0) {
                  const int last_stage = (num_k_iters - 1) % NUM_STAGES;
                  mbarrier_wait(mma_mbar[last_stage], (num_k_iters - 1) / NUM_STAGES);
              }
      }
  }
  
  // ----------------------------------------------------------------
  // SYNCHRONIZATION
  // ----------------------------------------------------------------
  __syncthreads(); // Wait for all warps to finish
  
  // Ensure MMA results visible
  asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

  // Epilogue
  epilogue_store<LOW_M_EPILOGUE>(prob, m_offset, n_offset, tid, warp_id, lane_id);

  __syncthreads();
  
    // Loop control: continue (persistent) or break (non-persistent)
    if constexpr (PERSISTENT) {
      // Fetch next work
      if (tid == 0) {
        shared_work_idx = atomicAdd(work_counter, 1);
      }
      __syncthreads();
      work_idx = shared_work_idx;
    } else {
      // Non-persistent: exit after single tile
      break;
    }
  } // End main loop
  
  // Deallocate TMEM once at end
  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};

  // cuTensorMapEncodeTiled is relatively expensive on the host; amortize by caching
  // a per-shape template and then patching only the base address each call.
  struct Key {
    uint64_t gh, gw;
    uint32_t sh, sw;
  };
  struct KeyHash {
    size_t operator()(const Key& k) const noexcept {
      uint64_t h = k.gh;
      h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
      h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
      return (size_t)h;
    }
  };
  struct KeyEq {
    bool operator()(const Key& a, const Key& b) const noexcept {
      return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw;
    }
  };
  static std::unordered_map<Key, CUtensorMap, KeyHash, KeyEq> tmpl_cache;

  Key key{global_height, global_width, shared_height, shared_width};
  auto it = tmpl_cache.find(key);
  if (it == tmpl_cache.end()) {
    CUtensorMap tmp;
    auto err = cuTensorMapEncodeTiled(
      &tmp,
      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 template");
    it = tmpl_cache.emplace(key, tmp).first;
  }

  // Copy template then patch base address.
  *tmap = it->second;
  auto err = cuTensorMapReplaceAddress(tmap, (void*)ptr);
  TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress failed for AB");
}

// Device-side padding/alignment for AB tensors.
//
// TMA with CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE requires all accessed elements to be
// in-bounds, so we pad the leading dimension to 128 (tile height) and zero-fill
// the tail. We cache by source data_ptr() + padded_M to make repeated calls cheap.
static at::Tensor pad_u4_tensor_m128(
    const at::Tensor& src,
    int64_t padded_m
) {
    TORCH_CHECK(src.is_cuda(), "pad_u4_tensor_m128: src must be CUDA");
    TORCH_CHECK(src.numel() > 0, "pad_u4_tensor_m128: empty tensor");
    TORCH_CHECK(padded_m >= src.size(0), "padded_m must be >= src.size(0)");

    const bool needs_pad = (src.size(0) != padded_m);
    const bool needs_align = (((uintptr_t)src.data_ptr() & 0xF) != 0);
    if (!needs_pad && !needs_align) return src;

    auto new_sizes = src.sizes().vec();
    new_sizes[0] = padded_m;
    at::Tensor dst = at::empty(new_sizes, src.options());

    // src is created by the benchmark generator and is contiguous; use a single
    // D2D memcpy and then zero the padded tail.
    const size_t copy_bytes = (size_t)src.nbytes();
    const size_t total_bytes = (size_t)dst.nbytes();
    TORCH_CHECK(copy_bytes <= total_bytes, "pad_u4_tensor_m128: size mismatch");

    CUDA_CHECK(cudaMemcpyAsync(dst.data_ptr(), src.data_ptr(), copy_bytes, cudaMemcpyDeviceToDevice));
    if (total_bytes > copy_bytes) {
        CUDA_CHECK(cudaMemsetAsync((char*)dst.data_ptr() + copy_bytes, 0, total_bytes - copy_bytes));
    }

    return dst;
}

// 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>();
    
    // Set shared memory attributes for all kernel variants (once)
    static bool attrs_set = false;
    if (!attrs_set) {
      auto set_attr = [&](auto kernel, int smem_size) {
          CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
      };
      set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, false>, SMEM_SIZE_MAIN);
      set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, true>,  SMEM_SIZE_MAIN);
      set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, false>, SMEM_SIZE_LOW_M);
      set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, true>,  SMEM_SIZE_LOW_M);
      attrs_set = true;
    }

    std::vector<ProblemInfo> problem_infos(G);
    std::vector<WorkItem> work_items_main;
    std::vector<WorkItem> work_items_low;
    work_items_main.reserve(G * 64);
    work_items_low.reserve(G * 64);
    
    // Track min K for persistence heuristic
    int min_K = INT_MAX;

    std::vector<at::Tensor> keepers;
    keepers.reserve(G * 2);

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

        // TMA requires K to be compatible and AB pointers to be aligned.
        // For partial tiles along M/N, we pad the leading dimension to 128 rows
        // and zero-fill, but keep prob_info.M/N as the *true* sizes for epilogue.
        const int64_t padded_M = ((M + TMA_BLOCK_M - 1) / TMA_BLOCK_M) * TMA_BLOCK_M;
        const int64_t padded_N = ((N + TMA_BLOCK_N - 1) / TMA_BLOCK_N) * TMA_BLOCK_N;
        
        A = pad_u4_tensor_m128(A, padded_M);
        B = pad_u4_tensor_m128(B, padded_N);
        keepers.push_back(A);
        keepers.push_back(B);

        ProblemInfo& prob_info = problem_infos[prob_idx];
        prob_info.M = (int)M; 
        prob_info.N = (int)N; 
        prob_info.K = (int)K;
        min_K = std::min(min_K, (int)K);
        prob_info.Cs0 = C.stride(0);
        prob_info.Cs1 = C.stride(1);
        prob_info.Cs2 = C.stride(2);
        prob_info.C_ptr = (half*)C.data_ptr();

        init_AB_tmap_u4(&prob_info.A_tmap, A.data_ptr(), (uint64_t)A.size(0), (uint64_t)K, TMA_BLOCK_M, TMA_BLOCK_K);
        init_AB_tmap_u4(&prob_info.B_tmap, B.data_ptr(), (uint64_t)B.size(0), (uint64_t)K, TMA_BLOCK_N, TMA_BLOCK_K);
        // The provided SF tensors are a (non-contiguous) view into a contiguous backing
        // storage laid out as [l=1, mn/128, (k/16)/4, 32, 4, 4]. We access the backing
        // storage directly via data_ptr().
        prob_info.SFA_ptr = (const char*)sfa.data_ptr();
        prob_info.SFB_ptr = (const char*)sfb.data_ptr();

        int num_tiles_m = ceil_div((int)M, TMA_BLOCK_M);
        int num_tiles_n = ceil_div((int)N, TMA_BLOCK_N);

        std::vector<WorkItem>& target = (M <= LOW_M_THRESHOLD) ? work_items_low : work_items_main;
        for (int tm = 0; tm < num_tiles_m; tm++) {
            for (int tn = 0; tn < num_tiles_n; tn++) {
                target.push_back({(int)prob_idx, tm, tn});
            }
        }
    }

    if (work_items_main.empty() && work_items_low.empty()) return C_list;

    // Device allocations
    auto options = at::TensorOptions().dtype(at::kByte).device(dev);
    at::Tensor d_probs = at::empty({(int64_t)(G * sizeof(ProblemInfo))}, options);
    
    CUDA_CHECK(cudaMemcpyAsync(d_probs.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));

    // Helper for launching kernels to reduce duplication
    // We capture relevant context (dev, options, min_K) by reference
    auto launch_batch = [&](std::vector<WorkItem>& items, bool is_low_m) {
        if (items.empty()) return;

        int num_items = (int)items.size();
        at::Tensor d_work = at::empty({(int64_t)(num_items * sizeof(WorkItem))}, options);
        CUDA_CHECK(cudaMemcpyAsync(d_work.data_ptr(), items.data(), num_items * sizeof(WorkItem), cudaMemcpyHostToDevice));

        // Persistence heuristic
        constexpr int MAX_PERSISTENT_CTAS = 264; 
        constexpr int MIN_K_FOR_PERSISTENCE = 4096;
        bool use_persistence = (num_items > MAX_PERSISTENT_CTAS) && (min_K >= MIN_K_FOR_PERSISTENCE);

        if (use_persistence) {
            at::Tensor d_counter = at::zeros({1}, options.dtype(at::kInt));
            int* counter_ptr = (int*)d_counter.data_ptr();

            if (is_low_m) {
                grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, true>
                    <<<MAX_PERSISTENT_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_LOW_M>>>(
                    (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, counter_ptr);
            } else {
                grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, true>
                    <<<MAX_PERSISTENT_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_MAIN>>>(
                    (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, counter_ptr);
            }
        } else {
            if (is_low_m) {
                grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, false>
                    <<<num_items, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_LOW_M>>>(
                    (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, nullptr);
            } else {
                grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, false>
                    <<<num_items, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_MAIN>>>(
                    (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, nullptr);
            }
        }
        CUDA_CHECK(cudaGetLastError());
    };

    launch_batch(work_items_main, false);
    launch_batch(work_items_low, true);
    
    // No explicit cleanup needed for host vectors or device tensors (PyTorch handles device tensor lifetime)
    
    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

# _sizes_cpu_cache = {}

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

    A_list = [t[0] for t in abc_tensors]
    B_list = [t[1] for t in abc_tensors]
    C_list = [t[2] for t in abc_tensors]
    
    sfa_list = [t[0] for t in sfasfb_reordered_tensors]
    sfb_list = [t[1] for t in sfasfb_reordered_tensors]

    # A/B padding/alignment is handled inside the C++ host path
    
    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 · 929 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 491660.

⋯ 6 unchanged lines
"""
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
- ⏱ 313 ± 0.4 µs
- ⚡ 310 µs 🐌 331 µs
+ ⏱ 169 ± 0.1 µs
+ ⚡ 169 µs 🐌 170 µ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
- ⏱ 286 ± 0.3 µs
- ⚡ 283 µs 🐌 302 µs
+ ⏱ 159 ± 0.1 µs
+ ⚡ 159 µs 🐌 159 µs
g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
- ⏱ 111 ± 0.4 µs
- ⚡ 108 µs 🐌 132 µs
+ ⏱ 45.6 ± 0.04 µs
+ ⚡ 45.5 µs 🐌 45.7 µs
g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
- ⏱ 63.7 ± 0.10 µs
- ⚡ 61.7 µs 🐌 69.1 µs
+ ⏱ 20.0 ± 0.02 µs
+ ⚡ 19.9 µs 🐌 20.0 µs
"""
CUDA_SRC = """
#include <vector>
+ #include <unordered_map>
#include <cstdint>
#include <cstdio>
#include <cuda.h>
⋯ 31 unchanged lines
struct __align__(128) ProblemInfo {
CUtensorMap A_tmap;
CUtensorMap B_tmap;
- CUtensorMap SFA_tmap;
- CUtensorMap SFB_tmap;
+ const char* SFA_ptr; // points to underlying contiguous storage in [l, mn/128, (k/16)/4, 32, 4, 4]
+ const char* SFB_ptr;
half* C_ptr;
int M, N, K;
int64_t Cs0, Cs1, Cs2;
⋯ 12 unchanged lines
:: "r"(mbar_addr), "r"(size) : "memory");
}
+ __device__ inline void mbarrier_inval(int mbar_addr) {
+ asm volatile("mbarrier.inval.shared::cta.b64 [%0];" :: "r"(mbar_addr) : "memory");
+ }
+
__device__ void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680;
asm volatile(
⋯ 31 unchanged lines
);
}
+ __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)
+ : "memory"
+ );
+ }
+
template <int CTA_GROUP = 1>
__device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
asm volatile(
⋯ 71 unchanged lines
);
}
- // Kernel Configuration
+ // SMEM Layout:
+ // [ProblemInfo] (aligned to 128)
+ // [Stage 0]
+ // [Stage 1]
+ // [Stage 2]
+ // [Mbarriers] (6 mbarriers: 3 for TMA complete, 3 for MMA commit)
+
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 NUM_STAGES_MAIN = 3; // Triple buffering
+ constexpr int NUM_STAGES_LOW_M = 2; // Lower SMEM footprint for small-M tiles
constexpr int TMA_NUM_WARPS = 4;
constexpr int TMEM_COLS = TMA_BLOCK_N * 2;
⋯ 1 unchanged lines
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
+ constexpr int STAGE_SIZE = TMA_A_SMEM_BYTES + TMA_B_SMEM_BYTES + TMA_SFA_SMEM_BYTES + TMA_SFB_SMEM_BYTES; // ~36KB
+ constexpr int MBAR_BYTES_MAIN = ((2 * NUM_STAGES_MAIN * 8 + 63) & ~63);
+ constexpr int MBAR_BYTES_LOW_M = ((2 * NUM_STAGES_LOW_M * 8 + 63) & ~63);
+ constexpr int SMEM_SIZE_MAIN = STAGE_SIZE * NUM_STAGES_MAIN + MBAR_BYTES_MAIN;
+ constexpr int SMEM_SIZE_LOW_M = STAGE_SIZE * NUM_STAGES_LOW_M + MBAR_BYTES_LOW_M;
+ constexpr int LOW_M_THRESHOLD = 96;
+
+ template <bool LOW_M_EPILOGUE>
+ __device__ __forceinline__ void epilogue_store(
+ const ProblemInfo& prob,
+ int m_offset,
+ int n_offset,
+ int tid,
+ int warp_id,
+ int lane_id
+ ) {
+ if (tid >= TMA_BLOCK_M) return;
+
+ const int M = prob.M;
+ const int N = prob.N;
+ half* C_ptr = prob.C_ptr;
+ const int64_t Cs0 = prob.Cs0;
+ const int64_t Cs1 = prob.Cs1;
+
+ const bool full_n = (n_offset + TMA_BLOCK_N <= N);
+ const bool full_m = (m_offset + TMA_BLOCK_M <= M);
+ const bool full_tile = full_n && full_m;
+ const bool contiguous = (Cs1 == 1);
+
+ const int warp_row_base = m_offset + warp_id * 32;
+ if (LOW_M_EPILOGUE && warp_row_base >= M) return;
+
+ int m_iters = 2;
+ if (LOW_M_EPILOGUE) {
+ const int remaining = M - warp_row_base;
+ m_iters = (remaining <= 16) ? 1 : 2;
+ }
+
+ const int lane_row = lane_id >> 2;
+ const int lane_col = (lane_id & 3) * 2;
+
+ for (int m = 0; m < m_iters; ++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");
+
+ const int row0 = warp_row_base + m * 16 + lane_row;
+ const int row1 = row0 + 8;
+
+ if (contiguous) {
+ if (full_tile) {
+ half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);
+ half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset);
+ #pragma unroll
+ for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
+ const int idx = i * 4;
+ const int col = i * 8 + lane_col;
+ const int h2_idx = col >> 1;
+ row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));
+ row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));
+ }
+ continue;
+ }
+
+ const bool row0_in = row0 < M;
+ const bool row1_in = row1 < M;
+ if (full_n) {
+ half2* row0_ptr = row0_in ? reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset) : nullptr;
+ half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset) : nullptr;
+ #pragma unroll
+ for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
+ const int idx = i * 4;
+ const int col = i * 8 + lane_col;
+ const int h2_idx = col >> 1;
+ if (row0_in) {
+ row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));
+ }
+ if (row1_in) {
+ row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));
+ }
+ }
+ } else {
+ #pragma unroll
+ for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
+ const int idx = i * 4;
+ const int col = n_offset + i * 8 + lane_col;
+ if (col < N) {
+ const half h00 = __float2half_rn(tmp[idx + 0]);
+ const half h01 = __float2half_rn(tmp[idx + 1]);
+ const half h10 = __float2half_rn(tmp[idx + 2]);
+ const half h11 = __float2half_rn(tmp[idx + 3]);
+ if (row0_in) {
+ if (col + 1 < N) {
+ reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = __halves2half2(h00, h01);
+ } else {
+ C_ptr[row0 * Cs0 + col] = h00;
+ }
+ }
+ if (row1_in) {
+ if (col + 1 < N) {
+ reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = __halves2half2(h10, h11);
+ } else {
+ C_ptr[row1 * Cs0 + col] = h10;
+ }
+ }
+ }
+ }
+ }
+ } else {
+ const bool row0_in = row0 < M;
+ const bool row1_in = row1 < M;
+ if (full_n) {
+ #pragma unroll
+ for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
+ const int idx = i * 4;
+ const int col = n_offset + i * 8 + lane_col;
+ const half h00 = __float2half_rn(tmp[idx + 0]);
+ const half h01 = __float2half_rn(tmp[idx + 1]);
+ const half h10 = __float2half_rn(tmp[idx + 2]);
+ const half h11 = __float2half_rn(tmp[idx + 3]);
+ if (row0_in) {
+ C_ptr[row0 * Cs0 + col * Cs1] = h00;
+ C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
+ }
+ if (row1_in) {
+ C_ptr[row1 * Cs0 + col * Cs1] = h10;
+ C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
+ }
+ }
+ } else {
+ #pragma unroll
+ for (int i = 0; i < TMA_BLOCK_N / 8; i++) {
+ const int idx = i * 4;
+ const int col = n_offset + i * 8 + lane_col;
+ if (col < N) {
+ const half h00 = __float2half_rn(tmp[idx + 0]);
+ const half h01 = __float2half_rn(tmp[idx + 1]);
+ const half h10 = __float2half_rn(tmp[idx + 2]);
+ const half h11 = __float2half_rn(tmp[idx + 3]);
+ if (row0_in) {
+ C_ptr[row0 * Cs0 + col * Cs1] = h00;
+ if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
+ }
+ if (row1_in) {
+ C_ptr[row1 * Cs0 + col * Cs1] = h10;
+ if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+
// 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,
+ const ProblemInfo* prob,
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);
+ tma_3d_gmem2smem<1>(A_smem, &prob->A_tmap, 0, m_offset, off_k / 256, mbar_addr, EVICT_NORMAL);
+ tma_3d_gmem2smem<1>(B_smem, &prob->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);
+ // Scale factors live in the underlying contiguous storage order:
+ // [l=1, mn/128, (k/16)/4, 32, 4, 4], with each (32,4,4) tile = 512 bytes.
+ // For BLOCK_K=256 we need 4 consecutive 512B tiles per k_iter (total 2048B).
+ const int rest_k = prob->K / 64; // (K/16)/4
+ const int k_blk = off_k / 64; // (off_k/16)/4
+ const char* SFA_src = prob->SFA_ptr + (int64_t)(m_tile_idx * rest_k + k_blk) * 512;
+ const char* SFB_src = prob->SFB_ptr + (int64_t)(n_tile_idx * rest_k + k_blk) * 512;
+ tma_gmem2smem(SFA_smem, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, EVICT_NORMAL);
+ tma_gmem2smem(SFB_smem, SFB_src, TMA_SFB_SMEM_BYTES, 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);
- }
-
- // Single Persistent Kernel
+ // Single Kernel (template for main/low-M variants, with optional persistence)
+ template <int NUM_STAGES, bool LOW_M_EPILOGUE, bool PERSISTENT>
__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)
void grouped_gemm_tcgen_tma_v3_persistent(
const ProblemInfo* __restrict__ global_probs,
const WorkItem* __restrict__ work_items,
- int num_items
+ int num_items,
+ int* __restrict__ work_counter // Only used when PERSISTENT=true
) {
- const int global_idx = blockIdx.x;
- if (global_idx >= num_items) return;
-
- const WorkItem& work = work_items[global_idx];
- const ProblemInfo& prob = global_probs[work.problem_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
+ // Shared Memory Setup (constant addresses)
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};
+ const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
- // 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");
+ // Pipeline Buffers
+ const int stages_base = smem_base;
+ int A_smem[NUM_STAGES];
+ #pragma unroll
+ for (int i = 0; i < NUM_STAGES; ++i) {
+ A_smem[i] = stages_base + STAGE_SIZE * i;
}
- __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));
+ const int B_off = TMA_A_SMEM_BYTES;
+ const int SFA_off = B_off + TMA_B_SMEM_BYTES;
+ const int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;
+
+ // Mbarriers (constant addresses)
+ const int mbar_base = stages_base + STAGE_SIZE * NUM_STAGES;
+ int tma_mbar[NUM_STAGES];
+ int mma_mbar[NUM_STAGES];
+ #pragma unroll
+ for (int i = 0; i < NUM_STAGES; ++i) {
+ tma_mbar[i] = mbar_base + i * 8;
+ mma_mbar[i] = mbar_base + (NUM_STAGES + i) * 8;
}
- __syncthreads();
+ // TMEM addresses (constant)
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],
- &prob.A_tmap, &prob.B_tmap, &prob.SFA_tmap, &prob.SFB_tmap,
- m_offset, n_offset, 0,
- m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);
+ // Allocate TMEM ONCE
+ if (warp_id == 0) {
+ asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(stages_base), "r"(TMEM_COLS));
}
__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
+
+ // Shared work index (only used in persistent mode)
+ __shared__ int shared_work_idx;
+
+ // ================================================================
+ // WORK DISPATCH: Persistent loop vs single-tile
+ // ================================================================
+ int work_idx;
+ if constexpr (PERSISTENT) {
+ // Fetch work atomically
if (tid == 0) {
- mbarrier_wait(tma_mbar[stage], k_iter / NUM_STAGES);
+ shared_work_idx = atomicAdd(work_counter, 1);
}
__syncthreads();
+ work_idx = shared_work_idx;
+ } else {
+ // Non-persistent: each CTA handles one tile via blockIdx
+ work_idx = blockIdx.x;
+ }
+
+ // Main processing loop (runs once for non-persistent, loops for persistent)
+ while (work_idx < num_items) {
+ const WorkItem& work = work_items[work_idx];
+ const ProblemInfo& prob = global_probs[work.problem_idx];
+
+ const int m_offset = work.tile_m * TMA_BLOCK_M;
+ const int n_offset = work.tile_n * TMA_BLOCK_N;
+ const int K = prob.K;
+ const int num_k_iters = (K + TMA_BLOCK_K - 1) / TMA_BLOCK_K;
+
+ // 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;
+
+ // Initialize or reinit mbarriers
+ // Note: mbarrier_inval is not strictly needed if we ensure all threads synced and state is clean via init
- // 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],
- &prob.A_tmap, &prob.B_tmap, &prob.SFA_tmap, &prob.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);
+ for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(tma_mbar[i], 1);
+ for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(mma_mbar[i], 1);
+ asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
__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);
+ // ----------------------------------------------------------------
+ // PRODUCER WARP (Warp 0): Issues TMA
+ // ----------------------------------------------------------------
+ if (warp_id == 0) {
+ if (lane_id == 0) {
+ for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {
+ const int stage = k_iter % NUM_STAGES;
+ // Wait for buffer to be free (consumed by MMA)
+ // Ideally, mma_mbar tracks "buffer consumed".
+ // Init count 1.
+ // Logic:
+ // k=0: buffer is free (init state).
+ // But we need to sync with consumer?
+ // Let's assume mma_mbar is signaled when MMA is done using the buffer.
+ // Initial state: buffers are free. mma_mbar should NOT block for first use.
+ // But mbarrier logic: wait() blocks until phase flips.
+ // We need to manage phases carefully.
+
+ // Correct logic:
+ // TMA thread waits for buffer to be available.
+ // For k < NUM_STAGES, buffers are initially available.
+ // For k >= NUM_STAGES, wait for previous usage to complete.
+
+ if (k_iter >= NUM_STAGES) {
+ // Wait for stage to be released by MMA
+ // Corresponding k was k_iter - NUM_STAGES
+ mbarrier_wait(mma_mbar[stage], (k_iter - NUM_STAGES) / NUM_STAGES); // Wait for phase flip?
+ // Or just use the same phase logic as coupled.
+ // Coupled used: mbarrier_wait(mma_mbar[next_stage], ...)
+ }
+
+ // Issue TMA
+ const int stage_base = A_smem[stage];
+ issue_tma_loads(stage_base, stage_base + B_off, stage_base + SFA_off, stage_base + SFB_off,
+ tma_mbar[stage], &prob, m_offset, n_offset, k_iter,
+ m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);
+ }
+ }
}
- __syncthreads();
-
- asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
- // Epilogue
- half* C_ptr = prob.C_ptr;
- int64_t Cs0 = prob.Cs0;
- int64_t Cs1 = prob.Cs1;
- // Cs2 is usually 1, but we should use it if needed or assume packed?
- // The store logic below assumes normal strided layout.
+ // ----------------------------------------------------------------
+ // CONSUMER WARP (Warp 1): Issues MMA
+ // ----------------------------------------------------------------
+ else if (warp_id == 1) {
+ if (lane_id == 0) {
+ for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {
+ const int stage = k_iter % NUM_STAGES;
+
+ // Wait for data ready (TMA complete)
+ mbarrier_wait(tma_mbar[stage], k_iter / NUM_STAGES);
+
+ const int stage_base = A_smem[stage];
+
+ // Issue MMA
+ constexpr uint64_t SF_desc = (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
+ const uint64_t SFA_desc = SF_desc | ((uint64_t)(stage_base + SFA_off) >> 4ULL);
+ const uint64_t SFB_desc = SF_desc | ((uint64_t)(stage_base + SFB_off) >> 4ULL);
- 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");
+ constexpr uint64_t AB_desc = (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
+ const uint64_t A_desc = AB_desc | ((uint64_t)stage_base >> 4ULL);
+ const uint64_t B_desc = AB_desc | ((uint64_t)(stage_base + B_off) >> 4ULL);
- 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;
+ #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));
+ }
- 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 (Cs1 == 1) {
- if (col + 1 < N) {
- reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = __halves2half2(h00, h01);
- } else {
- C_ptr[row0 * Cs0 + col] = h00;
- }
- } else {
- C_ptr[row0 * Cs0 + col * Cs1] = h00;
- if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
+ #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 = (k_iter == 0 && 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);
+ }
+
+ // Commit MMA -> Signals mma_mbar[stage]
+ // When commit reaches mma_mbar, it allows TMA producer to reuse buffer for next phase
+ tcgen05_commit<1>(mma_mbar[stage]);
}
- }
-
- if (row1 < M && col < N) {
- if (Cs1 == 1) {
- if (col + 1 < N) {
- reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = __halves2half2(h10, h11);
- } else {
- C_ptr[row1 * Cs0 + col] = h10;
- }
- } else {
- C_ptr[row1 * Cs0 + col * Cs1] = h10;
- if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
- }
- }
+
+ // Final Wait
+ // Wait for last commit to ensure all instructions retired?
+ // Actually, we must ensure all work is done before exiting.
+ // Wait for the last commit to finish
+ if (num_k_iters > 0) {
+ const int last_stage = (num_k_iters - 1) % NUM_STAGES;
+ mbarrier_wait(mma_mbar[last_stage], (num_k_iters - 1) / NUM_STAGES);
+ }
}
- }
}
+
+ // ----------------------------------------------------------------
+ // SYNCHRONIZATION
+ // ----------------------------------------------------------------
+ __syncthreads(); // Wait for all warps to finish
+
+ // Ensure MMA results visible
+ asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
+ // Epilogue
+ epilogue_store<LOW_M_EPILOGUE>(prob, m_offset, n_offset, tid, warp_id, lane_id);
+
__syncthreads();
+
+ // Loop control: continue (persistent) or break (non-persistent)
+ if constexpr (PERSISTENT) {
+ // Fetch next work
+ if (tid == 0) {
+ shared_work_idx = atomicAdd(work_counter, 1);
+ }
+ __syncthreads();
+ work_idx = shared_work_idx;
+ } else {
+ // Non-persistent: exit after single tile
+ break;
+ }
+ } // End main loop
+
+ // Deallocate TMEM once at end
if (warp_id == 0) {
tcgen05_dealloc_cols_cta1(tmem_base, TMEM_COLS);
}
}
+
// Tensor Map Initialization
void init_AB_tmap_u4(
CUtensorMap *tmap,
⋯ 11 unchanged lines
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");
+ // cuTensorMapEncodeTiled is relatively expensive on the host; amortize by caching
+ // a per-shape template and then patching only the base address each call.
+ struct Key {
+ uint64_t gh, gw;
+ uint32_t sh, sw;
+ };
+ struct KeyHash {
+ size_t operator()(const Key& k) const noexcept {
+ uint64_t h = k.gh;
+ h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
+ h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
+ return (size_t)h;
+ }
+ };
+ struct KeyEq {
+ bool operator()(const Key& a, const Key& b) const noexcept {
+ return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw;
+ }
+ };
+ static std::unordered_map<Key, CUtensorMap, KeyHash, KeyEq> tmpl_cache;
+
+ Key key{global_height, global_width, shared_height, shared_width};
+ auto it = tmpl_cache.find(key);
+ if (it == tmpl_cache.end()) {
+ CUtensorMap tmp;
+ auto err = cuTensorMapEncodeTiled(
+ &tmp,
+ 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 template");
+ it = tmpl_cache.emplace(key, tmp).first;
+ }
+
+ // Copy template then patch base address.
+ *tmap = it->second;
+ auto err = cuTensorMapReplaceAddress(tmap, (void*)ptr);
+ TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress 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};
+ // Device-side padding/alignment for AB tensors.
+ //
+ // TMA with CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE requires all accessed elements to be
+ // in-bounds, so we pad the leading dimension to 128 (tile height) and zero-fill
+ // the tail. We cache by source data_ptr() + padded_M to make repeated calls cheap.
+ static at::Tensor pad_u4_tensor_m128(
+ const at::Tensor& src,
+ int64_t padded_m
+ ) {
+ TORCH_CHECK(src.is_cuda(), "pad_u4_tensor_m128: src must be CUDA");
+ TORCH_CHECK(src.numel() > 0, "pad_u4_tensor_m128: empty tensor");
+ TORCH_CHECK(padded_m >= src.size(0), "padded_m must be >= src.size(0)");
- 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");
+ const bool needs_pad = (src.size(0) != padded_m);
+ const bool needs_align = (((uintptr_t)src.data_ptr() & 0xF) != 0);
+ if (!needs_pad && !needs_align) return src;
+
+ auto new_sizes = src.sizes().vec();
+ new_sizes[0] = padded_m;
+ at::Tensor dst = at::empty(new_sizes, src.options());
+
+ // src is created by the benchmark generator and is contiguous; use a single
+ // D2D memcpy and then zero the padded tail.
+ const size_t copy_bytes = (size_t)src.nbytes();
+ const size_t total_bytes = (size_t)dst.nbytes();
+ TORCH_CHECK(copy_bytes <= total_bytes, "pad_u4_tensor_m128: size mismatch");
+
+ CUDA_CHECK(cudaMemcpyAsync(dst.data_ptr(), src.data_ptr(), copy_bytes, cudaMemcpyDeviceToDevice));
+ if (total_bytes > copy_bytes) {
+ CUDA_CHECK(cudaMemsetAsync((char*)dst.data_ptr() + copy_bytes, 0, total_bytes - copy_bytes));
+ }
+
+ return dst;
}
// Host Entry Point
⋯ 16 unchanged lines
auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();
- CUDA_CHECK(cudaFuncSetAttribute(
- grouped_gemm_tcgen_tma_v3_persistent,
- cudaFuncAttributeMaxDynamicSharedMemorySize,
- SMEM_SIZE
- ));
+ // Set shared memory attributes for all kernel variants (once)
+ static bool attrs_set = false;
+ if (!attrs_set) {
+ auto set_attr = [&](auto kernel, int smem_size) {
+ CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
+ };
+ set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, false>, SMEM_SIZE_MAIN);
+ set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, true>, SMEM_SIZE_MAIN);
+ set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, false>, SMEM_SIZE_LOW_M);
+ set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, true>, SMEM_SIZE_LOW_M);
+ attrs_set = true;
+ }
std::vector<ProblemInfo> problem_infos(G);
- std::vector<WorkItem> all_work_items;
- all_work_items.reserve(G * 64); // heuristic reserve
+ std::vector<WorkItem> work_items_main;
+ std::vector<WorkItem> work_items_low;
+ work_items_main.reserve(G * 64);
+ work_items_low.reserve(G * 64);
+
+ // Track min K for persistence heuristic
+ int min_K = INT_MAX;
+ std::vector<at::Tensor> keepers;
+ keepers.reserve(G * 2);
+
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];
⋯ 8 unchanged lines
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);
+ // TMA requires K to be compatible and AB pointers to be aligned.
+ // For partial tiles along M/N, we pad the leading dimension to 128 rows
+ // and zero-fill, but keep prob_info.M/N as the *true* sizes for epilogue.
+ const int64_t padded_M = ((M + TMA_BLOCK_M - 1) / TMA_BLOCK_M) * TMA_BLOCK_M;
+ const int64_t padded_N = ((N + TMA_BLOCK_N - 1) / TMA_BLOCK_N) * TMA_BLOCK_N;
+
+ A = pad_u4_tensor_m128(A, padded_M);
+ B = pad_u4_tensor_m128(B, padded_N);
+ keepers.push_back(A);
+ keepers.push_back(B);
+
ProblemInfo& prob_info = problem_infos[prob_idx];
prob_info.M = (int)M;
prob_info.N = (int)N;
prob_info.K = (int)K;
+ min_K = std::min(min_K, (int)K);
prob_info.Cs0 = C.stride(0);
prob_info.Cs1 = C.stride(1);
prob_info.Cs2 = C.stride(2);
prob_info.C_ptr = (half*)C.data_ptr();
- init_AB_tmap_u4(&prob_info.A_tmap, A.data_ptr(), A.size(0), K, TMA_BLOCK_M, TMA_BLOCK_K);
- init_AB_tmap_u4(&prob_info.B_tmap, B.data_ptr(), B.size(0), K, TMA_BLOCK_N, TMA_BLOCK_K);
- init_SF_tmap_reordered(&prob_info.SFA_tmap, sfa.data_ptr(), sfa.numel() * sfa.element_size(), TMA_SFA_SMEM_BYTES);
- init_SF_tmap_reordered(&prob_info.SFB_tmap, sfb.data_ptr(), sfb.numel() * sfb.element_size(), TMA_SFB_SMEM_BYTES);
+ init_AB_tmap_u4(&prob_info.A_tmap, A.data_ptr(), (uint64_t)A.size(0), (uint64_t)K, TMA_BLOCK_M, TMA_BLOCK_K);
+ init_AB_tmap_u4(&prob_info.B_tmap, B.data_ptr(), (uint64_t)B.size(0), (uint64_t)K, TMA_BLOCK_N, TMA_BLOCK_K);
+ // The provided SF tensors are a (non-contiguous) view into a contiguous backing
+ // storage laid out as [l=1, mn/128, (k/16)/4, 32, 4, 4]. We access the backing
+ // storage directly via data_ptr().
+ prob_info.SFA_ptr = (const char*)sfa.data_ptr();
+ prob_info.SFB_ptr = (const char*)sfb.data_ptr();
int num_tiles_m = ceil_div((int)M, TMA_BLOCK_M);
int num_tiles_n = ceil_div((int)N, TMA_BLOCK_N);
-
+
+ std::vector<WorkItem>& target = (M <= LOW_M_THRESHOLD) ? work_items_low : work_items_main;
for (int tm = 0; tm < num_tiles_m; tm++) {
for (int tn = 0; tn < num_tiles_n; tn++) {
- all_work_items.push_back({(int)prob_idx, tm, tn});
+ target.push_back({(int)prob_idx, tm, tn});
}
}
}
- if (all_work_items.empty()) return C_list;
+ if (work_items_main.empty() && work_items_low.empty()) return C_list;
- ProblemInfo* d_problem_infos;
- WorkItem* d_work_items;
+ // Device allocations
+ auto options = at::TensorOptions().dtype(at::kByte).device(dev);
+ at::Tensor d_probs = at::empty({(int64_t)(G * sizeof(ProblemInfo))}, options);
- CUDA_CHECK(cudaMalloc(&d_problem_infos, G * sizeof(ProblemInfo)));
- CUDA_CHECK(cudaMemcpy(d_problem_infos, problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));
-
- CUDA_CHECK(cudaMalloc(&d_work_items, all_work_items.size() * sizeof(WorkItem)));
- CUDA_CHECK(cudaMemcpy(d_work_items, all_work_items.data(), all_work_items.size() * sizeof(WorkItem), cudaMemcpyHostToDevice));
+ CUDA_CHECK(cudaMemcpyAsync(d_probs.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));
- int num_items = (int)all_work_items.size();
- grouped_gemm_tcgen_tma_v3_persistent<<<num_items, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE>>>(
- d_problem_infos, d_work_items, num_items
- );
+ // Helper for launching kernels to reduce duplication
+ // We capture relevant context (dev, options, min_K) by reference
+ auto launch_batch = [&](std::vector<WorkItem>& items, bool is_low_m) {
+ if (items.empty()) return;
+
+ int num_items = (int)items.size();
+ at::Tensor d_work = at::empty({(int64_t)(num_items * sizeof(WorkItem))}, options);
+ CUDA_CHECK(cudaMemcpyAsync(d_work.data_ptr(), items.data(), num_items * sizeof(WorkItem), cudaMemcpyHostToDevice));
+
+ // Persistence heuristic
+ constexpr int MAX_PERSISTENT_CTAS = 264;
+ constexpr int MIN_K_FOR_PERSISTENCE = 4096;
+ bool use_persistence = (num_items > MAX_PERSISTENT_CTAS) && (min_K >= MIN_K_FOR_PERSISTENCE);
+
+ if (use_persistence) {
+ at::Tensor d_counter = at::zeros({1}, options.dtype(at::kInt));
+ int* counter_ptr = (int*)d_counter.data_ptr();
+
+ if (is_low_m) {
+ grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, true>
+ <<<MAX_PERSISTENT_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_LOW_M>>>(
+ (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, counter_ptr);
+ } else {
+ grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, true>
+ <<<MAX_PERSISTENT_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_MAIN>>>(
+ (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, counter_ptr);
+ }
+ } else {
+ if (is_low_m) {
+ grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, false>
+ <<<num_items, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_LOW_M>>>(
+ (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, nullptr);
+ } else {
+ grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, false>
+ <<<num_items, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_MAIN>>>(
+ (ProblemInfo*)d_probs.data_ptr(), (WorkItem*)d_work.data_ptr(), num_items, nullptr);
+ }
+ }
+ CUDA_CHECK(cudaGetLastError());
+ };
+
+ launch_batch(work_items_main, false);
+ launch_batch(work_items_low, true);
- // We remove the per-problem sync. We still need to manage the lifetime of d_problem_infos and d_work_items.
- // Ideally we would use async free or cached allocator, but for simplicity/safety we sync once at the end.
- CUDA_CHECK(cudaDeviceSynchronize());
- CUDA_CHECK(cudaFree(d_problem_infos));
- CUDA_CHECK(cudaFree(d_work_items));
+ // No explicit cleanup needed for host vectors or device tensors (PyTorch handles device tensor lifetime)
CUDA_CHECK(cudaGetLastError());
return C_list;
⋯ 25 unchanged lines
)
group_gemm = torch.ops.my_module.group_gemm
+ # _sizes_cpu_cache = {}
+
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 = []
+ A_list = [t[0] for t in abc_tensors]
+ B_list = [t[1] for t in abc_tensors]
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)
+ sfa_list = [t[0] for t in sfasfb_reordered_tensors]
+ sfb_list = [t[1] for t in sfasfb_reordered_tensors]
- 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)
+ # A/B padding/alignment is handled inside the C++ host path
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)
scrolls · 1028 diff lines total

Best evidence level for this revision: reported

JSON