Skip to content
KernelIndex
Search⌘K

submission 104755

mysfi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-104755?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.5µs
#156 of 678
2025-11-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:627a9d5c25144bcf8bbb27a3ebe17cd88101be8d4b4009b2fbe902cfb8b49bdd
license declaredunknown
license concludedunknown
authorsmysfi
imported2026-08-15

Techniques

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

async-copycutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);
fp4using ElementA = cutlass::float_e2m1_t;
fp8using ElementSFA = cutlass::float_e4m3_t;
fused-epiloguevoid apply_epilogue(
shared-memory__align__(16) ElementA smem_A[kBufferCount][kStageCount][kSmemPerStageA];

Kernel source

submission_inline.py647 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

CUTLASS_PATHS = [
    "/workspace/cutlass/include",
    "/workspace/cutlass/tools/include",
]

# CUDA source code with CUTLASS/CuTe header probes
cuda_src = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include "cutlass/cutlass.h"
#include "cutlass/layout/matrix.h"
#include "cutlass/matrix_coord.h"
#include "cutlass/numeric_types.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/arch/memory.h"
#include "cutlass/arch/memory_sm80.h"    
#include "cutlass/fast_math.h"
#include <cstdint> 
namespace nvfp4_gemv_sketch {
using ElementA  = cutlass::float_e2m1_t;    
using ElementB  = cutlass::float_e2m1_t;    
using ElementC  = cutlass::half_t;       
using ElementSFA = cutlass::float_e4m3_t;   
using ElementSFB = cutlass::float_e4m3_t;   
using ElementAccumulator = float;           
static constexpr int kElementsPerAccess = 32;
static constexpr int kSFVecSize        = 16;
static constexpr int kSFPerAccess      = 2;
static constexpr int kThreadCount   = 128;
static constexpr int kThreadsPerRow = 8;
static constexpr int kThreadsPerCol = kThreadCount / kThreadsPerRow; 
static constexpr int kStageCount  = 4;
static constexpr int kBufferCount = 2;
static constexpr int kPackedElementsA = 2;   
static constexpr int kPackedElementsB = 2;   
static constexpr int kPackedElements   = 2;  
static constexpr int kSmemPerStageA = kThreadCount   * (kElementsPerAccess / kPackedElementsA);
static constexpr int kSmemPerStageB = kThreadsPerRow * (kElementsPerAccess / kPackedElementsB);
using FragmentA       = cutlass::Array<ElementA,  kElementsPerAccess>;
using FragmentB       = cutlass::Array<ElementB,  kElementsPerAccess>;
using FragmentCompute = cutlass::Array<ElementAccumulator, kElementsPerAccess>;
using FragmentSFA     = cutlass::Array<ElementSFA, kSFPerAccess>;
using FragmentSFB     = cutlass::Array<ElementSFB, kSFPerAccess>;
using FragmentPackedA = cutlass::Array<ElementA,  kPackedElements>;
using FragmentPackedB = cutlass::Array<ElementB,  kPackedElements>;
static constexpr cutlass::FloatRoundStyle kRound = cutlass::FloatRoundStyle::round_to_nearest;
using A2AccConverter   = cutlass::NumericConverter<ElementAccumulator, ElementA,  kRound>;
using B2AccConverter   = cutlass::NumericConverter<ElementAccumulator, ElementB,  kRound>;
using SFA2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementSFA, kRound>;
using SFB2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementSFB, kRound>;
struct Params {
  int M = 0;
  int K = 0;
  int batch_count = 0;
  ElementA const* ptr_A        = nullptr;
  int64_t         stride_A      = 0;  
  int64_t         batch_stride_A = 0; 
  ElementB const* ptr_B         = nullptr;
  int64_t         batch_stride_B = 0; 
  ElementC* ptr_C         = nullptr; 
  int64_t         batch_stride_C = 0;      
  ElementSFA const* ptr_SFA        = nullptr;
  ElementSFB const* ptr_SFB        = nullptr;
  int64_t           batch_stride_SFA = 0;
  int64_t           batch_stride_SFB = 0;
  Params() = default;
};
struct SharedStorage {
  __align__(16) ElementA  smem_A[kBufferCount][kStageCount][kSmemPerStageA];
  __align__(16) ElementB  smem_B[kBufferCount][kStageCount][kSmemPerStageB];
  __align__(16) ElementSFA smem_SFA[kBufferCount][kStageCount][kThreadCount   * kSFPerAccess];
  __align__(16) ElementSFB smem_SFB[kBufferCount][kStageCount][kThreadsPerRow * kSFPerAccess];
};
__device__ __forceinline__
int compute_sf_offset_from_k(int global_k) {
  int SF_idx = global_k / kSFVecSize;
  int block   = SF_idx >> 2;      
  int offset  = SF_idx & 0x3;     
  int result  = (block << 9) + offset;  
  return result;
}
  CUTLASS_DEVICE
  ElementAccumulator blockscaled_multiply_add(
      FragmentA const& fragA,
      FragmentB const& fragB, 
      FragmentSFA const& fragSFA,
      FragmentSFB const& fragSFB) {
        uint16_t const& src_fragSFA_packed = reinterpret_cast<uint16_t const&>(fragSFA);
        uint16_t const& src_fragSFB_packed = reinterpret_cast<uint16_t const&>(fragSFB);
        uint32_t const* src_fragA_packed = reinterpret_cast<uint32_t const*>(&fragA);
        uint32_t const* src_fragB_packed = reinterpret_cast<uint32_t const*>(&fragB);
        __half out_h;
        uint16_t* out_fp16 = reinterpret_cast<uint16_t*>(&out_h);
        asm volatile( \
            "{\n" \
            ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n" \
            ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n" \
            ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n" \
            ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n" \
            ".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\n" \
            ".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\n" \
            ".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\n" \
            ".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\n" \
            ".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\n" \
            ".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\n" \
            ".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\n" \
            ".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\n" \
            ".reg .f16x2 sfa_f16x2;\n" \
            ".reg .f16x2 sfb_f16x2;\n" \
            ".reg .f16x2 sf_f16x2;\n" \
            ".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\n" \
            ".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\n" \
            ".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\n" \
            ".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\n" \
            ".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\n" \
            ".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\n" \
            ".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\n" \
            ".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\n" \
            ".reg .f16 result_f16, lane0, lane1;\n" \
            ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n" \
            "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %1;\n" \
            "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %2;\n" \
            "mov.b32 accum_0_0, 0;\n" \
            "mov.b32 accum_0_1, 0;\n" \
            "mov.b32 accum_0_2, 0;\n" \
            "mov.b32 accum_0_3, 0;\n" \
            "mov.b32 accum_1_0, 0;\n" \
            "mov.b32 accum_1_1, 0;\n" \
            "mov.b32 accum_1_2, 0;\n" \
            "mov.b32 accum_1_3, 0;\n" \
            "mov.b32 accum_2_0, 0;\n" \
            "mov.b32 accum_2_1, 0;\n" \
            "mov.b32 accum_2_2, 0;\n" \
            "mov.b32 accum_2_3, 0;\n" \
            "mov.b32 accum_3_0, 0;\n" \
            "mov.b32 accum_3_1, 0;\n" \
            "mov.b32 accum_3_2, 0;\n" \
            "mov.b32 accum_3_3, 0;\n" \
            "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n" \
            "mov.b32 {lane0, lane1}, sf_f16x2;\n" \
            "mov.b32 mul_f16x2_0, {lane0, lane0};\n" \
            "mov.b32 mul_f16x2_1, {lane1, lane1};\n" \
            "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %3;\n" \
            "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %4;\n" \
            "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n" \
            "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n" \
            "mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %7;\n" \
            "mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %8;\n" \
            "mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %9;\n" \
            "mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %10;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_0, byte0_0;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_1, byte0_1;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_2, byte0_2;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_3, byte0_3;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_4, byte0_4;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_5, byte0_5;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_6, byte0_6;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_0_7, byte0_7;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_0, byte1_0;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_1, byte1_1;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_2, byte1_2;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_3, byte1_3;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_4, byte1_4;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_5, byte1_5;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_6, byte1_6;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_1_7, byte1_7;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\n" \
            "cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\n" \
            "fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\n" \
            "fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\n" \
            "fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\n" \
            "fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\n" \
            "fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\n" \
            "fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\n" \
            "fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\n" \
            "fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\n" \
            "fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\n" \
            "fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\n" \
            "fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\n" \
            "fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\n" \
            "fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\n" \
            "fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\n" \
            "fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\n" \
            "fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\n" \
            "add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\n" \
            "add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\n" \
            "add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\n" \
            "add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\n" \
            "add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\n" \
            "add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\n" \
            "add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\n" \
            "add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\n" \
            "add.rn.f16x2 accum_0_0, accum_0_0, accum_0_2;\n" \
            "add.rn.f16x2 accum_1_0, accum_1_0, accum_1_2;\n" \
            "add.rn.f16x2 accum_2_0, accum_2_0, accum_2_2;\n" \
            "add.rn.f16x2 accum_3_0, accum_3_0, accum_3_2;\n" \
            "add.rn.f16x2 accum_0_0, accum_0_0, accum_1_0;\n" \
            "add.rn.f16x2 accum_2_0, accum_2_0, accum_3_0;\n" \
            "mul.rn.f16x2 accum_0_0, mul_f16x2_0, accum_0_0;\n" \
            "mul.rn.f16x2 accum_2_0, mul_f16x2_1, accum_2_0;\n" \
            "add.rn.f16x2 accum_0_0, accum_0_0, accum_2_0;\n" \
            "mov.b32 {lane0, lane1}, accum_0_0;\n" \
            "add.rn.f16 result_f16, lane0, lane1;\n" \
            "mov.b16 %0, result_f16;\n" \
            "}\n"
            : "=h"(out_fp16[0])                                     
            : "h"(src_fragSFA_packed), "h"(src_fragSFB_packed),     
              "r"(src_fragA_packed[0]), "r"(src_fragB_packed[0]),   
              "r"(src_fragA_packed[1]), "r"(src_fragB_packed[1]),   
              "r"(src_fragA_packed[2]), "r"(src_fragB_packed[2]),   
              "r"(src_fragA_packed[3]), "r"(src_fragB_packed[3])    
            : "memory"
        );
        return __half2float(out_h);
  }
__device__
ElementAccumulator process_tail_elements(
    int unrolled_k,
    int idx_col_k,
    int K,
    ElementA const* ptr_A,
    ElementB const* ptr_B,
    ElementSFA const* ptr_SFA,
    ElementSFB const* ptr_SFB)
{
  A2AccConverter   A_converter;
  B2AccConverter   B_converter;
  SFA2AccConverter SFA_converter;
  SFB2AccConverter SFB_converter;
  ElementAccumulator accum = ElementAccumulator(0);
  for (int k = unrolled_k + idx_col_k * kPackedElementsA;
       k < K;
       k += kThreadsPerRow * kPackedElementsA) {
    int SF_offset_by_k = compute_sf_offset_from_k(k);
    ElementSFA sfa = *(ptr_SFA + SF_offset_by_k);
    ElementSFB sfb = *(ptr_SFB + SF_offset_by_k);
    FragmentPackedA fragA;
    FragmentPackedB fragB;
    cutlass::arch::global_load<FragmentPackedA, sizeof(FragmentPackedA),
                               cutlass::arch::CacheOperation::Always>(
        fragA,
        ptr_A - (idx_col_k * kElementsPerAccess - k) / kPackedElementsA,
        true);
    cutlass::arch::global_load<FragmentPackedB, sizeof(FragmentPackedB),
                               cutlass::arch::CacheOperation::Always>(
        fragB,
        ptr_B - (idx_col_k * kElementsPerAccess - k) / kPackedElementsB,
        true);
    ElementAccumulator accum_SF_packed = ElementAccumulator(0);
    for (int e = 0; e < kPackedElements; ++e) {
      accum_SF_packed += A_converter(fragA[e]) * B_converter(fragB[e]);
    }
    accum_SF_packed *= SFA_converter(sfa) * SFB_converter(sfb);
    accum += accum_SF_packed;
  }
  return accum;
}
__device__
void load_stages_gmem_to_smem(
    int buffer_idx,         
    int num_stages,         
    int& unrolled_k,        
    int& global_k,          
    int tileA_k_local,      
    int smem_offset_A,      
    int smem_offset_B,      
    int smem_sf_write_offset, 
    bool is_even_thread,    
    bool load_b,            
    int K_limit,        
    ElementA const* ptr_A,  
    ElementB const* ptr_B,  
    ElementSFA const* ptr_SFA, 
    ElementSFB const* ptr_SFB, 
    SharedStorage& shared) {
  for (int s = 0; s < num_stages; ++s) {
    if (unrolled_k >= K_limit) {
      return;
    }
    int sf_offset = compute_sf_offset_from_k(global_k);
    bool stage_valid = (unrolled_k < K_limit);
    if (is_even_thread) {
      void* smem_ptr = &shared.smem_SFA[buffer_idx][s][smem_sf_write_offset];
      ElementSFA const* gmem_ptr = ptr_SFA + sf_offset;
      cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);
    }
    if (load_b && is_even_thread) {
      int sf_sfb_offset = (threadIdx.x / 2) * 4; 
      void* smem_ptr = &shared.smem_SFB[buffer_idx][s][sf_sfb_offset];
      ElementSFB const* gmem_ptr = ptr_SFB + sf_offset;
      cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);
    }
    {
      void* smem_ptr = &shared.smem_A[buffer_idx][s][smem_offset_A];
      int packed_offset = unrolled_k / kPackedElementsA;
      ElementA const* gmem_ptr = ptr_A + packed_offset;
      cutlass::arch::cp_async<sizeof(FragmentA)>(smem_ptr, gmem_ptr, stage_valid);
    }
    if (load_b) {
      void* smem_ptr = &shared.smem_B[buffer_idx][s][smem_offset_B];
      int packed_offset = unrolled_k / kPackedElementsB;
      ElementB const* gmem_ptr = ptr_B + packed_offset;
      cutlass::arch::cp_async<sizeof(FragmentB)>(smem_ptr, gmem_ptr, stage_valid);
    }
    unrolled_k += tileA_k_local;
    global_k   += tileA_k_local;
  }
}
__device__
void load_smem_fragments(
    FragmentA& fragA,
    FragmentB& fragB,
    FragmentSFA& fragSFA,
    FragmentSFB& fragSFB,
    int smem_pipe_idx,
    int k_block,
    int smem_offset_A,
    int smem_offset_B,
    int smem_sf_offset,
    SharedStorage& shared) {
  cutlass::arch::shared_load(fragA, &shared.smem_A[smem_pipe_idx][k_block][smem_offset_A]);
  cutlass::arch::shared_load(fragB, &shared.smem_B[smem_pipe_idx][k_block][smem_offset_B]);
  {
    uint32_t smem_ptr = cutlass::arch::cutlass_get_smem_pointer(
        &shared.smem_SFA[smem_pipe_idx][k_block][smem_sf_offset]);
    cutlass::arch::shared_load<2>(&fragSFA, smem_ptr);
  }
  {
    int sfb_offset = threadIdx.x * kSFPerAccess;
    uint32_t smem_ptr = cutlass::arch::cutlass_get_smem_pointer(
        &shared.smem_SFB[smem_pipe_idx][k_block][sfb_offset]);
    cutlass::arch::shared_load<2>(&fragSFB, smem_ptr);
  }
}
__device__
ElementAccumulator warp_reduce_over_k(ElementAccumulator accum) {
  for (int mask = (kThreadsPerRow >> 1); mask > 0; mask >>= 1) {
    float other = __shfl_xor_sync(0xFFFFFFFFu, accum, mask, 32);
    accum += other;
  }
  return accum;
}
__device__
void apply_epilogue(
    ElementAccumulator accum,
    ElementC* ptr_c_row)  
{
  *ptr_c_row = static_cast<ElementC>(__float2half(accum)); 
}
__global__
void nvfp4_gemv_kernel(Params params) {
  extern __shared__ uint8_t smem_raw[];
  SharedStorage& shared = *reinterpret_cast<SharedStorage*>(smem_raw);
  int M = params.M;
  int K = params.K;
  int batch_count = params.batch_count;
  for (int batch_idx = blockIdx.z; batch_idx < batch_count; batch_idx += gridDim.z) {
    int idx_col_k = threadIdx.x;                             
    int idx_row_m = blockIdx.x * blockDim.y + threadIdx.y;   
    if (idx_row_m >= M) {
      continue;
    }
    ElementA const* ptr_A = params.ptr_A
      + batch_idx * params.batch_stride_A
      + idx_row_m * params.stride_A;
    ptr_A += idx_col_k * (kElementsPerAccess / kPackedElementsA);
    ElementB const* ptr_B = params.ptr_B
      + batch_idx * params.batch_stride_B;
    ptr_B += idx_col_k * (kElementsPerAccess / kPackedElementsB);
    ElementC* ptr_C_row = params.ptr_C
      + batch_idx * params.batch_stride_C
      + idx_row_m;
    ElementSFA const* ptr_SFA = params.ptr_SFA + batch_idx * params.batch_stride_SFA;
    ElementSFB const* ptr_SFB = params.ptr_SFB + batch_idx * params.batch_stride_SFB;
    [[maybe_unused]] int SF_blocks_by_M = (M + 127) / 128;               
    int SF_blocks_by_K = (K / kSFVecSize + 3) / 4;                       
    int row_block   = idx_row_m >> 7;       
    int row_in_bloc = idx_row_m & 0x7f;     
    int row_low32   = row_in_bloc & 0x1f;   
    int row_hi      = row_in_bloc >> 5;     
    ptr_SFA += (row_block * SF_blocks_by_K << 9)  
            + (row_low32 << 4)                   
            + (row_hi << 2);                     
    ElementAccumulator accum = ElementAccumulator(0);
    const int tileA_k_local   = kThreadsPerRow * kElementsPerAccess; 
    const int tiles_per_group = kStageCount;                         
    const int macro_k         = tileA_k_local * tiles_per_group;     
    const int full_groups = K / macro_k;                 
    const int total_tiles = full_groups * tiles_per_group; 
    const int K_main      = total_tiles * tileA_k_local; 
    int unrolled_k = 0;                                  
    int global_k   = idx_col_k * kElementsPerAccess;     
    const int thread_id = threadIdx.y * kThreadsPerRow + threadIdx.x;
    const bool is_even_thread = ((threadIdx.x % 2) == 0);
    const bool load_b = (threadIdx.y == 0);
    const int smem_sf_write_offset = (thread_id / 2) * 4;  
    const int smem_sf_offset       = thread_id * kSFPerAccess;
    if (total_tiles == 0) {
      accum += process_tail_elements(
          0,
          idx_col_k,
          K,
          ptr_A,
          ptr_B,
          ptr_SFA,
          ptr_SFB);
    } else {
      int smem_offset_A = thread_id    * (kElementsPerAccess / kPackedElementsA);
      int smem_offset_B = threadIdx.x  * (kElementsPerAccess / kPackedElementsB);
      for (int buf = 0; buf < kBufferCount - 1; ++buf) {
        load_stages_gmem_to_smem(
            buf,
            kStageCount,
            unrolled_k,
            global_k,
            tileA_k_local,
            smem_offset_A,
            smem_offset_B,
            smem_sf_write_offset,
            is_even_thread,
            load_b,
            K_main,
            ptr_A,
            ptr_B,
            ptr_SFA,
            ptr_SFB,
            shared);
      }
      cutlass::arch::cp_async_fence();
      cutlass::arch::cp_async_wait<kBufferCount - 2>();
      __syncthreads();
      FragmentA   fragA_reg[2];
      FragmentB   fragB_reg[2];
      FragmentSFA fragSFA_reg[2];
      FragmentSFB fragSFB_reg[2];
      int smem_pipe_read  = 0;
      int smem_pipe_write = kBufferCount - 1;
      if (kStageCount > 1) {
        int frag_idx = 0;
        load_smem_fragments(
            fragA_reg[frag_idx],
            fragB_reg[frag_idx],
            fragSFA_reg[frag_idx],
            fragSFB_reg[frag_idx],
            smem_pipe_read,
            0,
            smem_offset_A,
            smem_offset_B,
            smem_sf_offset,
            shared);
      }
      int tile_idx = 0;
      while (tile_idx < total_tiles) {
        int smem_pipe_read_curr = smem_pipe_read;
        for (int k_block = 0; k_block < kStageCount; ++k_block) {
          if (k_block == kStageCount - 1) {
            cutlass::arch::cp_async_wait<kBufferCount - 2>();
            __syncthreads();
            smem_pipe_read_curr = smem_pipe_read;
          }
          int k_block_next = (k_block + 1) % kStageCount;
          int frag_idx_next = (k_block + 1) & 1;
          load_smem_fragments(
              fragA_reg[frag_idx_next],
              fragB_reg[frag_idx_next],
              fragSFA_reg[frag_idx_next],
              fragSFB_reg[frag_idx_next],
              smem_pipe_read_curr,
              k_block_next,
              smem_offset_A,
              smem_offset_B,
              smem_sf_offset,
              shared);
          if (k_block == 0) {
            load_stages_gmem_to_smem(
                smem_pipe_write,
                kStageCount,
                unrolled_k,
                global_k,
                tileA_k_local,
                smem_offset_A,
                smem_offset_B,
                smem_sf_write_offset,
                is_even_thread,
                load_b,
                K_main,
                ptr_A,
                ptr_B,
                ptr_SFA,
                ptr_SFB,
                shared);
            cutlass::arch::cp_async_fence();
            smem_pipe_write = smem_pipe_read;
            ++smem_pipe_read;
            if (smem_pipe_read == kBufferCount) {
              smem_pipe_read = 0;
            }
          }
          int frag_idx = k_block & 1;
          accum += blockscaled_multiply_add(
              fragA_reg[frag_idx],
              fragB_reg[frag_idx],
              fragSFA_reg[frag_idx],
              fragSFB_reg[frag_idx]);
        }
        tile_idx += kStageCount;
      }
      cutlass::arch::cp_async_wait<0>();
      __syncthreads();
      if (unrolled_k < K) {
        accum += process_tail_elements(
            unrolled_k,
            idx_col_k,
            K,
            ptr_A,
            ptr_B,
            ptr_SFA,
            ptr_SFB);
      }
    }
    accum = warp_reduce_over_k(accum);
    if (threadIdx.x == 0) {
      apply_epilogue(accum, ptr_C_row);
    }
  }
}
torch::Tensor nvfp4_gemv(
    torch::Tensor a,        
    torch::Tensor b,        
    torch::Tensor sfa_perm, 
    torch::Tensor sfb_perm, 
    torch::Tensor c)        
{
  TORCH_CHECK(a.is_cuda() && b.is_cuda() && c.is_cuda(),
              "All tensors must be CUDA tensors");
  int64_t M        = a.size(0);
  int64_t K_packed = a.size(1);       
  int64_t L        = a.size(2);       
  int64_t K        = K_packed * 2;    
  Params params;
  params.M           = static_cast<int>(M);
  params.K           = static_cast<int>(K);
  params.batch_count = static_cast<int>(L);
  auto a_strides = a.strides();  
  params.ptr_A          = reinterpret_cast<ElementA const*>(a.data_ptr());
  params.stride_A       = a_strides[0];  
  params.batch_stride_A = a_strides[2];  
  auto b_strides = b.strides();
  params.ptr_B          = reinterpret_cast<ElementB const*>(b.data_ptr());
  params.batch_stride_B = b_strides[2];  
  auto c_strides = c.strides();
  params.ptr_C          = reinterpret_cast<ElementC*>(c.data_ptr());
  params.batch_stride_C = c_strides[2];  
  params.ptr_SFA          = reinterpret_cast<ElementSFA const*>(sfa_perm.data_ptr());
  params.ptr_SFB          = reinterpret_cast<ElementSFB const*>(sfb_perm.data_ptr());
  params.batch_stride_SFA = sfa_perm.numel() / L;  
  params.batch_stride_SFB = sfb_perm.numel() / L;
  dim3 block(kThreadsPerRow, kThreadsPerCol, 1); 
  dim3 grid(
      (M + block.y - 1) / block.y,  
      1,
      params.batch_count);          
  size_t smem_size = sizeof(SharedStorage);
  cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
  nvfp4_gemv_kernel<<<grid, block, smem_size, stream>>>(params);
  return c;
}
} 
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def(
      "nvfp4_gemv",
      &nvfp4_gemv_sketch::nvfp4_gemv,
      "NVFP4 block-scaled GEMV (procedural sketch)");
}
"""

cpp_src = r"""
#include <torch/extension.h>

torch::Tensor nvfp4_gemv(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa_perm,
    torch::Tensor sfb_perm,
    torch::Tensor c
);
"""

# Global cache for compiled extension
_nvfp4_ext = None


def get_extension():
    """Compile and cache the CUDA extension."""
    global _nvfp4_ext
    if _nvfp4_ext is not None:
        return _nvfp4_ext

    _nvfp4_ext = load_inline(
        name="nvfp4_gemv_cutlass_probe",
        cpp_sources=[cpp_src],
        cuda_sources=[cuda_src],
        extra_cflags=["-O3", "-std=c++17"],
        extra_cuda_cflags=[
            "-O3",
            "-std=c++17",
            "--expt-relaxed-constexpr",
            "-use_fast_math",
            "-gencode=arch=compute_100a,code=sm_100a"
        ],
        extra_include_paths=CUTLASS_PATHS,
        verbose=True,
    )
    return _nvfp4_ext


def custom_kernel(data: input_t) -> output_t:
    """
    Skeleton implementation that zeros the output.
    This validates the PyTorch <-> CUDA extension pipeline with CUTLASS headers.
    """
    a, b, _, _, sfa_permuted, sfb_permuted, c = data

    # Get compiled extension and call the kernel
    ext = get_extension()
    return ext.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)
scrolls · 647 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