Skip to content
KernelIndex
Search⌘K

submission 489022

kathsucurry · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
61.1µs
#212 of 310
2026-02-12

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f1348f6c1d4aef997ccaebacebd1edb6ba8bf5239e6580416ff3b0428a76885c
license declaredunknown
license concludedunknown
authorskathsucurry
imported2026-08-15

Techniques

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

fp4constexpr int MMA_K = 64; // FP4 MMA K-dimension size.
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
shared-memoryextern __shared__ __align__(1024) char smem[];
stages = 4constexpr int NUM_STAGES = 4;
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tmaTORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);

Kernel source

submission.py501 lines
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

cuda_source = r"""
#include <cuda_fp16.h>
#include <cudaTypedefs.h>

#include <torch/extension.h>
#include <torch/library.h>


#define WARP_SIZE 32


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


template <const int NUM_ELEMENTS>
inline void create_tmap_descriptor(
    CUtensorMap *tmap,
    const char *ptr,
    uint64_t global_height, uint64_t global_width,
    uint32_t shared_height, uint32_t shared_width,
    CUtensorMapSwizzle swizzle_type
) {
    /*
    The goal is to transfer multiple of [shared_height, NUM_ELEMENTS] spanning
    [shared_height, shared_width] --> [shared_width / NUM_ELEMENTS, shared_height, NUM_ELEMENTS].

    Code taken and modified from:
    - https://docs.nvidia.com/cuda/cuda-programming-guide/04-special-topics/async-copies.html#using-tma-to-transfer-multi-dimensional-arrays.
    - https://gau-nernst.github.io/tcgen05/ 
    */
    constexpr int rank{3};
    uint64_t global_dim[rank] = {NUM_ELEMENTS, global_height, global_width / (uint64_t) NUM_ELEMENTS};
    // 4 bits would be 1/2 bytes.
    uint64_t global_strides[rank - 1] = {global_width / 2, NUM_ELEMENTS / 2}; 
    uint32_t box_dim[rank] = {NUM_ELEMENTS, shared_height, shared_width / NUM_ELEMENTS};
    uint32_t element_strides[rank] = {1, 1, 1};

    auto error = cuTensorMapEncodeTiled(
        tmap,
        CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        (void *)ptr,
        global_dim,
        global_strides,
        box_dim,
        element_strides,
        // Interleave patterns can be used to accelerate loading of values that
        // are less than 4 bytes long.
        CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
        swizzle_type,
        // L2 Promotion can be used to widen the effect of a cache-policy to a wider
        // set of L2 cache lines.
        CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
        // Any element that is outside of bounds will be set to zero by the TMA transfer.
        CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu_error(error);
}


// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__device__ inline 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__ inline void mbarrier_init(int mbar_addr, int count) {
    asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" ::"r"(mbar_addr), "r"(count));
}


// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
__device__ inline void mbarrier_wait(int mbar_addr, int phase) {
    uint32_t ticks = 0x989680; // arbitrarily large timer value.
    asm volatile(
        "{\n\t"
        ".reg .pred P1;\n\t"
        "LAB_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
        "@P1 bra.uni DONE;\n\t"
        "bra.uni LAB_WAIT;\n\t"
        "DONE:\n\t"
        "}" ::"r"(mbar_addr),
        "r"(phase), "r"(ticks));
}


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


template <int CTA_GROUP = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr)
{
    // when CTA_GROUP=1, we can use .shared::cta instead.
    // but .shared::cluster doesn't seem to be slower, so always use it unconditionally here.
    // .cta_group::2 allows mbar_addr and dst to be in different CTA's smem.
    asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "
                 "[%0], [%1, {%2, %3, %4}], [%5];" ::"r"(dst),
                 "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(CTA_GROUP)
                 : "memory");
}


// Encodes the matrix descriptor and ensures 64 bits.
__device__ inline
constexpr uint64_t encode_descriptor(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }


// Copy scale factors from shared memory to tensor memory.
// .32x128b = 32 rows x 16 bytes = one scale factor tile for one MMA.
// .warpx4 duplicates data across all 32-lane groups.
__device__ inline
void copy_sf_smem2tmem(int taddr, uint64_t s_desc) {
    asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}

// Issue FP4 MMA instruction with block scaling.
// d_tmem=0: accumulator always starts at TMEM column 0.
// enable_input_d: 0 = clear accumulator, nonzero = accumulate.
__device__ inline
void run_mma_nvfp4(
    uint64_t a_desc,
    uint64_t b_desc,
    uint32_t i_desc,
    int scale_A_tmem,
    int scale_B_tmem,
    int enable_input_d
) {
    const int d_tmem = 0;
    asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
    );
}

constexpr int MMA_K = 64;  // FP4 MMA K-dimension size.
constexpr int NUM_STAGES = 4;


template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>
__global__
__launch_bounds__(NUM_THREADS) void kernel_v05_warp_spec(
    const __grid_constant__ CUtensorMap A_tmap,
    const __grid_constant__ CUtensorMap B_tmap,
    const char *SFA,
    const char *SFB,
    half *C,
    int M,
    int N,
    int K
) {
    const int thread_idx{static_cast<int>(threadIdx.x)};
    const int block_idx{static_cast<int>(blockIdx.x)};

    const int warp_idx{thread_idx / WARP_SIZE};

    const int grid_dim_n{(N + BLOCK_N - 1) / BLOCK_N};

    const int block_idx_m{block_idx / grid_dim_n};
    const int block_idx_n{block_idx % grid_dim_n};

    const int offset_m{block_idx_m * BLOCK_M};
    const int offset_n{block_idx_n * BLOCK_N};

    // Multi-buffered shared memory layout:
    // [buf0: A | B | SFA | SFB | buf1: ... | buf2: ... | buf3: ...]
    constexpr int SF_size = 512 * BLOCK_K / MMA_K;
    constexpr int BUF_SIZE = BLOCK_M * BLOCK_K / 2 + BLOCK_N * BLOCK_K / 2 + 2 * SF_size;

    extern __shared__ __align__(1024) char smem[];
    const int smem_base{static_cast<int>(__cvta_generic_to_shared(smem))};

    int A_smem[NUM_STAGES], B_smem[NUM_STAGES], SFA_smem[NUM_STAGES], SFB_smem[NUM_STAGES];
    for (int s{0}; s < NUM_STAGES; ++s) {
        A_smem[s] = smem_base + s * BUF_SIZE;
        B_smem[s] = A_smem[s] + BLOCK_M * BLOCK_K / 2;
        SFA_smem[s] = B_smem[s] + BLOCK_N * BLOCK_K / 2;
        SFB_smem[s] = SFA_smem[s] + SF_size;
    }

#pragma nv_diag_suppress static_var_with_dynamic_init
    // Two sets of mbarriers:
    // tma_mbars: producer (warp 0) signals when TMA into a buffer is done; consumer (warp 1) waits.
    // mma_mbars: consumer signals when MMA from a buffer is done (via tcgen05.commit); producer waits
    //            before reusing that buffer.
    __shared__ uint64_t tma_mbars[NUM_STAGES];
    __shared__ uint64_t mma_mbars[NUM_STAGES];
    int tma_mbar_addrs[NUM_STAGES], mma_mbar_addrs[NUM_STAGES];
    for (int s{0}; s < NUM_STAGES; ++s) {
        tma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&tma_mbars[s]));
        mma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mma_mbars[s]));
    }
    __shared__ int tmem_addr[1];

    // TMEM layout:
    // Columns [0, BLOCK_N)                          : accumulator D
    // Columns [BLOCK_N, BLOCK_N + 4*BLOCK_K/MMA_K)  : SFA
    // Columns [BLOCK_N + 4*BLOCK_K/MMA_K, ...)       : SFB
    constexpr int SFA_tmem_start_col = BLOCK_N;
    constexpr int SFB_tmem_start_col = SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K);
    constexpr int TMEM_COLS = BLOCK_N * 2;

    if (warp_idx == 0 && elect_sync()) {
        for (int s{0}; s < NUM_STAGES; ++s) {
            mbarrier_init(tma_mbar_addrs[s], 1);
            mbarrier_init(mma_mbar_addrs[s], 1);
        }
        asm volatile("fence.mbarrier_init.release.cluster;");
    } else if (warp_idx == 1) {
        // Allocate TMEM for accumulator + scale factors.
        const int addr{static_cast<int>(__cvta_generic_to_shared(tmem_addr))};
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
            ::"r"(addr), "r"(TMEM_COLS));
    }
    __syncthreads();

    const int taddr{tmem_addr[0]};

    // Instruction descriptor for tcgen05.mma.kind::mxf4nvf4
    constexpr uint32_t i_desc = (1U << 7U)                         // atype=E2M1
                                | (1U << 10U)                      // btype=E2M1
                                | ((uint32_t)BLOCK_N >> 3U << 17U) // MMA_N
                                | ((uint32_t)BLOCK_M >> 7U << 27U) // MMA_M
                                ;

    const int num_iters{K / BLOCK_K};
    const int rest_k{K / 16 / 4};
    constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;

    // =========================================================================
    // Warp 0: TMA Producer
    // =========================================================================
    if (warp_idx == 0) {
        // Phase tracking for mma_mbars (producer waits on these before reusing a buffer).
        int mma_prod_phase[NUM_STAGES] = {};

        for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
            const int s{iter_k % NUM_STAGES};

            // Wait for consumer to finish MMA from buf[s] before overwriting it.
            // First NUM_STAGES iterations don't need to wait (buffers haven't been used yet).
            if (iter_k >= NUM_STAGES) {
                mbarrier_wait(mma_mbar_addrs[s], mma_prod_phase[s]);
                mma_prod_phase[s] ^= 1;
            }

            // Issue TMA into buf[s].
            if (elect_sync()) {
                const int off_k{iter_k * BLOCK_K};
                tma_3d_gmem2smem(A_smem[s], &A_tmap, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);
                tma_3d_gmem2smem(B_smem[s], &B_tmap, 0, offset_n, off_k / 256, tma_mbar_addrs[s]);
                const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;
                const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;
                tma_gmem2smem(SFA_smem[s], SFA_src, SF_size, tma_mbar_addrs[s]);
                tma_gmem2smem(SFB_smem[s], SFB_src, SF_size, tma_mbar_addrs[s]);
                asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                    ::"r"(tma_mbar_addrs[s]), "r"(cp_size) : "memory");
            }
        }

    // =========================================================================
    // Warp 1: MMA Consumer
    // =========================================================================
    } else if (warp_idx == 1) {
        // Phase tracking for tma_mbars (consumer waits on these for data).
        int tma_cons_phase[NUM_STAGES] = {};
        // Track mma_mbar phases for the final wait.
        int mma_done_phase[NUM_STAGES] = {};

        for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
            const int s{iter_k % NUM_STAGES};

            // Wait for producer to fill buf[s].
            mbarrier_wait(tma_mbar_addrs[s], tma_cons_phase[s]);
            tma_cons_phase[s] ^= 1;

            // Issue SF copy + MMA from buf[s].
            if (elect_sync()) {
                auto make_desc_AB = [](int addr) -> uint64_t
                {
                    constexpr int SBO = 8 * 256 / 2;
                    return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
                };

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

                // Copy scale factors from shared memory to tensor memory.
                for (int k{0}; k < BLOCK_K / MMA_K; ++k) {
                    uint64_t sfa_desc = make_desc_SF(SFA_smem[s] + k * 512);
                    uint64_t sfb_desc = make_desc_SF(SFB_smem[s] + k * 512);
                    copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);
                    copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);
                }

                // Issue MMA instructions. One swizzle tile = 256 FP4 elements = 128 bytes.
                for (int k1{0}; k1 < BLOCK_K / 256; ++k1) {
                    for (int k2{0}; k2 < 256 / MMA_K; ++k2) {
                        uint64_t a_desc{make_desc_AB(A_smem[s] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
                        uint64_t b_desc{make_desc_AB(B_smem[s] + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};

                        int k{k1 * 256 / MMA_K + k2};
                        const int scale_A_tmem{SFA_tmem_start_col + k * 4};
                        const int scale_B_tmem{SFB_tmem_start_col + k * 4};

                        const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
                        run_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
                    }
                }

                // Signal that MMA is done reading from buf[s].
                // tcgen05.commit waits for all prior tcgen05 ops to finish, then arrives on the mbarrier.
                asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                    ::"r"(mma_mbar_addrs[s]) : "memory");
            }
            mma_done_phase[s] ^= 1;
        }

        // Wait for the last MMA to complete before the epilogue reads the accumulator.
        if (num_iters > 0) {
            const int last_s{(num_iters - 1) % NUM_STAGES};
            mbarrier_wait(mma_mbar_addrs[last_s], mma_done_phase[last_s] ^ 1);
        }
    }
    // Warps 2-3 skip directly here.

    // Synchronize all warps before the epilogue.
    __syncthreads();

    // === Epilogue: Read accumulator from TMEM and store to global memory ===
    asm volatile("tcgen05.fence::after_thread_sync;");

    // Each thread handles one row. 4 warps * 32 threads = 128 rows = BLOCK_M.
    const int row{offset_m + thread_idx};
    for (int n{0}; n < BLOCK_N / 8; ++n) {
        float tmp[8];
        const int addr = taddr + ((warp_idx * 32) << 16) + (n * 8);
        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));
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        if (row >= M) continue;

        const int col{offset_n + n * 8};

        half2 out[4];
        for (int i{0}; i < 4; ++i)
            out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});

        half *out_ptr = C + row * N + col;
        if (col + 8 <= N) {
            reinterpret_cast<int4 *>(out_ptr)[0] = reinterpret_cast<int4 *>(out)[0];
        } else {
            const half *out_half = reinterpret_cast<const half *>(out);
            for (int i{0}; i < 8 && col + i < N; ++i)
                out_ptr[i] = out_half[i];
        }
    }
    __syncthreads();

    if (warp_idx == 0) {
        asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(TMEM_COLS));
    }
}


torch::Tensor launch_kernel_warp_spec_05(
    const torch::Tensor& A,
    const torch::Tensor& B,
    const torch::Tensor& sfa,
    const torch::Tensor& sfb,
    int M,
    int N,
    int K
) {
    auto C = torch::empty({M, N}, torch::dtype(torch::kFloat16).device(A.device()));

    constexpr int BLOCK_M{128};
    constexpr int BLOCK_N{128};
    constexpr int BLOCK_K{256};
    constexpr int NUM_THREADS{4 * WARP_SIZE};

    auto A_ptr{reinterpret_cast<const char *>(A.data_ptr())};
    auto B_ptr{reinterpret_cast<const char *>(B.data_ptr())};
    auto SFA_ptr{reinterpret_cast<const char *>(sfa.data_ptr())};
    auto SFB_ptr{reinterpret_cast<const char *>(sfb.data_ptr())};

    CUtensorMap A_tmap{}, B_tmap{};
    create_tmap_descriptor<256>(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
    create_tmap_descriptor<256>(&B_tmap, B_ptr, N, K, BLOCK_N, BLOCK_K,
        CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);

    constexpr int AB_SHARED_SIZE{(BLOCK_M + BLOCK_N) * BLOCK_K / 2};
    constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};
    constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};

    dim3 num_threads(NUM_THREADS);
    dim3 num_blocks(((M + BLOCK_M - 1) / BLOCK_M) * ((N + BLOCK_N - 1) / BLOCK_N));

    auto kernel{kernel_v05_warp_spec<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};

    if (SHARED_SIZE > 48'000)
        cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);

    kernel<<<num_blocks, num_threads, SHARED_SIZE>>>(
        A_tmap,
        B_tmap,
        SFA_ptr,
        SFB_ptr,
        reinterpret_cast<half *>(C.data_ptr<at::Half>()),
        M, N, K
    );
    return C;
}


"""

