Skip to content
KernelIndex
Search⌘K

submission 97474

CatsRCool · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

s_attempt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-97474?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
30.1µs
#147 of 678
2025-11-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f367021859a005605a780c61945a9ba38e144aeb6750bdeae713fb23e1a754b1
license declaredunknown
license concludedunknown
authorsCatsRCool
imported2026-08-15

Techniques

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

clustertemplate <typename ClusterShape>
fp4using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing FusionOperation = cutlass::epilogue::fusion::LinearCombination<ElementD, ElementCompute>;

Kernel source

s_attempt.py368 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils.blockscaled_layout as blockscaled_utils

# ---------------------------------------------------------------------------
# CUTLASS Tensor Core GEMV - Optimized with inline column extraction
# ---------------------------------------------------------------------------

cuda_source = """
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <stdexcept>
#include <map>
#include <tuple>
#include <cstdint>

// Prevent CUTLASS synclog from being included, then stub out the functions
#define CUTLASS_ARCH_SYNCLOG_H_

namespace cutlass {
namespace arch {
    template<typename... Args>
    __device__ __host__ inline void synclog_emit_tma_load(Args...) {}
    template<typename... Args>
    __device__ __host__ inline void synclog_emit_tma_store(Args...) {}
    template<typename... Args>
    __device__ __host__ inline void synclog_emit_fence_view_async_shared(Args...) {}
    template<typename... Args>
    __device__ __host__ inline void synclog_emit_tma_store_arrive(Args...) {}
    template<typename... Args>
    __device__ __host__ inline void synclog_emit_tma_store_wait(Args...) {}
}
}

#include "cutlass/cutlass.h"
#include <ATen/cuda/CUDAContext.h>
#include "cute/tensor.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/detail/sm100_blockscaled_layout.hpp"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/util/packed_stride.hpp"
#include "cutlass/util/device_memory.h"
#include "cutlass/cuda_host_adapter.hpp"

using namespace cute;

// ---------------------------------------------------------------------------
// Optimized Column Extraction Kernel (vectorized loads/stores)
// ---------------------------------------------------------------------------

__global__ void extract_col0_vectorized_kernel(
    const __half* __restrict__ gemm_output,  // [m, n=32, l] RowMajor layout
    __half* __restrict__ gemv_output,        // [m, l]
    int m, int n, int l
) {
    // Thread handles one element at a time
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int total = m * l;

    if (idx < total) {
        // Decode output position: idx = batch * m + row
        int batch = idx / m;
        int row = idx % m;

        // CUTLASS LayoutD is RowMajor, so stride: (n, 1, m*n)
        // Index for (row, col=0, batch): row*n + 0 + batch*m*n
        int gemm_idx = row * n + batch * m * n;
        gemv_output[idx] = gemm_output[gemm_idx];
    }
}

// ---------------------------------------------------------------------------
// Global Configurations
// ---------------------------------------------------------------------------

using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutATag = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;

using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutBTag = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 32;

using ElementC = cutlass::half_t;
using ElementD = cutlass::half_t;
using LayoutCTag = cutlass::layout::RowMajor;
using LayoutDTag = cutlass::layout::RowMajor;  // RowMajor for easier column extraction
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;

using ElementAccumulator = float;
using ElementCompute = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using MmaTileShape = Shape<_128, _64, _256>;  // M=128 is minimum, try larger K for better arithmetic intensity

using FusionOperation = cutlass::epilogue::fusion::LinearCombination<ElementD, ElementCompute>;

// ---------------------------------------------------------------------------
// Templated Gemm Config
// ---------------------------------------------------------------------------

template <typename ClusterShape>
struct GemmConfig {
    using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        MmaTileShape, ClusterShape,
        cutlass::epilogue::collective::EpilogueTileAuto,
        ElementAccumulator, ElementAccumulator,
        ElementC, LayoutCTag, AlignmentC,
        ElementD, LayoutDTag, AlignmentD,
        cutlass::epilogue::collective::EpilogueScheduleAuto,
        FusionOperation
    >::CollectiveOp;

    using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        ElementA, LayoutATag, AlignmentA,
        ElementB, LayoutBTag, AlignmentB,
        ElementAccumulator,
        MmaTileShape, ClusterShape,
        cutlass::gemm::collective::StageCountAutoCarveout<
            static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
        cutlass::gemm::collective::KernelScheduleAuto
    >::CollectiveOp;

    using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
        Shape<int, int, int, int>,
        CollectiveMainloop,
        CollectiveEpilogue,
        void>;

    using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
};

// ---------------------------------------------------------------------------
// Helper to run GEMM
// ---------------------------------------------------------------------------

static std::map<int, torch::Tensor> workspace_cache;

template <typename Config>
void run_gemm(
    torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB,
    torch::Tensor C_out,
    int m, int n_problem, int k, int l, int n_phys_padded
) {
    using Gemm = typename Config::Gemm;
    using StrideA = typename Gemm::GemmKernel::StrideA;
    using StrideB = typename Gemm::GemmKernel::StrideB;
    using StrideC = typename Gemm::GemmKernel::StrideC;
    using StrideD = typename Gemm::GemmKernel::StrideD;
    using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
    using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
    using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;

    StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, l});
    StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n_phys_padded, k, l});
    StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {m, n_problem, l});

    StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n_problem, l});

    LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n_phys_padded, k, l));
    LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n_phys_padded, k, l));

    typename Gemm::Arguments arguments{
        cutlass::gemm::GemmUniversalMode::kGemm,
        {m, n_problem, k, l},
        {
            reinterpret_cast<ElementA::DataType*>(A.data_ptr()), stride_A,
            reinterpret_cast<ElementB::DataType*>(B.data_ptr()), stride_B,
            reinterpret_cast<ElementA::ScaleFactorType*>(SFA.data_ptr()), layout_SFA,
            reinterpret_cast<ElementB::ScaleFactorType*>(SFB.data_ptr()), layout_SFB
        },
        {
            {ElementCompute(1.0f), ElementCompute(0.0f)},
            nullptr, stride_C,
            reinterpret_cast<ElementD*>(C_out.data_ptr()), stride_D
        }
    };

    Gemm gemm;
    size_t workspace_size = Gemm::get_workspace_size(arguments);
    int device_id = A.device().index();
    void* workspace_ptr = nullptr;

    if (workspace_size > 0) {
        auto ws_it = workspace_cache.find(device_id);
        if (ws_it == workspace_cache.end() || ws_it->second.numel() * ws_it->second.element_size() < workspace_size) {
            torch::Tensor workspace_tensor = torch::empty({(int64_t)workspace_size},
                torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));
            workspace_cache[device_id] = workspace_tensor;
            workspace_ptr = workspace_tensor.data_ptr();
        } else {
            workspace_ptr = ws_it->second.data_ptr();
        }
    }

    cutlass::Status status = gemm.can_implement(arguments);
    if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel cannot implement");
    status = gemm.initialize(arguments, workspace_ptr);
    if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel init failed");
    status = gemm.run(at::cuda::getCurrentCUDAStream());
    if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel run failed");
}

// ---------------------------------------------------------------------------
// Temporary output cache
// ---------------------------------------------------------------------------

static std::map<std::tuple<int, int, int>, torch::Tensor> temp_output_cache;

void fp4_gemv_cutlass(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C_out,
    int m, int k, int l,
    int sfa_s0, int sfa_s1, int sfa_s2, int sfa_s3, int sfa_s4, int sfa_s5,
    int sfb_s0, int sfb_s1, int sfb_s2, int sfb_s3, int sfb_s4, int sfb_s5
) {
    constexpr int n_phys_padded = 128;
    constexpr int n_problem = 32;

    // Get or allocate cached temporary output buffer [m, n_problem, l]
    int device_id = A.device().index();
    auto temp_key = std::make_tuple(m, l, device_id);

    torch::Tensor temp_output_tensor;
    auto temp_it = temp_output_cache.find(temp_key);
    if (temp_it == temp_output_cache.end()) {
        temp_output_tensor = torch::empty({m, n_problem, l},
            torch::TensorOptions().dtype(torch::kFloat16).device(A.device()));
        temp_output_cache[temp_key] = temp_output_tensor;
    } else {
        temp_output_tensor = temp_it->second;
    }

    // Run GEMM to temporary buffer
    if (m >= 512) {
        run_gemm<GemmConfig<Shape<_4, _1, _1>>>(A, B, SFA, SFB, temp_output_tensor, m, n_problem, k, l, n_phys_padded);
    } else {
        run_gemm<GemmConfig<Shape<_1, _1, _1>>>(A, B, SFA, SFB, temp_output_tensor, m, n_problem, k, l, n_phys_padded);
    }

    // Extract column 0 using optimized kernel
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    int total_elements = m * l;
    int threads_per_block = 256;
    int num_blocks = (total_elements + threads_per_block - 1) / threads_per_block;

    extract_col0_vectorized_kernel<<<num_blocks, threads_per_block, 0, stream>>>(
        reinterpret_cast<__half*>(temp_output_tensor.data_ptr()),
        reinterpret_cast<__half*>(C_out.data_ptr()),
        m, n_problem, l
    );

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error("Extract kernel launch failed");
    }
}
"""

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

