Skip to content
KernelIndex
Search⌘K

submission 106649

yue · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submit_v9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-106649?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
22.7µs
#50 of 678
2025-11-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6c4620cae9f7a5b0f25555022e4ae6c7ac20bf2e6b4b8c58475ad256d9c6080e
license declaredunknown
license concludedunknown
authorsyue
imported2026-08-15

Techniques

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

fp4- Combined FP4 operations for 2 uint32_t at once
tile-k = 64constexpr int TILE_K = 64;
vector-width = float4float4 A_data0_t0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1)));

Kernel source

submit_v9.py618 lines
import torch
import sys
from torch.utils.cpp_extension import load_inline
from typing import Tuple

gemv_cuda_src = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cstdint>

extern "C" __global__ __launch_bounds__(128, 8)
void block_scaled_gemv_fp4_fp8_fp16_optimized(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    __half* __restrict__ C,
    int M, int K, int L)
{
    constexpr int ROWS_PER_BLOCK = 8;
    constexpr int THREADS_PER_ROW = 16;
    constexpr int TILE_K = 64;

    const int tidx = threadIdx.x;
    const int tidy = threadIdx.y;
    const int block_row = blockIdx.x * ROWS_PER_BLOCK;
    const int batch_idx = blockIdx.z;
    const int global_row = block_row + tidy;

    if (global_row >= M) return;

    const int num_k_tiles = K >> 6;

    const uint8_t* B_base = B + batch_idx * (128 * (K >> 1));
    const uint8_t* SFB_base = SFB + batch_idx * (128 * (K >> 4));
    const uint8_t* A_row = A + batch_idx * (M * (K >> 1)) + global_row * (K >> 1);
    const uint8_t* SFA_row = SFA + batch_idx * (M * (K >> 4)) + global_row * (K >> 4);

    float local_sum = 0.0f;

    // Main loop - process 2 tiles per iteration, fully inlined
    int tile = tidx;
    
    #pragma unroll 1
    for (; tile + THREADS_PER_ROW < num_k_tiles; tile += 2 * THREADS_PER_ROW) {
        // Load tile 0
        const int k_offset_0 = tile * TILE_K;
        float4 A_data0_t0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1)));
        float4 B_data0_t0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1)));
        float4 A_data1_t0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1) + 16));
        float4 B_data1_t0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1) + 16));
        uint32_t sfa_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_0 >> 4)));
        uint32_t sfb_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_0 >> 4)));

        // Load tile 1
        const int k_offset_1 = (tile + THREADS_PER_ROW) * TILE_K;
        float4 A_data0_t1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1)));
        float4 B_data0_t1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1)));
        float4 A_data1_t1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1) + 16));
        float4 B_data1_t1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1) + 16));
        uint32_t sfa_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_1 >> 4)));
        uint32_t sfb_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_1 >> 4)));

        // Decode FP8 scales for tile 0
        uint32_t sfa_h2_0, sfa_h2_1, sfb_h2_0, sfb_h2_1;
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
            : "r"(sfa_vec_t0)
        );
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
            : "r"(sfb_vec_t0)
        );
        __half2 sfa_scales_t0_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
        __half2 sfa_scales_t0_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
        __half2 sfb_scales_t0_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
        __half2 sfb_scales_t0_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);

        // Process tile 0 - all 4 SF blocks
        const uint32_t* A_u32 = reinterpret_cast<const uint32_t*>(&A_data0_t0);
        const uint32_t* B_u32 = reinterpret_cast<const uint32_t*>(&B_data0_t0);
        float tile_sum_0 = 0.0f;
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_t0_0.x : sfa_scales_t0_0.y,
                sf == 0 ? sfb_scales_t0_0.x : sfb_scales_t0_0.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum_0 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_0);
        }
        
        A_u32 = reinterpret_cast<const uint32_t*>(&A_data1_t0);
        B_u32 = reinterpret_cast<const uint32_t*>(&B_data1_t0);
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_t0_1.x : sfa_scales_t0_1.y,
                sf == 0 ? sfb_scales_t0_1.x : sfb_scales_t0_1.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum_0 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_0);
        }
        
        local_sum += tile_sum_0;

        // Decode and process tile 1 (same as tile 0)
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
            : "r"(sfa_vec_t1)
        );
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
            : "r"(sfb_vec_t1)
        );
        __half2 sfa_scales_t1_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
        __half2 sfa_scales_t1_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
        __half2 sfb_scales_t1_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
        __half2 sfb_scales_t1_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);

        A_u32 = reinterpret_cast<const uint32_t*>(&A_data0_t1);
        B_u32 = reinterpret_cast<const uint32_t*>(&B_data0_t1);
        float tile_sum_1 = 0.0f;
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_t1_0.x : sfa_scales_t1_0.y,
                sf == 0 ? sfb_scales_t1_0.x : sfb_scales_t1_0.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum_1 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_1);
        }
        
        A_u32 = reinterpret_cast<const uint32_t*>(&A_data1_t1);
        B_u32 = reinterpret_cast<const uint32_t*>(&B_data1_t1);
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_t1_1.x : sfa_scales_t1_1.y,
                sf == 0 ? sfb_scales_t1_1.x : sfb_scales_t1_1.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum_1 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_1);
        }
        
        local_sum += tile_sum_1;
    }
    
    // Handle remaining single tile
    if (tile < num_k_tiles) {
        const int k_offset = tile * TILE_K;
        float4 A_data0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset >> 1)));
        float4 B_data0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset >> 1)));
        float4 A_data1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset >> 1) + 16));
        float4 B_data1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset >> 1) + 16));
        uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset >> 4)));
        uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset >> 4)));

        uint32_t sfa_h2_0, sfa_h2_1, sfb_h2_0, sfb_h2_1;
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
            : "r"(sfa_vec)
        );
        asm volatile (
            "{"
            "  .reg .b16 %%low, %%high;\\n"
            "  mov.b32 {%%low, %%high}, %2;\\n"
            "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
            "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
            "}"
            : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
            : "r"(sfb_vec)
        );
        __half2 sfa_scales_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
        __half2 sfa_scales_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
        __half2 sfb_scales_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
        __half2 sfb_scales_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);

        const uint32_t* A_u32 = reinterpret_cast<const uint32_t*>(&A_data0);
        const uint32_t* B_u32 = reinterpret_cast<const uint32_t*>(&B_data0);
        float tile_sum = 0.0f;
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_0.x : sfa_scales_0.y,
                sf == 0 ? sfb_scales_0.x : sfb_scales_0.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
        }
        
        A_u32 = reinterpret_cast<const uint32_t*>(&A_data1);
        B_u32 = reinterpret_cast<const uint32_t*>(&B_data1);
        
        #pragma unroll
        for (int sf = 0; sf < 2; sf++) {
            __half scale = __hmul(
                sf == 0 ? sfa_scales_1.x : sfa_scales_1.y,
                sf == 0 ? sfb_scales_1.x : sfb_scales_1.y
            );
            float block_sum;
            asm volatile (
                "{"
                "  .reg .b8 %%ab<4>, %%bb<4>;\\n"
                "  .reg .b32 %%a<4>, %%b<4>;\\n"
                "  .reg .b32 %%p0, %%p1;\\n"
                "  .reg .f16 %%h0, %%h1;\\n"
                "  .reg .f32 %%f0, %%f1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
                "  mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
                "  mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
                "  cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
                "  mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
                "  fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
                "  fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
                "  add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
                "  mov.b32 {%%h0, %%h1}, %%p0;\\n"
                "  cvt.f32.f16 %%f0, %%h0;\\n"
                "  cvt.f32.f16 %%f1, %%h1;\\n"
                "  add.f32 %0, %%f0, %%f1;\\n"
                "}"
                : "=f"(block_sum)
                : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]), 
                  "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
            );
            tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
        }
        
        local_sum += tile_sum;
    }

    // Warp-level reduction
    #pragma unroll
    for (int offset = 8; offset > 0; offset >>= 1) {
        local_sum += __shfl_xor_sync(0xffff, local_sum, offset, 16);
    }

    if (tidx == 0) {
        C[batch_idx * M + global_row] = __float2half(local_sum);
    }
}

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L)
{
    TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
    TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
    TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");
    TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");
    TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");

    constexpr int ROWS_PER_BLOCK = 8;
    constexpr int THREADS_PER_ROW = 16;

    const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);
    const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);

    block_scaled_gemv_fp4_fp8_fp16_optimized<<<grid, block>>>(
        reinterpret_cast<const uint8_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B.data_ptr()),
        reinterpret_cast<const uint8_t*>(SFA.data_ptr()),
        reinterpret_cast<const uint8_t*>(SFB.data_ptr()),
        reinterpret_cast<__half*>(C.data_ptr()),
        M, K, L);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess)
        throw std::runtime_error(cudaGetErrorString(err));

    return C;
}
"""

gemv_cpp_src = """
#include <torch/extension.h>

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L);
"""

_gemm_module = load_inline(
    name="block_scaled_gemv_inline_v1",
    cpp_sources=gemv_cpp_src,
    cuda_sources=gemv_cuda_src,
    functions=["gemv_fp4_fp8_fp16"],
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '-std=c++17',
        '--expt-relaxed-constexpr',
        '--maxrregcount=255',
        '--prec-div=false', 
        '--fmad=true',
        '--ftz=true',
        '-gencode=arch=compute_100a,code=sm_100a',
    ],
    verbose=True,
)

def custom_kernel(data):
    """
    Fully inlined version with 2 tiles:
    - No function calls, everything inlined directly
    - 2 tiles to keep register pressure low
    - Combined FP4 operations for 2 uint32_t at once
    - Should have minimal overhead
    """
    a, b, sfa, sfb, _, _, c = data
    m, k_packed, l = a.shape
    k = k_packed * 2

    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)
    sfa_uint8 = sfa.view(torch.uint8)
    sfb_uint8 = sfb.view(torch.uint8)

    _gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)

    return c

scrolls · 618 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 105562.

⋯ 7 unchanged lines
#include <cuda_runtime.h>
#include <cstdint>
- __device__ __forceinline__ float decode_mul_accumulate_fp4x8(
- const uint32_t a_packed,
- const uint32_t b_packed,
- float acc)
- {
- float result;
-
- asm volatile (
- "{"
- " .reg .b8 %%ab<4>, %%bb<4>;\\n"
- " .reg .b32 %%a<4>, %%b<4>;\\n"
- " .reg .b32 %%p0, %%p1;\\n"
- " .reg .f16 %%h0, %%h1;\\n"
- " .reg .f32 %%f0, %%f1;\\n"
- " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
- " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
- " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
- " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
- " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
- " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
- " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
- " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
- " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
- " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
- " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
- " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
- " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
- " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
- " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
- " mov.b32 {%%h0, %%h1}, %%p0;\\n"
- " cvt.f32.f16 %%f0, %%h0;\\n"
- " cvt.f32.f16 %%f1, %%h1;\\n"
- " add.f32 %%f0, %%f0, %%f1;\\n"
- " add.f32 %0, %%f0, %3;\\n"
- "}"
- : "=f"(result)
- : "r"(a_packed), "r"(b_packed), "f"(acc)
- );
-
- return result;
- }
-
- __device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
- {
- uint32_t out_low, out_high;
-
- asm volatile (
- "{"
- " .reg .b16 %%low, %%high;\\n"
- " mov.b32 {%%low, %%high}, %2;\\n"
- " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
- " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
- "}"
- : "=r"(out_low), "=r"(out_high)
- : "r"(packed)
- );
-
- __half2 h_low = *reinterpret_cast<const __half2*>(&out_low);
- __half2 h_high = *reinterpret_cast<const __half2*>(&out_high);
- h0 = h_low.x;
- h1 = h_low.y;
- h2 = h_high.x;
- h3 = h_high.y;
- }
-
- // Process one tile and return partial sum
- __device__ __forceinline__ float process_tile(
- const uint32_t* A_u32_0,
- const uint32_t* B_u32_0,
- const uint32_t* A_u32_1,
- const uint32_t* B_u32_1,
- const __half* sfa_scales,
- const __half* sfb_scales)
- {
- float tile_sum = 0.0f;
-
- // SF block 0
- {
- __half scale = __hmul(sfa_scales[0], sfb_scales[0]);
- float block_sum = 0.0f;
- block_sum = decode_mul_accumulate_fp4x8(A_u32_0[0], B_u32_0[0], block_sum);
- block_sum = decode_mul_accumulate_fp4x8(A_u32_0[1], B_u32_0[1], block_sum);
- tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
- }
-
- // SF block 1
- {
- __half scale = __hmul(sfa_scales[1], sfb_scales[1]);
- float block_sum = 0.0f;
- block_sum = decode_mul_accumulate_fp4x8(A_u32_0[2], B_u32_0[2], block_sum);
- block_sum = decode_mul_accumulate_fp4x8(A_u32_0[3], B_u32_0[3], block_sum);
- tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
- }
-
- // SF block 2
- {
- __half scale = __hmul(sfa_scales[2], sfb_scales[2]);
- float block_sum = 0.0f;
- block_sum = decode_mul_accumulate_fp4x8(A_u32_1[0], B_u32_1[0], block_sum);
- block_sum = decode_mul_accumulate_fp4x8(A_u32_1[1], B_u32_1[1], block_sum);
- tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
- }
-
- // SF block 3
- {
- __half scale = __hmul(sfa_scales[3], sfb_scales[3]);
- float block_sum = 0.0f;
- block_sum = decode_mul_accumulate_fp4x8(A_u32_1[2], B_u32_1[2], block_sum);
- block_sum = decode_mul_accumulate_fp4x8(A_u32_1[3], B_u32_1[3], block_sum);
- tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
- }
-
- return tile_sum;
- }
-
extern "C" __global__ __launch_bounds__(128, 8)
void block_scaled_gemv_fp4_fp8_fp16_optimized(
const uint8_t* __restrict__ A,
⋯ 24 unchanged lines
float local_sum = 0.0f;
- // Main loop - process 2 tiles per iteration for better ILP
+ // Main loop - process 2 tiles per iteration, fully inlined
int tile = tidx;
- #pragma unroll
+ #pragma unroll 1
for (; tile + THREADS_PER_ROW < num_k_tiles; tile += 2 * THREADS_PER_ROW) {
// Load tile 0
const int k_offset_0 = tile * TILE_K;
- const float4* A_ptr_0 = reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1));
- const float4* B_ptr_0 = reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1));
-
- float4 A_data0_t0 = __ldg(A_ptr_0);
- float4 B_data0_t0 = __ldg(B_ptr_0);
- float4 A_data1_t0 = __ldg(A_ptr_0 + 1);
- float4 B_data1_t0 = __ldg(B_ptr_0 + 1);
+ float4 A_data0_t0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1)));
+ float4 B_data0_t0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1)));
+ float4 A_data1_t0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1) + 16));
+ float4 B_data1_t0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1) + 16));
uint32_t sfa_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_0 >> 4)));
uint32_t sfb_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_0 >> 4)));
// Load tile 1
const int k_offset_1 = (tile + THREADS_PER_ROW) * TILE_K;
- const float4* A_ptr_1 = reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1));
- const float4* B_ptr_1 = reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1));
-
- float4 A_data0_t1 = __ldg(A_ptr_1);
- float4 B_data0_t1 = __ldg(B_ptr_1);
- float4 A_data1_t1 = __ldg(A_ptr_1 + 1);
- float4 B_data1_t1 = __ldg(B_ptr_1 + 1);
+ float4 A_data0_t1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1)));
+ float4 B_data0_t1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1)));
+ float4 A_data1_t1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1) + 16));
+ float4 B_data1_t1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1) + 16));
uint32_t sfa_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_1 >> 4)));
uint32_t sfb_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_1 >> 4)));
- // Process tile 0
- __half sfa_scales_t0[4], sfb_scales_t0[4];
- decode_fp8x4_e4m3fn_half4(sfa_vec_t0, sfa_scales_t0[0], sfa_scales_t0[1], sfa_scales_t0[2], sfa_scales_t0[3]);
- decode_fp8x4_e4m3fn_half4(sfb_vec_t0, sfb_scales_t0[0], sfb_scales_t0[1], sfb_scales_t0[2], sfb_scales_t0[3]);
+ // Decode FP8 scales for tile 0
+ uint32_t sfa_h2_0, sfa_h2_1, sfb_h2_0, sfb_h2_1;
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
+ : "r"(sfa_vec_t0)
+ );
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
+ : "r"(sfb_vec_t0)
+ );
+ __half2 sfa_scales_t0_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
+ __half2 sfa_scales_t0_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
+ __half2 sfb_scales_t0_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
+ __half2 sfb_scales_t0_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);
+
+ // Process tile 0 - all 4 SF blocks
+ const uint32_t* A_u32 = reinterpret_cast<const uint32_t*>(&A_data0_t0);
+ const uint32_t* B_u32 = reinterpret_cast<const uint32_t*>(&B_data0_t0);
+ float tile_sum_0 = 0.0f;
- local_sum += process_tile(
- reinterpret_cast<const uint32_t*>(&A_data0_t0),
- reinterpret_cast<const uint32_t*>(&B_data0_t0),
- reinterpret_cast<const uint32_t*>(&A_data1_t0),
- reinterpret_cast<const uint32_t*>(&B_data1_t0),
- sfa_scales_t0, sfb_scales_t0);
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_t0_0.x : sfa_scales_t0_0.y,
+ sf == 0 ? sfb_scales_t0_0.x : sfb_scales_t0_0.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum_0 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_0);
+ }
+
+ A_u32 = reinterpret_cast<const uint32_t*>(&A_data1_t0);
+ B_u32 = reinterpret_cast<const uint32_t*>(&B_data1_t0);
+
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_t0_1.x : sfa_scales_t0_1.y,
+ sf == 0 ? sfb_scales_t0_1.x : sfb_scales_t0_1.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum_0 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_0);
+ }
+
+ local_sum += tile_sum_0;
- // Process tile 1
- __half sfa_scales_t1[4], sfb_scales_t1[4];
- decode_fp8x4_e4m3fn_half4(sfa_vec_t1, sfa_scales_t1[0], sfa_scales_t1[1], sfa_scales_t1[2], sfa_scales_t1[3]);
- decode_fp8x4_e4m3fn_half4(sfb_vec_t1, sfb_scales_t1[0], sfb_scales_t1[1], sfb_scales_t1[2], sfb_scales_t1[3]);
+ // Decode and process tile 1 (same as tile 0)
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
+ : "r"(sfa_vec_t1)
+ );
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
+ : "r"(sfb_vec_t1)
+ );
+ __half2 sfa_scales_t1_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
+ __half2 sfa_scales_t1_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
+ __half2 sfb_scales_t1_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
+ __half2 sfb_scales_t1_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);
+
+ A_u32 = reinterpret_cast<const uint32_t*>(&A_data0_t1);
+ B_u32 = reinterpret_cast<const uint32_t*>(&B_data0_t1);
+ float tile_sum_1 = 0.0f;
- local_sum += process_tile(
- reinterpret_cast<const uint32_t*>(&A_data0_t1),
- reinterpret_cast<const uint32_t*>(&B_data0_t1),
- reinterpret_cast<const uint32_t*>(&A_data1_t1),
- reinterpret_cast<const uint32_t*>(&B_data1_t1),
- sfa_scales_t1, sfb_scales_t1);
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_t1_0.x : sfa_scales_t1_0.y,
+ sf == 0 ? sfb_scales_t1_0.x : sfb_scales_t1_0.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum_1 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_1);
+ }
+
+ A_u32 = reinterpret_cast<const uint32_t*>(&A_data1_t1);
+ B_u32 = reinterpret_cast<const uint32_t*>(&B_data1_t1);
+
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_t1_1.x : sfa_scales_t1_1.y,
+ sf == 0 ? sfb_scales_t1_1.x : sfb_scales_t1_1.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum_1 = __fmaf_rn(__half2float(scale), block_sum, tile_sum_1);
+ }
+
+ local_sum += tile_sum_1;
}
- // Handle remaining tile if odd number
+ // Handle remaining single tile
if (tile < num_k_tiles) {
const int k_offset = tile * TILE_K;
- const float4* A_ptr = reinterpret_cast<const float4*>(A_row + (k_offset >> 1));
- const float4* B_ptr = reinterpret_cast<const float4*>(B_base + (k_offset >> 1));
-
- float4 A_data0 = __ldg(A_ptr);
- float4 B_data0 = __ldg(B_ptr);
- float4 A_data1 = __ldg(A_ptr + 1);
- float4 B_data1 = __ldg(B_ptr + 1);
+ float4 A_data0 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset >> 1)));
+ float4 B_data0 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset >> 1)));
+ float4 A_data1 = __ldg(reinterpret_cast<const float4*>(A_row + (k_offset >> 1) + 16));
+ float4 B_data1 = __ldg(reinterpret_cast<const float4*>(B_base + (k_offset >> 1) + 16));
uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset >> 4)));
uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset >> 4)));
- __half sfa_scales[4], sfb_scales[4];
- decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);
- decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);
+ uint32_t sfa_h2_0, sfa_h2_1, sfb_h2_0, sfb_h2_1;
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfa_h2_0), "=r"(sfa_h2_1)
+ : "r"(sfa_vec)
+ );
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(sfb_h2_0), "=r"(sfb_h2_1)
+ : "r"(sfb_vec)
+ );
+ __half2 sfa_scales_0 = *reinterpret_cast<const __half2*>(&sfa_h2_0);
+ __half2 sfa_scales_1 = *reinterpret_cast<const __half2*>(&sfa_h2_1);
+ __half2 sfb_scales_0 = *reinterpret_cast<const __half2*>(&sfb_h2_0);
+ __half2 sfb_scales_1 = *reinterpret_cast<const __half2*>(&sfb_h2_1);
- local_sum += process_tile(
- reinterpret_cast<const uint32_t*>(&A_data0),
- reinterpret_cast<const uint32_t*>(&B_data0),
- reinterpret_cast<const uint32_t*>(&A_data1),
- reinterpret_cast<const uint32_t*>(&B_data1),
- sfa_scales, sfb_scales);
+ const uint32_t* A_u32 = reinterpret_cast<const uint32_t*>(&A_data0);
+ const uint32_t* B_u32 = reinterpret_cast<const uint32_t*>(&B_data0);
+ float tile_sum = 0.0f;
+
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_0.x : sfa_scales_0.y,
+ sf == 0 ? sfb_scales_0.x : sfb_scales_0.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
+ }
+
+ A_u32 = reinterpret_cast<const uint32_t*>(&A_data1);
+ B_u32 = reinterpret_cast<const uint32_t*>(&B_data1);
+
+ #pragma unroll
+ for (int sf = 0; sf < 2; sf++) {
+ __half scale = __hmul(
+ sf == 0 ? sfa_scales_1.x : sfa_scales_1.y,
+ sf == 0 ? sfb_scales_1.x : sfb_scales_1.y
+ );
+ float block_sum;
+ asm volatile (
+ "{"
+ " .reg .b8 %%ab<4>, %%bb<4>;\\n"
+ " .reg .b32 %%a<4>, %%b<4>;\\n"
+ " .reg .b32 %%p0, %%p1;\\n"
+ " .reg .f16 %%h0, %%h1;\\n"
+ " .reg .f32 %%f0, %%f1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
+ " mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %3;\\n"
+ " mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
+ " mul.rn.f16x2 %%p1, %%a0, %%b0;\\n"
+ " fma.rn.f16x2 %%p1, %%a1, %%b1, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a2, %%b2, %%p1;\\n"
+ " fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
+ " add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
+ " mov.b32 {%%h0, %%h1}, %%p0;\\n"
+ " cvt.f32.f16 %%f0, %%h0;\\n"
+ " cvt.f32.f16 %%f1, %%h1;\\n"
+ " add.f32 %0, %%f0, %%f1;\\n"
+ "}"
+ : "=f"(block_sum)
+ : "r"(A_u32[sf * 2]), "r"(B_u32[sf * 2]),
+ "r"(A_u32[sf * 2 + 1]), "r"(B_u32[sf * 2 + 1])
+ );
+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
+ }
+
+ local_sum += tile_sum;
}
// Warp-level reduction
⋯ 56 unchanged lines
"""
_gemm_module = load_inline(
- name="block_scaled_gemv_ilp_v1",
+ name="block_scaled_gemv_inline_v1",
cpp_sources=gemv_cpp_src,
cuda_sources=gemv_cuda_src,
functions=["gemv_fp4_fp8_fp16"],
⋯ 13 unchanged lines
def custom_kernel(data):
"""
- Optimized version focusing on ILP without shared memory:
- - Process 2 tiles per iteration to increase ILP
- - All loads issued together, then all computes
- - __launch_bounds__ for occupancy hint
- - No shared memory overhead
- - Direct global -> register path via __ldg
+ Fully inlined version with 2 tiles:
+ - No function calls, everything inlined directly
+ - 2 tiles to keep register pressure low
+ - Combined FP4 operations for 2 uint32_t at once
+ - Should have minimal overhead
"""
a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
⋯ 6 unchanged lines
_gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)
- return c
No newline at end of file
+ return c
+
scrolls · 693 diff lines total

Best evidence level for this revision: reported

JSON