Skip to content
KernelIndex
Search⌘K

submission 114783

agokrani · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_gemv_v17_swizzled.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-114783?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
70.3µs
#328 of 678
2025-11-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ff6228eb62074071cf7ceb58f03c17ed732c553e787d685cc04ced30cbea4d5f
license declaredunknown
license concludedunknown
authorsagokrani
imported2026-08-26

Techniques

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

async-copy- "Uncoalesced shared accesses" in cp_async_cg_16 affects ALL shapes
fp4NVFP4 Block-Scaled GEMV with XOR Swizzled Shared Memory Access.
shared-memoryextern __shared__ uint8_t smem_raw[];
vector-width = half2half2 fallback = __float22half2_rn(make_float2(0.0f, 0.0f));

Kernel source

nvfp4_gemv_v17_swizzled.py747 lines
"""
NVFP4 Block-Scaled GEMV with XOR Swizzled Shared Memory Access.

VERSION: v17_swizzled
Building on v15, this version eliminates shared memory bank conflicts
using XOR-based swizzling for A/B matrices and padding for scale arrays.

Key insight from profiling:
- "Uncoalesced shared accesses" in cp_async_cg_16 affects ALL shapes
- 16-byte stride creates 4-way bank conflicts (threads 0,8,16,24 conflict)
- For L=1, this cascades into occupancy drops and scoreboard stalls

Fix:
- XOR swizzle: index ^ (index >> 3) spreads threads across bank groups
- Padding: +4 floats per row in SFA/SFB eliminates row-based conflicts

Bank conflict analysis (before fix):
  Thread 0:  banks 0-3
  Thread 8:  banks 0-3  <- CONFLICT!
  Thread 16: banks 0-3  <- CONFLICT!
  Thread 24: banks 0-3  <- CONFLICT!

After XOR swizzle (thread_id ^ (thread_id >> 3)):
  Thread 0  -> swizzled 0:  banks 0-3
  Thread 8  -> swizzled 9:  banks 4-7   <- NO CONFLICT!
  Thread 16 -> swizzled 18: banks 8-11  <- NO CONFLICT!
  Thread 24 -> swizzled 27: banks 12-15 <- NO CONFLICT!
"""

import torch
import os
from pathlib import Path
from torch.utils.cpp_extension import load_inline

KERNEL_SOURCE = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_pipeline.h>

// ============================================================================
// Configuration Constants
// ============================================================================

// Standard config (for batched shapes with L >= 2)
constexpr int THREADS_PER_BLOCK = 128;
constexpr int THREADS_PER_ROW_STD = 16;
constexpr int ROWS_PER_BLOCK_STD = THREADS_PER_BLOCK / THREADS_PER_ROW_STD;  // 8
constexpr int ELEMENTS_PER_ACCESS = 32;
constexpr int BYTES_PER_ACCESS = 16;
constexpr int SF_VEC_SIZE = 16;

// Wide-K config (for L=1 large-K shapes - better tail utilization)
constexpr int THREADS_PER_ROW_WIDE = 32;
constexpr int ROWS_PER_BLOCK_WIDE = THREADS_PER_BLOCK / THREADS_PER_ROW_WIDE;  // 4

// Small K config (for K < 512)
constexpr int THREADS_PER_ROW_SMALL = 4;
constexpr int ROWS_PER_BLOCK_SMALL = THREADS_PER_BLOCK / THREADS_PER_ROW_SMALL;  // 32

// Padding for scale arrays to avoid bank conflicts between rows
// Each row gets +4 floats (16 bytes) padding to shift bank alignment
constexpr int SFA_PAD = 4;

// ============================================================================
// Helper Functions
// ============================================================================

__device__ __forceinline__ void cp_async_cg_16(void* dst, const void* src) {
    uint32_t dst_smem = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
    asm volatile(
        "cp.async.cg.shared.global.L2::128B [%0], [%1], 16;"
        : : "r"(dst_smem), "l"(src)
        : "memory");
}

__device__ __forceinline__ void commit_group() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void wait_group_0() { asm volatile("cp.async.wait_group 0;"); }

