Skip to content
KernelIndex
Search⌘K

submission 377876

nrehiew · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-377876?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
19.3µs
#118 of 161
2026-01-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9114eb0b12174c295cbbdbdd2c949d84f133271163db5040c7ed843ccc4f9efb
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-15

Techniques

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

async-copy__device__ __forceinline__ void cp_async_bulk_gmem2smem(void* dst, const void* src, uint32_t size, uint64_t* mbar) {
cluster__cluster_dims__(2, 1, 1)
fused-epilogueTAG_EPILOGUE = 4, // Time for writeback to global memory
mbarrierasm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];"
num-warps = 8constexpr int NUM_WARPS = 8;
shared-memory__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
tcgen05"tcgen05.alloc.cta_group::1.sync.aligned.b32 [%0], %1;\n"
tile-n = 64constexpr int STORE_BLOCK_N = 64; // 128 bytes / 2 bytes per bf16 for 128B swizzle
tmaasm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
vector-width = half2half2 packed[4];

Kernel source

submission.py1023 lines
# popcorn-cli submit submission.py --no-tui --leaderboard nvfp4_dual_gemm --gpu NVIDIA --mode test
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

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

torch::Tensor cuda_entry(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C);
"""

CUDA_SRC = r"""
#include <cuda.h>
#include <torch/torch.h>
#include <cuda_bf16.h>
#include <cooperative_groups.h>
#include <cuda/barrier>
#include <cuda/ptx>
#include <cudaTypedefs.h>

#include <cuda.h>
#include <cstdint>

namespace ptx = cuda::ptx;

#define PROFILE_TAGS_DEFINED
enum GemmProfileTags {
    TAG_TMA_ISSUE = 0,      // Time to issue TMA commands
    TAG_TMA_WAIT = 1,       // Time waiting for TMA
    TAG_MMA_ISSUE = 2,      // Time to issue MMA commands  
    TAG_MMA_WAIT = 3,       // Time waiting for MMA
    TAG_EPILOGUE = 4,      // Time for writeback to global memory
    TAG_SCALE_LOADING = 5,  // Time for scale loading
    TAG_SETUP = 6,          // Time for setup
};

#ifndef PROFILE_MAX_EVENTS
#define PROFILE_MAX_EVENTS 65536
#endif

#define PROFILE_BUFFER_SIZE (1 + PROFILE_MAX_EVENTS * 4)

// Buffer layout: profile[0] = count, then for event i:
//   profile[1+i*4+0] = start_ns, [1+i*4+1] = duration_ns, [1+i*4+2] = tag, [1+i*4+3] = tid

struct IntraKernelProfiler {
    int64_t* profile;
    int current_event_id;
    int tid;

    __device__ __forceinline__ void init(int64_t* profile_buffer, int thread_id) {
        profile = profile_buffer;
        current_event_id = -1;
        tid = thread_id;
    }

    __device__ __forceinline__ int64_t read_globaltimer() {
        int64_t time;
        asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(time));
        return time;
    }

    __device__ __forceinline__ void start(bool condition = true) {
        if (!condition) {
            current_event_id = -1;
            return;
        }
        int i = atomicAdd(reinterpret_cast<int*>(profile), 1);
        if (i >= PROFILE_MAX_EVENTS) {
            current_event_id = -1;
            return;
        }
        current_event_id = i;
        profile[1 + i * 4] = read_globaltimer();
    }

    __device__ __forceinline__ void stop(int tag, bool condition = true) {
        if (!condition || current_event_id < 0 || current_event_id >= PROFILE_MAX_EVENTS) return;
        int64_t end_time = read_globaltimer();
        profile[1 + current_event_id * 4 + 1] = end_time - profile[1 + current_event_id * 4];
        profile[1 + current_event_id * 4 + 2] = tag;
        profile[1 + current_event_id * 4 + 3] = tid;
    }
};

// Cache hint constants for L2 cache policy
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cute/arch/copy_sm90_desc.hpp#L193-L197
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;

constexpr int SF_VEC_SIZE = 16;
constexpr int ELEMENTS_PER_BYTE = 2;

constexpr int MMA_K = 64; // nvf4 processes 64 elements per MMA
constexpr int MMA_K_SCALES = MMA_K / SF_VEC_SIZE; // 4
constexpr int SF_NUM_COLS_PER_ITER = 4;
constexpr int MMA_K_IN_BYTES = MMA_K / ELEMENTS_PER_BYTE; // 32 bytes
constexpr int K_LOAD_SIZE_IN_BYTES = 128;

constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 8;
constexpr int NUM_THREADS = NUM_WARPS * WARP_SIZE;
constexpr int SF_TILE_SIZE_BYTES = 512;

constexpr int STORE_BLOCK_N = 64; // 128 bytes / 2 bytes per bf16 for 128B swizzle

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

inline unsigned int cdiv(unsigned int a, unsigned int b) {
    return (a + b - 1) / b;
}


__device__ __forceinline__
uint32_t cvta_shared(const void *ptr) { 
    return static_cast<uint32_t>(__cvta_generic_to_shared(ptr)); 
}
__forceinline__ __device__ uint32_t get_tmem_addr(uint32_t base_addr, int row_offset, int col_offset) {
    uint32_t offset = (row_offset << 16) | (col_offset & 0xFFFF);
    return base_addr + offset;
}

// PTX
__device__ inline void allocate_tensor_memory(uint32_t* ptr, int n_cols) {
    asm volatile(
        "tcgen05.alloc.cta_group::1.sync.aligned.b32 [%0], %1;\n"
        :: "l"(ptr), "r"(n_cols) 
    );
}
__device__ inline void deallocate_tensor_memory(uint32_t tmem_addr, int n_cols) {
    asm volatile(
        "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
        :: "r"(tmem_addr), "r"(n_cols)
    );
}

__device__ __forceinline__ void tcgen05_commit_group(uint64_t *bar) {
    uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];"
        :: "r"(mbar_ptr)  // 32-bit shared address
    );
}

