Skip to content
KernelIndex
Search⌘K

submission 485603

jiab_85281 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_static.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-485603?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
19.4µs
#77 of 310
2026-02-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2570c0e7712a212921e6d3ca149ab2955856779392971d46ba6e7a5f994274f6
license declaredunknown
license concludedunknown
authorsjiab_85281
imported2026-08-15

Techniques

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

fused-epilogue__device__ inline void do_epilogue(int warp_id, int lane_id, int done_mbar, int d_tmem_base,
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
shared-memory__device__ inline void fence_proxy_tensormap(const void *smem_ptr) {
stages = 6constexpr int NUM_STAGES = 6;
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tile-k = 256constexpr int BLOCK_K = 256;
tile-m = 128constexpr int BLOCK_M = 128;
tile-n = 128constexpr int BLOCK_N = 128;
tmaCUtensorMap A_full[MAX_GROUPS];
vector-width = half2reinterpret_cast<half2 *>(c_ptr + out_row0 * N + out_col)[0] =

Kernel source

sub_static.py576 lines
#!POPCORN gpu NVIDIA

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

cuda_src = """
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cstdint>

// ============================================================================
// Constants
// ============================================================================
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int A_SIZE  = BLOCK_M * BLOCK_K / 2;   // 16384
constexpr int B_SIZE  = BLOCK_N * BLOCK_K / 2;   // 16384
constexpr int SFA_SIZE = 128 * BLOCK_K / 16;     // 2048
constexpr int SFB_SIZE = 128 * BLOCK_K / 16;     // 2048
constexpr int MAIN_STAGE = A_SIZE + B_SIZE;       // 32768
constexpr int SF_STAGE   = SFA_SIZE + SFB_SIZE;   // 4096
constexpr int TMAP_SMEM  = 4 * 128;              // 512

constexpr int NUM_STAGES = 6;
constexpr int SMEM_SIZE  = TMAP_SMEM + MAIN_STAGE * NUM_STAGES + SF_STAGE * NUM_STAGES;
constexpr int NUM_MBAR = NUM_STAGES * 2 + 2;      // tma + mma + 2xdone

constexpr int NUM_EP_WARPS = 4;
constexpr int NUM_WARPS = NUM_EP_WARPS + 2;       // 6
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;     // 192
constexpr int TMEM_COLS = 512;
constexpr int MAX_LAUNCH_CTAS = 148;

constexpr int D_TMEM0  = 0;
constexpr int D_TMEM1  = BLOCK_N;                 // 128
constexpr int SFA_TMEM = 2 * BLOCK_N;             // 256
constexpr int SFB_TMEM = SFA_TMEM + 4 * (BLOCK_K / MMA_K); // 272

constexpr uint32_t I_DESC = (1U << 7U) | (1U << 10U) |
    ((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)BLOCK_M >> 7U << 27U);

constexpr int MAX_GROUPS = 8;
constexpr int TMAPS_PER_GROUP = 5;
constexpr int TMAP_A_FULL = 0;
constexpr int TMAP_A_TAIL = 1;
constexpr int TMAP_B = 2;
constexpr int TMAP_SFA = 3;
constexpr int TMAP_SFB = 4;

// ============================================================================
// Device structures
// ============================================================================
struct GroupInfo {
    half* c_ptr;
    int M, N, K;
    int tile_offset;
    int m_tiles, n_tiles;
};

struct KernelParams {
    GroupInfo groups[MAX_GROUPS];
    int num_groups;
    int total_tiles;
    int launch_ctas;
};

struct TmapParamPackG8 {
    CUtensorMap A_full[MAX_GROUPS];
    CUtensorMap A_tail[MAX_GROUPS];
    CUtensorMap B[MAX_GROUPS];
    CUtensorMap SFA[MAX_GROUPS];
    CUtensorMap SFB[MAX_GROUPS];
};

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

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

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

__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size) {
    asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                 :: "r"(mbar_addr), "r"(size) : "memory");
}

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

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

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

__device__ inline void tcgen05_mma_nvfp4(uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,
    int scale_A_tmem, int scale_B_tmem, int enable_input_d, int d_tmem) {
    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));
}

__device__ inline void tcgen05_commit(int mbar_addr) {
    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                 :: "r"(mbar_addr) : "memory");
}

static constexpr char SHAPE_16x256b[] = ".16x256b";
static constexpr char NUM_x1[] = ".x1";

template <const char *SHAPE, const char *NUM>
__device__ inline void tcgen05_ld_4regs(float *tmp, int row, int col) {
    asm volatile("tcgen05.ld.sync.aligned%5%6.b32 "
        "{ %0, %1, %2, %3 }, [%4];"
        : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3])
        : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

__device__ inline void tcgen05_ld_16x256bx1(float *tmp, int row, int col) {
    tcgen05_ld_4regs<SHAPE_16x256b, NUM_x1>(tmp, row, col);
}

__device__ inline void fence_proxy_tensormap(const void *smem_ptr) {
    uint64_t addr = reinterpret_cast<uint64_t>(smem_ptr);
    asm volatile("fence.proxy.tensormap::generic.acquire.gpu [%0], 128;" :: "l"(addr));
}

__device__ inline void fence_proxy_tensormap_release_gpu() {
    asm volatile("fence.proxy.tensormap::generic.release.gpu;" ::: "memory");
}

__device__ inline void tmap_replace_global_address(CUtensorMap *tmap_ptr, uint64_t new_addr) {
    asm volatile("tensormap.replace.tile.global_address.global.b1024.b64 [%0], %1;"
                 :: "l"(tmap_ptr), "l"(new_addr) : "memory");
}

__device__ inline void do_epilogue(int warp_id, int lane_id, int done_mbar, int d_tmem_base,
    half* c_ptr, int M, int N, int off_m, int off_n) {
    mbarrier_wait(done_mbar, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");

    const int col_lane = (lane_id % 4) * 2;
    const int row_lane = lane_id / 4;
    const int residue_m = M - off_m;

    #pragma unroll
    for (int m = 0; m < 2; m++) {
        const int tm = warp_id * 32 + m * 16;
        const int out_row0 = off_m + tm + row_lane;
        const int out_row1 = out_row0 + 8;

        #pragma unroll
        for (int chunk = 0; chunk < BLOCK_N / 8; chunk++) {
            float vals[4];
            tcgen05_ld_16x256bx1(vals, tm, d_tmem_base + chunk * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");

            const int out_col = off_n + chunk * 8 + col_lane;

            if (tm + row_lane < residue_m) {
                reinterpret_cast<half2 *>(c_ptr + out_row0 * N + out_col)[0] =
                    __float22half2_rn({vals[0], vals[1]});
            }
            if (tm + row_lane + 8 < residue_m) {
                reinterpret_cast<half2 *>(c_ptr + out_row1 * N + out_col)[0] =
                    __float22half2_rn({vals[2], vals[3]});
            }
        }
    }
}

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

void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t height, uint64_t width,
                  uint32_t box_h, uint32_t box_w) {
    constexpr uint32_t rank = 3;
    uint64_t globalDim[rank] = {256, height, width / 256};
    uint64_t globalStrides[rank - 1] = {width / 2, 128};
    uint32_t boxDim[rank] = {256, box_h, box_w / 256};
    uint32_t elementStrides[rank] = {1, 1, 1};
    check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B, rank, (void *)ptr,
        globalDim, globalStrides, boxDim, elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
        CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}

// SF reordered tensors have logical shape [32, 4, rest_m, 4, rest_k, L] but are
// a permuted view of a contiguous [L, rest_m, rest_k, 32, 4, 4] allocation.
// Physical memory is thus [rest_m][rest_k][512 bytes], i.e. each 512-byte SF tile
// (covering 128 M-rows x 1 MMA_K=64 step) is already contiguous.
// We encode this as a 3D TMA: dim0 = 256 uint16 (=512B block), dim1 = mn_blocks, dim2 = k_blocks.
void init_SF_tmap(CUtensorMap *tmap, const char *ptr, uint64_t mn, uint64_t K) {
    constexpr uint32_t rank = 3;
    const uint64_t k_blocks = K / 64;
    const uint64_t mn_blocks = (mn + 127) / 128;
    const uint32_t tile_k_blocks = BLOCK_K / 64;          // 4
    constexpr uint64_t SF_BLOCK_BYTES = 512;
    constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t);  // 256
    uint64_t globalDim[rank]       = {X_ELEMS, mn_blocks, k_blocks};
    uint64_t globalStrides[rank-1] = {k_blocks * SF_BLOCK_BYTES, SF_BLOCK_BYTES};
    uint32_t boxDim[rank]          = {(uint32_t)X_ELEMS, 1, tile_k_blocks};
    uint32_t elementStrides[rank]  = {1, 1, 1};
    check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_UINT16, rank, (void *)ptr,
        globalDim, globalStrides, boxDim, elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_NONE,
        CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}

// ============================================================================
// Kernel
// ============================================================================
__global__ __launch_bounds__(TB_SIZE)
void grouped_gemm_kernel(
    const __grid_constant__ KernelParams params,
    const __grid_constant__ TmapParamPackG8 tmap_pack_g8
) {
    struct EpMeta {
        half* c_ptr;
        int M, N;
        int off_m, off_n;
    };

    const int tid = threadIdx.x;
    const int warp_id = tid / WARP_SIZE;
    const int lane_id = tid % WARP_SIZE;
    const int bid = blockIdx.x;
    if (bid >= params.launch_ctas) return;

    const int base_tiles = params.total_tiles / params.launch_ctas;
    const int rem_tiles = params.total_tiles % params.launch_ctas;
    const int my_count = base_tiles + (bid < rem_tiles ? 1 : 0);
    const int my_start = bid * base_tiles + (bid < rem_tiles ? bid : rem_tiles);
    if (my_count <= 0) return;

    // --- SMEM setup ---
    extern __shared__ __align__(1024) char smem_raw[];
    const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));
    const int smem_main = smem + TMAP_SMEM;
    const int smem_sf   = smem_main + MAIN_STAGE * NUM_STAGES;

    #pragma nv_diag_suppress static_var_with_dynamic_init
    __shared__ int64_t mbars[NUM_MBAR];
    __shared__ int32_t tmem_alloc_buf;
    __shared__ EpMeta ep_meta[2];
    const int mbar_base = static_cast<int>(__cvta_generic_to_shared(mbars));
    const int tma_mbar  = mbar_base;
    const int mma_mbar  = tma_mbar + NUM_STAGES * 8;
    const int done_mbar0 = mma_mbar + NUM_STAGES * 8;
    const int done_mbar1 = done_mbar0 + 8;

    // Allocate TMEM once for this CTA.
    if (warp_id == 1) {
        int alloc_addr = static_cast<int>(__cvta_generic_to_shared(&tmem_alloc_buf));
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
                     :: "r"(alloc_addr), "r"(TMEM_COLS));
    }
    __syncthreads();

    // --- Descriptor helpers ---
    auto make_desc_AB = [](int addr) -> uint64_t {
        return desc_encode(addr) | (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
    };
    constexpr int SF_K_PER_BLOCK = BLOCK_K / 64;  // 4
    for (int tile_iter = 0; tile_iter < my_count; tile_iter++) {
        const int slot = tile_iter & 1;
        const int prev_slot = slot ^ 1;
        const int d_tmem_base = slot ? D_TMEM1 : D_TMEM0;
        const int done_mbar = slot ? done_mbar1 : done_mbar0;
        const int prev_d_tmem_base = prev_slot ? D_TMEM1 : D_TMEM0;
        const int prev_done_mbar = prev_slot ? done_mbar1 : done_mbar0;

        const int tile_id = my_start + tile_iter;
        int gidx = 0;
        #pragma unroll
        for (int g = 1; g < MAX_GROUPS; g++) {
            if (g < params.num_groups && tile_id >= params.groups[g].tile_offset)
                gidx = g;
        }
        const GroupInfo& gi = params.groups[gidx];
        const int local_tile = tile_id - gi.tile_offset;
        const int coord_x = local_tile % gi.m_tiles;
        const int coord_y = local_tile / gi.m_tiles;
        const int M = gi.M, N = gi.N, K = gi.K;
        const int num_k = K / BLOCK_K;
        const int off_m = coord_x * BLOCK_M;
        const int off_n = coord_y * BLOCK_N;

        if (warp_id == 0 && lane_id == 0) {
            ep_meta[slot].c_ptr = gi.c_ptr;
            ep_meta[slot].M = M;
            ep_meta[slot].N = N;
            ep_meta[slot].off_m = off_m;
            ep_meta[slot].off_n = off_n;
        }

        // Reset tile-local pipeline barriers before this tile starts.
        if (warp_id == 0 && elect_sync()) {
            #pragma unroll
            for (int i = 0; i < NUM_STAGES; i++) {
                mbarrier_init(tma_mbar + i * 8, 1);
                mbarrier_init(mma_mbar + i * 8, 1);
            }
            mbarrier_init(done_mbar, 1);
            asm volatile("fence.mbarrier_init.release.cluster;");
        }
        __syncthreads();

        const int m_tail = M % BLOCK_M;
        const bool use_A_tail = (coord_x == gi.m_tiles - 1) && (m_tail != 0);
        const int a_box_h = use_A_tail ? m_tail : BLOCK_M;
        const int a_bytes = a_box_h * BLOCK_K / 2;
        const int tma_expect_bytes = a_bytes + B_SIZE + SF_STAGE;

        const void *A_tmap = static_cast<const void *>(
            &(use_A_tail ? tmap_pack_g8.A_tail[gidx] : tmap_pack_g8.A_full[gidx]));
        const void *B_tmap = static_cast<const void *>(&tmap_pack_g8.B[gidx]);
        const void *SFA_tmap = static_cast<const void *>(&tmap_pack_g8.SFA[gidx]);
        const void *SFB_tmap = static_cast<const void *>(&tmap_pack_g8.SFB[gidx]);
        if (warp_id == 0 && lane_id == 0) {
            fence_proxy_tensormap(A_tmap);
            fence_proxy_tensormap(B_tmap);
            fence_proxy_tensormap(SFA_tmap);
            fence_proxy_tensormap(SFB_tmap);
        }
        __syncthreads();

        // TMA producer warp.
        if (warp_id == NUM_WARPS - 2 && elect_sync()) {
            #pragma unroll
            for (int ik = 0; ik < NUM_STAGES && ik < num_k; ik++) {
                int s = ik;
                int A_s = smem_main + s * MAIN_STAGE;
                int B_s = A_s + A_SIZE;
                int SFA_s = smem_sf + s * SF_STAGE;
                int SFB_s = SFA_s + SFA_SIZE;

                tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
                tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);

                int z_sf = ik * SF_K_PER_BLOCK;
                tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
                tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);

                mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
            }

            for (int ik = NUM_STAGES; ik < num_k; ik++) {
                int s = ik % NUM_STAGES;
                mbarrier_wait(mma_mbar + s * 8, (ik / NUM_STAGES - 1) % 2);

                int A_s = smem_main + s * MAIN_STAGE;
                int B_s = A_s + A_SIZE;
                int SFA_s = smem_sf + s * SF_STAGE;
                int SFB_s = SFA_s + SFA_SIZE;

                tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
                tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);

                int z_sf = ik * SF_K_PER_BLOCK;
                tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
                tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);

                mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
            }
        }

        // MMA consumer warp.
        if (warp_id == NUM_WARPS - 1 && elect_sync()) {
            #pragma unroll 1
            for (int ik = 0; ik < num_k; ik++) {
                int s = ik % NUM_STAGES;
                mbarrier_wait(tma_mbar + s * 8, (ik / NUM_STAGES) % 2);

                int A_s   = smem_main + s * MAIN_STAGE;
                int B_s   = A_s + A_SIZE;
                int SFA_s = smem_sf + s * SF_STAGE;
                int SFB_s = SFA_s + SFA_SIZE;

                constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
                uint64_t sfa_desc = sf_base + ((uint64_t)SFA_s >> 4ULL);
                uint64_t sfb_desc = sf_base + ((uint64_t)SFB_s >> 4ULL);

                #pragma unroll
                for (int k = 0; k < BLOCK_K / MMA_K; k++) {
                    tcgen05_cp_nvfp4(SFA_TMEM + k * 4, sfa_desc + (uint64_t)k * (512ULL >> 4ULL));
                    tcgen05_cp_nvfp4(SFB_TMEM + k * 4, sfb_desc + (uint64_t)k * (512ULL >> 4ULL));
                }

                #pragma unroll
                for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
                    uint64_t a_desc = make_desc_AB(A_s + k2 * 32);
                    uint64_t b_desc = make_desc_AB(B_s + k2 * 32);
                    int enable_d = (ik == 0 && k2 == 0) ? 0 : 1;
                    tcgen05_mma_nvfp4(a_desc, b_desc, I_DESC,
                        SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d, d_tmem_base);
                }

                tcgen05_commit(mma_mbar + s * 8);
            }
            tcgen05_commit(done_mbar);
        }

        // Overlap epilogue for previous tile with compute on current tile.
        if (warp_id < NUM_EP_WARPS && tile_iter > 0) {
            EpMeta meta = ep_meta[prev_slot];
            do_epilogue(warp_id, lane_id, prev_done_mbar, prev_d_tmem_base,
                meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);
        }
        __syncthreads();
    }

    // Drain last tile epilogue.
    if (warp_id < NUM_EP_WARPS) {
        int final_slot = (my_count - 1) & 1;
        int final_done_mbar = final_slot ? done_mbar1 : done_mbar0;
        int final_d_tmem_base = final_slot ? D_TMEM1 : D_TMEM0;
        EpMeta meta = ep_meta[final_slot];
        do_epilogue(warp_id, lane_id, final_done_mbar, final_d_tmem_base,
            meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);
    }
    __syncthreads();

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

// ============================================================================
// Host launch
// ============================================================================
void grouped_gemm_impl(
    at::TensorList A_list,
    at::TensorList B_list,
    at::TensorList C_list,
    at::TensorList SFA_list,
    at::TensorList SFB_list
) {
    int G = A_list.size();
    TORCH_CHECK(G <= MAX_GROUPS, "num groups exceeds MAX_GROUPS");
    if (G == 0) return;
    KernelParams params = {};
    params.num_groups = G;

    int total_tiles = 0;
    for (int g = 0; g < G; g++) {
        int Mi = A_list[g].size(0);
        int Ki = A_list[g].size(1) * 2;
        int Ni = B_list[g].size(0);
        int mt = (Mi + BLOCK_M - 1) / BLOCK_M;
        int nt = (Ni + BLOCK_N - 1) / BLOCK_N;
        params.groups[g] = {(half *)C_list[g].data_ptr(), Mi, Ni, Ki, total_tiles, mt, nt};
        total_tiles += mt * nt;
    }
    params.total_tiles = total_tiles;
    params.launch_ctas = total_tiles < MAX_LAUNCH_CTAS ? total_tiles : MAX_LAUNCH_CTAS;

    TmapParamPackG8 tmap_pack_g8 = {};
    for (int g = 0; g < G; g++) {
        int Mi = A_list[g].size(0);
        int Ki = A_list[g].size(1) * 2;
        int Ni = B_list[g].size(0);
        int tail_h = Mi % BLOCK_M;
        if (tail_h == 0) tail_h = BLOCK_M;
        init_AB_tmap(&tmap_pack_g8.A_full[g], (const char *)A_list[g].data_ptr(), Mi, Ki, BLOCK_M, BLOCK_K);
        init_AB_tmap(&tmap_pack_g8.A_tail[g], (const char *)A_list[g].data_ptr(), Mi, Ki, tail_h, BLOCK_K);
        init_AB_tmap(&tmap_pack_g8.B[g], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_N, BLOCK_K);
        init_SF_tmap(&tmap_pack_g8.SFA[g], (const char *)SFA_list[g].data_ptr(), Mi, Ki);
        init_SF_tmap(&tmap_pack_g8.SFB[g], (const char *)SFB_list[g].data_ptr(), Ni, Ki);
    }

    auto kernel = grouped_gemm_kernel;
    static int smem_size = 0;
    if (!smem_size) {
        int dev; cudaGetDevice(&dev);
        int smem_max;
        cudaDeviceGetAttribute(&smem_max, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
        smem_size = smem_max - 1024;
        cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
        cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared);
    }
    kernel<<<params.launch_ctas, TB_SIZE, smem_size>>>(params, tmap_pack_g8);
}

TORCH_LIBRARY(gg_v2_merged_nomemcpy, m) {
    m.def("run(Tensor[] A, Tensor[] B, Tensor[] C, Tensor[] SFA, Tensor[] SFB) -> ()");
    m.impl("run", &grouped_gemm_impl);
}
"""

load_inline(
    "grouped_gemm_v2_merged_nomemcpy_v1",
    cpp_sources="",
    cuda_sources=cuda_src,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3", "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math", "--expt-relaxed-constexpr",
        "--relocatable-device-code=false", "-lineinfo",
    ],
    extra_ldflags=["-lcuda"],
)

_run = torch.ops.gg_v2_merged_nomemcpy.run

def custom_kernel(data: input_t) -> output_t:
    # data = (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
    # sfasfb_reordered has logical shape [32, 4, rest_m, 4, rest_k, L] but is a permuted
    # view of contiguous [L, rest_m, rest_k, 32, 4, 4]. Physical memory is already
    # [rest_m][rest_k][512B tiles] — no host-side permute/contiguous needed.
    abc, _, sf_reordered, _ = data
    a, b, c = zip(*abc)
    sfa, sfb = zip(*sf_reordered)
    _run(list(a), list(b), list(c), list(sfa), list(sfb))
    return list(c)
scrolls · 576 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 483444.

-
#!POPCORN gpu NVIDIA
import torch
⋯ 25 unchanged lines
constexpr int SF_STAGE = SFA_SIZE + SFB_SIZE; // 4096
constexpr int TMAP_SMEM = 4 * 128; // 512
- constexpr int NUM_STAGES = 5;
+ constexpr int NUM_STAGES = 6;
constexpr int SMEM_SIZE = TMAP_SMEM + MAIN_STAGE * NUM_STAGES + SF_STAGE * NUM_STAGES;
- constexpr int NUM_MBAR = NUM_STAGES * 2 + 1; // 11
+ constexpr int NUM_MBAR = NUM_STAGES * 2 + 2; // tma + mma + 2xdone
constexpr int NUM_EP_WARPS = 4;
constexpr int NUM_WARPS = NUM_EP_WARPS + 2; // 6
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE; // 192
constexpr int TMEM_COLS = 512;
+ constexpr int MAX_LAUNCH_CTAS = 148;
- constexpr int D_TMEM = 0;
- constexpr int SFA_TMEM = BLOCK_N; // 128
- constexpr int SFB_TMEM = SFA_TMEM + 4 * (BLOCK_K / MMA_K); // 144
+ constexpr int D_TMEM0 = 0;
+ constexpr int D_TMEM1 = BLOCK_N; // 128
+ constexpr int SFA_TMEM = 2 * BLOCK_N; // 256
+ constexpr int SFB_TMEM = SFA_TMEM + 4 * (BLOCK_K / MMA_K); // 272
constexpr uint32_t I_DESC = (1U << 7U) | (1U << 10U) |
((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)BLOCK_M >> 7U << 27U);
⋯ 19 unchanged lines
struct KernelParams {
GroupInfo groups[MAX_GROUPS];
int num_groups;
+ int total_tiles;
+ int launch_ctas;
};
+ struct TmapParamPackG8 {
+ CUtensorMap A_full[MAX_GROUPS];
+ CUtensorMap A_tail[MAX_GROUPS];
+ CUtensorMap B[MAX_GROUPS];
+ CUtensorMap SFA[MAX_GROUPS];
+ CUtensorMap SFB[MAX_GROUPS];
+ };
+
// ============================================================================
// Inline PTX helpers
// ============================================================================
⋯ 92 unchanged lines
asm volatile("fence.proxy.tensormap::generic.acquire.gpu [%0], 128;" :: "l"(addr));
}
- __device__ inline void do_epilogue(int warp_id, int lane_id, int done_mbar,
- int M, int N, int off_m, int off_n, const GroupInfo& gi) {
+ __device__ inline void fence_proxy_tensormap_release_gpu() {
+ asm volatile("fence.proxy.tensormap::generic.release.gpu;" ::: "memory");
+ }
+
+ __device__ inline void tmap_replace_global_address(CUtensorMap *tmap_ptr, uint64_t new_addr) {
+ asm volatile("tensormap.replace.tile.global_address.global.b1024.b64 [%0], %1;"
+ :: "l"(tmap_ptr), "l"(new_addr) : "memory");
+ }
+
+ __device__ inline void do_epilogue(int warp_id, int lane_id, int done_mbar, int d_tmem_base,
+ half* c_ptr, int M, int N, int off_m, int off_n) {
mbarrier_wait(done_mbar, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
⋯ 10 unchanged lines
#pragma unroll
for (int chunk = 0; chunk < BLOCK_N / 8; chunk++) {
float vals[4];
- tcgen05_ld_16x256bx1(vals, tm, D_TMEM + chunk * 8);
+ tcgen05_ld_16x256bx1(vals, tm, d_tmem_base + chunk * 8);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int out_col = off_n + chunk * 8 + col_lane;
if (tm + row_lane < residue_m) {
- reinterpret_cast<half2 *>(gi.c_ptr + out_row0 * N + out_col)[0] =
+ reinterpret_cast<half2 *>(c_ptr + out_row0 * N + out_col)[0] =
__float22half2_rn({vals[0], vals[1]});
}
if (tm + row_lane + 8 < residue_m) {
- reinterpret_cast<half2 *>(gi.c_ptr + out_row1 * N + out_col)[0] =
+ reinterpret_cast<half2 *>(c_ptr + out_row1 * N + out_col)[0] =
__float22half2_rn({vals[2], vals[3]});
}
}
⋯ 51 unchanged lines
__global__ __launch_bounds__(TB_SIZE)
void grouped_gemm_kernel(
const __grid_constant__ KernelParams params,
- const CUtensorMap* __restrict__ d_tmaps
+ const __grid_constant__ TmapParamPackG8 tmap_pack_g8
) {
+ struct EpMeta {
+ half* c_ptr;
+ int M, N;
+ int off_m, off_n;
+ };
+
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid % WARP_SIZE;
-
- // --- Derive group and tile from blockIdx.x ---
const int bid = blockIdx.x;
- int gidx = 0;
- #pragma unroll
- for (int g = 1; g < MAX_GROUPS; g++) {
- if (g < params.num_groups && bid >= params.groups[g].tile_offset)
- gidx = g;
- }
- const GroupInfo& gi = params.groups[gidx];
- const int local_tile = bid - gi.tile_offset;
- const int coord_x = local_tile % gi.m_tiles;
- const int coord_y = local_tile / gi.m_tiles;
- const int M = gi.M, N = gi.N, K = gi.K;
- const int num_k = K / BLOCK_K;
- const int off_m = coord_x * BLOCK_M;
- const int off_n = coord_y * BLOCK_N;
+ if (bid >= params.launch_ctas) return;
+ const int base_tiles = params.total_tiles / params.launch_ctas;
+ const int rem_tiles = params.total_tiles % params.launch_ctas;
+ const int my_count = base_tiles + (bid < rem_tiles ? 1 : 0);
+ const int my_start = bid * base_tiles + (bid < rem_tiles ? bid : rem_tiles);
+ if (my_count <= 0) return;
+
// --- SMEM setup ---
extern __shared__ __align__(1024) char smem_raw[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));
⋯ 3 unchanged lines
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_MBAR];
__shared__ int32_t tmem_alloc_buf;
+ __shared__ EpMeta ep_meta[2];
const int mbar_base = static_cast<int>(__cvta_generic_to_shared(mbars));
const int tma_mbar = mbar_base;
const int mma_mbar = tma_mbar + NUM_STAGES * 8;
- const int done_mbar = mma_mbar + NUM_STAGES * 8;
+ const int done_mbar0 = mma_mbar + NUM_STAGES * 8;
+ const int done_mbar1 = done_mbar0 + 8;
- // --- INIT: mbarriers + TMEM alloc + tmap copy ---
- if (warp_id == 0 && elect_sync()) {
- #pragma unroll
- for (int i = 0; i < NUM_MBAR; i++)
- mbarrier_init(mbar_base + i * 8, 1);
- asm volatile("fence.mbarrier_init.release.cluster;");
- }
- else if (warp_id == 1) {
+ // Allocate TMEM once for this CTA.
+ if (warp_id == 1) {
int alloc_addr = static_cast<int>(__cvta_generic_to_shared(&tmem_alloc_buf));
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"(alloc_addr), "r"(TMEM_COLS));
}
-
__syncthreads();
- // Tmap pointers in global memory (tensor maps must reside in .param/.const/.global)
- const int m_tail = M % BLOCK_M;
- const bool use_A_tail = (coord_x == gi.m_tiles - 1) && (m_tail != 0);
- const int a_box_h = use_A_tail ? m_tail : BLOCK_M;
- const int a_bytes = a_box_h * BLOCK_K / 2;
- const int tma_expect_bytes = a_bytes + B_SIZE + SF_STAGE;
-
- const CUtensorMap *g_tmaps = d_tmaps + gidx * TMAPS_PER_GROUP;
- const void *A_tmap = static_cast<const void *>(g_tmaps + (use_A_tail ? TMAP_A_TAIL : TMAP_A_FULL));
- const void *B_tmap = static_cast<const void *>(g_tmaps + TMAP_B);
- const void *SFA_tmap = static_cast<const void *>(g_tmaps + TMAP_SFA);
- const void *SFB_tmap = static_cast<const void *>(g_tmaps + TMAP_SFB);
-
// --- Descriptor helpers ---
auto make_desc_AB = [](int addr) -> uint64_t {
return desc_encode(addr) | (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
- auto make_desc_SF = [](int addr) -> uint64_t {
- return desc_encode(addr) | (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
- };
-
constexpr int SF_K_PER_BLOCK = BLOCK_K / 64; // 4
+ for (int tile_iter = 0; tile_iter < my_count; tile_iter++) {
+ const int slot = tile_iter & 1;
+ const int prev_slot = slot ^ 1;
+ const int d_tmem_base = slot ? D_TMEM1 : D_TMEM0;
+ const int done_mbar = slot ? done_mbar1 : done_mbar0;
+ const int prev_d_tmem_base = prev_slot ? D_TMEM1 : D_TMEM0;
+ const int prev_done_mbar = prev_slot ? done_mbar1 : done_mbar0;
- // ========================================================================
- // TMA Producer Warp (warp 4)
- // ========================================================================
- if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- // Prefill
+ const int tile_id = my_start + tile_iter;
+ int gidx = 0;
#pragma unroll
- for (int ik = 0; ik < NUM_STAGES && ik < num_k; ik++) {
- int s = ik;
- int A_s = smem_main + s * MAIN_STAGE;
- int B_s = A_s + A_SIZE;
- int SFA_s = smem_sf + s * SF_STAGE;
- int SFB_s = SFA_s + SFA_SIZE;
+ for (int g = 1; g < MAX_GROUPS; g++) {
+ if (g < params.num_groups && tile_id >= params.groups[g].tile_offset)
+ gidx = g;
+ }
+ const GroupInfo& gi = params.groups[gidx];
+ const int local_tile = tile_id - gi.tile_offset;
+ const int coord_x = local_tile % gi.m_tiles;
+ const int coord_y = local_tile / gi.m_tiles;
+ const int M = gi.M, N = gi.N, K = gi.K;
+ const int num_k = K / BLOCK_K;
+ const int off_m = coord_x * BLOCK_M;
+ const int off_n = coord_y * BLOCK_N;
- tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
- tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);
+ if (warp_id == 0 && lane_id == 0) {
+ ep_meta[slot].c_ptr = gi.c_ptr;
+ ep_meta[slot].M = M;
+ ep_meta[slot].N = N;
+ ep_meta[slot].off_m = off_m;
+ ep_meta[slot].off_n = off_n;
+ }
- int z_sf = ik * SF_K_PER_BLOCK;
- tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
- tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);
+ // Reset tile-local pipeline barriers before this tile starts.
+ if (warp_id == 0 && elect_sync()) {
+ #pragma unroll
+ for (int i = 0; i < NUM_STAGES; i++) {
+ mbarrier_init(tma_mbar + i * 8, 1);
+ mbarrier_init(mma_mbar + i * 8, 1);
+ }
+ mbarrier_init(done_mbar, 1);
+ asm volatile("fence.mbarrier_init.release.cluster;");
+ }
+ __syncthreads();
- mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
+ const int m_tail = M % BLOCK_M;
+ const bool use_A_tail = (coord_x == gi.m_tiles - 1) && (m_tail != 0);
+ const int a_box_h = use_A_tail ? m_tail : BLOCK_M;
+ const int a_bytes = a_box_h * BLOCK_K / 2;
+ const int tma_expect_bytes = a_bytes + B_SIZE + SF_STAGE;
+
+ const void *A_tmap = static_cast<const void *>(
+ &(use_A_tail ? tmap_pack_g8.A_tail[gidx] : tmap_pack_g8.A_full[gidx]));
+ const void *B_tmap = static_cast<const void *>(&tmap_pack_g8.B[gidx]);
+ const void *SFA_tmap = static_cast<const void *>(&tmap_pack_g8.SFA[gidx]);
+ const void *SFB_tmap = static_cast<const void *>(&tmap_pack_g8.SFB[gidx]);
+ if (warp_id == 0 && lane_id == 0) {
+ fence_proxy_tensormap(A_tmap);
+ fence_proxy_tensormap(B_tmap);
+ fence_proxy_tensormap(SFA_tmap);
+ fence_proxy_tensormap(SFB_tmap);
}
+ __syncthreads();
- // Steady state
- for (int ik = NUM_STAGES; ik < num_k; ik++) {
- int s = ik % NUM_STAGES;
- mbarrier_wait(mma_mbar + s * 8, (ik / NUM_STAGES - 1) % 2);
+ // TMA producer warp.
+ if (warp_id == NUM_WARPS - 2 && elect_sync()) {
+ #pragma unroll
+ for (int ik = 0; ik < NUM_STAGES && ik < num_k; ik++) {
+ int s = ik;
+ int A_s = smem_main + s * MAIN_STAGE;
+ int B_s = A_s + A_SIZE;
+ int SFA_s = smem_sf + s * SF_STAGE;
+ int SFB_s = SFA_s + SFA_SIZE;
- int A_s = smem_main + s * MAIN_STAGE;
- int B_s = A_s + A_SIZE;
- int SFA_s = smem_sf + s * SF_STAGE;
- int SFB_s = SFA_s + SFA_SIZE;
+ tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
+ tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);
- tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
- tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);
+ int z_sf = ik * SF_K_PER_BLOCK;
+ tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
+ tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);
- int z_sf = ik * SF_K_PER_BLOCK;
- tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
- tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);
+ mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
+ }
- mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
- }
- }
+ for (int ik = NUM_STAGES; ik < num_k; ik++) {
+ int s = ik % NUM_STAGES;
+ mbarrier_wait(mma_mbar + s * 8, (ik / NUM_STAGES - 1) % 2);
- // ========================================================================
- // MMA Consumer Warp (warp 5)
- // ========================================================================
- if (warp_id == NUM_WARPS - 1 && elect_sync()) {
- #pragma unroll 1
- for (int ik = 0; ik < num_k; ik++) {
- int s = ik % NUM_STAGES;
- mbarrier_wait(tma_mbar + s * 8, (ik / NUM_STAGES) % 2);
+ int A_s = smem_main + s * MAIN_STAGE;
+ int B_s = A_s + A_SIZE;
+ int SFA_s = smem_sf + s * SF_STAGE;
+ int SFB_s = SFA_s + SFA_SIZE;
- int A_s = smem_main + s * MAIN_STAGE;
- int B_s = A_s + A_SIZE;
- int SFA_s = smem_sf + s * SF_STAGE;
- int SFB_s = SFA_s + SFA_SIZE;
+ tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, 0);
+ tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, 0);
- // Copy scale factors smem -> tmem
- constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
- uint64_t sfa_desc = sf_base + ((uint64_t)SFA_s >> 4ULL);
- uint64_t sfb_desc = sf_base + ((uint64_t)SFB_s >> 4ULL);
+ int z_sf = ik * SF_K_PER_BLOCK;
+ tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, 0);
+ tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, 0);
- #pragma unroll
- for (int k = 0; k < BLOCK_K / MMA_K; k++) {
- tcgen05_cp_nvfp4(SFA_TMEM + k * 4, sfa_desc + (uint64_t)k * (512ULL >> 4ULL));
- tcgen05_cp_nvfp4(SFB_TMEM + k * 4, sfb_desc + (uint64_t)k * (512ULL >> 4ULL));
+ mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);
}
+ }
- // MMA: BLOCK_K=256 = 1 × 256, so k1=0 only, k2=0..3
- #pragma unroll
- for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
- uint64_t a_desc = make_desc_AB(A_s + k2 * 32);
- uint64_t b_desc = make_desc_AB(B_s + k2 * 32);
- int enable_d = (ik == 0 && k2 == 0) ? 0 : 1;
- tcgen05_mma_nvfp4(a_desc, b_desc, I_DESC,
- SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d, D_TMEM);
+ // MMA consumer warp.
+ if (warp_id == NUM_WARPS - 1 && elect_sync()) {
+ #pragma unroll 1
+ for (int ik = 0; ik < num_k; ik++) {
+ int s = ik % NUM_STAGES;
+ mbarrier_wait(tma_mbar + s * 8, (ik / NUM_STAGES) % 2);
+
+ int A_s = smem_main + s * MAIN_STAGE;
+ int B_s = A_s + A_SIZE;
+ int SFA_s = smem_sf + s * SF_STAGE;
+ int SFB_s = SFA_s + SFA_SIZE;
+
+ constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);
+ uint64_t sfa_desc = sf_base + ((uint64_t)SFA_s >> 4ULL);
+ uint64_t sfb_desc = sf_base + ((uint64_t)SFB_s >> 4ULL);
+
+ #pragma unroll
+ for (int k = 0; k < BLOCK_K / MMA_K; k++) {
+ tcgen05_cp_nvfp4(SFA_TMEM + k * 4, sfa_desc + (uint64_t)k * (512ULL >> 4ULL));
+ tcgen05_cp_nvfp4(SFB_TMEM + k * 4, sfb_desc + (uint64_t)k * (512ULL >> 4ULL));
+ }
+
+ #pragma unroll
+ for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
+ uint64_t a_desc = make_desc_AB(A_s + k2 * 32);
+ uint64_t b_desc = make_desc_AB(B_s + k2 * 32);
+ int enable_d = (ik == 0 && k2 == 0) ? 0 : 1;
+ tcgen05_mma_nvfp4(a_desc, b_desc, I_DESC,
+ SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d, d_tmem_base);
+ }
+
+ tcgen05_commit(mma_mbar + s * 8);
}
+ tcgen05_commit(done_mbar);
+ }
- tcgen05_commit(mma_mbar + s * 8);
+ // Overlap epilogue for previous tile with compute on current tile.
+ if (warp_id < NUM_EP_WARPS && tile_iter > 0) {
+ EpMeta meta = ep_meta[prev_slot];
+ do_epilogue(warp_id, lane_id, prev_done_mbar, prev_d_tmem_base,
+ meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);
}
- tcgen05_commit(done_mbar);
+ __syncthreads();
}
- // ========================================================================
- // Epilogue: warps 0-3, f32 -> f16, predicated store
- // ========================================================================
+ // Drain last tile epilogue.
if (warp_id < NUM_EP_WARPS) {
- do_epilogue(warp_id, lane_id, done_mbar, M, N, off_m, off_n, gi);
- mbarrier_wait(done_mbar, 0);
+ int final_slot = (my_count - 1) & 1;
+ int final_done_mbar = final_slot ? done_mbar1 : done_mbar0;
+ int final_d_tmem_base = final_slot ? D_TMEM1 : D_TMEM0;
+ EpMeta meta = ep_meta[final_slot];
+ do_epilogue(warp_id, lane_id, final_done_mbar, final_d_tmem_base,
+ meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);
}
-
__syncthreads();
if (warp_id == 0)
⋯ 3 unchanged lines
// ============================================================================
// Host launch
// ============================================================================
- constexpr int TMAP_CACHE_CAPACITY = 32;
-
- struct TmapCacheEntry {
- uint64_t key;
- int groups;
- CUtensorMap* d_tmaps;
- bool valid;
- };
-
- static TmapCacheEntry s_tmap_cache[TMAP_CACHE_CAPACITY] = {};
- static int s_tmap_cache_rr = 0;
-
- inline uint64_t hash_u64(uint64_t h, uint64_t v) {
- return h ^ (v + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
- }
-
void grouped_gemm_impl(
at::TensorList A_list,
at::TensorList B_list,
⋯ 6 unchanged lines
if (G == 0) return;
KernelParams params = {};
params.num_groups = G;
- uint64_t key = 0x9e3779b97f4a7c15ULL;
- key = hash_u64(key, static_cast<uint64_t>(G));
int total_tiles = 0;
for (int g = 0; g < G; g++) {
⋯ 4 unchanged lines
int nt = (Ni + BLOCK_N - 1) / BLOCK_N;
params.groups[g] = {(half *)C_list[g].data_ptr(), Mi, Ni, Ki, total_tiles, mt, nt};
total_tiles += mt * nt;
-
- key = hash_u64(key, static_cast<uint64_t>(Mi));
- key = hash_u64(key, static_cast<uint64_t>(Ni));
- key = hash_u64(key, static_cast<uint64_t>(Ki));
- key = hash_u64(key, static_cast<uint64_t>(reinterpret_cast<uintptr_t>(A_list[g].data_ptr())));
- key = hash_u64(key, static_cast<uint64_t>(reinterpret_cast<uintptr_t>(B_list[g].data_ptr())));
- key = hash_u64(key, static_cast<uint64_t>(reinterpret_cast<uintptr_t>(SFA_list[g].data_ptr())));
- key = hash_u64(key, static_cast<uint64_t>(reinterpret_cast<uintptr_t>(SFB_list[g].data_ptr())));
}
+ params.total_tiles = total_tiles;
+ params.launch_ctas = total_tiles < MAX_LAUNCH_CTAS ? total_tiles : MAX_LAUNCH_CTAS;
- int cache_idx = -1;
- for (int i = 0; i < TMAP_CACHE_CAPACITY; i++) {
- if (s_tmap_cache[i].valid &&
- s_tmap_cache[i].groups == G &&
- s_tmap_cache[i].key == key) {
- cache_idx = i;
- break;
- }
+ TmapParamPackG8 tmap_pack_g8 = {};
+ for (int g = 0; g < G; g++) {
+ int Mi = A_list[g].size(0);
+ int Ki = A_list[g].size(1) * 2;
+ int Ni = B_list[g].size(0);
+ int tail_h = Mi % BLOCK_M;
+ if (tail_h == 0) tail_h = BLOCK_M;
+ init_AB_tmap(&tmap_pack_g8.A_full[g], (const char *)A_list[g].data_ptr(), Mi, Ki, BLOCK_M, BLOCK_K);
+ init_AB_tmap(&tmap_pack_g8.A_tail[g], (const char *)A_list[g].data_ptr(), Mi, Ki, tail_h, BLOCK_K);
+ init_AB_tmap(&tmap_pack_g8.B[g], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_N, BLOCK_K);
+ init_SF_tmap(&tmap_pack_g8.SFA[g], (const char *)SFA_list[g].data_ptr(), Mi, Ki);
+ init_SF_tmap(&tmap_pack_g8.SFB[g], (const char *)SFB_list[g].data_ptr(), Ni, Ki);
}
- if (cache_idx < 0) {
- for (int i = 0; i < TMAP_CACHE_CAPACITY; i++) {
- if (!s_tmap_cache[i].valid) {
- cache_idx = i;
- break;
- }
- }
- if (cache_idx < 0) {
- cache_idx = s_tmap_cache_rr;
- s_tmap_cache_rr = (s_tmap_cache_rr + 1) % TMAP_CACHE_CAPACITY;
- }
-
- TmapCacheEntry &entry = s_tmap_cache[cache_idx];
- if (!entry.d_tmaps) {
- cudaMalloc(&entry.d_tmaps, MAX_GROUPS * TMAPS_PER_GROUP * sizeof(CUtensorMap));
- }
-
- CUtensorMap h_tmaps[MAX_GROUPS * TMAPS_PER_GROUP];
- for (int g = 0; g < G; g++) {
- int Mi = A_list[g].size(0);
- int Ki = A_list[g].size(1) * 2;
- int Ni = B_list[g].size(0);
- int tail_h = Mi % BLOCK_M;
- if (tail_h == 0) tail_h = BLOCK_M;
-
- init_AB_tmap(&h_tmaps[g * TMAPS_PER_GROUP + TMAP_A_FULL], (const char *)A_list[g].data_ptr(), Mi, Ki, BLOCK_M, BLOCK_K);
- init_AB_tmap(&h_tmaps[g * TMAPS_PER_GROUP + TMAP_A_TAIL], (const char *)A_list[g].data_ptr(), Mi, Ki, tail_h, BLOCK_K);
- init_AB_tmap(&h_tmaps[g * TMAPS_PER_GROUP + TMAP_B], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_N, BLOCK_K);
- init_SF_tmap(&h_tmaps[g * TMAPS_PER_GROUP + TMAP_SFA], (const char *)SFA_list[g].data_ptr(), Mi, Ki);
- init_SF_tmap(&h_tmaps[g * TMAPS_PER_GROUP + TMAP_SFB], (const char *)SFB_list[g].data_ptr(), Ni, Ki);
- }
-
- cudaMemcpy(entry.d_tmaps, h_tmaps, G * TMAPS_PER_GROUP * sizeof(CUtensorMap), cudaMemcpyHostToDevice);
- entry.key = key;
- entry.groups = G;
- entry.valid = true;
- }
-
- CUtensorMap* d_tmaps = s_tmap_cache[cache_idx].d_tmaps;
-
auto kernel = grouped_gemm_kernel;
static int smem_size = 0;
if (!smem_size) {
⋯ 4 unchanged lines
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared);
}
- kernel<<<total_tiles, TB_SIZE, smem_size>>>(params, d_tmaps);
+ kernel<<<params.launch_ctas, TB_SIZE, smem_size>>>(params, tmap_pack_g8);
}
- TORCH_LIBRARY(gg_v2_baseline, m) {
+ TORCH_LIBRARY(gg_v2_merged_nomemcpy, m) {
m.def("run(Tensor[] A, Tensor[] B, Tensor[] C, Tensor[] SFA, Tensor[] SFB) -> ()");
m.impl("run", &grouped_gemm_impl);
}
"""
load_inline(
- "grouped_gemm_v2_baseline_v1",
+ "grouped_gemm_v2_merged_nomemcpy_v1",
cpp_sources="",
cuda_sources=cuda_src,
is_python_module=False,
⋯ 6 unchanged lines
extra_ldflags=["-lcuda"],
)
- _run = torch.ops.gg_v2_baseline.run
+ _run = torch.ops.gg_v2_merged_nomemcpy.run
def custom_kernel(data: input_t) -> output_t:
# data = (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
scrolls · 557 diff lines total

Best evidence level for this revision: reported

JSON