Skip to content
KernelIndex
Search⌘K

submission 111178

tomaszki · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-111178?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
26.6µs
#123 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:719a22f31f0f7e72bcf59574b4836be58a8b65c35379173a4546c3615a20109b
license declaredunknown
license concludedunknown
authorstomaszki
imported2026-08-15

Techniques

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

fp4PyTorch reference implementation of NVFP4 block-scaled GEMV.
fp8__nv_fp8x2_storage_t sfa_fp8x2,
shared-memoryextern __shared__ unsigned char shared_storage[];
vector-width = int4int4 a_packed,

Kernel source

submission.py652 lines
#!POPCORN leaderboard nvfp4_gemv

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


# CUDA SOURCE CODE

cuda_source = """
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_pipeline.h>


#define FULL_MASK 0xffffffff

__inline__ __device__ void multiply_and_accumulate(
    int4 a_packed,
    int4 b_packed,
    __nv_fp8x2_storage_t sfa_fp8x2,
    __nv_fp8x2_storage_t sfb_fp8x2,
    int* result_0,
    int* result_1,
    int* result_2,
    int* result_3
) {
    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
    );
}


__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];

        multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
    }


    // 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_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];

        multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
    }


    // 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
__maxnreg__(146)
gemv_kernel_7168_16384(
    const int4* __restrict__ a,
    const int4* __restrict__ b,
    const int* __restrict__ sfa,
    const int* __restrict__ sfb,
    __half* __restrict__ c
) {
    const int M = 7168;
    const int K = 16384;
    const int Q_SIZE = 2;
    const int active_warps = (blockIdx.x < 32) ? 26 : 25;

    __shared__ int4 a_shared[Q_SIZE + 1][25][2][32];
    __shared__ int sfa_shared[Q_SIZE + 1][25][32];

    // We will load all b and sfb, because we can, it simplifies the logic
    __shared__ int4 b_shared[8][2][32];
    __shared__ int sfb_shared[8][32];

    // L = 1 so we don't have to bother to offset b or sfb

    int offset = (blockIdx.x * 24 + min(blockIdx.x, 32) + threadIdx.y - 1) * 2 * (K / 2);
    a += offset / 16;
    sfa += offset / 32;

    // Prologue
    #pragma unroll
    for (int prefetch_idx = 0; prefetch_idx < Q_SIZE; prefetch_idx++) {
        int col_idx = prefetch_idx / 2;
        int row_idx = prefetch_idx % 2;
        if (threadIdx.y == 0 and row_idx == 0) {
            __pipeline_memcpy_async(&b_shared[col_idx][0][threadIdx.x], &b[col_idx * 64 + threadIdx.x], sizeof(int4));
            __pipeline_memcpy_async(&b_shared[col_idx][1][threadIdx.x], &b[col_idx * 64 + 32 + threadIdx.x], sizeof(int4));
            __pipeline_memcpy_async(&sfb_shared[col_idx][threadIdx.x], &sfb[col_idx * 32 + threadIdx.x], sizeof(int));
        } else if (threadIdx.y > 0 && threadIdx.y < active_warps) {
            __pipeline_memcpy_async(
                &a_shared[prefetch_idx][threadIdx.y - 1][0][threadIdx.x],
                &a[row_idx * (K / 32) + col_idx * 64 + threadIdx.x],
                sizeof(int4)
            );
            __pipeline_memcpy_async(
                &a_shared[prefetch_idx][threadIdx.y - 1][1][threadIdx.x],
                &a[row_idx * (K / 32) + col_idx * 64 + 32 + threadIdx.x],
                sizeof(int4)
            );
            __pipeline_memcpy_async(
                &sfa_shared[prefetch_idx][threadIdx.y - 1][threadIdx.x],
                &sfa[row_idx * (K / 64) + col_idx * 32 + threadIdx.x], sizeof(int)
            );
        }
        __pipeline_commit();
    }

    int result[2][4] = {0};
    #pragma unroll
    for (int load_idx = 0; load_idx + Q_SIZE < 8 * 2; load_idx++) {
        int prefetch_idx = load_idx + Q_SIZE;
        int col_idx = prefetch_idx / 2;
        int row_idx = prefetch_idx % 2;
        if (threadIdx.y == 0 and row_idx == 0) {
            __pipeline_memcpy_async(&b_shared[col_idx][0][threadIdx.x], &b[col_idx * 64 + threadIdx.x], sizeof(int4));
            __pipeline_memcpy_async(&b_shared[col_idx][1][threadIdx.x], &b[col_idx * 64 + 32 + threadIdx.x], sizeof(int4));
            __pipeline_memcpy_async(&sfb_shared[col_idx][threadIdx.x], &sfb[col_idx * 32 + threadIdx.x], sizeof(int));
        } else if (threadIdx.y > 0 && threadIdx.y < active_warps) {
            __pipeline_memcpy_async(
                &a_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x],
                &a[row_idx * (K / 32) + col_idx * 64 + threadIdx.x],
                sizeof(int4)
            );
            __pipeline_memcpy_async(
                &a_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x],
                &a[row_idx * (K / 32) + col_idx * 64 + 32 + threadIdx.x],
                sizeof(int4)
            );
            __pipeline_memcpy_async(
                &sfa_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][threadIdx.x],
                &sfa[row_idx * (K / 64) + col_idx * 32 + threadIdx.x], sizeof(int)
            );
        }
        __pipeline_commit();
        __pipeline_wait_prior(Q_SIZE);
        if (load_idx % 2 == 0) {
            __syncthreads();
        }

        if (threadIdx.y > 0 && threadIdx.y < active_warps) {
            int load_col_idx = load_idx / 2;
            int load_row_idx = load_idx % 2;
            int4 a_packed_0 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x]; // [Q_SIZE + 1][13][2][32]
            int4 b_packed_0 = b_shared[load_col_idx][0][threadIdx.x]; // [8][2][32]
            __nv_fp8x2_storage_t sfa_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x]; // [Q_SIZE + 1][13][32]
            __nv_fp8x2_storage_t sfb_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x]; // [8][32]
            multiply_and_accumulate(
                a_packed_0, b_packed_0, sfa_fp8x2_0, sfb_fp8x2_0,
                &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
            );

            // SECOND ITERATION
            int4 a_packed_1 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x];
            int4 b_packed_1 = b_shared[load_col_idx][1][threadIdx.x];
            __nv_fp8x2_storage_t sfa_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x + 32]; // [Q_SIZE + 1][13][32]
            __nv_fp8x2_storage_t sfb_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x + 32]; // [8][32]
            multiply_and_accumulate(
                a_packed_1, b_packed_1, sfa_fp8x2_1, sfb_fp8x2_1,
                &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
            );
        }
    }

    // Epilogue
    #pragma unroll
    for (int load_idx = 16 - Q_SIZE; load_idx < 8 * 2; load_idx++) {
        __pipeline_wait_prior(15 - load_idx);
        if (load_idx % 2 == 0) {
            __syncthreads();
        }

        if (threadIdx.y > 0 && threadIdx.y < active_warps) {
            int load_col_idx = load_idx / 2;
            int load_row_idx = load_idx % 2;
            int4 a_packed_0 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x]; // [Q_SIZE + 1][13][2][32]
            int4 b_packed_0 = b_shared[load_col_idx][0][threadIdx.x]; // [8][2][32]
            __nv_fp8x2_storage_t sfa_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x]; // [Q_SIZE + 1][13][32]
            __nv_fp8x2_storage_t sfb_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x]; // [8][32]
            multiply_and_accumulate(
                a_packed_0, b_packed_0, sfa_fp8x2_0, sfb_fp8x2_0,
                &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
            );

            // SECOND ITERATION
            int4 a_packed_1 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x];
            int4 b_packed_1 = b_shared[load_col_idx][1][threadIdx.x];
            __nv_fp8x2_storage_t sfa_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x + 32]; // [Q_SIZE + 1][13][32]
            __nv_fp8x2_storage_t sfb_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x + 32]; // [8][32]
            multiply_and_accumulate(
                a_packed_1, b_packed_1, sfa_fp8x2_1, sfb_fp8x2_1,
                &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
            );
        }
    }
    float final_result_f[2];
    for (int i = 0; i < 2; i++) {
        // Reduce the result and store it in shared memory
        __half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result[i][0]),
                reinterpret_cast<const __half2&>(result[i][1]));
        __half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result[i][2]),
                reinterpret_cast<const __half2&>(result[i][3]));
        reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
        final_result_f[i] = __half22float2(reduction_result_0).x + __half22float2(reduction_result_0).y;
    }
    for (int offset = 16; offset > 0; offset /= 2) {
        for (int i = 0; i < 2; i++) {
            final_result_f[i] += __shfl_down_sync(FULL_MASK, final_result_f[i], offset);
        }
    }
    if (threadIdx.x == 0 && threadIdx.y > 0 && threadIdx.y < active_warps) {
        __half final_result[2];
        for (int i = 0; i < 2; i++) {
            final_result[i] = __float2half_rn(final_result_f[i]);
        }
        int c_offset = (blockIdx.x * 24 + min((int)blockIdx.x, 32) + threadIdx.y - 1);
        reinterpret_cast<int*>(c)[c_offset] = reinterpret_cast<int&>(final_result);
    }
}



__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];

        multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
    }


    // 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) {
        grid_dim = dim3(148, 1, 1);
        block_dim = dim3(32, 26, 1);
        gemv_kernel_7168_16384<<<grid_dim, block_dim>>>(
            reinterpret_cast<const int4*>(a.data_ptr()),
            reinterpret_cast<const int4*>(b.data_ptr()),
            reinterpret_cast<const int*>(sfa.data_ptr()),
            reinterpret_cast<const int*>(sfb.data_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 · 652 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 106593.

⋯ 3 unchanged lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
- # Kernel configuration parameters
- sf_vec_size = 16
+ # CUDA SOURCE CODE
- # Helper function for ceiling division
- def ceil_div(a, b):
- return (a + b - 1) // b
+ cuda_source = """
+ #include <cuda_fp4.h>
+ #include <cuda_fp8.h>
+ #include <cuda_fp16.h>
+ #include <cuda_pipeline.h>
- # Helper function to convert scale factor tensor to blocked format
- def to_blocked(input_matrix):
- rows, cols = input_matrix.shape
+ #define FULL_MASK 0xffffffff
- # 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)
+ __inline__ __device__ void multiply_and_accumulate(
+ int4 a_packed,
+ int4 b_packed,
+ __nv_fp8x2_storage_t sfa_fp8x2,
+ __nv_fp8x2_storage_t sfb_fp8x2,
+ int* result_0,
+ int* result_1,
+ int* result_2,
+ int* result_3
+ ) {
+ 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" \\
- 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)
+ // 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" \\
- return rearranged.flatten()
+ // 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" \\
- 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
+ // 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" \\
- # Get dimensions from MxNxL layout
- _, _, l = c_ref.shape
+ // 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" \\
- # 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
+ // convert A and B tensors from fp4 to f16x2
- # CUDA SOURCE CODE
+ // 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" \\
- cuda_source = """
- #include <cuda_fp4.h>
- #include <cuda_fp8.h>
- #include <cuda_fp16.h>
+ // 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" \\
- #define FULL_MASK 0xffffffff
+ // 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
+ );
+ }
+
+
__global__ void gemv_kernel_4096_7168(
const __nv_fp4x2_storage_t* __restrict__ a,
const __nv_fp4x2_storage_t* __restrict__ b,
⋯ 36 unchanged lines
__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
- );
+ multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
}
⋯ 8 unchanged lines
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);
}
⋯ 42 unchanged lines
__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
- );
+ multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
}
⋯ 8 unchanged lines
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);
}
⋯ 1 unchanged lines
- __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,
+ __global__ void
+ __maxnreg__(146)
+ gemv_kernel_7168_16384(
+ const int4* __restrict__ a,
+ const int4* __restrict__ b,
+ const int* __restrict__ sfa,
+ const int* __restrict__ sfb,
__half* __restrict__ c
) {
const int M = 7168;
const int K = 16384;
+ const int Q_SIZE = 2;
+ const int active_warps = (blockIdx.x < 32) ? 26 : 25;
- 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];
+ __shared__ int4 a_shared[Q_SIZE + 1][25][2][32];
+ __shared__ int sfa_shared[Q_SIZE + 1][25][32];
- b += blockIdx.y * (K / 2) * 128;
- sfb += blockIdx.y * (K / 16) * 128;
+ // We will load all b and sfb, because we can, it simplifies the logic
+ __shared__ int4 b_shared[8][2][32];
+ __shared__ int sfb_shared[8][32];
- 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();
+ // L = 1 so we don't have to bother to offset b or sfb
- // 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];
+ int offset = (blockIdx.x * 24 + min(blockIdx.x, 32) + threadIdx.y - 1) * 2 * (K / 2);
+ a += offset / 16;
+ sfa += offset / 32;
- 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" \\
+ // Prologue
+ #pragma unroll
+ for (int prefetch_idx = 0; prefetch_idx < Q_SIZE; prefetch_idx++) {
+ int col_idx = prefetch_idx / 2;
+ int row_idx = prefetch_idx % 2;
+ if (threadIdx.y == 0 and row_idx == 0) {
+ __pipeline_memcpy_async(&b_shared[col_idx][0][threadIdx.x], &b[col_idx * 64 + threadIdx.x], sizeof(int4));
+ __pipeline_memcpy_async(&b_shared[col_idx][1][threadIdx.x], &b[col_idx * 64 + 32 + threadIdx.x], sizeof(int4));
+ __pipeline_memcpy_async(&sfb_shared[col_idx][threadIdx.x], &sfb[col_idx * 32 + threadIdx.x], sizeof(int));
+ } else if (threadIdx.y > 0 && threadIdx.y < active_warps) {
+ __pipeline_memcpy_async(
+ &a_shared[prefetch_idx][threadIdx.y - 1][0][threadIdx.x],
+ &a[row_idx * (K / 32) + col_idx * 64 + threadIdx.x],
+ sizeof(int4)
+ );
+ __pipeline_memcpy_async(
+ &a_shared[prefetch_idx][threadIdx.y - 1][1][threadIdx.x],
+ &a[row_idx * (K / 32) + col_idx * 64 + 32 + threadIdx.x],
+ sizeof(int4)
+ );
+ __pipeline_memcpy_async(
+ &sfa_shared[prefetch_idx][threadIdx.y - 1][threadIdx.x],
+ &sfa[row_idx * (K / 64) + col_idx * 32 + threadIdx.x], sizeof(int)
+ );
+ }
+ __pipeline_commit();
+ }
- // 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" \\
+ int result[2][4] = {0};
+ #pragma unroll
+ for (int load_idx = 0; load_idx + Q_SIZE < 8 * 2; load_idx++) {
+ int prefetch_idx = load_idx + Q_SIZE;
+ int col_idx = prefetch_idx / 2;
+ int row_idx = prefetch_idx % 2;
+ if (threadIdx.y == 0 and row_idx == 0) {
+ __pipeline_memcpy_async(&b_shared[col_idx][0][threadIdx.x], &b[col_idx * 64 + threadIdx.x], sizeof(int4));
+ __pipeline_memcpy_async(&b_shared[col_idx][1][threadIdx.x], &b[col_idx * 64 + 32 + threadIdx.x], sizeof(int4));
+ __pipeline_memcpy_async(&sfb_shared[col_idx][threadIdx.x], &sfb[col_idx * 32 + threadIdx.x], sizeof(int));
+ } else if (threadIdx.y > 0 && threadIdx.y < active_warps) {
+ __pipeline_memcpy_async(
+ &a_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x],
+ &a[row_idx * (K / 32) + col_idx * 64 + threadIdx.x],
+ sizeof(int4)
+ );
+ __pipeline_memcpy_async(
+ &a_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x],
+ &a[row_idx * (K / 32) + col_idx * 64 + 32 + threadIdx.x],
+ sizeof(int4)
+ );
+ __pipeline_memcpy_async(
+ &sfa_shared[prefetch_idx % (Q_SIZE + 1)][threadIdx.y - 1][threadIdx.x],
+ &sfa[row_idx * (K / 64) + col_idx * 32 + threadIdx.x], sizeof(int)
+ );
+ }
+ __pipeline_commit();
+ __pipeline_wait_prior(Q_SIZE);
+ if (load_idx % 2 == 0) {
+ __syncthreads();
+ }
- // 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" \\
+ if (threadIdx.y > 0 && threadIdx.y < active_warps) {
+ int load_col_idx = load_idx / 2;
+ int load_row_idx = load_idx % 2;
+ int4 a_packed_0 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x]; // [Q_SIZE + 1][13][2][32]
+ int4 b_packed_0 = b_shared[load_col_idx][0][threadIdx.x]; // [8][2][32]
+ __nv_fp8x2_storage_t sfa_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x]; // [Q_SIZE + 1][13][32]
+ __nv_fp8x2_storage_t sfb_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x]; // [8][32]
+ multiply_and_accumulate(
+ a_packed_0, b_packed_0, sfa_fp8x2_0, sfb_fp8x2_0,
+ &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
+ );
- // 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" \\
+ // SECOND ITERATION
+ int4 a_packed_1 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x];
+ int4 b_packed_1 = b_shared[load_col_idx][1][threadIdx.x];
+ __nv_fp8x2_storage_t sfa_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x + 32]; // [Q_SIZE + 1][13][32]
+ __nv_fp8x2_storage_t sfb_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x + 32]; // [8][32]
+ multiply_and_accumulate(
+ a_packed_1, b_packed_1, sfa_fp8x2_1, sfb_fp8x2_1,
+ &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
+ );
+ }
+ }
- // 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" \\
+ // Epilogue
+ #pragma unroll
+ for (int load_idx = 16 - Q_SIZE; load_idx < 8 * 2; load_idx++) {
+ __pipeline_wait_prior(15 - load_idx);
+ if (load_idx % 2 == 0) {
+ __syncthreads();
+ }
- // convert A and B tensors from fp4 to f16x2
+ if (threadIdx.y > 0 && threadIdx.y < active_warps) {
+ int load_col_idx = load_idx / 2;
+ int load_row_idx = load_idx % 2;
+ int4 a_packed_0 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][0][threadIdx.x]; // [Q_SIZE + 1][13][2][32]
+ int4 b_packed_0 = b_shared[load_col_idx][0][threadIdx.x]; // [8][2][32]
+ __nv_fp8x2_storage_t sfa_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x]; // [Q_SIZE + 1][13][32]
+ __nv_fp8x2_storage_t sfb_fp8x2_0 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x]; // [8][32]
+ multiply_and_accumulate(
+ a_packed_0, b_packed_0, sfa_fp8x2_0, sfb_fp8x2_0,
+ &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
+ );
- // 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
- );
+ // SECOND ITERATION
+ int4 a_packed_1 = a_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1][1][threadIdx.x];
+ int4 b_packed_1 = b_shared[load_col_idx][1][threadIdx.x];
+ __nv_fp8x2_storage_t sfa_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfa_shared[load_idx % (Q_SIZE + 1)][threadIdx.y - 1])[threadIdx.x + 32]; // [Q_SIZE + 1][13][32]
+ __nv_fp8x2_storage_t sfb_fp8x2_1 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared[load_col_idx])[threadIdx.x + 32]; // [8][32]
+ multiply_and_accumulate(
+ a_packed_1, b_packed_1, sfa_fp8x2_1, sfb_fp8x2_1,
+ &result[load_row_idx][0], &result[load_row_idx][1], &result[load_row_idx][2], &result[load_row_idx][3]
+ );
+ }
}
-
-
- // 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;
+ float final_result_f[2];
+ for (int i = 0; i < 2; i++) {
+ // Reduce the result and store it in shared memory
+ __half2 reduction_result_0 = __hadd2(reinterpret_cast<const __half2&>(result[i][0]),
+ reinterpret_cast<const __half2&>(result[i][1]));
+ __half2 reduction_result_1 = __hadd2(reinterpret_cast<const __half2&>(result[i][2]),
+ reinterpret_cast<const __half2&>(result[i][3]));
+ reduction_result_0 = __hadd2(reduction_result_0, reduction_result_1);
+ final_result_f[i] = __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);
+ for (int i = 0; i < 2; i++) {
+ final_result_f[i] += __shfl_down_sync(FULL_MASK, final_result_f[i], offset);
+ }
}
- if (threadIdx.x == 0) {
- int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.y;
- c[c_offset] = __float2half_rn(final_result_f);
+ if (threadIdx.x == 0 && threadIdx.y > 0 && threadIdx.y < active_warps) {
+ __half final_result[2];
+ for (int i = 0; i < 2; i++) {
+ final_result[i] = __float2half_rn(final_result_f[i]);
+ }
+ int c_offset = (blockIdx.x * 24 + min((int)blockIdx.x, 32) + threadIdx.y - 1);
+ reinterpret_cast<int*>(c)[c_offset] = reinterpret_cast<int&>(final_result);
}
}
⋯ 40 unchanged lines
__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
- );
+ multiply_and_accumulate(a_packed, b_packed, sfa_fp8x2, sfb_fp8x2, &result_0, &result_1, &result_2, &result_3);
}
⋯ 56 unchanged lines
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,
+ grid_dim = dim3(148, 1, 1);
+ block_dim = dim3(32, 26, 1);
+ gemv_kernel_7168_16384<<<grid_dim, block_dim>>>(
+ reinterpret_cast<const int4*>(a.data_ptr()),
+ reinterpret_cast<const int4*>(b.data_ptr()),
+ reinterpret_cast<const int*>(sfa.data_ptr()),
+ reinterpret_cast<const int*>(sfb.data_ptr()),
c_ptr
);
} else {
scrolls · 1196 diff lines total

Best evidence level for this revision: reported

JSON