cpp_source = """
#include <torch/extension.h>

torch::Tensor launch_kernel_warp_spec_05(
    const torch::Tensor& A,
    const torch::Tensor& B,
    const torch::Tensor& sfa,
    const torch::Tensor& sfb,
    int M,
    int N,
    int K);
"""

module = load_inline(
    name='kernel',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['launch_kernel_warp_spec_05'],
    verbose=True,
    is_python_module=True,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
        # "--keep",
        # "--keep-dir",
        # f"{Path(__file__).parent}/tmp",
    ],
    extra_ldflags=["-lcuda"],
)


def custom_kernel(data: input_t) -> output_t:
    results = []
    for (a, b, c), _, (sfa_reordered, sfb_reordered), (m, n, k, _) in zip(*data):
        c[:, :, 0] = module.launch_kernel_warp_spec_05(
            a[:, :, 0], b[:, :, 0], sfa_reordered, sfb_reordered, m, n, k
        )

        results.append(c)

    return results
scrolls · 501 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 489003.

⋯ 167 unchanged lines
template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>
__global__
- __launch_bounds__(NUM_THREADS) void kernel_v04_pipeline(
+ __launch_bounds__(NUM_THREADS) void kernel_v05_warp_spec(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B_tmap,
const char *SFA,
⋯ 16 unchanged lines
const int offset_m{block_idx_m * BLOCK_M};
const int offset_n{block_idx_n * BLOCK_N};
- // Double-buffered shared memory layout:
- // [buf0: A | B | SFA | SFB | buf1: A | B | SFA | SFB]
+ // Multi-buffered shared memory layout:
+ // [buf0: A | B | SFA | SFB | buf1: ... | buf2: ... | buf3: ...]
constexpr int SF_size = 512 * BLOCK_K / MMA_K;
constexpr int BUF_SIZE = BLOCK_M * BLOCK_K / 2 + BLOCK_N * BLOCK_K / 2 + 2 * SF_size;
⋯ 9 unchanged lines
}
#pragma nv_diag_suppress static_var_with_dynamic_init
- __shared__ uint64_t mbars[NUM_STAGES];
- int mbar_addrs[NUM_STAGES];
- for (int s{0}; s < NUM_STAGES; ++s)
- mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mbars[s]));
+ // Two sets of mbarriers:
+ // tma_mbars: producer (warp 0) signals when TMA into a buffer is done; consumer (warp 1) waits.
+ // mma_mbars: consumer signals when MMA from a buffer is done (via tcgen05.commit); producer waits
+ // before reusing that buffer.
+ __shared__ uint64_t tma_mbars[NUM_STAGES];
+ __shared__ uint64_t mma_mbars[NUM_STAGES];
+ int tma_mbar_addrs[NUM_STAGES], mma_mbar_addrs[NUM_STAGES];
+ for (int s{0}; s < NUM_STAGES; ++s) {
+ tma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&tma_mbars[s]));
+ mma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mma_mbars[s]));
+ }
__shared__ int tmem_addr[1];
// TMEM layout:
⋯ 2 unchanged lines
// Columns [BLOCK_N + 4*BLOCK_K/MMA_K, ...) : SFB
constexpr int SFA_tmem_start_col = BLOCK_N;
constexpr int SFB_tmem_start_col = SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K);
- constexpr int TMEM_COLS = BLOCK_N * 2; // BLOCK_N + 8 * (BLOCK_K / MMA_K), but it has to be a power of 2.
+ constexpr int TMEM_COLS = BLOCK_N * 2;
if (warp_idx == 0 && elect_sync()) {
- for (int s{0}; s < NUM_STAGES; ++s)
- mbarrier_init(mbar_addrs[s], 1);
+ for (int s{0}; s < NUM_STAGES; ++s) {
+ mbarrier_init(tma_mbar_addrs[s], 1);
+ mbarrier_init(mma_mbar_addrs[s], 1);
+ }
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_idx == 1) {
// Allocate TMEM for accumulator + scale factors.
⋯ 4 unchanged lines
__syncthreads();
const int taddr{tmem_addr[0]};
- int phase[NUM_STAGES] = {0, 0};
// Instruction descriptor for tcgen05.mma.kind::mxf4nvf4
- // atype=E2M1 (1), btype=E2M1 (1), MMA_N and MMA_M encoded in upper bits
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)BLOCK_N >> 3U << 17U) // MMA_N
⋯ 1 unchanged lines
;
const int num_iters{K / BLOCK_K};
- const int rest_k{K / 16 / 4}; // number of K-atoms (each covers 64 K-elements)
+ const int rest_k{K / 16 / 4};
constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;
- // === Prologue: load first tile into buf[0] ===
- if (warp_idx == 0 && elect_sync()) {
- const int off_k{0};
- tma_3d_gmem2smem(A_smem[0], &A_tmap, 0, offset_m, off_k / 256, mbar_addrs[0]);
- tma_3d_gmem2smem(B_smem[0], &B_tmap, 0, offset_n, off_k / 256, mbar_addrs[0]);
- const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;
- const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;
- tma_gmem2smem(SFA_smem[0], SFA_src, SF_size, mbar_addrs[0]);
- tma_gmem2smem(SFB_smem[0], SFB_src, SF_size, mbar_addrs[0]);
- asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
- ::"r"(mbar_addrs[0]), "r"(cp_size) : "memory");
- }
+ // =========================================================================
+ // Warp 0: TMA Producer
+ // =========================================================================
+ if (warp_idx == 0) {
+ // Phase tracking for mma_mbars (producer waits on these before reusing a buffer).
+ int mma_prod_phase[NUM_STAGES] = {};
- // Wait for first TMA to complete.
- mbarrier_wait(mbar_addrs[0], phase[0]);
- phase[0] ^= 1;
+ for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
+ const int s{iter_k % NUM_STAGES};
- // === Main loop: overlap TMA load of next tile with MMA compute of current tile ===
- for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
- const int s{iter_k % NUM_STAGES};
+ // Wait for consumer to finish MMA from buf[s] before overwriting it.
+ // First NUM_STAGES iterations don't need to wait (buffers haven't been used yet).
+ if (iter_k >= NUM_STAGES) {
+ mbarrier_wait(mma_mbar_addrs[s], mma_prod_phase[s]);
+ mma_prod_phase[s] ^= 1;
+ }
- if (warp_idx == 0 && elect_sync()) {
- // Prefetch next tile into the other buffer (overlaps with MMA below).
- if (iter_k + 1 < num_iters) {
- const int ns{(iter_k + 1) % NUM_STAGES};
- const int next_off_k{(iter_k + 1) * BLOCK_K};
- tma_3d_gmem2smem(A_smem[ns], &A_tmap, 0, offset_m, next_off_k / 256, mbar_addrs[ns]);
- tma_3d_gmem2smem(B_smem[ns], &B_tmap, 0, offset_n, next_off_k / 256, mbar_addrs[ns]);
- const char *SFA_src = SFA + ((offset_m / 128) * rest_k + next_off_k / (16 * 4)) * 512;
- const char *SFB_src = SFB + ((offset_n / 128) * rest_k + next_off_k / (16 * 4)) * 512;
- tma_gmem2smem(SFA_smem[ns], SFA_src, SF_size, mbar_addrs[ns]);
- tma_gmem2smem(SFB_smem[ns], SFB_src, SF_size, mbar_addrs[ns]);
+ // Issue TMA into buf[s].
+ if (elect_sync()) {
+ const int off_k{iter_k * BLOCK_K};
+ tma_3d_gmem2smem(A_smem[s], &A_tmap, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);
+ tma_3d_gmem2smem(B_smem[s], &B_tmap, 0, offset_n, off_k / 256, tma_mbar_addrs[s]);
+ const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;
+ const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;
+ tma_gmem2smem(SFA_smem[s], SFA_src, SF_size, tma_mbar_addrs[s]);
+ tma_gmem2smem(SFB_smem[s], SFB_src, SF_size, tma_mbar_addrs[s]);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
- ::"r"(mbar_addrs[ns]), "r"(cp_size) : "memory");
+ ::"r"(tma_mbar_addrs[s]), "r"(cp_size) : "memory");
}
+ }
- // MMA from current buffer.
- auto make_desc_AB = [](int addr) -> uint64_t
- {
- constexpr int SBO = 8 * 256 / 2; // in bytes.
- return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
- };
+ // =========================================================================
+ // Warp 1: MMA Consumer
+ // =========================================================================
+ } else if (warp_idx == 1) {
+ // Phase tracking for tma_mbars (consumer waits on these for data).
+ int tma_cons_phase[NUM_STAGES] = {};
+ // Track mma_mbar phases for the final wait.
+ int mma_done_phase[NUM_STAGES] = {};
- auto make_desc_SF = [](int addr) -> uint64_t
- {
- const int SBO = 8 * 16; // = 128 bytes
- return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);
- };
+ for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
+ const int s{iter_k % NUM_STAGES};
- // Copy scale factors from shared memory to tensor memory.
- for (int k{0}; k < BLOCK_K / MMA_K; ++k) {
- uint64_t sfa_desc = make_desc_SF(SFA_smem[s] + k * 512);
- uint64_t sfb_desc = make_desc_SF(SFB_smem[s] + k * 512);
- copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);
- copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);
- }
+ // Wait for producer to fill buf[s].
+ mbarrier_wait(tma_mbar_addrs[s], tma_cons_phase[s]);
+ tma_cons_phase[s] ^= 1;
- // One swizzle tile = 256 FP4 elements = 128 bytes.
- for (int k1{0}; k1 < BLOCK_K / 256; ++k1) {
- for (int k2{0}; k2 < 256 / MMA_K; ++k2) {
- uint64_t a_desc{make_desc_AB(A_smem[s] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
- uint64_t b_desc{make_desc_AB(B_smem[s] + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};
+ // Issue SF copy + MMA from buf[s].
+ if (elect_sync()) {
+ auto make_desc_AB = [](int addr) -> uint64_t
+ {
+ constexpr int SBO = 8 * 256 / 2;
+ return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
+ };
- int k{k1 * 256 / MMA_K + k2};
- const int scale_A_tmem{SFA_tmem_start_col + k * 4};
- const int scale_B_tmem{SFB_tmem_start_col + k * 4};
+ auto make_desc_SF = [](int addr) -> uint64_t
+ {
+ const int SBO = 8 * 16;
+ return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);
+ };
- const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
- run_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
+ // Copy scale factors from shared memory to tensor memory.
+ for (int k{0}; k < BLOCK_K / MMA_K; ++k) {
+ uint64_t sfa_desc = make_desc_SF(SFA_smem[s] + k * 512);
+ uint64_t sfb_desc = make_desc_SF(SFB_smem[s] + k * 512);
+ copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);
+ copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);
}
- }
- // Signal MMA completion on the current buffer's mbarrier.
- asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
- ::"r"(mbar_addrs[s]) : "memory");
- }
+ // Issue MMA instructions. One swizzle tile = 256 FP4 elements = 128 bytes.
+ for (int k1{0}; k1 < BLOCK_K / 256; ++k1) {
+ for (int k2{0}; k2 < 256 / MMA_K; ++k2) {
+ uint64_t a_desc{make_desc_AB(A_smem[s] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
+ uint64_t b_desc{make_desc_AB(B_smem[s] + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};
- // Wait for MMA on current buffer to complete.
- // This also ensures buf[s] is free to reuse for TMA 2 iterations later.
- mbarrier_wait(mbar_addrs[s], phase[s]);
- phase[s] ^= 1;
+ int k{k1 * 256 / MMA_K + k2};
+ const int scale_A_tmem{SFA_tmem_start_col + k * 4};
+ const int scale_B_tmem{SFB_tmem_start_col + k * 4};
- // Wait for next tile's TMA to complete (so next iteration can MMA from it).
- if (iter_k + 1 < num_iters) {
- const int ns{(iter_k + 1) % NUM_STAGES};
- mbarrier_wait(mbar_addrs[ns], phase[ns]);
- phase[ns] ^= 1;
+ const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
+ run_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
+ }
+ }
+
+ // Signal that MMA is done reading from buf[s].
+ // tcgen05.commit waits for all prior tcgen05 ops to finish, then arrives on the mbarrier.
+ asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
+ ::"r"(mma_mbar_addrs[s]) : "memory");
+ }
+ mma_done_phase[s] ^= 1;
}
+
+ // Wait for the last MMA to complete before the epilogue reads the accumulator.
+ if (num_iters > 0) {
+ const int last_s{(num_iters - 1) % NUM_STAGES};
+ mbarrier_wait(mma_mbar_addrs[last_s], mma_done_phase[last_s] ^ 1);
+ }
}
+ // Warps 2-3 skip directly here.
+ // Synchronize all warps before the epilogue.
+ __syncthreads();
+
// === Epilogue: Read accumulator from TMEM and store to global memory ===
- // PTX docs require this fence before tcgen05.ld, after tcgen05.mma.
asm volatile("tcgen05.fence::after_thread_sync;");
// Each thread handles one row. 4 warps * 32 threads = 128 rows = BLOCK_M.
- // Load 8 columns at a time from TMEM.
const int row{offset_m + thread_idx};
for (int n{0}; n < BLOCK_N / 8; ++n) {
float tmp[8];
- // TMEM address: 16 MSBs = row offset, 16 LSBs = column offset.
const int addr = taddr + ((warp_idx * 32) << 16) + (n * 8);
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]),
⋯ 5 unchanged lines
const int col{offset_n + n * 8};
- // Convert f32 pairs to f16 pairs and write to global memory.
half2 out[4];
for (int i{0}; i < 4; ++i)
out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});
- // Each thread writes 16 bytes (8 half values) to its row.
half *out_ptr = C + row * N + col;
if (col + 8 <= N) {
reinterpret_cast<int4 *>(out_ptr)[0] = reinterpret_cast<int4 *>(out)[0];
} else {
- // Partial column write for the last tile.
const half *out_half = reinterpret_cast<const half *>(out);
for (int i{0}; i < 8 && col + i < N; ++i)
out_ptr[i] = out_half[i];
⋯ 2 unchanged lines
__syncthreads();
if (warp_idx == 0) {
- // Deallocate TMEM (accumulator + scale factor columns).
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(TMEM_COLS));
}
}
- torch::Tensor launch_kernel_pipeline_04(
+ torch::Tensor launch_kernel_warp_spec_05(
const torch::Tensor& A,
const torch::Tensor& B,
const torch::Tensor& sfa,
⋯ 4 unchanged lines
) {
auto C = torch::empty({M, N}, torch::dtype(torch::kFloat16).device(A.device()));
- // Tile sizes.
- // BLOCK_M=128; 1 CTA for .kind::mxf4nvf4.
- // BLOCK_N=128; ensure one SF atom covers all N rows in the tile.
constexpr int BLOCK_M{128};
constexpr int BLOCK_N{128};
constexpr int BLOCK_K{256};
- constexpr int NUM_THREADS{4 * WARP_SIZE}; // 4 warps, 128 threads
+ constexpr int NUM_THREADS{4 * WARP_SIZE};
auto A_ptr{reinterpret_cast<const char *>(A.data_ptr())};
auto B_ptr{reinterpret_cast<const char *>(B.data_ptr())};
auto SFA_ptr{reinterpret_cast<const char *>(sfa.data_ptr())};
auto SFB_ptr{reinterpret_cast<const char *>(sfb.data_ptr())};
- // Create 3D TMA tensor maps for A and B.
CUtensorMap A_tmap{}, B_tmap{};
create_tmap_descriptor<256>(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
create_tmap_descriptor<256>(&B_tmap, B_ptr, N, K, BLOCK_N, BLOCK_K,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
- // Shared memory: 2x (A tile + B tile + SFA tile + SFB tile) for double buffering.
constexpr int AB_SHARED_SIZE{(BLOCK_M + BLOCK_N) * BLOCK_K / 2};
constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};
constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};
⋯ 1 unchanged lines
dim3 num_threads(NUM_THREADS);
dim3 num_blocks(((M + BLOCK_M - 1) / BLOCK_M) * ((N + BLOCK_N - 1) / BLOCK_N));
- auto kernel{kernel_v04_pipeline<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};
+ auto kernel{kernel_v05_warp_spec<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};
if (SHARED_SIZE > 48'000)
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);
⋯ 9 unchanged lines
return C;
}
+
"""
cpp_source = """
#include <torch/extension.h>
- torch::Tensor launch_kernel_pipeline_04(
+ torch::Tensor launch_kernel_warp_spec_05(
const torch::Tensor& A,
const torch::Tensor& B,
const torch::Tensor& sfa,
⋯ 7 unchanged lines
name='kernel',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
- functions=['launch_kernel_pipeline_04'],
+ functions=['launch_kernel_warp_spec_05'],
verbose=True,
is_python_module=True,
no_implicit_headers=True,
⋯ 14 unchanged lines
def custom_kernel(data: input_t) -> output_t:
- abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
-
- BLOCK_M = 128
- BLOCK_N = 128
-
results = []
- for (a, b, c), (sfa_reordered, sfb_reordered), (m, n, k, l) in zip(
- abc_tensors, sfasfb_reordered_tensors, problem_sizes
- ):
- for l_idx in range(l):
- # a_ptr = a[:, :, l_idx] # [m, k//2]
- # b_ptr = b[:, :, l_idx] # [n, k//2]
+ for (a, b, c), _, (sfa_reordered, sfb_reordered), (m, n, k, _) in zip(*data):
+ c[:, :, 0] = module.launch_kernel_warp_spec_05(
+ a[:, :, 0], b[:, :, 0], sfa_reordered, sfb_reordered, m, n, k
+ )
- # # Pad M and N to multiples of BLOCK_M and BLOCK_N.
- # padded_m = ((m + BLOCK_M - 1) // BLOCK_M) * BLOCK_M
- # padded_n = ((n + BLOCK_N - 1) // BLOCK_N) * BLOCK_N
-
- # if padded_m != m:
- # a_padded = torch.zeros(padded_m, k // 2, dtype=torch.uint8, device=a_ptr.device).view(a_ptr.dtype)
- # a_padded[:m] = a_ptr
- # a_ptr = a_padded
-
- # if padded_n != n:
- # b_padded = torch.zeros(padded_n, k // 2, dtype=torch.uint8, device=b_ptr.device).view(b_ptr.dtype)
- # b_padded[:n] = b_ptr
- # b_ptr = b_padded
-
- c_out = module.launch_kernel_pipeline_04(
- a[:, :, l_idx], b[:, :, l_idx], sfa_reordered, sfb_reordered, m, n, k
- )
-
- # Trim back to original size.
- c[:, :, l_idx] = c_out[:m, :n]
-
results.append(c)
return results
scrolls · 406 diff lines total

Best evidence level for this revision: reported

JSON