__device__ __forceinline__ uint32_t decode_fp4_hw(uint32_t packed_byte) {
    uint32_t result;
    #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 100
        asm volatile(
            "{"
            ".reg .b8 packed_b8;"
            "cvt.u8.u32 packed_b8, %1;"
            "cvt.rn.f16x2.e2m1x2 %0, packed_b8;"
            "}"
            : "=r"(result)
            : "r"(packed_byte)
        );
    #else
        half2 fallback = __float22half2_rn(make_float2(0.0f, 0.0f));
        result = *reinterpret_cast<uint32_t*>(&fallback);
    #endif
    return result;
}

__device__ __forceinline__ half2 to_half2(uint32_t x) {
    return *reinterpret_cast<half2*>(&x);
}

// XOR swizzle to eliminate bank conflicts
// Maps threads 0,8,16,24 to different bank groups instead of all hitting banks 0-3
__device__ __forceinline__ int xor_swizzle(int idx) {
    return idx ^ (idx >> 3);
}

// ============================================================================
// Shared Memory Structures with Padding
// ============================================================================

// Standard config shared memory
// A and B use XOR swizzle for access, no structural change needed
// SFA/SFB get padding to avoid row conflicts
struct alignas(128) SmemBufferStdSwizzled {
    uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS];  // 2048 bytes
    uint8_t B[THREADS_PER_ROW_STD * BYTES_PER_ACCESS];    // 256 bytes
    float SFA[ROWS_PER_BLOCK_STD][THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 8 x 36
    float SFB[THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 36
};

// Wide-K config shared memory
struct alignas(128) SmemBufferWideKSwizzled {
    uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS];  // 2048 bytes
    uint8_t B[THREADS_PER_ROW_WIDE * BYTES_PER_ACCESS];   // 512 bytes
    float SFA[ROWS_PER_BLOCK_WIDE][THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 4 x 68
    float SFB[THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 68
};

// Small K config shared memory
struct alignas(128) SmemBufferSmallKSwizzled {
    uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS];
    uint8_t B[THREADS_PER_ROW_SMALL * BYTES_PER_ACCESS];
    float SFA[ROWS_PER_BLOCK_SMALL][THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 32 x 12
    float SFB[THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD];  // 12
};

// ============================================================================
// Small K Kernel (K < 512) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_small_k(
    const uint8_t* __restrict__ A_packed,
    const uint8_t* __restrict__ B_packed,
    const float* __restrict__ SFA,
    const float* __restrict__ SFB,
    half* __restrict__ C,
    const int M,
    const int K,
    const int L,
    const int num_m_tiles,
    const int stride_A_m,
    const int stride_A_l,
    const int stride_B_l,
    const int stride_SFA_0,
    const int stride_SFA_1,
    const int stride_SFA_2,
    const int stride_SFA_3,
    const int stride_SFA_4,
    const int stride_SFA_5,
    const int stride_SFB_3,
    const int stride_SFB_4,
    const int stride_SFB_5,
    const int stride_C_m,
    const int stride_C_l
) {
    extern __shared__ uint8_t smem_raw[];
    SmemBufferSmallKSwizzled* smem = reinterpret_cast<SmemBufferSmallKSwizzled*>(smem_raw);

    const int thread_id = threadIdx.x;
    const int k_thread = thread_id % THREADS_PER_ROW_SMALL;
    const int m_local = thread_id / THREADS_PER_ROW_SMALL;

    // XOR swizzle indices for bank conflict avoidance
    const int sw_tid = xor_swizzle(thread_id);
    const int sw_k = xor_swizzle(k_thread);

    const int batch_idx = (gridDim.z > 1) ? blockIdx.z : (blockIdx.x / num_m_tiles);
    const int m_tile_idx = (gridDim.z > 1) ? blockIdx.x : (blockIdx.x % num_m_tiles);

    if (m_tile_idx >= num_m_tiles) return;

    const int m_base = m_tile_idx * ROWS_PER_BLOCK_SMALL;
    const int m_row = m_base + m_local;

    if (m_row >= M) return;

    const uint8_t* ptr_A = A_packed + batch_idx * stride_A_l;
    const uint8_t* ptr_B = B_packed + batch_idx * stride_B_l;
    half* ptr_C = C + batch_idx * stride_C_l + m_row * stride_C_m;

    const int mm32 = m_row % 32;
    const int mm4 = (m_row / 32) % 4;
    const int mm = m_row / 128;
    const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 +
                                mm * stride_SFA_2 + batch_idx * stride_SFA_5;
    const int base_sfb_offset = batch_idx * stride_SFB_5;

    const int tile_k = THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS;  // 128
    const int tile_k_bytes = tile_k / 2;
    const int total_k_tiles = K / tile_k;
    const int scales_per_tile = tile_k / SF_VEC_SIZE;

    float accum = 0.0f;

    for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
        const int k_byte_offset = tile_k_idx * tile_k_bytes;
        const int scale_k_base = tile_k_idx * scales_per_tile;

        // Load A with XOR swizzle
        const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
        cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);

        // Load B with XOR swizzle
        if (m_local == 0) {
            cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
        }

        // Load SFA (with padding in array)
        if (k_thread < 2) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
            cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
        }

        // Load SFB (with padding in array)
        if (m_local == 0 && k_thread < 2) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
            cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
        }

        commit_group();
        wait_group_0();
        __syncthreads();

        // Read with XOR swizzle (same indices as write)
        const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
        const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
        uint4 va = *vec_a;
        uint4 vb = *vec_b;

        half2 acc_lo = __float2half2_rn(0.0f);
        half2 acc_hi = __float2half2_rn(0.0f);

        #define PROC(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_lo = __hfma2(ha, hb, acc_lo); \
        }
        PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
        PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
        #undef PROC

        #define PROC_HI(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_hi = __hfma2(ha, hb, acc_hi); \
        }
        PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
        PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
        #undef PROC_HI

        float sfa0 = smem->SFA[m_local][k_thread * 2];
        float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
        float sfb0 = smem->SFB[k_thread * 2];
        float sfb1 = smem->SFB[k_thread * 2 + 1];

        float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
        float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
        accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;

        __syncthreads();
    }

    #pragma unroll
    for (int mask = THREADS_PER_ROW_SMALL / 2; mask > 0; mask >>= 1) {
        accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
    }

    if (k_thread == 0) {
        *ptr_C = __float2half(accum);
    }
}

