Skip to content
KernelIndex
Search⌘K

submission 499797

kathsucurry · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-499797?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
37.8µs
#39 of 145
2026-02-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:86cb1594fd1c2c8f1cc19ff02844dbb6275afad10da1681ff50cfd2bb7a0acf1
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.
fused-epilogueconstexpr int NUM_STAGES_EPILOGUE{2};
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
persistent-kernel__launch_bounds__(NUM_THREADS) void kernel_v10_persistent(
shared-memoryvoid tma_gmem2smem_cluster(int dst, const void *src, int size, int mbar_addr) {
tcgen05asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(NUM_CTA));
tmaTORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
vector-width = half2half2 out[WIDTH / 2];

Kernel source

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

cuda_common_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 zero-filled 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));
}


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


template <int NUM_CTA = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr)
{
    // when NUM_CTA=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"(NUM_CTA)
                 : "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.
template <int NUM_CTA=1>
__device__ inline
void copy_sf_smem2tmem(int taddr, uint64_t s_desc) {
    asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(NUM_CTA));
}

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

template <int NUM_CTA=1>
__device__ inline
void run_mma_nvfp4_custom_addr(
    const int d_tmem,
    uint64_t a_desc,
    uint64_t b_desc,
    uint32_t i_desc,
    int scale_A_tmem,
    int scale_B_tmem,
    int enable_input_d
) {
    asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %6, 0;\n\t"
        "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
        "}"
        :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
           "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d), "n"(NUM_CTA)
    );
}

__device__ inline
void load_from_tmem_32x32b_x8(float *tmp, const 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));
}


__device__ inline
void load_from_tmem_32x32b_x16(float *tmp, const int addr) {
    asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
                 "{%0, %1,  %2,  %3,  %4,  %5,  %6,  %7, "
                 " %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"
                    : "=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])
                    : "r"(addr));
}

__device__ inline
void load_from_tmem_32x32b_x32(float *tmp, const 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__ inline
void load_from_tmem_32x32b_x64(float *tmp, const 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));
}


__device__ inline
void load_from_tmem_32x32b_x128(float *tmp, const int addr) {
    asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.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,  %65,  %66,  %67,  %68,  %69,  %70,  %71, "
                 "  %72,  %73,  %74,  %75,  %76,  %77,  %78,  %79, "
                 "  %80,  %81,  %82,  %83,  %84,  %85,  %86,  %87, "
                 "  %88,  %89,  %90,  %91,  %92,  %93,  %94,  %95, "
                 "  %96,  %97,  %98,  %99, %100, %101, %102, %103, "
                 " %104, %105, %106, %107, %108, %109, %110, %111, "
                 " %112, %113, %114, %115, %116, %117, %118, %119, "
                 " %120, %121, %122, %123, %124, %125, %126, %127}, [%128];"
                    : "=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]),
                      "=f"(tmp[ 64]), "=f"(tmp[ 65]), "=f"(tmp[ 66]), "=f"(tmp[ 67]),
                      "=f"(tmp[ 68]), "=f"(tmp[ 69]), "=f"(tmp[ 70]), "=f"(tmp[ 71]),
                      "=f"(tmp[ 72]), "=f"(tmp[ 73]), "=f"(tmp[ 74]), "=f"(tmp[ 75]),
                      "=f"(tmp[ 76]), "=f"(tmp[ 77]), "=f"(tmp[ 78]), "=f"(tmp[ 79]),
                      "=f"(tmp[ 80]), "=f"(tmp[ 81]), "=f"(tmp[ 82]), "=f"(tmp[ 83]),
                      "=f"(tmp[ 84]), "=f"(tmp[ 85]), "=f"(tmp[ 86]), "=f"(tmp[ 87]),
                      "=f"(tmp[ 88]), "=f"(tmp[ 89]), "=f"(tmp[ 90]), "=f"(tmp[ 91]),
                      "=f"(tmp[ 92]), "=f"(tmp[ 93]), "=f"(tmp[ 94]), "=f"(tmp[ 95]),
                      "=f"(tmp[ 96]), "=f"(tmp[ 97]), "=f"(tmp[ 98]), "=f"(tmp[ 99]),
                      "=f"(tmp[100]), "=f"(tmp[101]), "=f"(tmp[102]), "=f"(tmp[103]),
                      "=f"(tmp[104]), "=f"(tmp[105]), "=f"(tmp[106]), "=f"(tmp[107]),
                      "=f"(tmp[108]), "=f"(tmp[109]), "=f"(tmp[110]), "=f"(tmp[111]),
                      "=f"(tmp[112]), "=f"(tmp[113]), "=f"(tmp[114]), "=f"(tmp[115]),
                      "=f"(tmp[116]), "=f"(tmp[117]), "=f"(tmp[118]), "=f"(tmp[119]),
                      "=f"(tmp[120]), "=f"(tmp[121]), "=f"(tmp[122]), "=f"(tmp[123]),
                      "=f"(tmp[124]), "=f"(tmp[125]), "=f"(tmp[126]), "=f"(tmp[127])
                    : "r"(addr));
}

"""

cuda_kernel_source = r"""
constexpr int MMA_K{64};  // FP4 MMA K-dimension size.
constexpr int NUM_STAGES{4};
constexpr int NUM_STAGES_EPILOGUE{2};
constexpr int MAX_GROUPS{8};