__device__ static __forceinline__ void init_barrier(uint64_t* bar, int thread_count) {
    uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
    asm volatile (
        "mbarrier.init.shared::cta.b64 [%0], %1;\n"
        :: "r"(bar_ptr), "r"(thread_count)
    );
}

__device__ static __forceinline__ void expect_bytes(uint64_t* bar, uint32_t bytes) {
    uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
    asm volatile(
        "mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n"
        :: "r"(bar_ptr), "r"(bytes)
    );
}

__device__ static __forceinline__ void wait(uint64_t& bar, int phase_bit) {
    uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&bar));
    asm volatile (
        "{\n"
        ".reg .pred P1;\n"
        "LAB_WAIT:\n"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1;\n"
        "@P1 bra.uni DONE;\n"
        "bra.uni LAB_WAIT;\n"
        "DONE:\n"
        "}\n"
        :: "r"(mbar_ptr), "r"(phase_bit)
    );
}


template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_gmem2smem(void* dst, const void* src, uint32_t size, uint64_t* mbar) {
    uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
    uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
    asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
                 :: "r"(dst_addr), "l"(src), "r"(size), "r"(mbar_addr), "l"(CACHE_POLICY));
}

template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
    uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
    uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
    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_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)
                 : "memory");
}

template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
    uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
    uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
    asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
                 "[%0], [%1, {%2, %3, %4}], [%5], %6;"
                 :: "r"(dst_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)
                 : "memory");
}

__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {
    asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(__cvta_generic_to_shared(ptr)), "r"(x), "r"(y), "r"(z), "r"(w));
}


__device__ inline
uint64_t matrix_descriptor_encode(const uint64_t index) {
    return (index & 0x3FFFF) >> 4;
}
__device__
uint64_t make_matrix_descriptor_sf(void* x) {
    constexpr int sbo = 16 * 8; 
    constexpr int lbo = 16;
    // Blackwell reference 9.7.16.4.1
    // Both A and B are contiguous along K (K Major)
    uint64_t descriptor = 0;
    descriptor |= matrix_descriptor_encode(cvta_shared(x)); // Matrix start address
    // leading dimension byte offset (row size in bytes for K-major layout)
    // Only used when swizzle_mode = 0 (no swizzle)
    descriptor |= matrix_descriptor_encode(lbo) << 16;
    // stride dimension byte offset (offset from the first 8 columns to the next 8 columns)
    // For FP4 packed: stride = 8 rows × K_BYTES bytes/row
    descriptor |= matrix_descriptor_encode(sbo) << 32;
    descriptor |= 0b001ULL << 46; // Fixed constant value of 0b001
    return descriptor;
}

__device__
uint64_t make_matrix_descriptor_matmul(void* x) {
    // Blackwell reference 9.7.16.4.1
    // Both A and B are contiguous along K (K Major)
    uint64_t descriptor = 0;
    constexpr int sbo = K_LOAD_SIZE_IN_BYTES * 8;
    descriptor |= matrix_descriptor_encode(cvta_shared(x)); // Matrix start address
    // leading dimension byte offset (row size in bytes for K-major layout)
    // stride dimension byte offset (offset from the first 8 columns to the next 8 columns)
    // For FP4 packed: stride = 8 rows × K_BYTES bytes/row
    descriptor |= matrix_descriptor_encode(sbo) << 32;
    descriptor |= 0b001ULL << 46; // Fixed constant value of 0b001
    descriptor |= (2llu) << 61; //  2. 128-Byte swizzling 
    return descriptor;
}

__device__ __forceinline__ void tcgen05_ld_128(float* results, uint32_t base_addr){
    asm volatile(
        "tcgen05.ld.sync.aligned.16x256b.x16.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"(results[0]), "=f"(results[1]), "=f"(results[2]), "=f"(results[3]), "=f"(results[4]), "=f"(results[5]), "=f"(results[6]), "=f"(results[7]),
        "=f"(results[8]), "=f"(results[9]), "=f"(results[10]), "=f"(results[11]), "=f"(results[12]), "=f"(results[13]), "=f"(results[14]), "=f"(results[15]),
        "=f"(results[16]), "=f"(results[17]), "=f"(results[18]), "=f"(results[19]), "=f"(results[20]), "=f"(results[21]), "=f"(results[22]), "=f"(results[23]),
        "=f"(results[24]), "=f"(results[25]), "=f"(results[26]), "=f"(results[27]), "=f"(results[28]), "=f"(results[29]), "=f"(results[30]), "=f"(results[31]),
        "=f"(results[32]), "=f"(results[33]), "=f"(results[34]), "=f"(results[35]), "=f"(results[36]), "=f"(results[37]), "=f"(results[38]), "=f"(results[39]),
        "=f"(results[40]), "=f"(results[41]), "=f"(results[42]), "=f"(results[43]), "=f"(results[44]), "=f"(results[45]), "=f"(results[46]), "=f"(results[47]),
        "=f"(results[48]), "=f"(results[49]), "=f"(results[50]), "=f"(results[51]), "=f"(results[52]), "=f"(results[53]), "=f"(results[54]), "=f"(results[55]),
        "=f"(results[56]), "=f"(results[57]), "=f"(results[58]), "=f"(results[59]), "=f"(results[60]), "=f"(results[61]), "=f"(results[62]), "=f"(results[63])
        : "r"(base_addr)
    );
}

__device__ __forceinline__ void tcgen05_ld_64(float* results, uint32_t base_addr){
    asm volatile(
        "tcgen05.ld.sync.aligned.16x256b.x8.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"(results[0]), "=f"(results[1]), "=f"(results[2]), "=f"(results[3]), "=f"(results[4]), "=f"(results[5]), "=f"(results[6]), "=f"(results[7]),
        "=f"(results[8]), "=f"(results[9]), "=f"(results[10]), "=f"(results[11]), "=f"(results[12]), "=f"(results[13]), "=f"(results[14]), "=f"(results[15]),
        "=f"(results[16]), "=f"(results[17]), "=f"(results[18]), "=f"(results[19]), "=f"(results[20]), "=f"(results[21]), "=f"(results[22]), "=f"(results[23]),
        "=f"(results[24]), "=f"(results[25]), "=f"(results[26]), "=f"(results[27]), "=f"(results[28]), "=f"(results[29]), "=f"(results[30]), "=f"(results[31])
        : "r"(base_addr)
    );

}