// ============================================================================
// Wide-K Kernel (K >= 8K, L=1) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_wide_k(
    const uint8_t* __restrict__ A_packed,
    const uint8_t* __restrict__ B_packed,
    const float* __restrict__ SFA,
    const float* __restrict__ SFB,
    half* __restrict__ C,
    const int M,
    const int K,
    const int L,
    const int num_m_tiles,
    const int stride_A_m,
    const int stride_A_l,
    const int stride_B_l,
    const int stride_SFA_0,
    const int stride_SFA_1,
    const int stride_SFA_2,
    const int stride_SFA_3,
    const int stride_SFA_4,
    const int stride_SFA_5,
    const int stride_SFB_3,
    const int stride_SFB_4,
    const int stride_SFB_5,
    const int stride_C_m,
    const int stride_C_l
) {
    extern __shared__ uint8_t smem_raw[];
    SmemBufferWideKSwizzled* smem = reinterpret_cast<SmemBufferWideKSwizzled*>(smem_raw);

    const int thread_id = threadIdx.x;
    const int k_thread = thread_id % THREADS_PER_ROW_WIDE;  // 0-31
    const int m_local = thread_id / THREADS_PER_ROW_WIDE;   // 0-3

    // XOR swizzle indices
    const int sw_tid = xor_swizzle(thread_id);
    const int sw_k = xor_swizzle(k_thread);

    const int m_tile_idx = blockIdx.x;
    if (m_tile_idx >= num_m_tiles) return;

    const int m_base = m_tile_idx * ROWS_PER_BLOCK_WIDE;
    const int m_row = m_base + m_local;
    if (m_row >= M) return;

    const uint8_t* ptr_A = A_packed;
    const uint8_t* ptr_B = B_packed;
    half* ptr_C = C + m_row * stride_C_m;

    const int mm32 = m_row % 32;
    const int mm4 = (m_row / 32) % 4;
    const int mm = m_row / 128;
    const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 + mm * stride_SFA_2;
    const int base_sfb_offset = 0;

    const int tile_k = THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS;  // 1024
    const int tile_k_bytes = tile_k / 2;
    const int total_k_tiles = K / tile_k;
    const int scales_per_tile = tile_k / SF_VEC_SIZE;  // 64

    float accum = 0.0f;

    for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
        const int k_byte_offset = tile_k_idx * tile_k_bytes;
        const int scale_k_base = tile_k_idx * scales_per_tile;

        // Load A with XOR swizzle
        const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
        cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);

        // Load B with XOR swizzle
        if (m_local == 0) {
            cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
        }

        // Load SFA (16 threads load 4 floats each = 64 scales)
        if (k_thread < 16) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
            cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
        }

        // Load SFB
        if (m_local == 0 && k_thread < 16) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
            cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
        }

        commit_group();
        wait_group_0();
        __syncthreads();

        // Read with XOR swizzle
        const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
        const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
        uint4 va = *vec_a;
        uint4 vb = *vec_b;

        half2 acc_lo = __float2half2_rn(0.0f);
        half2 acc_hi = __float2half2_rn(0.0f);

        #define PROC(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_lo = __hfma2(ha, hb, acc_lo); \
        }
        PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
        PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
        #undef PROC

        #define PROC_HI(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_hi = __hfma2(ha, hb, acc_hi); \
        }
        PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
        PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
        #undef PROC_HI

        float sfa0 = smem->SFA[m_local][k_thread * 2];
        float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
        float sfb0 = smem->SFB[k_thread * 2];
        float sfb1 = smem->SFB[k_thread * 2 + 1];

        float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
        float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
        accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;

        __syncthreads();
    }

    // Warp reduction
    #pragma unroll
    for (int mask = THREADS_PER_ROW_WIDE / 2; mask > 0; mask >>= 1) {
        accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
    }

    if (k_thread == 0) {
        *ptr_C = __float2half(accum);
    }
}

