Skip to content
KernelIndex
Search⌘K

submission 276247

binsquare · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
15.8µs
#115 of 420
2026-01-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:37f67c7ba16a6c64d4cd27a1fbe6eaf267667f6a4a183e62b2375042f59d1daf
license declaredunknown
license concludedunknown
authorsbinsquare
imported2026-08-26

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) {
shared-memory__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint16_t ctamask, uint64_t cache_policy) {
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tmaasm 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));

Kernel source

submission.py538 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
from utils import make_match_reference

sf_vec_size = 16

def ceil_div(a, b):
    return (a + b - 1) // b

def to_blocked(input_matrix):
    rows, cols = input_matrix.shape
    n_row_blocks = (rows + 127) // 128
    n_col_blocks = (cols + 3) // 4
    blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3).contiguous()
    rearranged = blocks.view(-1, 4, 32, 4).transpose(1, 2).contiguous().view(-1, 32, 16)
    return rearranged.flatten()

def ref_kernel(data: input_t) -> output_t:
    a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
    m, n, l = c_ref.shape
    ref1 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
    ref2 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
    for l_idx in range(l):
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
        scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])
        res1 = torch._scaled_mm(a_ref[:, :, l_idx], b1_ref[:, :, l_idx].transpose(0, 1),
                                scale_a.cuda(), scale_b1.cuda(), bias=None, out_dtype=torch.float32)
        ref1[:, :, l_idx] = res1
        res2 = torch._scaled_mm(a_ref[:, :, l_idx], b2_ref[:, :, l_idx].transpose(0, 1),
                                scale_a.cuda(), scale_b2.cuda(), bias=None, out_dtype=torch.float32)
        ref2[:, :, l_idx] = res2
    c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
    return c_ref

def generate_input(m: int, n: int, k: int, l: int, seed: int):
    torch.manual_seed(seed)
    def create_fp4_tensors(l, mn, k):
        ref_i8 = torch.randint(255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda")
        ref_i8 = ref_i8 & 0b1011_1011
        return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
    a_ref = create_fp4_tensors(l, m, k).view(torch.float4_e2m1fn_x2)
    b1_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
    b2_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
    c_ref = torch.randn((l, m, n), dtype=torch.float16, device="cuda").permute(1, 2, 0)
    def create_scale_factor_tensors(l, mn, sf_k):
        ref_shape = (l, mn, sf_k)
        ref_f8_random_fp32 = torch.rand(ref_shape, dtype=torch.float32, device='cuda')
        ref_f8_torch_tensor = ref_f8_random_fp32.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
        atom_m, atom_k = (32, 4), 4
        mma_shape = (l, ceil_div(mn, atom_m[0] * atom_m[1]), ceil_div(sf_k, atom_k), atom_m[0], atom_m[1], atom_k)
        rand_int_tensor = torch.empty(mma_shape, dtype=torch.int8, device='cuda')
        reordered = rand_int_tensor.to(dtype=torch.float8_e4m3fn).permute(3, 4, 1, 5, 2, 0)
        i_idx, j_idx, b_idx = torch.arange(mn, device='cuda'), torch.arange(sf_k, device='cuda'), torch.arange(l, device='cuda')
        i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing='ij')
        mm, mm32, mm4 = i_grid // (atom_m[0] * atom_m[1]), i_grid % atom_m[0], (i_grid % 128) // atom_m[0]
        kk, kk4 = j_grid // atom_k, j_grid % atom_k
        reordered[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_torch_tensor[i_grid, j_grid, b_grid]
        return ref_f8_torch_tensor.cpu(), reordered
    sf_k = ceil_div(k, sf_vec_size)
    sfa_ref_cpu, sfa_perm = create_scale_factor_tensors(l, m, sf_k)
    sfb1_ref_cpu, sfb1_perm = create_scale_factor_tensors(l, n, sf_k)
    sfb2_ref_cpu, sfb2_perm = create_scale_factor_tensors(l, n, sf_k)
    return (a_ref, b1_ref, b2_ref, sfa_ref_cpu.to("cuda"), sfb1_ref_cpu.to("cuda"), sfb2_ref_cpu.to("cuda"), sfa_perm, sfb1_perm, sfb2_perm, c_ref)

check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)