struct GroupParams {
    const char *SFA;
    const char *SFB;
    half *C;
    int M, N, K;
    int block_offset;   // cumulative block count before this group
    int grid_dim_n;
    int rest_k;         // K / 16 / 4
    int num_iters;      // K / BLOCK_K
};

struct GroupedKernelArgs {
    CUtensorMap tmaps[MAX_GROUPS * 2];  // A_tmap, B_tmap per group
    GroupParams params[MAX_GROUPS];
    int num_groups;
};



template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>
__global__
__launch_bounds__(NUM_THREADS) void kernel_v10_persistent(
    const __grid_constant__ GroupedKernelArgs args,
    const int total_num_tiles
) {
    const int thread_idx{static_cast<int>(threadIdx.x)};
    const int global_block_idx{static_cast<int>(blockIdx.x)};
    const int warp_idx{thread_idx / WARP_SIZE};
    const int lane_idx{thread_idx % WARP_SIZE};
    const int num_blocks{gridDim.x};

    // =========================================================================
    // Shared memory setup
    // =========================================================================
    // 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;
    }

    // =========================================================================
    // Mbarrier and tensor memory setup
    // =========================================================================
#pragma nv_diag_suppress static_var_with_dynamic_init
    __shared__ uint64_t tma_mbars[NUM_STAGES];
    __shared__ uint64_t mma_mbars[NUM_STAGES];
    __shared__ uint64_t mainloop_mbars[NUM_STAGES_EPILOGUE];
    __shared__ uint64_t epilogue_mbars[NUM_STAGES_EPILOGUE];
    int tma_mbar_addrs[NUM_STAGES], mma_mbar_addrs[NUM_STAGES];
    int mainloop_mbar_addrs[NUM_STAGES_EPILOGUE], epilogue_mbar_addrs[NUM_STAGES_EPILOGUE];
    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]));
    }
    for (int s{0}; s < NUM_STAGES_EPILOGUE; ++s) {
        mainloop_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mainloop_mbars[s]));
        epilogue_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&epilogue_mbars[s]));
    }

    __shared__ int tmem_addr[1];

    constexpr int SFA_tmem_start_col{NUM_STAGES_EPILOGUE * BLOCK_N};
    constexpr int SFB_tmem_start_col{SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K)};
    constexpr int TMEM_COLS{NUM_STAGES_EPILOGUE * BLOCK_N * 2}; // Has to be a power of 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);
        }
        for (int s{0}; s < NUM_STAGES_EPILOGUE; ++s) {
            mbarrier_init(mainloop_mbar_addrs[s], 1);
            mbarrier_init(epilogue_mbar_addrs[s], 4);
        }
        
        asm volatile("fence.mbarrier_init.release.cluster;"); // Make it visible to async proxy.
    } else if (warp_idx == 1) {
        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();

    // Make sure tcgen05.alloc has completed.
    const int taddr{tmem_addr[0]};

    constexpr uint32_t i_desc = (1U << 7U)
                                | (1U << 10U)
                                | ((uint32_t)BLOCK_N >> 3U << 17U)
                                | ((uint32_t)BLOCK_M >> 7U << 27U);
    constexpr int cp_size{(BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size};

    // =========================================================================
    // Define functions
    // =========================================================================
    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);
    };

    auto load = [&](const int stage, const int mma_phase, const int iter_k, const int group_idx, const int block_idx) {
        const GroupParams &gp = args.params[group_idx];
        const int group_block_idx{block_idx - gp.block_offset};
        const int block_idx_m{group_block_idx / gp.grid_dim_n};
        const int block_idx_n{group_block_idx % gp.grid_dim_n};

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

        const CUtensorMap *A_tmap_ptr = &args.tmaps[group_idx * 2];
        const CUtensorMap *B_tmap_ptr = &args.tmaps[group_idx * 2 + 1];
        
        mbarrier_wait(mma_mbar_addrs[stage], mma_phase);

        tma_3d_gmem2smem(A_smem[stage], A_tmap_ptr, 0, offset_m, offset_k / 256, tma_mbar_addrs[stage]);
        tma_3d_gmem2smem(B_smem[stage], B_tmap_ptr, 0, offset_n, offset_k / 256, tma_mbar_addrs[stage]);

        const char *SFA_src = gp.SFA + ((offset_m / 128) * gp.rest_k + offset_k / (16 * 4)) * 512;
        const char *SFB_src = gp.SFB + ((offset_n / 128) * gp.rest_k + offset_k / (16 * 4)) * 512;
        tma_gmem2smem(SFA_smem[stage], SFA_src, SF_size, tma_mbar_addrs[stage]);
        tma_gmem2smem(SFB_smem[stage], SFB_src, SF_size, tma_mbar_addrs[stage]);
        asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
            ::"r"(tma_mbar_addrs[stage]), "r"(cp_size) : "memory");
    };

    auto compute = [&](const int tma_stage, const int tma_phase, const int iter_k, const int mainloop_stage,
                       const int group_idx, const int block_idx) {
        const GroupParams &gp = args.params[group_idx];
        const int group_block_idx{block_idx - gp.block_offset};
        const int block_idx_m{group_block_idx / gp.grid_dim_n};
        const int block_idx_n{group_block_idx % gp.grid_dim_n};
        
        mbarrier_wait(tma_mbar_addrs[tma_stage], tma_phase);

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

        // Perform MMA.
        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[tma_stage] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
                uint64_t b_desc{make_desc_AB(B_smem[tma_stage] + 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 + (block_idx_m % (128 / BLOCK_M)) * (BLOCK_M / 32)};
                const int scale_B_tmem{SFB_tmem_start_col + k * 4 + (block_idx_n % (128 / BLOCK_N)) * (BLOCK_N / 32)};
                const int enable_input_d{(k1 == 0 && k2 == 0) ? iter_k : 1};
                run_mma_nvfp4_custom_addr(taddr + mainloop_stage * BLOCK_N, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
            }
        }

        asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
            ::"r"(mma_mbar_addrs[tma_stage]) : "memory");
    };

    auto epilogue = [&](int mainloop_stage, const int group_idx, const int block_idx) {
        const GroupParams &gp = args.params[group_idx];
        const int group_block_idx{block_idx - gp.block_offset};
        const int block_idx_m{group_block_idx / gp.grid_dim_n};
        const int block_idx_n{group_block_idx % gp.grid_dim_n};

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

        // Remap to access tmem since we use warp 2 - 5 for epilogue.
        const int epilogue_warp_idx{warp_idx % 4};
        const int epilogue_thread_idx{epilogue_warp_idx * WARP_SIZE + lane_idx};

        const int row{offset_m + epilogue_thread_idx};
        constexpr int WIDTH{16};
        constexpr int NUM_CHUNKS{BLOCK_N / WIDTH};

        // Issue all loads.
        float tmp[NUM_CHUNKS][WIDTH];
        for (int n{0}; n < NUM_CHUNKS; ++n) {
            const int addr = taddr + (mainloop_stage * BLOCK_N) + ((epilogue_warp_idx * 32) << 16) + (n * WIDTH);
            load_from_tmem_32x32b_x16(tmp[n], addr);
        }
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        if (row < gp.M){
            for (int n{0}; n < NUM_CHUNKS; ++n) {
                const int col{offset_n + n * WIDTH};
                half2 out[WIDTH / 2];
                for (int i{0}; i < WIDTH / 2; ++i)
                    out[i] = __float22half2_rn({tmp[n][i * 2], tmp[n][i * 2 + 1]});
            
                half *out_ptr = gp.C + row * gp.N + col;
                if (col + WIDTH <= gp.N) {
                    for (int i{0}; i < WIDTH / 8; ++i)
                        reinterpret_cast<int4 *>(out_ptr)[i] = reinterpret_cast<int4 *>(out)[i];
                } else {
                    const half *out_half = reinterpret_cast<const half *>(out);
                    for (int i{0}; i < WIDTH && col + i < gp.N; ++i)
                        out_ptr[i] = out_half[i];
                }
            }
        }
    };

    auto compute_group_idx = [&](int block_idx) -> int {
        int group_idx = 0;
        for (int g{1}; g < args.num_groups; ++g)
            if (block_idx >= args.params[g].block_offset)
                group_idx = g;
        return group_idx;
    };

    // =========================================================================
    // Warp 0: TMA Producer
    // =========================================================================
    if (warp_idx == 0 && elect_sync()) {
        int tma_stage{0};
        int mma_phase{1};

        for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
            const int group_idx{compute_group_idx(this_block_idx)};
            const int num_iters{args.params[group_idx].K / BLOCK_K};

            for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
                load(tma_stage, mma_phase, iter_k, group_idx, this_block_idx);

                // Flip phase when we have cycled through all TMA buffers.
                tma_stage = (tma_stage + 1) % NUM_STAGES;
                if (tma_stage == 0)
                    mma_phase ^= 1;
            }
        }
    // ========================================================================
    // Warp 1: MMA Consumer
    // ========================================================================
    } else if (warp_idx == 1 && elect_sync()) {
        int tma_stage{0};
        int tma_phase{0};
        int mainloop_stage{0};
        int epilogue_phase{1};

        for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
            const int group_idx{compute_group_idx(this_block_idx)};
            const int num_iters{args.params[group_idx].K / BLOCK_K};

            // Wait for epilogue to finish since we'll be reusing the tmem.
            mbarrier_wait(epilogue_mbar_addrs[mainloop_stage], epilogue_phase);

            for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
                compute(tma_stage, tma_phase, iter_k, mainloop_stage, group_idx, this_block_idx);

                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_addrs[mainloop_stage]) : "memory");
            
            mainloop_stage = (mainloop_stage + 1) % NUM_STAGES_EPILOGUE;
            if (mainloop_stage == 0)
                epilogue_phase ^= 1;
        }
    // ========================================================================
    // Warp 2 - 5: Epilogue
    // ========================================================================
    } else if (warp_idx >= 2) {
        int mainloop_stage{0};
        int mainloop_phase{0};

        for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
            // Wait for mainloop to finish.
            mbarrier_wait(mainloop_mbar_addrs[mainloop_stage], mainloop_phase);
            // PTX doc says we need to add this before tcgen05.ld, after tcgen05.mma
            asm volatile("tcgen05.fence::after_thread_sync;");

            const int group_idx{compute_group_idx(this_block_idx)};
            epilogue(mainloop_stage, group_idx, this_block_idx);

            if (elect_sync()) {
                asm volatile("mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];"
                :: "r"(epilogue_mbar_addrs[mainloop_stage]) : "memory");
            }

            mainloop_stage = (mainloop_stage + 1) % 2;
            if (mainloop_stage == 0)
                mainloop_phase ^= 1;
        }
    }

    __syncthreads();

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