template <const int BLOCK_N>
__device__ __forceinline__ void tcgen05_ld(float* results, uint32_t base_addr){
    // results is always BLOCK_N/2 registers
    if constexpr (BLOCK_N == 128) {
        tcgen05_ld_128(results, base_addr);
    } else if constexpr (BLOCK_N == 64) {
        tcgen05_ld_64(results, base_addr);
    } else {
        static_assert(BLOCK_N != 128 && BLOCK_N != 64, "Invalid block size");
    }
}

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

__device__ __forceinline__ void tcgen05_cp(uint32_t tmem_addr, uint64_t sdesc) {
    asm volatile(
        "tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" // shape: [no. lanes in TMEM, size in bits across columns]
        :: "r"(tmem_addr), "l"(sdesc)
    );
}


// nvf4 MMA with block scaling - matrices in shared memory, scale factors in tensor memory
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4(
    uint32_t d_tmem_addr,
    uint32_t i_desc,
    void* sA,
    void* sB,
    uint32_t scale_a_tmem,
    uint32_t scale_b_tmem
) {
    uint64_t a_desc = make_matrix_descriptor_matmul(sA); 
    uint64_t b_desc = make_matrix_descriptor_matmul(sB);

    if (init) {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    } else {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    }
}

// nvf4 MMA with collector buffer fill - fills collector with matrix A
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4_fill(
    uint32_t d_tmem_addr,
    uint32_t i_desc,
    void* sA,
    void* sB,
    uint32_t scale_a_tmem,
    uint32_t scale_b_tmem
) {
    uint64_t a_desc = make_matrix_descriptor_matmul(sA); 
    uint64_t b_desc = make_matrix_descriptor_matmul(sB);

    if (init) {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    } else {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    }
}

// nvf4 MMA with collector buffer lastuse - uses cached A from collector and discards
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4_lastuse(
    uint32_t d_tmem_addr,
    uint32_t i_desc,
    void* sA,
    void* sB,
    uint32_t scale_a_tmem,
    uint32_t scale_b_tmem
) {
    uint64_t a_desc = make_matrix_descriptor_matmul(sA); 
    uint64_t b_desc = make_matrix_descriptor_matmul(sB);

    if (init) {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    } else {
        asm volatile(
            "{\n"
            ".reg .pred p;\n"
            "setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
            "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
            "[%0], %1, %2, %3, [%4], [%5], p;\n"
            "}\n"
            :: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
               "r"(scale_a_tmem), "r"(scale_b_tmem)
        );
    }
}

__device__ static __forceinline__ void named_barrier_sync(int barrier_id, int thread_count) {
    asm volatile("barrier.sync %0, %1;" :: "r"(barrier_id), "r"(thread_count));
}

__device__ __forceinline__ float silu(float x) {
    // return x / (1.0f + expf(-x));
    return x * (1.0f / (1.0f + __expf(-x)));
}

constexpr size_t align_up(size_t x, size_t a) {
    return (x + a - 1) & ~(a - 1);
}

