Skip to content
KernelIndex
Search⌘K

submission 78065

gau.nernst · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v2_cpasync.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-78065?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
25.4µs
#107 of 678
2025-11-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2eacadee22c6ae6fdff26a4240b5836b4974bae33d6a39d1141a14823aaa1e25
license declaredunknown
license concludedunknown
authorsgau.nernst
imported2026-08-15

Techniques

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

async-copyvoid cp_async_2d(int dst, const T *src, int src_stride, int tid) {
fp4"cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"
fp8const __nv_fp8_e4m3 *SFA_ptr, // [L, M, K/8]
num-warps = 4constexpr int NUM_WARPS = 4;
shared-memoryextern __shared__ char smem[];
stages = 1template <int THREAD_K, int NUM_STAGES = 1>
vector-width = float2float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;

Kernel source

submission_v2_cpasync.py342 lines
#!POPCORN leaderboard nvfp4_gemv

from pathlib import Path

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

CUDA_SRC = r"""
#include <cuda_fp16.h>
#include <cuda_fp8.h>

#include <torch/library.h>
#include <ATen/ATen.h>
#include <ATen/core/Tensor.h>
#include <ATen/cuda/CUDAUtils.h>
#include <ATen/cuda/CUDAContext.h>

constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 4;
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;
constexpr int THREAD_M = 4;

__device__
void fp4x8_to_fp32x2x4(int in, int64_t *out) {
  int tmp[4];
  asm volatile(
    "{\n"
    ".reg .b8 tmp0, tmp1, tmp2, tmp3;\n"
    "mov.b32 {tmp0, tmp1, tmp2, tmp3}, %4; // unpack 32-bit register to 4x fp4x2\n"
    "cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"
    "cvt.rn.f16x2.e2m1x2 %1, tmp1;\n"
    "cvt.rn.f16x2.e2m1x2 %2, tmp2;\n"
    "cvt.rn.f16x2.e2m1x2 %3, tmp3;\n"
    "}\n"
    : "=r"(tmp[0]), "=r"(tmp[1]), "=r"(tmp[2]), "=r"(tmp[3])
    : "r"(in)
  );

  for (int i = 0; i < 4; i++)
    asm volatile(
      "{\n"
      ".reg .b16 b16_0, b16_1;\n"
      ".reg .b32 f32_0, f32_1;\n"
      "mov.b32 {b16_0, b16_1}, %1;  // unpack\n"
      "cvt.f32.f16 f32_0, b16_0;\n"
      "cvt.f32.f16 f32_1, b16_1;\n"
      "mov.b64 %0, {f32_0, f32_1};  // pack\n"
      "}\n"
      : "=l"(out[i])
      : "r"(tmp[i])
    );
}

template <int HEIGHT, int WIDTH, int TB_SIZE, typename T>
__device__
void cp_async_2d(int dst, const T *src, int src_stride, int tid) {
  auto load = [&](int idx) {
    const int row = idx / WIDTH;
    const int col = idx % WIDTH;

    const int dst_addr = dst + idx * sizeof(T);
    const T *src_addr = src + (row * src_stride + col);
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst_addr), "l"(src_addr));
  };

  constexpr int num_elems = 16 / sizeof(T);
  constexpr int num_iters = HEIGHT * WIDTH / (TB_SIZE * num_elems);

  for (int iter = 0; iter < num_iters; iter++)
    load((iter * TB_SIZE + tid) * num_elems);

  // handle the case when tile size is not divisible by threadblock size
  if constexpr ((HEIGHT * WIDTH) % (TB_SIZE * num_elems) != 0) {
    const int idx = (num_iters * TB_SIZE + tid) * num_elems;
    if (idx < HEIGHT * WIDTH)
      load(idx);
  }
}

// to make our calculations simple, let's treat fp4x2 as a unit.
// hence, K = number of fp4x2 elements, and 8 elements share
// the same scale.
template <int THREAD_K, int NUM_STAGES = 1>
__global__
__launch_bounds__(NUM_WARPS * WARP_SIZE)
void kernel(
  const char          *A_ptr,    // [L,   M, K]
  const char          *B_ptr,    // [L, 128, K]
  const __nv_fp8_e4m3 *SFA_ptr,  // [L,   M, K/8]
  const __nv_fp8_e4m3 *SFB_ptr,  // [L, 128, K/8]
        half          *C_ptr,    // [L,   M]
  int L, int M, int K
) {
  // to ensure coalesced access, we need at least 8 threads per row (16B x 8 = 128B)
  // each thread reads 16B, which covers 2 scaled groups. hence, we only need within
  // thread reduction during the main loop.
  static_assert(THREAD_M == 4);
  static_assert(THREAD_K >= 8);
  static_assert(THREAD_K <= TB_SIZE);
  constexpr int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M;
  constexpr int BLOCK_K = THREAD_K * 16;
  constexpr int SF_BLOCK_K = BLOCK_K / 8;

  const int tid = threadIdx.x;
  const int bid = blockIdx.x;
  const int batch_id = blockIdx.y;

  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  const int off_m = bid * BLOCK_M;
  const int off_k = (tid % THREAD_K) * 16;  // each thread reads 16 fp4x2 values at a time

  A_ptr += (batch_id *   M * K) + (off_m * K);
  B_ptr += (batch_id * 128 * K);

  SFA_ptr += (batch_id *   M * (K / 8)) + (off_m * (K / 8));
  SFB_ptr += (batch_id * 128 * (K / 8));

  // set up smem
  extern __shared__ char smem[];
  const int smem_u32 = static_cast<int>(__cvta_generic_to_shared(smem));
  constexpr int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K;

  char   *A_smem = smem;
  char   *B_smem =   A_smem + BLOCK_M * BLOCK_K;
  char *SFA_smem =   B_smem + BLOCK_K;
  char *SFB_smem = SFA_smem + BLOCK_M * SF_BLOCK_K;

  // to be used for smem->rmem load
  char   *A_smem_ld =   A_smem + (tid / THREAD_K) * THREAD_M *    BLOCK_K + off_k;
  char   *B_smem_ld =   B_smem                                            + off_k;
  char *SFA_smem_ld = SFA_smem + (tid / THREAD_K) * THREAD_M * SF_BLOCK_K + (off_k / 8);
  char *SFB_smem_ld = SFB_smem +                                          + (off_k / 8);

  float acc[THREAD_M] = {};

  auto load = [&](int iter_k) {
    // NOTE: since B, SFA, and SFB does not require the whole threadblock to load, we can partition it within the threadblock.
    const int buffer = smem_u32 + (iter_k % NUM_STAGES) * TOTAL_SMEM;
    const int A_buf   = buffer;
    const int B_buf   = A_buf + BLOCK_M * BLOCK_K;
    const int SFA_buf = B_buf + BLOCK_K;
    const int SFB_buf = SFA_buf + BLOCK_M * SF_BLOCK_K;

    cp_async_2d<BLOCK_M,    BLOCK_K, TB_SIZE>(  A_buf,   A_ptr,     K, tid);
    cp_async_2d<      1,    BLOCK_K, TB_SIZE>(  B_buf,   B_ptr,     K, tid);
    cp_async_2d<BLOCK_M, SF_BLOCK_K, TB_SIZE>(SFA_buf, SFA_ptr, K / 8, tid);
    cp_async_2d<      1, SF_BLOCK_K, TB_SIZE>(SFB_buf, SFB_ptr, K / 8, tid);

    asm volatile("cp.async.commit_group;\n");

    A_ptr += BLOCK_K;
    B_ptr += BLOCK_K;
    SFA_ptr += BLOCK_K / 8;
    SFB_ptr += BLOCK_K / 8;
  };

  for (int iter_k = 0; iter_k < NUM_STAGES - 1; iter_k++)
    load(iter_k);

  const int num_iters = K / BLOCK_K;

  for (int iter_k = 0; iter_k < num_iters; iter_k++) {
    // gmem -> smem
    if (iter_k + NUM_STAGES - 1 < num_iters) {
      __syncthreads();  // make sure previous compute finish using the buffer
      load(iter_k + NUM_STAGES - 1);
    } else {
      asm volatile("cp.async.commit_group;\n");
    }

    // smem -> rmem
    asm volatile("cp.async.wait_group %0;\n" :: "n"(NUM_STAGES - 1));
    __syncthreads();  // memory barrier

    int A_fp4x8[THREAD_M][4], B_fp4x8[4];
    float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;
    int buf_offset = (iter_k % NUM_STAGES) * TOTAL_SMEM;

    for (int m = 0; m < THREAD_M; m++) {
      reinterpret_cast<int4 *>(A_fp4x8[m])[0] = reinterpret_cast<const int4 *>(A_smem_ld + buf_offset + m * BLOCK_K)[0];
      SFA_fp32x2[m] = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFA_smem_ld + buf_offset + m * SF_BLOCK_K)[0]);
    }

    reinterpret_cast<int4 *>(B_fp4x8)[0] = reinterpret_cast<const int4 *>(B_smem_ld + buf_offset)[0];
    SFB_fp32x2 = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFB_smem_ld + buf_offset)[0]);

    // unpack to FP32
    int64_t A_fp32x2[THREAD_M][16], B_fp32x2[16];

    for (int m = 0; m < THREAD_M; m++)
      for (int i = 0; i < 4; i++)
        fp4x8_to_fp32x2x4(A_fp4x8[m][i], A_fp32x2[m] + i * 4);

    for (int i = 0; i < 4; i++)
      fp4x8_to_fp32x2x4(B_fp4x8[i], B_fp32x2 + i * 4);

    for (int m = 0; m < THREAD_M; m++)
      for (int group_id = 0; group_id < 2; group_id++) {
        // FMA. manually unroll the 1st iteration
        int64_t sub_acc;
        asm volatile("mul.rn.f32x2 %0, %1, %2;\n"
                    : "=l"(sub_acc)
                    : "l"(A_fp32x2[m][group_id * 8]), "l"(B_fp32x2[group_id * 8]));
        for (int i = 1; i < 8; i++)
          asm volatile("fma.rn.f32x2 %0, %1, %2, %0;\n"
                      : "+l"(sub_acc)
                      : "l"(A_fp32x2[m][group_id * 8 + i]), "l"(B_fp32x2[group_id * 8 + i]));

        float tmp[2];
        std::memcpy(tmp, &sub_acc, sizeof(sub_acc));

        float sfa = reinterpret_cast<float *>(SFA_fp32x2 + m)[group_id];
        float sfb = reinterpret_cast<float *>(&SFB_fp32x2)[group_id];
        acc[m] += (tmp[0] + tmp[1]) * sfa * sfb;
      }
  }

  // this is so cursed
  long2 acc_fp32x2x2;
  std::memcpy(&acc_fp32x2x2, acc, sizeof(acc_fp32x2x2));

  // threadblock reduction
  if constexpr (THREAD_K > WARP_SIZE) {
    __shared__ long2 smem[TB_SIZE];
    smem[tid] = acc_fp32x2x2;
    __syncthreads();

    for (int stride = THREAD_K / 2; stride >= WARP_SIZE; stride /= 2) {
      if ((tid % THREAD_K) < stride) {
        long2 tmp = smem[tid + stride];
        asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp.x));
        asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp.y));
        smem[tid] = acc_fp32x2x2;
      }
      __syncthreads();
    }
  }

  // warp reduction
  constexpr int start_stride = std::min(THREAD_K, WARP_SIZE) / 2;
  for (int stride = start_stride; stride > 0; stride /= 2) {
    long tmp[2];
    tmp[0] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.x, stride);
    tmp[1] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.y, stride);
    asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp[0]));
    asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp[1]));
  }

  if (tid % THREAD_K == 0) {
    half2 out[2];
    out[0] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[0]);
    out[1] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[1]);
    reinterpret_cast<int2 *>(C_ptr + (batch_id * M + off_m + (tid / THREAD_K) * THREAD_M))[0] = reinterpret_cast<int2 *>(out)[0];
  }
}

void gemv(
  const at::Tensor& A,
  const at::Tensor& B,
  const at::Tensor& SFA,
  const at::Tensor& SFB,
        at::Tensor& C
) {
  const int M = A.size(0);
  const int K = A.size(1);
  const int L = A.size(2);

  auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
  auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
  auto SFA_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFA.data_ptr());
  auto SFB_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFB.data_ptr());
  auto C_ptr = reinterpret_cast<half *>(C.data_ptr());

  auto stream = at::cuda::getCurrentCUDAStream();
  constexpr int NUM_STAGES = 2;

#define launch(THREAD_K) { \
  int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M; \
  int BLOCK_K = THREAD_K * 16; \
  int SF_BLOCK_K = BLOCK_K / 8; \
  int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K; \
  dim3 grid(M / BLOCK_M, L); \
  int smem_size = TOTAL_SMEM * NUM_STAGES; \
  kernel<THREAD_K, NUM_STAGES><<<grid, TB_SIZE, smem_size, stream>>>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, L, M, K); \
}

  if (false) {}
  else if (K % (128 * 16) == 0) launch(128)  // benchmark.0
  else if (K % (32 * 16) == 0) launch(32)    // benchmark.1 and benchmark.2
  else launch(8)                             // the rest

#undef launch
}

TORCH_LIBRARY(my_module, m) {
  m.def("gemv(Tensor A, Tensor B, Tensor SFA, Tensor SFB, Tensor(a!) C) -> ()");
  m.impl("gemv", &gemv);
}
"""

load_inline(
    "gemv_c0",
    cpp_sources="",
    cuda_sources=CUDA_SRC,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "-gencode=arch=compute_120a,code=sm_120a",
        "-lineinfo",
    ],
)