// ============================================================================
// Standard Kernel (K >= 512, L >= 2) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_std(
    const uint8_t* __restrict__ A_packed,
    const uint8_t* __restrict__ B_packed,
    const float* __restrict__ SFA,
    const float* __restrict__ SFB,
    half* __restrict__ C,
    const int M,
    const int K,
    const int L,
    const int num_m_tiles,
    const int stride_A_m,
    const int stride_A_l,
    const int stride_B_l,
    const int stride_SFA_0,
    const int stride_SFA_1,
    const int stride_SFA_2,
    const int stride_SFA_3,
    const int stride_SFA_4,
    const int stride_SFA_5,
    const int stride_SFB_3,
    const int stride_SFB_4,
    const int stride_SFB_5,
    const int stride_C_m,
    const int stride_C_l
) {
    extern __shared__ uint8_t smem_raw[];
    SmemBufferStdSwizzled* smem = reinterpret_cast<SmemBufferStdSwizzled*>(smem_raw);

    const int thread_id = threadIdx.x;
    const int k_thread = thread_id % THREADS_PER_ROW_STD;
    const int m_local = thread_id / THREADS_PER_ROW_STD;

    // XOR swizzle indices
    const int sw_tid = xor_swizzle(thread_id);
    const int sw_k = xor_swizzle(k_thread);

    const int m_tile_idx = blockIdx.x;
    const int batch_idx = blockIdx.z;

    if (m_tile_idx >= num_m_tiles) return;

    const int m_base = m_tile_idx * ROWS_PER_BLOCK_STD;
    const int m_row = m_base + m_local;

    if (m_row >= M) return;

    const uint8_t* ptr_A = A_packed + batch_idx * stride_A_l;
    const uint8_t* ptr_B = B_packed + batch_idx * stride_B_l;
    half* ptr_C = C + batch_idx * stride_C_l + m_row * stride_C_m;

    const int mm32 = m_row % 32;
    const int mm4 = (m_row / 32) % 4;
    const int mm = m_row / 128;
    const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 +
                                mm * stride_SFA_2 + batch_idx * stride_SFA_5;
    const int base_sfb_offset = batch_idx * stride_SFB_5;

    const int tile_k = THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS;  // 512
    const int tile_k_bytes = tile_k / 2;
    const int total_k_tiles = K / tile_k;
    const int scales_per_tile = tile_k / SF_VEC_SIZE;

    float accum = 0.0f;

    for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
        const int k_byte_offset = tile_k_idx * tile_k_bytes;
        const int scale_k_base = tile_k_idx * scales_per_tile;

        // Load A with XOR swizzle
        const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
        cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);

        // Load B with XOR swizzle
        if (m_local == 0) {
            cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
        }

        // Load SFA
        if (k_thread < 8) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
            cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
        }

        // Load SFB
        if (m_local == 0 && k_thread < 8) {
            int vec_idx = k_thread * 4;
            int scale_idx = scale_k_base + vec_idx;
            int kk4 = scale_idx % 4;
            int kk = scale_idx / 4;
            int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
            cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
        }

        commit_group();
        wait_group_0();
        __syncthreads();

        // Read with XOR swizzle
        const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
        const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
        uint4 va = *vec_a;
        uint4 vb = *vec_b;

        half2 acc_lo = __float2half2_rn(0.0f);
        half2 acc_hi = __float2half2_rn(0.0f);

        #define PROC(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_lo = __hfma2(ha, hb, acc_lo); \
        }
        PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
        PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
        #undef PROC

        #define PROC_HI(Ra, Rb, S) { \
            half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
            half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
            acc_hi = __hfma2(ha, hb, acc_hi); \
        }
        PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
        PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
        #undef PROC_HI

        float sfa0 = smem->SFA[m_local][k_thread * 2];
        float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
        float sfb0 = smem->SFB[k_thread * 2];
        float sfb1 = smem->SFB[k_thread * 2 + 1];

        float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
        float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
        accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;

        __syncthreads();
    }

    #pragma unroll
    for (int mask = THREADS_PER_ROW_STD / 2; mask > 0; mask >>= 1) {
        accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
    }

    if (k_thread == 0) {
        *ptr_C = __float2half(accum);
    }
}

