Skip to content
KernelIndex
Search⌘K

submission 109980

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109980?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
#122 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6b570c84858031bba9935e2cb012f78c345de7f1edda0390da304074bb46e4c3
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

async-copyasm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \
shared-memory__shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];
vector-width = half2__device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {

Kernel source

submission.py452 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# Combine both v13 (direct) and v14 (pipelined) kernels with selection logic
CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>

// =============================================================================
// Common Utilities
// =============================================================================

enum class CacheHint { CS, CA };

template<CacheHint hint>
__device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {
    if constexpr (hint == CacheHint::CS) {
        asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
            : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
    } else {
        asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
            : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
    }
}

template<CacheHint hint>
__device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {
    if constexpr (hint == CacheHint::CS) {
        asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
    } else {
        asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
    }
}

__device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {
    asm volatile(
        "{\n\t"
        ".reg .b8 b0, b1, b2, b3;\n\t"
        "mov.b32 {b0, b1, b2, b3}, %4;\n\t"
        "cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
        "cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
        "cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
        "cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
        "}"
        : "=r"(reinterpret_cast<int&>(out[0])),
          "=r"(reinterpret_cast<int&>(out[1])),
          "=r"(reinterpret_cast<int&>(out[2])),
          "=r"(reinterpret_cast<int&>(out[3]))
        : "r"(in)
    );
}

__device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {
    int out;
    asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
    return *reinterpret_cast<half2*>(&out);
}

// =============================================================================
// DIRECT KERNEL (from v13)
// =============================================================================

__device__ __forceinline__ float dot_scaled_direct(
    const uint32_t* a_data,
    const uint32_t* b_data,
    half2 scale
) {
    half2 a_f16[16], b_f16[16];
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
        cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
    }
    
    half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
    half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
    
    #pragma unroll
    for (int i = 1; i < 8; i++) {
        acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
        acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
    }
    
    half sum0 = __hadd(acc0.x, acc0.y);
    half sum1 = __hadd(acc1.x, acc1.y);
    
    float result = 0.0f;
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" 
        : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" 
        : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
    
    return result;
}

template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CacheHint A_HINT, CacheHint B_HINT>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_direct_kernel(
    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, int N_pad
) {
    const int row_base = blockIdx.x * ROWS_PER_BLOCK;
    const int batch = blockIdx.y;
    const int tid = threadIdx.x;
    
    const int warp_id = tid / THREADS_PER_ROW;
    const int lane = tid % THREADS_PER_ROW;
    
    const int row = row_base + warp_id;
    if (row >= M) return;
    
    const int K_bytes = K / 2;
    const int K_sf = K / 16;
    
    const size_t a_off = (size_t)batch * M * K_bytes + row * K_bytes;
    const size_t b_off = (size_t)batch * N_pad * K_bytes;
    const size_t sfa_off = (size_t)batch * M * K_sf + row * K_sf;
    const size_t sfb_off = (size_t)batch * N_pad * K_sf;
    
    const uint8_t* a_ptr = a + a_off;
    const uint8_t* b_ptr = b + b_off;
    const uint8_t* sfa_ptr = sfa + sfa_off;
    const uint8_t* sfb_ptr = sfb + sfb_off;
    
    float acc = 0.0f;
    const int num_sf_pairs = K_sf / 2;
    
    #pragma unroll 4
    for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
        const int sf_idx = sp * 2;
        const int byte_idx = sf_idx * 8;
        
        uint16_t sfa_packed, sfb_packed;
        load_u16<A_HINT>(&sfa_packed, sfa_ptr + sf_idx);
        load_u16<B_HINT>(&sfb_packed, sfb_ptr + sf_idx);
        half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
        
        uint32_t a_regs[4], b_regs[4];
        load_vec4<A_HINT>(a_regs, a_ptr + byte_idx);
        load_vec4<B_HINT>(b_regs, b_ptr + byte_idx);
        
        acc += dot_scaled_direct(a_regs, b_regs, scale);
    }
    
    #pragma unroll
    for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }
    
    if (lane == 0) {
        c[(size_t)row + (size_t)batch * M] = __float2half(acc);
    }
}

