submission 106593
tomaszki · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1092 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-106593?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:32260cd0483b8f932c13ba3badabb23bdda66196d3ffc5be1501a102e3d6d983
license declaredunknown
license concludedunknown
authorstomaszki
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled GEMV.fp8
const __nv_fp8_e4m3* __restrict__ sfa,shared-memory
extern __shared__ unsigned char shared_storage[];vector-width = int4
reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];Kernel source
submission.py1092 lines
#!POPCORN leaderboard nvfp4_gemv
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# Kernel configuration parameters
sf_vec_size = 16
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
# Helper function to convert scale factor tensor to blocked format
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
# Please ensure rows and cols are multiples of 128 and 4 respectively
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def naive_pytorch(data: input_t) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data
# Get dimensions from MxNxL layout
_, _, l = c_ref.shape
# Call torch._scaled_mm to compute the GEMV result
for l_idx in range(l):
# Convert the scale factor tensor to blocked format
scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])
# (m, k) @ (n, k).T -> (m, n)
res = torch._scaled_mm(
a_ref[:, :, l_idx],
b_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(),
scale_b.cuda(),
bias=None,
out_dtype=torch.float16,
)
c_ref[:, 0, l_idx] = res[:, 0]
return c_ref
# CUDA SOURCE CODE
cuda_source = """
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#define FULL_MASK 0xffffffff
__global__ void gemv_kernel_4096_7168(
const __nv_fp4x2_storage_t* __restrict__ a,
const __nv_fp4x2_storage_t* __restrict__ b,
const __nv_fp8_e4m3* __restrict__ sfa,
const __nv_fp8_e4m3* __restrict__ sfb,
__half* __restrict__ c
) {
const int M = 4096;
const int K = 7168;
extern __shared__ unsigned char shared_storage[];
auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);
auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));
__shared__ __half c_shared[32];
b += blockIdx.y * (K / 2) * 128;
sfb += blockIdx.y * (K / 16) * 128;
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];
}
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];
}
__syncthreads();
// Each warp computes one result and saves it to shared memory
int result_0 = 0;
int result_1 = 0;
int result_2 = 0;
int result_3 = 0;
int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
a += offset;
sfa += offset / 8;
for (int i = threadIdx.x; i < K / 32; i += 32) {
int4 a_packed = reinterpret_cast<const int4*>(a)[i];
int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];
__nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
__nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
asm volatile( \\
"{\\n" \\
// declare registers for A / B tensors
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\
".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\
".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\
".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\
".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\
// declare registers for accumulators
".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\
".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\
".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\
".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\
// declare registers for scaling factors
".reg .f16x2 sfa_f16x2;\\n" \\
".reg .f16x2 sfb_f16x2;\\n" \\
".reg .f16x2 sf_f16x2;\\n" \\
// declare registers for conversion
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\
".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\
".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\
".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\
".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\
".reg .f16 result_f16, lane0, lane1;\\n" \\
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\
// convert scaling factors from fp8 to f16x2
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\
// clear accumulators
"mov.b32 accum_0_0, 0;\\n" \\
"mov.b32 accum_0_1, 0;\\n" \\
"mov.b32 accum_0_2, 0;\\n" \\
"mov.b32 accum_0_3, 0;\\n" \\
"mov.b32 accum_1_0, 0;\\n" \\
"mov.b32 accum_1_1, 0;\\n" \\
"mov.b32 accum_1_2, 0;\\n" \\
"mov.b32 accum_1_3, 0;\\n" \\
"mov.b32 accum_2_0, 0;\\n" \\
"mov.b32 accum_2_1, 0;\\n" \\
"mov.b32 accum_2_2, 0;\\n" \\
"mov.b32 accum_2_3, 0;\\n" \\
"mov.b32 accum_3_0, 0;\\n" \\
"mov.b32 accum_3_1, 0;\\n" \\
"mov.b32 accum_3_2, 0;\\n" \\
"mov.b32 accum_3_3, 0;\\n" \\
// multiply, unpacking and permuting scale factors
"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" \\
// unpacking A and B tensors
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\
"mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\
"mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\
"mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\
"mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\
// convert A and B tensors from fp4 to f16x2
// A[0 - 7] and B[0 - 7]
"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" \\
// A[8 - 15] and B[8 - 15]
"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" \\
// A[16 - 23] and B[16 - 23]
"cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\
// A[24 - 31] and B[24 - 31]
"cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\
// fma for A[0 - 7] and B[0 - 7]
"fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\
"fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\
"fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\
"fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\
// fma for A[8 - 15] and B[8 - 15]
"fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\
"fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\
"fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\
"fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\
// fma for A[16 - 23] and B[16 - 23]
"fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\
"fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\
"fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\
"fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\
// fma for A[24 - 31] and B[24 - 31]
"fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\
"fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\
"fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\
"fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\
// tree reduction for accumulators
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\
"add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\
"add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\
"add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\
"add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\
"fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\
"fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\
"fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\
"fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\
"fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\
"fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\
"fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\
"fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\
"}\\n"
: "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3
: "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5
"r"(a_packed.x), "r"(b_packed.x), // 6, 7
"r"(a_packed.y), "r"(b_packed.y), // 8, 9
"r"(a_packed.z), "r"(b_packed.z), // 10, 11
"r"(a_packed.w), "r"(b_packed.w) // 12, 13
);
}
// Reduce the result and store it in shared memory
__half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),
reinterpret_cast<const __half2&>(result_1));
__half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),
reinterpret_cast<const __half2&>(result_3));
reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;
for (int offset = 16; offset > 0; offset /= 2) {
final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);
}
if (threadIdx.x == 0) {
c_shared[threadIdx.y] = __float2half_rn(final_result_f);
}
__syncthreads();
// Write the result to global memory
if (threadIdx.x == 0) {
int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;
c[c_offset] = __float2half_rn(final_result_f);
}
}
__global__ void gemv_kernel_7168_2048(
const __nv_fp4x2_storage_t* __restrict__ a,
const __nv_fp4x2_storage_t* __restrict__ b,
const __nv_fp8_e4m3* __restrict__ sfa,
const __nv_fp8_e4m3* __restrict__ sfb,
__half* __restrict__ c
) {
const int M = 7168;
const int K = 2048;
extern __shared__ unsigned char shared_storage[];
auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);
auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));
__shared__ __half c_shared[32];
b += blockIdx.y * (K / 2) * 128;
sfb += blockIdx.y * (K / 16) * 128;
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];
}
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];
}
__syncthreads();
// Each warp computes one result and saves it to shared memory
int result_0 = 0;
int result_1 = 0;
int result_2 = 0;
int result_3 = 0;
int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
a += offset;
sfa += offset / 8;
for (int i = threadIdx.x; i < K / 32; i += 32) {
int4 a_packed = reinterpret_cast<const int4*>(a)[i];
int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];
__nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
__nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
asm volatile( \\
"{\\n" \\
// declare registers for A / B tensors
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\
".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\
".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\
".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\
".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\
// declare registers for accumulators
".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\
".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\
".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\
".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\
// declare registers for scaling factors
".reg .f16x2 sfa_f16x2;\\n" \\
".reg .f16x2 sfb_f16x2;\\n" \\
".reg .f16x2 sf_f16x2;\\n" \\
// declare registers for conversion
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\
".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\
".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\
".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\
".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\
".reg .f16 result_f16, lane0, lane1;\\n" \\
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\
// convert scaling factors from fp8 to f16x2
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\
// clear accumulators
"mov.b32 accum_0_0, 0;\\n" \\
"mov.b32 accum_0_1, 0;\\n" \\
"mov.b32 accum_0_2, 0;\\n" \\
"mov.b32 accum_0_3, 0;\\n" \\
"mov.b32 accum_1_0, 0;\\n" \\
"mov.b32 accum_1_1, 0;\\n" \\
"mov.b32 accum_1_2, 0;\\n" \\
"mov.b32 accum_1_3, 0;\\n" \\
"mov.b32 accum_2_0, 0;\\n" \\
"mov.b32 accum_2_1, 0;\\n" \\
"mov.b32 accum_2_2, 0;\\n" \\
"mov.b32 accum_2_3, 0;\\n" \\
"mov.b32 accum_3_0, 0;\\n" \\
"mov.b32 accum_3_1, 0;\\n" \\
"mov.b32 accum_3_2, 0;\\n" \\
"mov.b32 accum_3_3, 0;\\n" \\
// multiply, unpacking and permuting scale factors
"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" \\
// unpacking A and B tensors
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\
"mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\
"mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\
"mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\
"mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\
// convert A and B tensors from fp4 to f16x2
// A[0 - 7] and B[0 - 7]
"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" \\
// A[8 - 15] and B[8 - 15]
"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" \\
// A[16 - 23] and B[16 - 23]
"cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\
// A[24 - 31] and B[24 - 31]
"cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\
// fma for A[0 - 7] and B[0 - 7]
"fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\
"fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\
"fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\
"fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\
// fma for A[8 - 15] and B[8 - 15]
"fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\
"fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\
"fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\
"fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\
// fma for A[16 - 23] and B[16 - 23]
"fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\
"fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\
"fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\
"fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\
// fma for A[24 - 31] and B[24 - 31]
"fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\
"fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\
"fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\
"fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\
// tree reduction for accumulators
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\
"add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\
"add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\
"add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\
"add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\
"fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\
"fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\
"fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\
"fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\
"fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\
"fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\
"fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\
"fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\
"}\\n"
: "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3
: "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5
"r"(a_packed.x), "r"(b_packed.x), // 6, 7
"r"(a_packed.y), "r"(b_packed.y), // 8, 9
"r"(a_packed.z), "r"(b_packed.z), // 10, 11
"r"(a_packed.w), "r"(b_packed.w) // 12, 13
);
}
// Reduce the result and store it in shared memory
__half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),
reinterpret_cast<const __half2&>(result_1));
__half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),
reinterpret_cast<const __half2&>(result_3));
reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;
for (int offset = 16; offset > 0; offset /= 2) {
final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);
}
if (threadIdx.x == 0) {
c_shared[threadIdx.y] = __float2half_rn(final_result_f);
}
__syncthreads();
// Write the result to global memory
if (threadIdx.x == 0) {
int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;
c[c_offset] = __float2half_rn(final_result_f);
}
}
__global__ void gemv_kernel_7168_16384(
const __nv_fp4x2_storage_t* __restrict__ a,
const __nv_fp4x2_storage_t* __restrict__ b,
const __nv_fp8_e4m3* __restrict__ sfa,
const __nv_fp8_e4m3* __restrict__ sfb,
__half* __restrict__ c
) {
const int M = 7168;
const int K = 16384;
extern __shared__ unsigned char shared_storage[];
auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);
auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));
__shared__ __half c_shared[32];
b += blockIdx.y * (K / 2) * 128;
sfb += blockIdx.y * (K / 16) * 128;
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];
}
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];
}
__syncthreads();
// Each warp computes one result and saves it to shared memory
int result_0 = 0;
int result_1 = 0;
int result_2 = 0;
int result_3 = 0;
int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
a += offset;
sfa += offset / 8;
for (int i = threadIdx.x; i < K / 32; i += 32) {
int4 a_packed = reinterpret_cast<const int4*>(a)[i];
int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];
__nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
__nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
asm volatile( \\
"{\\n" \\
// declare registers for A / B tensors
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\
".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\
".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\
".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\
".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\
// declare registers for accumulators
".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\
".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\
".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\
".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\
// declare registers for scaling factors
".reg .f16x2 sfa_f16x2;\\n" \\
".reg .f16x2 sfb_f16x2;\\n" \\
".reg .f16x2 sf_f16x2;\\n" \\
// declare registers for conversion
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\
".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\
".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\
".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\
".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\
".reg .f16 result_f16, lane0, lane1;\\n" \\
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\
// convert scaling factors from fp8 to f16x2
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\
// clear accumulators
"mov.b32 accum_0_0, 0;\\n" \\
"mov.b32 accum_0_1, 0;\\n" \\
"mov.b32 accum_0_2, 0;\\n" \\
"mov.b32 accum_0_3, 0;\\n" \\
"mov.b32 accum_1_0, 0;\\n" \\
"mov.b32 accum_1_1, 0;\\n" \\
"mov.b32 accum_1_2, 0;\\n" \\
"mov.b32 accum_1_3, 0;\\n" \\
"mov.b32 accum_2_0, 0;\\n" \\
"mov.b32 accum_2_1, 0;\\n" \\
"mov.b32 accum_2_2, 0;\\n" \\
"mov.b32 accum_2_3, 0;\\n" \\
"mov.b32 accum_3_0, 0;\\n" \\
"mov.b32 accum_3_1, 0;\\n" \\
"mov.b32 accum_3_2, 0;\\n" \\
"mov.b32 accum_3_3, 0;\\n" \\
// multiply, unpacking and permuting scale factors
"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" \\
// unpacking A and B tensors
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\
"mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\
"mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\
"mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\
"mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\
// convert A and B tensors from fp4 to f16x2
// A[0 - 7] and B[0 - 7]
"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" \\
// A[8 - 15] and B[8 - 15]
"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" \\
// A[16 - 23] and B[16 - 23]
"cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\
// A[24 - 31] and B[24 - 31]
"cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\
// fma for A[0 - 7] and B[0 - 7]
"fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\
"fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\
"fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\
"fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\
// fma for A[8 - 15] and B[8 - 15]
"fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\
"fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\
"fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\
"fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\
// fma for A[16 - 23] and B[16 - 23]
"fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\
"fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\
"fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\
"fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\
// fma for A[24 - 31] and B[24 - 31]
"fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\
"fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\
"fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\
"fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\
// tree reduction for accumulators
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\
"add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\
"add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\
"add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\
"add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\
"fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\
"fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\
"fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\
"fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\
"fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\
"fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\
"fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\
"fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\
"}\\n"
: "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3
: "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5
"r"(a_packed.x), "r"(b_packed.x), // 6, 7
"r"(a_packed.y), "r"(b_packed.y), // 8, 9
"r"(a_packed.z), "r"(b_packed.z), // 10, 11
"r"(a_packed.w), "r"(b_packed.w) // 12, 13
);
}
// Reduce the result and store it in shared memory
__half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),
reinterpret_cast<const __half2&>(result_1));
__half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),
reinterpret_cast<const __half2&>(result_3));
reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;
for (int offset = 16; offset > 0; offset /= 2) {
final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);
}
if (threadIdx.x == 0) {
int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;
c[c_offset] = __float2half_rn(final_result_f);
}
}
__global__ void gemv_kernel(
const __nv_fp4x2_storage_t* __restrict__ a,
const __nv_fp4x2_storage_t* __restrict__ b,
const __nv_fp8_e4m3* __restrict__ sfa,
const __nv_fp8_e4m3* __restrict__ sfb,
__half* __restrict__ c,
int M,
int K
) {
extern __shared__ unsigned char shared_storage[];
auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);
auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));
__shared__ __half c_shared[32];
b += blockIdx.y * (K / 2) * 128;
sfb += blockIdx.y * (K / 16) * 128;
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];
}
for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {
reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];
}
__syncthreads();
// Each warp computes one result and saves it to shared memory
int result_0 = 0;
int result_1 = 0;
int result_2 = 0;
int result_3 = 0;
int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
a += offset;
sfa += offset / 8;
for (int i = threadIdx.x; i < K / 32; i += 32) {
int4 a_packed = reinterpret_cast<const int4*>(a)[i];
int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];
__nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
__nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
asm volatile( \\
"{\\n" \\
// declare registers for A / B tensors
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\
".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\
".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\
".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\
".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\
// declare registers for accumulators
".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\
".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\
".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\
".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\
// declare registers for scaling factors
".reg .f16x2 sfa_f16x2;\\n" \\
".reg .f16x2 sfb_f16x2;\\n" \\
".reg .f16x2 sf_f16x2;\\n" \\
// declare registers for conversion
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\
".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\
".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\
".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\
".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\
".reg .f16 result_f16, lane0, lane1;\\n" \\
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\
// convert scaling factors from fp8 to f16x2
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\
// clear accumulators
"mov.b32 accum_0_0, 0;\\n" \\
"mov.b32 accum_0_1, 0;\\n" \\
"mov.b32 accum_0_2, 0;\\n" \\
"mov.b32 accum_0_3, 0;\\n" \\
"mov.b32 accum_1_0, 0;\\n" \\
"mov.b32 accum_1_1, 0;\\n" \\
"mov.b32 accum_1_2, 0;\\n" \\
"mov.b32 accum_1_3, 0;\\n" \\
"mov.b32 accum_2_0, 0;\\n" \\
"mov.b32 accum_2_1, 0;\\n" \\
"mov.b32 accum_2_2, 0;\\n" \\
"mov.b32 accum_2_3, 0;\\n" \\
"mov.b32 accum_3_0, 0;\\n" \\
"mov.b32 accum_3_1, 0;\\n" \\
"mov.b32 accum_3_2, 0;\\n" \\
"mov.b32 accum_3_3, 0;\\n" \\
// multiply, unpacking and permuting scale factors
"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" \\
// unpacking A and B tensors
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\
"mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\
"mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\
"mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\
"mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\
// convert A and B tensors from fp4 to f16x2
// A[0 - 7] and B[0 - 7]
"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" \\
// A[8 - 15] and B[8 - 15]
"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" \\
// A[16 - 23] and B[16 - 23]
"cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\
// A[24 - 31] and B[24 - 31]
"cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\
// fma for A[0 - 7] and B[0 - 7]
"fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\
"fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\
"fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\
"fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\
// fma for A[8 - 15] and B[8 - 15]
"fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\
"fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\
"fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\
"fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\
// fma for A[16 - 23] and B[16 - 23]
"fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\
"fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\
"fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\
"fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\
// fma for A[24 - 31] and B[24 - 31]
"fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\
"fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\
"fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\
"fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\
// tree reduction for accumulators
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\
"add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\
"add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\
"add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\
"add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\
"fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\
"fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\
"fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\
"fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\
"fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\
"fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\
"fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\
"fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\
"}\\n"
: "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3
: "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5
"r"(a_packed.x), "r"(b_packed.x), // 6, 7
"r"(a_packed.y), "r"(b_packed.y), // 8, 9
"r"(a_packed.z), "r"(b_packed.z), // 10, 11
"r"(a_packed.w), "r"(b_packed.w) // 12, 13
);
}
// Reduce the result and store it in shared memory
__half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),
reinterpret_cast<const __half2&>(result_1));
__half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),
reinterpret_cast<const __half2&>(result_3));
reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;
for (int offset = 16; offset > 0; offset /= 2) {
final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);
}
if (threadIdx.x == 0) {
c_shared[threadIdx.y] = __float2half_rn(final_result_f);
}
__syncthreads();
// Write the result to global memory
if (threadIdx.y == 0) {
int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.x;
c[c_offset] = c_shared[threadIdx.x];
}
}
torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
const int64_t M = a.size(0);
const int64_t K = a.size(1) * 2;
const int64_t L = a.size(2);
dim3 block_dim(32, 32, 1);
dim3 grid_dim(M / 32, L, 1);
const auto* a_ptr = reinterpret_cast<const __nv_fp4x2_storage_t*>(a.data_ptr());
const auto* b_ptr = reinterpret_cast<const __nv_fp4x2_storage_t*>(b.data_ptr());
const auto* sfa_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());
const auto* sfb_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());
auto* c_ptr = reinterpret_cast<__half*>(c.data_ptr<c10::Half>());
size_t shared_mem_bytes =
(static_cast<size_t>(K) / 2) * sizeof(__nv_fp4x2_storage_t) +
(static_cast<size_t>(K) / 16) * sizeof(__nv_fp8_e4m3);
if (M == 4096 && K == 7168) {
gemv_kernel_4096_7168<<<grid_dim, block_dim, shared_mem_bytes>>>(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr
);
} else if (M == 7168 && K == 2048) {
gemv_kernel_7168_2048<<<grid_dim, block_dim, shared_mem_bytes>>>(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr
);
} else if (M == 7168 && K == 16384) {
gemv_kernel_7168_16384<<<grid_dim, block_dim, shared_mem_bytes>>>(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr
);
} else {
gemv_kernel<<<grid_dim, block_dim, shared_mem_bytes>>>(
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
static_cast<int>(M),
static_cast<int>(K)
);
}
return c;
}
"""
cpp_source = """
#include <torch/extension.h>
torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""
gemv_module = load_inline(
name='gemv_cuda',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['gemv_cuda'],
verbose=True,
extra_cuda_cflags=['-arch=compute_100a', '-code=sm_100a', '-O3'],
)
def custom_kernel(
data: input_t,
) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a, b, sfa, sfb, _, _, c = data
return gemv_module.gemv_cuda(a, b, sfa, sfb, c)scrolls · 1092 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 104178.
⋯ 292 unchanged lines__syncthreads();// Write the result to global memory- if (threadIdx.y == 0) {- int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.x;- c[c_offset] = c_shared[threadIdx.x];+ if (threadIdx.x == 0) {+ int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;+ c[c_offset] = __float2half_rn(final_result_f);}}+ __global__ void gemv_kernel_7168_2048(+ const __nv_fp4x2_storage_t* __restrict__ a,+ const __nv_fp4x2_storage_t* __restrict__ b,+ const __nv_fp8_e4m3* __restrict__ sfa,+ const __nv_fp8_e4m3* __restrict__ sfb,+ __half* __restrict__ c+ ) {+ const int M = 7168;+ const int K = 2048;+ extern __shared__ unsigned char shared_storage[];+ auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);+ auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));+ __shared__ __half c_shared[32];++ b += blockIdx.y * (K / 2) * 128;+ sfb += blockIdx.y * (K / 16) * 128;++ for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {+ reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];+ }+ for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {+ reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];+ }+ __syncthreads();++ // Each warp computes one result and saves it to shared memory+ int result_0 = 0;+ int result_1 = 0;+ int result_2 = 0;+ int result_3 = 0;+ int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);+ a += offset;+ sfa += offset / 8;++ for (int i = threadIdx.x; i < K / 32; i += 32) {+ int4 a_packed = reinterpret_cast<const int4*>(a)[i];+ int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];++ __nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];+ __nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];++ asm volatile( \\+ "{\\n" \\+ // declare registers for A / B tensors+ ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\+ ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\+ ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\+ ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\+ ".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\+ ".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\+ ".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\+ ".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\++ // declare registers for accumulators+ ".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\+ ".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\+ ".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\+ ".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\++ // declare registers for scaling factors+ ".reg .f16x2 sfa_f16x2;\\n" \\+ ".reg .f16x2 sfb_f16x2;\\n" \\+ ".reg .f16x2 sf_f16x2;\\n" \\++ // declare registers for conversion+ ".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\+ ".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\+ ".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\+ ".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\+ ".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\+ ".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\+ ".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\+ ".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\+ ".reg .f16 result_f16, lane0, lane1;\\n" \\+ ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\++ // convert scaling factors from fp8 to f16x2+ "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\+ "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\++ // clear accumulators+ "mov.b32 accum_0_0, 0;\\n" \\+ "mov.b32 accum_0_1, 0;\\n" \\+ "mov.b32 accum_0_2, 0;\\n" \\+ "mov.b32 accum_0_3, 0;\\n" \\+ "mov.b32 accum_1_0, 0;\\n" \\+ "mov.b32 accum_1_1, 0;\\n" \\+ "mov.b32 accum_1_2, 0;\\n" \\+ "mov.b32 accum_1_3, 0;\\n" \\+ "mov.b32 accum_2_0, 0;\\n" \\+ "mov.b32 accum_2_1, 0;\\n" \\+ "mov.b32 accum_2_2, 0;\\n" \\+ "mov.b32 accum_2_3, 0;\\n" \\+ "mov.b32 accum_3_0, 0;\\n" \\+ "mov.b32 accum_3_1, 0;\\n" \\+ "mov.b32 accum_3_2, 0;\\n" \\+ "mov.b32 accum_3_3, 0;\\n" \\++ // multiply, unpacking and permuting scale factors+ "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" \\++ // unpacking A and B tensors+ "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\+ "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\+ "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\+ "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\+ "mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\+ "mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\+ "mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\+ "mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\++ // convert A and B tensors from fp4 to f16x2++ // A[0 - 7] and B[0 - 7]+ "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" \\++ // A[8 - 15] and B[8 - 15]+ "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" \\++ // A[16 - 23] and B[16 - 23]+ "cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\++ // A[24 - 31] and B[24 - 31]+ "cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\++ // fma for A[0 - 7] and B[0 - 7]+ "fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\+ "fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\+ "fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\+ "fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\++ // fma for A[8 - 15] and B[8 - 15]+ "fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\+ "fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\+ "fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\+ "fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\++ // fma for A[16 - 23] and B[16 - 23]+ "fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\+ "fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\+ "fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\+ "fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\++ // fma for A[24 - 31] and B[24 - 31]+ "fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\+ "fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\+ "fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\+ "fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\++ // tree reduction for accumulators+ "add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\+ "add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\+ "add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\+ "add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\+ "add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\+ "add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\+ "add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\+ "add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\++ "fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\+ "fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\+ "fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\+ "fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\+++ "fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\+ "fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\+ "fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\+ "fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\++ "}\\n"+ : "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3+ : "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5+ "r"(a_packed.x), "r"(b_packed.x), // 6, 7+ "r"(a_packed.y), "r"(b_packed.y), // 8, 9+ "r"(a_packed.z), "r"(b_packed.z), // 10, 11+ "r"(a_packed.w), "r"(b_packed.w) // 12, 13+ );+ }+++ // Reduce the result and store it in shared memory+ __half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),+ reinterpret_cast<const __half2&>(result_1));+ __half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),+ reinterpret_cast<const __half2&>(result_3));+ reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);+ float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;+ for (int offset = 16; offset > 0; offset /= 2) {+ final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);+ }+ if (threadIdx.x == 0) {+ c_shared[threadIdx.y] = __float2half_rn(final_result_f);+ }+ __syncthreads();++ // Write the result to global memory+ if (threadIdx.x == 0) {+ int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;+ c[c_offset] = __float2half_rn(final_result_f);+ }+ }++++ __global__ void gemv_kernel_7168_16384(+ const __nv_fp4x2_storage_t* __restrict__ a,+ const __nv_fp4x2_storage_t* __restrict__ b,+ const __nv_fp8_e4m3* __restrict__ sfa,+ const __nv_fp8_e4m3* __restrict__ sfb,+ __half* __restrict__ c+ ) {+ const int M = 7168;+ const int K = 16384;++ extern __shared__ unsigned char shared_storage[];+ auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);+ auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));+ __shared__ __half c_shared[32];++ b += blockIdx.y * (K / 2) * 128;+ sfb += blockIdx.y * (K / 16) * 128;++ for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 32; i += blockDim.y * blockDim.x) {+ reinterpret_cast<int4*>(b_shared)[i] = reinterpret_cast<const int4*>(b)[i];+ }+ for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 256; i += blockDim.y * blockDim.x) {+ reinterpret_cast<int4*>(sfb_shared)[i] = reinterpret_cast<const int4*>(sfb)[i];+ }+ __syncthreads();++ // Each warp computes one result and saves it to shared memory+ int result_0 = 0;+ int result_1 = 0;+ int result_2 = 0;+ int result_3 = 0;+ int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);+ a += offset;+ sfa += offset / 8;++ for (int i = threadIdx.x; i < K / 32; i += 32) {+ int4 a_packed = reinterpret_cast<const int4*>(a)[i];+ int4 b_packed = reinterpret_cast<int4*>(b_shared)[i];++ __nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];+ __nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];++ asm volatile( \\+ "{\\n" \\+ // declare registers for A / B tensors+ ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\\n" \\+ ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\\n" \\+ ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\\n" \\+ ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\\n" \\+ ".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\\n" \\+ ".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\\n" \\+ ".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\\n" \\+ ".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\\n" \\++ // declare registers for accumulators+ ".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\\n" \\+ ".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\\n" \\+ ".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\\n" \\+ ".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\\n" \\++ // declare registers for scaling factors+ ".reg .f16x2 sfa_f16x2;\\n" \\+ ".reg .f16x2 sfb_f16x2;\\n" \\+ ".reg .f16x2 sf_f16x2;\\n" \\++ // declare registers for conversion+ ".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\\n" \\+ ".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\\n" \\+ ".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\\n" \\+ ".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\\n" \\+ ".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\\n" \\+ ".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\\n" \\+ ".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\\n" \\+ ".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\\n" \\+ ".reg .f16 result_f16, lane0, lane1;\\n" \\+ ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\\n" \\++ // convert scaling factors from fp8 to f16x2+ "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %4;\\n" \\+ "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %5;\\n" \\++ // clear accumulators+ "mov.b32 accum_0_0, 0;\\n" \\+ "mov.b32 accum_0_1, 0;\\n" \\+ "mov.b32 accum_0_2, 0;\\n" \\+ "mov.b32 accum_0_3, 0;\\n" \\+ "mov.b32 accum_1_0, 0;\\n" \\+ "mov.b32 accum_1_1, 0;\\n" \\+ "mov.b32 accum_1_2, 0;\\n" \\+ "mov.b32 accum_1_3, 0;\\n" \\+ "mov.b32 accum_2_0, 0;\\n" \\+ "mov.b32 accum_2_1, 0;\\n" \\+ "mov.b32 accum_2_2, 0;\\n" \\+ "mov.b32 accum_2_3, 0;\\n" \\+ "mov.b32 accum_3_0, 0;\\n" \\+ "mov.b32 accum_3_1, 0;\\n" \\+ "mov.b32 accum_3_2, 0;\\n" \\+ "mov.b32 accum_3_3, 0;\\n" \\++ // multiply, unpacking and permuting scale factors+ "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" \\++ // unpacking A and B tensors+ "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %6;\\n" \\+ "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %7;\\n" \\+ "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %8;\\n" \\+ "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %9;\\n" \\+ "mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %10;\\n" \\+ "mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %11;\\n" \\+ "mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %12;\\n" \\+ "mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %13;\\n" \\++ // convert A and B tensors from fp4 to f16x2++ // A[0 - 7] and B[0 - 7]+ "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" \\++ // A[8 - 15] and B[8 - 15]+ "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" \\++ // A[16 - 23] and B[16 - 23]+ "cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\\n" \\++ // A[24 - 31] and B[24 - 31]+ "cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\\n" \\+ "cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\\n" \\++ // fma for A[0 - 7] and B[0 - 7]+ "fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\\n" \\+ "fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\\n" \\+ "fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\\n" \\+ "fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\\n" \\++ // fma for A[8 - 15] and B[8 - 15]+ "fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\\n" \\+ "fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\\n" \\+ "fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\\n" \\+ "fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\\n" \\++ // fma for A[16 - 23] and B[16 - 23]+ "fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\\n" \\+ "fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\\n" \\+ "fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\\n" \\+ "fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\\n" \\++ // fma for A[24 - 31] and B[24 - 31]+ "fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\\n" \\+ "fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\\n" \\+ "fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\\n" \\+ "fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\\n" \\++ // tree reduction for accumulators+ "add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\\n" \\+ "add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\\n" \\+ "add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\\n" \\+ "add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\\n" \\+ "add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\\n" \\+ "add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\\n" \\+ "add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\\n" \\+ "add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\\n" \\++ "fma.rn.f16x2 %0, accum_0_0, mul_f16x2_0, %0;\\n" \\+ "fma.rn.f16x2 %1, accum_0_2, mul_f16x2_0, %1;\\n" \\+ "fma.rn.f16x2 %2, accum_1_0, mul_f16x2_0, %2;\\n" \\+ "fma.rn.f16x2 %3, accum_1_2, mul_f16x2_0, %3;\\n" \\+++ "fma.rn.f16x2 %0, accum_2_0, mul_f16x2_1, %0;\\n" \\+ "fma.rn.f16x2 %1, accum_2_2, mul_f16x2_1, %1;\\n" \\+ "fma.rn.f16x2 %2, accum_3_0, mul_f16x2_1, %2;\\n" \\+ "fma.rn.f16x2 %3, accum_3_2, mul_f16x2_1, %3;\\n" \\++ "}\\n"+ : "+r"(result_0), "+r"(result_1), "+r"(result_2), "+r"(result_3) // 0, 1, 2, 3+ : "h"(sfa_fp8x2), "h"(sfb_fp8x2), // 4, 5+ "r"(a_packed.x), "r"(b_packed.x), // 6, 7+ "r"(a_packed.y), "r"(b_packed.y), // 8, 9+ "r"(a_packed.z), "r"(b_packed.z), // 10, 11+ "r"(a_packed.w), "r"(b_packed.w) // 12, 13+ );+ }+++ // Reduce the result and store it in shared memory+ __half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result_0),+ reinterpret_cast<const __half2&>(result_1));+ __half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result_2),+ reinterpret_cast<const __half2&>(result_3));+ reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);+ float final_result_f = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;+ for (int offset = 16; offset > 0; offset /= 2) {+ final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);+ }+ if (threadIdx.x == 0) {+ int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;+ c[c_offset] = __float2half_rn(final_result_f);+ }+ }+++__global__ void gemv_kernel(const __nv_fp4x2_storage_t* __restrict__ a,const __nv_fp4x2_storage_t* __restrict__ b,⋯ 256 unchanged linessfb_ptr,c_ptr);+ } else if (M == 7168 && K == 2048) {+ gemv_kernel_7168_2048<<<grid_dim, block_dim, shared_mem_bytes>>>(+ a_ptr,+ b_ptr,+ sfa_ptr,+ sfb_ptr,+ c_ptr+ );+ } else if (M == 7168 && K == 16384) {+ gemv_kernel_7168_16384<<<grid_dim, block_dim, shared_mem_bytes>>>(+ a_ptr,+ b_ptr,+ sfa_ptr,+ sfb_ptr,+ c_ptr+ );} else {gemv_kernel<<<grid_dim, block_dim, shared_mem_bytes>>>(a_ptr,
scrolls · 508 diff lines total
Best evidence level for this revision: reported
JSON