// ============================================================================
// Launcher with Shape-Specialized Dispatch
// ============================================================================
torch::Tensor nvfp4_gemv_v17(
    torch::Tensor A_packed,
    torch::Tensor B_packed,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C
) {
    const int M = C.size(0);
    const int K = A_packed.size(1) * 2;
    const int L = C.size(2);

    const int stride_A_m = A_packed.stride(0);
    const int stride_A_l = A_packed.stride(2);
    const int stride_B_l = B_packed.stride(2);

    const int stride_SFA_0 = SFA.stride(0);
    const int stride_SFA_1 = SFA.stride(1);
    const int stride_SFA_2 = SFA.stride(2);
    const int stride_SFA_3 = SFA.stride(3);
    const int stride_SFA_4 = SFA.stride(4);
    const int stride_SFA_5 = SFA.stride(5);

    const int stride_SFB_3 = SFB.stride(3);
    const int stride_SFB_4 = SFB.stride(4);
    const int stride_SFB_5 = SFB.stride(5);

    const int stride_C_m = C.stride(0);
    const int stride_C_l = C.stride(2);

    constexpr int K_SMALL_THRESHOLD = 512;
    constexpr int K_WIDE_THRESHOLD = 8192;

    if (K < K_SMALL_THRESHOLD) {
        const int num_m_tiles = (M + ROWS_PER_BLOCK_SMALL - 1) / ROWS_PER_BLOCK_SMALL;
        dim3 grid(num_m_tiles, 1, L);
        dim3 block(THREADS_PER_BLOCK);
        int smem_bytes = sizeof(SmemBufferSmallKSwizzled);

        nvfp4_gemv_kernel_v17_small_k<<<grid, block, smem_bytes>>>(
            A_packed.data_ptr<uint8_t>(),
            B_packed.data_ptr<uint8_t>(),
            SFA.data_ptr<float>(),
            SFB.data_ptr<float>(),
            reinterpret_cast<half*>(C.data_ptr<at::Half>()),
            M, K, L,
            num_m_tiles,
            stride_A_m, stride_A_l, stride_B_l,
            stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
            stride_SFB_3, stride_SFB_4, stride_SFB_5,
            stride_C_m, stride_C_l
        );
    } else if (L == 1 && K >= K_WIDE_THRESHOLD) {
        const int num_m_tiles = (M + ROWS_PER_BLOCK_WIDE - 1) / ROWS_PER_BLOCK_WIDE;
        dim3 grid(num_m_tiles);
        dim3 block(THREADS_PER_BLOCK);
        int smem_bytes = sizeof(SmemBufferWideKSwizzled);

        nvfp4_gemv_kernel_v17_wide_k<<<grid, block, smem_bytes>>>(
            A_packed.data_ptr<uint8_t>(),
            B_packed.data_ptr<uint8_t>(),
            SFA.data_ptr<float>(),
            SFB.data_ptr<float>(),
            reinterpret_cast<half*>(C.data_ptr<at::Half>()),
            M, K, L,
            num_m_tiles,
            stride_A_m, stride_A_l, stride_B_l,
            stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
            stride_SFB_3, stride_SFB_4, stride_SFB_5,
            stride_C_m, stride_C_l
        );
    } else {
        const int num_m_tiles = (M + ROWS_PER_BLOCK_STD - 1) / ROWS_PER_BLOCK_STD;
        dim3 grid(num_m_tiles, 1, L);
        dim3 block(THREADS_PER_BLOCK);
        int smem_bytes = sizeof(SmemBufferStdSwizzled);

        nvfp4_gemv_kernel_v17_std<<<grid, block, smem_bytes>>>(
            A_packed.data_ptr<uint8_t>(),
            B_packed.data_ptr<uint8_t>(),
            SFA.data_ptr<float>(),
            SFB.data_ptr<float>(),
            reinterpret_cast<half*>(C.data_ptr<at::Half>()),
            M, K, L,
            num_m_tiles,
            stride_A_m, stride_A_l, stride_B_l,
            stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
            stride_SFB_3, stride_SFB_4, stride_SFB_5,
            stride_C_m, stride_C_l
        );
    }

    return C;
}