// =============================================================================
// PIPELINED KERNEL (from v14)
// =============================================================================

#define ASYNC_CP_16(dst, src) \
    asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \
        :: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_CP_4(dst, src) \
    asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" \
        :: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
#define ASYNC_WAIT(n) asm volatile("cp.async.wait_group %0;" :: "n"(n))

#define NUM_STAGES 2

__device__ __forceinline__ float dot_scaled_pipelined(
    const uint8_t* a_smem,
    const uint8_t* b_smem,
    half2 scale
) {
    const uint32_t* a_data = reinterpret_cast<const uint32_t*>(a_smem);
    const uint32_t* b_data = reinterpret_cast<const uint32_t*>(b_smem);
    
    half2 a_f16[16], b_f16[16];
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
        cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
    }
    
    half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
    half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
    
    #pragma unroll
    for (int i = 1; i < 8; i++) {
        acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
        acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
    }
    
    half sum0 = __hadd(acc0.x, acc0.y);
    half sum1 = __hadd(acc1.x, acc1.y);
    
    float result = 0.0f;
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" 
        : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" 
        : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
    
    return result;
}

template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, int K_TILE>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_pipelined_kernel(
    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, int N_pad
) {
    constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW;
    constexpr int K_TILE_BYTES = K_TILE / 2;
    constexpr int K_TILE_SF = K_TILE / 16;
    
    __shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];
    __shared__ __align__(16) uint8_t sh_b[NUM_STAGES][K_TILE_BYTES];
    __shared__ __align__(16) uint8_t sh_sfa[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_SF];
    __shared__ __align__(16) uint8_t sh_sfb[NUM_STAGES][K_TILE_SF];
    
    const int row_base = blockIdx.x * ROWS_PER_BLOCK;
    const int batch = blockIdx.y;
    const int tid = threadIdx.x;
    
    const int warp_id = tid / THREADS_PER_ROW;
    const int lane = tid % THREADS_PER_ROW;
    
    const int row = row_base + warp_id;
    if (row >= M) return;
    
    const int K_bytes = K / 2;
    const int K_sf = K / 16;
    const int num_tiles = (K_bytes + K_TILE_BYTES - 1) / K_TILE_BYTES;
    
    const size_t a_base = (size_t)batch * M * K_bytes;
    const size_t b_base = (size_t)batch * N_pad * K_bytes;
    const size_t sfa_base = (size_t)batch * M * K_sf;
    const size_t sfb_base = (size_t)batch * N_pad * K_sf;
    
    auto issue_tile = [&](int tile, int stage) {
        const int k_off_bytes = tile * K_TILE_BYTES;
        const int k_off_sf = tile * K_TILE_SF;
        const int bytes_this = min(K_TILE_BYTES, K_bytes - k_off_bytes);
        const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
        
        if (warp_id < ROWS_PER_BLOCK && (row_base + warp_id) < M) {
            const uint8_t* a_row = a + a_base + (row_base + warp_id) * K_bytes + k_off_bytes;
            for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
                if (i + 16 <= bytes_this) {
                    ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);
                }
            }
            const uint8_t* sfa_row = sfa + sfa_base + (row_base + warp_id) * K_sf + k_off_sf;
            for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
                if (i + 4 <= sf_this) {
                    ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);
                }
            }
        }
        
        if (warp_id == 0) {
            const uint8_t* b_ptr = b + b_base + k_off_bytes;
            for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
                if (i + 16 <= bytes_this) {
                    ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);
                }
            }
            const uint8_t* sfb_ptr = sfb + sfb_base + k_off_sf;
            for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
                if (i + 4 <= sf_this) {
                    ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);
                }
            }
        }
        ASYNC_COMMIT();
    };
    
    issue_tile(0, 0);
    ASYNC_WAIT(0);
    __syncthreads();
    
    float acc = 0.0f;
    
    for (int tile = 0; tile < num_tiles; tile++) {
        const int stage = tile % NUM_STAGES;
        const int next_stage = (tile + 1) % NUM_STAGES;
        
        if (tile + 1 < num_tiles) {
            issue_tile(tile + 1, next_stage);
        }
        
        const int k_off_sf = tile * K_TILE_SF;
        const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
        const int num_sf_pairs = sf_this / 2;
        
        #pragma unroll 2
        for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
            const int sf_idx = sp * 2;
            const int byte_idx = sf_idx * 8;
            
            uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);
            uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);
            half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
            
            acc += dot_scaled_pipelined(
                &sh_a[stage][warp_id][byte_idx],
                &sh_b[stage][byte_idx],
                scale
            );
        }
        
        if (tile + 1 < num_tiles) {
            ASYNC_WAIT(0);
            __syncthreads();
        }
    }
    
    #pragma unroll
    for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }
    
    if (lane == 0) {
        c[(size_t)row + (size_t)batch * M] = __float2half(acc);
    }
}

