submission 85852
s.am._ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 451 lines, June 9 Researcher Reciprocity License v1.0.
finetuned_naive.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-85852?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:0185e2ca04468d611f39ba5acd72f25ed21731bc5891a98e6b7dd5a6c98bc337
license declaredunknown
license concludedunknown
authorss.am._
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
if (params.k <= 256) { // <= 512 FP4 valuesfp8
const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;shared-memory
__shared__ float sdata[THREADS_PER_ROW];Kernel source
finetuned_naive.py451 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# ---- C++ stub: declare the function so load_inline can bind it ----
gemv_cpp = r"""
#include <torch/extension.h>
// Forward declaration so PyTorch can bind it (definition is in the CUDA source).
torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB);
"""
# ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----
gemv_cuda = r"""
#include <assert.h>
#include <cuda.h>
#include <stdio.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp4.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>
// ---- gemv.h ----
struct Gemv_params {
using index_t = uint64_t;
int b, m, k, real_k;
void *__restrict__ a_ptr;
void *__restrict__ b_ptr;
void *__restrict__ sfa_ptr;
void *__restrict__ sfb_ptr;
void *__restrict__ o_ptr;
index_t a_batch_stride;
index_t b_batch_stride;
index_t sfa_batch_stride;
index_t sfb_batch_stride;
index_t o_batch_stride;
index_t a_row_stride;
index_t b_row_stride;
index_t sfa_row_stride;
index_t sfb_row_stride;
index_t o_row_stride;
};
static constexpr int BLOCK_SIZE = 128; // 128
__device__ __forceinline__ void load_block_16x2fp4(
const __nv_fp4x2_e2m1* rowA,
const __nv_fp4x2_e2m1* vecB,
const uint16_t* rowS_u16,
const uint16_t* vecS_u16,
int elem_base,
int block_base,
uint64_t (&a_regs)[2],
uint64_t (&b_regs)[2],
uint16_t &sfa_regs,
uint16_t &sfb_regs)
{
uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);
asm volatile(
"ld.global.u64.v2 {%0, %1}, [%4];\n\t"
"ld.global.u64.v2 {%2, %3}, [%5];\n\t"
: "=l"(a_regs[0]), "=l"(a_regs[1]),
"=l"(b_regs[0]), "=l"(b_regs[1])
: "l"(rowA_addr), "l"(vecB_addr)
);
asm volatile(
"ld.global.u16 %0, [%2];\n\t"
"ld.global.u16 %1, [%3];\n\t"
: "=h"(sfa_regs), "=h"(sfb_regs)
: "l"(rowS_addr), "l"(vecS_addr)
);
}
__device__ __forceinline__ float block_scaled_fma_16x2fp4(
const uint64_t (&a_regs)[2],
const uint64_t (&b_regs)[2],
uint16_t sfa_regs,
uint16_t sfb_regs)
{
uint32_t const* a_regs_packed = reinterpret_cast<uint32_t const*>(&a_regs);
uint32_t const* b_regs_packed = reinterpret_cast<uint32_t const*>(&b_regs);
float out_f32;
asm volatile(
"{\n"
// 8 bytes of A and B at a time (reused for upper half)
".reg .b8 a0_0, a0_1, a0_2, a0_3;\n"
".reg .b8 a0_4, a0_5, a0_6, a0_7;\n"
".reg .b8 b0_0, b0_1, b0_2, b0_3;\n"
".reg .b8 b0_4, b0_5, b0_6, b0_7;\n"
// scales and accumulators
".reg .f16x2 sfa_f16x2, sfb_f16x2, sf_f16x2;\n"
".reg .f16x2 scale0_f16x2, scale1_f16x2;\n"
".reg .f16x2 accum_total, accum_group;\n"
// converted fp4 -> f16x2 (only 8 per vector kept live)
".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 .f16 lane0, lane1, result_f16;\n"
".reg .f32 result_f32;\n"
// scales
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %5;\n"
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %6;\n"
"mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
"mov.b32 {lane0, lane1}, sf_f16x2;\n"
"mov.b32 scale0_f16x2, {lane0, lane0};\n"
"mov.b32 scale1_f16x2, {lane1, lane1};\n"
"mov.b32 accum_total, 0;\n"
//----------------------------------------------------------------------
// First 8×(2×FP4) -> uses scale0
//----------------------------------------------------------------------
"mov.b32 {a0_0, a0_1, a0_2, a0_3}, %1;\n"
"mov.b32 {a0_4, a0_5, a0_6, a0_7}, %2;\n"
"mov.b32 {b0_0, b0_1, b0_2, b0_3}, %3;\n"
"mov.b32 {b0_4, b0_5, b0_6, b0_7}, %4;\n"
// all conversions first
"cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_7;\n"
// then all FMAs into one group accumulator
"mov.b32 accum_group, 0;\n"
"fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"
"mul.rn.f16x2 accum_group, scale0_f16x2, accum_group;\n"
"add.rn.f16x2 accum_total, accum_total, accum_group;\n"
//----------------------------------------------------------------------
// Second 8×(2×FP4) -> uses scale1, reusing all the same regs
//----------------------------------------------------------------------
"mov.b32 {a0_0, a0_1, a0_2, a0_3}, %7;\n"
"mov.b32 {a0_4, a0_5, a0_6, a0_7}, %8;\n"
"mov.b32 {b0_0, b0_1, b0_2, b0_3}, %9;\n"
"mov.b32 {b0_4, b0_5, b0_6, b0_7}, %10;\n"
// conversions for upper half
"cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_7;\n"
"mov.b32 accum_group, 0;\n"
"fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"
"fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"
"mul.rn.f16x2 accum_group, scale1_f16x2, accum_group;\n"
"add.rn.f16x2 accum_total, accum_total, accum_group;\n"
// final reduction to scalar f16 -> then upconvert to f32
"mov.b32 {lane0, lane1}, accum_total;\n"
"add.rn.f16 result_f16, lane0, lane1;\n"
"cvt.f32.f16 result_f32, result_f16;\n"
"mov.b32 %0, result_f32;\n"
"}\n"
: "=f"(out_f32)
: "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
"r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
"h"(sfa_regs), "h"(sfb_regs),
"r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
"r"(b_regs_packed[2]), "r"(b_regs_packed[3])
: "memory"
);
return out_f32;
}
template <int ROWS_PER_BLOCK, int THREADS_PER_ROW>
__global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)
gemv_kernel_shared(const __grid_constant__ Gemv_params params)
{
const int tid = threadIdx.x;
const int rib = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int batch = blockIdx.z;
const int row = blockIdx.x * ROWS_PER_BLOCK + rib;
const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;
const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;
const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;
const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;
const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;
const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);
const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);
float sum = 0.f;
int iters = params.k / (THREADS_PER_ROW * 16);
for (int idx = 0; idx < iters; ++idx) {
int block_base = idx * THREADS_PER_ROW + lane;
int elem_base = block_base * 16;
uint64_t a_regs[2], b_regs[2];
uint16_t sfa_regs, sfb_regs;
load_block_16x2fp4(
rowA, vecB,
rowS_u16, vecS_u16,
elem_base, block_base,
a_regs, b_regs,
sfa_regs, sfb_regs);
sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);;
}
__shared__ float sdata[THREADS_PER_ROW];
sdata[lane] = sum;
__syncthreads();
if (tid < 64) sdata[lane] += sdata[lane+64];
__syncthreads();
if (lane < 32) {
float val = sdata[lane] + sdata[lane + 32];
val += __shfl_down_sync(0xffffffff, val, 16);
val += __shfl_down_sync(0xffffffff, val, 8);
val += __shfl_down_sync(0xffffffff, val, 4);
val += __shfl_down_sync(0xffffffff, val, 2);
val += __shfl_down_sync(0xffffffff, val, 1);
if (lane == 0) {
__half* out = (__half*)params.o_ptr + C_batch_base + row;
out[0] = __float2half(val);
}
}
}
template <int ROWS_PER_BLOCK, int THREADS_PER_ROW>
__global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)
gemv_kernel(const __grid_constant__ Gemv_params params)
{
const int tid = threadIdx.x;
const int rib = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int batch = blockIdx.z;
const int row = blockIdx.x * ROWS_PER_BLOCK + rib;
const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;
const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;
const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;
const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;
const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;
const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);
const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);
float sum = 0.f;
int iters = params.k / (THREADS_PER_ROW * 16);
for (int idx = 0; idx < iters; ++idx) {
int block_base = idx * THREADS_PER_ROW + lane;
int elem_base = block_base * 16;
uint64_t a_regs[2], b_regs[2];
uint16_t sfa_regs, sfb_regs;
load_block_16x2fp4(
rowA, vecB,
rowS_u16, vecS_u16,
elem_base, block_base,
a_regs, b_regs,
sfa_regs, sfb_regs);
sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);;
}
#pragma unroll
for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
sum += __shfl_down_sync(0xffffffffu, sum, offset, THREADS_PER_ROW);
}
if (lane == 0) {
__half* out = (__half*)params.o_ptr + C_batch_base + row;
out[0] = __float2half(sum);
}
}
torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB)
{
const auto sizes = A.sizes();
const int M = sizes[0];
const int K = sizes[1];
const int L = sizes[2];
Gemv_params params{};
params.b = L;
params.m = M;
params.k = K;
params.a_ptr = A.data_ptr();
params.b_ptr = B.data_ptr();
params.sfa_ptr= SFA.data_ptr();
params.sfb_ptr= SFB.data_ptr();
params.o_ptr = C.data_ptr();
params.a_batch_stride = A.stride(2);
params.b_batch_stride = B.stride(2);
params.sfa_batch_stride= SFA.stride(2);
params.sfb_batch_stride= SFB.stride(2);
params.o_batch_stride = C.stride(2);
params.a_row_stride = A.stride(0);
params.b_row_stride = B.stride(0);
params.sfa_row_stride= SFA.stride(0);
params.sfb_row_stride= SFB.stride(0);
params.o_row_stride = C.stride(0);
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (params.k <= 256) { // <= 512 FP4 values
dim3 grid(params.m / 16, 1, params.b);
dim3 block(128, 1, 1);
gemv_kernel<16, 8><<<grid, block, 0, stream>>>(params);
} else if (params.k == 3584) {
dim3 block(32, 1, 1);
dim3 grid(params.m / 4, 1, params.b);
gemv_kernel<4, 8><<<grid, block, 0, stream>>>(params);
} else if (params.k == 8192) {
dim3 block(128, 1, 1);
dim3 grid(params.m, 1, params.b);
gemv_kernel_shared<1, 128><<<grid, block>>>(params);
} else if (params.k == 1024) {
dim3 block(64, 1, 1);
dim3 grid(params.m / 8, 1, params.b);
gemv_kernel<8, 8><<<grid, block, 0, stream>>>(params);
} else {
dim3 block(128, 1, 1);
dim3 grid(params.m / 8, 1, params.b);
gemv_kernel<8, 16><<<grid, block, 0, stream>>>(params);
}
return C;
}
"""
# ---- build the module ----
nvfp4_module = load_inline(
name="nvfp4_gemv",
cpp_sources=[gemv_cpp],
cuda_sources=[gemv_cuda],
functions=["cuda_nvfp4_gemv"], # this exposes the function to Python
extra_cuda_cflags=[
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",
"--ptxas-options=--gpu-name=sm_100a",
"-O3",
"-w",
"--use_fast_math",
"-allow-unsupported-compiler",
],
extra_ldflags=["-lcuda", "-lcublas"],
verbose=True,
)
def custom_kernel(data: input_t) -> output_t:
return nvfp4_module.cuda_nvfp4_gemv(data[0], data[1], data[6], data[2], data[3])
scrolls · 451 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 80484.
⋯ 5 unchanged linesgemv_cpp = r"""#include <torch/extension.h>- // Single Python-visible entry point that dispatches between kernels- torch::Tensor nvfp4_gemv_dispatch(torch::Tensor A,- torch::Tensor B,- torch::Tensor C,- torch::Tensor SFA,- torch::Tensor SFB);+ // Forward declaration so PyTorch can bind it (definition is in the CUDA source).+ torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,+ torch::Tensor B,+ torch::Tensor C,+ torch::Tensor SFA,+ torch::Tensor SFB);"""- # ---- CUDA source: params struct, both kernels, launchers, and Python-facing wrapper ----+ # ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----gemv_cuda = r"""#include <assert.h>#include <cuda.h>⋯ 8 unchanged lines#include <cuda_bf16.h>#include <cuda_fp8.h>- // ---- gemv.h-ish: params struct ----+ // ---- gemv.h ----struct Gemv_params {using index_t = uint64_t;⋯ 18 unchanged linesindex_t o_row_stride;};- static constexpr int ROWS_PER_BLOCK = 8;- static constexpr int THREADS_PER_ROW = 16;- static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128- static constexpr int BUFFER = 2;- static constexpr int STAGES = 4;+ static constexpr int BLOCK_SIZE = 128; // 128- template <int LoadBytes>- __device__ __forceinline__ void cp_async_ca(uint32_t smem_addr, void const *global_ptr, bool pred_guard=true) {- asm volatile(- "{\n"- " .reg .pred p;\n"- " setp.ne.b32 p, %0, 0;\n"- " @p cp.async.ca.shared.global.L2::128B [%1], [%2], %3;\n"- "}\n"- :- :"r"((int)pred_guard), "r"(smem_addr), "l"(global_ptr), "n"(LoadBytes)- );- }-- struct SharedStorage {- alignas(128) __nv_fp4x2_e2m1 A[BUFFER][STAGES][128 * 16];- alignas(128) __nv_fp4x2_e2m1 B[BUFFER][STAGES][THREADS_PER_ROW * 16];- alignas(128) __nv_fp8_e4m3 SFA[BUFFER][STAGES][128 * 2];- alignas(128) __nv_fp8_e4m3 SFB[BUFFER][STAGES][THREADS_PER_ROW * 2];- };--- __device__ __forceinline__ void cpasync_load(- int buffer_idx,- int &global_k, // should be in 2xfp4- int smem_offset_A,- int smem_offset_B,- int smem_sf_write_offset,- int lane,- bool is_even,- bool load_b,- bool valid,- const __nv_fp4x2_e2m1* rowA,- const __nv_fp4x2_e2m1* vecB,- const __nv_fp8_e4m3* rowS,- const __nv_fp8_e4m3* vecS,- SharedStorage& smem)- {- for (int s = 0; s < STAGES; ++s) {- int SF_offset = (global_k >> 3) + lane * 2;-- uint32_t smem_addr_A = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[buffer_idx][s][smem_offset_A]));- uint32_t smem_addr_B = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[buffer_idx][s][smem_offset_B]));- uint32_t smem_addr_SFA = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[buffer_idx][s][smem_sf_write_offset]));- uint32_t smem_addr_SFB = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[buffer_idx][s][(lane / 2) * 4]));-- cp_async_ca<16>(smem_addr_A, rowA + global_k, valid);- cp_async_ca<16>(smem_addr_B, vecB + global_k, valid && load_b);- cp_async_ca<4>(smem_addr_SFA, rowS + SF_offset, valid && is_even);- cp_async_ca<4>(smem_addr_SFB, vecS + SF_offset, valid && is_even && load_b);-- global_k += 256;- }- }-__device__ __forceinline__ void load_block_16x2fp4(const __nv_fp4x2_e2m1* rowA,const __nv_fp4x2_e2m1* vecB,⋯ 27 unchanged lines);}- __device__ __forceinline__ void load_fragments(- uint64_t (&a_regs)[2],- uint64_t (&b_regs)[2],- uint16_t &sfa_regs,- uint16_t &sfb_regs,- int lane,- int smem_pipe_read,- int k_block,- int smem_offset_A,- int smem_offset_B,- int smem_sf_offset,- SharedStorage& smem)- {- uint32_t smem_addr_a = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[smem_pipe_read][k_block][smem_offset_A]));- uint32_t smem_addr_b = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[smem_pipe_read][k_block][smem_offset_B]));-- asm volatile(- "ld.shared.v2.u64 {%0, %1}, [%4];\n\t"- "ld.shared.v2.u64 {%2, %3}, [%5];\n\t"- : "=l"(a_regs[0]), "=l"(a_regs[1]),- "=l"(b_regs[0]), "=l"(b_regs[1])- : "r"(smem_addr_a), "r"(smem_addr_b)- );- uint32_t smem_addr_sfa = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[smem_pipe_read][k_block][smem_sf_offset]));- uint32_t smem_addr_sfb = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[smem_pipe_read][k_block][lane * 2]));-- asm volatile(- "ld.shared.u16 %0, [%2];\n\t"- "ld.shared.u16 %1, [%3];\n\t"- : "=h"(sfa_regs), "=h"(sfb_regs)- : "r"(smem_addr_sfa), "r"(smem_addr_sfb)- );- }---- __device__ __forceinline__ __half block_scaled_fma_16x2fp4(+ __device__ __forceinline__ float block_scaled_fma_16x2fp4(const uint64_t (&a_regs)[2],const uint64_t (&b_regs)[2],uint16_t sfa_regs,⋯ 2 unchanged linesuint32_t const* a_regs_packed = reinterpret_cast<uint32_t const*>(&a_regs);uint32_t const* b_regs_packed = reinterpret_cast<uint32_t const*>(&b_regs);- uint16_t out_half_bits;+ float out_f32;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 byte0_8, byte0_9, byte0_10, byte0_11;\n"- ".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\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 byte1_8, byte1_9, byte1_10, byte1_11;\n"- ".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"+ // 8 bytes of A and B at a time (reused for upper half)+ ".reg .b8 a0_0, a0_1, a0_2, a0_3;\n"+ ".reg .b8 a0_4, a0_5, a0_6, a0_7;\n"+ ".reg .b8 b0_0, b0_1, b0_2, b0_3;\n"+ ".reg .b8 b0_4, b0_5, b0_6, b0_7;\n"- ".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"- ".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"- ".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"- ".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"+ // scales and accumulators+ ".reg .f16x2 sfa_f16x2, sfb_f16x2, sf_f16x2;\n"+ ".reg .f16x2 scale0_f16x2, scale1_f16x2;\n"+ ".reg .f16x2 accum_total, accum_group;\n"- ".reg .f16x2 sfa_f16x2;\n"- ".reg .f16x2 sfb_f16x2;\n"- ".reg .f16x2 sf_f16x2;\n"-+ // converted fp4 -> f16x2 (only 8 per vector kept live)".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_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"- ".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\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_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"- ".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"- ".reg .f16 result_f16, lane0, lane1;\n"- ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"- "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"- "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"+ ".reg .f16 lane0, lane1, result_f16;\n"+ ".reg .f32 result_f32;\n"- "mov.b32 accum_0, 0;\n"- "mov.b32 accum_1, 0;\n"- "mov.b32 accum_2, 0;\n"- "mov.b32 accum_3, 0;\n"- "mov.b32 accum_4, 0;\n"- "mov.b32 accum_5, 0;\n"- "mov.b32 accum_6, 0;\n"- "mov.b32 accum_7, 0;\n"- "mov.b32 accum_8, 0;\n"- "mov.b32 accum_9, 0;\n"- "mov.b32 accum_10, 0;\n"- "mov.b32 accum_11, 0;\n"- "mov.b32 accum_12, 0;\n"- "mov.b32 accum_13, 0;\n"- "mov.b32 accum_14, 0;\n"- "mov.b32 accum_15, 0;\n"-+ // scales+ "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %5;\n"+ "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %6;\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 scale0_f16x2, {lane0, lane0};\n"+ "mov.b32 scale1_f16x2, {lane1, lane1};\n"- "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"- "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"- "mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"- "mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %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 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"- "mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"+ "mov.b32 accum_total, 0;\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"+ //----------------------------------------------------------------------+ // First 8×(2×FP4) -> uses scale0+ //----------------------------------------------------------------------+ "mov.b32 {a0_0, a0_1, a0_2, a0_3}, %1;\n"+ "mov.b32 {a0_4, a0_5, a0_6, a0_7}, %2;\n"+ "mov.b32 {b0_0, b0_1, b0_2, b0_3}, %3;\n"+ "mov.b32 {b0_4, b0_5, b0_6, b0_7}, %4;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_8, byte0_8;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_9, byte0_9;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"- "cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"+ // all conversions first+ "cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_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"+ // then all FMAs into one group accumulator+ "mov.b32 accum_group, 0;\n"+ "fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_8, byte1_8;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_9, byte1_9;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"- "cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"+ "mul.rn.f16x2 accum_group, scale0_f16x2, accum_group;\n"+ "add.rn.f16x2 accum_total, accum_total, accum_group;\n"- "fma.rn.f16x2 accum_0, cvt_0_0, cvt_1_0, accum_0;\n"- "fma.rn.f16x2 accum_1, cvt_0_1, cvt_1_1, accum_1;\n"- "fma.rn.f16x2 accum_2, cvt_0_2, cvt_1_2, accum_2;\n"- "fma.rn.f16x2 accum_3, cvt_0_3, cvt_1_3, accum_3;\n"- "fma.rn.f16x2 accum_4, cvt_0_4, cvt_1_4, accum_4;\n"- "fma.rn.f16x2 accum_5, cvt_0_5, cvt_1_5, accum_5;\n"- "fma.rn.f16x2 accum_6, cvt_0_6, cvt_1_6, accum_6;\n"- "fma.rn.f16x2 accum_7, cvt_0_7, cvt_1_7, accum_7;\n"+ //----------------------------------------------------------------------+ // Second 8×(2×FP4) -> uses scale1, reusing all the same regs+ //----------------------------------------------------------------------+ "mov.b32 {a0_0, a0_1, a0_2, a0_3}, %7;\n"+ "mov.b32 {a0_4, a0_5, a0_6, a0_7}, %8;\n"+ "mov.b32 {b0_0, b0_1, b0_2, b0_3}, %9;\n"+ "mov.b32 {b0_4, b0_5, b0_6, b0_7}, %10;\n"- "fma.rn.f16x2 accum_8, cvt_0_8, cvt_1_8, accum_8;\n"- "fma.rn.f16x2 accum_9, cvt_0_9, cvt_1_9, accum_9;\n"- "fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"- "fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"- "fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"- "fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"- "fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"- "fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"+ // conversions for upper half+ "cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_7;\n"- "add.rn.f16x2 accum_0, accum_0, accum_1;\n"- "add.rn.f16x2 accum_2, accum_2, accum_3;\n"- "add.rn.f16x2 accum_4, accum_4, accum_5;\n"- "add.rn.f16x2 accum_6, accum_6, accum_7;\n"- "add.rn.f16x2 accum_8, accum_8, accum_9;\n"- "add.rn.f16x2 accum_10, accum_10, accum_11;\n"- "add.rn.f16x2 accum_12, accum_12, accum_13;\n"- "add.rn.f16x2 accum_14, accum_14, accum_15;\n"+ "mov.b32 accum_group, 0;\n"+ "fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"+ "fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"- "add.rn.f16x2 accum_0, accum_0, accum_2;\n"- "add.rn.f16x2 accum_4, accum_4, accum_6;\n"- "add.rn.f16x2 accum_8, accum_8, accum_10;\n"- "add.rn.f16x2 accum_12, accum_12, accum_14;\n"+ "mul.rn.f16x2 accum_group, scale1_f16x2, accum_group;\n"+ "add.rn.f16x2 accum_total, accum_total, accum_group;\n"- "add.rn.f16x2 accum_0, accum_0, accum_4;\n"- "add.rn.f16x2 accum_8, accum_8, accum_12;\n"-- "mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"- "mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"-- "add.rn.f16x2 accum_0, accum_0, accum_8;\n"-- "mov.b32 {lane0, lane1}, accum_0;\n"+ // final reduction to scalar f16 -> then upconvert to f32+ "mov.b32 {lane0, lane1}, accum_total;\n""add.rn.f16 result_f16, lane0, lane1;\n"+ "cvt.f32.f16 result_f32, result_f16;\n"+ "mov.b32 %0, result_f32;\n"- "mov.b16 %0, result_f16;\n""}\n"- : "=h"(out_half_bits)+ : "=f"(out_f32): "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),- "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),"r"(b_regs_packed[0]), "r"(b_regs_packed[1]),- "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),- "h"(sfa_regs), "h"(sfb_regs)+ "h"(sfa_regs), "h"(sfb_regs),+ "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),+ "r"(b_regs_packed[2]), "r"(b_regs_packed[3]): "memory");- union { uint16_t u; __half h; } conv;- conv.u = out_half_bits;- return conv.h;+ return out_f32;}-- __global__ void __launch_bounds__(BLOCK_SIZE, 8)- gemv_kernel_fast(const __grid_constant__ Gemv_params params)+ template <int ROWS_PER_BLOCK, int THREADS_PER_ROW>+ __global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)+ gemv_kernel_shared(const __grid_constant__ Gemv_params params){const int tid = threadIdx.x;const int rib = tid / THREADS_PER_ROW;⋯ 1 unchanged linesconst int batch = blockIdx.z;const int row = blockIdx.x * ROWS_PER_BLOCK + rib;- extern __shared__ __align__(128) uint8_t shared_storage[];- SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;⋯ 1 unchanged linesconst size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;+const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;- rowA += lane * 16;- vecB += lane * 16;- const int k = params.k;+ const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);+ const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);+float sum = 0.f;- int global_k = 0;- const int smem_offset_A = tid * 16;- const int smem_offset_B = lane * 16;- const int smem_sf_write_offset = (tid / 2) * 4;- const int smem_sf_offset = tid * 2;- const bool load_b = rib == 0;- const bool is_even = (lane % 2 == 0);- // prefetch to buffer 0- cpasync_load(- 0,- global_k,- smem_offset_A,- smem_offset_B,- smem_sf_write_offset,- lane,- is_even,- load_b,- true,- rowA,- vecB,- rowS,- vecS,- smem- );- asm volatile("cp.async.commit_group;\n" ::);- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();- uint64_t a_regs[2][2], b_regs[2][2];- uint16_t sfa_regs[2], sfb_regs[2];-- int smem_pipe_read = 0;- int smem_pipe_write = 1;-- // prefetch buffer 0 stage 0 to registers 0- load_fragments(- a_regs[0],- b_regs[0],- sfa_regs[0],- sfb_regs[0],- lane,- 0, //smem_pipe_read- 0, //k_block- smem_offset_A,- smem_offset_B,- smem_sf_offset,- smem- );- asm volatile("cp.async.commit_group;\n" ::);-int iters = params.k / (THREADS_PER_ROW * 16);- int idx = 0;- while (idx < iters) {- int smem_pipe_read_curr = smem_pipe_read;- for (int k_block = 0; k_block < STAGES; ++k_block)- {- if (k_block == STAGES-1)- {- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();+ for (int idx = 0; idx < iters; ++idx) {+ int block_base = idx * THREADS_PER_ROW + lane;+ int elem_base = block_base * 16;- smem_pipe_read_curr = smem_pipe_read;- }- auto k_block_next = (k_block + 1) % STAGES;- int frag_idx_next = (k_block + 1) & 1;+ uint64_t a_regs[2], b_regs[2];+ uint16_t sfa_regs, sfb_regs;- load_fragments(- a_regs[frag_idx_next],- b_regs[frag_idx_next],- sfa_regs[frag_idx_next],- sfb_regs[frag_idx_next],- lane,- smem_pipe_read_curr, //smem_pipe_read- k_block_next, //k_block- smem_offset_A,- smem_offset_B,- smem_sf_offset,- smem- );- if (k_block == 0)- {- bool valid = (global_k < k);- cpasync_load(- smem_pipe_write,- global_k,- smem_offset_A,- smem_offset_B,- smem_sf_write_offset,- lane,- is_even,- load_b,- valid,- rowA,- vecB,- rowS,- vecS,- smem);- asm volatile("cp.async.commit_group;\n" ::);-- smem_pipe_write = smem_pipe_read;- smem_pipe_read = (smem_pipe_read + 1) & 1;- }--- int frag_idx = k_block & 1;- __half h = block_scaled_fma_16x2fp4(- a_regs[frag_idx],- b_regs[frag_idx],- sfa_regs[frag_idx],- sfb_regs[frag_idx]);- sum += __half2float(h);-- }- idx += STAGES;+ load_block_16x2fp4(+ rowA, vecB,+ rowS_u16, vecS_u16,+ elem_base, block_base,+ a_regs, b_regs,+ sfa_regs, sfb_regs);+ sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);;}- asm volatile("cp.async.wait_all;\n" ::);++ __shared__ float sdata[THREADS_PER_ROW];+ sdata[lane] = sum;__syncthreads();- unsigned mask = 0xffffffffu;- sum += __shfl_down_sync(mask, sum, 8, 16);- sum += __shfl_down_sync(mask, sum, 4, 16);- sum += __shfl_down_sync(mask, sum, 2, 16);- sum += __shfl_down_sync(mask, sum, 1, 16);+ if (tid < 64) sdata[lane] += sdata[lane+64];+ __syncthreads();- if (lane == 0) {- __half* out = (__half*)params.o_ptr + C_batch_base + row;- out[0] = __float2half(sum);- }- }- static inline void launch_kernel_fast(Gemv_params ¶ms, cudaStream_t stream)- {- const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;- size_t smem_size = sizeof(SharedStorage);- dim3 grid(grid_x, 1, params.b);- dim3 block(BLOCK_SIZE, 1, 1);- gemv_kernel_fast<<<grid, block, smem_size>>>(params);- }-- __global__ void __launch_bounds__(BLOCK_SIZE, 8)- gemv_kernel_k1024(const __grid_constant__ Gemv_params params)- {- const int tid = threadIdx.x;- const int rib = tid / THREADS_PER_ROW;- const int lane = tid % THREADS_PER_ROW;- const int batch = blockIdx.z;- const int row = blockIdx.x * ROWS_PER_BLOCK + rib;-- const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;- const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;- const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;- const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;- const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;-- const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;- const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;- const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;- const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;-- const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);- const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);-- float sum0 = 0.f;- float sum1 = 0.f;-- // we know iters == 4, so two iterations here;- // then force unroll to get all 4 “stages” laid out- #pragma unroll- for (int idx = 0; idx < 4; idx += 2) {- // Stage idx- {- int block_base = (idx + 0) * THREADS_PER_ROW + lane;- int elem_base = block_base * 16;-- uint64_t a_regs0[2], b_regs0[2];- uint16_t sfa_regs0, sfb_regs0;-- load_block_16x2fp4(- rowA, vecB,- rowS_u16, vecS_u16,- elem_base, block_base,- a_regs0, b_regs0,- sfa_regs0, sfb_regs0);-- __half h0 = block_scaled_fma_16x2fp4(a_regs0, b_regs0, sfa_regs0, sfb_regs0);- sum0 += __half2float(h0);+ if (lane < 32) {+ float val = sdata[lane] + sdata[lane + 32];+ val += __shfl_down_sync(0xffffffff, val, 16);+ val += __shfl_down_sync(0xffffffff, val, 8);+ val += __shfl_down_sync(0xffffffff, val, 4);+ val += __shfl_down_sync(0xffffffff, val, 2);+ val += __shfl_down_sync(0xffffffff, val, 1);++ if (lane == 0) {+ __half* out = (__half*)params.o_ptr + C_batch_base + row;+ out[0] = __float2half(val);}-- // Stage idx + 1- {- int block_base = (idx + 1) * THREADS_PER_ROW + lane;- int elem_base = block_base * 16;-- uint64_t a_regs1[2], b_regs1[2];- uint16_t sfa_regs1, sfb_regs1;-- load_block_16x2fp4(- rowA, vecB,- rowS_u16, vecS_u16,- elem_base, block_base,- a_regs1, b_regs1,- sfa_regs1, sfb_regs1);-- __half h1 = block_scaled_fma_16x2fp4(a_regs1, b_regs1, sfa_regs1, sfb_regs1);- sum1 += __half2float(h1);- }}-- float sum = sum0 + sum1;---- unsigned mask = 0xffffffffu;- sum += __shfl_down_sync(mask, sum, 8, 16);- sum += __shfl_down_sync(mask, sum, 4, 16);- sum += __shfl_down_sync(mask, sum, 2, 16);- sum += __shfl_down_sync(mask, sum, 1, 16);-- if (lane == 0) {- __half* out = (__half*)params.o_ptr + C_batch_base + row;- out[0] = __float2half(sum);- }}- static inline void launch_kernel_k1024(Gemv_params ¶ms, cudaStream_t stream)+ template <int ROWS_PER_BLOCK, int THREADS_PER_ROW>+ __global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)+ gemv_kernel(const __grid_constant__ Gemv_params params){- const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;- size_t smem_size = sizeof(SharedStorage);- dim3 grid(grid_x, 1, params.b);- dim3 block(BLOCK_SIZE, 1, 1);- gemv_kernel_k1024<<<grid, block, 0, stream>>>(params);- }-- // ============================================================================- // K=3584-specialized kernel- // ============================================================================-- __global__ void __launch_bounds__(BLOCK_SIZE, 8)- gemv_kernel_k3584(const __grid_constant__ Gemv_params params)- {const int tid = threadIdx.x;const int rib = tid / THREADS_PER_ROW;const int lane = tid % THREADS_PER_ROW;const int batch = blockIdx.z;const int row = blockIdx.x * ROWS_PER_BLOCK + rib;- extern __shared__ __align__(128) uint8_t shared_storage[];- SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);-const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;⋯ 5 unchanged linesconst __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;- rowA += lane * 16;- vecB += lane * 16;+ const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);+ const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);- const int k = params.k;-float sum = 0.f;- int global_k = 0;- const int smem_offset_A = tid * 16;- const int smem_offset_B = lane * 16;- const int smem_sf_write_offset = (tid / 2) * 4;- const int smem_sf_offset = tid * 2;- const bool load_b = rib == 0;- const bool is_even = (lane % 2 == 0);- // prefetch to buffer 0- cpasync_load(- 0,- global_k,- smem_offset_A,- smem_offset_B,- smem_sf_write_offset,- lane,- is_even,- load_b,- true,- rowA,- vecB,- rowS,- vecS,- smem- );- asm volatile("cp.async.commit_group;\n" ::);- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();+ int iters = params.k / (THREADS_PER_ROW * 16);- uint64_t a_regs[2][2], b_regs[2][2];- uint16_t sfa_regs[2], sfb_regs[2];+ for (int idx = 0; idx < iters; ++idx) {+ int block_base = idx * THREADS_PER_ROW + lane;+ int elem_base = block_base * 16;- int smem_pipe_read = 0;- int smem_pipe_write = 1;+ uint64_t a_regs[2], b_regs[2];+ uint16_t sfa_regs, sfb_regs;- // prefetch buffer 0 stage 0 to registers 0- load_fragments(- a_regs[0],- b_regs[0],- sfa_regs[0],- sfb_regs[0],- lane,- 0, //smem_pipe_read- 0, //k_block- smem_offset_A,- smem_offset_B,- smem_sf_offset,- smem- );- asm volatile("cp.async.commit_group;\n" ::);-- int total_stages = params.k / (THREADS_PER_ROW * 16);- int full_iters = total_stages / STAGES; // Number of complete STAGES-groups- int tail_stages = total_stages % STAGES; // Remaining stages-- // Main loop: process complete groups of STAGES- for (int idx = 0; idx < full_iters; ++idx) {- int smem_pipe_read_curr = smem_pipe_read;-- for (int k_block = 0; k_block < STAGES; ++k_block)- {- if (k_block == STAGES-1)- {- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();-- smem_pipe_read_curr = smem_pipe_read;- }-- auto k_block_next = (k_block + 1) % STAGES;- int frag_idx_next = (k_block + 1) & 1;-- load_fragments(- a_regs[frag_idx_next],- b_regs[frag_idx_next],- sfa_regs[frag_idx_next],- sfb_regs[frag_idx_next],- lane,- smem_pipe_read_curr,- k_block_next,- smem_offset_A,- smem_offset_B,- smem_sf_offset,- smem- );-- if (k_block == 0)- {- bool valid = (global_k < k);- cpasync_load(- smem_pipe_write,- global_k,- smem_offset_A,- smem_offset_B,- smem_sf_write_offset,- lane,- is_even,- load_b,- valid,- rowA,- vecB,- rowS,- vecS,- smem);- asm volatile("cp.async.commit_group;\n" ::);-- smem_pipe_write = smem_pipe_read;- smem_pipe_read = (smem_pipe_read + 1) & 1;- }-- int frag_idx = k_block & 1;- __half h = block_scaled_fma_16x2fp4(- a_regs[frag_idx],- b_regs[frag_idx],- sfa_regs[frag_idx],- sfb_regs[frag_idx]);- sum += __half2float(h);- }+ load_block_16x2fp4(+ rowA, vecB,+ rowS_u16, vecS_u16,+ elem_base, block_base,+ a_regs, b_regs,+ sfa_regs, sfb_regs);+ sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);;}- if (tail_stages > 0) {- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();-- for (int k_block = 0; k_block < tail_stages; ++k_block)- {- int frag_idx = k_block & 1;- int frag_idx_next = (k_block + 1) & 1;-- if (k_block < tail_stages - 1) {- load_fragments(- a_regs[frag_idx_next],- b_regs[frag_idx_next],- sfa_regs[frag_idx_next],- sfb_regs[frag_idx_next],- lane,- smem_pipe_read,- k_block + 1,- smem_offset_A,- smem_offset_B,- smem_sf_offset,- smem- );- }-- __half h = block_scaled_fma_16x2fp4(- a_regs[frag_idx],- b_regs[frag_idx],- sfa_regs[frag_idx],- sfb_regs[frag_idx]);- sum += __half2float(h);- }+ #pragma unroll+ for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {+ sum += __shfl_down_sync(0xffffffffu, sum, offset, THREADS_PER_ROW);}- asm volatile("cp.async.wait_all;\n" ::);- __syncthreads();--- unsigned mask = 0xffffffffu;- sum += __shfl_down_sync(mask, sum, 8, 16);- sum += __shfl_down_sync(mask, sum, 4, 16);- sum += __shfl_down_sync(mask, sum, 2, 16);- sum += __shfl_down_sync(mask, sum, 1, 16);-if (lane == 0) {__half* out = (__half*)params.o_ptr + C_batch_base + row;out[0] = __float2half(sum);}}- static inline void launch_kernel_k3584(Gemv_params ¶ms, cudaStream_t stream)- {- const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;- size_t smem_size = sizeof(SharedStorage);- dim3 grid(grid_x, 1, params.b);- dim3 block(BLOCK_SIZE, 1, 1);- gemv_kernel_k3584<<<grid, block, smem_size>>>(params);- }- // ============================================================================- // K=256-specialized kernel (your previous v1) -> gemv_kernel_k256 / launch_kernel_k256- // ============================================================================-- __global__ void __launch_bounds__(BLOCK_SIZE, 8)- gemv_kernel_k256(const __grid_constant__ Gemv_params params)+ torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,+ torch::Tensor B,+ torch::Tensor C,+ torch::Tensor SFA,+ torch::Tensor SFB){- const int tid = threadIdx.x;- const int rib = tid / THREADS_PER_ROW;- const int lane = tid % THREADS_PER_ROW;- const int batch = blockIdx.z;- const int row = blockIdx.x * ROWS_PER_BLOCK + rib;- const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;- const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;- const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;- const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;- const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;+ const auto sizes = A.sizes();+ const int M = sizes[0];+ const int K = sizes[1];+ const int L = sizes[2];- const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;- const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;-- const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;- const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;-- float sum = 0.f;-- // Each thread does 1 16 group FP4 or (8 2xFP4)- for (int idx = 0; idx < params.k / THREADS_PER_ROW / 8; ++idx) {- int base = idx * 16;- const int base_id = (idx * THREADS_PER_ROW + lane) * 8;-- __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);- __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);- __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);- __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);- __half scale = __hmul(sfa, sfb);-- __half2 acc = __float2half2_rn(0.0f);-- #pragma unroll- for (int i = 0; i < 8; ++i) { // go over each individual 2xFP4- const int id = base_id + i;- __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);- __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);- __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);- __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);-- const __half2 a_h2 = __half2(a_raw);- const __half2 b_h2 = __half2(b_raw);-- acc = __hfma2(a_h2, b_h2, acc);- }-- __half fin = __hadd(__low2half(acc), __high2half(acc));- __half h = __hmul(fin, scale);-- sum += __half2float(h);- }-- // Reduce within the 16-thread subgroup (one output row)- unsigned mask = 0xffffffffu;- sum += __shfl_down_sync(mask, sum, 8, 16);- sum += __shfl_down_sync(mask, sum, 4, 16);- sum += __shfl_down_sync(mask, sum, 2, 16);- sum += __shfl_down_sync(mask, sum, 1, 16);-- if (lane == 0) {- __half* out = (__half*)params.o_ptr + C_batch_base + row;- out[0] = __float2half(sum);- }- }-- static inline void launch_kernel_k256(Gemv_params ¶ms, cudaStream_t stream)- {- const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;- dim3 grid(grid_x, 1, params.b);- dim3 block(BLOCK_SIZE, 1, 1);- gemv_kernel_k256<<<grid, block, 0, stream>>>(params);- }-- // ============================================================================- // Python-facing function: sets up params and dispatches based on K- // ============================================================================-- torch::Tensor nvfp4_gemv_dispatch(torch::Tensor A,- torch::Tensor B,- torch::Tensor C,- torch::Tensor SFA,- torch::Tensor SFB)- {- auto sizes = A.sizes();- const int64_t M = sizes[0];- const int64_t K = sizes[1];- const int64_t L = sizes[2];-Gemv_params params{};- params.b = static_cast<int>(L);- params.m = static_cast<int>(M);- params.k = static_cast<int>(K);- params.real_k = static_cast<int>(K * 2);+ params.b = L;+ params.m = M;+ params.k = K;- params.a_ptr = A.data_ptr();- params.b_ptr = B.data_ptr();- params.sfa_ptr = SFA.data_ptr();- params.sfb_ptr = SFB.data_ptr();- params.o_ptr = C.data_ptr();+ params.a_ptr = A.data_ptr();+ params.b_ptr = B.data_ptr();+ params.sfa_ptr= SFA.data_ptr();+ params.sfb_ptr= SFB.data_ptr();+ params.o_ptr = C.data_ptr();- params.a_batch_stride = static_cast<uint64_t>(A.stride(2));- params.b_batch_stride = static_cast<uint64_t>(B.stride(2));- params.sfa_batch_stride = static_cast<uint64_t>(SFA.stride(2));- params.sfb_batch_stride = static_cast<uint64_t>(SFB.stride(2));- params.o_batch_stride = static_cast<uint64_t>(C.stride(2));+ params.a_batch_stride = A.stride(2);+ params.b_batch_stride = B.stride(2);+ params.sfa_batch_stride= SFA.stride(2);+ params.sfb_batch_stride= SFB.stride(2);+ params.o_batch_stride = C.stride(2);- params.a_row_stride = static_cast<uint64_t>(A.stride(0));- params.b_row_stride = static_cast<uint64_t>(B.stride(0));- params.sfa_row_stride = static_cast<uint64_t>(SFA.stride(0));- params.sfb_row_stride = static_cast<uint64_t>(SFB.stride(0));- params.o_row_stride = static_cast<uint64_t>(C.stride(0));+ params.a_row_stride = A.stride(0);+ params.b_row_stride = B.stride(0);+ params.sfa_row_stride= SFA.stride(0);+ params.sfb_row_stride= SFB.stride(0);+ params.o_row_stride = C.stride(0);auto stream = at::cuda::getCurrentCUDAStream().stream();-- // Tiny host-side dispatch: effectively zero overhead vs kernel time- if (params.k < 512) {- launch_kernel_k256(params, stream);- }- else if (params.k == 1024) {- launch_kernel_k1024(params, stream);+ if (params.k <= 256) { // <= 512 FP4 values+ dim3 grid(params.m / 16, 1, params.b);+ dim3 block(128, 1, 1);+ gemv_kernel<16, 8><<<grid, block, 0, stream>>>(params);+ } else if (params.k == 3584) {+ dim3 block(32, 1, 1);+ dim3 grid(params.m / 4, 1, params.b);+ gemv_kernel<4, 8><<<grid, block, 0, stream>>>(params);+ } else if (params.k == 8192) {+ dim3 block(128, 1, 1);+ dim3 grid(params.m, 1, params.b);+ gemv_kernel_shared<1, 128><<<grid, block>>>(params);+ } else if (params.k == 1024) {+ dim3 block(64, 1, 1);+ dim3 grid(params.m / 8, 1, params.b);+ gemv_kernel<8, 8><<<grid, block, 0, stream>>>(params);+ } else {+ dim3 block(128, 1, 1);+ dim3 grid(params.m / 8, 1, params.b);+ gemv_kernel<8, 16><<<grid, block, 0, stream>>>(params);}- else if (params.k % 1024) {- launch_kernel_k3584(params, stream);- }- else {- launch_kernel_fast(params, stream);- }return C;}⋯ 4 unchanged linesname="nvfp4_gemv",cpp_sources=[gemv_cpp],cuda_sources=[gemv_cuda],- functions=["nvfp4_gemv_dispatch"], # single Python-visible entry point+ functions=["cuda_nvfp4_gemv"], # this exposes the function to Pythonextra_cuda_cflags=["-std=c++17","-gencode=arch=compute_100a,code=sm_100a",⋯ 9 unchanged linesdef custom_kernel(data: input_t) -> output_t:- return nvfp4_module.nvfp4_gemv_dispatch(data[0], data[1], data[6], data[2], data[3])+ return nvfp4_module.cuda_nvfp4_gemv(data[0], data[1], data[6], data[2], data[3])
scrolls · 1145 diff lines total
Best evidence level for this revision: reported
JSON