# Optimized dual GEMM kernel with improved memory access
CUDA_SRC = r"""
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;

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

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

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

__device__ void mbarrier_wait(int mbar_addr, int phase) {
  uint32_t ticks = 0x989680;
  asm volatile("{\n\t.reg .pred P1;\n\tLAB_WAIT:\n\tmbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t@P1 bra.uni DONE;\n\tbra.uni LAB_WAIT;\n\tDONE:\n\t}" :: "r"(mbar_addr), "r"(phase), "r"(ticks));
}

__device__ inline void mbarrier_inval(int mbar_addr) {
  asm volatile("mbarrier.inval.shared::cta.b64 [%0];" :: "r"(mbar_addr));
}

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

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

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

// Get CTA rank within cluster
__device__ inline uint32_t get_cta_rank_in_cluster() {
  uint32_t rank;
  asm volatile("mov.u32 %0, %%clusterid;" : "=r"(rank));
  return rank;
}

// Cluster-wide barrier arrive
__device__ inline void cluster_arrive() {
  asm volatile("barrier.cluster.arrive;");
}

// Cluster-wide barrier wait
__device__ inline void cluster_wait() {
  asm volatile("barrier.cluster.wait;");
}

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

__device__ inline void tcgen05_mma_nvfp4(uint64_t a_desc, uint64_t b_desc, uint32_t i_desc, int scale_A, int scale_B, int enable_d, int d_tmem = 0) {
  asm volatile("{\n\t.reg .pred p;\n\tsetp.ne.b32 p, %6, 0;\n\ttcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t}" :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc), "r"(scale_A), "r"(scale_B), "r"(enable_d));
}

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

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

void check_cu(CUresult err) { if (err != CUDA_SUCCESS) { const char *msg; cuGetErrorString(err, &msg); TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", msg); } }

void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t gh, uint64_t gw, uint32_t sh, uint32_t sw) {
  uint64_t globalDim[3] = {256, gh, gw / 256};
  uint64_t globalStrides[2] = {gw / 2, 128};
  uint32_t boxDim[3] = {256, sh, sw / 256};
  uint32_t elementStrides[3] = {1, 1, 1};
  check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B, 3, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides, CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}

// Non-persistent dual GEMM kernel - optimized version
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_kernel(
    const __grid_constant__ CUtensorMap A_tmap,
    const __grid_constant__ CUtensorMap B1_tmap,
    const __grid_constant__ CUtensorMap B2_tmap,
    const char *SFA_ptr, const char *SFB1_ptr, const char *SFB2_ptr,
    half *C_ptr, int M, int N, int grid_n) {

  const int tid = threadIdx.x, bid_m = blockIdx.x, bid_n = blockIdx.y;
  const int lane_id = tid % WARP_SIZE, warp_id = tid / WARP_SIZE;
  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
  const int off_m = bid_m * BLOCK_M, off_n = bid_n * BLOCK_N;

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

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

  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));
  const int mma_mbar = tma_mbar + NUM_STAGES * 8, mainloop_mbar = mma_mbar + NUM_STAGES * 8;

  constexpr int ACC1 = 0, ACC2 = BLOCK_N, SFA_tm = 2 * BLOCK_N;
  constexpr int SFB1_tm = SFA_tm + 4 * (BLOCK_K / MMA_K), SFB2_tm = SFB1_tm + 4 * (BLOCK_K / MMA_K);
  const int num_iters = K / BLOCK_K;
  const int sfb_n_offset = (BLOCK_N == 64) ? (bid_n % 2) * 2 : 0;

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

  // TMA warp
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
    auto issue_tma = [&](int iter_k, int stage) {
      const int mbar = tma_mbar + stage * 8;
      const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
      const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem(A_sm, &A_tmap, 0, off_m, off_k / 256, mbar, cache_A);
      tma_3d_gmem2smem(B1_sm, &B1_tmap, 0, off_n, off_k / 256, mbar, cache_B);
      tma_3d_gmem2smem(B2_sm, &B2_tmap, 0, off_n, off_k / 256, mbar, cache_B);
      const int rest_k = K / 16 / 4;
      tma_gmem2smem(SFA_sm, SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512, SFA_size, mbar, cache_A);
      tma_gmem2smem(SFB1_sm, SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
      tma_gmem2smem(SFB2_sm, SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
      asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(mbar), "r"(STAGE_SIZE) : "memory");
    };
    for (int k = 0; k < NUM_STAGES; k++) issue_tma(k, k);
    for (int k = NUM_STAGES; k < num_iters; k++) {
      int stage = k % NUM_STAGES;
      mbarrier_wait(mma_mbar + stage * 8, (k / NUM_STAGES - 1) % 2);
      issue_tma(k, stage);
    }
  }
  // MMA warp
  else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (1U << 27U);
    for (int iter_k = 0; iter_k < num_iters; iter_k++) {
      int stage = iter_k % NUM_STAGES;
      mbarrier_wait(tma_mbar + stage * 8, (iter_k / NUM_STAGES) % 2);
      const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
      const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
      auto make_AB = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL); };
      auto make_SF = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL); };
      uint64_t SFA_d = make_SF(0) + ((uint64_t)SFA_sm >> 4ULL);
      uint64_t SFB1_d = make_SF(0) + ((uint64_t)SFB1_sm >> 4ULL);
      uint64_t SFB2_d = make_SF(0) + ((uint64_t)SFB2_sm >> 4ULL);
      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        tcgen05_cp_nvfp4(SFA_tm + k * 4, SFA_d + (uint64_t)k * 32ULL);
        tcgen05_cp_nvfp4(SFB1_tm + k * 4, SFB1_d + (uint64_t)k * 32ULL);
        tcgen05_cp_nvfp4(SFB2_tm + k * 4, SFB2_d + (uint64_t)k * 32ULL);
      }
      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        uint64_t a_d = make_AB(A_sm + k * 32);
        uint64_t b1_d = make_AB(B1_sm + k * 32);
        uint64_t b2_d = make_AB(B2_sm + k * 32);
        int sA = SFA_tm + k * 4;
        int sB1 = SFB1_tm + k * 4 + sfb_n_offset;
        int sB2 = SFB2_tm + k * 4 + sfb_n_offset;
        int en = (k == 0) ? iter_k : 1;
        tcgen05_mma_nvfp4(a_d, b1_d, i_desc, sA, sB1, en, ACC1);
        tcgen05_mma_nvfp4(a_d, b2_d, i_desc, sA, sB2, en, ACC2);
      }
      asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mma_mbar + stage * 8) : "memory");
    }
    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mainloop_mbar) : "memory");
  }
  // Epilogue warps - 4 warps for better parallelism
  else if (tid < BLOCK_M) {
    mbarrier_wait(mainloop_mbar, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");
    const int lane_row = lane_id / 4, lane_col = (lane_id % 4) * 2;
    half* __restrict__ C_base = C_ptr + off_m * N + off_n + lane_col;
    constexpr int CHUNK = 64;
    constexpr int NUM_CHUNKS = BLOCK_N / CHUNK;
    #pragma unroll
    for (int m = 0; m < 2; m++) {
      const int row_base = warp_id * 32 + m * 16;
      half* __restrict__ C_row0 = C_base + (row_base + lane_row) * N;
      half* __restrict__ C_row8 = C_row0 + 8 * N;
      #pragma unroll
      for (int c = 0; c < NUM_CHUNKS; c++) {
        float tmp1[32], tmp2[32];
        tcgen05_ld_16x256bx8(tmp1, row_base, ACC1 + c * CHUNK);
        tcgen05_ld_16x256bx8(tmp2, row_base, ACC2 + c * CHUNK);
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        #pragma unroll
        for (int i = 0; i < 8; i++) {
          float x0 = tmp1[i*4], x1 = tmp1[i*4+1], x2 = tmp1[i*4+2], x3 = tmp1[i*4+3];
          float y0 = tmp2[i*4], y1 = tmp2[i*4+1], y2 = tmp2[i*4+2], y3 = tmp2[i*4+3];
          float e0 = __expf(-x0), e1 = __expf(-x1), e2 = __expf(-x2), e3 = __expf(-x3);
          float s0 = x0 / (1.f + e0) * y0, s1 = x1 / (1.f + e1) * y1;
          float s2 = x2 / (1.f + e2) * y2, s3 = x3 / (1.f + e3) * y3;
          *reinterpret_cast<half2*>(C_row0 + c * CHUNK + i * 8) = __float22half2_rn({s0, s1});
          *reinterpret_cast<half2*>(C_row8 + c * CHUNK + i * 8) = __float22half2_rn({s2, s3});
        }
      }
    }
  }

  // Cleanup
  if (warp_id == 0 && tid == 0) {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
  }
}

// Cluster-aware kernel - each CTA in cluster processes different N tile
// Cluster shape: (1, CLUSTER_N, 1) - CTAs along N dimension
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_N>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_cluster_kernel(
    const __grid_constant__ CUtensorMap A_tmap,
    const __grid_constant__ CUtensorMap B1_tmap,
    const __grid_constant__ CUtensorMap B2_tmap,
    const char *SFA_ptr, const char *SFB1_ptr, const char *SFB2_ptr,
    half *C_ptr, int M, int N, int grid_n) {

  const int tid = threadIdx.x;
  const int lane_id = tid % WARP_SIZE, warp_id = tid / WARP_SIZE;
  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

  // Use standard block indices - no cluster complexity for now
  // Just verify the kernel works with standard launch
  const int bid_m = blockIdx.x;
  const int bid_n = blockIdx.y;
  const int off_m = bid_m * BLOCK_M, off_n = bid_n * BLOCK_N;

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

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

  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));
  const int mma_mbar = tma_mbar + NUM_STAGES * 8, mainloop_mbar = mma_mbar + NUM_STAGES * 8;

  constexpr int ACC1 = 0, ACC2 = BLOCK_N, SFA_tm = 2 * BLOCK_N;
  constexpr int SFB1_tm = SFA_tm + 4 * (BLOCK_K / MMA_K), SFB2_tm = SFB1_tm + 4 * (BLOCK_K / MMA_K);
  const int num_iters = K / BLOCK_K;
  const int sfb_n_offset = (BLOCK_N == 64) ? (bid_n % 2) * 2 : 0;

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

  // TMA warp
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
    auto issue_tma = [&](int iter_k, int stage) {
      const int mbar = tma_mbar + stage * 8;
      const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
      const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
      const int off_k = iter_k * BLOCK_K;

      tma_3d_gmem2smem(A_sm, &A_tmap, 0, off_m, off_k / 256, mbar, cache_A);
      tma_3d_gmem2smem(B1_sm, &B1_tmap, 0, off_n, off_k / 256, mbar, cache_B);
      tma_3d_gmem2smem(B2_sm, &B2_tmap, 0, off_n, off_k / 256, mbar, cache_B);
      const int rest_k = K / 16 / 4;
      tma_gmem2smem(SFA_sm, SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512, SFA_size, mbar, cache_A);
      tma_gmem2smem(SFB1_sm, SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
      tma_gmem2smem(SFB2_sm, SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
      asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(mbar), "r"(STAGE_SIZE) : "memory");
    };
    for (int k = 0; k < NUM_STAGES; k++) issue_tma(k, k);
    for (int k = NUM_STAGES; k < num_iters; k++) {
      int stage = k % NUM_STAGES;
      mbarrier_wait(mma_mbar + stage * 8, (k / NUM_STAGES - 1) % 2);
      issue_tma(k, stage);
    }
  }
  // MMA warp
  else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (1U << 27U);
    for (int iter_k = 0; iter_k < num_iters; iter_k++) {
      int stage = iter_k % NUM_STAGES;
      mbarrier_wait(tma_mbar + stage * 8, (iter_k / NUM_STAGES) % 2);
      const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
      const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
      auto make_AB = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL); };
      auto make_SF = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL); };
      uint64_t SFA_d = make_SF(0) + ((uint64_t)SFA_sm >> 4ULL);
      uint64_t SFB1_d = make_SF(0) + ((uint64_t)SFB1_sm >> 4ULL);
      uint64_t SFB2_d = make_SF(0) + ((uint64_t)SFB2_sm >> 4ULL);
      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        tcgen05_cp_nvfp4(SFA_tm + k * 4, SFA_d + (uint64_t)k * 32ULL);
        tcgen05_cp_nvfp4(SFB1_tm + k * 4, SFB1_d + (uint64_t)k * 32ULL);
        tcgen05_cp_nvfp4(SFB2_tm + k * 4, SFB2_d + (uint64_t)k * 32ULL);
      }
      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        uint64_t a_d = make_AB(A_sm + k * 32);
        uint64_t b1_d = make_AB(B1_sm + k * 32);
        uint64_t b2_d = make_AB(B2_sm + k * 32);
        int sA = SFA_tm + k * 4;
        int sB1 = SFB1_tm + k * 4 + sfb_n_offset;
        int sB2 = SFB2_tm + k * 4 + sfb_n_offset;
        int en = (k == 0) ? iter_k : 1;
        tcgen05_mma_nvfp4(a_d, b1_d, i_desc, sA, sB1, en, ACC1);
        tcgen05_mma_nvfp4(a_d, b2_d, i_desc, sA, sB2, en, ACC2);
      }
      asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mma_mbar + stage * 8) : "memory");
    }
    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mainloop_mbar) : "memory");
  }
  // Epilogue warps
  else if (tid < BLOCK_M) {
    mbarrier_wait(mainloop_mbar, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");
    const int lane_row = lane_id / 4, lane_col = (lane_id % 4) * 2;
    half* __restrict__ C_base = C_ptr + off_m * N + off_n + lane_col;
    constexpr int CHUNK = 64;
    constexpr int NUM_CHUNKS = BLOCK_N / CHUNK;
    #pragma unroll
    for (int m = 0; m < 2; m++) {
      const int row_base = warp_id * 32 + m * 16;
      half* __restrict__ C_row0 = C_base + (row_base + lane_row) * N;
      half* __restrict__ C_row8 = C_row0 + 8 * N;
      #pragma unroll
      for (int c = 0; c < NUM_CHUNKS; c++) {
        float tmp1[32], tmp2[32];
        tcgen05_ld_16x256bx8(tmp1, row_base, ACC1 + c * CHUNK);
        tcgen05_ld_16x256bx8(tmp2, row_base, ACC2 + c * CHUNK);
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        #pragma unroll
        for (int i = 0; i < 8; i++) {
          float x0 = tmp1[i*4], x1 = tmp1[i*4+1], x2 = tmp1[i*4+2], x3 = tmp1[i*4+3];
          float y0 = tmp2[i*4], y1 = tmp2[i*4+1], y2 = tmp2[i*4+2], y3 = tmp2[i*4+3];
          float e0 = __expf(-x0), e1 = __expf(-x1), e2 = __expf(-x2), e3 = __expf(-x3);
          float s0 = x0 / (1.f + e0) * y0, s1 = x1 / (1.f + e1) * y1;
          float s2 = x2 / (1.f + e2) * y2, s3 = x3 / (1.f + e3) * y3;
          *reinterpret_cast<half2*>(C_row0 + c * CHUNK + i * 8) = __float22half2_rn({s0, s1});
          *reinterpret_cast<half2*>(C_row8 + c * CHUNK + i * 8) = __float22half2_rn({s2, s3});
        }
      }
    }
  }

  // Cleanup
  if (warp_id == 0 && tid == 0) {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
  }
}

// Launch function for cluster kernel
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_N>
at::Tensor launch_cluster(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2, const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
  int M = A.size(0), N = B1.size(0);
  CUtensorMap A_tm, B1_tm, B2_tm;
  init_AB_tmap(&A_tm, (const char*)A.data_ptr(), M, K, BLOCK_M, BLOCK_K);
  init_AB_tmap(&B1_tm, (const char*)B1.data_ptr(), N, K, BLOCK_N, BLOCK_K);
  init_AB_tmap(&B2_tm, (const char*)B2.data_ptr(), N, K, BLOCK_N, BLOCK_K);

  // Grid is in CTAs, not clusters. With __cluster_dims__, runtime groups CTAs.
  int grid_m = M / BLOCK_M;
  int grid_n = N / BLOCK_N;  // Total CTAs, runtime groups into clusters
  dim3 grid(grid_m, grid_n);
  int tb = BLOCK_M + 2 * WARP_SIZE;
  int smem = ((BLOCK_M + 2 * BLOCK_N) * (BLOCK_K / 2) + 128 * (BLOCK_K / 16) * 3) * NUM_STAGES;

  auto kernel = dual_gemm_cluster_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, CLUSTER_N>;
  if (smem > 48000) cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);

  // Use cudaLaunchKernelEx for cluster launch
  cudaLaunchConfig_t config = {};
  config.gridDim = grid;
  config.blockDim = dim3(tb, 1, 1);
  config.dynamicSmemBytes = smem;
  cudaLaunchAttribute attrs[1];
  attrs[0].id = cudaLaunchAttributeClusterDimension;
  attrs[0].val.clusterDim = {1, (unsigned)CLUSTER_N, 1};
  config.attrs = attrs;
  config.numAttrs = 1;
  cudaLaunchKernelEx(&config, kernel, A_tm, B1_tm, B2_tm, (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, grid_n);
  return C;
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
at::Tensor launch(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2, const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
  int M = A.size(0), N = B1.size(0);
  CUtensorMap A_tm, B1_tm, B2_tm;
  init_AB_tmap(&A_tm, (const char*)A.data_ptr(), M, K, BLOCK_M, BLOCK_K);
  init_AB_tmap(&B1_tm, (const char*)B1.data_ptr(), N, K, BLOCK_N, BLOCK_K);
  init_AB_tmap(&B2_tm, (const char*)B2.data_ptr(), N, K, BLOCK_N, BLOCK_K);

  int grid_m = M / BLOCK_M, grid_n = N / BLOCK_N;
  dim3 grid(grid_m, grid_n);
  int tb = BLOCK_M + 2 * WARP_SIZE;
  int smem = ((BLOCK_M + 2 * BLOCK_N) * (BLOCK_K / 2) + 128 * (BLOCK_K / 16) * 3) * NUM_STAGES;

  auto kernel = dual_gemm_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
  if (smem > 48000) cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
  kernel<<<grid, tb, smem>>>(A_tm, B1_tm, B2_tm, (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, grid_n);
  return C;
}

at::Tensor dual_gemm(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2, const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
  int M = A.size(0), K = A.size(1) * 2;

  // Use non-cluster kernel - 5 stages for all cases
#define L64_5(K_) else if (K == K_) C = launch<K_, 128, 64, 256, 5>(A, B1, B2, SFA, SFB1, SFB2, C);
#define L128_4(K_) else if (K == K_) C = launch<K_, 128, 128, 256, 4>(A, B1, B2, SFA, SFB1, SFB2, C);

  if (M <= 256) {
    if (false) {}
    L64_5(16384) L64_5(7168) L64_5(4096) L64_5(2048) L64_5(2304) L64_5(1536) L64_5(512) L64_5(256)
  } else {
    if (false) {}
    L128_4(16384) L128_4(7168) L128_4(4096) L128_4(2048) L128_4(2304) L128_4(1536) L128_4(512) L128_4(256)
  }
#undef L64_5
#undef L128_4
  return C;
}

TORCH_LIBRARY(dual_gemm_mod, m) {
  m.def("dual_gemm(Tensor A, Tensor B1, Tensor B2, Tensor SFA, Tensor SFB1, Tensor SFB2, Tensor(a!) C) -> Tensor");
  m.impl("dual_gemm", &dual_gemm);
}
"""

load_inline("dual_gemm_0", 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"], extra_ldflags=["-lcuda"])

dual_gemm = torch.ops.dual_gemm_mod.dual_gemm

def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c = data
    m, n, l = c.shape
    if l == 1:
        out = torch.empty((m, n), dtype=torch.float16, device="cuda")
        return dual_gemm(a[:,:,0], b1[:,:,0], b2[:,:,0], sfa_perm[:,:,:,:,:,0], sfb1_perm[:,:,:,:,:,0], sfb2_perm[:,:,:,:,:,0], out).unsqueeze(2)
    else:
        out = torch.empty((m, n, l), dtype=torch.float16, device="cuda")
        for i in range(l):
            dual_gemm(a[:,:,i], b1[:,:,i], b2[:,:,i], sfa_perm[:,:,:,:,:,i], sfb1_perm[:,:,:,:,:,i], sfb2_perm[:,:,:,:,:,i], out[:,:,i])
        return out
scrolls · 538 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