def custom_kernel(data: input_t) -> output_t:
    # a:   [  M, K, L],                   natural shape [L,   M, K]
    # b:   [128, K, L],                   natural shape [L, 128, K] - only the 1st row is used
    # sfa: [32, 4, rest_m, 4, rest_k, L], natural shape [L, rest_m, rest_k, 32, 4, 4]
    # sfb: [32, 4,      1, 4, rest_k, L], natural shape [L,      1, rest_k, 32, 4, 4]
    # c:   [  M, 1, L],                   natural shape [L, M, 1]
    a, b, sfa, sfb, _, _, c_ref = data
    torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)

    if False:
        M, K, L = a.shape
        path = Path(f"profile_data/{M=}_K={K * 2}_{L=}.json.gz")
        if not path.exists():
            a.new_zeros(int(1e8), dtype=torch.uint8)  # 100 MB

            with torch.profiler.profile() as prof:
                torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)

            path.parent.mkdir(exist_ok=True)
            prof.export_chrome_trace(str(path))

    return c_ref
scrolls · 342 lines total

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

Changes from previous submission

Against this author's previous submission submission 73354.

⋯ 5 unchanged lines
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
- # https://github.com/NVIDIA/cutlass/blob/v4.2.1/examples/72_blackwell_narrow_precision_gemm/72b_blackwell_nvfp4_nvfp4_gemm.cu
CUDA_SRC = r"""
- #include "cutlass/cutlass.h"
+ #include <cuda_fp16.h>
+ #include <cuda_fp8.h>
- #include "cute/tensor.hpp"
- #include "cutlass/tensor_ref.h"
- #include "cutlass/epilogue/thread/linear_combination.h"
- #include "cutlass/gemm/dispatch_policy.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/gemm/kernel/tile_scheduler_params.h"
-
- #include "cutlass/util/packed_stride.hpp"
-
#include <torch/library.h>
#include <ATen/ATen.h>
#include <ATen/core/Tensor.h>
#include <ATen/cuda/CUDAUtils.h>
#include <ATen/cuda/CUDAContext.h>
- #define STRINGIFY(x) #x
- #define CUTLASS_CHECK(call) \
- do { \
- auto status = call; \
- TORCH_CHECK(status == cutlass::Status::kSuccess, STRINGIFY(call), ": ", status, " - ", cutlassGetStatusString(status)); \
- } while (0)
+ constexpr int WARP_SIZE = 32;
+ constexpr int NUM_WARPS = 4;
+ constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;
+ constexpr int THREAD_M = 4;
- using namespace cute;
+ __device__
+ void fp4x8_to_fp32x2x4(int in, int64_t *out) {
+ int tmp[4];
+ asm volatile(
+ "{\n"
+ ".reg .b8 tmp0, tmp1, tmp2, tmp3;\n"
+ "mov.b32 {tmp0, tmp1, tmp2, tmp3}, %4; // unpack 32-bit register to 4x fp4x2\n"
+ "cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"
+ "cvt.rn.f16x2.e2m1x2 %1, tmp1;\n"
+ "cvt.rn.f16x2.e2m1x2 %2, tmp2;\n"
+ "cvt.rn.f16x2.e2m1x2 %3, tmp3;\n"
+ "}\n"
+ : "=r"(tmp[0]), "=r"(tmp[1]), "=r"(tmp[2]), "=r"(tmp[3])
+ : "r"(in)
+ );
- using ElementAB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
- using ElementC = cutlass::half_t;
- using ElementAcc = float;
+ for (int i = 0; i < 4; i++)
+ asm volatile(
+ "{\n"
+ ".reg .b16 b16_0, b16_1;\n"
+ ".reg .b32 f32_0, f32_1;\n"
+ "mov.b32 {b16_0, b16_1}, %1; // unpack\n"
+ "cvt.f32.f16 f32_0, b16_0;\n"
+ "cvt.f32.f16 f32_1, b16_1;\n"
+ "mov.b64 %0, {f32_0, f32_1}; // pack\n"
+ "}\n"
+ : "=l"(out[i])
+ : "r"(tmp[i])
+ );
+ }
- constexpr int AlignmentAB = 128 / 4; // 32
- constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value; // 8
+ template <int HEIGHT, int WIDTH, int TB_SIZE, typename T>
+ __device__
+ void cp_async_2d(int dst, const T *src, int src_stride, int tid) {
+ auto load = [&](int idx) {
+ const int row = idx / WIDTH;
+ const int col = idx % WIDTH;
- using LayoutATag = cutlass::layout::RowMajor;
- using LayoutBTag = cutlass::layout::ColumnMajor;
- using LayoutCTag = cutlass::layout::RowMajor;
+ const int dst_addr = dst + idx * sizeof(T);
+ const T *src_addr = src + (row * src_stride + col);
+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst_addr), "l"(src_addr));
+ };
- using ArchTag = cutlass::arch::Sm100;
- using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
+ constexpr int num_elems = 16 / sizeof(T);
+ constexpr int num_iters = HEIGHT * WIDTH / (TB_SIZE * num_elems);
- // Kernel Perf config
- using MmaTileShape = Shape<_128,_128,_256>;
- using ClusterShape = Shape<_1,_1,_1>;
+ for (int iter = 0; iter < num_iters; iter++)
+ load((iter * TB_SIZE + tid) * num_elems);
- using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
- ArchTag, OperatorClass,
- MmaTileShape, ClusterShape,
- cutlass::epilogue::collective::EpilogueTileAuto,
- ElementAcc, ElementAcc,
- ElementC, LayoutCTag, AlignmentC,
- ElementC, LayoutCTag, AlignmentC,
- cutlass::epilogue::collective::EpilogueScheduleAuto
- >::CollectiveOp;
+ // handle the case when tile size is not divisible by threadblock size
+ if constexpr ((HEIGHT * WIDTH) % (TB_SIZE * num_elems) != 0) {
+ const int idx = (num_iters * TB_SIZE + tid) * num_elems;
+ if (idx < HEIGHT * WIDTH)
+ load(idx);
+ }
+ }
- using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
- ArchTag, OperatorClass,
- ElementAB, LayoutATag, AlignmentAB,
- ElementAB, LayoutBTag, AlignmentAB,
- ElementAcc,
- MmaTileShape, ClusterShape,
- cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
- cutlass::gemm::collective::KernelScheduleAuto
- >::CollectiveOp;
+ // to make our calculations simple, let's treat fp4x2 as a unit.
+ // hence, K = number of fp4x2 elements, and 8 elements share
+ // the same scale.
+ template <int THREAD_K, int NUM_STAGES = 1>
+ __global__
+ __launch_bounds__(NUM_WARPS * WARP_SIZE)
+ void kernel(
+ const char *A_ptr, // [L, M, K]
+ const char *B_ptr, // [L, 128, K]
+ const __nv_fp8_e4m3 *SFA_ptr, // [L, M, K/8]
+ const __nv_fp8_e4m3 *SFB_ptr, // [L, 128, K/8]
+ half *C_ptr, // [L, M]
+ int L, int M, int K
+ ) {
+ // to ensure coalesced access, we need at least 8 threads per row (16B x 8 = 128B)
+ // each thread reads 16B, which covers 2 scaled groups. hence, we only need within
+ // thread reduction during the main loop.
+ static_assert(THREAD_M == 4);
+ static_assert(THREAD_K >= 8);
+ static_assert(THREAD_K <= TB_SIZE);
+ constexpr int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M;
+ constexpr int BLOCK_K = THREAD_K * 16;
+ constexpr int SF_BLOCK_K = BLOCK_K / 8;
- using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
- Shape<int, int, int, int>,
- CollectiveMainloop,
- CollectiveEpilogue,
- void>;
+ const int tid = threadIdx.x;
+ const int bid = blockIdx.x;
+ const int batch_id = blockIdx.y;
- using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
+ const int lane_id = tid % WARP_SIZE;
+ const int warp_id = tid / WARP_SIZE;
+ const int off_m = bid * BLOCK_M;
+ const int off_k = (tid % THREAD_K) * 16; // each thread reads 16 fp4x2 values at a time
+
+ A_ptr += (batch_id * M * K) + (off_m * K);
+ B_ptr += (batch_id * 128 * K);
+
+ SFA_ptr += (batch_id * M * (K / 8)) + (off_m * (K / 8));
+ SFB_ptr += (batch_id * 128 * (K / 8));
+
+ // set up smem
+ extern __shared__ char smem[];
+ const int smem_u32 = static_cast<int>(__cvta_generic_to_shared(smem));
+ constexpr int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K;
+
+ char *A_smem = smem;
+ char *B_smem = A_smem + BLOCK_M * BLOCK_K;
+ char *SFA_smem = B_smem + BLOCK_K;
+ char *SFB_smem = SFA_smem + BLOCK_M * SF_BLOCK_K;
+
+ // to be used for smem->rmem load
+ char *A_smem_ld = A_smem + (tid / THREAD_K) * THREAD_M * BLOCK_K + off_k;
+ char *B_smem_ld = B_smem + off_k;
+ char *SFA_smem_ld = SFA_smem + (tid / THREAD_K) * THREAD_M * SF_BLOCK_K + (off_k / 8);
+ char *SFB_smem_ld = SFB_smem + + (off_k / 8);
+
+ float acc[THREAD_M] = {};
+
+ auto load = [&](int iter_k) {
+ // NOTE: since B, SFA, and SFB does not require the whole threadblock to load, we can partition it within the threadblock.
+ const int buffer = smem_u32 + (iter_k % NUM_STAGES) * TOTAL_SMEM;
+ const int A_buf = buffer;
+ const int B_buf = A_buf + BLOCK_M * BLOCK_K;
+ const int SFA_buf = B_buf + BLOCK_K;
+ const int SFB_buf = SFA_buf + BLOCK_M * SF_BLOCK_K;
+
+ cp_async_2d<BLOCK_M, BLOCK_K, TB_SIZE>( A_buf, A_ptr, K, tid);
+ cp_async_2d< 1, BLOCK_K, TB_SIZE>( B_buf, B_ptr, K, tid);
+ cp_async_2d<BLOCK_M, SF_BLOCK_K, TB_SIZE>(SFA_buf, SFA_ptr, K / 8, tid);
+ cp_async_2d< 1, SF_BLOCK_K, TB_SIZE>(SFB_buf, SFB_ptr, K / 8, tid);
+
+ asm volatile("cp.async.commit_group;\n");
+
+ A_ptr += BLOCK_K;
+ B_ptr += BLOCK_K;
+ SFA_ptr += BLOCK_K / 8;
+ SFB_ptr += BLOCK_K / 8;
+ };
+
+ for (int iter_k = 0; iter_k < NUM_STAGES - 1; iter_k++)
+ load(iter_k);
+
+ const int num_iters = K / BLOCK_K;
+
+ for (int iter_k = 0; iter_k < num_iters; iter_k++) {
+ // gmem -> smem
+ if (iter_k + NUM_STAGES - 1 < num_iters) {
+ __syncthreads(); // make sure previous compute finish using the buffer
+ load(iter_k + NUM_STAGES - 1);
+ } else {
+ asm volatile("cp.async.commit_group;\n");
+ }
+
+ // smem -> rmem
+ asm volatile("cp.async.wait_group %0;\n" :: "n"(NUM_STAGES - 1));
+ __syncthreads(); // memory barrier
+
+ int A_fp4x8[THREAD_M][4], B_fp4x8[4];
+ float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;
+ int buf_offset = (iter_k % NUM_STAGES) * TOTAL_SMEM;
+
+ for (int m = 0; m < THREAD_M; m++) {
+ reinterpret_cast<int4 *>(A_fp4x8[m])[0] = reinterpret_cast<const int4 *>(A_smem_ld + buf_offset + m * BLOCK_K)[0];
+ SFA_fp32x2[m] = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFA_smem_ld + buf_offset + m * SF_BLOCK_K)[0]);
+ }
+
+ reinterpret_cast<int4 *>(B_fp4x8)[0] = reinterpret_cast<const int4 *>(B_smem_ld + buf_offset)[0];
+ SFB_fp32x2 = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFB_smem_ld + buf_offset)[0]);
+
+ // unpack to FP32
+ int64_t A_fp32x2[THREAD_M][16], B_fp32x2[16];
+
+ for (int m = 0; m < THREAD_M; m++)
+ for (int i = 0; i < 4; i++)
+ fp4x8_to_fp32x2x4(A_fp4x8[m][i], A_fp32x2[m] + i * 4);
+
+ for (int i = 0; i < 4; i++)
+ fp4x8_to_fp32x2x4(B_fp4x8[i], B_fp32x2 + i * 4);
+
+ for (int m = 0; m < THREAD_M; m++)
+ for (int group_id = 0; group_id < 2; group_id++) {
+ // FMA. manually unroll the 1st iteration
+ int64_t sub_acc;
+ asm volatile("mul.rn.f32x2 %0, %1, %2;\n"
+ : "=l"(sub_acc)
+ : "l"(A_fp32x2[m][group_id * 8]), "l"(B_fp32x2[group_id * 8]));
+ for (int i = 1; i < 8; i++)
+ asm volatile("fma.rn.f32x2 %0, %1, %2, %0;\n"
+ : "+l"(sub_acc)
+ : "l"(A_fp32x2[m][group_id * 8 + i]), "l"(B_fp32x2[group_id * 8 + i]));
+
+ float tmp[2];
+ std::memcpy(tmp, &sub_acc, sizeof(sub_acc));
+
+ float sfa = reinterpret_cast<float *>(SFA_fp32x2 + m)[group_id];
+ float sfb = reinterpret_cast<float *>(&SFB_fp32x2)[group_id];
+ acc[m] += (tmp[0] + tmp[1]) * sfa * sfb;
+ }
+ }
+
+ // this is so cursed
+ long2 acc_fp32x2x2;
+ std::memcpy(&acc_fp32x2x2, acc, sizeof(acc_fp32x2x2));
+
+ // threadblock reduction
+ if constexpr (THREAD_K > WARP_SIZE) {
+ __shared__ long2 smem[TB_SIZE];
+ smem[tid] = acc_fp32x2x2;
+ __syncthreads();
+
+ for (int stride = THREAD_K / 2; stride >= WARP_SIZE; stride /= 2) {
+ if ((tid % THREAD_K) < stride) {
+ long2 tmp = smem[tid + stride];
+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp.x));
+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp.y));
+ smem[tid] = acc_fp32x2x2;
+ }
+ __syncthreads();
+ }
+ }
+
+ // warp reduction
+ constexpr int start_stride = std::min(THREAD_K, WARP_SIZE) / 2;
+ for (int stride = start_stride; stride > 0; stride /= 2) {
+ long tmp[2];
+ tmp[0] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.x, stride);
+ tmp[1] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.y, stride);
+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp[0]));
+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp[1]));
+ }
+
+ if (tid % THREAD_K == 0) {
+ half2 out[2];
+ out[0] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[0]);
+ out[1] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[1]);
+ reinterpret_cast<int2 *>(C_ptr + (batch_id * M + off_m + (tid / THREAD_K) * THREAD_M))[0] = reinterpret_cast<int2 *>(out)[0];
+ }
+ }
+
void gemv(
const at::Tensor& A,
const at::Tensor& B,
⋯ 2 unchanged lines
at::Tensor& C
) {
const int M = A.size(0);
- const int N = 128;
- const int K = A.size(1) * 2;
+ const int K = A.size(1);
const int L = A.size(2);
- using ABType = typename ElementAB::DataType;
- using SFType = typename ElementAB::ScaleFactorType;
+ auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
+ auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
+ auto SFA_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFA.data_ptr());
+ auto SFB_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFB.data_ptr());
+ auto C_ptr = reinterpret_cast<half *>(C.data_ptr());
- auto stride_A = cutlass::make_cute_packed_stride(typename GemmKernel::StrideA{}, {M, K, L});
- auto stride_B = cutlass::make_cute_packed_stride(typename GemmKernel::StrideB{}, {N, K, L});
- auto stride_C = cutlass::make_cute_packed_stride(typename GemmKernel::StrideC{}, {M, N, L});
+ auto stream = at::cuda::getCurrentCUDAStream();
+ constexpr int NUM_STAGES = 2;
- using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
- auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, L));
- auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, L));
+ #define launch(THREAD_K) { \
+ int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M; \
+ int BLOCK_K = THREAD_K * 16; \
+ int SF_BLOCK_K = BLOCK_K / 8; \
+ int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K; \
+ dim3 grid(M / BLOCK_M, L); \
+ int smem_size = TOTAL_SMEM * NUM_STAGES; \
+ kernel<THREAD_K, NUM_STAGES><<<grid, TB_SIZE, smem_size, stream>>>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, L, M, K); \
+ }
- auto *A_ptr = reinterpret_cast<const ABType *>(A.data_ptr());
- auto *B_ptr = reinterpret_cast<const ABType *>(B.data_ptr());
- auto *SFA_ptr = reinterpret_cast<const SFType *>(SFA.data_ptr());
- auto *SFB_ptr = reinterpret_cast<const SFType *>(SFB.data_ptr());
- auto *C_ptr = reinterpret_cast<ElementC *>(C.data_ptr());
+ if (false) {}
+ else if (K % (128 * 16) == 0) launch(128) // benchmark.0
+ else if (K % (32 * 16) == 0) launch(32) // benchmark.1 and benchmark.2
+ else launch(8) // the rest
- typename Gemm::Arguments arguments{
- cutlass::gemm::GemmUniversalMode::kGemm,
- {M, N, K, L},
- {
- A_ptr, stride_A,
- B_ptr, stride_B,
- SFA_ptr, layout_SFA,
- SFB_ptr, layout_SFB,
- },
- {
- {1.0f, 0.0f}, // alpha and beta
- C_ptr, stride_C,
- C_ptr, stride_C,
- }
- };
-
- Gemm gemm;
- //CUTLASS_CHECK(gemm.can_implement(arguments));
-
- //long workspace_size = Gemm::get_workspace_size(arguments);
- //at::Tensor workspace = at::empty({workspace_size}, A.options().dtype(at::kByte));
- auto stream = at::cuda::getCurrentCUDAStream();
-
- //CUTLASS_CHECK(gemm.initialize(arguments, workspace.data_ptr(), stream));
- CUTLASS_CHECK(gemm.initialize(arguments, 0, stream));
- CUTLASS_CHECK(gemm.run(stream));
+ #undef launch
}
TORCH_LIBRARY(my_module, m) {
⋯ 12 unchanged lines
extra_cuda_cflags=[
"-O3",
"-gencode=arch=compute_100a,code=sm_100a",
+ "-gencode=arch=compute_120a,code=sm_120a",
+ "-lineinfo",
],
)
⋯ 4 unchanged lines
# sfa: [32, 4, rest_m, 4, rest_k, L], natural shape [L, rest_m, rest_k, 32, 4, 4]
# sfb: [32, 4, 1, 4, rest_k, L], natural shape [L, 1, rest_k, 32, 4, 4]
# c: [ M, 1, L], natural shape [L, M, 1]
- a, b, _, _, sfa, sfb, c_ref = data
+ a, b, sfa, sfb, _, _, c_ref = data
+ torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)
- M = a.shape[0]
- N = 128
- K = a.shape[1] * 2
- L = a.shape[2]
-
- big_c = c_ref.new_empty(L, M, N)
- torch.ops.my_module.gemv(a, b, sfa, sfb, big_c)
-
if False:
- path = Path(f"profile_data/{M=}_{K=}_{L=}.json.gz")
+ M, K, L = a.shape
+ path = Path(f"profile_data/{M=}_K={K * 2}_{L=}.json.gz")
if not path.exists():
a.new_zeros(int(1e8), dtype=torch.uint8) # 100 MB
with torch.profiler.profile() as prof:
- torch.ops.my_module.gemv(a, b, sfa, sfb, big_c)
+ torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)
path.parent.mkdir(exist_ok=True)
prof.export_chrome_trace(str(path))
- return big_c[..., :1].permute(1, 2, 0) # convert to [M, 1, L]
+ return c_ref
scrolls · 434 diff lines total

Best evidence level for this revision: reported

JSON