void fp4_gemv_cutlass(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C_out,
    int m, int k, int l,
    int sfa_s0, int sfa_s1, int sfa_s2, int sfa_s3, int sfa_s4, int sfa_s5,
    int sfb_s0, int sfb_s1, int sfb_s2, int sfb_s3, int sfb_s4, int sfb_s5
);
"""

_cuda_module = None

def get_cuda_module():
    """Compile and cache the CUDA module."""
    global _cuda_module
    if _cuda_module is None:
        import os

        # Get CUTLASS include path
        cutlass_include = os.environ.get('CUTLASS_PATH', '/tmp/cutlass')
        cutlass_include_dir = f"{cutlass_include}/include"

        _cuda_module = load_inline(
            name='fp4_gemv_cutlass_opt',
            cpp_sources=[cpp_source],
            cuda_sources=[cuda_source],
            functions=['fp4_gemv_cutlass'],
            extra_cuda_cflags=[
                '-O3',
                '--use_fast_math',
                '-gencode=arch=compute_100a,code=sm_100a',
                '-std=c++17',
                f'-I{cutlass_include_dir}',
                '-DCUTLASS_ARCH_MMA_SM100_SUPPORTED',
                '-DCUTLASS_ARCH_MMA_SM100_ENABLED',
                '-DCUTLASS_ARCH_MMA_SM100A_ENABLED',
                '-DCUTE_ARCH_TMA_SM90_ENABLED',
                '-DCUTE_ARCH_TCGEN05_TMEM_ENABLED',
                '-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1',
                '-DCUTLASS_DEBUG_TRACE_LEVEL=0',
                '-DCUTLASS_ENABLE_DEBUG_SYNCLOG=0',
                '--expt-extended-lambda',
                '--expt-relaxed-constexpr',
                '-Xcudafe', '--diag_suppress=177'
            ],
            verbose=True
        )
    return _cuda_module


def custom_kernel(data: input_t) -> output_t:
    a, b, _, _, sfa_permuted, sfb_permuted, c = data

    m, k_half, l = a.shape
    k = k_half * 2
    n_problem = 32  # Must match CUDA code

    cuda_module = get_cuda_module()

    # The scale factor strides are not actually used by CUTLASS (it computes its own layout)
    # but we pass them for interface compatibility
    sfa_strides = [0, 0, 0, 0, 0, 0]
    sfb_strides = [0, 0, 0, 0, 0, 0]

    # C_out will be filled directly by the extraction kernel in C++ layer
    # Ensure c has the right shape for output
    if c.dim() == 3 and c.size(1) == 1:
        # Reshape to (m, l) for kernel, then restore view
        c_reshaped = c.view(m, l)
        cuda_module.fp4_gemv_cutlass(
            a, b,
            sfa_permuted, sfb_permuted,
            c_reshaped,
            m, k, l,
            *sfa_strides,
            *sfb_strides
        )
    else:
        # Direct write to c (already (m, l) shape)
        cuda_module.fp4_gemv_cutlass(
            a, b,
            sfa_permuted, sfb_permuted,
            c,
            m, k, l,
            *sfa_strides,
            *sfb_strides
        )

    return c
scrolls · 368 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