template<const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE, const bool DO_PROFILE>
__launch_bounds__(NUM_THREADS)
__cluster_dims__(2, 1, 1)
__global__ void cuda_kernel(
    const int M,
    const int N,
    const int K,
    const __grid_constant__ CUtensorMap a_tensor_map,
    const __grid_constant__ CUtensorMap b1_tensor_map,
    const __grid_constant__ CUtensorMap b2_tensor_map,
    const __grid_constant__ CUtensorMap c_tensor_map,
    char* sfa_ptr,
    char* sfb1_ptr,
    char* sfb2_ptr,
    half* __restrict__ C,
    int64_t* profile_buf
){
    constexpr int BLOCK_K_IN_BYTES = BLOCK_K / ELEMENTS_PER_BYTE;
    constexpr int NUM_MMA_ITERS = BLOCK_K / MMA_K;
    constexpr int SFA_COLS = SF_NUM_COLS_PER_ITER * NUM_MMA_ITERS;
    constexpr int SFB_COLS = SF_NUM_COLS_PER_ITER * NUM_MMA_ITERS;
    constexpr int A_bytes = BLOCK_M * BLOCK_K_IN_BYTES;
    constexpr int B_bytes = BLOCK_N * BLOCK_K_IN_BYTES;
    constexpr int SF_bytes = SF_TILE_SIZE_BYTES * NUM_MMA_ITERS;
    constexpr int matrix_bytes = A_bytes + B_bytes * 2;  // A, B1, B2
    constexpr int sf_tma_bytes = SF_bytes * 3;           // SFA, SFB1, SFB2
    constexpr int stage_size = matrix_bytes + sf_tma_bytes;
    constexpr int shared_size = stage_size * NUM_STAGES;

    constexpr int SFA_PER_TILE_SIZE = BLOCK_M * MMA_K_SCALES; // for one mma iter
    constexpr int SFB_PER_TILE_SIZE = SFA_PER_TILE_SIZE; // for one mma iter
    constexpr uint32_t idesc_nvf4 =
        (1 << 7) | (1 << 10) | ((BLOCK_N >> 3) << 17) | ((BLOCK_M >> 7) << 27);
    
    
    const int warp_id = threadIdx.x / WARP_SIZE;
    const int block_idx = blockIdx.x;
    uint32_t cluster_ctaid;
    asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cluster_ctaid));
    // const int block_row = block_idx / (N / BLOCK_N);
    // const int block_col = block_idx % (N / BLOCK_N);
    const int cluster_idx = blockIdx.x / 2;
    const int blocks_per_row = N / BLOCK_N;
    const int cluster_row = (cluster_idx / blocks_per_row) * 2;
    const int cluster_col = cluster_idx % blocks_per_row;
    const int block_row = cluster_row + cluster_ctaid;
    const int block_col = cluster_col;
    

    IntraKernelProfiler profiler;
    bool should_profile = false;
    if constexpr (DO_PROFILE) {
        profiler.init(profile_buf, block_idx * 10 + warp_id);
        should_profile = (blockIdx.x < 20) && (threadIdx.x % WARP_SIZE == 0);
    }

    if constexpr (DO_PROFILE) profiler.start(should_profile);

    C += block_row * BLOCK_M * N + block_col * BLOCK_N;
    extern __shared__ __align__(128) uint8_t shared_mem[];
    uint8_t* smem_ptr = shared_mem;

    uint8_t* As = smem_ptr;
    smem_ptr += align_up(A_bytes * NUM_STAGES, 128);
    uint8_t* B1s = smem_ptr;
    smem_ptr += align_up(B_bytes * NUM_STAGES, 128);
    uint8_t* B2s = smem_ptr;
    smem_ptr += align_up(B_bytes * NUM_STAGES, 128);
    uint8_t* SFA = smem_ptr;
    smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
    uint8_t* SFB1 = smem_ptr;
    smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
    uint8_t* SFB2 = smem_ptr;
    smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
    half* Cs = reinterpret_cast<half*>(As);

    __shared__ uint32_t tmem_addr_base_shared;

    if (warp_id == 1) {
        allocate_tensor_memory(&tmem_addr_base_shared, 512);
    }

    // Tensor memory layout: SFA * NUM_STAGES, SFB1 * NUM_STAGES, SFB2 * NUM_STAGES
    constexpr int SF_TOTAL_COLS = SFA_COLS * NUM_STAGES + SFB_COLS * NUM_STAGES * 2;

    __shared__ __align__(8) uint64_t is_empty_bars[NUM_STAGES]; 
    __shared__ __align__(8) uint64_t mma_ready_bars[NUM_STAGES];
    __shared__ __align__(8) uint64_t sf_tma_ready_bars[NUM_STAGES];
    __shared__ __align__(8) uint64_t last_stage_bar;

    if (threadIdx.x < NUM_STAGES) {
        init_barrier(&is_empty_bars[threadIdx.x], 2);
        init_barrier(&mma_ready_bars[threadIdx.x], 2);
        init_barrier(&sf_tma_ready_bars[threadIdx.x], 1);
    }
    if (threadIdx.x == 0) {
        init_barrier(&last_stage_bar, 2);
    }
    __syncthreads();

    const uint32_t tmem_addr_base = tmem_addr_base_shared;
    const uint32_t sfa_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 0);
    const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS * NUM_STAGES);
    const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS * NUM_STAGES);
    const uint32_t d1_tmem_addr = get_tmem_addr(sfb2_tmem_addr, 0, SFB_COLS * NUM_STAGES);
    const uint32_t d2_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 512 - BLOCK_N);


    int row_start = block_row * BLOCK_M;
    int col_start = block_col * BLOCK_N;
    int K_BYTES = K / 2;
    int rest_k = K / SF_VEC_SIZE / 4;
    int num_iters = K_BYTES / BLOCK_K_IN_BYTES;

    sfa_ptr += row_start/128 * SF_TILE_SIZE_BYTES * rest_k; 
    sfb1_ptr += col_start/128 * SF_TILE_SIZE_BYTES * rest_k;
    sfb2_ptr += col_start/128 * SF_TILE_SIZE_BYTES * rest_k;

    if constexpr (DO_PROFILE) profiler.stop(TAG_SETUP, should_profile);
    if (warp_id == 4 && elect_sync()) {
        auto issue_tma = [&](int iter, int stage_id) {
            if constexpr (DO_PROFILE) profiler.start(should_profile);
            
            constexpr int NUM_K_CHUNKS = BLOCK_K_IN_BYTES / K_LOAD_SIZE_IN_BYTES;
            const int k_chunk_start = iter * NUM_K_CHUNKS;
            int32_t a_tensor_coords[3] = {0, row_start, k_chunk_start};
            int32_t b12_tensor_coords[3] = {0, col_start, k_chunk_start};

            cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
            cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
            cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
            expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);
            
            cp_async_bulk_tensor_gmem2smem<EVICT_LAST>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);
            cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B1s[stage_id * B_bytes], &b1_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);
            cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B2s[stage_id * B_bytes], &b2_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);

            expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);

            if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_ISSUE, should_profile);
        };

        #pragma unroll
        for (int iter = 0; iter < NUM_STAGES; iter++) {
            issue_tma(iter, iter);
        }

        #pragma unroll
        for (int iter = NUM_STAGES; iter < num_iters; iter++) {
            const int stage_id = iter % NUM_STAGES;
            const int mma_phase = (iter / NUM_STAGES - 1) % 2;
            
            if constexpr (DO_PROFILE) profiler.start(should_profile);
            wait(is_empty_bars[stage_id], mma_phase);
            if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_WAIT, should_profile);
            
            issue_tma(iter, stage_id);
        }
    } else if (warp_id == 5 && elect_sync()) {
        uint32_t sf_phase[NUM_STAGES] = {};
        int current_stage_idx = 0;
        for (int iter = 0; iter < num_iters; iter++) {

            uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
            uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);
            uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);

            if constexpr (DO_PROFILE) profiler.start(should_profile);
            wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);
            sf_phase[current_stage_idx] ^= 1;
            if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
            if constexpr (DO_PROFILE) profiler.start(should_profile);
            
            for (int i = 0; i < NUM_MMA_ITERS; i++) {
                uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
                uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
                uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
                tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);
                tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);
                tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);
            }
            if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);
            
            tcgen05_commit_group(&mma_ready_bars[current_stage_idx]);
            current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
        }
    } else if (warp_id == 6 && elect_sync()) {
        uint32_t mma_phase[NUM_STAGES] = {};
        int current_stage_idx = 0;

        const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
        const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
        for (int iter = 0; iter < num_iters; iter++) {
            uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
            uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);

            if constexpr (DO_PROFILE) profiler.start(should_profile);
            wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
            if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
            mma_phase[current_stage_idx] ^= 1;

            if constexpr (DO_PROFILE) profiler.start(should_profile);

            static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");

            constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;

            {
                uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
                uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);

                if (iter == 0) {
                    tcgen05_mma_nvf4<true>(d1_tmem_addr, idesc_nvf4,
                        &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
                        sfa_addr, sfb1_addr);
                } else {
                    tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
                        &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
                        sfa_addr, sfb1_addr);
                }
            }

            #pragma unroll
            for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
                uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
                uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);
                tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
                    &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
                    sfa_addr, sfb1_addr);
            }
            tcgen05_commit_group(&is_empty_bars[current_stage_idx]);
            if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);
            current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
        }
        tcgen05_commit_group(&last_stage_bar);
    } else if (warp_id == 7 && elect_sync()) {
        uint32_t mma_phase[NUM_STAGES] = {};
        int current_stage_idx = 0;

        const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
        const int sfb2_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
        for (int iter = 0; iter < num_iters; iter++) {
            uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
            uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);

            if constexpr (DO_PROFILE) profiler.start(should_profile);
            wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
            if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
            mma_phase[current_stage_idx] ^= 1;

            if constexpr (DO_PROFILE) profiler.start(should_profile);

            static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");

            constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;

            {
                uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
                uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, sfb2_offset);

                if (iter == 0) {
                    tcgen05_mma_nvf4<true>(d2_tmem_addr, idesc_nvf4,
                        &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
                        sfa_addr, sfb2_addr);
                } else {
                    tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
                        &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
                        sfa_addr, sfb2_addr);
                }
            }

            #pragma unroll
            for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
                uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
                uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb2_offset);

                tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
                    &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
                    sfa_addr, sfb2_addr);
            }
            tcgen05_commit_group(&is_empty_bars[current_stage_idx]);
            if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);
            current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
        }
        tcgen05_commit_group(&last_stage_bar);
    } else if (warp_id < 4) {
        wait(last_stage_bar, 0);
        asm volatile("tcgen05.fence::after_thread_sync;");
        if constexpr (DO_PROFILE) profiler.start(should_profile);
        
        if constexpr (USE_TMA_EPILOGUE) {
            constexpr int kSwizzleCDMode = 128;
            constexpr int kNumBankGroupBytes = 16;
            constexpr int kNumElemsPerBankGroup = kNumBankGroupBytes / sizeof(half); // 8 elements
            constexpr int kNumStores = BLOCK_N / STORE_BLOCK_N;
    
            const int lane_id = threadIdx.x % WARP_SIZE;
    
            int row = threadIdx.x;
            uint8_t* cs_base = reinterpret_cast<uint8_t*>(&Cs[0]);
            int row_byte_offset = row * STORE_BLOCK_N * sizeof(half);
            uint32_t row_smem_addr = cvta_shared(cs_base + row_byte_offset);
            int swizzle_xor = (row_smem_addr >> 7) & 0x7;
    
            for (int chunk = 0; chunk < kNumStores; ++chunk) {
                int chunk_smem_offset_base = chunk * (BLOCK_M * STORE_BLOCK_N * sizeof(half));
                int chunk_col_offset = chunk * STORE_BLOCK_N;
    
                float d1_buf[2][8];
                float d2_buf[2][8];
                int stage = 0;
                
                constexpr int num_iterations = STORE_BLOCK_N / kNumElemsPerBankGroup;
                
                uint32_t d1_base_addr_first = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset);
                uint32_t d2_base_addr_first = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset);
                tcgen05_ld_8(d1_buf[stage], d1_base_addr_first);
                tcgen05_ld_8(d2_buf[stage], d2_base_addr_first);
    
                uint32_t d1_base_addr_next = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset + kNumElemsPerBankGroup);
                uint32_t d2_base_addr_next = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset + kNumElemsPerBankGroup);
    
                for (int iter = 0; iter < num_iterations; ++iter) {
                    int col_in_chunk = iter * kNumElemsPerBankGroup;
                    int bank_group_index = col_in_chunk / kNumElemsPerBankGroup;
                    
                    asm volatile("tcgen05.wait::ld.sync.aligned;");
                    
                    if (iter + 1 < num_iterations) {
                        int next_stage = stage ^ 1;
                        tcgen05_ld_8(d1_buf[next_stage], d1_base_addr_next);
                        tcgen05_ld_8(d2_buf[next_stage], d2_base_addr_next);
                        d1_base_addr_next = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset + (iter + 2) * kNumElemsPerBankGroup);
                        d2_base_addr_next = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset + (iter + 2) * kNumElemsPerBankGroup);
                    }
    
                    half2 packed[4];

                    #pragma unroll 2
                    for (int i = 0; i < 4; i++) {
                        float2 d1_pair = {d1_buf[stage][i*2], d1_buf[stage][i*2+1]};
                        float2 d2_pair = {d2_buf[stage][i*2], d2_buf[stage][i*2+1]};
                        
                        // Vectorized SiLU
                        float2 sigmoid = {1.0f / (1.0f + __expf(-d1_pair.x)), 
                                          1.0f / (1.0f + __expf(-d1_pair.y))};
                        float2 result = {d1_pair.x * sigmoid.x * d2_pair.x,
                                         d1_pair.y * sigmoid.y * d2_pair.y};
                        
                        packed[i] = __float22half2_rn(result);
                    }
                    
    
                    int col = bank_group_index ^ swizzle_xor;
                    auto smem_ptr = cs_base + row_byte_offset + col * kNumBankGroupBytes + chunk_smem_offset_base;
    
                    st_shared(smem_ptr,
                                 *reinterpret_cast<uint32_t*>(&packed[0]),
                                 *reinterpret_cast<uint32_t*>(&packed[1]),
                                 *reinterpret_cast<uint32_t*>(&packed[2]),
                                 *reinterpret_cast<uint32_t*>(&packed[3]));
                    
                    stage ^= 1;
                }
                
                ptx::fence_proxy_async(ptx::space_shared);
                named_barrier_sync(0, 128);
                if (warp_id == 0 && elect_sync()) {
                    int32_t c_tensor_coords[2] = {block_col * BLOCK_N + chunk_col_offset, block_row * BLOCK_M};
                    ptx::cp_async_bulk_tensor(ptx::space_global, ptx::space_shared, &c_tensor_map, c_tensor_coords, Cs + chunk * (BLOCK_M * STORE_BLOCK_N));
                    ptx::cp_async_bulk_commit_group();
                }
                
            }
            if (warp_id == 0 && elect_sync()) {
                ptx::cp_async_bulk_wait_group_read(ptx::n32_t<0>());
            }
    
        } else {
            const int lane_id = threadIdx.x % WARP_SIZE;

            float d1_results[2][BLOCK_N / 2];
            float d2_results[2][BLOCK_N / 2];
            
            uint32_t d1_base_addr = get_tmem_addr(d1_tmem_addr, warp_id * 32, 0);
            uint32_t d2_base_addr = get_tmem_addr(d2_tmem_addr, warp_id * 32, 0);
            tcgen05_ld<BLOCK_N>(d1_results[0], d1_base_addr);
            tcgen05_ld<BLOCK_N>(d2_results[0], d2_base_addr);
            
            for (int batch = 0; batch < 2; batch++) {
                asm volatile("tcgen05.wait::ld.sync.aligned;");
                
                if (batch < 1) {
                    uint32_t d1_next_addr = get_tmem_addr(d1_tmem_addr, warp_id * 32 + 16, 0);
                    uint32_t d2_next_addr = get_tmem_addr(d2_tmem_addr, warp_id * 32 + 16, 0);
                    tcgen05_ld<BLOCK_N>(d1_results[1], d1_next_addr);
                    tcgen05_ld<BLOCK_N>(d2_results[1], d2_next_addr);
                }
                
                const int base_row = warp_id * 32 + batch * 16 + lane_id / 4;
                const int col_base = (lane_id % 4) * 2;
                
                #pragma unroll
                for (int i = 0; i < BLOCK_N / 8; i++) {
                    const int col = i * 8 + col_base;
                    const int idx = i * 4;
                    
                    float a = silu(d1_results[batch][idx + 0]) * d2_results[batch][idx + 0];
                    float b = silu(d1_results[batch][idx + 1]) * d2_results[batch][idx + 1];
                    float c = silu(d1_results[batch][idx + 2]) * d2_results[batch][idx + 2];
                    float d = silu(d1_results[batch][idx + 3]) * d2_results[batch][idx + 3];
                    
                    *reinterpret_cast<half2*>(&C[(base_row + 0) * N + col]) = __float22half2_rn({a, b});
                    *reinterpret_cast<half2*>(&C[(base_row + 8) * N + col]) = __float22half2_rn({c, d});
                }
            }        
        }
        named_barrier_sync(0, 128);
        if constexpr (DO_PROFILE) profiler.stop(TAG_EPILOGUE, should_profile);

        if (warp_id == 1) {
            deallocate_tensor_memory(tmem_addr_base, 512);
        }
    }
    
}