// =============================================================================
// Dispatch
// =============================================================================

void run_nvfp4_gemv(
    torch::Tensor C, torch::Tensor A, torch::Tensor B,
    torch::Tensor SFA, torch::Tensor SFB,
    int m, int k, int l, int n_pad
) {
    constexpr auto CS = CacheHint::CS;
    constexpr auto CA = CacheHint::CA;
    
    // Selection heuristic based on benchmarks:
    // - k=16384, l=1: pipelined faster (24.6 vs 30.8)
    // - k=7168, l=8:  direct faster (38.5 vs 51.0)
    // - k=2048, l=4:  same (16.4)
    
    const bool use_pipelined = (k >= 8192 && l <= 2) ||  // Large K, small batch
                               (k <= 4096 && l >= 4);    // Small K, larger batch
    
    const int k_bytes = k / 2;
    const bool b_fits_l1 = (k_bytes <= 32 * 1024);
    
    auto* a_ptr = A.data_ptr<uint8_t>();
    auto* b_ptr = B.data_ptr<uint8_t>();
    auto* sfa_ptr = SFA.data_ptr<uint8_t>();
    auto* sfb_ptr = SFB.data_ptr<uint8_t>();
    auto* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
    
    if (use_pipelined) {
        // Pipelined kernel - configs from v14
        if (k >= 8192) {
            dim3 grid((m + 4 - 1) / 4, l);
            dim3 block(4 * 32);
            gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(
                a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
        } else if (k >= 4096) {
            dim3 grid((m + 4 - 1) / 4, l);
            dim3 block(4 * 32);
            gemv_pipelined_kernel<4, 32, 1024><<<grid, block>>>(
                a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
        } else {
            dim3 grid((m + 8 - 1) / 8, l);
            dim3 block(8 * 16);
            gemv_pipelined_kernel<8, 16, 2048><<<grid, block>>>(
                a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
        }
    } else {
        // Direct kernel - configs from v13
        if (k >= 8192) {
            dim3 grid((m + 8 - 1) / 8, l);
            dim3 block(8 * 32);
            if (b_fits_l1) {
                gemv_direct_kernel<8, 32, CS, CA><<<grid, block>>>(
                    a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
            } else {
                gemv_direct_kernel<8, 32, CS, CS><<<grid, block>>>(
                    a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
            }
        } else if (k >= 4096) {
            dim3 grid((m + 4 - 1) / 4, l);
            dim3 block(4 * 32);
            gemv_direct_kernel<4, 32, CS, CA><<<grid, block>>>(
                a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
        } else {
            dim3 grid((m + 8 - 1) / 8, l);
            dim3 block(8 * 16);
            gemv_direct_kernel<8, 16, CS, CA><<<grid, block>>>(
                a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
        }
    }
}
'''

CPP_SRC = r'''
#include <torch/extension.h>
void run_nvfp4_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,
                    torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);
'''

_module = None

def _load():
    global _module
    if _module is None:
        _module = load_inline(
            name='nvfp4_gemv_v15_hybrid',
            cpp_sources=CPP_SRC,
            cuda_sources=CUDA_SRC,
            functions=['run_nvfp4_gemv'],
            extra_cuda_cflags=[
                '-O3', '--use_fast_math', '-std=c++17',
                '-gencode=arch=compute_100a,code=sm_100a',
            ],
            verbose=False,
        )
    return _module

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, _, _, c = data
    m, k_packed, l = a.shape
    k = k_packed * 2
    n_pad = b.shape[0]
    
    if not sfa.is_cuda:
        sfa = sfa.to(a.device)
    if not sfb.is_cuda:
        sfb = sfb.to(a.device)
    
    _load().run_nvfp4_gemv(
        c.squeeze(1), a.view(torch.uint8), b.view(torch.uint8),
        sfa.view(torch.uint8), sfb.view(torch.uint8), m, k, l, n_pad)
    return c
scrolls · 452 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 109911.

⋯ 1 unchanged lines
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
+ # Combine both v13 (direct) and v14 (pipelined) kernels with selection logic
CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
+ // =============================================================================
+ // Common Utilities
+ // =============================================================================
- __device__ __forceinline__ void ldcs_u32x4(uint32_t* dst, const void* src) {
- asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
- : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
- : "l"(src));
- }
+ enum class CacheHint { CS, CA };
- // Cached load for B (keep in L1, reused across M rows)
- __device__ __forceinline__ void ldca_u32x4(uint32_t* dst, const void* src) {
- asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
- : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
- : "l"(src));
+ template<CacheHint hint>
+ __device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {
+ if constexpr (hint == CacheHint::CS) {
+ asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
+ } else {
+ asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
+ }
}
- __device__ __forceinline__ void ldcs_u16(uint16_t* dst, const void* src) {
- asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
+ template<CacheHint hint>
+ __device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {
+ if constexpr (hint == CacheHint::CS) {
+ asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
+ } else {
+ asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
+ }
}
- // Cached load for scale factors B
- __device__ __forceinline__ void ldca_u16(uint16_t* dst, const void* src) {
- asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
- }
-
-
- // Convert 4 bytes of FP4x2 to 4 x half2 (8 FP16 values)
- __device__ __forceinline__ void fp4x8_to_half2x4(half2* out, uint32_t in) {
+ __device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {
asm volatile(
"{\n\t"
".reg .b8 b0, b1, b2, b3;\n\t"
⋯ 11 unchanged lines
);
}
- // Convert 2 FP8 scale factors to half2
- __device__ __forceinline__ half2 fp8x2_to_half2(uint16_t in) {
+ __device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {
int out;
asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
return *reinterpret_cast<half2*>(&out);
}
- __device__ __forceinline__ float blockscaled_dot_16bytes(
- const uint32_t* a_regs, // 4 x uint32 = 16 bytes = 32 FP4
- const uint32_t* b_regs, // 4 x uint32 = 16 bytes = 32 FP4
- half2 combined_scale // SFA * SFB already fused
+ // =============================================================================
+ // DIRECT KERNEL (from v13)
+ // =============================================================================
+
+ __device__ __forceinline__ float dot_scaled_direct(
+ const uint32_t* a_data,
+ const uint32_t* b_data,
+ half2 scale
) {
- // Unpack to half2
- half2 a_h2[16], b_h2[16];
- fp4x8_to_half2x4(a_h2 + 0, a_regs[0]);
- fp4x8_to_half2x4(a_h2 + 4, a_regs[1]);
- fp4x8_to_half2x4(a_h2 + 8, a_regs[2]);
- fp4x8_to_half2x4(a_h2 + 12, a_regs[3]);
+ half2 a_f16[16], b_f16[16];
+ #pragma unroll
+ for (int i = 0; i < 4; i++) {
+ cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
+ cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
+ }
- fp4x8_to_half2x4(b_h2 + 0, b_regs[0]);
- fp4x8_to_half2x4(b_h2 + 4, b_regs[1]);
- fp4x8_to_half2x4(b_h2 + 8, b_regs[2]);
- fp4x8_to_half2x4(b_h2 + 12, b_regs[3]);
+ half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
+ half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
- // Dot product WITHOUT scale (deferred scaling)
- // First 8 half2 use scale.x, next 8 use scale.y
- half2 acc0 = __hmul2(a_h2[0], b_h2[0]);
- half2 acc1 = __hmul2(a_h2[8], b_h2[8]);
+ #pragma unroll
+ for (int i = 1; i < 8; i++) {
+ acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
+ acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
+ }
+ half sum0 = __hadd(acc0.x, acc0.y);
+ half sum1 = __hadd(acc1.x, acc1.y);
+
+ float result = 0.0f;
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
+ : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
+ : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
+
+ return result;
+ }
+
+ template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CacheHint A_HINT, CacheHint B_HINT>
+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
+ void gemv_direct_kernel(
+ 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, int N_pad
+ ) {
+ const int row_base = blockIdx.x * ROWS_PER_BLOCK;
+ const int batch = blockIdx.y;
+ const int tid = threadIdx.x;
+
+ const int warp_id = tid / THREADS_PER_ROW;
+ const int lane = tid % THREADS_PER_ROW;
+
+ const int row = row_base + warp_id;
+ if (row >= M) return;
+
+ const int K_bytes = K / 2;
+ const int K_sf = K / 16;
+
+ const size_t a_off = (size_t)batch * M * K_bytes + row * K_bytes;
+ const size_t b_off = (size_t)batch * N_pad * K_bytes;
+ const size_t sfa_off = (size_t)batch * M * K_sf + row * K_sf;
+ const size_t sfb_off = (size_t)batch * N_pad * K_sf;
+
+ const uint8_t* a_ptr = a + a_off;
+ const uint8_t* b_ptr = b + b_off;
+ const uint8_t* sfa_ptr = sfa + sfa_off;
+ const uint8_t* sfb_ptr = sfb + sfb_off;
+
+ float acc = 0.0f;
+ const int num_sf_pairs = K_sf / 2;
+
+ #pragma unroll 4
+ for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
+ const int sf_idx = sp * 2;
+ const int byte_idx = sf_idx * 8;
+
+ uint16_t sfa_packed, sfb_packed;
+ load_u16<A_HINT>(&sfa_packed, sfa_ptr + sf_idx);
+ load_u16<B_HINT>(&sfb_packed, sfb_ptr + sf_idx);
+ half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
+
+ uint32_t a_regs[4], b_regs[4];
+ load_vec4<A_HINT>(a_regs, a_ptr + byte_idx);
+ load_vec4<B_HINT>(b_regs, b_ptr + byte_idx);
+
+ acc += dot_scaled_direct(a_regs, b_regs, scale);
+ }
+
#pragma unroll
+ for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
+ acc += __shfl_down_sync(0xffffffff, acc, offset);
+ }
+
+ if (lane == 0) {
+ c[(size_t)row + (size_t)batch * M] = __float2half(acc);
+ }
+ }
+
+ // =============================================================================
+ // PIPELINED KERNEL (from v14)
+ // =============================================================================
+
+ #define ASYNC_CP_16(dst, src) \
+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \
+ :: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
+ #define ASYNC_CP_4(dst, src) \
+ asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" \
+ :: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
+ #define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
+ #define ASYNC_WAIT(n) asm volatile("cp.async.wait_group %0;" :: "n"(n))
+
+ #define NUM_STAGES 2
+
+ __device__ __forceinline__ float dot_scaled_pipelined(
+ const uint8_t* a_smem,
+ const uint8_t* b_smem,
+ half2 scale
+ ) {
+ const uint32_t* a_data = reinterpret_cast<const uint32_t*>(a_smem);
+ const uint32_t* b_data = reinterpret_cast<const uint32_t*>(b_smem);
+
+ half2 a_f16[16], b_f16[16];
+ #pragma unroll
+ for (int i = 0; i < 4; i++) {
+ cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
+ cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
+ }
+
+ half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
+ half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
+
+ #pragma unroll
for (int i = 1; i < 8; i++) {
- acc0 = __hfma2(a_h2[i], b_h2[i], acc0);
- acc1 = __hfma2(a_h2[8 + i], b_h2[8 + i], acc1);
+ acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
+ acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
}
- // Reduce each half2 to half
half sum0 = __hadd(acc0.x, acc0.y);
half sum1 = __hadd(acc1.x, acc1.y);
float result = 0.0f;
- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum0)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.x)));
- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum1)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.y)));
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
+ : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
+ : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
return result;
}
- template<int BLOCK_M, int THREADS_PER_ROW>
- __global__ __launch_bounds__(BLOCK_M * THREADS_PER_ROW)
- void gemv_optimized_kernel(
+ template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, int K_TILE>
+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
+ void gemv_pipelined_kernel(
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, int N_rows
+ int M, int K, int L, int N_pad
) {
- constexpr int BLOCK_SIZE = BLOCK_M * THREADS_PER_ROW;
+ constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW;
+ constexpr int K_TILE_BYTES = K_TILE / 2;
+ constexpr int K_TILE_SF = K_TILE / 16;
- const int m_base = blockIdx.x * BLOCK_M;
- const int batch_id = blockIdx.y;
+ __shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];
+ __shared__ __align__(16) uint8_t sh_b[NUM_STAGES][K_TILE_BYTES];
+ __shared__ __align__(16) uint8_t sh_sfa[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_SF];
+ __shared__ __align__(16) uint8_t sh_sfb[NUM_STAGES][K_TILE_SF];
+
+ const int row_base = blockIdx.x * ROWS_PER_BLOCK;
+ const int batch = blockIdx.y;
const int tid = threadIdx.x;
const int warp_id = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
- const int m = m_base + warp_id;
- if (m >= M) return;
+ const int row = row_base + warp_id;
+ if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
+ const int num_tiles = (K_bytes + K_TILE_BYTES - 1) / K_TILE_BYTES;
- const size_t a_batch_stride = (size_t)M * K_bytes;
- const size_t b_batch_stride = (size_t)N_rows * K_bytes;
- const size_t sfa_batch_stride = (size_t)M * K_sf;
- const size_t sfb_batch_stride = (size_t)N_rows * K_sf;
+ const size_t a_base = (size_t)batch * M * K_bytes;
+ const size_t b_base = (size_t)batch * N_pad * K_bytes;
+ const size_t sfa_base = (size_t)batch * M * K_sf;
+ const size_t sfb_base = (size_t)batch * N_pad * K_sf;
- const uint8_t* row_a = a + batch_id * a_batch_stride + m * K_bytes;
- const uint8_t* batch_b = b + batch_id * b_batch_stride;
- const uint8_t* row_sfa = sfa + batch_id * sfa_batch_stride + m * K_sf;
- const uint8_t* batch_sfb = sfb + batch_id * sfb_batch_stride;
+ auto issue_tile = [&](int tile, int stage) {
+ const int k_off_bytes = tile * K_TILE_BYTES;
+ const int k_off_sf = tile * K_TILE_SF;
+ const int bytes_this = min(K_TILE_BYTES, K_bytes - k_off_bytes);
+ const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
+
+ if (warp_id < ROWS_PER_BLOCK && (row_base + warp_id) < M) {
+ const uint8_t* a_row = a + a_base + (row_base + warp_id) * K_bytes + k_off_bytes;
+ for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
+ if (i + 16 <= bytes_this) {
+ ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);
+ }
+ }
+ const uint8_t* sfa_row = sfa + sfa_base + (row_base + warp_id) * K_sf + k_off_sf;
+ for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
+ if (i + 4 <= sf_this) {
+ ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);
+ }
+ }
+ }
+
+ if (warp_id == 0) {
+ const uint8_t* b_ptr = b + b_base + k_off_bytes;
+ for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
+ if (i + 16 <= bytes_this) {
+ ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);
+ }
+ }
+ const uint8_t* sfb_ptr = sfb + sfb_base + k_off_sf;
+ for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
+ if (i + 4 <= sf_this) {
+ ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);
+ }
+ }
+ }
+ ASYNC_COMMIT();
+ };
+ issue_tile(0, 0);
+ ASYNC_WAIT(0);
+ __syncthreads();
+
float acc = 0.0f;
- // Process 2 scale groups per iteration (32 FP4 values = 16 bytes)
- // Each thread handles multiple scale pairs
- const int num_scale_pairs = K_sf / 2;
-
- #pragma unroll 4
- for (int sp = lane; sp < num_scale_pairs; sp += THREADS_PER_ROW) {
- int sf_base = sp * 2;
- int byte_base = sf_base * 8; // 8 bytes per scale factor
+ for (int tile = 0; tile < num_tiles; tile++) {
+ const int stage = tile % NUM_STAGES;
+ const int next_stage = (tile + 1) % NUM_STAGES;
- // Load scale factors with cache hints
- uint16_t sfa_raw, sfb_raw;
- ldcs_u16(&sfa_raw, row_sfa + sf_base);
- ldca_u16(&sfb_raw, batch_sfb + sf_base);
+ if (tile + 1 < num_tiles) {
+ issue_tile(tile + 1, next_stage);
+ }
- // Convert FP8x2 to half2 and fuse scales
- half2 scale_a = fp8x2_to_half2(sfa_raw);
- half2 scale_b = fp8x2_to_half2(sfb_raw);
- half2 combined_scale = __hmul2(scale_a, scale_b);
+ const int k_off_sf = tile * K_TILE_SF;
+ const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
+ const int num_sf_pairs = sf_this / 2;
- // Load 16 bytes of A and B with cache hints
- uint32_t a_regs[4], b_regs[4];
- ldcs_u32x4(a_regs, row_a + byte_base);
- ldca_u32x4(b_regs, batch_b + byte_base);
+ #pragma unroll 2
+ for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
+ const int sf_idx = sp * 2;
+ const int byte_idx = sf_idx * 8;
+
+ uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);
+ uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);
+ half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
+
+ acc += dot_scaled_pipelined(
+ &sh_a[stage][warp_id][byte_idx],
+ &sh_b[stage][byte_idx],
+ scale
+ );
+ }
- // Compute with deferred scaling
- acc += blockscaled_dot_16bytes(a_regs, b_regs, combined_scale);
+ if (tile + 1 < num_tiles) {
+ ASYNC_WAIT(0);
+ __syncthreads();
+ }
}
- // Warp reduction
#pragma unroll
for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
if (lane == 0) {
- c[(size_t)m + (size_t)batch_id * M] = __float2half(acc);
+ c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
}
+ // =============================================================================
+ // Dispatch
+ // =============================================================================
+
void run_nvfp4_gemv(
torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
- // K-specialized dispatch
- if (k >= 8192) {
- // Large K: 8 rows per block, 32 threads per row
- constexpr int BLOCK_M = 8;
- constexpr int THREADS_PER_ROW = 32;
- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
- dim3 grid(num_m_blocks, l);
- dim3 block(BLOCK_M * THREADS_PER_ROW);
- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
- A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
- SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
- reinterpret_cast<half*>(C.data_ptr<at::Half>()),
- m, k, l, n_pad);
- } else if (k >= 4096) {
- // Medium-large K: 4 rows per block, 32 threads per row
- constexpr int BLOCK_M = 4;
- constexpr int THREADS_PER_ROW = 32;
- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
- dim3 grid(num_m_blocks, l);
- dim3 block(BLOCK_M * THREADS_PER_ROW);
- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
- A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
- SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
- reinterpret_cast<half*>(C.data_ptr<at::Half>()),
- m, k, l, n_pad);
+ constexpr auto CS = CacheHint::CS;
+ constexpr auto CA = CacheHint::CA;
+
+ // Selection heuristic based on benchmarks:
+ // - k=16384, l=1: pipelined faster (24.6 vs 30.8)
+ // - k=7168, l=8: direct faster (38.5 vs 51.0)
+ // - k=2048, l=4: same (16.4)
+
+ const bool use_pipelined = (k >= 8192 && l <= 2) || // Large K, small batch
+ (k <= 4096 && l >= 4); // Small K, larger batch
+
+ const int k_bytes = k / 2;
+ const bool b_fits_l1 = (k_bytes <= 32 * 1024);
+
+ auto* a_ptr = A.data_ptr<uint8_t>();
+ auto* b_ptr = B.data_ptr<uint8_t>();
+ auto* sfa_ptr = SFA.data_ptr<uint8_t>();
+ auto* sfb_ptr = SFB.data_ptr<uint8_t>();
+ auto* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
+
+ if (use_pipelined) {
+ // Pipelined kernel - configs from v14
+ if (k >= 8192) {
+ dim3 grid((m + 4 - 1) / 4, l);
+ dim3 block(4 * 32);
+ gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ } else if (k >= 4096) {
+ dim3 grid((m + 4 - 1) / 4, l);
+ dim3 block(4 * 32);
+ gemv_pipelined_kernel<4, 32, 1024><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ } else {
+ dim3 grid((m + 8 - 1) / 8, l);
+ dim3 block(8 * 16);
+ gemv_pipelined_kernel<8, 16, 2048><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ }
} else {
- // Small K: 8 rows per block, 16 threads per row
- constexpr int BLOCK_M = 8;
- constexpr int THREADS_PER_ROW = 16;
- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
- dim3 grid(num_m_blocks, l);
- dim3 block(BLOCK_M * THREADS_PER_ROW);
- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
- A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
- SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
- reinterpret_cast<half*>(C.data_ptr<at::Half>()),
- m, k, l, n_pad);
+ // Direct kernel - configs from v13
+ if (k >= 8192) {
+ dim3 grid((m + 8 - 1) / 8, l);
+ dim3 block(8 * 32);
+ if (b_fits_l1) {
+ gemv_direct_kernel<8, 32, CS, CA><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ } else {
+ gemv_direct_kernel<8, 32, CS, CS><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ }
+ } else if (k >= 4096) {
+ dim3 grid((m + 4 - 1) / 4, l);
+ dim3 block(4 * 32);
+ gemv_direct_kernel<4, 32, CS, CA><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ } else {
+ dim3 grid((m + 8 - 1) / 8, l);
+ dim3 block(8 * 16);
+ gemv_direct_kernel<8, 16, CS, CA><<<grid, block>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
+ }
}
}
'''
⋯ 4 unchanged lines
torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);
'''
- _cuda_module = None
+ _module = None
- def get_cuda_module():
- global _cuda_module
- if _cuda_module is None:
- _cuda_module = load_inline(
- name='nvfp4_gemv_optimized_v13',
+ def _load():
+ global _module
+ if _module is None:
+ _module = load_inline(
+ name='nvfp4_gemv_v15_hybrid',
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=['run_nvfp4_gemv'],
extra_cuda_cflags=[
- '-O3',
- '--use_fast_math',
- '-std=c++17',
+ '-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=False,
)
- return _cuda_module
+ return _module
def custom_kernel(data: input_t) -> output_t:
- a, b, sfa_ref, sfb_ref, _, _, c = data
- module = get_cuda_module()
+ a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
k = k_packed * 2
n_pad = b.shape[0]
- if not sfa_ref.is_cuda:
- sfa_ref = sfa_ref.to(a.device)
- if not sfb_ref.is_cuda:
- sfb_ref = sfb_ref.to(a.device)
- c_out = c.squeeze(1)
- module.run_nvfp4_gemv(c_out, a.view(torch.uint8), b.view(torch.uint8),
- sfa_ref.view(torch.uint8), sfb_ref.view(torch.uint8),
- m, k, l, n_pad)
+
+ if not sfa.is_cuda:
+ sfa = sfa.to(a.device)
+ if not sfb.is_cuda:
+ sfb = sfb.to(a.device)
+
+ _load().run_nvfp4_gemv(
+ c.squeeze(1), a.view(torch.uint8), b.view(torch.uint8),
+ sfa.view(torch.uint8), sfb.view(torch.uint8), m, k, l, n_pad)
return c
-
scrolls · 581 diff lines total

Best evidence level for this revision: reported

JSON