submission 824705
weltschmerz007 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1195 lines, June 9 Researcher Reciprocity License v1.0.
sub19.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-824705?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:c8499a0c4eb4a22073e3b4544c2b8f5cddbbe8244797e11813fe07751feac9e1
license declaredunknown
license concludedunknown
authorsweltschmerz007
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float red[THREADS];Kernel source
sub19.py1195 lines
# custom_kernel.py
import torch
from torch.utils.cpp_extension import load_inline
cpp_source = r"""
#include <torch/extension.h>
std::tuple<torch::Tensor, torch::Tensor> batched_qr_forward(torch::Tensor A);
"""
cuda_source = r"""
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <ATen/Context.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <tuple>
#include <stdint.h>
#include <limits.h>
#include <mutex>
#include <algorithm>
#include <array>
#define CHECK_INPUT(x) TORCH_CHECK(x.is_cuda(), #x " must be CUDA")
#define CHECK_FLOAT(x) TORCH_CHECK(x.scalar_type() == at::kFloat, #x " must be float32")
static inline void check_launch(const char* msg) {
cudaError_t e = cudaGetLastError();
TORCH_CHECK(e == cudaSuccess, msg, ": ", cudaGetErrorString(e));
}
static inline const char* cublas_status_string(cublasStatus_t s) {
switch (s) {
case CUBLAS_STATUS_SUCCESS: return "CUBLAS_STATUS_SUCCESS";
case CUBLAS_STATUS_NOT_INITIALIZED: return "CUBLAS_STATUS_NOT_INITIALIZED";
case CUBLAS_STATUS_ALLOC_FAILED: return "CUBLAS_STATUS_ALLOC_FAILED";
case CUBLAS_STATUS_INVALID_VALUE: return "CUBLAS_STATUS_INVALID_VALUE";
case CUBLAS_STATUS_ARCH_MISMATCH: return "CUBLAS_STATUS_ARCH_MISMATCH";
case CUBLAS_STATUS_MAPPING_ERROR: return "CUBLAS_STATUS_MAPPING_ERROR";
case CUBLAS_STATUS_EXECUTION_FAILED: return "CUBLAS_STATUS_EXECUTION_FAILED";
case CUBLAS_STATUS_INTERNAL_ERROR: return "CUBLAS_STATUS_INTERNAL_ERROR";
#if defined(CUBLAS_STATUS_NOT_SUPPORTED)
case CUBLAS_STATUS_NOT_SUPPORTED: return "CUBLAS_STATUS_NOT_SUPPORTED";
#endif
#if defined(CUBLAS_STATUS_LICENSE_ERROR)
case CUBLAS_STATUS_LICENSE_ERROR: return "CUBLAS_STATUS_LICENSE_ERROR";
#endif
default: return "CUBLAS_STATUS_UNKNOWN";
}
}
static inline void check_cublas(cublasStatus_t s, const char* msg) {
TORCH_CHECK(s == CUBLAS_STATUS_SUCCESS, msg, ": ", cublas_status_string(s));
}
static inline cublasHandle_t blas_handle() {
int dev = 0;
cudaError_t ce = cudaGetDevice(&dev);
TORCH_CHECK(ce == cudaSuccess, "cudaGetDevice failed: ", cudaGetErrorString(ce));
TORCH_CHECK(dev >= 0 && dev < 32, "bad device index");
static thread_local cublasHandle_t handles[32] = {};
if (handles[dev] == nullptr) {
ce = cudaFree(0);
TORCH_CHECK(ce == cudaSuccess, "cudaFree(0) failed: ", cudaGetErrorString(ce));
check_cublas(cublasCreate(&handles[dev]), "cublasCreate");
check_cublas(cublasSetPointerMode(handles[dev], CUBLAS_POINTER_MODE_HOST), "cublasSetPointerMode");
check_cublas(cublasSetAtomicsMode(handles[dev], CUBLAS_ATOMICS_ALLOWED), "cublasSetAtomicsMode");
}
return handles[dev];
}
static inline void bgemm(
cublasHandle_t h,
cublasOperation_t opA,
cublasOperation_t opB,
int m,
int n,
int k,
float alpha,
const float* A,
int lda,
long long strideA,
const float* B,
int ldb,
long long strideB,
float beta,
float* C,
int ldc,
long long strideC,
int batch,
bool fast_tf32
) {
if (m <= 0 || n <= 0 || k <= 0 || batch <= 0) return;
cublasComputeType_t ct = fast_tf32 ? CUBLAS_COMPUTE_32F_FAST_TF32 : CUBLAS_COMPUTE_32F_PEDANTIC;
cublasGemmAlgo_t algo = fast_tf32 ? CUBLAS_GEMM_DEFAULT_TENSOR_OP : CUBLAS_GEMM_DEFAULT;
check_cublas(
cublasGemmStridedBatchedEx(
h,
opA,
opB,
m,
n,
k,
&alpha,
A,
CUDA_R_32F,
lda,
strideA,
B,
CUDA_R_32F,
ldb,
strideB,
&beta,
C,
CUDA_R_32F,
ldc,
strideC,
batch,
ct,
algo
),
"cublasGemmStridedBatchedEx"
);
}
__device__ __forceinline__ float warp_sum(float v) {
v += __shfl_down_sync(0xffffffffu, v, 16);
v += __shfl_down_sync(0xffffffffu, v, 8);
v += __shfl_down_sync(0xffffffffu, v, 4);
v += __shfl_down_sync(0xffffffffu, v, 2);
v += __shfl_down_sync(0xffffffffu, v, 1);
return v;
}
__device__ __forceinline__ float warp_max(float v) {
v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 16));
v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 8));
v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 4));
v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 2));
v = fmaxf(v, __shfl_down_sync(0xffffffffu, v, 1));
return v;
}
template<int THREADS>
__device__ __forceinline__ float block_sum(float v, float* red) {
constexpr int WARPS = THREADS / 32;
int tid = threadIdx.x;
int lane = tid & 31;
int wid = tid >> 5;
v = warp_sum(v);
if (lane == 0) red[wid] = v;
__syncthreads();
float out = 0.0f;
if (wid == 0) {
out = (lane < WARPS) ? red[lane] : 0.0f;
out = warp_sum(out);
if (lane == 0) red[0] = out;
}
__syncthreads();
return red[0];
}
template<int THREADS>
__device__ __forceinline__ float block_max(float v, float* red) {
constexpr int WARPS = THREADS / 32;
int tid = threadIdx.x;
int lane = tid & 31;
int wid = tid >> 5;
v = warp_max(v);
if (lane == 0) red[wid] = v;
__syncthreads();
float out = 0.0f;
if (wid == 0) {
out = (lane < WARPS) ? red[lane] : 0.0f;
out = warp_max(out);
if (lane == 0) red[0] = out;
}
__syncthreads();
return red[0];
}
__global__ void lastmax512_kernel(const float* __restrict__ A, float* __restrict__ out, int B) {
constexpr int N = 512;
constexpr int THREADS = 256;
int tid = threadIdx.x;
__shared__ float red[THREADS];
float mx = 0.0f;
int total = B * N;
for (int idx = tid; idx < total; idx += THREADS) {
int b = idx / N;
int r = idx - b * N;
float v = A[((int64_t)b * N + r) * N + (N - 1)];
mx = fmaxf(mx, fabsf(v));
}
red[tid] = mx;
__syncthreads();
for (int s = THREADS / 2; s > 0; s >>= 1) {
if (tid < s) red[tid] = fmaxf(red[tid], red[tid + s]);
__syncthreads();
}
if (tid == 0) out[0] = red[0];
}
__global__ void colmax512_kernel(const float* __restrict__ A, float* __restrict__ colmax, int B) {
constexpr int N = 512;
constexpr int THREADS = 256;
int c = blockIdx.x;
int tid = threadIdx.x;
__shared__ float red[THREADS];
float mx = 0.0f;
int total = B * N;
for (int idx = tid; idx < total; idx += THREADS) {
int b = idx / N;
int r = idx - b * N;
float v = A[((int64_t)b * N + r) * N + c];
mx = fmaxf(mx, fabsf(v));
}
red[tid] = mx;
__syncthreads();
for (int s = THREADS / 2; s > 0; s >>= 1) {
if (tid < s) red[tid] = fmaxf(red[tid], red[tid + s]);
__syncthreads();
}
if (tid == 0) colmax[c] = red[0];
}
static int detect_active_cols_512(torch::Tensor A, int B) {
auto one = torch::empty({1}, A.options());
lastmax512_kernel<<<1, 256>>>(A.data_ptr<float>(), one.data_ptr<float>(), B);
check_launch("lastmax512");
float last_h = 0.0f;
cudaError_t e0 = cudaMemcpy(&last_h, one.data_ptr<float>(), sizeof(float), cudaMemcpyDeviceToHost);
TORCH_CHECK(e0 == cudaSuccess, "lastmax512 copy failed: ", cudaGetErrorString(e0));
if (last_h > 1.0e-4f) return 512;
auto tmp = torch::empty({512}, A.options());
colmax512_kernel<<<512, 256>>>(A.data_ptr<float>(), tmp.data_ptr<float>(), B);
check_launch("colmax512");
std::array<float, 512> h;
cudaError_t e = cudaMemcpy(h.data(), tmp.data_ptr<float>(), 512 * sizeof(float), cudaMemcpyDeviceToHost);
TORCH_CHECK(e == cudaSuccess, "colmax512 copy failed: ", cudaGetErrorString(e));
float gmax = 0.0f;
for (int i = 0; i < 512; ++i) gmax = std::max(gmax, h[i]);
if (gmax == 0.0f) return 0;
float thresh = 1.0e-3f * gmax;
int active = 0;
for (int c = 511; c >= 0; --c) {
if (h[c] > thresh) {
active = c + 1;
break;
}
}
return active < 496 ? active : 512;
}
template<int N, int THREADS>
__global__ void __launch_bounds__(THREADS, 1) small_qr_kernel(
const float* __restrict__ A_in,
float* __restrict__ H_out,
float* __restrict__ tau,
int B
) {
int b = blockIdx.x;
if (b >= B) return;
constexpr int S = N + 1;
extern __shared__ float smem[];
float* M = smem;
float* red = smem + N * S;
const float* A = A_in + (int64_t)b * N * N;
float* H = H_out + (int64_t)b * N * N;
float* t = tau + (int64_t)b * N;
int tid = threadIdx.x;
for (int idx = tid; idx < N * N; idx += THREADS) {
int i = idx / N;
int j = idx - i * N;
M[i * S + j] = A[idx];
}
__syncthreads();
for (int k = 0; k < N - 1; ++k) {
int len = N - k;
float mx0 = 0.0f;
for (int i = tid; i < len; i += THREADS) {
mx0 = fmaxf(mx0, fabsf(M[(k + i) * S + k]));
}
float mx = block_max<THREADS>(mx0, red);
if (mx > 1.0e-12f) {
float sq = 0.0f;
for (int i = tid + 1; i < len; i += THREADS) {
float v = M[(k + i) * S + k] / mx;
sq += v * v;
}
float tail = block_sum<THREADS>(sq, red);
if (tid == 0) {
float alpha = M[k * S + k] / mx;
float normx = sqrtf(alpha * alpha + tail);
float beta_s = (alpha >= 0.0f) ? -normx : normx;
red[0] = (beta_s - alpha) / beta_s;
red[1] = beta_s * mx;
red[2] = 1.0f / (M[k * S + k] - red[1]);
}
__syncthreads();
float tauk = red[0];
float beta = red[1];
float scale = red[2];
for (int i = tid + 1; i < len; i += THREADS) {
M[(k + i) * S + k] *= scale;
}
__syncthreads();
if (tauk > 0.0f) {
for (int j = k + 1 + tid; j < N; j += THREADS) {
float w = M[k * S + j];
#pragma unroll 4
for (int i = 1; i < len; ++i) {
w += M[(k + i) * S + k] * M[(k + i) * S + j];
}
float c = tauk * w;
M[k * S + j] -= c;
#pragma unroll 4
for (int i = 1; i < len; ++i) {
M[(k + i) * S + j] -= c * M[(k + i) * S + k];
}
}
}
__syncthreads();
if (tid == 0) {
M[k * S + k] = beta;
t[k] = tauk;
}
__syncthreads();
} else {
if (tid == 0) t[k] = 0.0f;
__syncthreads();
}
}
if (tid == 0) t[N - 1] = 0.0f;
__syncthreads();
for (int idx = tid; idx < N * N; idx += THREADS) {
int i = idx / N;
int j = idx - i * N;
H[idx] = M[i * S + j];
}
}
template<int THREADS>
__global__ void __launch_bounds__(THREADS, 1) panel_smem_kernel(
float* __restrict__ A,
float* __restrict__ tau_out,
float* __restrict__ T_out,
float* __restrict__ V_out,
int B,
int n,
int k0,
int m,
int jb,
int nb
) {
int b = blockIdx.x;
if (b >= B) return;
int SV = jb + 1;
int ST = jb + 1;
extern __shared__ float smem[];
float* P = smem;
float* TT = P + m * SV;
float* red = TT + jb * ST;
float* A_b = A + (int64_t)b * n * n;
float* tau_b = tau_out + (int64_t)b * n;
float* T_b = T_out + (int64_t)b * nb * nb;
float* V_b = V_out + (int64_t)b * n * nb;
int tid = threadIdx.x;
for (int idx = tid; idx < m * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
P[i * SV + j] = A_b[(int64_t)(k0 + i) * n + (k0 + j)];
}
__syncthreads();
for (int j = 0; j < jb; ++j) {
int len = m - j;
float mx0 = 0.0f;
for (int i = tid; i < len; i += THREADS) {
mx0 = fmaxf(mx0, fabsf(P[(j + i) * SV + j]));
}
float mx = block_max<THREADS>(mx0, red);
if (mx > 1.0e-12f) {
float sq = 0.0f;
for (int i = tid + 1; i < len; i += THREADS) {
float v = P[(j + i) * SV + j] / mx;
sq += v * v;
}
float tail = block_sum<THREADS>(sq, red);
float tauj, beta, scale;
if (tid == 0) {
float alpha = P[j * SV + j] / mx;
float normx = sqrtf(alpha * alpha + tail);
float beta_s = (alpha >= 0.0f) ? -normx : normx;
tauj = (beta_s - alpha) / beta_s;
beta = beta_s * mx;
scale = 1.0f / (P[j * SV + j] - beta);
red[0] = tauj;
red[1] = beta;
red[2] = scale;
}
__syncthreads();
tauj = red[0];
beta = red[1];
scale = red[2];
for (int i = tid + 1; i < len; i += THREADS) {
P[(j + i) * SV + j] *= scale;
}
__syncthreads();
if (tauj > 0.0f) {
for (int c = j + 1 + tid; c < jb; c += THREADS) {
float w = P[j * SV + c];
#pragma unroll 4
for (int i = 1; i < len; ++i) {
w += P[(j + i) * SV + j] * P[(j + i) * SV + c];
}
float wt = tauj * w;
P[j * SV + c] -= wt;
#pragma unroll 4
for (int i = 1; i < len; ++i) {
P[(j + i) * SV + c] -= wt * P[(j + i) * SV + j];
}
}
}
__syncthreads();
if (tid == 0) {
P[j * SV + j] = beta;
tau_b[k0 + j] = tauj;
}
__syncthreads();
if (tauj > 0.0f) {
for (int i = tid; i < j; i += THREADS) {
float w = P[j * SV + i];
#pragma unroll 4
for (int r = 1; r < len; ++r) {
w += P[(j + r) * SV + i] * P[(j + r) * SV + j];
}
TT[j * ST + i] = w;
}
__syncthreads();
for (int i = tid; i < j; i += THREADS) {
float acc = 0.0f;
#pragma unroll 4
for (int l = i; l < j; ++l) {
acc += TT[i * ST + l] * TT[j * ST + l];
}
TT[i * ST + j] = -tauj * acc;
}
} else {
for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
}
} else {
if (tid == 0) {
tau_b[k0 + j] = 0.0f;
red[0] = 0.0f;
}
__syncthreads();
for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
}
if (tid == 0) TT[j * ST + j] = (mx > 1.0e-12f) ? red[0] : 0.0f;
__syncthreads();
}
for (int idx = tid; idx < m * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
float val = P[i * SV + j];
A_b[(int64_t)(k0 + i) * n + (k0 + j)] = val;
if (i > j) V_b[i * nb + j] = val;
else if (i == j) V_b[i * nb + j] = 1.0f;
else V_b[i * nb + j] = 0.0f;
}
for (int idx = tid; idx < jb * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
T_b[i * nb + j] = (i <= j) ? TT[i * ST + j] : 0.0f;
}
}
template<int THREADS>
__global__ void __launch_bounds__(THREADS, 1) panel_inplace_v_kernel(
float* __restrict__ A,
float* __restrict__ tau_out,
float* __restrict__ T_out,
float* __restrict__ R_out,
int B,
int n,
int k0,
int m,
int jb,
int nb
) {
int b = blockIdx.x;
if (b >= B) return;
int SV = jb + 1;
int ST = jb + 1;
extern __shared__ float smem[];
float* P = smem;
float* TT = P + m * SV;
float* red = TT + jb * ST;
float* A_b = A + (int64_t)b * n * n;
float* tau_b = tau_out + (int64_t)b * n;
float* T_b = T_out + (int64_t)b * nb * nb;
float* R_b = R_out + (int64_t)b * nb * nb;
int tid = threadIdx.x;
for (int idx = tid; idx < m * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
P[i * SV + j] = A_b[(int64_t)(k0 + i) * n + (k0 + j)];
}
__syncthreads();
for (int j = 0; j < jb; ++j) {
int len = m - j;
float mx0 = 0.0f;
for (int i = tid; i < len; i += THREADS) {
mx0 = fmaxf(mx0, fabsf(P[(j + i) * SV + j]));
}
float mx = block_max<THREADS>(mx0, red);
if (mx > 1.0e-12f) {
float sq = 0.0f;
for (int i = tid + 1; i < len; i += THREADS) {
float v = P[(j + i) * SV + j] / mx;
sq += v * v;
}
float tail = block_sum<THREADS>(sq, red);
float tauj, beta, scale;
if (tid == 0) {
float alpha = P[j * SV + j] / mx;
float normx = sqrtf(alpha * alpha + tail);
float beta_s = (alpha >= 0.0f) ? -normx : normx;
tauj = (beta_s - alpha) / beta_s;
beta = beta_s * mx;
scale = 1.0f / (P[j * SV + j] - beta);
red[0] = tauj;
red[1] = beta;
red[2] = scale;
}
__syncthreads();
tauj = red[0];
beta = red[1];
scale = red[2];
for (int i = tid + 1; i < len; i += THREADS) {
P[(j + i) * SV + j] *= scale;
}
__syncthreads();
if (tauj > 0.0f) {
for (int c = j + 1 + tid; c < jb; c += THREADS) {
float w = P[j * SV + c];
#pragma unroll 4
for (int i = 1; i < len; ++i) {
w += P[(j + i) * SV + j] * P[(j + i) * SV + c];
}
float wt = tauj * w;
P[j * SV + c] -= wt;
#pragma unroll 4
for (int i = 1; i < len; ++i) {
P[(j + i) * SV + c] -= wt * P[(j + i) * SV + j];
}
}
}
__syncthreads();
if (tid == 0) {
P[j * SV + j] = beta;
tau_b[k0 + j] = tauj;
}
__syncthreads();
if (tauj > 0.0f) {
for (int i = tid; i < j; i += THREADS) {
float w = P[j * SV + i];
#pragma unroll 4
for (int r = 1; r < len; ++r) {
w += P[(j + r) * SV + i] * P[(j + r) * SV + j];
}
TT[j * ST + i] = w;
}
__syncthreads();
for (int i = tid; i < j; i += THREADS) {
float acc = 0.0f;
#pragma unroll 4
for (int l = i; l < j; ++l) {
acc += TT[i * ST + l] * TT[j * ST + l];
}
TT[i * ST + j] = -tauj * acc;
}
} else {
for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
}
} else {
if (tid == 0) {
tau_b[k0 + j] = 0.0f;
red[0] = 0.0f;
}
__syncthreads();
for (int i = tid; i < j; i += THREADS) TT[i * ST + j] = 0.0f;
}
if (tid == 0) TT[j * ST + j] = (mx > 1.0e-12f) ? red[0] : 0.0f;
__syncthreads();
}
for (int idx = tid; idx < m * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
float val = P[i * SV + j];
if (i < jb && i <= j) R_b[i * nb + j] = val;
float vout;
if (i > j) vout = val;
else if (i == j) vout = 1.0f;
else vout = 0.0f;
A_b[(int64_t)(k0 + i) * n + (k0 + j)] = vout;
}
for (int idx = tid; idx < jb * jb; idx += THREADS) {
int i = idx / jb;
int j = idx - i * jb;
T_b[i * nb + j] = (i <= j) ? TT[i * ST + j] : 0.0f;
}
}
__global__ void restore_r_kernel(
float* __restrict__ A,
const float* __restrict__ R,
int B,
int n,
int k0,
int jb,
int nb
) {
int b = blockIdx.x;
int tid = threadIdx.x;
if (b >= B) return;
float* A_b = A + (int64_t)b * n * n;
const float* R_b = R + (int64_t)b * nb * nb;
for (int idx = tid; idx < jb * jb; idx += blockDim.x) {
int i = idx / jb;
int j = idx - i * jb;
if (i <= j) {
A_b[(int64_t)(k0 + i) * n + (k0 + j)] = R_b[i * nb + j];
}
}
}
std::tuple<torch::Tensor, torch::Tensor> batched_qr_forward(torch::Tensor A) {
torch::NoGradGuard no_grad;
CHECK_INPUT(A);
CHECK_FLOAT(A);
TORCH_CHECK(A.dim() == 3, "A must have shape batch x n x n");
TORCH_CHECK(A.size(1) == A.size(2), "A must be square");
TORCH_CHECK(A.size(0) <= INT_MAX, "batch too large");
TORCH_CHECK(A.size(1) <= INT_MAX, "n too large");
TORCH_CHECK(A.is_contiguous(), "A must be contiguous");
int dev = A.get_device();
cudaSetDevice(dev);
int64_t B64 = A.size(0);
int64_t N64 = A.size(1);
int B = (int)B64;
int n = (int)N64;
if (B == 0) {
auto H0 = torch::empty_like(A);
auto tau0 = torch::empty({B64, N64}, A.options());
return std::make_tuple(H0, tau0);
}
if (n == 1) {
auto H1 = A.clone();
auto tau1 = torch::zeros({B64, N64}, A.options());
return std::make_tuple(H1, tau1);
}
static std::once_flag flag;
static int max_smem = 0;
std::call_once(flag, [&]() {
cudaDeviceGetAttribute(&max_smem, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
int s = max_smem > 0 ? max_smem : 155000;
cudaFuncSetAttribute(
small_qr_kernel<32, 64>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
std::min(s, 160000)
);
cudaFuncSetAttribute(
small_qr_kernel<176, 256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
std::min(s, 160000)
);
cudaFuncSetAttribute(
panel_smem_kernel<256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
s
);
cudaFuncSetAttribute(
panel_inplace_v_kernel<256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
s
);
});
if (n == 32) {
auto H = torch::empty_like(A);
auto tau = torch::empty({B64, N64}, A.options());
small_qr_kernel<32, 64><<<B, 64, (32 * 33 + 8) * sizeof(float)>>>(
A.data_ptr<float>(),
H.data_ptr<float>(),
tau.data_ptr<float>(),
B
);
check_launch("qr32");
return std::make_tuple(H, tau);
}
if (n == 176) {
auto H = torch::empty_like(A);
auto tau = torch::empty({B64, N64}, A.options());
small_qr_kernel<176, 256><<<B, 256, (176 * 177 + 8) * sizeof(float)>>>(
A.data_ptr<float>(),
H.data_ptr<float>(),
tau.data_ptr<float>(),
B
);
check_launch("qr176");
return std::make_tuple(H, tau);
}
if (n <= 1024) {
if (n == 1024 && B <= 8) {
auto res = at::geqrf(A.contiguous());
return std::make_tuple(std::get<0>(res), std::get<1>(res));
}
int active_n = n;
if (n == 512 && B > 64) {
active_n = detect_active_cols_512(A, B);
}
auto H = A.clone();
torch::Tensor tau;
if (active_n < n) tau = torch::zeros({B64, N64}, H.options());
else tau = torch::empty({B64, N64}, H.options());
int nb = (n <= 512) ? 64 : 48;
float* H_ptr = H.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
cublasHandle_t h = blas_handle();
int64_t t_count = B64 * nb * nb;
int64_t w_count = B64 * nb * N64;
int64_t z_count = B64 * nb * N64;
int64_t x_count = (n < 1024) ? (B64 * N64 * nb) : (B64 * nb * nb);
auto work = torch::empty({t_count + w_count + z_count + x_count}, H.options());
float* base = work.data_ptr<float>();
float* T_ptr = base;
float* W_ptr = T_ptr + t_count;
float* Z_ptr = W_ptr + w_count;
float* X_ptr = Z_ptr + z_count;
if (n < 512 || (n == 512 && B <= 64)) {
float* V_ptr = X_ptr;
for (int k = 0; k < active_n; k += nb) {
int jb = std::min(nb, active_n - k);
int m = n - k;
int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);
panel_smem_kernel<256><<<B, 256, smem>>>(
H_ptr,
tau_ptr,
T_ptr,
V_ptr,
B,
n,
k,
m,
jb,
nb
);
if (k + jb < active_n) {
int cols = active_n - k - jb;
float* Atr = H_ptr + (int64_t)k * n + (k + jb);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
m,
1.0f,
Atr,
n,
(long long)n * n,
V_ptr,
nb,
(long long)n * nb,
0.0f,
W_ptr,
n,
(long long)nb * n,
B,
false
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
jb,
1.0f,
W_ptr,
n,
(long long)nb * n,
T_ptr,
nb,
(long long)nb * nb,
0.0f,
Z_ptr,
n,
(long long)nb * n,
B,
false
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_N,
cols,
m,
jb,
-1.0f,
Z_ptr,
n,
(long long)nb * n,
V_ptr,
nb,
(long long)n * nb,
1.0f,
Atr,
n,
(long long)n * n,
B,
false
);
}
}
check_launch("exact LARFB");
return std::make_tuple(H, tau);
}
if (n == 512) {
float* V_ptr = X_ptr;
for (int k = 0; k < active_n; k += nb) {
int jb = std::min(nb, active_n - k);
int m = n - k;
int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);
panel_smem_kernel<256><<<B, 256, smem>>>(
H_ptr,
tau_ptr,
T_ptr,
V_ptr,
B,
n,
k,
m,
jb,
nb
);
if (k + jb < active_n) {
int cols = active_n - k - jb;
float* Atr = H_ptr + (int64_t)k * n + (k + jb);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
m,
1.0f,
Atr,
n,
(long long)n * n,
V_ptr,
nb,
(long long)n * nb,
0.0f,
W_ptr,
n,
(long long)nb * n,
B,
true
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
jb,
1.0f,
W_ptr,
n,
(long long)nb * n,
T_ptr,
nb,
(long long)nb * nb,
0.0f,
Z_ptr,
n,
(long long)nb * n,
B,
false
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_N,
cols,
m,
jb,
-1.0f,
Z_ptr,
n,
(long long)nb * n,
V_ptr,
nb,
(long long)n * nb,
1.0f,
Atr,
n,
(long long)n * n,
B,
true
);
}
}
check_launch("fast n512");
return std::make_tuple(H, tau);
}
float* R_ptr = X_ptr;
for (int k = 0; k < active_n; k += nb) {
int jb = std::min(nb, active_n - k);
int m = n - k;
int smem = (m * (jb + 1) + jb * (jb + 1) + 8) * sizeof(float);
panel_inplace_v_kernel<256><<<B, 256, smem>>>(
H_ptr,
tau_ptr,
T_ptr,
R_ptr,
B,
n,
k,
m,
jb,
nb
);
float* Vp = H_ptr + (int64_t)k * n + k;
if (k + jb < active_n) {
int cols = active_n - k - jb;
float* Atr = H_ptr + (int64_t)k * n + (k + jb);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
m,
1.0f,
Atr,
n,
(long long)n * n,
Vp,
n,
(long long)n * n,
0.0f,
W_ptr,
n,
(long long)nb * n,
B,
true
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_T,
cols,
jb,
jb,
1.0f,
W_ptr,
n,
(long long)nb * n,
T_ptr,
nb,
(long long)nb * nb,
0.0f,
Z_ptr,
n,
(long long)nb * n,
B,
true
);
bgemm(
h,
CUBLAS_OP_N,
CUBLAS_OP_N,
cols,
m,
jb,
-1.0f,
Z_ptr,
n,
(long long)nb * n,
Vp,
n,
(long long)n * n,
1.0f,
Atr,
n,
(long long)n * n,
B,
true
);
}
restore_r_kernel<<<B, 256>>>(
H_ptr,
R_ptr,
B,
n,
k,
jb,
nb
);
}
check_launch("fast n1024");
return std::make_tuple(H, tau);
}
auto res = at::geqrf(A.contiguous());
return std::make_tuple(std::get<0>(res), std::get<1>(res));
}
"""
_b200_qr_module = load_inline(
name="b200_qr_tf32_middle1024_v1",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["batched_qr_forward"],
with_cuda=True,
extra_cflags=[
"-O3",
"-DNDEBUG",
],
extra_cuda_cflags=[
"-O3",
"-DNDEBUG",
"--use_fast_math",
"--extra-device-vectorization",
"--fmad=true",
"--ftz=true",
"--prec-div=false",
"--prec-sqrt=false",
"-Xptxas=-O3",
"-Xptxas=-dlcm=ca",
],
extra_ldflags=[
"-lcublas",
],
verbose=False,
)
def custom_kernel(A: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
return _b200_qr_module.batched_qr_forward(A)scrolls · 1195 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