template<const int SMEM_WIDTH, const int SMEM_HEIGHT >
CUtensorMap create_tensor_map_ab(void* global_address, int GMEM_WIDTH, int GMEM_HEIGHT) {
    CUtensorMap tensor_map{};
    constexpr int rank = 3;
    uint32_t elem_stride[rank] = {1, 1, 1};
    uint64_t size[rank] = {K_LOAD_SIZE_IN_BYTES, static_cast<uint64_t>(GMEM_HEIGHT), static_cast<uint64_t>(GMEM_WIDTH / K_LOAD_SIZE_IN_BYTES)};
    uint32_t box_size[rank] = {K_LOAD_SIZE_IN_BYTES, SMEM_HEIGHT, SMEM_WIDTH / K_LOAD_SIZE_IN_BYTES};
    uint64_t stride[rank - 1] = {static_cast<uint64_t>(GMEM_WIDTH), K_LOAD_SIZE_IN_BYTES};


    CUresult res = cuTensorMapEncodeTiled(
      &tensor_map,
      CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_UINT8,
      rank,
      (void*)global_address,
      size,
      stride,
      box_size,
      elem_stride,
      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
    );
    return tensor_map;
}

template<const int SMEM_WIDTH, const int SMEM_HEIGHT>
CUtensorMap create_tensor_map_c(void* global_address, int M, int N) {
    CUtensorMap tensor_map{};
    constexpr int rank = 2;
    uint32_t elem_stride[rank] = {1, 1};
    uint64_t size[rank] = {static_cast<uint64_t>(N), static_cast<uint64_t>(M)};
    uint32_t box_size[rank] = {SMEM_WIDTH, SMEM_HEIGHT};
    uint64_t stride[1] = {static_cast<uint64_t>(N * sizeof(half))};

    CUresult res = cuTensorMapEncodeTiled(
      &tensor_map,
      CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
      rank,
      (void*)global_address,
      size,
      stride,
      box_size,
      elem_stride,
      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
    );
    return tensor_map;
}

