Skip to content
KernelIndex
Search⌘K

submission 227573

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
34.3µs
#266 of 420
2025-12-28

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

mbarriervoid mbarrier_init(int mbar_addr, int count) {
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;"
tile-n = 64constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
tma"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "

Kernel source

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

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

# ============================================================================
# Fused Dual GEMM: Load A ONCE, compute both GEMMs in parallel
# C = silu(A @ B1) * (A @ B2)
#
# Key optimization: A is loaded only once per K-iteration, then used for
# both B1 and B2 matrix multiplications.
# ============================================================================

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

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

constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

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

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

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

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

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

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

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

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

// MMA with explicit destination TMEM address
__device__ __forceinline__
void tcgen05_mma_nvfp4_at(
    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::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__ __forceinline__
void tcgen05_ld_32x32bx32(float *tmp, int row, int col) {
    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"((row << 16) | col)
    );
}

__device__ __forceinline__
void tcgen05_ld_32x32bx64(float *tmp, int row, int col) {
    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"((row << 16) | col)
    );
}

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

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

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

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

// ============================================================================
// Fused Dual GEMM Kernel - Load A ONCE, compute both GEMMs
// 
// SMEM Layout per stage:
//   [A tile][B1 tile][B2 tile][SFA][SFB1][SFB2]
//
// TMEM Layout:
//   [ACC1: 0..BLOCK_N-1][ACC2: BLOCK_N..2*BLOCK_N-1][SFA][SFB1][SFB2]
// ============================================================================

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
__global__
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_fused_kernel(
    const __grid_constant__ CUtensorMap A_tmap,
    const __grid_constant__ CUtensorMap B1_tmap,
    const __grid_constant__ CUtensorMap B2_tmap,
    const char *SFA_ptr,
    const char *SFB1_ptr,
    const char *SFB2_ptr,
    half *C_ptr,
    int M, int N
) {
    const int tid = threadIdx.x;
    const int bid = blockIdx.x;
    const int warp_id = tid / WARP_SIZE;

    const int grid_n = N / BLOCK_N;
    const int bid_m = bid / grid_n;
    const int bid_n = bid % grid_n;
    const int off_m = bid_m * BLOCK_M;
    const int off_n = bid_n * BLOCK_N;

    constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

    extern __shared__ __align__(1024) char smem_ptr[];
    const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
    
    // SMEM sizes - now with A, B1, B2, SFA, SFB1, SFB2 per stage
    constexpr int A_size    = BLOCK_M * BLOCK_K / 2;
    constexpr int B1_size   = BLOCK_N * BLOCK_K / 2;
    constexpr int B2_size   = BLOCK_N * BLOCK_K / 2;
    constexpr int SFA_size  = 128 * BLOCK_K / 16;
    constexpr int SFB1_size = 128 * BLOCK_K / 16;
    constexpr int SFB2_size = 128 * BLOCK_K / 16;
    constexpr int STAGE_SIZE = A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size;

    #pragma nv_diag_suppress static_var_with_dynamic_init
    __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
    const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
    const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
    const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

    // TMEM layout: 
    //   ACC1: columns [0, BLOCK_N)
    //   ACC2: columns [BLOCK_N, 2*BLOCK_N)
    //   SFA:  columns [2*BLOCK_N, 2*BLOCK_N + 16)
    //   SFB1: columns [2*BLOCK_N + 16, 2*BLOCK_N + 32)
    //   SFB2: columns [2*BLOCK_N + 32, 2*BLOCK_N + 48)
    constexpr int ACC1_tmem = 0;
    constexpr int ACC2_tmem = BLOCK_N;
    constexpr int SFA_tmem  = 2 * BLOCK_N;
    constexpr int SFB1_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
    constexpr int SFB2_tmem = SFB1_tmem + 4 * (BLOCK_K / MMA_K);
    // TMEM allocation must be power of 2: 64*2 + 48 = 176 -> round up to 256
    constexpr int TOTAL_TMEM_COLS = 256;

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

    constexpr int num_iters = K / BLOCK_K;

    // ========================================================================
    // TMA Warp: Load A, B1, B2, SFA, SFB1, SFB2 per iteration
    // Key: A is loaded ONCE and used for BOTH GEMMs!
    // ========================================================================
    if (warp_id == NUM_WARPS - 2 && elect_sync()) {
        uint64_t cache_A = EVICT_FIRST;   // A reused twice, evict early
        uint64_t cache_B = EVICT_LAST;    // B tiles used once each

        auto issue_tma = [&](int iter_k, int stage_id) {
            const int mbar_addr = tma_mbar_addr + stage_id * 8;
            const int A_smem   = smem + stage_id * STAGE_SIZE;
            const int B1_smem  = A_smem + A_size;
            const int B2_smem  = B1_smem + B1_size;
            const int SFA_smem = B2_smem + B2_size;
            const int SFB1_smem = SFA_smem + SFA_size;
            const int SFB2_smem = SFB1_smem + SFB1_size;

            const int off_k = iter_k * BLOCK_K;

            // Load A (shared between both GEMMs)
            tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
            
            // Load B1 and B2
            tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
            tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

            // Load scale factors
            const int rest_k = K / 16 / 4;
            const char *SFA_src  = SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512;
            const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
            const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
            
            tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
            tma_gmem2smem(SFB1_smem, SFB1_src, SFB1_size, mbar_addr, cache_B);
            tma_gmem2smem(SFB2_smem, SFB2_src, SFB2_size, mbar_addr, cache_B);

            asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                        :: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
        };

        for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++)
            issue_tma(iter_k, iter_k);

        for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
            const int stage_id = iter_k % NUM_STAGES;
            const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
            mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
            issue_tma(iter_k, stage_id);
        }
    }
    // ========================================================================
    // MMA Warp: Execute BOTH GEMMs using the SAME A tile
    // ========================================================================
    else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
        constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U)
                                  | ((uint32_t)BLOCK_N >> 3U << 17U)
                                  | ((uint32_t)128 >> 7U << 27U);

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
            const int stage_id = iter_k % NUM_STAGES;
            const int tma_phase = (iter_k / NUM_STAGES) % 2;
            mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

            const int A_smem   = smem + stage_id * STAGE_SIZE;
            const int B1_smem  = A_smem + A_size;
            const int B2_smem  = B1_smem + B1_size;
            const int SFA_smem = B2_smem + B2_size;
            const int SFB1_smem = SFA_smem + SFA_size;
            const int SFB2_smem = SFB1_smem + SFB1_size;

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

            constexpr uint64_t SF_desc = make_desc_SF(0);
            const uint64_t SFA_desc  = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
            const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
            const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

            // Copy ALL scale factors to TMEM
            for (int k = 0; k < BLOCK_K / MMA_K; k++) {
                tcgen05_cp_nvfp4(SFA_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
                tcgen05_cp_nvfp4(SFB1_tmem + k * 4, SFB1_desc + (uint64_t)k * (512ULL >> 4ULL));
                tcgen05_cp_nvfp4(SFB2_tmem + k * 4, SFB2_desc + (uint64_t)k * (512ULL >> 4ULL));
            }

            // Execute BOTH MMAs using the SAME A tile
            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 + k1 * BLOCK_M * 128 + k2 * 32);
                    uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
                    uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

                    int k_sf = k1 * 4 + k2;
                    const int scale_A  = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
                    const int scale_B1 = SFB1_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
                    const int scale_B2 = SFB2_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);

                    const int enable_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
                    
                    // GEMM1: A @ B1 -> ACC1
                    tcgen05_mma_nvfp4_at(ACC1_tmem, a_desc, b1_desc, i_desc, scale_A, scale_B1, enable_d);
                    
                    // GEMM2: A @ B2 -> ACC2 (using same A descriptor!)
                    tcgen05_mma_nvfp4_at(ACC2_tmem, a_desc, b2_desc, i_desc, scale_A, scale_B2, enable_d);
                }

            asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                        :: "r"(mma_mbar_addr + stage_id * 8) : "memory");
        }
        asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                    :: "r"(mainloop_mbar_addr) : "memory");
    }
    // ========================================================================
    // Epilogue Warps: Load ACC1 and ACC2, apply SiLU + multiply, store
    // ========================================================================
    else if (tid < BLOCK_M) {
        mbarrier_wait(mainloop_mbar_addr, 0);
        asm volatile("tcgen05.fence::after_thread_sync;");

        constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
        
        for (int n = 0; n < BLOCK_N / WIDTH; n++) {
            float acc1[WIDTH];
            float acc2[WIDTH];
            
            // Load both accumulators
            if constexpr (WIDTH == 64) {
                tcgen05_ld_32x32bx64(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
                tcgen05_ld_32x32bx64(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
            } else {
                tcgen05_ld_32x32bx32(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
                tcgen05_ld_32x32bx32(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
            }
            asm volatile("tcgen05.wait::ld.sync.aligned;");

            // Fused SiLU + multiply: C = silu(acc1) * acc2
            for (int i = 0; i < WIDTH; i++) {
                float x = acc1[i];
                float silu = x / (1.0f + __expf(-x));
                float result = silu * acc2[i];
                
                const int row = off_m + tid;
                const int col = off_n + n * WIDTH + i;
                C_ptr[row * N + col] = __float2half(result);
            }
        }

        asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
        if (warp_id == 0)
            asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" 
                         :: "r"(0), "r"(256));  // Must match power-of-2 allocation
    }
}

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

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
at::Tensor dual_gemm_launch(
    const at::Tensor& A,
    const at::Tensor& B1,
    const at::Tensor& B2,
    const at::Tensor& SFA,
    const at::Tensor& SFB1,
    const at::Tensor& SFB2,
    at::Tensor& C
) {
    const int M = A.size(0);
    const int N = B1.size(0);

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

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

    dim3 grid((M / BLOCK_M) * (N / BLOCK_N), 1, 1);
    int tb_size = BLOCK_M + 2 * WARP_SIZE;
    
    // New SMEM size with dual B tiles
    constexpr int A_size    = BLOCK_M * BLOCK_K / 2;
    constexpr int B1_size   = BLOCK_N * BLOCK_K / 2;
    constexpr int B2_size   = BLOCK_N * BLOCK_K / 2;
    constexpr int SFA_size  = 128 * BLOCK_K / 16;
    constexpr int SFB1_size = 128 * BLOCK_K / 16;
    constexpr int SFB2_size = 128 * BLOCK_K / 16;
    int smem_size = (A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size) * NUM_STAGES;

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

    kernel_fn<<<grid, tb_size, smem_size>>>(
        A_tmap, B1_tmap, B2_tmap, SFA_ptr, SFB1_ptr, SFB2_ptr, C_ptr, M, N
    );

    return C;
}

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

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

    // Main benchmark configs - use NUM_STAGES=4 for dual GEMM (more SMEM per stage)
    LAUNCH(7168, 128, 64, 256, 4)
    LAUNCH(4096, 128, 64, 256, 4)
    
    // Other K values
    LAUNCH(256,  128, 64, 256, 1)
    LAUNCH(512,  128, 64, 256, 2)
    LAUNCH(1024, 128, 64, 256, 3)
    LAUNCH(1536, 128, 64, 256, 4)
    LAUNCH(2048, 128, 64, 256, 4)
    LAUNCH(2304, 128, 64, 256, 4)

#undef LAUNCH

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

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

_compiled_module = None

def _get_module():
    global _compiled_module
    if _compiled_module is None:
        _compiled_module = load_inline(
            "dual_gemm_cuda",
            cpp_sources="",
            cuda_sources=CUDA_SOURCE,
            verbose=True,
            is_python_module=False,
            extra_cuda_cflags=[
                "-O3",
                "-gencode=arch=compute_100a,code=sm_100a",
                "--use_fast_math",
                "--expt-relaxed-constexpr",
                "--relocatable-device-code=false",
                "-lineinfo",
            ],
            extra_ldflags=["-lcuda"],
        )
    return _compiled_module


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
    _get_module()
    result = torch.ops.dual_gemm_module.dual_gemm(
        a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c
    )
    return result
scrolls · 583 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 227504.

⋯ 5 unchanged lines
from torch.utils.cpp_extension import load_inline
# ============================================================================
- # Dual GEMM: Two separate single GEMMs + fused SiLU
- # Uses the proven winner's approach for each GEMM
+ # Fused Dual GEMM: Load A ONCE, compute both GEMMs in parallel
# C = silu(A @ B1) * (A @ B2)
+ #
+ # Key optimization: A is loaded only once per K-iteration, then used for
+ # both B1 and B2 matrix multiplications.
# ============================================================================
CUDA_SOURCE = r"""
⋯ 10 unchanged lines
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
// ============================================================================
- // PTX Helper Functions (from winner's implementation)
+ // PTX Helper Functions
// ============================================================================
__device__ __forceinline__
⋯ 66 unchanged lines
:: "r"(taddr), "l"(s_desc));
}
+ // MMA with explicit destination TMEM address
__device__ __forceinline__
- void tcgen05_mma_nvfp4(
+ void tcgen05_mma_nvfp4_at(
+ int d_tmem,
uint64_t a_desc,
uint64_t b_desc,
uint32_t i_desc,
⋯ 1 unchanged lines
int scale_B_tmem,
int enable_input_d
) {
- const int d_tmem = 0;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
⋯ 102 unchanged lines
}
// ============================================================================
- // Single GEMM Kernel (proven winner's approach)
+ // Fused Dual GEMM Kernel - Load A ONCE, compute both GEMMs
+ //
+ // SMEM Layout per stage:
+ // [A tile][B1 tile][B2 tile][SFA][SFB1][SFB2]
+ //
+ // TMEM Layout:
+ // [ACC1: 0..BLOCK_N-1][ACC2: BLOCK_N..2*BLOCK_N-1][SFA][SFB1][SFB2]
// ============================================================================
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
__global__
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
- void single_gemm_kernel(
+ void dual_gemm_fused_kernel(
const __grid_constant__ CUtensorMap A_tmap,
- const __grid_constant__ CUtensorMap B_tmap,
+ const __grid_constant__ CUtensorMap B1_tmap,
+ const __grid_constant__ CUtensorMap B2_tmap,
const char *SFA_ptr,
- const char *SFB_ptr,
- float *C_ptr, // Output is float32 for fused operation
+ const char *SFB1_ptr,
+ const char *SFB2_ptr,
+ half *C_ptr,
int M, int N
) {
const int tid = threadIdx.x;
⋯ 11 unchanged lines
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
- constexpr int A_size = BLOCK_M * BLOCK_K / 2;
- constexpr int B_size = BLOCK_N * BLOCK_K / 2;
- constexpr int SFA_size = 128 * BLOCK_K / 16;
- constexpr int SFB_size = 128 * BLOCK_K / 16;
- constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
+ // SMEM sizes - now with A, B1, B2, SFA, SFB1, SFB2 per stage
+ constexpr int A_size = BLOCK_M * BLOCK_K / 2;
+ constexpr int B1_size = BLOCK_N * BLOCK_K / 2;
+ constexpr int B2_size = BLOCK_N * BLOCK_K / 2;
+ constexpr int SFA_size = 128 * BLOCK_K / 16;
+ constexpr int SFB1_size = 128 * BLOCK_K / 16;
+ constexpr int SFB2_size = 128 * BLOCK_K / 16;
+ constexpr int STAGE_SIZE = A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size;
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
⋯ 1 unchanged lines
const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;
- constexpr int SFA_tmem = BLOCK_N;
- constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
+ // TMEM layout:
+ // ACC1: columns [0, BLOCK_N)
+ // ACC2: columns [BLOCK_N, 2*BLOCK_N)
+ // SFA: columns [2*BLOCK_N, 2*BLOCK_N + 16)
+ // SFB1: columns [2*BLOCK_N + 16, 2*BLOCK_N + 32)
+ // SFB2: columns [2*BLOCK_N + 32, 2*BLOCK_N + 48)
+ constexpr int ACC1_tmem = 0;
+ constexpr int ACC2_tmem = BLOCK_N;
+ constexpr int SFA_tmem = 2 * BLOCK_N;
+ constexpr int SFB1_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
+ constexpr int SFB2_tmem = SFB1_tmem + 4 * (BLOCK_K / MMA_K);
+ // TMEM allocation must be power of 2: 64*2 + 48 = 176 -> round up to 256
+ constexpr int TOTAL_TMEM_COLS = 256;
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NUM_STAGES * 2 + 1; i++)
⋯ 2 unchanged lines
}
else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
- :: "r"(smem), "r"(BLOCK_N * 2));
+ :: "r"(smem), "r"(TOTAL_TMEM_COLS));
}
__syncthreads();
constexpr int num_iters = K / BLOCK_K;
- // TMA Warp
+ // ========================================================================
+ // TMA Warp: Load A, B1, B2, SFA, SFB1, SFB2 per iteration
+ // Key: A is loaded ONCE and used for BOTH GEMMs!
+ // ========================================================================
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- uint64_t cache_A = (M > N) ? EVICT_FIRST : EVICT_LAST;
- uint64_t cache_B = (M > N) ? EVICT_LAST : EVICT_FIRST;
+ uint64_t cache_A = EVICT_FIRST; // A reused twice, evict early
+ uint64_t cache_B = EVICT_LAST; // B tiles used once each
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
- const int A_smem = smem + stage_id * STAGE_SIZE;
- const int B_smem = A_smem + A_size;
- const int SFA_smem = B_smem + B_size;
- const int SFB_smem = SFA_smem + SFA_size;
+ const int A_smem = smem + stage_id * STAGE_SIZE;
+ const int B1_smem = A_smem + A_size;
+ const int B2_smem = B1_smem + B1_size;
+ const int SFA_smem = B2_smem + B2_size;
+ const int SFB1_smem = SFA_smem + SFA_size;
+ const int SFB2_smem = SFB1_smem + SFB1_size;
const int off_k = iter_k * BLOCK_K;
+
+ // Load A (shared between both GEMMs)
tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
- tma_3d_gmem2smem(B_smem, &B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
+
+ // Load B1 and B2
+ tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
+ tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
+ // Load scale factors
const int rest_k = K / 16 / 4;
- const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512;
- const char *SFB_src = SFB_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
+ const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512;
+ const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
+ const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
+
tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
- tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);
+ tma_gmem2smem(SFB1_smem, SFB1_src, SFB1_size, mbar_addr, cache_B);
+ tma_gmem2smem(SFB2_smem, SFB2_src, SFB2_size, mbar_addr, cache_B);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
⋯ 9 unchanged lines
issue_tma(iter_k, stage_id);
}
}
- // MMA Warp
+ // ========================================================================
+ // MMA Warp: Execute BOTH GEMMs using the SAME A tile
+ // ========================================================================
else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U)
| ((uint32_t)BLOCK_N >> 3U << 17U)
⋯ 4 unchanged lines
const int tma_phase = (iter_k / NUM_STAGES) % 2;
mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);
- const int A_smem = smem + stage_id * STAGE_SIZE;
- const int B_smem = A_smem + A_size;
- const int SFA_smem = B_smem + B_size;
- const int SFB_smem = SFA_smem + SFA_size;
+ const int A_smem = smem + stage_id * STAGE_SIZE;
+ const int B1_smem = A_smem + A_size;
+ const int B2_smem = B1_smem + B1_size;
+ const int SFA_smem = B2_smem + B2_size;
+ const int SFB1_smem = SFA_smem + SFA_size;
+ const int SFB2_smem = SFB1_smem + SFB1_size;
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * 128;
⋯ 5 unchanged lines
};
constexpr uint64_t SF_desc = make_desc_SF(0);
- const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
- const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);
+ const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
+ const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
+ const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);
+ // Copy ALL scale factors to TMEM
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));
+ tcgen05_cp_nvfp4(SFB1_tmem + k * 4, SFB1_desc + (uint64_t)k * (512ULL >> 4ULL));
+ tcgen05_cp_nvfp4(SFB2_tmem + k * 4, SFB2_desc + (uint64_t)k * (512ULL >> 4ULL));
}
+ // Execute BOTH MMAs using the SAME A tile
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 + k1 * BLOCK_M * 128 + k2 * 32);
- uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);
+ uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
+ uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
+ uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);
int k_sf = k1 * 4 + k2;
- const int scale_A = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
- const int scale_B = SFB_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
+ const int scale_A = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
+ const int scale_B1 = SFB1_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
+ const int scale_B2 = SFB2_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
const int enable_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
- tcgen05_mma_nvfp4(a_desc, b_desc, i_desc, scale_A, scale_B, enable_d);
+
+ // GEMM1: A @ B1 -> ACC1
+ tcgen05_mma_nvfp4_at(ACC1_tmem, a_desc, b1_desc, i_desc, scale_A, scale_B1, enable_d);
+
+ // GEMM2: A @ B2 -> ACC2 (using same A descriptor!)
+ tcgen05_mma_nvfp4_at(ACC2_tmem, a_desc, b2_desc, i_desc, scale_A, scale_B2, enable_d);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
⋯ 2 unchanged lines
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mainloop_mbar_addr) : "memory");
}
- // Epilogue Warps - store as float32
+ // ========================================================================
+ // Epilogue Warps: Load ACC1 and ACC2, apply SiLU + multiply, store
+ // ========================================================================
else if (tid < BLOCK_M) {
mbarrier_wait(mainloop_mbar_addr, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
⋯ 1 unchanged lines
constexpr int WIDTH = (BLOCK_N <= 64) ? BLOCK_N : 64;
for (int n = 0; n < BLOCK_N / WIDTH; n++) {
- float tmp[WIDTH];
+ float acc1[WIDTH];
+ float acc2[WIDTH];
+
+ // Load both accumulators
if constexpr (WIDTH == 64) {
- tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH);
+ tcgen05_ld_32x32bx64(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
+ tcgen05_ld_32x32bx64(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
} else {
- tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH);
+ tcgen05_ld_32x32bx32(acc1, warp_id * 32, ACC1_tmem + n * WIDTH);
+ tcgen05_ld_32x32bx32(acc2, warp_id * 32, ACC2_tmem + n * WIDTH);
}
asm volatile("tcgen05.wait::ld.sync.aligned;");
+ // Fused SiLU + multiply: C = silu(acc1) * acc2
for (int i = 0; i < WIDTH; i++) {
+ float x = acc1[i];
+ float silu = x / (1.0f + __expf(-x));
+ float result = silu * acc2[i];
+
const int row = off_m + tid;
const int col = off_n + n * WIDTH + i;
- C_ptr[row * N + col] = tmp[i];
+ C_ptr[row * N + col] = __float2half(result);
}
}
asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
if (warp_id == 0)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
- :: "r"(0), "r"(BLOCK_N * 2));
+ :: "r"(0), "r"(256)); // Must match power-of-2 allocation
}
}
// ============================================================================
- // Fused SiLU+Multiply Kernel
- // ============================================================================
-
- __global__ void fused_silu_mul_kernel(
- const float *gemm1,
- const float *gemm2,
- half *output,
- int size
- ) {
- int idx = blockIdx.x * blockDim.x + threadIdx.x;
- if (idx < size) {
- float x = gemm1[idx];
- float silu = x / (1.0f + __expf(-x));
- float result = silu * gemm2[idx];
- output[idx] = __float2half(result);
- }
- }
-
- // ============================================================================
// Launch Wrapper
// ============================================================================
⋯ 10 unchanged lines
const int M = A.size(0);
const int N = B1.size(0);
- // Allocate temporary float32 buffers for intermediate results
- auto gemm1_buf = at::empty({M, N}, A.options().dtype(at::kFloat));
- auto gemm2_buf = at::empty({M, N}, A.options().dtype(at::kFloat));
-
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B1_ptr = reinterpret_cast<const char *>(B1.data_ptr());
auto B2_ptr = reinterpret_cast<const char *>(B2.data_ptr());
auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
auto SFB1_ptr = reinterpret_cast<const char *>(SFB1.data_ptr());
auto SFB2_ptr = reinterpret_cast<const char *>(SFB2.data_ptr());
+ auto C_ptr = reinterpret_cast<half *>(C.data_ptr());
CUtensorMap A_tmap, B1_tmap, B2_tmap;
init_AB_tmap(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K);
⋯ 3 unchanged lines
dim3 grid((M / BLOCK_M) * (N / BLOCK_N), 1, 1);
int tb_size = BLOCK_M + 2 * WARP_SIZE;
- constexpr int A_size = BLOCK_M * BLOCK_K / 2;
- constexpr int B_size = BLOCK_N * BLOCK_K / 2;
- constexpr int SFA_size = 128 * BLOCK_K / 16;
- constexpr int SFB_size = 128 * BLOCK_K / 16;
- int smem_size = (A_size + B_size + SFA_size + SFB_size) * NUM_STAGES;
+ // New SMEM size with dual B tiles
+ constexpr int A_size = BLOCK_M * BLOCK_K / 2;
+ constexpr int B1_size = BLOCK_N * BLOCK_K / 2;
+ constexpr int B2_size = BLOCK_N * BLOCK_K / 2;
+ constexpr int SFA_size = 128 * BLOCK_K / 16;
+ constexpr int SFB1_size = 128 * BLOCK_K / 16;
+ constexpr int SFB2_size = 128 * BLOCK_K / 16;
+ int smem_size = (A_size + B1_size + B2_size + SFA_size + SFB1_size + SFB2_size) * NUM_STAGES;
- auto kernel_fn = single_gemm_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
+ auto kernel_fn = dual_gemm_fused_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
if (smem_size > 48000)
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- // Run GEMM 1: A @ B1
kernel_fn<<<grid, tb_size, smem_size>>>(
- A_tmap, B1_tmap, SFA_ptr, SFB1_ptr, gemm1_buf.data_ptr<float>(), M, N
+ A_tmap, B1_tmap, B2_tmap, SFA_ptr, SFB1_ptr, SFB2_ptr, C_ptr, M, N
);
- // Run GEMM 2: A @ B2
- kernel_fn<<<grid, tb_size, smem_size>>>(
- A_tmap, B2_tmap, SFA_ptr, SFB2_ptr, gemm2_buf.data_ptr<float>(), M, N
- );
-
- // Fused SiLU + Multiply
- int total_size = M * N;
- int block_size = 256;
- int num_blocks = (total_size + block_size - 1) / block_size;
- fused_silu_mul_kernel<<<num_blocks, block_size>>>(
- gemm1_buf.data_ptr<float>(),
- gemm2_buf.data_ptr<float>(),
- reinterpret_cast<half*>(C.data_ptr()),
- total_size
- );
-
return C;
}
⋯ 12 unchanged lines
if (K == K_) return dual_gemm_launch<K_, BLOCK_M_, BLOCK_N_, BLOCK_K_, NUM_STAGES_>( \
A, B1, B2, SFA, SFB1, SFB2, C);
- // Main benchmark configs
- LAUNCH(7168, 128, 64, 256, 6)
- LAUNCH(4096, 128, 64, 256, 6)
+ // Main benchmark configs - use NUM_STAGES=4 for dual GEMM (more SMEM per stage)
+ LAUNCH(7168, 128, 64, 256, 4)
+ LAUNCH(4096, 128, 64, 256, 4)
- // Other K values from tests
+ // Other K values
LAUNCH(256, 128, 64, 256, 1)
LAUNCH(512, 128, 64, 256, 2)
- LAUNCH(1024, 128, 64, 256, 4)
- LAUNCH(1536, 128, 64, 256, 6)
- LAUNCH(2048, 128, 64, 256, 6)
- LAUNCH(2304, 128, 64, 256, 6)
+ LAUNCH(1024, 128, 64, 256, 3)
+ LAUNCH(1536, 128, 64, 256, 4)
+ LAUNCH(2048, 128, 64, 256, 4)
+ LAUNCH(2304, 128, 64, 256, 4)
#undef LAUNCH
scrolls · 417 diff lines total

Best evidence level for this revision: reported

JSON