Skip to content
KernelIndex
Search⌘K

submission 275764

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-275764?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
16.6µs
#142 of 420
2026-01-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e025044f0390085d99ad11d634f7f93127a1807a738166467b58795adefb10fb
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

cluster__cluster_dims__(2, 1, 1)
fused-epilogueconstexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 3; // 4 epilogue + 1 SF + 1 TMA + 1 MMA = 7
mbarriervoid mbarrier_init(int mbar_addr, int count) {
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
tcgen05asm volatile("tcgen05.cp.cta_group::2.32x128b.warpx4 [%0], %1;"
tile-n = 64constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
tma"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6.L2::cache_hint "
vector-width = st.global.v4"st.global.v4.b64 [%0], {%1, %2, %3, %4};"

Kernel source

submission.py803 lines
#!POPCORN leaderboard nvfp4_dual_gemm
#!POPCORN gpu NVIDIA

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

# ============================================================================
# 2-SM MMA Dual GEMM v6 - Dedicated SF Warp + Pipelining
# Optimization: Separate warp for SF TMA, overlapped with tensor TMA
# ============================================================================

CUDA_SOURCE = r"""
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cstdlib>
#include <cstdio>

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

// L2 Cache Hints (from 1st.py)
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000ULL;
constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

// ============================================================================
// PTX Helper Functions
// ============================================================================

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

// 32B global store (4x64b) to improve L1TEX sector utilization vs 16B stores.
__device__ __forceinline__
void stg_32b(const void* dst, unsigned long long v0, unsigned long long v1,
            unsigned long long v2, unsigned long long v3) {
    asm volatile(
        "st.global.v4.b64 [%0], {%1, %2, %3, %4};"
        :: "l"(dst), "l"(v0), "l"(v1), "l"(v2), "l"(v3)
        : "memory"
    );
}

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

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

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

// TMA with .cta_group::2 and L2 cache hint
// The .cta_group::2 modifier allows mbar_addr and dst to be in different CTA's smem
template <int CTA_GROUP>
__device__ __forceinline__
void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
    asm volatile(
        "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6.L2::cache_hint "
        "[%0], [%1, {%2, %3, %4}], [%5], %7;"
        :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(CTA_GROUP), "l"(cache_policy)
        : "memory"
    );
}

// Bulk TMA with cache hint for scale factors
__device__ __forceinline__
void tma_bulk_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"
    );
}

// Scale factor copy with cta_group::2
__device__ __forceinline__
void tcgen05_cp_cta2(int taddr, uint64_t s_desc) {
    asm volatile("tcgen05.cp.cta_group::2.32x128b.warpx4 [%0], %1;" 
                 :: "r"(taddr), "l"(s_desc));
}

__device__ __forceinline__
void tcgen05_mma_cta2(
    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::2.kind::mxf4nvf4.block_scale.block16 "
        "[%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
    );
}

__device__ __forceinline__
void tcgen05_ld_32x32bx8(float *tmp, int addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x8.b32 "
        "{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
        : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
          "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
        : "r"(addr)
    );
}

// Wider TMEM loads for faster epilogue - takes full address (taddr + (row << 16) + col)
__device__ __forceinline__
void tcgen05_ld_32x32bx32_addr(float *tmp, int addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x32.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"(addr)
    );
}

__device__ __forceinline__
void tcgen05_ld_32x32bx64_addr(float *tmp, int addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x64.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"(addr)
    );
}

// ============================================================================
// TensorMap Creation
// ============================================================================

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

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

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

// Scale-factor TensorMap (UINT16 view) for the permuted SF layout.
// We view SF as a tiled 3D tensor: (512 bytes, mn/128 blocks, K/64 blocks).
// This matches the existing pointer arithmetic:
//   block_index = (mn_block * (K/64) + k_block) * 512
// and lets us use tensor TMA (supports cta_group::2) instead of bulk TMA (does not).
void init_SF_tmap(
    CUtensorMap *tmap,
    const char *ptr,
    uint64_t mn,
    uint64_t K,
    uint32_t block_k  // == BLOCK_K
) {
    constexpr uint32_t rank = 3;
    const uint64_t k_blocks = K / 64;     // 64-element SF granularity
    const uint64_t mn_blocks = mn / 128;  // 128-row/col SF granularity
    const uint32_t tile_k_blocks = block_k / 64;

    // TensorMap has limits on the X dimension; represent a 512B SF block as 256xUINT16.
    constexpr uint64_t SF_BLOCK_BYTES = 512;
    constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t);  // 256
    uint64_t globalDim[rank]       = {X_ELEMS, mn_blocks, k_blocks};
    uint64_t globalStrides[rank-1] = {k_blocks * SF_BLOCK_BYTES, SF_BLOCK_BYTES};  // bytes
    uint32_t boxDim[rank]          = {(uint32_t)X_ELEMS, 1, tile_k_blocks};
    uint32_t elementStrides[rank]  = {1, 1, 1};

    auto err = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_UINT16,
        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
    );
    check_cu(err);
}

// ============================================================================
// 2-SM MMA Dual GEMM Kernel - Following reference pattern
// ============================================================================

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
__global__
__cluster_dims__(2, 1, 1)
__launch_bounds__(BLOCK_M + 3 * WARP_SIZE)
void dual_gemm_cta2_kernel(
    const __grid_constant__ CUtensorMap A_tmap,
    const __grid_constant__ CUtensorMap B1_tmap,
    const __grid_constant__ CUtensorMap B2_tmap,
    const __grid_constant__ CUtensorMap SFA_tmap,
    const __grid_constant__ CUtensorMap SFB1_tmap,
    const __grid_constant__ CUtensorMap SFB2_tmap,
    half *C_ptr,
    int M, int N
) {
    constexpr int CTA_GROUP = 2;
    constexpr int HALF_BLOCK_N = BLOCK_N / CTA_GROUP;
    // v6: Add dedicated SF warp (warp 5), so +3 instead of +2
    constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 3;  // 4 epilogue + 1 SF + 1 TMA + 1 MMA = 7
    
    const int tid = threadIdx.x;
    const int bid = blockIdx.x;
    const int warp_id = tid / WARP_SIZE;
    
    int cta_rank;
    asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
    
    // Grid indexing - M-mode first for cta_group::2
    const int cluster_idx = bid / CTA_GROUP;
    const int grid_n_clusters = N / BLOCK_N;
    const int cluster_m = cluster_idx / grid_n_clusters;
    const int cluster_n = cluster_idx % grid_n_clusters;
    const int off_m = cluster_m * (BLOCK_M * CTA_GROUP) + cta_rank * BLOCK_M;
    const int off_n = cluster_n * BLOCK_N;

    extern __shared__ __align__(1024) char smem_ptr[];
    const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
    
    // SMEM layout (must be identical across CTAs!)
    // In 2-SM MMA, the B operand is split across CTAs (each CTA holds HALF_BLOCK_N columns).
    constexpr int A_size    = BLOCK_M * BLOCK_K / 2;
    constexpr int B1_size   = HALF_BLOCK_N * BLOCK_K / 2;
    constexpr int B2_size   = HALF_BLOCK_N * BLOCK_K / 2;
    constexpr int SFA_size  = 128 * BLOCK_K / 16;
    constexpr int SFB1_size = 128 * BLOCK_K / 16;
    constexpr int SFB2_size = 128 * BLOCK_K / 16;
    constexpr int STAGE_SIZE = A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size;

    // Mbarrier layout:
    // - tma_mbar: count=CTA_GROUP*2
    //   - tensor warp issues expect_tx for tensor bytes (1 arrival per CTA)
    //   - sf warp issues expect_tx for SF bytes (1 arrival per CTA)
    //   Both report into CTA0's mbar (masked address), using .shared::cluster.
    // - mma_mbar: count=1, CTA0 multicasts to both CTAs (stage reuse)
    // - mainloop_mbar: count=1, CTA0 multicasts to both CTAs (epilogue start)
    #pragma nv_diag_suppress static_var_with_dynamic_init
    __shared__ uint64_t mbars[NUM_STAGES * 2 + 1];
    __shared__ int tmem_addr[1];
    const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
    const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
    const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

    // TMEM layout for cta_group::2
    constexpr int ACC1_tmem = 0;
    constexpr int ACC2_tmem = BLOCK_N;
    constexpr int SFA_COLS_PER_K = 8;  // 256 rows / 32
    constexpr int SFB_COLS_PER_K = 4;  // 128 cols / 32
    constexpr int SFA_tmem  = 2 * BLOCK_N;
    constexpr int SFB1_tmem = SFA_tmem + SFA_COLS_PER_K * (BLOCK_K / MMA_K);
    constexpr int SFB2_tmem = SFB1_tmem + SFB_COLS_PER_K * (BLOCK_K / MMA_K);
    // TMEM allocation must be a power-of-2 column count.
    // - For BLOCK_N=128 we need 512 cols (ACC1+ACC2 already consumes 256, plus scale factors).
    // - For BLOCK_N=64, 256 cols is sufficient and can reduce TMEM pressure.
    constexpr int TOTAL_TMEM_COLS = (BLOCK_N <= 64) ? 256 : 512;

    // ========================================================================
    // Initialization - following reference exactly
    // ========================================================================
    if (warp_id == 0 && elect_sync()) {
        for (int i = 0; i < NUM_STAGES; i++) {
            // 4 arrivals = 2 (tensor expect_tx) + 2 (SF expect_tx)
            mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP * 2);
            mbarrier_init(mma_mbar_addr + i * 8, 1);                // CTA0 multicasts to both
        }
        mbarrier_init(mainloop_mbar_addr, 1);  // CTA0 multicasts to both
        asm volatile("fence.mbarrier_init.release.cluster;");
    }
    else if (warp_id == 1) {
        const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
        asm volatile("tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;"
                    :: "r"(addr), "r"(TOTAL_TMEM_COLS));
    }
    
    // Cluster barrier - visible to all threads in cluster
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
    
    const int taddr = tmem_addr[0];

    // Instruction descriptor for MMA_M=256, MMA_N=BLOCK_N
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (2U << 27U);
    constexpr int SBO_AB = 8 * 128;
    constexpr int SBO_SF = 8 * 16;
    constexpr uint64_t AB_desc_base = (desc_encode(SBO_AB) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
    constexpr uint64_t SF_desc_base = (desc_encode(SBO_SF) << 32ULL) | (1ULL << 46ULL);

    const int scale_B_base_off = (cluster_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
    constexpr int num_iters = K / BLOCK_K;

    // L2 cache hints (winner pattern):
    // If M > N, keep B (evict A first); else keep A (evict B first).
    const uint64_t cache_A = (M > N) ? EVICT_FIRST : EVICT_LAST;
    const uint64_t cache_B = (M > N) ? EVICT_LAST  : EVICT_FIRST;
    
    // ========================================================================
    // SF Warp (warp 4) - Issues SF TMA loads in parallel with tensor TMA
    // ========================================================================
    if (warp_id == NUM_WARPS - 3 && elect_sync()) {
        int tma_stage = 0;
        int mma_phase = 1;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
            // Wait for MMA to release this buffer (skip for initial pipeline fill)
            if (iter_k >= NUM_STAGES)
                mbarrier_wait(mma_mbar_addr + tma_stage * 8, mma_phase);

            const int mbar_addr = (tma_mbar_addr + tma_stage * 8) & 0xFEFFFFFF;
            const int off_k = iter_k * BLOCK_K;
            
            // SMEM addresses for SF
            const int base_smem = smem + tma_stage * STAGE_SIZE;
            const int SFA_smem = base_smem + A_size + B1_size + B2_size;
            const int SFB1_smem = SFA_smem + SFA_size;
            const int SFB2_smem = SFB1_smem + SFB1_size;
            
            // Scale-factor tensor TMA: report directly into CTA0's stage mbarrier.
            constexpr int SF_TMA_SIZE = SFA_size + SFB1_size + SFB2_size;
            asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
                        :: "r"(mbar_addr), "r"(SF_TMA_SIZE) : "memory");

            // Scale factors via tensor TMA (supports remote mbarrier through cta_group::2).
            const int sf_y_A = off_m / 128;
            const int sf_y_B = off_n / 128;
            const int sf_z   = off_k / 64;

            tma_3d_gmem2smem<CTA_GROUP>(SFA_smem,  &SFA_tmap,  0, sf_y_A, sf_z, mbar_addr, cache_A);
            tma_3d_gmem2smem<CTA_GROUP>(SFB1_smem, &SFB1_tmap, 0, sf_y_B, sf_z, mbar_addr, cache_B);
            tma_3d_gmem2smem<CTA_GROUP>(SFB2_smem, &SFB2_tmap, 0, sf_y_B, sf_z, mbar_addr, cache_B);

            tma_stage = (tma_stage + 1) % NUM_STAGES;
            if (tma_stage == 0) mma_phase ^= 1;
        }
    }
    // ========================================================================
    // TMA Warp (warp 5) - Issues TENSOR TMA loads only (parallel with SF warp)
    // ========================================================================
    else if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        int tma_stage = 0;
        int mma_phase = 1;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
            // Wait for MMA to release this buffer (skip for initial pipeline fill)
            if (iter_k >= NUM_STAGES)
                mbarrier_wait(mma_mbar_addr + tma_stage * 8, mma_phase);

            const int mbar_addr = (tma_mbar_addr + tma_stage * 8) & 0xFEFFFFFF;
            const int off_k = iter_k * BLOCK_K;
            
            // SMEM addresses
            const int A_smem   = smem + tma_stage * STAGE_SIZE;
            const int B1_smem  = A_smem + A_size;
            const int B2_smem  = B1_smem + B1_size;

            // Arrive.expect_tx for this CTA's tensor TMAs, then issue loads.
            constexpr int TENSOR_TMA_SIZE = A_size + B1_size + B2_size;
            asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
                        :: "r"(mbar_addr), "r"(TENSOR_TMA_SIZE) : "memory");

            // Issue tensor TMA loads (A is not split; B1/B2 are split along N).
            tma_3d_gmem2smem<CTA_GROUP>(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
            const int B_col_offset = off_n + cta_rank * HALF_BLOCK_N;
            tma_3d_gmem2smem<CTA_GROUP>(B1_smem, &B1_tmap, 0, B_col_offset, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem<CTA_GROUP>(B2_smem, &B2_tmap, 0, B_col_offset, off_k / 256, mbar_addr, cache_B);

            tma_stage = (tma_stage + 1) % NUM_STAGES;
            if (tma_stage == 0) mma_phase ^= 1;
        }
    }
    // ========================================================================
    // MMA Warp (warp 6, CTA0 ONLY) - Wait for TMA, issue tcgen05.cp and tcgen05.mma
    // ========================================================================
    else if (cta_rank == 0 && warp_id == NUM_WARPS - 1 && elect_sync()) {
        int tma_stage = 0;
        int tma_phase = 0;
        
        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
            // Wait for ALL TMAs (count=4: 2 tensor expect_tx + 2 SF arrive)
            mbarrier_wait(tma_mbar_addr + tma_stage * 8, tma_phase);
            
            asm volatile("tcgen05.fence::after_thread_sync;");

            // SMEM addresses
            const int base_smem = smem + tma_stage * STAGE_SIZE;
            const int A_smem   = base_smem;
            const int B1_smem  = base_smem + A_size;
            const int B2_smem  = base_smem + A_size + B1_size;
            const int SFA_smem = base_smem + A_size + B1_size + B2_size;
            const int SFB1_smem = SFA_smem + SFA_size;
            const int SFB2_smem = SFB1_smem + SFB1_size;

            // tcgen05.cp - reads from BOTH CTAs' SMEM, writes to BOTH TMEMs
            const uint64_t SFA_desc  = SF_desc_base + ((uint64_t)SFA_smem >> 4ULL);
            const uint64_t SFB1_desc = SF_desc_base + ((uint64_t)SFB1_smem >> 4ULL);
            const uint64_t SFB2_desc = SF_desc_base + ((uint64_t)SFB2_smem >> 4ULL);
            
            #pragma unroll
            for (int k = 0; k < BLOCK_K / MMA_K; k++) {
                tcgen05_cp_cta2(SFA_tmem + k * SFA_COLS_PER_K, SFA_desc + (uint64_t)k * 32ULL);
                tcgen05_cp_cta2(SFB1_tmem + k * SFB_COLS_PER_K, SFB1_desc + (uint64_t)k * 32ULL);
                tcgen05_cp_cta2(SFB2_tmem + k * SFB_COLS_PER_K, SFB2_desc + (uint64_t)k * 32ULL);
            }
            
            // Fence to ensure tcgen05.cp completes before tcgen05.mma
            asm volatile("tcgen05.fence::before_thread_sync;");

            // MMA
            #pragma unroll
            for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
                #pragma unroll
                for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
                    const int a_off = k1 * BLOCK_M * 128 + k2 * 32;
                    const int b_off = k1 * HALF_BLOCK_N * 128 + k2 * 32;
                    
                    uint64_t a_desc  = AB_desc_base + desc_encode(A_smem + a_off);
                    uint64_t b1_desc = AB_desc_base + desc_encode(B1_smem + b_off);
                    uint64_t b2_desc = AB_desc_base + desc_encode(B2_smem + b_off);

                    const int k_sf = k1 * 4 + k2;
                    const int scale_A  = SFA_tmem + k_sf * SFA_COLS_PER_K;
                    const int scale_B1 = SFB1_tmem + k_sf * SFB_COLS_PER_K + scale_B_base_off;
                    const int scale_B2 = SFB2_tmem + k_sf * SFB_COLS_PER_K + scale_B_base_off;

                    const int enable_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
                    tcgen05_mma_cta2(ACC1_tmem, a_desc, b1_desc, i_desc, scale_A, scale_B1, enable_d);
                    tcgen05_mma_cta2(ACC2_tmem, a_desc, b2_desc, i_desc, scale_A, scale_B2, enable_d);
                }
            }

            // Commit MMA - multicast to BOTH CTAs (following reference)
            constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;  // 0b11
            asm volatile("tcgen05.commit.cta_group::2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                        :: "r"(mma_mbar_addr + tma_stage * 8), "h"(cta_mask) : "memory");

            // Flip phase when cycled through all stages
            tma_stage = (tma_stage + 1) % NUM_STAGES;
            if (tma_stage == 0) {
                tma_phase ^= 1;
            }
        }
        
        // Signal mainloop completion - multicast to BOTH CTAs
        constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
        asm volatile("tcgen05.commit.cta_group::2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                    :: "r"(mainloop_mbar_addr), "h"(cta_mask) : "memory");
    }

    // ========================================================================
    // Epilogue - BOTH CTAs wait for mainloop completion
    // Optimized with wider TMEM loads (64 columns at once instead of 8)
    // ========================================================================
    mbarrier_wait(mainloop_mbar_addr, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");

    if (tid < BLOCK_M) {
        // cta_group::2 MMA produces full BLOCK_N columns per CTA accumulator
        constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
        const int tmem_row = cta_rank * 128 + warp_id * 32;
        
        for (int n = 0; n < BLOCK_N / WIDTH; n++) {
            float acc1[WIDTH], acc2[WIDTH];
            
            // Compute full TMEM address: taddr + (row << 16) + col
            const int addr1 = taddr + (tmem_row << 16) + (ACC1_tmem + n * WIDTH);
            const int addr2 = taddr + (tmem_row << 16) + (ACC2_tmem + n * WIDTH);
            
            // Load with wider TMEM loads
            if constexpr (WIDTH == 64) {
                tcgen05_ld_32x32bx64_addr(acc1, addr1);
                tcgen05_ld_32x32bx64_addr(acc2, addr2);
            } else {
                tcgen05_ld_32x32bx32_addr(acc1, addr1);
                tcgen05_ld_32x32bx32_addr(acc2, addr2);
            }
            asm volatile("tcgen05.wait::ld.sync.aligned;");

            // Store C as (M, N) row-major (matches reference layout), vectorized per thread.
            half* row_ptr = C_ptr + (off_m + tid) * N + off_n + n * WIDTH;
            
            #pragma unroll
            for (int i = 0; i < WIDTH; i += 16) {
                float r0  = acc1[i+0]  / (1.0f + __expf(-acc1[i+0]))  * acc2[i+0];
                float r1  = acc1[i+1]  / (1.0f + __expf(-acc1[i+1]))  * acc2[i+1];
                float r2  = acc1[i+2]  / (1.0f + __expf(-acc1[i+2]))  * acc2[i+2];
                float r3  = acc1[i+3]  / (1.0f + __expf(-acc1[i+3]))  * acc2[i+3];
                float r4  = acc1[i+4]  / (1.0f + __expf(-acc1[i+4]))  * acc2[i+4];
                float r5  = acc1[i+5]  / (1.0f + __expf(-acc1[i+5]))  * acc2[i+5];
                float r6  = acc1[i+6]  / (1.0f + __expf(-acc1[i+6]))  * acc2[i+6];
                float r7  = acc1[i+7]  / (1.0f + __expf(-acc1[i+7]))  * acc2[i+7];
                float r8  = acc1[i+8]  / (1.0f + __expf(-acc1[i+8]))  * acc2[i+8];
                float r9  = acc1[i+9]  / (1.0f + __expf(-acc1[i+9]))  * acc2[i+9];
                float r10 = acc1[i+10] / (1.0f + __expf(-acc1[i+10])) * acc2[i+10];
                float r11 = acc1[i+11] / (1.0f + __expf(-acc1[i+11])) * acc2[i+11];
                float r12 = acc1[i+12] / (1.0f + __expf(-acc1[i+12])) * acc2[i+12];
                float r13 = acc1[i+13] / (1.0f + __expf(-acc1[i+13])) * acc2[i+13];
                float r14 = acc1[i+14] / (1.0f + __expf(-acc1[i+14])) * acc2[i+14];
                float r15 = acc1[i+15] / (1.0f + __expf(-acc1[i+15])) * acc2[i+15];

                half2 h0 = __float22half2_rn({r0,  r1});
                half2 h1 = __float22half2_rn({r2,  r3});
                half2 h2 = __float22half2_rn({r4,  r5});
                half2 h3 = __float22half2_rn({r6,  r7});
                half2 h4 = __float22half2_rn({r8,  r9});
                half2 h5 = __float22half2_rn({r10, r11});
                half2 h6 = __float22half2_rn({r12, r13});
                half2 h7 = __float22half2_rn({r14, r15});

                const uint32_t u0 = *reinterpret_cast<uint32_t*>(&h0);
                const uint32_t u1 = *reinterpret_cast<uint32_t*>(&h1);
                const uint32_t u2 = *reinterpret_cast<uint32_t*>(&h2);
                const uint32_t u3 = *reinterpret_cast<uint32_t*>(&h3);
                const uint32_t u4 = *reinterpret_cast<uint32_t*>(&h4);
                const uint32_t u5 = *reinterpret_cast<uint32_t*>(&h5);
                const uint32_t u6 = *reinterpret_cast<uint32_t*>(&h6);
                const uint32_t u7 = *reinterpret_cast<uint32_t*>(&h7);

                const unsigned long long q0 = (unsigned long long)u0 | ((unsigned long long)u1 << 32);
                const unsigned long long q1 = (unsigned long long)u2 | ((unsigned long long)u3 << 32);
                const unsigned long long q2 = (unsigned long long)u4 | ((unsigned long long)u5 << 32);
                const unsigned long long q3 = (unsigned long long)u6 | ((unsigned long long)u7 << 32);

                stg_32b((const void*)(row_ptr + i), q0, q1, q2, q3);
            }
        }
    }

    // Cluster barrier before deallocation (following reference)
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
    
    if (warp_id == 0) {
        asm volatile("tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;" 
                    :: "r"(taddr), "r"(TOTAL_TMEM_COLS));
    }
}

// ============================================================================
// Launch Wrapper
// ============================================================================

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
at::Tensor dual_gemm_cta2_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
) {
    constexpr int HALF_BLOCK_N = BLOCK_N / 2;
    
    const int M = A.size(0);
    const int N = B1.size(0);

    auto A_ptr    = reinterpret_cast<const char *>(A.data_ptr());
    auto B1_ptr   = reinterpret_cast<const char *>(B1.data_ptr());
    auto B2_ptr   = reinterpret_cast<const char *>(B2.data_ptr());
    auto SFA_ptr  = reinterpret_cast<const char *>(SFA.data_ptr());
    auto SFB1_ptr = reinterpret_cast<const char *>(SFB1.data_ptr());
    auto SFB2_ptr = reinterpret_cast<const char *>(SFB2.data_ptr());
    auto C_ptr    = reinterpret_cast<half *>(C.data_ptr());

    CUtensorMap A_tmap, B1_tmap, B2_tmap;
    init_AB_tmap(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K);
    init_AB_tmap(&B1_tmap, B1_ptr, N, K, HALF_BLOCK_N, BLOCK_K);
    init_AB_tmap(&B2_tmap, B2_ptr, N, K, HALF_BLOCK_N, BLOCK_K);

    CUtensorMap SFA_tmap, SFB1_tmap, SFB2_tmap;
    init_SF_tmap(&SFA_tmap,  SFA_ptr,  M, K, BLOCK_K);
    init_SF_tmap(&SFB1_tmap, SFB1_ptr, N, K, BLOCK_K);
    init_SF_tmap(&SFB2_tmap, SFB2_ptr, N, K, BLOCK_K);

    const int num_blocks = (M / BLOCK_M) * (N / BLOCK_N);
    dim3 grid(num_blocks, 1, 1);
    int tb_size = BLOCK_M + 3 * WARP_SIZE;  // +3 for SF, TMA, MMA warps
    
    constexpr int A_size_c    = BLOCK_M * BLOCK_K / 2;
    constexpr int B1_size_c   = HALF_BLOCK_N * BLOCK_K / 2;
    constexpr int B2_size_c   = HALF_BLOCK_N * BLOCK_K / 2;
    constexpr int SFA_size_c  = 128 * BLOCK_K / 16;
    constexpr int SFB1_size_c = 128 * BLOCK_K / 16;
    constexpr int SFB2_size_c = 128 * BLOCK_K / 16;
    int smem_size = (A_size_c + B1_size_c + B2_size_c + SFA_size_c + SFB1_size_c + SFB2_size_c) * NUM_STAGES;

    auto kernel_fn = dual_gemm_cta2_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
    if (smem_size > 48000)
        cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

    kernel_fn<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap, SFA_tmap, SFB1_tmap, SFB2_tmap, C_ptr, M, 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
) {
    const int K = A.size(1) * 2;
    const int M = A.size(0);
    const int N = B1.size(0);

#define LAUNCH(K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_) \
    if (K == K_) return dual_gemm_cta2_launch<K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_>( \
        A, B1, B2, SFA, SFB1, SFB2, C);

    // With MMA_M=256 (cta_group::2), the M=256 case has only 1 cluster in M,
    // so cluster count ~= N / BLOCK_N. Using BLOCK_N=64 increases cluster count vs 128.
    //
    // IMPORTANT: for M=512, reducing BLOCK_N often hurts (fixed overhead per tile dominates),
    // so we only enable the 64-wide path for M=256.
    bool use_block_n_64 = (M == 256);
    // Optional override for benchmarking: set NVFP4_M256_BLOCK_N to 64 or 128.
    if (M == 256) {
        if (const char* env = std::getenv("NVFP4_M256_BLOCK_N")) {
            const int v = std::atoi(env);
            if (v == 128) use_block_n_64 = false;
            else if (v == 64) use_block_n_64 = true;
        }
    }

    if (use_block_n_64) {
        // M=256 underfills with BLOCK_N=128; try BLOCK_N=64 for more clusters in N.
        LAUNCH(7168, 128,  64, 256, 7)
        LAUNCH(4096, 128,  64, 256, 7)
        LAUNCH(3072, 128,  64, 256, 7)
        LAUNCH(2304, 128,  64, 256, 7)
        LAUNCH(2048, 128,  64, 256, 7)
        LAUNCH(1536, 128,  64, 256, 6)
        LAUNCH(1024, 128,  64, 256, 4)
        LAUNCH(512,  128,  64, 256, 2)
        LAUNCH(256,  128,  64, 256, 1)
    } else {
        LAUNCH(7168, 128, 128, 256, 5)
        LAUNCH(4096, 128, 128, 256, 5)
        LAUNCH(3072, 128, 128, 256, 5)
        LAUNCH(2304, 128, 128, 256, 5)
        LAUNCH(2048, 128, 128, 256, 5)
        LAUNCH(1536, 128, 128, 256, 5)
        LAUNCH(1024, 128, 128, 256, 4)
        LAUNCH(512,  128, 128, 256, 2)
        LAUNCH(256,  128, 128, 256, 1)
    }

#undef LAUNCH

    TORCH_CHECK(false, "Unsupported K value: ", K);
}

TORCH_LIBRARY(dual_gemm_cta2_v6_module, 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);
}
"""

_compiled_module = None

def _get_module():
    global _compiled_module
    if _compiled_module is None:
        _compiled_module = load_inline(
            "dual_gemm_cta2_v6_cuda",
            cpp_sources="",
            cuda_sources=CUDA_SOURCE,
            verbose=True,
            is_python_module=False,
            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"],
        )
    return _compiled_module


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
    _get_module()
    result = torch.ops.dual_gemm_cta2_v6_module.dual_gemm(
        a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c
    )
    return result

scrolls · 803 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 227611.

⋯ 5 unchanged lines
from torch.utils.cpp_extension import load_inline
# ============================================================================
- # Fused Dual GEMM: Load A ONCE, compute both GEMMs in parallel
- # C = silu(A @ B1) * (A @ B2)
- #
- # Key optimization: A is loaded only once per K-iteration, then used for
- # both B1 and B2 matrix multiplications.
+ # 2-SM MMA Dual GEMM v6 - Dedicated SF Warp + Pipelining
+ # Optimization: Separate warp for SF TMA, overlapped with tensor TMA
# ============================================================================
CUDA_SOURCE = r"""
⋯ 2 unchanged lines
#include <cuda_fp8.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
+ #include <cstdlib>
+ #include <cstdio>
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
+ // L2 Cache Hints (from 1st.py)
+ constexpr uint64_t EVICT_NORMAL = 0x1000000000000000ULL;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
⋯ 6 unchanged lines
return (x & 0x3'FFFFULL) >> 4ULL;
}
+ // 32B global store (4x64b) to improve L1TEX sector utilization vs 16B stores.
__device__ __forceinline__
+ void stg_32b(const void* dst, unsigned long long v0, unsigned long long v1,
+ unsigned long long v2, unsigned long long v3) {
+ asm volatile(
+ "st.global.v4.b64 [%0], {%1, %2, %3, %4};"
+ :: "l"(dst), "l"(v0), "l"(v1), "l"(v2), "l"(v3)
+ : "memory"
+ );
+ }
+
+ __device__ __forceinline__
uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
⋯ 30 unchanged lines
);
}
+ // TMA with .cta_group::2 and L2 cache hint
+ // The .cta_group::2 modifier allows mbar_addr and dst to be in different CTA's smem
+ template <int CTA_GROUP>
__device__ __forceinline__
- void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
+ 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.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)
+ "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6.L2::cache_hint "
+ "[%0], [%1, {%2, %3, %4}], [%5], %7;"
+ :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(CTA_GROUP), "l"(cache_policy)
+ : "memory"
);
}
+ // Bulk TMA with cache hint for scale factors
__device__ __forceinline__
- void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z,
- int mbar_addr, uint64_t cache_policy) {
+ void tma_bulk_gmem2smem(int dst, const void *src, int size, 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"
+ "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"
);
}
+ // Scale factor copy with cta_group::2
__device__ __forceinline__
- void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
- asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;"
+ void tcgen05_cp_cta2(int taddr, uint64_t s_desc) {
+ asm volatile("tcgen05.cp.cta_group::2.32x128b.warpx4 [%0], %1;"
:: "r"(taddr), "l"(s_desc));
}
- // MMA with explicit destination TMEM address
__device__ __forceinline__
- void tcgen05_mma_nvfp4_at(
+ void tcgen05_mma_cta2(
int d_tmem,
uint64_t a_desc,
uint64_t b_desc,
⋯ 6 unchanged lines
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %6, 0;\n\t"
- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
+ "tcgen05.mma.cta_group::2.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),
⋯ 2 unchanged lines
}
__device__ __forceinline__
- void tcgen05_ld_32x32bx32(float *tmp, int row, int col) {
+ void tcgen05_ld_32x32bx8(float *tmp, int addr) {
asm volatile(
+ "tcgen05.ld.sync.aligned.32x32b.x8.b32 "
+ "{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
+ : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
+ "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
+ : "r"(addr)
+ );
+ }
+
+ // Wider TMEM loads for faster epilogue - takes full address (taddr + (row << 16) + col)
+ __device__ __forceinline__
+ void tcgen05_ld_32x32bx32_addr(float *tmp, int addr) {
+ asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x32.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
⋯ 7 unchanged lines
"=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)
+ : "r"(addr)
);
}
__device__ __forceinline__
- void tcgen05_ld_32x32bx64(float *tmp, int row, int col) {
+ void tcgen05_ld_32x32bx64_addr(float *tmp, int addr) {
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x64.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
⋯ 20 unchanged lines
"=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)
+ : "r"(addr)
);
}
⋯ 40 unchanged lines
check_cu(err);
}
+ // Scale-factor TensorMap (UINT16 view) for the permuted SF layout.
+ // We view SF as a tiled 3D tensor: (512 bytes, mn/128 blocks, K/64 blocks).
+ // This matches the existing pointer arithmetic:
+ // block_index = (mn_block * (K/64) + k_block) * 512
+ // and lets us use tensor TMA (supports cta_group::2) instead of bulk TMA (does not).
+ void init_SF_tmap(
+ CUtensorMap *tmap,
+ const char *ptr,
+ uint64_t mn,
+ uint64_t K,
+ uint32_t block_k // == BLOCK_K
+ ) {
+ constexpr uint32_t rank = 3;
+ const uint64_t k_blocks = K / 64; // 64-element SF granularity
+ const uint64_t mn_blocks = mn / 128; // 128-row/col SF granularity
+ const uint32_t tile_k_blocks = block_k / 64;
+
+ // TensorMap has limits on the X dimension; represent a 512B SF block as 256xUINT16.
+ constexpr uint64_t SF_BLOCK_BYTES = 512;
+ constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t); // 256
+ uint64_t globalDim[rank] = {X_ELEMS, mn_blocks, k_blocks};
+ uint64_t globalStrides[rank-1] = {k_blocks * SF_BLOCK_BYTES, SF_BLOCK_BYTES}; // bytes
+ uint32_t boxDim[rank] = {(uint32_t)X_ELEMS, 1, tile_k_blocks};
+ uint32_t elementStrides[rank] = {1, 1, 1};
+
+ auto err = cuTensorMapEncodeTiled(
+ tmap,
+ CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_UINT16,
+ 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
+ );
+ check_cu(err);
+ }
+
// ============================================================================
- // Fused Dual GEMM Kernel - Load A ONCE, compute both GEMMs
- //
- // SMEM Layout per stage:
- // [A tile][B1 tile][B2 tile][SFA][SFB1][SFB2]
- //
- // TMEM Layout:
- // [ACC1: 0..BLOCK_N-1][ACC2: BLOCK_N..2*BLOCK_N-1][SFA][SFB1][SFB2]
+ // 2-SM MMA Dual GEMM Kernel - Following reference pattern
// ============================================================================
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_fused_kernel(
+ __cluster_dims__(2, 1, 1)
+ __launch_bounds__(BLOCK_M + 3 * WARP_SIZE)
+ void dual_gemm_cta2_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,
+ const __grid_constant__ CUtensorMap SFA_tmap,
+ const __grid_constant__ CUtensorMap SFB1_tmap,
+ const __grid_constant__ CUtensorMap SFB2_tmap,
half *C_ptr,
int M, int N
) {
+ constexpr int CTA_GROUP = 2;
+ constexpr int HALF_BLOCK_N = BLOCK_N / CTA_GROUP;
+ // v6: Add dedicated SF warp (warp 5), so +3 instead of +2
+ constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 3; // 4 epilogue + 1 SF + 1 TMA + 1 MMA = 7
+
const int tid = threadIdx.x;
const int bid = blockIdx.x;
const int warp_id = tid / WARP_SIZE;
+
+ int cta_rank;
+ asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
+
+ // Grid indexing - M-mode first for cta_group::2
+ const int cluster_idx = bid / CTA_GROUP;
+ const int grid_n_clusters = N / BLOCK_N;
+ const int cluster_m = cluster_idx / grid_n_clusters;
+ const int cluster_n = cluster_idx % grid_n_clusters;
+ const int off_m = cluster_m * (BLOCK_M * CTA_GROUP) + cta_rank * BLOCK_M;
+ const int off_n = cluster_n * BLOCK_N;
- const int grid_n = N / BLOCK_N;
- const int bid_m = bid / grid_n;
- const int bid_n = bid % grid_n;
- const int off_m = bid_m * BLOCK_M;
- const int off_n = bid_n * BLOCK_N;
-
- constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
-
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
- // SMEM sizes - now with A, B1, B2, SFA, SFB1, SFB2 per stage
+ // SMEM layout (must be identical across CTAs!)
+ // In 2-SM MMA, the B operand is split across CTAs (each CTA holds HALF_BLOCK_N columns).
constexpr int A_size = BLOCK_M * BLOCK_K / 2;
- constexpr int B1_size = BLOCK_N * BLOCK_K / 2;
- constexpr int B2_size = BLOCK_N * BLOCK_K / 2;
+ constexpr int B1_size = HALF_BLOCK_N * BLOCK_K / 2;
+ constexpr int B2_size = HALF_BLOCK_N * BLOCK_K / 2;
constexpr int SFA_size = 128 * BLOCK_K / 16;
constexpr int SFB1_size = 128 * BLOCK_K / 16;
constexpr int SFB2_size = 128 * BLOCK_K / 16;
constexpr int STAGE_SIZE = A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size;
+ // Mbarrier layout:
+ // - tma_mbar: count=CTA_GROUP*2
+ // - tensor warp issues expect_tx for tensor bytes (1 arrival per CTA)
+ // - sf warp issues expect_tx for SF bytes (1 arrival per CTA)
+ // Both report into CTA0's mbar (masked address), using .shared::cluster.
+ // - mma_mbar: count=1, CTA0 multicasts to both CTAs (stage reuse)
+ // - mainloop_mbar: count=1, CTA0 multicasts to both CTAs (epilogue start)
#pragma nv_diag_suppress static_var_with_dynamic_init
- __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
+ __shared__ uint64_t mbars[NUM_STAGES * 2 + 1];
+ __shared__ int tmem_addr[1];
const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;
- // TMEM layout:
- // ACC1: columns [0, BLOCK_N)
- // ACC2: columns [BLOCK_N, 2*BLOCK_N)
- // SFA: columns [2*BLOCK_N, 2*BLOCK_N + 16)
- // SFB1: columns [2*BLOCK_N + 16, 2*BLOCK_N + 32)
- // SFB2: columns [2*BLOCK_N + 32, 2*BLOCK_N + 48)
+ // TMEM layout for cta_group::2
constexpr int ACC1_tmem = 0;
constexpr int ACC2_tmem = BLOCK_N;
+ constexpr int SFA_COLS_PER_K = 8; // 256 rows / 32
+ constexpr int SFB_COLS_PER_K = 4; // 128 cols / 32
constexpr int SFA_tmem = 2 * BLOCK_N;
- constexpr int SFB1_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
- constexpr int SFB2_tmem = SFB1_tmem + 4 * (BLOCK_K / MMA_K);
- // TMEM allocation must be power of 2: 64*2 + 48 = 176 -> round up to 256
- constexpr int TOTAL_TMEM_COLS = 256;
+ constexpr int SFB1_tmem = SFA_tmem + SFA_COLS_PER_K * (BLOCK_K / MMA_K);
+ constexpr int SFB2_tmem = SFB1_tmem + SFB_COLS_PER_K * (BLOCK_K / MMA_K);
+ // TMEM allocation must be a power-of-2 column count.
+ // - For BLOCK_N=128 we need 512 cols (ACC1+ACC2 already consumes 256, plus scale factors).
+ // - For BLOCK_N=64, 256 cols is sufficient and can reduce TMEM pressure.
+ constexpr int TOTAL_TMEM_COLS = (BLOCK_N <= 64) ? 256 : 512;
+ // ========================================================================
+ // Initialization - following reference exactly
+ // ========================================================================
if (warp_id == 0 && elect_sync()) {
- for (int i = 0; i < NUM_STAGES * 2 + 1; i++)
- mbarrier_init(tma_mbar_addr + i * 8, 1);
+ for (int i = 0; i < NUM_STAGES; i++) {
+ // 4 arrivals = 2 (tensor expect_tx) + 2 (SF expect_tx)
+ mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP * 2);
+ mbarrier_init(mma_mbar_addr + i * 8, 1); // CTA0 multicasts to both
+ }
+ mbarrier_init(mainloop_mbar_addr, 1); // CTA0 multicasts to both
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"(TOTAL_TMEM_COLS));
+ const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
+ asm volatile("tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;"
+ :: "r"(addr), "r"(TOTAL_TMEM_COLS));
}
- __syncthreads();
+
+ // Cluster barrier - visible to all threads in cluster
+ asm volatile("barrier.cluster.arrive.release.aligned;");
+ asm volatile("barrier.cluster.wait.acquire.aligned;");
+
+ const int taddr = tmem_addr[0];
+ // Instruction descriptor for MMA_M=256, MMA_N=BLOCK_N
+ constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (2U << 27U);
+ constexpr int SBO_AB = 8 * 128;
+ constexpr int SBO_SF = 8 * 16;
+ constexpr uint64_t AB_desc_base = (desc_encode(SBO_AB) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
+ constexpr uint64_t SF_desc_base = (desc_encode(SBO_SF) << 32ULL) | (1ULL << 46ULL);
+
+ const int scale_B_base_off = (cluster_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
constexpr int num_iters = K / BLOCK_K;
+ // L2 cache hints (winner pattern):
+ // If M > N, keep B (evict A first); else keep A (evict B first).
+ const uint64_t cache_A = (M > N) ? EVICT_FIRST : EVICT_LAST;
+ const uint64_t cache_B = (M > N) ? EVICT_LAST : EVICT_FIRST;
+
// ========================================================================
- // TMA Warp: Load A, B1, B2, SFA, SFB1, SFB2 per iteration
- // Key: A is loaded ONCE and used for BOTH GEMMs!
+ // SF Warp (warp 4) - Issues SF TMA loads in parallel with tensor TMA
// ========================================================================
- if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- // Cache policy based on tile reuse patterns
- // A is reused for BOTH B1 and B2, so keep it longer (EVICT_LAST)
- // B tiles are used once, so can evict earlier (EVICT_FIRST)
- // Also consider M vs N: if M < N, more B tiles so evict B first
- uint64_t cache_A = (M < N) ? EVICT_LAST : EVICT_FIRST;
- uint64_t cache_B = (M < N) ? EVICT_FIRST : EVICT_LAST;
+ if (warp_id == NUM_WARPS - 3 && elect_sync()) {
+ int tma_stage = 0;
+ int mma_phase = 1;
- auto issue_tma = [&](int iter_k, int stage_id) {
- const int mbar_addr = tma_mbar_addr + stage_id * 8;
- const int A_smem = smem + stage_id * STAGE_SIZE;
- const int B1_smem = A_smem + A_size;
- const int B2_smem = B1_smem + B1_size;
- const int SFA_smem = B2_smem + B2_size;
- const int SFB1_smem = SFA_smem + SFA_size;
- const int SFB2_smem = SFB1_smem + SFB1_size;
+ for (int iter_k = 0; iter_k < num_iters; iter_k++) {
+ // Wait for MMA to release this buffer (skip for initial pipeline fill)
+ if (iter_k >= NUM_STAGES)
+ mbarrier_wait(mma_mbar_addr + tma_stage * 8, mma_phase);
+ const int mbar_addr = (tma_mbar_addr + tma_stage * 8) & 0xFEFFFFFF;
const int off_k = iter_k * BLOCK_K;
-
- // Load A (shared between both GEMMs)
- tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
- // Load B1 and B2
- tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
- tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
-
- // Load scale factors
- const int rest_k = K / 16 / 4;
- const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512;
- const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
- const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
+ // SMEM addresses for SF
+ const int base_smem = smem + tma_stage * STAGE_SIZE;
+ const int SFA_smem = base_smem + A_size + B1_size + B2_size;
+ const int SFB1_smem = SFA_smem + SFA_size;
+ const int SFB2_smem = SFB1_smem + SFB1_size;
- // Use same cache hints as the matrices
- tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
- tma_gmem2smem(SFB1_smem, SFB1_src, SFB1_size, mbar_addr, cache_B);
- tma_gmem2smem(SFB2_smem, SFB2_src, SFB2_size, mbar_addr, cache_B);
+ // Scale-factor tensor TMA: report directly into CTA0's stage mbarrier.
+ constexpr int SF_TMA_SIZE = SFA_size + SFB1_size + SFB2_size;
+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
+ :: "r"(mbar_addr), "r"(SF_TMA_SIZE) : "memory");
- asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
- :: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
- };
+ // Scale factors via tensor TMA (supports remote mbarrier through cta_group::2).
+ const int sf_y_A = off_m / 128;
+ const int sf_y_B = off_n / 128;
+ const int sf_z = off_k / 64;
- for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++)
- issue_tma(iter_k, iter_k);
+ tma_3d_gmem2smem<CTA_GROUP>(SFA_smem, &SFA_tmap, 0, sf_y_A, sf_z, mbar_addr, cache_A);
+ tma_3d_gmem2smem<CTA_GROUP>(SFB1_smem, &SFB1_tmap, 0, sf_y_B, sf_z, mbar_addr, cache_B);
+ tma_3d_gmem2smem<CTA_GROUP>(SFB2_smem, &SFB2_tmap, 0, sf_y_B, sf_z, mbar_addr, cache_B);
- for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
- const int stage_id = iter_k % NUM_STAGES;
- const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
- mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
- issue_tma(iter_k, stage_id);
+ tma_stage = (tma_stage + 1) % NUM_STAGES;
+ if (tma_stage == 0) mma_phase ^= 1;
}
}
// ========================================================================
- // MMA Warp: Execute BOTH GEMMs using the SAME A tile
+ // TMA Warp (warp 5) - Issues TENSOR TMA loads only (parallel with SF 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)
- | ((uint32_t)128 >> 7U << 27U);
+ else if (warp_id == NUM_WARPS - 2 && elect_sync()) {
+ int tma_stage = 0;
+ int mma_phase = 1;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
- const int stage_id = iter_k % NUM_STAGES;
- const int tma_phase = (iter_k / NUM_STAGES) % 2;
- mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);
+ // Wait for MMA to release this buffer (skip for initial pipeline fill)
+ if (iter_k >= NUM_STAGES)
+ mbarrier_wait(mma_mbar_addr + tma_stage * 8, mma_phase);
- const int A_smem = smem + stage_id * STAGE_SIZE;
+ const int mbar_addr = (tma_mbar_addr + tma_stage * 8) & 0xFEFFFFFF;
+ const int off_k = iter_k * BLOCK_K;
+
+ // SMEM addresses
+ const int A_smem = smem + tma_stage * STAGE_SIZE;
const int B1_smem = A_smem + A_size;
const int B2_smem = B1_smem + B1_size;
- const int SFA_smem = B2_smem + B2_size;
- const int SFB1_smem = SFA_smem + SFA_size;
- const int SFB2_smem = SFB1_smem + SFB1_size;
- auto make_desc_AB = [](int addr) -> uint64_t {
- const int SBO = 8 * 128;
- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
- };
- auto make_desc_SF = [](int addr) -> uint64_t {
- const int SBO = 8 * 16;
- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
- };
+ // Arrive.expect_tx for this CTA's tensor TMAs, then issue loads.
+ constexpr int TENSOR_TMA_SIZE = A_size + B1_size + B2_size;
+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
+ :: "r"(mbar_addr), "r"(TENSOR_TMA_SIZE) : "memory");
- constexpr uint64_t SF_desc = make_desc_SF(0);
- const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
- const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
- const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);
+ // Issue tensor TMA loads (A is not split; B1/B2 are split along N).
+ tma_3d_gmem2smem<CTA_GROUP>(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
+ const int B_col_offset = off_n + cta_rank * HALF_BLOCK_N;
+ tma_3d_gmem2smem<CTA_GROUP>(B1_smem, &B1_tmap, 0, B_col_offset, off_k / 256, mbar_addr, cache_B);
+ tma_3d_gmem2smem<CTA_GROUP>(B2_smem, &B2_tmap, 0, B_col_offset, off_k / 256, mbar_addr, cache_B);
- // Copy ALL scale factors to TMEM
+ tma_stage = (tma_stage + 1) % NUM_STAGES;
+ if (tma_stage == 0) mma_phase ^= 1;
+ }
+ }
+ // ========================================================================
+ // MMA Warp (warp 6, CTA0 ONLY) - Wait for TMA, issue tcgen05.cp and tcgen05.mma
+ // ========================================================================
+ else if (cta_rank == 0 && warp_id == NUM_WARPS - 1 && elect_sync()) {
+ int tma_stage = 0;
+ int tma_phase = 0;
+
+ for (int iter_k = 0; iter_k < num_iters; iter_k++) {
+ // Wait for ALL TMAs (count=4: 2 tensor expect_tx + 2 SF arrive)
+ mbarrier_wait(tma_mbar_addr + tma_stage * 8, tma_phase);
+
+ asm volatile("tcgen05.fence::after_thread_sync;");
+
+ // SMEM addresses
+ const int base_smem = smem + tma_stage * STAGE_SIZE;
+ const int A_smem = base_smem;
+ const int B1_smem = base_smem + A_size;
+ const int B2_smem = base_smem + A_size + B1_size;
+ const int SFA_smem = base_smem + A_size + B1_size + B2_size;
+ const int SFB1_smem = SFA_smem + SFA_size;
+ const int SFB2_smem = SFB1_smem + SFB1_size;
+
+ // tcgen05.cp - reads from BOTH CTAs' SMEM, writes to BOTH TMEMs
+ const uint64_t SFA_desc = SF_desc_base + ((uint64_t)SFA_smem >> 4ULL);
+ const uint64_t SFB1_desc = SF_desc_base + ((uint64_t)SFB1_smem >> 4ULL);
+ const uint64_t SFB2_desc = SF_desc_base + ((uint64_t)SFB2_smem >> 4ULL);
+
+ #pragma unroll
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
- tcgen05_cp_nvfp4(SFA_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
- tcgen05_cp_nvfp4(SFB1_tmem + k * 4, SFB1_desc + (uint64_t)k * (512ULL >> 4ULL));
- tcgen05_cp_nvfp4(SFB2_tmem + k * 4, SFB2_desc + (uint64_t)k * (512ULL >> 4ULL));
+ tcgen05_cp_cta2(SFA_tmem + k * SFA_COLS_PER_K, SFA_desc + (uint64_t)k * 32ULL);
+ tcgen05_cp_cta2(SFB1_tmem + k * SFB_COLS_PER_K, SFB1_desc + (uint64_t)k * 32ULL);
+ tcgen05_cp_cta2(SFB2_tmem + k * SFB_COLS_PER_K, SFB2_desc + (uint64_t)k * 32ULL);
}
+
+ // Fence to ensure tcgen05.cp completes before tcgen05.mma
+ asm volatile("tcgen05.fence::before_thread_sync;");
- // Execute BOTH MMAs using the SAME A tile
- for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
+ // MMA
+ #pragma unroll
+ for (int k1 = 0; k1 < BLOCK_K / 256; k1++) {
+ #pragma unroll
for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
- uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
- uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
- uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
+ const int a_off = k1 * BLOCK_M * 128 + k2 * 32;
+ const int b_off = k1 * HALF_BLOCK_N * 128 + k2 * 32;
+
+ uint64_t a_desc = AB_desc_base + desc_encode(A_smem + a_off);
+ uint64_t b1_desc = AB_desc_base + desc_encode(B1_smem + b_off);
+ uint64_t b2_desc = AB_desc_base + desc_encode(B2_smem + b_off);
- int k_sf = k1 * 4 + k2;
- const int scale_A = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
- const int scale_B1 = SFB1_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
- const int scale_B2 = SFB2_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
+ const int k_sf = k1 * 4 + k2;
+ const int scale_A = SFA_tmem + k_sf * SFA_COLS_PER_K;
+ const int scale_B1 = SFB1_tmem + k_sf * SFB_COLS_PER_K + scale_B_base_off;
+ const int scale_B2 = SFB2_tmem + k_sf * SFB_COLS_PER_K + scale_B_base_off;
const int enable_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
-
- // GEMM1: A @ B1 -> ACC1
- tcgen05_mma_nvfp4_at(ACC1_tmem, a_desc, b1_desc, i_desc, scale_A, scale_B1, enable_d);
-
- // GEMM2: A @ B2 -> ACC2 (using same A descriptor!)
- tcgen05_mma_nvfp4_at(ACC2_tmem, a_desc, b2_desc, i_desc, scale_A, scale_B2, enable_d);
+ tcgen05_mma_cta2(ACC1_tmem, a_desc, b1_desc, i_desc, scale_A, scale_B1, enable_d);
+ tcgen05_mma_cta2(ACC2_tmem, a_desc, b2_desc, i_desc, scale_A, scale_B2, enable_d);
}
+ }
- asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
- :: "r"(mma_mbar_addr + stage_id * 8) : "memory");
+ // Commit MMA - multicast to BOTH CTAs (following reference)
+ constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1; // 0b11
+ asm volatile("tcgen05.commit.cta_group::2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
+ :: "r"(mma_mbar_addr + tma_stage * 8), "h"(cta_mask) : "memory");
+
+ // Flip phase when cycled through all stages
+ tma_stage = (tma_stage + 1) % NUM_STAGES;
+ if (tma_stage == 0) {
+ tma_phase ^= 1;
+ }
}
- asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
- :: "r"(mainloop_mbar_addr) : "memory");
+
+ // Signal mainloop completion - multicast to BOTH CTAs
+ constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
+ asm volatile("tcgen05.commit.cta_group::2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
+ :: "r"(mainloop_mbar_addr), "h"(cta_mask) : "memory");
}
+
// ========================================================================
- // Epilogue Warps: Load ACC1 and ACC2, apply SiLU + multiply, store
- // Optimized with vectorized stores
+ // Epilogue - BOTH CTAs wait for mainloop completion
+ // Optimized with wider TMEM loads (64 columns at once instead of 8)
// ========================================================================
- else if (tid < BLOCK_M) {
- mbarrier_wait(mainloop_mbar_addr, 0);
- asm volatile("tcgen05.fence::after_thread_sync;");
+ mbarrier_wait(mainloop_mbar_addr, 0);
+ asm volatile("tcgen05.fence::after_thread_sync;");
+ if (tid < BLOCK_M) {
+ // cta_group::2 MMA produces full BLOCK_N columns per CTA accumulator
constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
+ const int tmem_row = cta_rank * 128 + warp_id * 32;
for (int n = 0; n < BLOCK_N / WIDTH; n++) {
- float acc1[WIDTH];
- float acc2[WIDTH];
+ float acc1[WIDTH], acc2[WIDTH];
- // Load both accumulators
+ // Compute full TMEM address: taddr + (row << 16) + col
+ const int addr1 = taddr + (tmem_row << 16) + (ACC1_tmem + n * WIDTH);
+ const int addr2 = taddr + (tmem_row << 16) + (ACC2_tmem + n * WIDTH);
+
+ // Load with wider TMEM loads
if constexpr (WIDTH == 64) {
- tcgen05_ld_32x32bx64(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
- tcgen05_ld_32x32bx64(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
+ tcgen05_ld_32x32bx64_addr(acc1, addr1);
+ tcgen05_ld_32x32bx64_addr(acc2, addr2);
} else {
- tcgen05_ld_32x32bx32(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
- tcgen05_ld_32x32bx32(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
+ tcgen05_ld_32x32bx32_addr(acc1, addr1);
+ tcgen05_ld_32x32bx32_addr(acc2, addr2);
}
asm volatile("tcgen05.wait::ld.sync.aligned;");
- const int row = off_m + tid;
- half* row_ptr = C_ptr + row * N + off_n + n * WIDTH;
-
- // Fused SiLU + multiply with vectorized stores (8 halfs at a time)
+ // Store C as (M, N) row-major (matches reference layout), vectorized per thread.
+ half* row_ptr = C_ptr + (off_m + tid) * N + off_n + n * WIDTH;
+
#pragma unroll
- for (int i = 0; i < WIDTH; i += 8) {
- // Load and compute 8 elements
- float r0 = acc1[i+0] / (1.0f + __expf(-acc1[i+0])) * acc2[i+0];
- float r1 = acc1[i+1] / (1.0f + __expf(-acc1[i+1])) * acc2[i+1];
- float r2 = acc1[i+2] / (1.0f + __expf(-acc1[i+2])) * acc2[i+2];
- float r3 = acc1[i+3] / (1.0f + __expf(-acc1[i+3])) * acc2[i+3];
- float r4 = acc1[i+4] / (1.0f + __expf(-acc1[i+4])) * acc2[i+4];
- float r5 = acc1[i+5] / (1.0f + __expf(-acc1[i+5])) * acc2[i+5];
- float r6 = acc1[i+6] / (1.0f + __expf(-acc1[i+6])) * acc2[i+6];
- float r7 = acc1[i+7] / (1.0f + __expf(-acc1[i+7])) * acc2[i+7];
-
- // Convert and store as 8 halfs (uint4 = 16 bytes = 8 halfs)
- half2 h0 = __halves2half2(__float2half(r0), __float2half(r1));
- half2 h1 = __halves2half2(__float2half(r2), __float2half(r3));
- half2 h2 = __halves2half2(__float2half(r4), __float2half(r5));
- half2 h3 = __halves2half2(__float2half(r6), __float2half(r7));
- *reinterpret_cast<uint4*>(row_ptr + i) = make_uint4(
- *reinterpret_cast<uint32_t*>(&h0),
- *reinterpret_cast<uint32_t*>(&h1),
- *reinterpret_cast<uint32_t*>(&h2),
- *reinterpret_cast<uint32_t*>(&h3)
- );
+ for (int i = 0; i < WIDTH; i += 16) {
+ float r0 = acc1[i+0] / (1.0f + __expf(-acc1[i+0])) * acc2[i+0];
+ float r1 = acc1[i+1] / (1.0f + __expf(-acc1[i+1])) * acc2[i+1];
+ float r2 = acc1[i+2] / (1.0f + __expf(-acc1[i+2])) * acc2[i+2];
+ float r3 = acc1[i+3] / (1.0f + __expf(-acc1[i+3])) * acc2[i+3];
+ float r4 = acc1[i+4] / (1.0f + __expf(-acc1[i+4])) * acc2[i+4];
+ float r5 = acc1[i+5] / (1.0f + __expf(-acc1[i+5])) * acc2[i+5];
+ float r6 = acc1[i+6] / (1.0f + __expf(-acc1[i+6])) * acc2[i+6];
+ float r7 = acc1[i+7] / (1.0f + __expf(-acc1[i+7])) * acc2[i+7];
+ float r8 = acc1[i+8] / (1.0f + __expf(-acc1[i+8])) * acc2[i+8];
+ float r9 = acc1[i+9] / (1.0f + __expf(-acc1[i+9])) * acc2[i+9];
+ float r10 = acc1[i+10] / (1.0f + __expf(-acc1[i+10])) * acc2[i+10];
+ float r11 = acc1[i+11] / (1.0f + __expf(-acc1[i+11])) * acc2[i+11];
+ float r12 = acc1[i+12] / (1.0f + __expf(-acc1[i+12])) * acc2[i+12];
+ float r13 = acc1[i+13] / (1.0f + __expf(-acc1[i+13])) * acc2[i+13];
+ float r14 = acc1[i+14] / (1.0f + __expf(-acc1[i+14])) * acc2[i+14];
+ float r15 = acc1[i+15] / (1.0f + __expf(-acc1[i+15])) * acc2[i+15];
+
+ half2 h0 = __float22half2_rn({r0, r1});
+ half2 h1 = __float22half2_rn({r2, r3});
+ half2 h2 = __float22half2_rn({r4, r5});
+ half2 h3 = __float22half2_rn({r6, r7});
+ half2 h4 = __float22half2_rn({r8, r9});
+ half2 h5 = __float22half2_rn({r10, r11});
+ half2 h6 = __float22half2_rn({r12, r13});
+ half2 h7 = __float22half2_rn({r14, r15});
+
+ const uint32_t u0 = *reinterpret_cast<uint32_t*>(&h0);
+ const uint32_t u1 = *reinterpret_cast<uint32_t*>(&h1);
+ const uint32_t u2 = *reinterpret_cast<uint32_t*>(&h2);
+ const uint32_t u3 = *reinterpret_cast<uint32_t*>(&h3);
+ const uint32_t u4 = *reinterpret_cast<uint32_t*>(&h4);
+ const uint32_t u5 = *reinterpret_cast<uint32_t*>(&h5);
+ const uint32_t u6 = *reinterpret_cast<uint32_t*>(&h6);
+ const uint32_t u7 = *reinterpret_cast<uint32_t*>(&h7);
+
+ const unsigned long long q0 = (unsigned long long)u0 | ((unsigned long long)u1 << 32);
+ const unsigned long long q1 = (unsigned long long)u2 | ((unsigned long long)u3 << 32);
+ const unsigned long long q2 = (unsigned long long)u4 | ((unsigned long long)u5 << 32);
+ const unsigned long long q3 = (unsigned long long)u6 | ((unsigned long long)u7 << 32);
+
+ stg_32b((const void*)(row_ptr + i), q0, q1, q2, q3);
}
}
+ }
- asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
- if (warp_id == 0)
- asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
- :: "r"(0), "r"(256)); // Must match power-of-2 allocation
+ // Cluster barrier before deallocation (following reference)
+ asm volatile("barrier.cluster.arrive.release.aligned;");
+ asm volatile("barrier.cluster.wait.acquire.aligned;");
+
+ if (warp_id == 0) {
+ asm volatile("tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;"
+ :: "r"(taddr), "r"(TOTAL_TMEM_COLS));
}
}
⋯ 2 unchanged lines
// ============================================================================
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
- at::Tensor dual_gemm_launch(
+ at::Tensor dual_gemm_cta2_launch(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
⋯ 2 unchanged lines
const at::Tensor& SFB2,
at::Tensor& C
) {
+ constexpr int HALF_BLOCK_N = BLOCK_N / 2;
+
const int M = A.size(0);
const int N = B1.size(0);
⋯ 7 unchanged lines
CUtensorMap A_tmap, B1_tmap, B2_tmap;
init_AB_tmap(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K);
- init_AB_tmap(&B1_tmap, B1_ptr, N, K, BLOCK_N, BLOCK_K);
- init_AB_tmap(&B2_tmap, B2_ptr, N, K, BLOCK_N, BLOCK_K);
+ init_AB_tmap(&B1_tmap, B1_ptr, N, K, HALF_BLOCK_N, BLOCK_K);
+ init_AB_tmap(&B2_tmap, B2_ptr, N, K, HALF_BLOCK_N, BLOCK_K);
- dim3 grid((M / BLOCK_M) * (N / BLOCK_N), 1, 1);
- int tb_size = BLOCK_M + 2 * WARP_SIZE;
+ CUtensorMap SFA_tmap, SFB1_tmap, SFB2_tmap;
+ init_SF_tmap(&SFA_tmap, SFA_ptr, M, K, BLOCK_K);
+ init_SF_tmap(&SFB1_tmap, SFB1_ptr, N, K, BLOCK_K);
+ init_SF_tmap(&SFB2_tmap, SFB2_ptr, N, K, BLOCK_K);
+
+ const int num_blocks = (M / BLOCK_M) * (N / BLOCK_N);
+ dim3 grid(num_blocks, 1, 1);
+ int tb_size = BLOCK_M + 3 * WARP_SIZE; // +3 for SF, TMA, MMA warps
- // New SMEM size with dual B tiles
- constexpr int A_size = BLOCK_M * BLOCK_K / 2;
- constexpr int B1_size = BLOCK_N * BLOCK_K / 2;
- constexpr int B2_size = BLOCK_N * BLOCK_K / 2;
- constexpr int SFA_size = 128 * BLOCK_K / 16;
- constexpr int SFB1_size = 128 * BLOCK_K / 16;
- constexpr int SFB2_size = 128 * BLOCK_K / 16;
- int smem_size = (A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size) * NUM_STAGES;
+ constexpr int A_size_c = BLOCK_M * BLOCK_K / 2;
+ constexpr int B1_size_c = HALF_BLOCK_N * BLOCK_K / 2;
+ constexpr int B2_size_c = HALF_BLOCK_N * BLOCK_K / 2;
+ constexpr int SFA_size_c = 128 * BLOCK_K / 16;
+ constexpr int SFB1_size_c = 128 * BLOCK_K / 16;
+ constexpr int SFB2_size_c = 128 * BLOCK_K / 16;
+ int smem_size = (A_size_c + B1_size_c + B2_size_c + SFA_size_c + SFB1_size_c + SFB2_size_c) * NUM_STAGES;
- auto kernel_fn = dual_gemm_fused_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
+ auto kernel_fn = dual_gemm_cta2_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
if (smem_size > 48000)
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
kernel_fn<<<grid, tb_size, smem_size>>>(
- A_tmap, B1_tmap, B2_tmap, SFA_ptr, SFB1_ptr, SFB2_ptr, C_ptr, M, N
+ A_tmap, B1_tmap, B2_tmap, SFA_tmap, SFB1_tmap, SFB2_tmap, C_ptr, M, N
);
return C;
⋯ 9 unchanged lines
at::Tensor& C
) {
const int K = A.size(1) * 2;
+ const int M = A.size(0);
+ const int N = B1.size(0);
#define LAUNCH(K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_) \
- if (K == K_) return dual_gemm_launch<K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_>( \
+ if (K == K_) return dual_gemm_cta2_launch<K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_>( \
A, B1, B2, SFA, SFB1, SFB2, C);
- // Main benchmark configs - 5 stages with BLOCK_N=64
- LAUNCH(7168, 128, 64, 256, 5)
- LAUNCH(4096, 128, 64, 256, 5)
-
- // Other K values
- LAUNCH(256, 128, 64, 256, 1)
- LAUNCH(512, 128, 64, 256, 2)
- LAUNCH(1024, 128, 64, 256, 4)
- LAUNCH(1536, 128, 64, 256, 5)
- LAUNCH(2048, 128, 64, 256, 5)
- LAUNCH(2304, 128, 64, 256, 5)
+ // With MMA_M=256 (cta_group::2), the M=256 case has only 1 cluster in M,
+ // so cluster count ~= N / BLOCK_N. Using BLOCK_N=64 increases cluster count vs 128.
+ //
+ // IMPORTANT: for M=512, reducing BLOCK_N often hurts (fixed overhead per tile dominates),
+ // so we only enable the 64-wide path for M=256.
+ bool use_block_n_64 = (M == 256);
+ // Optional override for benchmarking: set NVFP4_M256_BLOCK_N to 64 or 128.
+ if (M == 256) {
+ if (const char* env = std::getenv("NVFP4_M256_BLOCK_N")) {
+ const int v = std::atoi(env);
+ if (v == 128) use_block_n_64 = false;
+ else if (v == 64) use_block_n_64 = true;
+ }
+ }
+ if (use_block_n_64) {
+ // M=256 underfills with BLOCK_N=128; try BLOCK_N=64 for more clusters in N.
+ LAUNCH(7168, 128, 64, 256, 7)
+ LAUNCH(4096, 128, 64, 256, 7)
+ LAUNCH(3072, 128, 64, 256, 7)
+ LAUNCH(2304, 128, 64, 256, 7)
+ LAUNCH(2048, 128, 64, 256, 7)
+ LAUNCH(1536, 128, 64, 256, 6)
+ LAUNCH(1024, 128, 64, 256, 4)
+ LAUNCH(512, 128, 64, 256, 2)
+ LAUNCH(256, 128, 64, 256, 1)
+ } else {
+ LAUNCH(7168, 128, 128, 256, 5)
+ LAUNCH(4096, 128, 128, 256, 5)
+ LAUNCH(3072, 128, 128, 256, 5)
+ LAUNCH(2304, 128, 128, 256, 5)
+ LAUNCH(2048, 128, 128, 256, 5)
+ LAUNCH(1536, 128, 128, 256, 5)
+ LAUNCH(1024, 128, 128, 256, 4)
+ LAUNCH(512, 128, 128, 256, 2)
+ LAUNCH(256, 128, 128, 256, 1)
+ }
+
#undef LAUNCH
TORCH_CHECK(false, "Unsupported K value: ", K);
}
- TORCH_LIBRARY(dual_gemm_module, m) {
+ TORCH_LIBRARY(dual_gemm_cta2_v6_module, 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);
}
⋯ 5 unchanged lines
global _compiled_module
if _compiled_module is None:
_compiled_module = load_inline(
- "dual_gemm_cuda",
+ "dual_gemm_cta2_v6_cuda",
cpp_sources="",
cuda_sources=CUDA_SOURCE,
verbose=True,
⋯ 14 unchanged lines
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
_get_module()
- result = torch.ops.dual_gemm_module.dual_gemm(
+ result = torch.ops.dual_gemm_cta2_v6_module.dual_gemm(
a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c
)
return result
+
scrolls · 887 diff lines total

Best evidence level for this revision: reported

JSON