template<const bool DO_PROFILE, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE>
torch::Tensor matmul_impl(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, int64_t* profile_buf) {
    constexpr int BLOCK_K_IN_BYTES = BLOCK_K / ELEMENTS_PER_BYTE;
    
    int M = A.size(0);
    int N = B1.size(0);
    int K_BYTES = A.size(1); // This is in bytes 
    int K = K_BYTES * 2;
    constexpr int L = 1;

    CUtensorMap a_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_M>(
        A.data_ptr(), K_BYTES, M);  // M rows, K_BYTES cols (bytes)
    CUtensorMap b1_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_N>(
        B1.data_ptr(), K_BYTES, N);
    CUtensorMap b2_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_N>(
        B2.data_ptr(), K_BYTES, N);
    CUtensorMap c_tensor_map = create_tensor_map_c<STORE_BLOCK_N, BLOCK_M>(C.data_ptr(), M, N);

    const int NUM_BLOCKS = cdiv(N, BLOCK_N) * cdiv(M, BLOCK_M);
    dim3 gridDim(NUM_BLOCKS);
    dim3 blockDim(NUM_THREADS);

    auto kernel = cuda_kernel<BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, USE_TMA_EPILOGUE, DO_PROFILE>;
    constexpr int A_bytes = BLOCK_M * BLOCK_K_IN_BYTES;
    constexpr int B_bytes = BLOCK_N * BLOCK_K_IN_BYTES;
    constexpr int NUM_MMA_ITERS = BLOCK_K / MMA_K;
    constexpr int SF_bytes = SF_TILE_SIZE_BYTES * NUM_MMA_ITERS;
    constexpr int stage_size = A_bytes + B_bytes * 2 + SF_bytes * 3;
    constexpr int shared_size = stage_size * NUM_STAGES;
    cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, shared_size);

    kernel<<<gridDim, blockDim, shared_size>>>(
        M,
        N,
        K,  // K in elements
        a_tensor_map,
        b1_tensor_map,
        b2_tensor_map,
        c_tensor_map,
        reinterpret_cast<char*>(SFA_permuted.data_ptr()),
        reinterpret_cast<char*>(SFB1_permuted.data_ptr()),
        reinterpret_cast<char*>(SFB2_permuted.data_ptr()),
        reinterpret_cast<half*>(C.data_ptr()),
        profile_buf
    );
    return C;
}

