submission 79075
txacvalh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 285 lines, June 9 Researcher Reciprocity License v1.0.
kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-79075?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:67ec5cfd8e24c70327751674945ed8e9970bd20ba0914c563f7dcfb6718b78ee
license declaredunknown
license concludedunknown
authorstxacvalh
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled GEMV.fp8
using fp8e4m3 = __nv_fp8_e4m3;shared-memory
extern __shared__ __align__(16) uint64_t sB_u64[];tile-k = 16
constexpr int TILE_K = 16;tile-m = 16
constexpr int TILE_M = 16;tile-n = 8
constexpr int TILE_N = 8;vector-width = half2
half2 rA_float4 = static_cast<half2>(rA_value.small_vec[i]);Kernel source
kernel.py285 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t
import os
custom_func_cuda_source = """
// m16n8k64
using byte = uint8_t;
using int16 = int16_t;
using int32 = int32_t;
using int64 = int64_t;
using fp32 = float;
using fp16 = half;
using bf16 = nv_bfloat16;
using fp8e4m3 = __nv_fp8_e4m3;
using fp8e5m2 = __nv_fp8_e5m2;
using fp4e2m1 = __nv_fp4_e2m1;
using fp4e2m1x2 = __nv_fp4x2_e2m1;
using fp4e2m1x4 = __nv_fp4x4_e2m1;
using e8m0_t = uint8_t;
constexpr int WARP_SIZE = 32;
constexpr int WARP_COUNT = 16;
constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;
constexpr int SM_COUNT = 148;
constexpr int FP4X8_SIZE_IN_BYTES = 4;
constexpr int FP4X16_SIZE_IN_BYTES = 8;
constexpr int BLOCK_M = WARP_COUNT;
// constexpr int BLOCK_M = 128;
constexpr int TILE_M = 16;
constexpr int TILE_K = 16;
constexpr int TILE_N = 8;
constexpr int N_PADDED_128 = 128;
constexpr int N = 1;
template <typename T>
constexpr T DIVUP(const T &x, const T &y) {
return (((x) + ((y)-1)) / (y));
}
__inline__ __device__
int scale_idx(int mn_idx, int k_idx, int l_idx, int MN, int K) {
constexpr int ATOM_SIZE = 128 * 4;
int TOTAL_NUM_ELEMENTS_PER_L = MN * K / 16;
int rest_k = k_idx / 4;
int sub_k_idx = k_idx % 4;
int rest_mn = mn_idx / 128;
int sub_mn_idx1 = (mn_idx % 128) / 32;
int sub_mn_idx0 = mn_idx % 32;
return l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + rest_mn * (128 * K / 16) + rest_k * ATOM_SIZE + sub_mn_idx0 * 16 + sub_mn_idx1 * 4 + sub_k_idx;
}
union fp4vec {
uint64_t vec;
fp4e2m1x2 small_vec[8];
};
// Warp-level reduction: sums 'val' across the active threads in the warp
__inline__ __device__
float warpReduceSum(float val) {
// Mask of active lanes in this warp
unsigned int mask = __activemask();
// Do tree-reduction within the warp
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
val += __shfl_down_sync(mask, val, offset);
}
return val; // After this, lane 0 holds the sum of the warp
}
template <int GRID_DIM_X>
__global__ void custom_func_kernel(const uint64_t* const __restrict__ A,
const uint64_t* const __restrict__ B,
const fp8e4m3 * const __restrict__ scale_A,
const uint32_t * const __restrict__ scale_B,
half* const __restrict__ C,
const int L, const int M, const int K, const int sB_size_in_bytes) {
const int block_m_idx = blockIdx.x;
const int l_idx = blockIdx.y;
const int warp_id = threadIdx.x / WARP_SIZE;
const int lane_id = threadIdx.x % WARP_SIZE;
int m_idx = block_m_idx * BLOCK_M + warp_id;
extern __shared__ __align__(16) uint64_t sB_u64[];
fp8e4m3 * sB_scale = reinterpret_cast<fp8e4m3 *>(
reinterpret_cast<uint8_t *>(sB_u64) + sB_size_in_bytes);
// Load B into shared mem
int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2 * sizeof(uint64_t));
int b_bound = K / (2 * sizeof(uint64_t));
#pragma unroll
for (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {
sB_u64[i] = B[b_l_base_off + i];
}
int b_scale_bound = K / (16 * sizeof(uint32_t));
#pragma unroll
for (int i = threadIdx.x; i < b_scale_bound; i += THREADS_PER_BLOCK) {
reinterpret_cast<uint32_t *>(sB_scale)[i] = scale_B[scale_idx(0, i * sizeof(uint32_t), l_idx, N_PADDED_128, K) / sizeof(uint32_t)];
}
__syncthreads();
for (;m_idx < M; m_idx += GRID_DIM_X * BLOCK_M) {
const int a_l_base_off_u64 = (l_idx * (M * K) + m_idx * (K)) / (2 * sizeof(uint64_t));
float acc{};
#pragma unroll
for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * sizeof(uint64_t)); tile_idx_k += WARP_SIZE) {
fp4vec rA_value;
fp4vec rB_value;
rA_value.vec = A[a_l_base_off_u64 + tile_idx_k];
rB_value.vec = sB_u64[tile_idx_k];
half tile_acc{};
#pragma unroll
for (int i = 0; i < 8; i++) {
half2 rA_float4 = static_cast<half2>(rA_value.small_vec[i]);
half2 rB_float4 = static_cast<half2>(rB_value.small_vec[i]);
tile_acc += rA_float4.x * rB_float4.x;
tile_acc += rA_float4.y * rB_float4.y;
}
/* if (lane_id == 0) {
//printf("hello");
printf("L %d M %d K-tile %d A %d B %d acc %f %f %f %f\\n", l_idx, m_idx, tile_idx_k, rA_value.vec, rB_value.vec, tile_accs.x, tile_accs.y, tile_accs.z, tile_accs.w);
} */
const float a_scale = static_cast<float>(scale_A[scale_idx(m_idx, tile_idx_k * (2 * sizeof(uint64_t)) / 16, l_idx, M, K)]);
// const float b_scale = static_cast<float>(sB_scale[tile_idx_k * (2 * sizeof(uint64_t)) / 16]);
const float b_scale = static_cast<float>(sB_scale[tile_idx_k]);
acc += static_cast<float>(tile_acc) * a_scale * b_scale;
}
acc = warpReduceSum(acc);
if (lane_id == 0) {
const int c_l_base_off_half = l_idx * (M) + m_idx;
C[c_l_base_off_half] = static_cast<half>(acc);
}
}
}
void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C) {
const int L = A.size(0);
const int M = A.size(1);
const int K = A.size(2) * 2;
const int threads = THREADS_PER_BLOCK;
// const int blocks = (N + threads - 1) / threads;
const int sB_size_in_bytes = K / 2;
const int sB_scale_size_in_bytes = K / 16;
const int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes;
constexpr int GRID_DIM_X = SM_COUNT;
dim3 blocks(GRID_DIM_X, L, 1);
custom_func_kernel<GRID_DIM_X><<<blocks, threads, shared_mem_size_in_bytes>>>(
reinterpret_cast<uint64_t *>(A.data_ptr()),
reinterpret_cast<uint64_t *>(B.data_ptr()),
reinterpret_cast<fp8e4m3 *>(scale_A.data_ptr()),
reinterpret_cast<uint32_t *>(scale_B.data_ptr()),
reinterpret_cast<half *>(C.data_ptr()),
L, M, K, sB_size_in_bytes);
/* cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
} */
}
"""
# print(custom_func_cuda_source)
custom_func_cpp_source = """
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_fp4.h>
#include <torch/extension.h>
void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C);
"""
os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0a;12.0a"
module = load_inline(
name='custom_func',
cpp_sources=custom_func_cpp_source,
cuda_sources=custom_func_cuda_source,
functions=['custom_func'],
verbose=True,
extra_cuda_cflags=['-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF_OPERATORS__ ', '-U__CUDA_NO_HALF2_OPERATORS__', '-O3', '-Xptxas', '-O3', '-use_fast_math'],
)
def custom_kernel(
data: input_t,
) -> output_t:
"""
Reference implementation of block-scale fp8 gemv
Args:
data: Tuple that expands to:
a: torch.Tensor[float4e2m1fn] of shape [m, k, l],
b: torch.Tensor[float4e2m1fn] of shape [1, k, l],
sfa: torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l], used by reference implementation
sfb: torch.Tensor[float8_e4m3fnuz] of shape [1, k // 16, l], used by reference implementation
sfa_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_m, 4, rest_k, l],
sfb_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_n, 4, rest_k, l],
c: torch.Tensor[float16] of shape [m, 1, l]
Returns:
Tensor containing output in float16
c: torch.Tensor[float16] of shape [m, 1, l]
"""
"""
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data
# Get dimensions from MxNxL layout
# _, _, l = c_ref.shape
# # Call torch._scaled_mm to compute the GEMV result
# for l_idx in range(l):
# # Convert the scale factor tensor to blocked format
# scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
# scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])
# # (m, k) @ (n, k).T -> (m, n)
# res = torch._scaled_mm(
# a_ref[:, :, l_idx],
# b_ref[:, :, l_idx].transpose(0, 1),
# scale_a.cuda(),
# scale_b.cuda(),
# bias=None,
# out_dtype=torch.float16,
# )
# c_ref[:, 0, l_idx] = res[:, 0]
# a = [l, m, k]
# b = [l, 1, k]
# [l, rest_n, rest_k, 32, 4, 4]
# for e in [a_ref, b_ref, sfa_permuted, sfb_permuted, c_ref]:
# print(e.stride(), e.shape)
args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())
# for a in args:
# assert a.is_cuda, "All input tensors must be on GPU"
# print(a.stride(), a.shape)
module.custom_func(*args)
# for l_idx in range(l):
# # (m, k) @ (n, k).T -> (m, n)
# res = torch._scaled_mm(
# a_ref[:, :, l_idx],
# b_ref[:, :, l_idx].transpose(0, 1),
# sfa_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),
# sfb_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),
# bias=None,
# out_dtype=torch.float16,
# )
# c_ref[:, 0, l_idx] = res[:, 0]
# module.custom_func
return c_refscrolls · 285 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 72157.
import torch+ from torch.utils.cpp_extension import load_inline+ from typing import Listfrom task import input_t, output_t+ import os- # Kernel configuration parameters- sf_vec_size = 16+ custom_func_cuda_source = """+ // m16n8k64+ using byte = uint8_t;+ using int16 = int16_t;+ using int32 = int32_t;+ using int64 = int64_t;+ using fp32 = float;+ using fp16 = half;+ using bf16 = nv_bfloat16;+ using fp8e4m3 = __nv_fp8_e4m3;+ using fp8e5m2 = __nv_fp8_e5m2;- # Helper function for ceiling division- def ceil_div(a, b):- return (a + b - 1) // b+ using fp4e2m1 = __nv_fp4_e2m1;+ using fp4e2m1x2 = __nv_fp4x2_e2m1;+ using fp4e2m1x4 = __nv_fp4x4_e2m1;+ using e8m0_t = uint8_t;+ constexpr int WARP_SIZE = 32;+ constexpr int WARP_COUNT = 16;+ constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;+ constexpr int SM_COUNT = 148;- # Helper function to convert scale factor tensor to blocked format- def to_blocked(input_matrix):- rows, cols = input_matrix.shape+ constexpr int FP4X8_SIZE_IN_BYTES = 4;+ constexpr int FP4X16_SIZE_IN_BYTES = 8;- # Please ensure rows and cols are multiples of 128 and 4 respectively- n_row_blocks = ceil_div(rows, 128)- n_col_blocks = ceil_div(cols, 4)+ constexpr int BLOCK_M = WARP_COUNT;+ // constexpr int BLOCK_M = 128;- padded = input_matrix- blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)- rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)- return rearranged.flatten()+ constexpr int TILE_M = 16;+ constexpr int TILE_K = 16;+ constexpr int TILE_N = 8;+ constexpr int N_PADDED_128 = 128;+ constexpr int N = 1;- @torch.compile- def to_blocked_2(input_matrix):- rows, cols, l = input_matrix.shape- # Please ensure rows and cols are multiples of 128 and 4 respectively- n_row_blocks = ceil_div(rows, 128)- n_col_blocks = ceil_div(cols, 4)+ template <typename T>+ constexpr T DIVUP(const T &x, const T &y) {+ return (((x) + ((y)-1)) / (y));+ }- padded = input_matrix- blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4, l).permute(4, 0, 2, 1, 3)- rearranged = blocks.reshape(l, -1, 4, 32, 4).transpose(2, 3).reshape(l, -1, 32, 16)+ __inline__ __device__+ int scale_idx(int mn_idx, int k_idx, int l_idx, int MN, int K) {+ constexpr int ATOM_SIZE = 128 * 4;+ int TOTAL_NUM_ELEMENTS_PER_L = MN * K / 16;+ int rest_k = k_idx / 4;+ int sub_k_idx = k_idx % 4;+ int rest_mn = mn_idx / 128;+ int sub_mn_idx1 = (mn_idx % 128) / 32;+ int sub_mn_idx0 = mn_idx % 32;- return rearranged.reshape(l, -1)++ return l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + rest_mn * (128 * K / 16) + rest_k * ATOM_SIZE + sub_mn_idx0 * 16 + sub_mn_idx1 * 4 + sub_k_idx;+ }+ union fp4vec {+ uint64_t vec;+ fp4e2m1x2 small_vec[8];+ };+ // Warp-level reduction: sums 'val' across the active threads in the warp+ __inline__ __device__+ float warpReduceSum(float val) {+ // Mask of active lanes in this warp+ unsigned int mask = __activemask();++ // Do tree-reduction within the warp+ for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {+ val += __shfl_down_sync(mask, val, offset);+ }+ return val; // After this, lane 0 holds the sum of the warp+ }+ template <int GRID_DIM_X>+ __global__ void custom_func_kernel(const uint64_t* const __restrict__ A,+ const uint64_t* const __restrict__ B,+ const fp8e4m3 * const __restrict__ scale_A,+ const uint32_t * const __restrict__ scale_B,+ half* const __restrict__ C,+ const int L, const int M, const int K, const int sB_size_in_bytes) {+ const int block_m_idx = blockIdx.x;+ const int l_idx = blockIdx.y;+ const int warp_id = threadIdx.x / WARP_SIZE;+ const int lane_id = threadIdx.x % WARP_SIZE;+ int m_idx = block_m_idx * BLOCK_M + warp_id;++ extern __shared__ __align__(16) uint64_t sB_u64[];+ fp8e4m3 * sB_scale = reinterpret_cast<fp8e4m3 *>(+ reinterpret_cast<uint8_t *>(sB_u64) + sB_size_in_bytes);++++ // Load B into shared mem+ int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2 * sizeof(uint64_t));+ int b_bound = K / (2 * sizeof(uint64_t));++ #pragma unroll+ for (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {+ sB_u64[i] = B[b_l_base_off + i];+ }++ int b_scale_bound = K / (16 * sizeof(uint32_t));+ #pragma unroll+ for (int i = threadIdx.x; i < b_scale_bound; i += THREADS_PER_BLOCK) {+ reinterpret_cast<uint32_t *>(sB_scale)[i] = scale_B[scale_idx(0, i * sizeof(uint32_t), l_idx, N_PADDED_128, K) / sizeof(uint32_t)];+ }+++ __syncthreads();++ for (;m_idx < M; m_idx += GRID_DIM_X * BLOCK_M) {++ const int a_l_base_off_u64 = (l_idx * (M * K) + m_idx * (K)) / (2 * sizeof(uint64_t));++ float acc{};++ #pragma unroll+ for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * sizeof(uint64_t)); tile_idx_k += WARP_SIZE) {++ fp4vec rA_value;+ fp4vec rB_value;++ rA_value.vec = A[a_l_base_off_u64 + tile_idx_k];+ rB_value.vec = sB_u64[tile_idx_k];+ half tile_acc{};+++ #pragma unroll+ for (int i = 0; i < 8; i++) {+ half2 rA_float4 = static_cast<half2>(rA_value.small_vec[i]);+ half2 rB_float4 = static_cast<half2>(rB_value.small_vec[i]);+ tile_acc += rA_float4.x * rB_float4.x;+ tile_acc += rA_float4.y * rB_float4.y;+ }++ /* if (lane_id == 0) {+ //printf("hello");+ printf("L %d M %d K-tile %d A %d B %d acc %f %f %f %f\\n", l_idx, m_idx, tile_idx_k, rA_value.vec, rB_value.vec, tile_accs.x, tile_accs.y, tile_accs.z, tile_accs.w);+ } */++ const float a_scale = static_cast<float>(scale_A[scale_idx(m_idx, tile_idx_k * (2 * sizeof(uint64_t)) / 16, l_idx, M, K)]);+ // const float b_scale = static_cast<float>(sB_scale[tile_idx_k * (2 * sizeof(uint64_t)) / 16]);+ const float b_scale = static_cast<float>(sB_scale[tile_idx_k]);++ acc += static_cast<float>(tile_acc) * a_scale * b_scale;++ }+ acc = warpReduceSum(acc);+ if (lane_id == 0) {+ const int c_l_base_off_half = l_idx * (M) + m_idx;+ C[c_l_base_off_half] = static_cast<half>(acc);+ }++ }+ }++ void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C) {+ const int L = A.size(0);+ const int M = A.size(1);+ const int K = A.size(2) * 2;++ const int threads = THREADS_PER_BLOCK;+ // const int blocks = (N + threads - 1) / threads;++ const int sB_size_in_bytes = K / 2;+ const int sB_scale_size_in_bytes = K / 16;+ const int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes;+++ constexpr int GRID_DIM_X = SM_COUNT;+ dim3 blocks(GRID_DIM_X, L, 1);+ custom_func_kernel<GRID_DIM_X><<<blocks, threads, shared_mem_size_in_bytes>>>(+ reinterpret_cast<uint64_t *>(A.data_ptr()),+ reinterpret_cast<uint64_t *>(B.data_ptr()),+ reinterpret_cast<fp8e4m3 *>(scale_A.data_ptr()),+ reinterpret_cast<uint32_t *>(scale_B.data_ptr()),+ reinterpret_cast<half *>(C.data_ptr()),+ L, M, K, sB_size_in_bytes);++++ /* cudaError_t err = cudaGetLastError();+ if (err != cudaSuccess) {+ throw std::runtime_error(cudaGetErrorString(err));+ } */+ }++ """+ # print(custom_func_cuda_source)+ custom_func_cpp_source = """+ #include <cuda.h>++ #include <cuda_bf16.h>+ #include <cuda_fp16.h>+ #include <cuda_fp8.h>+ #include <cuda_fp4.h>+ #include <torch/extension.h>++ void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C);+ """+ os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0a;12.0a"++ module = load_inline(+ name='custom_func',+ cpp_sources=custom_func_cpp_source,+ cuda_sources=custom_func_cuda_source,+ functions=['custom_func'],+ verbose=True,+ extra_cuda_cflags=['-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF_OPERATORS__ ', '-U__CUDA_NO_HALF2_OPERATORS__', '-O3', '-Xptxas', '-O3', '-use_fast_math'],+ )++def custom_kernel(data: input_t,) -> output_t:⋯ 18 unchanged linesa_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data# Get dimensions from MxNxL layout- m, k, l = a_ref.shape- n_padded_128 = 128+ # _, _, l = c_ref.shape- # a_ref = a_ref.view(torch.uint8).permute(2, 0, 1).contiguous().view(torch.float4_e2m1fn_x2) # [l, m, k]- # b_ref = b_ref.view(torch.uint8).permute(2, 1, 0).contiguous().view(torch.float4_e2m1fn_x2) # [l, k, 1]- # Convert the scale factor tensor to blocked format- # scale_a = to_blocked_2(sfa_ref_cpu.cuda())- # scale_b = to_blocked_2(sfb_ref_cpu.cuda())+ # # Call torch._scaled_mm to compute the GEMV result+ # for l_idx in range(l):+ # # Convert the scale factor tensor to blocked format+ # scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])+ # scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])+ # # (m, k) @ (n, k).T -> (m, n)+ # res = torch._scaled_mm(+ # a_ref[:, :, l_idx],+ # b_ref[:, :, l_idx].transpose(0, 1),+ # scale_a.cuda(),+ # scale_b.cuda(),+ # bias=None,+ # out_dtype=torch.float16,+ # )+ # c_ref[:, 0, l_idx] = res[:, 0]++ # a = [l, m, k]+ # b = [l, 1, k]+ # [l, rest_n, rest_k, 32, 4, 4]+ # for e in [a_ref, b_ref, sfa_permuted, sfb_permuted, c_ref]:+ # print(e.stride(), e.shape)+ args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())++ # for a in args:+ # assert a.is_cuda, "All input tensors must be on GPU"+ # print(a.stride(), a.shape)+ module.custom_func(*args)- # Call torch._scaled_mm to compute the GEMV result- tmp_c = torch.empty((l, 64, m), dtype=torch.float16, device="cuda")- for l_idx in range(l):-- # (m, k) @ (n, k).T -> (m, n)- torch._scaled_mm(- a_ref[:, :, l_idx],- b_ref[0:64, :, l_idx].transpose(0, 1),- sfa_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),- sfb_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),- bias=None,- out_dtype=torch.float16,- out=tmp_c[l_idx,:,:].transpose(0,1),- )- return tmp_c[:, 0:1, :].permute(2, 1, 0)No newline at end of file+ # for l_idx in range(l):+ # # (m, k) @ (n, k).T -> (m, n)+ # res = torch._scaled_mm(+ # a_ref[:, :, l_idx],+ # b_ref[:, :, l_idx].transpose(0, 1),+ # sfa_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),+ # sfb_permuted[:,:,:,:,:,l_idx].permute(2, 4, 0, 1, 3).flatten(),+ # bias=None,+ # out_dtype=torch.float16,+ # )+ # c_ref[:, 0, l_idx] = res[:, 0]+ # module.custom_func+ return c_refNo newline at end of file
scrolls · 317 diff lines total
Best evidence level for this revision: reported
JSON