std::vector<torch::Tensor> launch_kernel(
    std::vector<torch::Tensor> As,
    std::vector<torch::Tensor> Bs,
    std::vector<torch::Tensor> SFAs,
    std::vector<torch::Tensor> SFBs,
    std::vector<int64_t> Ms,
    std::vector<int64_t> Ns,
    std::vector<int64_t> Ks
) {
    constexpr int BLOCK_M{128};
    constexpr int BLOCK_N{128};
    constexpr int BLOCK_K{256};
    // 1 for TMA, 1 for MMA, 4 for epilogue.
    constexpr int NUM_THREADS{6 * WARP_SIZE};

    int num_groups = As.size();

    // Create output tensors.
    std::vector<torch::Tensor> Cs;
    for (int g = 0; g < num_groups; ++g) {
        Cs.push_back(torch::empty({Ms[g], Ns[g]},
            torch::dtype(torch::kFloat16).device(As[g].device())));
    }

    // Build kernel args on the host stack (~2.5 KB, fits in 4 KB param limit).
    GroupedKernelArgs host_args{};
    host_args.num_groups = num_groups;
    int total_blocks = 0;

    for (int g = 0; g < num_groups; ++g) {
        auto A_ptr = reinterpret_cast<const char *>(As[g].data_ptr());
        auto B_ptr = reinterpret_cast<const char *>(Bs[g].data_ptr());

        create_tmap_descriptor<256>(&host_args.tmaps[g * 2], A_ptr, Ms[g], Ks[g],
            BLOCK_M, BLOCK_K, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
        create_tmap_descriptor<256>(&host_args.tmaps[g * 2 + 1], B_ptr, Ns[g], Ks[g],
            BLOCK_N, BLOCK_K, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);

        int grid_m = (Ms[g] + BLOCK_M - 1) / BLOCK_M;
        int grid_n = (Ns[g] + BLOCK_N - 1) / BLOCK_N;

        host_args.params[g].SFA = reinterpret_cast<const char *>(SFAs[g].data_ptr());
        host_args.params[g].SFB = reinterpret_cast<const char *>(SFBs[g].data_ptr());
        host_args.params[g].C = reinterpret_cast<half *>(Cs[g].data_ptr<at::Half>());
        host_args.params[g].M = Ms[g];
        host_args.params[g].N = Ns[g];
        host_args.params[g].K = Ks[g];
        host_args.params[g].block_offset = total_blocks;
        host_args.params[g].grid_dim_n = grid_n;
        host_args.params[g].rest_k = Ks[g] / 16 / 4;
        host_args.params[g].num_iters = Ks[g] / BLOCK_K;

        total_blocks += grid_m * grid_n;
    }

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

    auto kernel = kernel_v10_persistent<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;

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

    // The number of SMs.
    constexpr int NUM_BLOCKS{148};
    // Using __grid_constant__ allows passing host_args via constant memory, no need to do allocation.
    kernel<<<NUM_BLOCKS, NUM_THREADS, SHARED_SIZE>>>(host_args, total_blocks);

    return Cs;
}
"""

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

std::vector<torch::Tensor> launch_kernel(
    std::vector<torch::Tensor> As,
    std::vector<torch::Tensor> Bs,
    std::vector<torch::Tensor> SFAs,
    std::vector<torch::Tensor> SFBs,
    std::vector<int64_t> Ms,
    std::vector<int64_t> Ns,
    std::vector<int64_t> Ks);
"""

module = load_inline(
    name='kernel',
    cpp_sources=cpp_source,
    cuda_sources=cuda_common_source + cuda_kernel_source,
    functions=['launch_kernel'],
    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:
    As, Bs, SFAs, SFBs = [], [], [], []
    Ms, Ns, Ks = [], [], []
    cs = []
    for (a, b, c), _, (sfa_reordered, sfb_reordered), (m, n, k, _) in zip(*data):
        As.append(a[:, :, 0])
        Bs.append(b[:, :, 0])
        SFAs.append(sfa_reordered)
        SFBs.append(sfb_reordered)
        Ms.append(m)
        Ns.append(n)
        Ks.append(k)
        cs.append(c)

    outputs = module.launch_kernel(As, Bs, SFAs, SFBs, Ms, Ns, Ks)

    for c, out in zip(cs, outputs):
        c[:, :, 0] = out

    return cs
scrolls · 775 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 498891.

⋯ 171 unchanged lines
);
}
+ template <int NUM_CTA=1>
__device__ inline
+ void run_mma_nvfp4_custom_addr(
+ const int d_tmem,
+ uint64_t a_desc,
+ uint64_t b_desc,
+ uint32_t i_desc,
+ int scale_A_tmem,
+ int scale_B_tmem,
+ int enable_input_d
+ ) {
+ asm volatile(
+ "{\n\t"
+ ".reg .pred p;\n\t"
+ "setp.ne.b32 p, %6, 0;\n\t"
+ "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
+ "}"
+ :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
+ "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d), "n"(NUM_CTA)
+ );
+ }
+
+ __device__ inline
void load_from_tmem_32x32b_x8(float *tmp, const 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]),
⋯ 13 unchanged lines
"=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
: "r"(addr));
}
+
+ __device__ inline
+ void load_from_tmem_32x32b_x32(float *tmp, const 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__ inline
+ void load_from_tmem_32x32b_x64(float *tmp, const 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));
+ }
+
+
+ __device__ inline
+ void load_from_tmem_32x32b_x128(float *tmp, const int addr) {
+ asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.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, %65, %66, %67, %68, %69, %70, %71, "
+ " %72, %73, %74, %75, %76, %77, %78, %79, "
+ " %80, %81, %82, %83, %84, %85, %86, %87, "
+ " %88, %89, %90, %91, %92, %93, %94, %95, "
+ " %96, %97, %98, %99, %100, %101, %102, %103, "
+ " %104, %105, %106, %107, %108, %109, %110, %111, "
+ " %112, %113, %114, %115, %116, %117, %118, %119, "
+ " %120, %121, %122, %123, %124, %125, %126, %127}, [%128];"
+ : "=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]),
+ "=f"(tmp[ 64]), "=f"(tmp[ 65]), "=f"(tmp[ 66]), "=f"(tmp[ 67]),
+ "=f"(tmp[ 68]), "=f"(tmp[ 69]), "=f"(tmp[ 70]), "=f"(tmp[ 71]),
+ "=f"(tmp[ 72]), "=f"(tmp[ 73]), "=f"(tmp[ 74]), "=f"(tmp[ 75]),
+ "=f"(tmp[ 76]), "=f"(tmp[ 77]), "=f"(tmp[ 78]), "=f"(tmp[ 79]),
+ "=f"(tmp[ 80]), "=f"(tmp[ 81]), "=f"(tmp[ 82]), "=f"(tmp[ 83]),
+ "=f"(tmp[ 84]), "=f"(tmp[ 85]), "=f"(tmp[ 86]), "=f"(tmp[ 87]),
+ "=f"(tmp[ 88]), "=f"(tmp[ 89]), "=f"(tmp[ 90]), "=f"(tmp[ 91]),
+ "=f"(tmp[ 92]), "=f"(tmp[ 93]), "=f"(tmp[ 94]), "=f"(tmp[ 95]),
+ "=f"(tmp[ 96]), "=f"(tmp[ 97]), "=f"(tmp[ 98]), "=f"(tmp[ 99]),
+ "=f"(tmp[100]), "=f"(tmp[101]), "=f"(tmp[102]), "=f"(tmp[103]),
+ "=f"(tmp[104]), "=f"(tmp[105]), "=f"(tmp[106]), "=f"(tmp[107]),
+ "=f"(tmp[108]), "=f"(tmp[109]), "=f"(tmp[110]), "=f"(tmp[111]),
+ "=f"(tmp[112]), "=f"(tmp[113]), "=f"(tmp[114]), "=f"(tmp[115]),
+ "=f"(tmp[116]), "=f"(tmp[117]), "=f"(tmp[118]), "=f"(tmp[119]),
+ "=f"(tmp[120]), "=f"(tmp[121]), "=f"(tmp[122]), "=f"(tmp[123]),
+ "=f"(tmp[124]), "=f"(tmp[125]), "=f"(tmp[126]), "=f"(tmp[127])
+ : "r"(addr));
+ }
+
"""
cuda_kernel_source = r"""
- constexpr int MMA_K = 64; // FP4 MMA K-dimension size.
- constexpr int NUM_STAGES = 4;
- constexpr int MAX_GROUPS = 8;
+ constexpr int MMA_K{64}; // FP4 MMA K-dimension size.
+ constexpr int NUM_STAGES{4};
+ constexpr int NUM_STAGES_EPILOGUE{2};
+ constexpr int MAX_GROUPS{8};
struct GroupParams {
⋯ 14 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_v09_improve_epilogue(
- const __grid_constant__ GroupedKernelArgs args
+ __launch_bounds__(NUM_THREADS) void kernel_v10_persistent(
+ const __grid_constant__ GroupedKernelArgs args,
+ const int total_num_tiles
) {
const int thread_idx{static_cast<int>(threadIdx.x)};
const int global_block_idx{static_cast<int>(blockIdx.x)};
const int warp_idx{thread_idx / WARP_SIZE};
+ const int lane_idx{thread_idx % WARP_SIZE};
+ const int num_blocks{gridDim.x};
- // Find which group this block belongs to (linear scan, num_groups <= 8).
- int group = 0;
- for (int g = 1; g < args.num_groups; ++g) {
- if (global_block_idx >= args.params[g].block_offset)
- group = g;
- }
-
- // Load per-group parameters.
- const GroupParams &gp = args.params[group];
- const int block_idx = global_block_idx - gp.block_offset;
-
- const int block_idx_m{block_idx / gp.grid_dim_n};
- const int block_idx_n{block_idx % gp.grid_dim_n};
-
- const int offset_m{block_idx_m * BLOCK_M};
- const int offset_n{block_idx_n * BLOCK_N};
-
- const int M = gp.M;
- const int N = gp.N;
- const int num_iters = gp.num_iters;
- const int rest_k = gp.rest_k;
- const char *SFA = gp.SFA;
- const char *SFB = gp.SFB;
- half *C = gp.C;
-
- // Tensor maps for this group (in __grid_constant__ / .param space).
- const CUtensorMap *A_tmap_ptr = &args.tmaps[group * 2];
- const CUtensorMap *B_tmap_ptr = &args.tmaps[group * 2 + 1];
-
+ // =========================================================================
+ // Shared memory setup
+ // =========================================================================
// Multi-buffered shared memory layout:
// [buf0: A | B | SFA | SFB | buf1: ... | buf2: ... | buf3: ...]
constexpr int SF_size = 512 * BLOCK_K / MMA_K;
⋯ 1 unchanged lines
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;
⋯ 2 unchanged lines
SFB_smem[s] = SFA_smem[s] + SF_size;
}
+ // =========================================================================
+ // Mbarrier and tensor memory setup
+ // =========================================================================
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ uint64_t tma_mbars[NUM_STAGES];
__shared__ uint64_t mma_mbars[NUM_STAGES];
+ __shared__ uint64_t mainloop_mbars[NUM_STAGES_EPILOGUE];
+ __shared__ uint64_t epilogue_mbars[NUM_STAGES_EPILOGUE];
int tma_mbar_addrs[NUM_STAGES], mma_mbar_addrs[NUM_STAGES];
+ int mainloop_mbar_addrs[NUM_STAGES_EPILOGUE], epilogue_mbar_addrs[NUM_STAGES_EPILOGUE];
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]));
}
+ for (int s{0}; s < NUM_STAGES_EPILOGUE; ++s) {
+ mainloop_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mainloop_mbars[s]));
+ epilogue_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&epilogue_mbars[s]));
+ }
+
__shared__ int tmem_addr[1];
- 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;
+ constexpr int SFA_tmem_start_col{NUM_STAGES_EPILOGUE * BLOCK_N};
+ constexpr int SFB_tmem_start_col{SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K)};
+ constexpr int TMEM_COLS{NUM_STAGES_EPILOGUE * BLOCK_N * 2}; // Has to be a power of 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;");
+ for (int s{0}; s < NUM_STAGES_EPILOGUE; ++s) {
+ mbarrier_init(mainloop_mbar_addrs[s], 1);
+ mbarrier_init(epilogue_mbar_addrs[s], 4);
+ }
+
+ asm volatile("fence.mbarrier_init.release.cluster;"); // Make it visible to async proxy.
} else if (warp_idx == 1) {
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;"
⋯ 1 unchanged lines
}
__syncthreads();
+ // Make sure tcgen05.alloc has completed.
const int taddr{tmem_addr[0]};
constexpr uint32_t i_desc = (1U << 7U)
| (1U << 10U)
| ((uint32_t)BLOCK_N >> 3U << 17U)
| ((uint32_t)BLOCK_M >> 7U << 27U);
+ constexpr int cp_size{(BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size};
- constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;
-
// =========================================================================
- // Warp 0: TMA Producer
+ // Define functions
// =========================================================================
- if (warp_idx == 0 && elect_sync()) {
- int mma_prod_phase[NUM_STAGES] = {};
+ 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);
+ };
- for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
- const int s{iter_k % NUM_STAGES};
+ auto make_desc_SF = [](int addr) -> uint64_t {
+ const int SBO = 8 * 16;
+ return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);
+ };
- if (iter_k >= NUM_STAGES) {
- mbarrier_wait(mma_mbar_addrs[s], mma_prod_phase[s]);
- mma_prod_phase[s] ^= 1;
- }
+ auto load = [&](const int stage, const int mma_phase, const int iter_k, const int group_idx, const int block_idx) {
+ const GroupParams &gp = args.params[group_idx];
+ const int group_block_idx{block_idx - gp.block_offset};
+ const int block_idx_m{group_block_idx / gp.grid_dim_n};
+ const int block_idx_n{group_block_idx % gp.grid_dim_n};
- const int off_k{iter_k * BLOCK_K};
- tma_3d_gmem2smem(A_smem[s], A_tmap_ptr, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);
- tma_3d_gmem2smem(B_smem[s], B_tmap_ptr, 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");
- }
+ const int offset_m{block_idx_m * BLOCK_M};
+ const int offset_n{block_idx_n * BLOCK_N};
+ const int offset_k{iter_k * BLOCK_K};
- // =========================================================================
- // Warp 1: MMA Consumer
- // =========================================================================
- } else if (warp_idx == 1 && elect_sync()) {
- int tma_cons_phase[NUM_STAGES] = {};
- int mma_done_phase[NUM_STAGES] = {};
+ const CUtensorMap *A_tmap_ptr = &args.tmaps[group_idx * 2];
+ const CUtensorMap *B_tmap_ptr = &args.tmaps[group_idx * 2 + 1];
+
+ mbarrier_wait(mma_mbar_addrs[stage], mma_phase);
- for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
- const int s{iter_k % NUM_STAGES};
+ tma_3d_gmem2smem(A_smem[stage], A_tmap_ptr, 0, offset_m, offset_k / 256, tma_mbar_addrs[stage]);
+ tma_3d_gmem2smem(B_smem[stage], B_tmap_ptr, 0, offset_n, offset_k / 256, tma_mbar_addrs[stage]);
- mbarrier_wait(tma_mbar_addrs[s], tma_cons_phase[s]);
- tma_cons_phase[s] ^= 1;
+ const char *SFA_src = gp.SFA + ((offset_m / 128) * gp.rest_k + offset_k / (16 * 4)) * 512;
+ const char *SFB_src = gp.SFB + ((offset_n / 128) * gp.rest_k + offset_k / (16 * 4)) * 512;
+ tma_gmem2smem(SFA_smem[stage], SFA_src, SF_size, tma_mbar_addrs[stage]);
+ tma_gmem2smem(SFB_smem[stage], SFB_src, SF_size, tma_mbar_addrs[stage]);
+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
+ ::"r"(tma_mbar_addrs[stage]), "r"(cp_size) : "memory");
+ };
- 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 compute = [&](const int tma_stage, const int tma_phase, const int iter_k, const int mainloop_stage,
+ const int group_idx, const int block_idx) {
+ const GroupParams &gp = args.params[group_idx];
+ const int group_block_idx{block_idx - gp.block_offset};
+ const int block_idx_m{group_block_idx / gp.grid_dim_n};
+ const int block_idx_n{group_block_idx % gp.grid_dim_n};
+
+ mbarrier_wait(tma_mbar_addrs[tma_stage], tma_phase);
- 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 into tensor memory.
+ for (int k{0}; k < BLOCK_K / MMA_K; ++k) {
+ uint64_t sfa_desc = make_desc_SF(SFA_smem[tma_stage] + k * 512);
+ uint64_t sfb_desc = make_desc_SF(SFB_smem[tma_stage] + k * 512);
+ copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);
+ copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);
+ }
- 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);
+ // Perform MMA.
+ 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[tma_stage] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
+ uint64_t b_desc{make_desc_AB(B_smem[tma_stage] + 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 + (block_idx_m % (128 / BLOCK_M)) * (BLOCK_M / 32)};
+ const int scale_B_tmem{SFB_tmem_start_col + k * 4 + (block_idx_n % (128 / BLOCK_N)) * (BLOCK_N / 32)};
+ const int enable_input_d{(k1 == 0 && k2 == 0) ? iter_k : 1};
+ run_mma_nvfp4_custom_addr(taddr + mainloop_stage * BLOCK_N, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
+ }
- 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)};
+ asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
+ ::"r"(mma_mbar_addrs[tma_stage]) : "memory");
+ };
- 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 epilogue = [&](int mainloop_stage, const int group_idx, const int block_idx) {
+ const GroupParams &gp = args.params[group_idx];
+ const int group_block_idx{block_idx - gp.block_offset};
+ const int block_idx_m{group_block_idx / gp.grid_dim_n};
+ const int block_idx_n{group_block_idx % gp.grid_dim_n};
- 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);
+ const int offset_m{block_idx_m * BLOCK_M};
+ const int offset_n{block_idx_n * BLOCK_N};
+
+ // Remap to access tmem since we use warp 2 - 5 for epilogue.
+ const int epilogue_warp_idx{warp_idx % 4};
+ const int epilogue_thread_idx{epilogue_warp_idx * WARP_SIZE + lane_idx};
+
+ const int row{offset_m + epilogue_thread_idx};
+ constexpr int WIDTH{16};
+ constexpr int NUM_CHUNKS{BLOCK_N / WIDTH};
+
+ // Issue all loads.
+ float tmp[NUM_CHUNKS][WIDTH];
+ for (int n{0}; n < NUM_CHUNKS; ++n) {
+ const int addr = taddr + (mainloop_stage * BLOCK_N) + ((epilogue_warp_idx * 32) << 16) + (n * WIDTH);
+ load_from_tmem_32x32b_x16(tmp[n], addr);
+ }
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
+
+ if (row < gp.M){
+ for (int n{0}; n < NUM_CHUNKS; ++n) {
+ const int col{offset_n + n * WIDTH};
+ half2 out[WIDTH / 2];
+ for (int i{0}; i < WIDTH / 2; ++i)
+ out[i] = __float22half2_rn({tmp[n][i * 2], tmp[n][i * 2 + 1]});
+
+ half *out_ptr = gp.C + row * gp.N + col;
+ if (col + WIDTH <= gp.N) {
+ for (int i{0}; i < WIDTH / 8; ++i)
+ reinterpret_cast<int4 *>(out_ptr)[i] = reinterpret_cast<int4 *>(out)[i];
+ } else {
+ const half *out_half = reinterpret_cast<const half *>(out);
+ for (int i{0}; i < WIDTH && col + i < gp.N; ++i)
+ out_ptr[i] = out_half[i];
}
}
-
- 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;
}
+ };
- 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);
+ auto compute_group_idx = [&](int block_idx) -> int {
+ int group_idx = 0;
+ for (int g{1}; g < args.num_groups; ++g)
+ if (block_idx >= args.params[g].block_offset)
+ group_idx = g;
+ return group_idx;
+ };
+
+ // =========================================================================
+ // Warp 0: TMA Producer
+ // =========================================================================
+ if (warp_idx == 0 && elect_sync()) {
+ int tma_stage{0};
+ int mma_phase{1};
+
+ for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
+ const int group_idx{compute_group_idx(this_block_idx)};
+ const int num_iters{args.params[group_idx].K / BLOCK_K};
+
+ for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
+ load(tma_stage, mma_phase, iter_k, group_idx, this_block_idx);
+
+ // Flip phase when we have cycled through all TMA buffers.
+ tma_stage = (tma_stage + 1) % NUM_STAGES;
+ if (tma_stage == 0)
+ mma_phase ^= 1;
+ }
}
- }
+ // ========================================================================
+ // Warp 1: MMA Consumer
+ // ========================================================================
+ } else if (warp_idx == 1 && elect_sync()) {
+ int tma_stage{0};
+ int tma_phase{0};
+ int mainloop_stage{0};
+ int epilogue_phase{1};
- __syncthreads();
+ for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
+ const int group_idx{compute_group_idx(this_block_idx)};
+ const int num_iters{args.params[group_idx].K / BLOCK_K};
- // === Epilogue: Read accumulator from TMEM and store to global memory ===
- asm volatile("tcgen05.fence::after_thread_sync;");
+ // Wait for epilogue to finish since we'll be reusing the tmem.
+ mbarrier_wait(epilogue_mbar_addrs[mainloop_stage], epilogue_phase);
- const int row{offset_m + thread_idx};
- constexpr int WIDTH{16};
- for (int n{0}; n < BLOCK_N / WIDTH; ++n) {
- float tmp[WIDTH];
- const int addr = taddr + ((warp_idx * 32) << 16) + (n * WIDTH);
- load_from_tmem_32x32b_x16(tmp, addr);
- asm volatile("tcgen05.wait::ld.sync.aligned;");
+ for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
+ compute(tma_stage, tma_phase, iter_k, mainloop_stage, group_idx, this_block_idx);
- if (row >= M) continue;
+ tma_stage = (tma_stage + 1) % NUM_STAGES;
+ if (tma_stage == 0)
+ tma_phase ^= 1;
+ }
- const int col{offset_n + n * WIDTH};
- half2 out[WIDTH / 2];
- for (int i{0}; i < WIDTH / 2; ++i)
- out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});
+ asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
+ ::"r"(mainloop_mbar_addrs[mainloop_stage]) : "memory");
+
+ mainloop_stage = (mainloop_stage + 1) % NUM_STAGES_EPILOGUE;
+ if (mainloop_stage == 0)
+ epilogue_phase ^= 1;
+ }
+ // ========================================================================
+ // Warp 2 - 5: Epilogue
+ // ========================================================================
+ } else if (warp_idx >= 2) {
+ int mainloop_stage{0};
+ int mainloop_phase{0};
- half *out_ptr = C + row * N + col;
- if (col + WIDTH <= N) {
- for (int i{0}; i < WIDTH / 8; ++i)
- reinterpret_cast<int4 *>(out_ptr)[i] = reinterpret_cast<int4 *>(out)[i];
- } else {
- const half *out_half = reinterpret_cast<const half *>(out);
- for (int i{0}; i < WIDTH && col + i < N; ++i)
- out_ptr[i] = out_half[i];
+ for (int this_block_idx{global_block_idx}; this_block_idx < total_num_tiles; this_block_idx += num_blocks) {
+ // Wait for mainloop to finish.
+ mbarrier_wait(mainloop_mbar_addrs[mainloop_stage], mainloop_phase);
+ // PTX doc says we need to add this before tcgen05.ld, after tcgen05.mma
+ asm volatile("tcgen05.fence::after_thread_sync;");
+
+ const int group_idx{compute_group_idx(this_block_idx)};
+ epilogue(mainloop_stage, group_idx, this_block_idx);
+
+ if (elect_sync()) {
+ asm volatile("mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];"
+ :: "r"(epilogue_mbar_addrs[mainloop_stage]) : "memory");
+ }
+
+ mainloop_stage = (mainloop_stage + 1) % 2;
+ if (mainloop_stage == 0)
+ mainloop_phase ^= 1;
}
}
-
+
__syncthreads();
if (warp_idx == 0) {
⋯ 2 unchanged lines
}
- std::vector<torch::Tensor> launch_grouped_kernel(
+ std::vector<torch::Tensor> launch_kernel(
std::vector<torch::Tensor> As,
std::vector<torch::Tensor> Bs,
std::vector<torch::Tensor> SFAs,
⋯ 5 unchanged lines
constexpr int BLOCK_M{128};
constexpr int BLOCK_N{128};
constexpr int BLOCK_K{256};
- constexpr int NUM_THREADS{4 * WARP_SIZE};
+ // 1 for TMA, 1 for MMA, 4 for epilogue.
+ constexpr int NUM_THREADS{6 * WARP_SIZE};
int num_groups = As.size();
⋯ 39 unchanged lines
constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};
constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};
- auto kernel = kernel_v09_improve_epilogue<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;
+ auto kernel = kernel_v10_persistent<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;
if (SHARED_SIZE > 48'000)
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);
+ // The number of SMs.
+ constexpr int NUM_BLOCKS{148};
// Using __grid_constant__ allows passing host_args via constant memory, no need to do allocation.
- kernel<<<total_blocks, NUM_THREADS, SHARED_SIZE>>>(host_args);
+ kernel<<<NUM_BLOCKS, NUM_THREADS, SHARED_SIZE>>>(host_args, total_blocks);
return Cs;
}
⋯ 2 unchanged lines
cpp_source = """
#include <torch/extension.h>
- std::vector<torch::Tensor> launch_grouped_kernel(
+ std::vector<torch::Tensor> launch_kernel(
std::vector<torch::Tensor> As,
std::vector<torch::Tensor> Bs,
std::vector<torch::Tensor> SFAs,
⋯ 7 unchanged lines
name='kernel',
cpp_sources=cpp_source,
cuda_sources=cuda_common_source + cuda_kernel_source,
- functions=['launch_grouped_kernel'],
+ functions=['launch_kernel'],
verbose=True,
is_python_module=True,
no_implicit_headers=True,
⋯ 27 unchanged lines
Ks.append(k)
cs.append(c)
- outputs = module.launch_grouped_kernel(As, Bs, SFAs, SFBs, Ms, Ns, Ks)
+ outputs = module.launch_kernel(As, Bs, SFAs, SFBs, Ms, Ns, Ks)
for c, out in zip(cs, outputs):
c[:, :, 0] = out
scrolls · 633 diff lines total

Best evidence level for this revision: reported

JSON