template<const bool DO_PROFILE>
torch::Tensor matmul_impl(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, int64_t* profile_buf) {
    int M = A.size(0);
    if (M == 256) {
        return matmul_impl<DO_PROFILE, 128, 64, 256, 5, true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, profile_buf);
    } else {
        return matmul_impl<DO_PROFILE, 128, 128, 256, 4, true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, profile_buf);
    }
    
}


torch::Tensor cuda_entry(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C) {
    return matmul_impl<false>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, nullptr);
}

torch::Tensor cuda_entry_with_profile(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, torch::Tensor profile_buf) {
    return matmul_impl<true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, reinterpret_cast<int64_t*>(profile_buf.data_ptr()));
}
"""

cuda_module = load_inline(
    name="cuda_kernel",
    cpp_sources=cpp_source,
    cuda_sources=CUDA_SRC,
    functions=["cuda_entry"],
    extra_cuda_cflags=[
            "-O3",
            "-lineinfo",
            "-Xptxas=-v",
            "-gencode=arch=compute_100a,code=sm_100a",
            "--use_fast_math",
        ],
    extra_ldflags=["-lcuda"],  # for cuTensorMapEncodeTiled() used by TMA
    verbose=True,
)
import torch

def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
    out = cuda_module.cuda_entry(a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c)
    return out
scrolls · 1023 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 341566.

⋯ 204 unchanged lines
: "memory");
}
+ template<uint64_t CACHE_POLICY>
+ __device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
+ uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
+ uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
+ asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
+ "[%0], [%1, {%2, %3, %4}], [%5], %6;"
+ :: "r"(dst_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)
+ : "memory");
+ }
+
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {
asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(__cvta_generic_to_shared(ptr)), "r"(x), "r"(y), "r"(z), "r"(w));
}
⋯ 181 unchanged lines
// nvf4 MMA with collector buffer lastuse - uses cached A from collector and discards
template <bool init>
- __device__ __forceinline__ void tcgen05_mma_nvf4_use(
+ __device__ __forceinline__ void tcgen05_mma_nvf4_lastuse(
uint32_t d_tmem_addr,
uint32_t i_desc,
void* sA,
⋯ 9 unchanged lines
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::use "
+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
⋯ 4 unchanged lines
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::use "
+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
⋯ 17 unchanged lines
template<const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE, const bool DO_PROFILE>
__launch_bounds__(NUM_THREADS)
+ __cluster_dims__(2, 1, 1)
__global__ void cuda_kernel(
const int M,
const int N,
⋯ 25 unchanged lines
constexpr uint32_t idesc_nvf4 =
(1 << 7) | (1 << 10) | ((BLOCK_N >> 3) << 17) | ((BLOCK_M >> 7) << 27);
+
const int warp_id = threadIdx.x / WARP_SIZE;
const int block_idx = blockIdx.x;
- const int block_row = block_idx / (N / BLOCK_N);
- const int block_col = block_idx % (N / BLOCK_N);
+ uint32_t cluster_ctaid;
+ asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cluster_ctaid));
+ // const int block_row = block_idx / (N / BLOCK_N);
+ // const int block_col = block_idx % (N / BLOCK_N);
+ const int cluster_idx = blockIdx.x / 2;
+ const int blocks_per_row = N / BLOCK_N;
+ const int cluster_row = (cluster_idx / blocks_per_row) * 2;
+ const int cluster_col = cluster_idx % blocks_per_row;
+ const int block_row = cluster_row + cluster_ctaid;
+ const int block_col = cluster_col;
+
IntraKernelProfiler profiler;
bool should_profile = false;
⋯ 28 unchanged lines
allocate_tensor_memory(&tmem_addr_base_shared, 512);
}
- // Each stage uses SFA_COLS + SFB_COLS * 2 columns
- constexpr int SF_STAGE_STRIDE = SFA_COLS + SFB_COLS * 2;
+ // Tensor memory layout: SFA * NUM_STAGES, SFB1 * NUM_STAGES, SFB2 * NUM_STAGES
+ constexpr int SF_TOTAL_COLS = SFA_COLS * NUM_STAGES + SFB_COLS * NUM_STAGES * 2;
__shared__ __align__(8) uint64_t is_empty_bars[NUM_STAGES];
__shared__ __align__(8) uint64_t mma_ready_bars[NUM_STAGES];
⋯ 1 unchanged lines
__shared__ __align__(8) uint64_t last_stage_bar;
if (threadIdx.x < NUM_STAGES) {
- init_barrier(&is_empty_bars[threadIdx.x], 1);
- init_barrier(&mma_ready_bars[threadIdx.x], 3);
+ init_barrier(&is_empty_bars[threadIdx.x], 2);
+ init_barrier(&mma_ready_bars[threadIdx.x], 2);
init_barrier(&sf_tma_ready_bars[threadIdx.x], 1);
}
if (threadIdx.x == 0) {
- init_barrier(&last_stage_bar, 1);
+ init_barrier(&last_stage_bar, 2);
}
__syncthreads();
const uint32_t tmem_addr_base = tmem_addr_base_shared;
const uint32_t sfa_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 0);
- const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS);
- const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS);
- const uint32_t d1_tmem_addr = get_tmem_addr(tmem_addr_base, 0, NUM_STAGES * SF_STAGE_STRIDE);
+ const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS * NUM_STAGES);
+ const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS * NUM_STAGES);
+ const uint32_t d1_tmem_addr = get_tmem_addr(sfb2_tmem_addr, 0, SFB_COLS * NUM_STAGES);
const uint32_t d2_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 512 - BLOCK_N);
⋯ 12 unchanged lines
auto issue_tma = [&](int iter, int stage_id) {
if constexpr (DO_PROFILE) profiler.start(should_profile);
- expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);
constexpr int NUM_K_CHUNKS = BLOCK_K_IN_BYTES / K_LOAD_SIZE_IN_BYTES;
const int k_chunk_start = iter * NUM_K_CHUNKS;
int32_t a_tensor_coords[3] = {0, row_start, k_chunk_start};
int32_t b12_tensor_coords[3] = {0, col_start, k_chunk_start};
- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
-
- expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);
+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
+ expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);
- cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);
+ cp_async_bulk_tensor_gmem2smem<EVICT_LAST>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);
cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B1s[stage_id * B_bytes], &b1_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);
cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B2s[stage_id * B_bytes], &b2_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);
+ expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);
+
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_ISSUE, should_profile);
};
⋯ 18 unchanged lines
int current_stage_idx = 0;
for (int iter = 0; iter < num_iters; iter++) {
- uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- // uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
+ uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);
+ uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);
⋯ 3 unchanged lines
for (int i = 0; i < NUM_MMA_ITERS; i++) {
uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
- // uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
+ uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);
- // tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);
+ tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);
tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);
}
if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);
⋯ 2 unchanged lines
current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
}
} else if (warp_id == 6 && elect_sync()) {
- uint32_t sf_phase[NUM_STAGES] = {};
+ uint32_t mma_phase[NUM_STAGES] = {};
int current_stage_idx = 0;
+
+ const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
+ const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
for (int iter = 0; iter < num_iters; iter++) {
+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
+ uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);
- // uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- // uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
-
if constexpr (DO_PROFILE) profiler.start(should_profile);
- wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);
- sf_phase[current_stage_idx] ^= 1;
+ wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
+ mma_phase[current_stage_idx] ^= 1;
+
if constexpr (DO_PROFILE) profiler.start(should_profile);
-
- for (int i = 0; i < NUM_MMA_ITERS; i++) {
- // uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
- uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
- // uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
- // tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);
- tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);
- // tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);
+
+ static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");
+
+ constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;
+
+ {
+ uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
+ uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);
+
+ if (iter == 0) {
+ tcgen05_mma_nvf4<true>(d1_tmem_addr, idesc_nvf4,
+ &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
+ sfa_addr, sfb1_addr);
+ } else {
+ tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
+ &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
+ sfa_addr, sfb1_addr);
+ }
}
- if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);
-
- tcgen05_commit_group(&mma_ready_bars[current_stage_idx]);
+
+ #pragma unroll
+ for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
+ uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
+ uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);
+ tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
+ &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
+ sfa_addr, sfb1_addr);
+ }
+ tcgen05_commit_group(&is_empty_bars[current_stage_idx]);
+ if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);
current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
}
+ tcgen05_commit_group(&last_stage_bar);
} else if (warp_id == 7 && elect_sync()) {
uint32_t mma_phase[NUM_STAGES] = {};
int current_stage_idx = 0;
const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
- const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
const int sfb2_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
for (int iter = 0; iter < num_iters; iter++) {
- uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
- uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);
+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
+ uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);
-
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
mma_phase[current_stage_idx] ^= 1;
-
if constexpr (DO_PROFILE) profiler.start(should_profile);
⋯ 3 unchanged lines
{
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
- uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);
uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, sfb2_offset);
if (iter == 0) {
- tcgen05_mma_nvf4_fill<true>(d1_tmem_addr, idesc_nvf4,
- &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
- sfa_addr, sfb1_addr);
- tcgen05_mma_nvf4_use<true>(d2_tmem_addr, idesc_nvf4,
+ tcgen05_mma_nvf4<true>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
} else {
- tcgen05_mma_nvf4_fill<false>(d1_tmem_addr, idesc_nvf4,
- &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
- sfa_addr, sfb1_addr);
- tcgen05_mma_nvf4_use<false>(d2_tmem_addr, idesc_nvf4,
+ tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
}
⋯ 2 unchanged lines
#pragma unroll
for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
- uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);
uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb2_offset);
- tcgen05_mma_nvf4_fill<false>(d1_tmem_addr, idesc_nvf4,
- &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
- sfa_addr, sfb1_addr);
- tcgen05_mma_nvf4_use<false>(d2_tmem_addr, idesc_nvf4,
+ tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
}
⋯ 55 unchanged lines
half2 packed[4];
+ #pragma unroll 2
for (int i = 0; i < 4; i++) {
float2 d1_pair = {d1_buf[stage][i*2], d1_buf[stage][i*2+1]};
float2 d2_pair = {d2_buf[stage][i*2], d2_buf[stage][i*2+1]};
⋯ 219 unchanged lines
extra_ldflags=["-lcuda"], # for cuTensorMapEncodeTiled() used by TMA
verbose=True,
)
+ import torch
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
scrolls · 317 diff lines total

Best evidence level for this revision: reported

JSON