'''

CPP_SOURCE = r'''
#include <torch/extension.h>

torch::Tensor nvfp4_gemv_v17(
    torch::Tensor A_packed,
    torch::Tensor B_packed,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C
);
'''

_extension = None

def _get_extension():
    global _extension
    if _extension is None:
        import hashlib
        code_hash = hashlib.md5(KERNEL_SOURCE.encode()).hexdigest()[:10]

        print(f"Building XOR-swizzled kernel v17: nvfp4_gemv_v17_{code_hash}")
        print("  Key optimization: XOR swizzle eliminates 4-way bank conflicts")
        print("  Before: threads 0,8,16,24 all hit banks 0-3 (4-way conflict)")
        print("  After:  threads map to banks 0-3, 4-7, 8-11, 12-15 (no conflict)")

        _extension = load_inline(
            name=f"nvfp4_gemv_v17_{code_hash}",
            cpp_sources=[CPP_SOURCE],
            cuda_sources=[KERNEL_SOURCE],
            functions=["nvfp4_gemv_v17"],
            extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo",
                             "-std=c++17", "-gencode", "arch=compute_100a,code=sm_100a"],
            extra_ldflags=["-lcuda"],
            with_cuda=True,
            verbose=False
        )
        print("Compilation successful!")
    return _extension

def custom_kernel(input_tuple):
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = input_tuple

    ext = _get_extension()

    a_bytes = a.view(torch.uint8)
    b_bytes = b.view(torch.uint8)
    sfa_f32 = sfa_permuted.to(torch.float32)
    sfb_f32 = sfb_permuted.to(torch.float32)

    c_out = c.clone()
    ext.nvfp4_gemv_v17(a_bytes, b_bytes, sfa_f32, sfb_f32, c_out)

    return c_out
scrolls · 747 lines total

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

Best evidence level for this revision: reported

JSON