Skip to content
KernelIndex
Search⌘K

submission 75298

sohail · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_tc_tmem.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-75298?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 GEMVsuite of 3 cases
NVIDIA B200
313.8µs
#587 of 678
2025-11-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a5086b6b664f3e20f27d0425786fba41c21139f9003c7b7321f4a5f0df6fca78
license declaredunknown
license concludedunknown
authorssohail
imported2026-08-26

Techniques

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

async-copy__device__ __forceinline__ uint32_t tcgen_cp_async(
fp4printf("[nvfp4] selected kernel path = %d\n", static_cast<int>(selected));
mbarrier__device__ __forceinline__ void mbarrier_init_cta(uint64_t* bar, uint32_t pending_count) {
mmanamespace wmma = nvcuda::wmma;
shared-memory__device__ __forceinline__ uint32_t smem_ptr(const void* ptr) {
tcgen05"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"

Kernel source

submission_tc_tmem.py1573 lines
# submission.py
# Optimized PyTorch extension for B200 NVFP4 operations with FP8 scales
import os
import torch
from torch.utils.cpp_extension import load_inline

_EXT = None  # lazily compiled extension
_FP4_LUT_CACHE = {}

def _get_fp4_lut(device):
    key = device.index if device.type == "cuda" else device
    lut = _FP4_LUT_CACHE.get(key)
    if lut is None:
        values = torch.tensor(
            [
                0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
                0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0,
            ],
            device=device,
            dtype=torch.float32,
        )
        _FP4_LUT_CACHE[key] = values
        lut = values
    return lut

def _decode_fp4_tensor(packed):
    lut = _get_fp4_lut(packed.device)
    low = (packed & 0xF).to(torch.long)
    high = ((packed >> 4) & 0xF).to(torch.long)
    decoded = torch.stack((lut[low], lut[high]), dim=-1)
    shape = packed.shape
    return decoded.view(shape[0], shape[1] * 2, shape[2]).contiguous()

def _decode_fp8_tensor(packed):
    if hasattr(torch, "float8_e4m3fn"):
        return packed.view(torch.float8_e4m3fn).to(torch.float32).contiguous()
    # Fallback decoding if float8 dtype is unavailable
    bytes_float = packed.to(torch.float32)
    sign = ((packed >> 7) & 0x1).to(torch.float32)
    exponent = ((packed >> 3) & 0xF).to(torch.float32)
    mantissa = (packed & 0x7).to(torch.float32)
    is_zero = (exponent == 0) & (mantissa == 0)
    norm = exponent >= 1
    value = torch.empty_like(bytes_float)
    value[norm] = (1.0 + mantissa[norm] / 8.0) * torch.pow(2.0, exponent[norm] - 7.0)
    value[~norm] = (mantissa[~norm] / 8.0) * torch.pow(2.0, -6.0)
    value[is_zero] = 0.0
    return torch.where(sign > 0, -value, value).contiguous()


def _get_ext():
    global _EXT
    if _EXT is not None:
        return _EXT

    # Minimal C++ source with forward declaration
    CPP_SRC = """
    #include <torch/extension.h>
    
    void bmv_launcher(
        torch::Tensor A,
        torch::Tensor B,
        torch::Tensor SFA,
        torch::Tensor SFB,
        torch::Tensor C);
    """

    enable_tcgen_build = os.environ.get("NVFP4_ENABLE_TCGEN_BUILD", "").lower() not in {"", "0", "false", "no"}
    experimental_tcgen = os.environ.get("NVFP4_EXPERIMENTAL_TCGEN_MMA", "").lower() not in {"", "0", "false", "no"}
    if experimental_tcgen:
        print("Enabling NVFP4_EXPERIMENTAL_TCGEN_MMA build flag (1)")
    if os.environ.get("NVFP4_CAPTURE_TMEM_DEBUG"):
        keep_flags = "-keep --keep-dir /tmp/nvcc_keep"
        existing = os.environ.get("TORCH_NVCC_FLAGS", "")
        if keep_flags not in existing:
            combined = f"{existing} {keep_flags}".strip()
            os.environ["TORCH_NVCC_FLAGS"] = combined
        os.makedirs("/tmp/nvcc_keep", exist_ok=True)

    # CUDA implementation
    CUDA_SRC = r"""
    #include <cuda.h>
    #include <cuda_runtime.h>
    #include <cuda_fp16.h>
    #include <mma.h>
    #include <torch/extension.h>
    #include <ATen/cuda/CUDAContext.h>
    #include <c10/cuda/CUDAStream.h>
    #include <stdint.h>
    #include <cstdlib>
    #include <cstring>
    #include <cctype>
    #include <cstdio>
    #include <cmath>
    #include <vector>

#ifndef NVFP4_ENABLE_TCGEN
#define NVFP4_ENABLE_TCGEN 0
#endif

#ifndef NVFP4_EXPERIMENTAL_TCGEN_MMA
#define NVFP4_EXPERIMENTAL_TCGEN_MMA 0
#endif

#ifndef ENABLE_WMMA_KERNEL
#define ENABLE_WMMA_KERNEL 0
#endif

#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
    #define CP_ASYNC_SUPPORTED 1
#else
    #define CP_ASYNC_SUPPORTED 0
#endif

    namespace wmma = nvcuda::wmma;

    constexpr int WMMA_M = 16;
    constexpr int WMMA_N = 8;
    constexpr int WMMA_K = 64;
    constexpr int CTA_M = WMMA_M * 8;
    constexpr int CTA_N = WMMA_N;
    constexpr int CTA_WARPS = CTA_M / WMMA_M;

    constexpr int TCGEN_CTA_M = WMMA_M * 4;           // 64 rows per CTA (one warpgroup)
    constexpr int TCGEN_CTA_WARPS = TCGEN_CTA_M / WMMA_M;
    constexpr int TCGEN_THREADS = TCGEN_CTA_WARPS * 32;

    enum class KernelPath : int {
        Auto = 0,
        Scalar = 1,
        WMMA = 2,
        TCGEN = 3
    };

    enum class SmallNMode : int {
        Auto = 0,
        Batch = 1,
        Pad = 2,
        Scalar = 3
    };

    inline bool equals_ignore_case(const char* lhs, const char* rhs) {
        if (lhs == nullptr || rhs == nullptr) {
            return false;
        }
        while (*lhs && *rhs) {
            if (std::tolower(*lhs) != std::tolower(*rhs)) {
                return false;
            }
            ++lhs;
            ++rhs;
        }
        return (*lhs == '\0') && (*rhs == '\0');
    }

    inline KernelPath kernel_path_override_from_env() {
        const char* env = std::getenv("NVFP4_FORCE_KERNEL");
        if (env == nullptr || env[0] == '\0') {
            return KernelPath::Auto;
        }
        if (equals_ignore_case(env, "scalar")) {
            return KernelPath::Scalar;
        }
        if (equals_ignore_case(env, "wmma")) {
            return KernelPath::WMMA;
        }
        if (equals_ignore_case(env, "tcgen") || equals_ignore_case(env, "tmemory")) {
            return KernelPath::TCGEN;
        }
        return KernelPath::Auto;
    }

    inline SmallNMode small_n_mode_override_from_env() {
        const char* env = std::getenv("NVFP4_SMALLN_MODE");
        if (env == nullptr || env[0] == '\0') {
            return SmallNMode::Auto;
        }
        if (equals_ignore_case(env, "batch")) {
            return SmallNMode::Batch;
        }
        if (equals_ignore_case(env, "pad")) {
            return SmallNMode::Pad;
        }
        if (equals_ignore_case(env, "scalar")) {
            return SmallNMode::Scalar;
        }
        return SmallNMode::Auto;
    }

    inline bool tcgen_auto_enabled() {
        const char* env = std::getenv("NVFP4_ENABLE_TCGEN_AUTO");
        return env && env[0] != '\0' && env[0] != '0';
    }

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
    inline bool capture_tmem_debug_enabled() {
        const char* env = std::getenv("NVFP4_CAPTURE_TMEM_DEBUG");
        return env && env[0] != '\0' && env[0] != '0';
    }
#endif

    inline bool arch_supports_wmma(const cudaDeviceProp* prop) {
        return ENABLE_WMMA_KERNEL && prop && prop->major >= 8;
    }

    inline bool arch_supports_tcgen(const cudaDeviceProp* prop) {
        return prop && prop->major >= 10;
    }

    // -----------------------------------------------
    // Tensor Memory helpers (alloc, dealloc, copies)
    // -----------------------------------------------

    __device__ __forceinline__ uint32_t smem_ptr(const void* ptr) {
        return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
    }

#if NVFP4_ENABLE_TCGEN
    __device__ __forceinline__ void tcgen_alloc_cols(uint32_t* dst_smem_ptr, uint32_t num_cols) {
        asm volatile(
            "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
            :
            : "r"(smem_ptr(dst_smem_ptr)), "r"(num_cols)
            : "memory"
        );
    }

    __device__ __forceinline__ void tcgen_dealloc_cols(uint32_t tmem_addr, uint32_t num_cols) {
        asm volatile(
            "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
            :
            : "r"(tmem_addr), "r"(num_cols)
            : "memory"
        );
    }
#else
    __device__ __forceinline__ void tcgen_alloc_cols(uint32_t*, uint32_t) {}
    __device__ __forceinline__ void tcgen_dealloc_cols(uint32_t, uint32_t) {}
#endif

    __device__ __forceinline__ uint64_t encode_matrix_component(uint32_t value_bytes) {
        return static_cast<uint64_t>((value_bytes & 0x3FFFFu) >> 4);
    }

    __device__ __forceinline__ uint64_t make_smem_matrix_desc(
        uint32_t smem_addr_bytes,
        uint32_t lead_dim_bytes,
        uint32_t stride_dim_bytes,
        bool lead_is_absolute,
        int swizzle_mode = 0)
    {
        uint64_t desc = 0;
        desc |= encode_matrix_component(smem_addr_bytes) << 0;
        desc |= encode_matrix_component(lead_dim_bytes) << 16;
        desc |= encode_matrix_component(stride_dim_bytes) << 32;
        desc |= (uint64_t)0x1 << 46;  // fixed constant 0b001
        desc |= (uint64_t)(lead_is_absolute ? 1 : 0) << 52;
        desc |= (uint64_t)0xB0 << 53; // fixed constant per spec
        desc |= (uint64_t)(swizzle_mode & 0x7) << 61;
        return desc;
    }

#if NVFP4_ENABLE_TCGEN
    __device__ __forceinline__ uint32_t tcgen_cp_async(
        uint32_t tmem_addr,
        uint64_t smem_desc,
        int shape_selector)
    {
        switch (shape_selector) {
            case 0:
                asm volatile(
                    "tcgen05.cp.cta_group::1.128x256b.b8x16.b4x16_p64 [%0], %1;\n"
                    :
                    : "r"(tmem_addr), "l"(smem_desc)
                    : "memory"
                );
                break;
            default:
                break;
        }
        return tmem_addr;
    }
#else
    __device__ __forceinline__ uint32_t tcgen_cp_async(uint32_t tmem_addr, uint64_t, int) {
        return tmem_addr;
    }
#endif

#if NVFP4_ENABLE_TCGEN
    __device__ __forceinline__ uint32_t tmem_add_lane_offset(uint32_t base_addr, uint32_t lane_offset) {
        const uint32_t column = base_addr & 0xFFFFu;
        const uint32_t lane = ((base_addr >> 16) + lane_offset) & 0xFFFFu;
        return (lane << 16) | column;
    }

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
    __device__ float g_debug_tm_frag[TCGEN_THREADS * 4];
    __device__ float g_debug_ref_frag[TCGEN_THREADS * 4];
#endif
#else
    __device__ __forceinline__ uint32_t tmem_add_lane_offset(uint32_t base_addr, uint32_t) {
        return base_addr;
    }
#endif

    __device__ __forceinline__ void mbarrier_init_cta(uint64_t* bar, uint32_t pending_count) {
#if defined(__CUDA_ARCH__)
        asm volatile(
            "mbarrier.init.shared::cta.b64 [%0], %1;\n"
            :
            : "r"(smem_ptr(bar)), "r"(pending_count)
            : "memory"
        );
#else
        (void)bar;
        (void)pending_count;
#endif
    }

    __device__ __forceinline__ void tcgen_commit_mbarrier(uint64_t* bar) {
#if defined(__CUDA_ARCH__) && defined(__CUDA_ARCH_FAMILY_SPECIFIC__)
        asm volatile(
            "tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];\n"
            :
            : "r"(smem_ptr(bar))
            : "memory"
        );
#else
        (void)bar;
#endif
    }

    __device__ __forceinline__ void mbarrier_wait_parity_cta(uint64_t* bar, uint32_t parity) {
#if defined(__CUDA_ARCH__)
        unsigned int complete = 0;
        do {
            asm volatile(
                "{\n\t"
                ".reg .pred p;\n\t"
                "mbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n\t"
                "selp.u32 %0, 1, 0, p;\n\t"
                "}\n"
                : "=r"(complete)
                : "r"(smem_ptr(bar)), "r"(parity)
                : "memory"
            );
        } while (!complete);
#else
        (void)bar;
        (void)parity;
#endif
    }

    __device__ __forceinline__ uint32_t make_mxf4nvf4_instr_desc(
        int m_dim,
        int n_dim,
        bool scale_is_ue4m3,
        int scaleA_id,
        int scaleB_id)
    {
        uint32_t desc = 0;
        desc |= (scaleB_id & 0x3) << 4;
        desc |= (1u) << 7;   // atype = E2M1
        desc |= (1u) << 10;  // btype = E2M1
        desc |= (uint32_t)((n_dim >> 3) & 0x3F) << 17;
        desc |= (scale_is_ue4m3 ? 0u : 1u) << 23;
        desc |= (uint32_t)((m_dim >> 7) & 0x3) << 27;
        desc |= (scaleA_id & 0x3) << 29;
        return desc;
    }


    inline bool tcgen_small_n_allowed(int L, SmallNMode mode) {
        switch (mode) {
            case SmallNMode::Batch:
                return L >= WMMA_N;
            case SmallNMode::Pad:
                return L >= 1;
            case SmallNMode::Scalar:
                return false;
            case SmallNMode::Auto:
            default:
                return L >= 4;
        }
    }

    inline KernelPath choose_kernel_path(
        const cudaDeviceProp* prop,
        KernelPath override,
        SmallNMode small_n_mode,
        int K_bytes,
        int L)
    {
        if (override != KernelPath::Auto) {
            return override;
        }

        const bool k_aligned = ((K_bytes * 2) % WMMA_K) == 0;
        if (tcgen_auto_enabled() && arch_supports_tcgen(prop) && k_aligned && tcgen_small_n_allowed(L, small_n_mode)) {
#if NVFP4_ENABLE_TCGEN
            return KernelPath::TCGEN;
#endif
        }

        if (arch_supports_wmma(prop) && (K_bytes % (WMMA_K / 2) == 0) && L >= CTA_N) {
            return KernelPath::WMMA;
        }

        return KernelPath::Scalar;
    }

    // ----------------------- FP4/FP8 converters -----------------------
    __device__ __constant__ float FP4_E2M1_LUT[16] = {
        0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
        0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
    };

    __device__ __constant__ float FP8_E4M3_LUT[256] = {
        0.0f, 0.0019531250f, 0.0039062500f, 0.0058593750f,
        0.0078125000f, 0.0097656250f, 0.0117187500f, 0.0136718750f,
        0.0156250000f, 0.0175781250f, 0.0195312500f, 0.0214843750f,
        0.0234375000f, 0.0253906250f, 0.0273437500f, 0.0292968750f,
        0.0312500000f, 0.0351562500f, 0.0390625000f, 0.0429687500f,
        0.0468750000f, 0.0507812500f, 0.0546875000f, 0.0585937500f,
        0.0625000000f, 0.0703125000f, 0.0781250000f, 0.0859375000f,
        0.0937500000f, 0.1015625000f, 0.1093750000f, 0.1171875000f,
        0.1250000000f, 0.1406250000f, 0.1562500000f, 0.1718750000f,
        0.1875000000f, 0.2031250000f, 0.2187500000f, 0.2343750000f,
        0.2500000000f, 0.2812500000f, 0.3125000000f, 0.3437500000f,
        0.3750000000f, 0.4062500000f, 0.4375000000f, 0.4687500000f,
        0.5000000000f, 0.5625000000f, 0.6250000000f, 0.6875000000f,
        0.7500000000f, 0.8125000000f, 0.8750000000f, 0.9375000000f,
        1.0000000000f, 1.1250000000f, 1.2500000000f, 1.3750000000f,
        1.5000000000f, 1.6250000000f, 1.7500000000f, 1.8750000000f,
        2.0000000000f, 2.2500000000f, 2.5000000000f, 2.7500000000f,
        3.0000000000f, 3.2500000000f, 3.5000000000f, 3.7500000000f,
        4.0000000000f, 4.5000000000f, 5.0000000000f, 5.5000000000f,
        6.0000000000f, 6.5000000000f, 7.0000000000f, 7.5000000000f,
        8.0000000000f, 9.0000000000f, 10.0000000000f, 11.0000000000f,
        12.0000000000f, 13.0000000000f, 14.0000000000f, 15.0000000000f,
        16.0000000000f, 18.0000000000f, 20.0000000000f, 22.0000000000f,
        24.0000000000f, 26.0000000000f, 28.0000000000f, 30.0000000000f,
        32.0000000000f, 36.0000000000f, 40.0000000000f, 44.0000000000f,
        48.0000000000f, 52.0000000000f, 56.0000000000f, 60.0000000000f,
        64.0000000000f, 72.0000000000f, 80.0000000000f, 88.0000000000f,
        96.0000000000f, 104.0000000000f, 112.0000000000f, 120.0000000000f,
        128.0000000000f, 144.0000000000f, 160.0000000000f, 176.0000000000f,
        192.0000000000f, 208.0000000000f, 224.0000000000f, 240.0000000000f,
        256.0000000000f, 288.0000000000f, 320.0000000000f, 352.0000000000f,
        384.0000000000f, 416.0000000000f, 448.0000000000f, 448.0000000000f,
        -0.0f, -0.0019531250f, -0.0039062500f, -0.0058593750f,
        -0.0078125000f, -0.0097656250f, -0.0117187500f, -0.0136718750f,
        -0.0156250000f, -0.0175781250f, -0.0195312500f, -0.0214843750f,
        -0.0234375000f, -0.0253906250f, -0.0273437500f, -0.0292968750f,
        -0.0312500000f, -0.0351562500f, -0.0390625000f, -0.0429687500f,
        -0.0468750000f, -0.0507812500f, -0.0546875000f, -0.0585937500f,
        -0.0625000000f, -0.0703125000f, -0.0781250000f, -0.0859375000f,
        -0.0937500000f, -0.1015625000f, -0.1093750000f, -0.1171875000f,
        -0.1250000000f, -0.1406250000f, -0.1562500000f, -0.1718750000f,
        -0.1875000000f, -0.2031250000f, -0.2187500000f, -0.2343750000f,
        -0.2500000000f, -0.2812500000f, -0.3125000000f, -0.3437500000f,
        -0.3750000000f, -0.4062500000f, -0.4375000000f, -0.4687500000f,
        -0.5000000000f, -0.5625000000f, -0.6250000000f, -0.6875000000f,
        -0.7500000000f, -0.8125000000f, -0.8750000000f, -0.9375000000f,
        -1.0000000000f, -1.1250000000f, -1.2500000000f, -1.3750000000f,
        -1.5000000000f, -1.6250000000f, -1.7500000000f, -1.8750000000f,
        -2.0000000000f, -2.2500000000f, -2.5000000000f, -2.7500000000f,
        -3.0000000000f, -3.2500000000f, -3.5000000000f, -3.7500000000f,
        -4.0000000000f, -4.5000000000f, -5.0000000000f, -5.5000000000f,
        -6.0000000000f, -6.5000000000f, -7.0000000000f, -7.5000000000f,
        -8.0000000000f, -9.0000000000f, -10.0000000000f, -11.0000000000f,
        -12.0000000000f, -13.0000000000f, -14.0000000000f, -15.0000000000f,
        -16.0000000000f, -18.0000000000f, -20.0000000000f, -22.0000000000f,
        -24.0000000000f, -26.0000000000f, -28.0000000000f, -30.0000000000f,
        -32.0000000000f, -36.0000000000f, -40.0000000000f, -44.0000000000f,
        -48.0000000000f, -52.0000000000f, -56.0000000000f, -60.0000000000f,
        -64.0000000000f, -72.0000000000f, -80.0000000000f, -88.0000000000f,
        -96.0000000000f, -104.0000000000f, -112.0000000000f, -120.0000000000f,
        -128.0000000000f, -144.0000000000f, -160.0000000000f, -176.0000000000f,
        -192.0000000000f, -208.0000000000f, -224.0000000000f, -240.0000000000f,
        -256.0000000000f, -288.0000000000f, -320.0000000000f, -352.0000000000f,
        -384.0000000000f, -416.0000000000f, -448.0000000000f, -448.0000000000f
    };


    __device__ __forceinline__ float decode_fp4_e2m1(uint8_t nibble) {
        return FP4_E2M1_LUT[nibble & 0xF];
    }

    __device__ __forceinline__ float decode_fp8_e4m3(uint8_t byte_val) {
        return FP8_E4M3_LUT[byte_val];
    }

    __device__ __forceinline__ int frag_a_row(int group_id, int elem_idx) {
        return ((elem_idx < 8) || (elem_idx >= 16 && elem_idx < 24)) ? group_id : (group_id + 8);
    }

    __device__ __forceinline__ int frag_a_col(int tid_in_group, int elem_idx) {
        int col = tid_in_group * 8 + (elem_idx & 7);
        if (elem_idx >= 16) {
            col += 32;
        }
        return col;
    }

    __device__ __forceinline__ int frag_b_row(int tid_in_group, int elem_idx) {
        int row = tid_in_group * 8 + (elem_idx & 7);
        if (elem_idx >= 8) {
            row += 32;
        }
        return row;
    }

    __device__ __forceinline__ int frag_b_col(int group_id) {
        return group_id;
    }

    __device__ __forceinline__ int frag_acc_row(int group_id, int acc_idx) {
        return (acc_idx < 2) ? group_id : (group_id + 8);
    }

    __device__ __forceinline__ int frag_acc_col(int tid_in_group, int acc_idx) {
        return tid_in_group * 2 + (acc_idx & 1);
    }

    __device__ __forceinline__ uint8_t load_stage_a_nibble(
        const uint8_t* tile,
        int row_local,
        int col_local)
    {
        const int bytes_per_row = WMMA_K / 2;
        const uint8_t byte_val = tile[row_local * bytes_per_row + (col_local >> 1)];
        return (col_local & 1) ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
    }

    __device__ __forceinline__ uint8_t load_stage_b_nibble(
        const uint8_t* tile,
        int col_local)
    {
        const uint8_t byte_val = tile[col_local >> 1];
        return (col_local & 1) ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
    }

    __device__ __forceinline__ uint8_t load_stage_sfa_byte(
        const uint8_t* tile,
        int row_local,
        int elem)
    {
        return tile[row_local * 4 + elem];
    }

    __device__ __forceinline__ uint8_t load_stage_sfb_byte(
        const uint8_t* tile,
        int elem)
    {
        return tile[elem];
    }

    __device__ __forceinline__ uint8_t load_fp4_nibble(
        const uint8_t* __restrict__ tensor,
        int dim_m,
        int dim_l,
        int K_vals,
        int global_row,
        int global_k,
        int global_l,
        int64_t stride_m,
        int64_t stride_k,
        int64_t stride_l)
    {
        if (global_row < 0 || global_row >= dim_m || global_l < 0 || global_l >= dim_l) {
            return 0;
        }
        if (global_k < 0 || global_k >= K_vals) {
            return 0;
        }
        const int byte_index = global_k >> 1;
        const bool high = (global_k & 1);
        const int64_t offset = (int64_t)global_row * stride_m +
                               (int64_t)byte_index * stride_k +
                               (int64_t)global_l * stride_l;
        const uint8_t byte_val = tensor[offset];
        return high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
    }

    __device__ __forceinline__ uint8_t load_fp4_nibble_linear(
        const uint8_t* data,
        int K_vals,
        int global_k)
    {
        if (global_k < 0 || global_k >= K_vals) {
            return 0;
        }
        const int byte_index = global_k >> 1;
        const bool high = (global_k & 1);
        const uint8_t byte_val = data[byte_index];
        return high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
    }

    __device__ __forceinline__ uint8_t load_scale_byte(
        const uint8_t* __restrict__ tensor,
        int dim_m,
        int dim_l,
        int K_scales,
        int global_row,
        int scale_idx,
        int global_l,
        int64_t stride_m,
        int64_t stride_k,
        int64_t stride_l)
    {
        if (global_row < 0 || global_row >= dim_m || global_l < 0 || global_l >= dim_l) {
            return 0;
        }
        if (scale_idx < 0 || scale_idx >= K_scales) {
            return 0;
        }
        const int64_t offset = (int64_t)global_row * stride_m +
                               (int64_t)scale_idx * stride_k +
                               (int64_t)global_l * stride_l;
        return tensor[offset];
    }

    __global__ void bmv_kernel_scalar(
        const uint8_t* __restrict__ A,
        const uint8_t* __restrict__ B,
        const uint8_t* __restrict__ SFA,
        const uint8_t* __restrict__ SFB,
        __half* __restrict__ C,
        int M, int M_B, int K_bytes, int L,
        int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
        int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
        int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
        int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
    {
        const int l = blockIdx.y;
        if (l >= L) {
            return;
        }

        const int m = blockIdx.x * blockDim.x + threadIdx.x;
        const int K_scales = K_bytes / 8;

        if (K_scales == 0) {
            return;
        }

        const int m_b = 0;  // competition inputs broadcast B/SFB across M

        extern __shared__ uint8_t shared_bytes[];
        uint8_t* sh_B = shared_bytes;
        const size_t sh_B_size = ((size_t)K_bytes + 15) & ~size_t(15);
        float* sh_SFB = reinterpret_cast<float*>(shared_bytes + sh_B_size);

        const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)l * b_stride_l;
        const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)l * sfb_stride_l;

#if CP_ASYNC_SUPPORTED
        const bool can_async = (b_stride_k == 1);
        if (can_async) {
            const int chunk = 16;
            const int stride = blockDim.x * chunk;
            for (int idx = threadIdx.x * chunk; idx + chunk <= K_bytes; idx += stride) {
                void* dst = sh_B + idx;
                const void* src = B + b_row_base + idx;
                unsigned smem_addr = static_cast<unsigned>(__cvta_generic_to_shared(dst));
                unsigned long long gmem_addr = __cvta_generic_to_global(src);
                asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(smem_addr), "l"(gmem_addr));
            }
            asm volatile("cp.async.commit_group;\n" ::);
            asm volatile("cp.async.wait_all;\n" ::);

            const int tail_start = (K_bytes & ~15);
            for (int idx = tail_start + threadIdx.x; idx < K_bytes; idx += blockDim.x) {
                sh_B[idx] = B[b_row_base + (int64_t)idx * b_stride_k];
            }
        } else
#endif
        {
            for (int idx = threadIdx.x; idx < K_bytes; idx += blockDim.x) {
                const int64_t g_idx = b_row_base + (int64_t)idx * b_stride_k;
                sh_B[idx] = B[g_idx];
            }
        }

        for (int idx = threadIdx.x; idx < K_scales; idx += blockDim.x) {
            const int64_t g_idx = sfb_row_base + (int64_t)idx * sfb_stride_k;
            sh_SFB[idx] = decode_fp8_e4m3(SFB[g_idx]);
        }

        __syncthreads();

        if (m >= M) {
            return;
        }

        const int64_t a_row_base = (int64_t)m * a_stride_m + (int64_t)l * a_stride_l;
        const int64_t sfa_row_base = (int64_t)m * sfa_stride_m + (int64_t)l * sfa_stride_l;

        float acc = 0.0f;

        for (int g = 0; g < K_scales; ++g) {
            const int64_t sfa_idx = sfa_row_base + (int64_t)g * sfa_stride_k;

            const float scale_a = decode_fp8_e4m3(SFA[sfa_idx]);
            const float scale_b = sh_SFB[g];
            const float block_scale = scale_a * scale_b;

            if (block_scale == 0.0f) {
                continue;
            }

            float group_sum = 0.0f;
            const int byte_start = g * 8;

            #pragma unroll 8
            for (int i = 0; i < 8; ++i) {
                const int byte_idx = byte_start + i;
                const int64_t a_idx = a_row_base + (int64_t)byte_idx * a_stride_k;

                const uint8_t a_byte = A[a_idx];
                const uint8_t b_byte = sh_B[byte_idx];

                const float a0 = decode_fp4_e2m1(a_byte & 0xF);
                const float a1 = decode_fp4_e2m1((a_byte >> 4) & 0xF);
                const float b0 = decode_fp4_e2m1(b_byte & 0xF);
                const float b1 = decode_fp4_e2m1((b_byte >> 4) & 0xF);

                group_sum += a0 * b0;
                group_sum += a1 * b1;
            }

            acc += group_sum * block_scale;
        }

        const int64_t c_idx = (int64_t)m * L + l;
        C[c_idx] = __float2half_rn(acc);
    }

    #if ENABLE_WMMA_KERNEL
    __global__ void bmv_kernel_tensor(
        const uint8_t* __restrict__ A,
        const uint8_t* __restrict__ B,
        const uint8_t* __restrict__ SFA,
        const uint8_t* __restrict__ SFB,
        __half* __restrict__ C,
        int M, int M_B, int K_bytes, int L,
        int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
        int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
        int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
        int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
    {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
        const int tile_m = blockIdx.x * CTA_M;
        const int tile_l = blockIdx.y * CTA_N;

        if (tile_m >= M || tile_l >= L) {
            return;
        }

        const int warp_id = threadIdx.x / warpSize;
        const int lane_id = threadIdx.x % warpSize;

        const int K_vals = K_bytes * 2;
        const int K_scales = K_bytes / 8;
        const int m_b = 0;

        extern __shared__ uint8_t shared_raw[];
        half* sh_A = reinterpret_cast<half*>(shared_raw);
        half* sh_B = sh_A + CTA_M * WMMA_K;
        float* sh_out = reinterpret_cast<float*>(sh_B + CTA_N * WMMA_K);

        wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float> acc_frag;
        wmma::fill_fragment(acc_frag, 0.0f);

        for (int k_base = 0; k_base < K_vals; k_base += WMMA_K) {
            // Stage A tile
            for (int idx = threadIdx.x; idx < CTA_M * WMMA_K; idx += blockDim.x) {
                const int row = idx / WMMA_K;
                const int k_inner = idx % WMMA_K;
                const int global_m = tile_m + row;
                const int global_k = k_base + k_inner;
                float val = 0.0f;
                if (global_m < M && global_k < K_vals) {
                    const int64_t a_row_base = (int64_t)global_m * a_stride_m + (int64_t)tile_l * a_stride_l;
                    const int64_t sfa_row_base = (int64_t)global_m * sfa_stride_m + (int64_t)tile_l * sfa_stride_l;
                    const int g = global_k / 16;
                    if (g < K_scales) {
                        const float scale_a = decode_fp8_e4m3(SFA[sfa_row_base + (int64_t)g * sfa_stride_k]);
                        const int byte_index = global_k >> 1;
                        const uint8_t byte_val = A[a_row_base + (int64_t)byte_index * a_stride_k];
                        const bool high = (global_k & 1);
                        const uint8_t nibble = high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
                        val = decode_fp4_e2m1(nibble) * scale_a;
                    }
                }
                sh_A[idx] = __float2half(val);
            }

            // Stage B tile
            const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)tile_l * b_stride_l;
            const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)tile_l * sfb_stride_l;
            for (int idx = threadIdx.x; idx < CTA_N * WMMA_K; idx += blockDim.x) {
                const int col = idx / WMMA_K;
                const int k_inner = idx % WMMA_K;
                const int global_l = tile_l + col;
                const int global_k = k_base + k_inner;
                float val = 0.0f;
                if (global_l < L && global_k < K_vals) {
                    const int g = global_k / 16;
                    if (g < K_scales) {
                        const float scale_b = decode_fp8_e4m3(SFB[sfb_row_base + (int64_t)g * sfb_stride_k]);
                        const int byte_index = global_k >> 1;
                        const uint8_t byte_val = B[b_row_base + (int64_t)byte_index * b_stride_k];
                        const bool high = (global_k & 1);
                        const uint8_t nibble = high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
                        val = decode_fp4_e2m1(nibble) * scale_b;
                    }
                }
                sh_B[idx] = __float2half(val);
            }

            __syncthreads();

            if (warp_id < (CTA_M / WMMA_M)) {
                const half* warp_A = sh_A + warp_id * WMMA_M * WMMA_K;
                wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, half, wmma::row_major> a_frag;
                wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, half, wmma::col_major> b_frag;
                wmma::load_matrix_sync(a_frag, warp_A, WMMA_K);
                wmma::load_matrix_sync(b_frag, sh_B, WMMA_K);
                wmma::mma_sync(acc_frag, a_frag, b_frag, acc_frag);
            }

            __syncthreads();
        }

        if (warp_id < (CTA_M / WMMA_M)) {
            float* warp_out = sh_out + warp_id * WMMA_M * CTA_N;
            wmma::store_matrix_sync(warp_out, acc_frag, CTA_N, wmma::mem_row_major);
        }

        __syncthreads();

        for (int idx = threadIdx.x; idx < CTA_M * CTA_N; idx += blockDim.x) {
            const int row = idx / CTA_N;
            const int col = idx % CTA_N;
            const int global_m = tile_m + row;
            const int global_l = tile_l + col;
            if (global_m < M && global_l < L) {
                const float val = sh_out[idx];
                C[global_m * L + global_l] = __float2half(val);
            }
        }
#else
        (void)A; (void)B; (void)SFA; (void)SFB; (void)C;
        (void)M; (void)M_B; (void)K_bytes; (void)L;
        (void)a_stride_m; (void)a_stride_k; (void)a_stride_l;
        (void)b_stride_m; (void)b_stride_k; (void)b_stride_l;
        (void)sfa_stride_m; (void)sfa_stride_k; (void)sfa_stride_l;
        (void)sfb_stride_m; (void)sfb_stride_k; (void)sfb_stride_l;
#endif
    }

#endif // ENABLE_WMMA_KERNEL

#if NVFP4_ENABLE_TCGEN
    __global__ void bmv_kernel_tcgen(
        const uint8_t* __restrict__ A,
        const uint8_t* __restrict__ B,
        const uint8_t* __restrict__ SFA,
        const uint8_t* __restrict__ SFB,
        __half* __restrict__ C,
        int M, int M_B, int K_bytes, int L,
        int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
        int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
        int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
        int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
    {
#if defined(__CUDA_ARCH__) && defined(__CUDA_ARCH_FAMILY_SPECIFIC__)
        constexpr uint32_t TMEM_COLS_ACC = 32;
        constexpr uint32_t TMEM_COLS_A   = 32;
        constexpr uint32_t TMEM_COLS_B   = 32;
        constexpr uint32_t TMEM_COLS_SFA = 32;
        constexpr uint32_t TMEM_COLS_SFB = 32;

        __shared__ uint32_t sh_tmem_acc;
        __shared__ uint32_t sh_tmem_a;
        __shared__ uint32_t sh_tmem_b;
        __shared__ uint32_t sh_tmem_sfa;
        __shared__ uint32_t sh_tmem_sfb;
        __shared__ uint64_t sh_mbarrier;

        const int tile_m = blockIdx.x * TCGEN_CTA_M;
        const int tile_l = blockIdx.y;

        if (tile_m >= M || tile_l >= L) {
            return;
        }

        const int warp_id = threadIdx.x / warpSize;
        const int lane_id = threadIdx.x % warpSize;
        if (warp_id >= TCGEN_CTA_WARPS) {
            return;
        }

        const int lane_group = lane_id >> 2;
        const int lane_tid4 = lane_id & 3;

        const int warp_row_base = tile_m + warp_id * WMMA_M;
        if (warp_row_base >= M) {
            return;
        }

        const int K_vals = K_bytes * 2;
        const int K_scales = K_bytes / 8;
        const int m_b = 0;
        const int b_row = 0;

        if (threadIdx.x == 0) {
            tcgen_alloc_cols(&sh_tmem_acc, TMEM_COLS_ACC);
            tcgen_alloc_cols(&sh_tmem_a, TMEM_COLS_A);
            tcgen_alloc_cols(&sh_tmem_b, TMEM_COLS_B);
            tcgen_alloc_cols(&sh_tmem_sfa, TMEM_COLS_SFA);
            tcgen_alloc_cols(&sh_tmem_sfb, TMEM_COLS_SFB);
        }
        __syncthreads();
        const uint32_t tmem_acc_addr = sh_tmem_acc;
        const uint32_t tmem_a_addr = sh_tmem_a;
        const uint32_t tmem_b_addr = sh_tmem_b;
        const uint32_t tmem_sfa_addr = sh_tmem_sfa;
        const uint32_t tmem_sfb_addr = sh_tmem_sfb;
        (void)tmem_acc_addr;
        if (threadIdx.x == 0) {
            mbarrier_init_cta(&sh_mbarrier, 0);
        }
        __syncthreads();
        uint32_t mbarrier_parity = 0;

        extern __shared__ uint8_t shared_tiles[];
        struct StageBuffers {
            uint8_t* A;
            uint8_t* B;
            uint8_t* SFA;
            uint8_t* SFB;
        };
        const size_t bytes_A_row = WMMA_K / 2;
        const size_t bytes_A_tile = (size_t)TCGEN_CTA_M * bytes_A_row;
        const size_t bytes_B_tile = WMMA_K / 2;
        const size_t bytes_SFA_tile = (size_t)TCGEN_CTA_M * 4;
        const size_t bytes_SFB_tile = 4;
        const size_t stage_bytes = bytes_A_tile + bytes_B_tile + bytes_SFA_tile + bytes_SFB_tile;

        auto stage_buffers = [&](int idx) {
            StageBuffers buf;
            uint8_t* base = shared_tiles + idx * stage_bytes;
            buf.A = base;
            buf.B = buf.A + bytes_A_tile;
            buf.SFA = buf.B + bytes_B_tile;
            buf.SFB = buf.SFA + bytes_SFA_tile;
            return buf;
        };
        auto make_desc = [&](const void* ptr, uint32_t lead_bytes, uint32_t stride_bytes) {
            return make_smem_matrix_desc(smem_ptr(ptr), lead_bytes, stride_bytes, true);
        };

        const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)tile_l * b_stride_l;
        const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)tile_l * sfb_stride_l;

        auto load_stage = [&](const StageBuffers& buf, int k_base_panel) {
            const int byte_offset = k_base_panel >> 1;
            for (int idx = threadIdx.x; idx < bytes_A_tile; idx += blockDim.x) {
                const int row_local = idx / bytes_A_row;
                const int byte_in_row = idx % bytes_A_row;
                const int global_row = tile_m + row_local;
                const int global_byte = byte_offset + byte_in_row;
                uint8_t val = 0;
                if (global_row < M && global_byte < K_bytes) {
                    const int64_t a_offset = (int64_t)global_row * a_stride_m +
                                             (int64_t)global_byte * a_stride_k +
                                             (int64_t)tile_l * a_stride_l;
                    val = A[a_offset];
                }
                buf.A[idx] = val;
            }

#if CP_ASYNC_SUPPORTED
            if (b_stride_k == 1) {
                const int chunk = 16;
                const int copy_bytes = bytes_B_tile & ~ (chunk - 1);
                for (int off = threadIdx.x * chunk; off < copy_bytes; off += blockDim.x * chunk) {
                    const int global_byte = byte_offset + off;
                    const uint32_t dst = smem_ptr(buf.B + off);
                    const unsigned long long src = __cvta_generic_to_global(B + b_row_base + global_byte);
                    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst), "l"(src));
                }
                asm volatile("cp.async.commit_group;\n" ::);
                asm volatile("cp.async.wait_all;\n" ::);
                for (int off = copy_bytes + threadIdx.x; off < bytes_B_tile; off += blockDim.x) {
                    const int global_byte = byte_offset + off;
                    uint8_t val = 0;
                    if (global_byte < K_bytes) {
                        const int64_t b_offset = b_row_base + (int64_t)global_byte * b_stride_k;
                        val = B[b_offset];
                    }
                    buf.B[off] = val;
                }
            } else
#endif
            {
                for (int idx = threadIdx.x; idx < bytes_B_tile; idx += blockDim.x) {
                    const int global_byte = byte_offset + idx;
                    uint8_t val = 0;
                    if (global_byte < K_bytes) {
                        const int64_t b_offset = b_row_base + (int64_t)global_byte * b_stride_k;
                        val = B[b_offset];
                    }
                    buf.B[idx] = val;
                }
            }

            const int scale_offset = k_base_panel >> 4;
            for (int idx = threadIdx.x; idx < bytes_SFA_tile; idx += blockDim.x) {
                const int row_local = idx / 4;
                const int elem = idx & 3;
                const int global_row = tile_m + row_local;
                const int scale_idx = scale_offset + elem;
                uint8_t val = 0;
                if (global_row < M && scale_idx < K_scales) {
                    const int64_t sfa_offset = (int64_t)global_row * sfa_stride_m +
                                               (int64_t)scale_idx * sfa_stride_k +
                                               (int64_t)tile_l * sfa_stride_l;
                    val = SFA[sfa_offset];
                }
                buf.SFA[idx] = val;
            }

            for (int idx = threadIdx.x; idx < bytes_SFB_tile; idx += blockDim.x) {
                const int scale_idx = scale_offset + idx;
                uint8_t val = 0;
                if (scale_idx < K_scales) {
                    const int64_t sfb_offset = sfb_row_base + (int64_t)scale_idx * sfb_stride_k;
                    val = SFB[sfb_offset];
                }
                buf.SFB[idx] = val;
            }
        };

        StageBuffers stage0 = stage_buffers(0);
        StageBuffers stage1 = stage_buffers(1);
        load_stage(stage0, 0);
        __syncthreads();
        int stage_idx = 0;

        float d0 = 0.0f, d1 = 0.0f, d2 = 0.0f, d3 = 0.0f;
        const uint16_t bidA = 0, tidA = 0;
        const uint16_t bidB = 0, tidB = 0;

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
        const uint32_t instr_desc_tcgen = make_mxf4nvf4_instr_desc(TCGEN_CTA_M, WMMA_N, true, 0, 0);
#endif

        for (int k_base = 0; k_base < K_vals; k_base += WMMA_K) {
            const StageBuffers& stage = (stage_idx == 0) ? stage0 : stage1;
            uint32_t a0 = 0u, a1 = 0u, a2 = 0u, a3 = 0u;
            uint32_t b0 = 0u, b1 = 0u;
            uint32_t scaleAData = 0u;
            uint32_t scaleBData = 0u;

            if (threadIdx.x == 0) {
                const uint64_t descA = make_desc(stage.A, bytes_A_row, bytes_A_tile);
                const uint64_t descB = make_desc(stage.B, bytes_B_tile, bytes_B_tile);
                const uint64_t descSFA = make_desc(stage.SFA, 4u, bytes_SFA_tile);
                const uint64_t descSFB = make_desc(stage.SFB, 4u, bytes_SFB_tile);
                tcgen_cp_async(tmem_a_addr, descA, 0);
                tcgen_cp_async(tmem_b_addr, descB, 0);
                tcgen_cp_async(tmem_sfa_addr, descSFA, 0);
                tcgen_cp_async(tmem_sfb_addr, descSFB, 0);
            }
            __syncthreads();
            if (threadIdx.x == 0) {
                tcgen_commit_mbarrier(&sh_mbarrier);
            }
            mbarrier_wait_parity_cta(&sh_mbarrier, mbarrier_parity);
            mbarrier_parity ^= 1;

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
            if (threadIdx.x == 0) {
                const uint64_t descB_tc = make_desc(stage.B, bytes_B_tile, bytes_B_tile);
                const uint32_t enable_flag = (k_base == 0) ? 0u : 1u;
                asm volatile(
                    "{\n\t"
                    ".reg .pred p_enable;\n\t"
                    "setp.ne.u32 p_enable, %6, 0;\n\t"
                    "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.scale_vec::4X "
                    "[%0], [%1], %2, %3, [%4], [%5], p_enable;\n\t"
                    "}\n"
                    :
                    : "r"(tmem_acc_addr),
                      "r"(tmem_a_addr),
                      "l"(descB_tc),
                      "r"(instr_desc_tcgen),
                      "r"(tmem_sfa_addr),
                      "r"(tmem_sfb_addr),
                      "r"(enable_flag)
                    : "memory"
                );
            }
#endif
#pragma unroll
            for (int elem = 0; elem < 32; ++elem) {
                const int row_local = frag_a_row(lane_group, elem);
                const int col_local = frag_a_col(lane_tid4, elem);
                const int row_cta = warp_id * WMMA_M + row_local;
                const uint8_t nibble = load_stage_a_nibble(stage.A, row_cta, col_local);
                uint32_t* dst = (elem < 8) ? &a0 : (elem < 16) ? &a1 : (elem < 24) ? &a2 : &a3;
                const int shift = (elem & 7) * 4;
                *dst |= uint32_t(nibble & 0xF) << shift;
            }

#pragma unroll
            for (int elem = 0; elem < 16; ++elem) {
                const int row_local = frag_b_row(lane_tid4, elem);
                const uint8_t nibble = load_stage_b_nibble(stage.B, row_local);
                uint32_t* dst = (elem < 8) ? &b0 : &b1;
                const int shift = (elem & 7) * 4;
                *dst |= uint32_t(nibble & 0xF) << shift;
            }

            if (lane_tid4 <= 1) {
                const int row_local = lane_group + (lane_tid4 ? 8 : 0);
                const int row_cta = warp_id * WMMA_M + row_local;
                uint8_t s0 = load_stage_sfa_byte(stage.SFA, row_cta, 0);
                uint8_t s1 = load_stage_sfa_byte(stage.SFA, row_cta, 1);
                uint8_t s2 = load_stage_sfa_byte(stage.SFA, row_cta, 2);
                uint8_t s3 = load_stage_sfa_byte(stage.SFA, row_cta, 3);
                scaleAData = uint32_t(s0) | (uint32_t(s1) << 8) | (uint32_t(s2) << 16) | (uint32_t(s3) << 24);
            }

            if (lane_group == 0 && lane_tid4 == 0) {
                uint8_t s0 = load_stage_sfb_byte(stage.SFB, 0);
                uint8_t s1 = load_stage_sfb_byte(stage.SFB, 1);
                uint8_t s2 = load_stage_sfb_byte(stage.SFB, 2);
                uint8_t s3 = load_stage_sfb_byte(stage.SFB, 3);
                scaleBData = uint32_t(s0) | (uint32_t(s1) << 8) | (uint32_t(s2) << 16) | (uint32_t(s3) << 24);
            }

            float c0 = d0, c1 = d1, c2 = d2, c3 = d3;
            asm volatile(
                "mma.sync.aligned.m16n8k64.row.col.kind::mxf4nvf4.block_scale.scale_vec::4X "
                ".f32.e2m1.e2m1.f32.ue4m3 "
                "{%0,%1,%2,%3}, "
                "{%4,%5,%6,%7}, "
                "{%8,%9}, "
                "{%10,%11,%12,%13}, "
                "%14, {%15,%16}, "
                "%17, {%18,%19};\n"
                : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
                : "r"(a0), "r"(a1), "r"(a2), "r"(a3),
                  "r"(b0), "r"(b1),
                  "f"(c0), "f"(c1), "f"(c2), "f"(c3),
                  "r"(scaleAData), "h"(bidA), "h"(tidA),
                  "r"(scaleBData), "h"(bidB), "h"(tidB)
            );
            stage_idx ^= 1;
            const int next_k = k_base + WMMA_K;
            if (next_k < K_vals) {
                const StageBuffers& next_stage = (stage_idx == 0) ? stage0 : stage1;
                load_stage(next_stage, next_k);
            }
            __syncthreads();
        }

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
        const float wmma_ref0 = d0;
        const float wmma_ref1 = d1;
        const float wmma_ref2 = d2;
        const float wmma_ref3 = d3;

        __syncthreads();
        if (threadIdx.x == 0) {
            tcgen_commit_mbarrier(&sh_mbarrier);
        }
        mbarrier_wait_parity_cta(&sh_mbarrier, mbarrier_parity);
        mbarrier_parity ^= 1;

        const uint32_t warp_lane_base = tmem_add_lane_offset(tmem_acc_addr, static_cast<uint32_t>(warp_id * 32));
        float tm_d0, tm_d1, tm_d2, tm_d3;
        asm volatile(
                    "tcgen05.ld.sync.aligned.16x64b.x4.b32 {%0,%1,%2,%3}, [%4];\n"
            : "=f"(tm_d0), "=f"(tm_d1), "=f"(tm_d2), "=f"(tm_d3)
            : "r"(warp_lane_base)
            : "memory"
        );
        d0 = tm_d0;
        d1 = tm_d1;
        d2 = tm_d2;
        d3 = tm_d3;
        if (blockIdx.x == 0 && blockIdx.y == 0) {
            const int base_idx = threadIdx.x * 4;
            g_debug_tm_frag[base_idx + 0] = tm_d0;
            g_debug_tm_frag[base_idx + 1] = tm_d1;
            g_debug_tm_frag[base_idx + 2] = tm_d2;
            g_debug_tm_frag[base_idx + 3] = tm_d3;
            g_debug_ref_frag[base_idx + 0] = wmma_ref0;
            g_debug_ref_frag[base_idx + 1] = wmma_ref1;
            g_debug_ref_frag[base_idx + 2] = wmma_ref2;
            g_debug_ref_frag[base_idx + 3] = wmma_ref3;
        }
#endif

        float row_sum_top = 0.0f;
        float row_sum_bottom = 0.0f;
#pragma unroll
        for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
            const float value = (acc_idx == 0) ? d0 : (acc_idx == 1) ? d1 : (acc_idx == 2) ? d2 : d3;
            const int row_local = frag_acc_row(lane_group, acc_idx);
            if (row_local == lane_group) {
                row_sum_top += value;
            } else {
                row_sum_bottom += value;
            }
        }

        for (int offset = 2; offset > 0; offset >>= 1) {
            row_sum_top += __shfl_xor_sync(0xFFFFFFFF, row_sum_top, offset, 4);
            row_sum_bottom += __shfl_xor_sync(0xFFFFFFFF, row_sum_bottom, offset, 4);
        }

        if (lane_tid4 == 0) {
            const int global_col = tile_l;
            const int row0 = warp_row_base + lane_group;
            if (row0 < M && global_col < L) {
                const int64_t out_idx = (int64_t)row0 * L + global_col;
                C[out_idx] = __float2half_rn(row_sum_top);
            }
            const int row1 = warp_row_base + lane_group + 8;
            if (row1 < M && global_col < L) {
                const int64_t out_idx = (int64_t)row1 * L + global_col;
                C[out_idx] = __float2half_rn(row_sum_bottom);
            }
        }

        __syncthreads();
        if (threadIdx.x == 0) {
            tcgen_dealloc_cols(tmem_sfb_addr, TMEM_COLS_SFB);
            tcgen_dealloc_cols(tmem_sfa_addr, TMEM_COLS_SFA);
            tcgen_dealloc_cols(tmem_b_addr, TMEM_COLS_B);
            tcgen_dealloc_cols(tmem_a_addr, TMEM_COLS_A);
            tcgen_dealloc_cols(tmem_acc_addr, TMEM_COLS_ACC);
        }
#else
        (void)A; (void)B; (void)SFA; (void)SFB; (void)C;
        (void)M; (void)M_B; (void)K_bytes; (void)L;
        (void)a_stride_m; (void)a_stride_k; (void)a_stride_l;
        (void)b_stride_m; (void)b_stride_k; (void)b_stride_l;
        (void)sfa_stride_m; (void)sfa_stride_k; (void)sfa_stride_l;
        (void)sfb_stride_m; (void)sfb_stride_k; (void)sfb_stride_l;
#endif
    }

#endif // NVFP4_ENABLE_TCGEN


    void bmv_launcher(
        torch::Tensor A,
        torch::Tensor B,
        torch::Tensor SFA,
        torch::Tensor SFB,
        torch::Tensor C)
    {
        TORCH_CHECK(A.dim() == 3, "A must be 3D");
        TORCH_CHECK(B.dim() == 3, "B must be 3D");
        TORCH_CHECK(SFA.dim() == 3, "SFA must be 3D");
        TORCH_CHECK(SFB.dim() == 3, "SFB must be 3D");

        const int64_t M = A.size(0);
        const int64_t K_bytes = A.size(1);
        const int64_t L = A.size(2);
        const int64_t M_B = B.size(0);

        TORCH_CHECK(M_B >= 1, "B must have at least one row");
        TORCH_CHECK(B.size(1) == K_bytes && B.size(2) == L, "B shape mismatch");
        TORCH_CHECK(SFA.size(0) == M && SFA.size(1) == K_bytes / 8 && SFA.size(2) == L, "SFA shape mismatch");
        TORCH_CHECK(SFB.size(0) >= 1 && SFB.size(1) == K_bytes / 8 && SFB.size(2) == L, "SFB shape mismatch");

        const auto a_stride = A.strides();
        const auto b_stride = B.strides();
        const auto sfa_stride = SFA.strides();
        const auto sfb_stride = SFB.strides();

        cudaStream_t stream = c10::cuda::getCurrentCUDAStream().stream();

        auto align_up = [](size_t value, size_t alignment) {
            return (value + alignment - 1) & ~(alignment - 1);
        };

        const cudaDeviceProp* device_prop = at::cuda::getCurrentDeviceProperties();
        const KernelPath override = kernel_path_override_from_env();
        const SmallNMode small_n_mode = small_n_mode_override_from_env();
        const KernelPath selected = choose_kernel_path(device_prop, override, small_n_mode, static_cast<int>(K_bytes), static_cast<int>(L));
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
        const bool capture_tmem_debug = capture_tmem_debug_enabled();
        if (capture_tmem_debug) {
            printf("[nvfp4] selected kernel path = %d\n", static_cast<int>(selected));
        }
#else
        const bool capture_tmem_debug = false;
#endif

#if !NVFP4_ENABLE_TCGEN
        TORCH_CHECK(
            selected != KernelPath::TCGEN,
            "TCGEN kernel requested but this build was compiled without NVFP4 tensor-core support. "
            "Set NVFP4_ENABLE_TCGEN_BUILD=1 before importing submission.py to enable it.");
#endif

        switch (selected) {
#if NVFP4_ENABLE_TCGEN
            case KernelPath::TCGEN: {
                TORCH_CHECK(
                    arch_supports_tcgen(device_prop),
                    "TCGEN path selected but device does not support it.");

                const int threads = TCGEN_THREADS;
                const dim3 grid_tcgen(
                    static_cast<unsigned int>((M + TCGEN_CTA_M - 1) / TCGEN_CTA_M),
                    static_cast<unsigned int>(L)
                );

                const size_t bytes_A_tile = static_cast<size_t>(TCGEN_CTA_M) * (WMMA_K / 2);
                const size_t bytes_B_tile = WMMA_K / 2;
                const size_t bytes_SFA_tile = static_cast<size_t>(TCGEN_CTA_M) * 4;
                const size_t bytes_SFB_tile = 4;
                const size_t stage_bytes = bytes_A_tile + bytes_B_tile + bytes_SFA_tile + bytes_SFB_tile;
                const size_t shared_tcgen = stage_bytes * 2;

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
                if (capture_tmem_debug) {
                    const size_t debug_bytes = static_cast<size_t>(TCGEN_THREADS) * 4 * sizeof(float);
                    cudaMemsetAsync(g_debug_tm_frag, 0, debug_bytes, stream);
                    cudaMemsetAsync(g_debug_ref_frag, 0, debug_bytes, stream);
                }
#endif

                bmv_kernel_tcgen<<<grid_tcgen, threads, shared_tcgen, stream>>>(
                    A.data_ptr<uint8_t>(),
                    B.data_ptr<uint8_t>(),
                    SFA.data_ptr<uint8_t>(),
                    SFB.data_ptr<uint8_t>(),
                    reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
                    static_cast<int>(M),
                    static_cast<int>(M_B),
                    static_cast<int>(K_bytes),
                    static_cast<int>(L),
                    a_stride[0], a_stride[1], a_stride[2],
                    b_stride[0], b_stride[1], b_stride[2],
                    sfa_stride[0], sfa_stride[1], sfa_stride[2],
                    sfb_stride[0], sfb_stride[1], sfb_stride[2]
                );
                break;
            }
#endif
#if ENABLE_WMMA_KERNEL
            case KernelPath::WMMA: {
                TORCH_CHECK(
                    arch_supports_wmma(device_prop),
                    "WMMA tensor path selected but device does not support it.");

                const dim3 grid_tensor(
                    static_cast<unsigned int>((M + CTA_M - 1) / CTA_M),
                    static_cast<unsigned int>((L + CTA_N - 1) / CTA_N)
                );

                const size_t shared_tensor =
                    (static_cast<size_t>(CTA_M) * WMMA_K +
                     static_cast<size_t>(CTA_N) * WMMA_K) * sizeof(half) +
                    static_cast<size_t>(CTA_M) * CTA_N * sizeof(float);

                bmv_kernel_tensor<<<grid_tensor, 256, shared_tensor, stream>>>(
                    A.data_ptr<uint8_t>(),
                    B.data_ptr<uint8_t>(),
                    SFA.data_ptr<uint8_t>(),
                    SFB.data_ptr<uint8_t>(),
                    reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
                    static_cast<int>(M),
                    static_cast<int>(M_B),
                    static_cast<int>(K_bytes),
                    static_cast<int>(L),
                    a_stride[0], a_stride[1], a_stride[2],
                    b_stride[0], b_stride[1], b_stride[2],
                    sfa_stride[0], sfa_stride[1], sfa_stride[2],
                    sfb_stride[0], sfb_stride[1], sfb_stride[2]
                );
                break;
            }
#endif
            case KernelPath::Scalar:
            default: {
                const int threads = 256;
                const dim3 grid_scalar(
                    static_cast<unsigned int>((M + threads - 1) / threads),
                    static_cast<unsigned int>(L)
                );

                const size_t shared_scalar = align_up(static_cast<size_t>(K_bytes), size_t(16)) +
                    static_cast<size_t>(K_bytes / 8) * sizeof(float);

                bmv_kernel_scalar<<<grid_scalar, threads, shared_scalar, stream>>>(
                    A.data_ptr<uint8_t>(),
                    B.data_ptr<uint8_t>(),
                    SFA.data_ptr<uint8_t>(),
                    SFB.data_ptr<uint8_t>(),
                    reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
                    static_cast<int>(M),
                    static_cast<int>(M_B),
                    static_cast<int>(K_bytes),
                    static_cast<int>(L),
                    a_stride[0], a_stride[1], a_stride[2],
                    b_stride[0], b_stride[1], b_stride[2],
                    sfa_stride[0], sfa_stride[1], sfa_stride[2],
                    sfb_stride[0], sfb_stride[1], sfb_stride[2]
                );
                break;
            }
        }


        cudaError_t err = cudaGetLastError();
        if (err != cudaSuccess) {
            TORCH_CHECK(false, "CUDA kernel launch failed: ", cudaGetErrorString(err));
        }

#if NVFP4_EXPERIMENTAL_TCGEN_MMA
        if (selected == KernelPath::TCGEN && capture_tmem_debug) {
            const size_t debug_elems = static_cast<size_t>(TCGEN_THREADS) * 4;
            const size_t debug_bytes = debug_elems * sizeof(float);
            std::vector<float> debug_tm(debug_elems, 0.0f);
            std::vector<float> debug_ref(debug_elems, 0.0f);
            cudaMemcpyFromSymbolAsync(debug_tm.data(), g_debug_tm_frag, debug_bytes, 0, cudaMemcpyDeviceToHost, stream);
            cudaMemcpyFromSymbolAsync(debug_ref.data(), g_debug_ref_frag, debug_bytes, 0, cudaMemcpyDeviceToHost, stream);
            cudaStreamSynchronize(stream);
            printf("==== TMEM debug (CTA 0, warp 0) ====\n");
            for (int thread = 0; thread < TCGEN_THREADS; ++thread) {
                const int warp = thread / 32;
                if (warp > 0) {
                    continue;
                }
                const int lane = thread % 32;
                const int group = lane >> 2;
                const int tid4 = lane & 3;
                for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
                    const float tm_val = debug_tm[thread * 4 + acc_idx];
                    const float ref_val = debug_ref[thread * 4 + acc_idx];
                    if (tm_val == 0.0f && ref_val == 0.0f) {
                        continue;
                    }
                    const int row_local = (acc_idx < 2) ? group : (group + 8);
                    const int col_local = tid4 * 2 + (acc_idx & 1);
                    const float diff = std::fabs(tm_val - ref_val);
                    printf("lane=%d acc=%d row=%d col=%d tm=%f ref=%f diff=%f\n",
                           lane, acc_idx, row_local, col_local, tm_val, ref_val, diff);
                }
            }
            printf("==== TMEM debug end ====\n");
        }
#endif
    }
    """

    # Build the extension
    extra_cuda_cflags = [
        "-O3",
        "-std=c++17",
        "-gencode=arch=compute_80,code=sm_80",
        "-gencode=arch=compute_100,code=sm_100",
        "-U__CUDA_NO_HALF_OPERATORS__",
        "-U__CUDA_NO_HALF_CONVERSIONS__",
        "-U__CUDA_NO_BFLOAT16_CONVERSIONS__",
        "--expt-relaxed-constexpr",
        "-use_fast_math",
        f"-DNVFP4_ENABLE_TCGEN={1 if enable_tcgen_build else 0}",
        f"-DNVFP4_EXPERIMENTAL_TCGEN_MMA={1 if experimental_tcgen else 0}",
    ]
    if os.environ.get("NVFP4_CAPTURE_TMEM_DEBUG"):
        extra_cuda_cflags.extend(["-keep", "--keep-dir=/tmp/nvcc_keep"])
    if enable_tcgen_build:
        generic_flag = "-gencode=arch=compute_100,code=sm_100"
        extra_cuda_cflags = [flag for flag in extra_cuda_cflags if flag != generic_flag]
        extra_cuda_cflags.insert(4, "-gencode=arch=compute_100a,code=sm_100a")
    _EXT = load_inline(
        name="bmv_nvfp4_b200_ext",
        cpp_sources=[CPP_SRC],
        cuda_sources=[CUDA_SRC],
        functions=["bmv_launcher"],
        with_cuda=True,
        extra_cflags=["-O3", "-std=c++17"],
        extra_cuda_cflags=extra_cuda_cflags,
        verbose=False,
    )
    return _EXT


def _convert_to_uint8(tensor):
    """Convert FP4/FP8 dtype tensors to uint8 tensor containing raw bytes."""
    if hasattr(torch, 'float4_e2m1fn') and tensor.dtype == torch.float4_e2m1fn:
        return tensor.view(torch.uint8)
    elif hasattr(torch, 'float4_e2m1fn_x2') and str(tensor.dtype) == 'torch.float4_e2m1fn_x2':
        return tensor.view(torch.uint8)
    elif tensor.dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
        return tensor.view(torch.uint8)
    elif tensor.dtype == torch.uint8:
        return tensor
    else:
        try:
            return tensor.view(torch.uint8)
        except:
            raise ValueError(f"Cannot convert dtype {tensor.dtype} to uint8")


@torch.inference_mode()
def custom_kernel(tensors):
    """
    NVFP4 batched matrix-vector multiplication with FP8 block scales.
    
    Always returns a NEW tensor of shape [M, 1, L] in FP16.
    
    Expected input: (a, b, sfa, sfb, c)
      a:  [M, K_bytes, L], FP4 packed (2 FP4 vals/byte)
      b:  [M_B, K_bytes, L], FP4 packed where M_B can be any value
      sfa:[M, K_bytes/8, L], FP8 scales
      sfb:[M_B, K_bytes/8, L], FP8 scales
      c:  Output tensor (any shape - only used for device)
      
    Returns: NEW tensor of shape [M, 1, L] in FP16
    """
    if len(tensors) >= 5:
        a, b, sfa, sfb, c = tensors[:5]
    else:
        raise ValueError(f"Expected at least 5 tensors, got {len(tensors)}")
    
    # Get dimensions
    M, K_bytes, L = a.shape
    M_B = b.shape[0]
    
    # Validate compatible K and L dimensions
    assert b.shape[1] == K_bytes and b.shape[2] == L, f"B shape {b.shape} incompatible with A shape {a.shape}"
    
    # Validate scale shapes
    K_scales = K_bytes // 8
    assert sfa.shape == (M, K_scales, L), f"SFA shape {sfa.shape} should be ({M}, {K_scales}, {L})"
    assert sfb.shape == (M_B, K_scales, L), f"SFB shape {sfb.shape} should be ({M_B}, {K_scales}, {L})"
    
    # Validate K alignment
    assert (K_bytes * 2) % 16 == 0, f"K must be multiple of 16 FP4 elements"
    
    # Convert to uint8 raw bytes without altering layout (preserve strides)
    a = _convert_to_uint8(a)
    b = _convert_to_uint8(b)
    sfa = _convert_to_uint8(sfa)
    sfb = _convert_to_uint8(sfb)

    # Create NEW output tensor with correct shape [M, 1, L] in FP16
    output = torch.zeros(M, 1, L, dtype=torch.float16, device=c.device).contiguous()
    
    # Call CUDA kernel
    ext = _get_ext()
    ext.bmv_launcher(a, b, sfa, sfb, output)
    
    # Ensure completion
    torch.cuda.synchronize()
    
    return output
scrolls · 1573 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON