Skip to content
KernelIndex
Search⌘K

submission 80995

Alex A · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

slow.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-80995?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
34.8µs
#188 of 678
2025-11-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f22607401b4b1988e4a09077eea9835aa3a6ba968647b44639088c6774a536fc
license declaredunknown
license concludedunknown
authorsAlex A
imported2026-08-15

Techniques

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

clusterusing ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
fp4using nvf4_t = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing EpilogueTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}));
vector-width = uint4uint4 packed[kTilesPerGroup];
warp-specializationcutlass::epilogue::TmaWarpSpecialized2Sm,

Kernel source

slow.py743 lines
#!POPCORN leaderboard nvfp4_gemv
import os
import torch
from torch.utils.cpp_extension import load_inline

###############################################################################
# 1. C++ binding (pybind11) – exposes the Python entry point
###############################################################################

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

void blockscaled_gemm_nvf4_f16(
    at::Tensor& A,
    at::Tensor& B,
    at::Tensor& C,
    at::Tensor& A_scales,
    at::Tensor& B_scales,
    int batch_size);

void to_blocked_and_gemm(
    at::Tensor& A_fp4, at::Tensor& B_fp4,
    at::Tensor& C,
    at::Tensor& SFA, at::Tensor& SFB,
    int batch_size);

void to_blocked_batched_cuda_wrapper(at::Tensor input, at::Tensor output);
void to_blocked_batched_dual_cuda_wrapper(
    at::Tensor input0, at::Tensor output0,
    at::Tensor input1, at::Tensor output1);

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("blockscaled_gemm_nvf4_f16",
        &blockscaled_gemm_nvf4_f16,
        "NVF4 blockscaled GEMV with FP16 output");
  m.def("to_blocked_batched_cuda",
        &to_blocked_batched_cuda_wrapper,
        "Batched to_blocked conversion (CUDA)");
  m.def("to_blocked_batched_dual_cuda",
        &to_blocked_batched_dual_cuda_wrapper,
        "Dual batched to_blocked conversion (CUDA)");
  m.def("to_blocked_and_gemm",
        &to_blocked_and_gemm,
        "Run scale blocking followed by blockscaled GEMM");
}
"""

###############################################################################
# 2. CUDA / CUTLASS kernel implementation (inline translation unit)
###############################################################################

cuda_src = r"""
#include <cuda.h>
#include <cuda_runtime.h>

#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDACachingAllocator.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/extension.h>

#include <iostream>
#include <type_traits>
#include <cstdlib>
#include <cstdint>
#include <limits>

#include <cutlass/cutlass.h>
#include <cutlass/numeric_types.h>
#include <cutlass/epilogue/fusion/operations.hpp>
#include <cutlass/epilogue/thread/linear_combination.h>
#include <cutlass/epilogue/collective/collective_builder.hpp>
#include <cutlass/gemm/collective/collective_builder.hpp>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#include <cutlass/gemm/kernel/tile_scheduler_params.h>
#include <cutlass/detail/sm100_blockscaled_layout.hpp>
#include <cutlass/util/packed_stride.hpp>

#include <cute/tensor.hpp>

// ============================================================================
// CUDA kernel for fast batched to_blocked (supports two tensors per launch)
// ============================================================================
struct TensorParams {
  const uint8_t* input;
  uint8_t* output;
  int rows;
  int cols;
  int batches;
  int n_row_blocks;
  int n_col_blocks;
  int n_col_groups;
  int64_t in_row_stride;
  int64_t in_col_stride;
  int64_t in_batch_stride;
  int64_t out_batch_stride;
};

constexpr int kWarpSize       = 32;
constexpr int kTilesPerCTA    = 4;
constexpr int kTilesPerGroup  = 4;

__global__ void to_blocked_batched_kernel(
    const TensorParams  params0,
    const TensorParams  params1,
    size_t total_tiles0,
    size_t total_tiles1) {

  const int warp_idx = threadIdx.y;
  const int lane     = threadIdx.x;

  const size_t tile_linear =
      static_cast<size_t>(blockIdx.x) * blockDim.y + warp_idx;
  const size_t total_tiles = total_tiles0 + total_tiles1;

  if (tile_linear >= total_tiles) {
    return;
  }

  const bool use_first = (tile_linear < total_tiles0);
  const TensorParams params =
      use_first ? params0 : params1;
  const size_t local_tile =
      use_first ? tile_linear : (tile_linear - total_tiles0);

  const int rows          = params.rows;
  const int batches       = params.batches;
  const int n_row_blocks  = params.n_row_blocks;
  const int n_col_blocks  = params.n_col_blocks;
  const int n_col_groups  = params.n_col_groups;

  const int64_t in_batch_stride = params.in_batch_stride;
  const int64_t in_row_stride   = params.in_row_stride;
  const int64_t out_batch_stride= params.out_batch_stride;

  const uint8_t*  input  = params.input;
  uint8_t*        output = params.output;

  const size_t tile_groups_per_batch =
      static_cast<size_t>(n_row_blocks) * n_col_groups;
  if (tile_groups_per_batch == 0) {
    return;
  }

  const int batch = static_cast<int>(local_tile / tile_groups_per_batch);
  if (batch >= batches) {
    return;
  }

  const size_t tile_in_batch = local_tile % tile_groups_per_batch;

  const int rb       = static_cast<int>(tile_in_batch / n_col_groups);
  const int cb_group = static_cast<int>(tile_in_batch % n_col_groups);
  const int cb_base  = cb_group * kTilesPerGroup;

  int tiles_in_group = n_col_blocks - cb_base;
  if (tiles_in_group <= 0) {
    return;
  }
  if (tiles_in_group > kTilesPerGroup) {
    tiles_in_group = kTilesPerGroup;
  }

  const int tile_row_base = rb * 128;
  if (tile_row_base >= rows) {
    return;
  }

  int row_limit = rows - tile_row_base;
  if (row_limit > 128) {
    row_limit = 128;
  }

  const size_t tile_idx_base =
      static_cast<size_t>(rb) * n_col_blocks + cb_base;
  const size_t tile_base_base =
      static_cast<size_t>(batch) * out_batch_stride + tile_idx_base * 512;

  // Precompute the base pointer to the first row of this row-block & batch.
  const int64_t col_offset_bytes = static_cast<int64_t>(cb_base) * 4;
  const uint8_t*  base_row_ptr =
      input +
      static_cast<int64_t>(batch) * in_batch_stride +
      static_cast<int64_t>(tile_row_base) * in_row_stride +
      col_offset_bytes;

  uint4 packed[kTilesPerGroup];
#pragma unroll
  for (int t = 0; t < kTilesPerGroup; ++t) {
    packed[t] = make_uint4(0u, 0u, 0u, 0u);
  }

  // Each warp covers 128 rows (4 groups * 32 lanes).
#pragma unroll
  for (int group = 0; group < 4; ++group) {
    const int logical_row = lane + group * kWarpSize;
    const bool valid      = (logical_row < row_limit);

    const uint8_t* row_ptr =
        base_row_ptr + static_cast<int64_t>(logical_row) * in_row_stride;

    // Masked 128b load via read-only path.
    uint4 chunk = valid
        ? __ldg(reinterpret_cast<const uint4*>(row_ptr))
        : make_uint4(0u, 0u, 0u, 0u);

    const uint32_t w0 = chunk.x;
    const uint32_t w1 = chunk.y;
    const uint32_t w2 = chunk.z;
    const uint32_t w3 = chunk.w;

    if (tiles_in_group > 0) reinterpret_cast<uint32_t*>(&packed[0])[group] = w0;
    if (tiles_in_group > 1) reinterpret_cast<uint32_t*>(&packed[1])[group] = w1;
    if (tiles_in_group > 2) reinterpret_cast<uint32_t*>(&packed[2])[group] = w2;
    if (tiles_in_group > 3) reinterpret_cast<uint32_t*>(&packed[3])[group] = w3;
  }

#pragma unroll
  for (int t = 0; t < kTilesPerGroup; ++t) {
    if (t >= tiles_in_group) {
      break;
    }

    const size_t tile_base =
        tile_base_base + static_cast<size_t>(t) * 512;
    uint8_t* tile_out_base =
        output + tile_base + static_cast<size_t>(lane) * 16;

    // Streaming 128-bit store: no L1 allocation, minimal L2 pollution.
    uint4 v = packed[t];
    asm volatile(
        "st.global.cs.v4.b32 [%0], {%1,%2,%3,%4};\n"
        :
        : "l"(tile_out_base),
          "r"(v.x), "r"(v.y), "r"(v.z), "r"(v.w)
    );
  }
}

namespace {

TensorParams make_params(const at::Tensor& input, const at::Tensor& output) {
  TensorParams params;
  params.input = input.data_ptr<uint8_t>();
  params.output = output.data_ptr<uint8_t>();
  params.rows = static_cast<int>(input.size(0));
  params.cols = static_cast<int>(input.size(1));
  params.batches = static_cast<int>(input.size(2));
  params.n_row_blocks = (params.rows + 127) / 128;
  params.n_col_blocks = (params.cols + 3) / 4;
  params.n_col_groups = (params.n_col_blocks + kTilesPerGroup - 1) / kTilesPerGroup;
  params.in_row_stride = input.stride(0);
  params.in_col_stride = input.stride(1);
  params.in_batch_stride = input.stride(2);
  params.out_batch_stride = static_cast<int64_t>(params.rows) * params.cols;
  return params;
}

size_t count_tile_groups(const TensorParams& params) {
  return static_cast<size_t>(params.batches) *
         params.n_row_blocks *
         params.n_col_groups;
}

void launch_to_blocked(const TensorParams& params0, size_t groups0,
                       const TensorParams& params1, size_t groups1) {
  size_t total_groups = groups0 + groups1;
  if (total_groups == 0) {
    return;
  }

  TORCH_CHECK(total_groups <= static_cast<size_t>(std::numeric_limits<int>::max()),
              "Total tile group count exceeds supported range");

  auto validate = [](const TensorParams& params, size_t groups) {
    if (groups == 0) {
      return;
    }
    TORCH_CHECK(params.in_col_stride == 1,
                "Expected unit column stride for blocked conversion");
    TORCH_CHECK((reinterpret_cast<uintptr_t>(params.input) & 0xF) == 0,
                "Input tensor must be 16-byte aligned");
    TORCH_CHECK((reinterpret_cast<uintptr_t>(params.output) & 0xF) == 0,
                "Output tensor must be 16-byte aligned");
  };

  validate(params0, groups0);
  validate(params1, groups1);

  auto stream = at::cuda::getCurrentCUDAStream();
  unsigned int grid_blocks = static_cast<unsigned int>((total_groups + kTilesPerCTA - 1) / kTilesPerCTA);
  dim3 grid(grid_blocks);
  dim3 block(kWarpSize, kTilesPerCTA);

  to_blocked_batched_kernel<<<grid, block, 0, stream>>>(
      params0, params1, groups0, groups1);
}

}  // namespace

void to_blocked_batched_cuda_wrapper(at::Tensor input, at::Tensor output) {
  int rows = input.size(0);
  int cols = input.size(1);

  TORCH_CHECK(rows % 128 == 0, "to_blocked_batched expects rows to be a multiple of 128");
  TORCH_CHECK(cols % 4 == 0, "to_blocked_batched expects cols to be a multiple of 4");

  TensorParams params0 = make_params(input, output);
  size_t tiles0 = count_tile_groups(params0);

  TensorParams empty{};
  launch_to_blocked(params0, tiles0, empty, 0);
}

void to_blocked_batched_dual_cuda_wrapper(
    at::Tensor input0, at::Tensor output0,
    at::Tensor input1, at::Tensor output1) {
  int rows0 = input0.size(0);
  int cols0 = input0.size(1);
  int rows1 = input1.size(0);
  int cols1 = input1.size(1);

  TORCH_CHECK(rows0 % 128 == 0, "A_scales rows must be multiple of 128");
  TORCH_CHECK(cols0 % 4 == 0, "A_scales cols must be multiple of 4");
  TORCH_CHECK(rows1 % 128 == 0, "B_scales rows must be multiple of 128");
  TORCH_CHECK(cols1 % 4 == 0, "B_scales cols must be multiple of 4");

  TensorParams params0 = make_params(input0, output0);
  TensorParams params1 = make_params(input1, output1);

  size_t tiles0 = count_tile_groups(params0);
  size_t tiles1 = count_tile_groups(params1);

  launch_to_blocked(params0, tiles0, params1, tiles1);
}

#define CUTLASS_CHECK(status)                                                                    \
  {                                                                                              \
    cutlass::Status error = status;                                                              \
    if (error != cutlass::Status::kSuccess) {                                                    \
      std::cerr << "Got CUTLASS error: " << cutlassGetStatusString(error) << " at line "         \
                << __LINE__ << std::endl;                                                        \
      exit(EXIT_FAILURE);                                                                        \
    }                                                                                            \
  }

struct GemmParams {
  int m, n, k, l;
  void *ptr_A, *ptr_B, *ptr_C;
  void *ptr_SFA, *ptr_SFB;
  float alpha, beta;
};

using nvf4_t = cutlass::nv_float4_t<cutlass::float_e2m1_t>;

template <typename ElementA, typename ElementB, typename ElementC,
          typename ElementAccum, int bM_, int bN_, int bK_>
struct CollectiveTemplates {
  using LayoutATag = cutlass::layout::RowMajor;
  using LayoutBTag = cutlass::layout::ColumnMajor;
  using LayoutCTag = cutlass::layout::RowMajor;

  static constexpr int AlignmentA = cute::is_same_v<ElementA, nvf4_t> ? 128 : 16;
  static constexpr int AlignmentB = cute::is_same_v<ElementB, nvf4_t> ? 128 : 16;
  static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;

  using ElementAccumulator = ElementAccum;
  using ArchTag = cutlass::arch::Sm100;
  using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;

  static constexpr int bM = bM_;
  static constexpr int bN = bN_;
  static constexpr int bK = bK_;
  using MmaTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}, cute::Int<bK>{}));
  using EpilogueTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}));
  using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
  using ProblemShape = cute::Shape<int, int, int, int>;

  using DefaultEpilogueOp = cutlass::epilogue::fusion::LinearCombination<ElementC, ElementAccum>;

  using DefaultCollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
      ArchTag,
      OperatorClass,
      MmaTileShape,
      ClusterShape,
      cutlass::epilogue::collective::EpilogueTileAuto,
      ElementAccumulator,
      ElementAccumulator,
      void,
      LayoutCTag,
      AlignmentC,
      ElementC,
      LayoutCTag,
      AlignmentC,
      cutlass::epilogue::TmaWarpSpecialized2Sm,
      DefaultEpilogueOp
  >::CollectiveOp;

  template <typename CollectiveEpilogue_>
  using DefaultCollectiveMainloop = 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::KernelTmaWarpSpecialized2SmBlockScaledSm100
  >::CollectiveOp;

  template <typename CollectiveMainloop_, typename CollectiveEpilogue_>
  using DefaultGemmKernel = cutlass::gemm::kernel::GemmUniversal<
      ProblemShape,
      CollectiveMainloop_,
      CollectiveEpilogue_,
      void>;
};

template <int bM_, int bN_, int bK_,
          typename ElementA, typename ElementB,
          typename ElementC, typename ElementAccum>
void run_blockscaled_gemm(GemmParams& params, cudaStream_t stream) {
  using Collectives = CollectiveTemplates<ElementA, ElementB, ElementC, ElementAccum, bM_, bN_, bK_>;
  using CollectiveEpilogue = typename Collectives::DefaultCollectiveEpilogue;
  using CollectiveMainloop = typename Collectives::template DefaultCollectiveMainloop<CollectiveEpilogue>;
  using GemmKernel = typename Collectives::template DefaultGemmKernel<CollectiveMainloop, CollectiveEpilogue>;
  using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
  using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;

  using StrideA = typename Gemm::GemmKernel::StrideA;
  using StrideB = typename Gemm::GemmKernel::StrideB;
  using StrideC = typename Gemm::GemmKernel::StrideC;
  using StrideD = typename Gemm::GemmKernel::StrideD;

  StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {params.m, params.k, params.l});
  StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {params.n, params.k, params.l});
  StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {params.m, params.n, params.l});
  StrideD stride_D = stride_C;

  auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(
      cute::make_shape(params.m, params.n, params.k, params.l));
  auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(
      cute::make_shape(params.m, params.n, params.k, params.l));

  cutlass::KernelHardwareInfo hw_info;
  hw_info.device_id = 0;
  hw_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);
  hw_info.cluster_shape = dim3(2, 1, 1);
  hw_info.cluster_shape_fallback = dim3(2, 1, 1);

  typename Gemm::GemmKernel::TileSchedulerArguments scheduler;

  typename Gemm::Arguments arguments;
  decltype(arguments.epilogue.thread) fusion_args;
  fusion_args.alpha = params.alpha;
  fusion_args.beta = params.beta;

  arguments = typename Gemm::Arguments{
      cutlass::gemm::GemmUniversalMode::kGemm,
      {params.m, params.n, params.k, params.l},
      {
          static_cast<typename ElementA::DataType*>(params.ptr_A), stride_A,
          static_cast<typename ElementB::DataType*>(params.ptr_B), stride_B,
          static_cast<typename ElementA::ScaleFactorType*>(params.ptr_SFA), layout_SFA,
          static_cast<typename ElementB::ScaleFactorType*>(params.ptr_SFB), layout_SFB
      },
      {
          fusion_args,
          nullptr, stride_C,
          static_cast<ElementC*>(params.ptr_C), stride_D
      },
      hw_info,
      scheduler
  };
  //
  // -------- FIX #1: Persistent workspace -------------
  //

  size_t workspace_size = Gemm::get_workspace_size(arguments);

  static thread_local void* workspace_ptr = nullptr;
  static thread_local size_t workspace_cap = 0;

  if (workspace_size > workspace_cap) {
      if (workspace_ptr) {
          c10::cuda::CUDACachingAllocator::raw_delete(workspace_ptr);
      }

      workspace_ptr = c10::cuda::CUDACachingAllocator::raw_alloc_with_stream(
          workspace_size, stream);

      workspace_cap = workspace_size;
  }

  //
  // -------- FIX #2: Persistent initialized GEMM object -------------
  //
  static thread_local Gemm gemm;
  static thread_local bool initialized = false;

  if (!initialized) {
      CUTLASS_CHECK(gemm.can_implement(arguments));
      CUTLASS_CHECK(gemm.initialize(arguments, workspace_ptr, stream));
      initialized = true;
  } else {
      // Only update problem sizes + pointers
      CUTLASS_CHECK(gemm.update(arguments));
  }

  //
  // -------- Execute GEMM -------------
  //
  CUTLASS_CHECK(gemm.run(stream));
}

template void run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(
    GemmParams&, cudaStream_t);

void blockscaled_gemm_nvf4_f16(
    at::Tensor& A,
    at::Tensor& B,
    at::Tensor& C,
    at::Tensor& A_scales,
    at::Tensor& B_scales,
    int batch_size) {
  TORCH_CHECK(A.is_cuda(), "A must be CUDA");
  TORCH_CHECK(B.is_cuda(), "B must be CUDA");
  TORCH_CHECK(C.is_cuda(), "C must be CUDA");
  TORCH_CHECK(A_scales.is_cuda(), "A_scales must be CUDA");
  TORCH_CHECK(B_scales.is_cuda(), "B_scales must be CUDA");

  GemmParams params;
  params.m = static_cast<int>(A.size(0));
  params.n = static_cast<int>(B.size(0));
  params.k = static_cast<int>(A.size(1)) * 2;
  params.l = static_cast<int>(A.size(2));
  params.ptr_A = A.data_ptr();
  params.ptr_B = B.data_ptr();
  params.ptr_C = C.data_ptr();
  params.ptr_SFA = A_scales.data_ptr();
  params.ptr_SFB = B_scales.data_ptr();
  params.alpha = 1.0f;
  params.beta = 0.0f;

  auto stream = at::cuda::getCurrentCUDAStream().stream();
  run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(params, stream);
}

at::Tensor _make_blocked_tensor_like(const at::Tensor& src) {
  auto options = src.options().dtype(at::kByte);
  auto rows = src.size(0);
  auto cols = src.size(1);
  auto batches = src.size(2);

  TORCH_CHECK(src.dim() == 3, "Expected scale tensor with 3 dimensions");

  auto total_elements = rows * cols * batches;
  auto buffer = at::empty({total_elements}, options);

  auto blocked = buffer.as_strided(
      {rows, cols, batches},
      {cols, 1, rows * cols});
  return blocked;
}

void to_blocked_and_gemm(
    at::Tensor& A_fp4, at::Tensor& B_fp4,
    at::Tensor& C,
    at::Tensor& SFA, at::Tensor& SFB,
    int batch_size) {
  TORCH_CHECK(A_fp4.is_cuda(), "A_fp4 must be CUDA");
  TORCH_CHECK(B_fp4.is_cuda(), "B_fp4 must be CUDA");
  TORCH_CHECK(C.is_cuda(), "C must be CUDA");
  TORCH_CHECK(SFA.is_cuda(), "SFA must be CUDA");
  TORCH_CHECK(SFB.is_cuda(), "SFB must be CUDA");

  TORCH_CHECK(A_fp4.dim() == 3, "A tensor must be [m, k_packed, l]");
  TORCH_CHECK(B_fp4.dim() == 3, "B tensor must be [n, k_packed, l]");
  TORCH_CHECK(C.dim() == 3, "C tensor must be [m, n, l]");

  TORCH_CHECK(SFA.size(2) == batch_size, "SFA batch dimension mismatch");
  TORCH_CHECK(SFB.size(2) == batch_size, "SFB batch dimension mismatch");

  c10::cuda::OptionalCUDAGuard guard(A_fp4.device());
  auto torch_stream = at::cuda::getCurrentCUDAStream();
  auto cuda_stream = torch_stream.stream();

  auto A_scales_blocked = _make_blocked_tensor_like(SFA);
  auto B_scales_blocked = _make_blocked_tensor_like(SFB);

  to_blocked_batched_dual_cuda_wrapper(SFA, A_scales_blocked, SFB, B_scales_blocked);

  GemmParams params;
  params.m = static_cast<int>(A_fp4.size(0));
  params.n = static_cast<int>(B_fp4.size(0));
  params.k = static_cast<int>(A_fp4.size(1)) * 2;
  params.l = batch_size;
  params.ptr_A = A_fp4.data_ptr();
  params.ptr_B = B_fp4.data_ptr();
  params.ptr_C = C.data_ptr();
  params.ptr_SFA = A_scales_blocked.data_ptr();
  params.ptr_SFB = B_scales_blocked.data_ptr();
  params.alpha = 1.0f;
  params.beta = 0.0f;

  run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(params, cuda_stream);
}
"""

###############################################################################
# 3. Build the extension
###############################################################################

_THIS_DIR = os.path.dirname(__file__)
_INCLUDE_DIRS = [
    "/usr/local/include",
    "/usr/include",
    "/opt/cutlass/include",
    os.path.join(_THIS_DIR, "cutlass", "include"),
    os.path.abspath(os.path.join(_THIS_DIR, "..", "..", "..", "..", "cutlass", "include")),
    os.path.join(_THIS_DIR, "cutlass", "tools", "util", "include"),
    os.path.abspath(os.path.join(_THIS_DIR, "..", "..", "..", "..", "cutlass", "tools", "util", "include")),
]

_CUDA_FLAGS = [
    "-O3",
    "--use_fast_math",
    "-std=c++17",
    "-gencode=arch=compute_100a,code=sm_100a",
]
_CUDA_FLAGS.extend(f"-I{path}" for path in _INCLUDE_DIRS)

ext = load_inline(
    name="fp4gemv_ext_v9",
    cpp_sources=[binding_cpp],
    cuda_sources=[cuda_src],
    extra_cflags=["-O3", "-std=c++17"],
    extra_cuda_cflags=_CUDA_FLAGS,
    verbose=False,
)

###############################################################################
# 4. Helper utilities shared with the kernel wrapper
###############################################################################

def ensure_packed_fp4(tensor: torch.Tensor, rows: int, k_elements: int, batches: int) -> torch.Tensor:
    """
    Ensure FP4 tensor is packed as NVF4 bytes (two 4-bit values per byte).
    Returns a tensor of shape [rows, k_elements//2, batches] in uint8.
    """
    bytes_view = tensor.view(torch.uint8).reshape(rows, -1, batches)
    bytes_per_row = bytes_view.size(1)
    target_bytes = k_elements // 2

    if bytes_per_row != target_bytes:
        raise RuntimeError(f"Expected packed NVF4 with {target_bytes} bytes/row, got {bytes_per_row}")

    return bytes_view


def _allocate_blocked_buffer(rows: int, cols: int, batches: int, device):
    buf = torch.empty((rows * cols * batches,), dtype=torch.uint8, device=device)
    return buf.as_strided((rows, cols, batches), (cols, 1, rows * cols))


def to_blocked_batched_pair(t0: torch.Tensor, t1: torch.Tensor):
    r0, c0, b0 = t0.shape
    r1, c1, b1 = t1.shape
    out0 = _allocate_blocked_buffer(r0, c0, b0, t0.device)
    out1 = _allocate_blocked_buffer(r1, c1, b1, t1.device)
    ext.to_blocked_batched_dual_cuda(t0, out0, t1, out1)
    return out0, out1


def to_blocked_batched_fast(t: torch.Tensor) -> torch.Tensor:
    r, c, b = t.shape
    out = _allocate_blocked_buffer(r, c, b, t.device)
    ext.to_blocked_batched_cuda(t, out)
    return out


def _require_profiler_stride(tensor: torch.Tensor, dim0: int, dim1: int, batches: int, name: str):
    expected_stride = (dim1, 1, dim0 * dim1)
    if tensor.stride() != expected_stride:
        raise RuntimeError(
            f"{name} must use profiler stride {expected_stride}, got {tensor.stride()}"
        )


def blockscaled_gemm(A: torch.Tensor,
                     B: torch.Tensor,
                     C: torch.Tensor,
                     A_scales: torch.Tensor,
                     B_scales: torch.Tensor,
                     batch_size: int = 1) -> torch.Tensor:
    m, k_packed, l = A.shape
    n = B.size(0)

    _require_profiler_stride(A, m, k_packed, l, "A")
    _require_profiler_stride(B, n, k_packed, l, "B")
    _require_profiler_stride(C, m, n, l, "C")

    ext.to_blocked_and_gemm(A, B, C, A_scales, B_scales, batch_size)

    return C


###############################################################################
# 5. GPU submission entry point expected by the leaderboard harness
###############################################################################

def custom_kernel(submission_input):
    a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_template = submission_input

    m, _, l = a_ref.shape
    tile_n = b_ref.size(0)

    # Create C with profiler-style stride to avoid conversion
    C_buf = torch.empty((m * tile_n * l,), dtype=torch.float16, device=c_template.device)
    C_full = C_buf.as_strided((m, tile_n, l), (tile_n, 1, m * tile_n))

    sfa_u8 = sfa_ref.view(torch.uint8)
    sfb_u8 = sfb_ref.view(torch.uint8)

    blockscaled_gemm(
        a_ref,
        b_ref,
        C_full,
        sfa_u8,
        sfb_u8,
        batch_size=l,
    )

    return C_full.narrow(1, 0, 1)

scrolls · 743 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