submission 821679
David Xia · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 403 lines, June 9 Researcher Reciprocity License v1.0.
blocked.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-821679?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
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:62ba7fb676a944f1f3cffc73db4764023255099cf56d942cdd5122135bc5c6de
license declaredunknown
license concludedunknown
authorsDavid Xia
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float smem_sum[32];Kernel source
blocked.py403 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
# =====================================================================
# 1. 纯手工 CUDA 内核:Panel、V矩阵提取
# =====================================================================
cuda_source = """
#include <cuda_runtime.h>
#include <math.h>
#include <stdint.h>
__inline__ __device__ float blockReduceSum(float val, float* shared) {
int lane = threadIdx.x & 31;
int wid = threadIdx.x >> 5;
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
if (lane == 0) shared[wid] = val;
__syncthreads();
int nwarps = (blockDim.x + 31) >> 5;
val = (wid == 0 && lane < nwarps) ? shared[lane] : 0.0f;
if (wid == 0) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
}
return val;
}
__inline__ __device__ float blockReduceMax(float val, float* shared) {
int lane = threadIdx.x & 31;
int wid = threadIdx.x >> 5;
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
}
if (lane == 0) shared[wid] = val;
__syncthreads();
int nwarps = (blockDim.x + 31) >> 5;
val = (wid == 0 && lane < nwarps) ? shared[lane] : 0.0f;
if (wid == 0) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
}
}
return val;
}
__global__ void unblocked_qr_kernel(float* __restrict__ A, float* __restrict__ tau, int batch, int n) {
int b_idx = blockIdx.x;
if (b_idx >= batch) return;
float* A_b = A + b_idx * n * n;
float* tau_b = tau + b_idx * n;
__shared__ float smem_sum[32];
__shared__ float smem_max[32];
if (n == 1) {
if (threadIdx.x == 0) tau_b[0] = 0.0f;
return;
}
for (int k = 0; k < n; ++k) {
float local_max = 0.0f;
for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) {
local_max = fmaxf(local_max, fabsf(A_b[i * n + k]));
}
float max_val = blockReduceMax(local_max, smem_max);
if (threadIdx.x == 0) smem_max[0] = max_val;
__syncthreads();
max_val = smem_max[0];
if (max_val == 0.0f) max_val = 1.0f;
float local_sum = 0.0f;
for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) {
float v = A_b[i * n + k] / max_val;
local_sum += v * v;
}
float tail_sq = blockReduceSum(local_sum, smem_sum);
__syncthreads();
__shared__ float s_tau, s_v0;
if (threadIdx.x == 0) {
tail_sq *= (max_val * max_val);
float alpha = A_b[k * n + k];
float xnorm = sqrtf(tail_sq);
if (xnorm == 0.0f) {
s_tau = 0.0f; s_v0 = 1.0f;
} else {
float beta = -copysignf(hypotf(alpha, xnorm), alpha);
s_tau = (beta - alpha) / beta;
s_v0 = alpha - beta;
A_b[k * n + k] = beta;
}
tau_b[k] = s_tau;
}
__syncthreads();
if (s_tau != 0.0f) {
float inv_v0 = 1.0f / s_v0;
for (int i = k + 1 + threadIdx.x; i < n; i += blockDim.x) A_b[i * n + k] *= inv_v0;
}
__syncthreads();
if (s_tau != 0.0f) {
for (int j = k + 1 + threadIdx.x; j < n; j += blockDim.x) {
float dot = A_b[k * n + j];
for (int i = k + 1; i < n; ++i) dot += A_b[i * n + k] * A_b[i * n + j];
float factor = s_tau * dot;
A_b[k * n + j] -= factor;
for (int i = k + 1; i < n; ++i) A_b[i * n + j] -= factor * A_b[i * n + k];
}
}
__syncthreads();
}
}
__global__ void panel_wy_kernel(float* __restrict__ A, float* __restrict__ tau, float* __restrict__ T_out,
int batch, int n, int k, int b, int block_size)
{
int b_idx = blockIdx.x;
if (b_idx >= batch) return;
float* A_b = A + b_idx * n * n;
float* tau_b = tau + b_idx * n;
float* T_b = T_out + b_idx * block_size * block_size;
__shared__ float smem_sum[32];
__shared__ float smem_max[32];
__shared__ float s_w[64];
__shared__ float s_T[4096];
for (int i = threadIdx.x; i < b * b; i += blockDim.x) s_T[i] = 0.0f;
__syncthreads();
for (int j = 0; j < b; ++j) {
int col_idx = k + j;
if (col_idx == n - 1) {
if (threadIdx.x == 0) { tau_b[col_idx] = 0.0f; s_T[j * b + j] = 0.0f; }
continue;
}
float local_max = 0.0f;
for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) {
local_max = fmaxf(local_max, fabsf(A_b[i * n + col_idx]));
}
float max_val = blockReduceMax(local_max, smem_max);
if (threadIdx.x == 0) smem_max[0] = max_val;
__syncthreads();
max_val = smem_max[0];
if (max_val == 0.0f) max_val = 1.0f;
float local_sum = 0.0f;
for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) {
float val = A_b[i * n + col_idx] / max_val;
local_sum += val * val;
}
float tail_sq = blockReduceSum(local_sum, smem_sum);
__syncthreads();
__shared__ float s_tau, s_v0;
if (threadIdx.x == 0) {
tail_sq *= (max_val * max_val);
float alpha = A_b[col_idx * n + col_idx];
float xnorm = sqrtf(tail_sq);
if (xnorm == 0.0f) {
s_tau = 0.0f; s_v0 = 1.0f;
} else {
float beta = -copysignf(hypotf(alpha, xnorm), alpha);
s_tau = (beta - alpha) / beta;
s_v0 = alpha - beta;
A_b[col_idx * n + col_idx] = beta;
}
tau_b[col_idx] = s_tau;
s_T[j * b + j] = s_tau;
}
__syncthreads();
if (s_tau != 0.0f) {
float inv_v0 = 1.0f / s_v0;
for (int i = col_idx + 1 + threadIdx.x; i < n; i += blockDim.x) A_b[i * n + col_idx] *= inv_v0;
}
__syncthreads();
if (s_tau != 0.0f) {
for (int p = j + 1 + threadIdx.x; p < b; p += blockDim.x) {
int target_col = k + p;
float dot = A_b[col_idx * n + target_col];
for (int i = col_idx + 1; i < n; ++i) dot += A_b[i * n + col_idx] * A_b[i * n + target_col];
float factor = s_tau * dot;
A_b[col_idx * n + target_col] -= factor;
for (int i = col_idx + 1; i < n; ++i) A_b[i * n + target_col] -= factor * A_b[i * n + col_idx];
}
}
__syncthreads();
if (s_tau != 0.0f && j > 0) {
if (threadIdx.x < j) {
int i = threadIdx.x;
float d = A_b[(k + j) * n + (k + i)] * 1.0f;
for (int r = k + j + 1; r < n; ++r) d += A_b[r * n + (k + i)] * A_b[r * n + (k + j)];
s_w[i] = d;
}
__syncthreads();
if (threadIdx.x < j) {
int i = threadIdx.x;
float val = 0.0f;
for (int m = i; m < j; ++m) val += s_T[i * b + m] * s_w[m];
s_T[i * b + j] = -s_tau * val;
}
__syncthreads();
}
}
for (int i = threadIdx.x; i < b * b; i += blockDim.x) {
int row = i / b;
int col = i % b;
T_b[row * block_size + col] = s_T[i];
}
}
// 快速提取纯净的 V 矩阵缓冲区
__global__ void extract_V_kernel(float* A, float* V_buf, int batch, int n, int k, int b, int block_size) {
int m_idx = blockIdx.x * blockDim.x + threadIdx.x;
int n_idx = blockIdx.y * blockDim.y + threadIdx.y;
int b_idx = blockIdx.z;
int M_v = n - k;
if (m_idx < M_v && n_idx < b && b_idx < batch) {
float* A_b = A + b_idx * n * n;
float* V_b = V_buf + b_idx * n * block_size;
float val = 0.0f;
if (m_idx > n_idx) {
val = A_b[(k + m_idx) * n + (k + n_idx)];
} else if (m_idx == n_idx) {
val = 1.0f;
}
// V_buf 是行主序: [n, block_size]
V_b[m_idx * block_size + n_idx] = val;
}
}
void launch_unblocked_qr(uintptr_t A_ptr, uintptr_t tau_ptr, int batch, int n) {
int threads = (n <= 64) ? 64 : ((n <= 128) ? 128 : 256);
unblocked_qr_kernel<<<batch, threads>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(tau_ptr), batch, n);
}
void launch_panel(uintptr_t A_ptr, uintptr_t tau_ptr, uintptr_t T_ptr, int batch, int n, int k, int b, int block_size) {
panel_wy_kernel<<<batch, 256>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(tau_ptr), reinterpret_cast<float*>(T_ptr), batch, n, k, b, block_size);
}
void launch_extract_V(uintptr_t A_ptr, uintptr_t V_ptr, int batch, int n, int k, int b, int block_size) {
dim3 threads(16, 16);
dim3 blocks((n - k + 15) / 16, (b + 15) / 16, batch);
extract_V_kernel<<<blocks, threads>>>(reinterpret_cast<float*>(A_ptr), reinterpret_cast<float*>(V_ptr), batch, n, k, b, block_size);
}
"""
# =====================================================================
# 2. 裸调 cuBLAS API:完全没有 ATen 对象,0 开销下发
# =====================================================================
cpp_source = """
#include <ATen/core/Tensor.h>
#include <pybind11/pybind11.h>
#include <torch/csrc/utils/pybind.h>
#include <cublas_v2.h>
void launch_unblocked_qr(uintptr_t A_ptr, uintptr_t tau_ptr, int batch, int n);
void launch_panel(uintptr_t A_ptr, uintptr_t tau_ptr, uintptr_t T_ptr, int batch, int n, int k, int b, int block_size);
void launch_extract_V(uintptr_t A_ptr, uintptr_t V_ptr, int batch, int n, int k, int b, int block_size);
// 全局唯一的 cuBLAS 句柄,初始化一次
static cublasHandle_t handle = nullptr;
void ultimate_cublas_qr(at::Tensor H, at::Tensor tau, at::Tensor T_buf, at::Tensor V_buf, at::Tensor M1_buf, at::Tensor M2_buf, int block_size) {
if (!handle) {
cublasCreate(&handle);
cublasSetMathMode(handle, CUBLAS_DEFAULT_MATH);
}
int batch = H.size(0);
int n = H.size(1);
if (n <= 176) {
launch_unblocked_qr((uintptr_t)H.data_ptr<float>(), (uintptr_t)tau.data_ptr<float>(), batch, n);
return;
}
// 将所有对象转换为裸指针,消灭在循环里创建 Tensor 的一切开销
float* d_H = H.data_ptr<float>();
float* d_tau = tau.data_ptr<float>();
float* d_T = T_buf.data_ptr<float>();
float* d_V = V_buf.data_ptr<float>();
float* d_M1 = M1_buf.data_ptr<float>();
float* d_M2 = M2_buf.data_ptr<float>();
// 预先计算 Batch Stride
long long stride_H = n * n;
long long stride_T = block_size * block_size;
long long stride_V = n * block_size;
long long stride_M = block_size * n;
float alpha1 = 1.0f, beta0 = 0.0f;
float alpha_sub = -1.0f, beta1 = 1.0f;
for (int k = 0; k < n; k += block_size) {
int b = std::min(block_size, n - k);
launch_panel((uintptr_t)d_H, (uintptr_t)d_tau, (uintptr_t)d_T, batch, n, k, b, block_size);
if (k + b < n) {
launch_extract_V((uintptr_t)d_H, (uintptr_t)d_V, batch, n, k, b, block_size);
int M_v = n - k;
int N_c = n - k - b;
float* C_ptr = d_H + k * n + (k + b);
// =========================================================================
// 极致魔法:利用行主序/列主序的 Stride 对偶性,零开销原地转置
// =========================================================================
// GEMM 1: M1 = V^T * C
cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_T,
N_c, b, M_v, &alpha1,
C_ptr, n, stride_H,
d_V, block_size, stride_V,
&beta0,
d_M1, n, stride_M, batch);
// GEMM 2: M2 = T^T * M1
cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_T,
N_c, b, b, &alpha1,
d_M1, n, stride_M,
d_T, block_size, stride_T,
&beta0,
d_M2, n, stride_M, batch);
// GEMM 3: C = C - V * M2
cublasSgemmStridedBatched(handle, CUBLAS_OP_N, CUBLAS_OP_N,
N_c, M_v, b, &alpha_sub,
d_M2, n, stride_M,
d_V, block_size, stride_V,
&beta1,
C_ptr, n, stride_H, batch);
}
}
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("ultimate_cublas_qr", &ultimate_cublas_qr, "Zero-Overhead Direct cuBLAS Call");
}
"""
# =====================================================================
# 3. 编译:必须链接 -lcublas
# =====================================================================
qr_module = load_inline(
name='ultimate_cublas_wy',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
verbose=False,
no_implicit_headers=True,
extra_cflags=['-O3'],
extra_cuda_cflags=['-O3', '--use_fast_math', '-Xptxas=-O3', '-arch=sm_100a'],
extra_ldflags=['-lcublas'] # 关键:链接底层的 cuBLAS 硬件加速库
)
# =====================================================================
# 4. 纯净 Python 入口
# =====================================================================
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
H = data.clone().contiguous()
tau = torch.zeros((batch, n), dtype=torch.float32, device=data.device)
# 彻底关闭 TF32 以保证恶劣病态矩阵的绝对精度
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = False
block_size = 64 if n >= 2048 else 32
# 所有中间内存全部在外层一次性静态分配好
T_buf = torch.zeros((batch, block_size, block_size), dtype=torch.float32, device=data.device)
V_buf = torch.empty((batch, n, block_size), dtype=torch.float32, device=data.device)
M1_buf = torch.empty((batch, block_size, n), dtype=torch.float32, device=data.device)
M2_buf = torch.empty((batch, block_size, n), dtype=torch.float32, device=data.device)
# 毫无波澜,一指下发,直接交由 C 底层瞬间执行完毕
qr_module.ultimate_cublas_qr(H, tau, T_buf, V_buf, M1_buf, M2_buf, block_size)
return H, tauscrolls · 403 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