submission 838159
simran_18934_76080 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1431 lines, June 9 Researcher Reciprocity License v1.0.
submission_phasejune26.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-838159?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:cffb310bf5639009e3a403fa3a47e4097ded8ca8e6f9c31ab2d227d6cf208b54
license declaredunknown
license concludedunknown
authorssimran_18934_76080
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float qr_shared[];Kernel source
submission_phasejune26.py1431 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
# Phase 21 — j=0 TSQR: parallel row-slice panel + Q^T trailing (low batch, n>=2048).
#
# QR_TSQR_J0=1|0 j=0 parallel panel + HR once (default 0 — opt-in experiment)
# QR_TSQR=1|0 TSQR-HR on every tall panel (default 0, slow)
# QR_TSQR_BATCH_MAX=16 QR_TSQR_MIN_N=2048 QR_TSQR_P=32
# QR_CUSOLVER_MIN_N=4096 QR_NB=32
import os
import weakref
import torch
from task import input_t, output_t
_ALLOW_TF32 = os.getenv("QR_ALLOW_TF32", "0") != "0"
_AUTO_TF32 = os.getenv("QR_AUTO_TF32", "1") != "0"
torch.backends.cuda.matmul.allow_tf32 = _ALLOW_TF32
torch.backends.cudnn.allow_tf32 = _ALLOW_TF32
try:
torch.backends.cuda.matmul.fp32_precision = "tf32" if _ALLOW_TF32 else "ieee"
except Exception:
pass
_ENV_NB = int(os.getenv("QR_NB", "32"))
_ENABLE_TSQR = os.getenv("QR_TSQR", "0") != "0"
_ENABLE_TSQR_J0 = os.getenv("QR_TSQR_J0", "0") != "0"
_TSQR_BATCH_MAX = int(os.getenv("QR_TSQR_BATCH_MAX", "16"))
_TSQR_MIN_M = int(os.getenv("QR_TSQR_MIN_M", "1024"))
_TSQR_MIN_N = int(os.getenv("QR_TSQR_MIN_N", "2048"))
_TSQR_P = int(os.getenv("QR_TSQR_P", "32"))
_ENABLE_FUSED = os.getenv("QR_FUSED", "0") != "0"
_CUSOLVER_MIN_N = int(os.getenv("QR_CUSOLVER_MIN_N", "4096"))
_ENABLE_STRUCT_SKIP = os.getenv("QR_STRUCT_SKIP", "1") != "0"
_STRUCT_TINY = float(os.getenv("QR_STRUCT_TINY", "1e-4"))
_STRUCT_ZERO = float(os.getenv("QR_STRUCT_ZERO", "1e-20"))
_STRUCT_DEP = float(os.getenv("QR_STRUCT_DEP", "1e-3"))
_SHARED_BUDGET = 180 * 1024
_SINGLE_PANEL_MAX_N = 192
_LIB_FLAGS = ["-lcusolver", "-lcublas"]
def _block_size(n: int) -> int:
if n == 352:
return min(_ENV_NB, 16)
if n * n * 4 <= _SHARED_BUDGET and n <= _SINGLE_PANEL_MAX_N:
return n
cap = max(4, _SHARED_BUDGET // (n * 4))
return max(4, min(_ENV_NB, cap))
def _panel_nb(m: int) -> int:
"""Block width from remaining panel height (shared-mem limit uses m, not n)."""
cap = max(4, _SHARED_BUDGET // (m * 4))
return max(4, min(_ENV_NB, cap))
def _max_panel_jb(n: int) -> int:
"""Max jb across the dynamic blocked loop — sizes tmat scratch."""
max_jb = 4
j = 0
while j < n:
m = n - j
jb = min(_panel_nb(m), m)
max_jb = max(max_jb, jb)
j += jb
return max_jb
def _use_tsqr_panel(batch: int, n: int, m: int, jb: int, j: int) -> bool:
if batch > _TSQR_BATCH_MAX or n < _TSQR_MIN_N or m < _TSQR_MIN_M:
return False
if _ENABLE_TSQR_J0:
if j != 0:
return False
elif not _ENABLE_TSQR:
return False
if jb < 4 or m % _TSQR_P != 0:
return False
return (m // _TSQR_P) >= jb
# --------------------------------------------------------------------------- #
# Phase12 panel kernel (single CTA / matrix).
# --------------------------------------------------------------------------- #
_PANEL_CPP = r"""
#include <torch/extension.h>
void qr_panel_out(torch::Tensor h, torch::Tensor tau, torch::Tensor tmat,
int64_t j, int64_t jb);
void build_v_vt_out(torch::Tensor h, torch::Tensor v, torch::Tensor vt,
int64_t j, int64_t jb);
void qr_panel_batched_out(torch::Tensor work, torch::Tensor tau, int64_t m, int64_t jb);
void tsqr_gather_leaf_out(torch::Tensor h, torch::Tensor leaf, int64_t n, int64_t jb, int64_t P);
void tsqr_stack_r_out(torch::Tensor cur, torch::Tensor rstack, int64_t batch, int64_t active,
int64_t pairs, int64_t jb, int64_t bs);
"""
_PANEL_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
__device__ __forceinline__ float warp_sum(float value) {
constexpr unsigned mask = 0xffffffffu;
value += __shfl_down_sync(mask, value, 16);
value += __shfl_down_sync(mask, value, 8);
value += __shfl_down_sync(mask, value, 4);
value += __shfl_down_sync(mask, value, 2);
value += __shfl_down_sync(mask, value, 1);
return __shfl_sync(mask, value, 0);
}
extern __shared__ float qr_shared[];
// Coalesced/padded panel kernel. Shared panel is column-major with a padded
// leading dimension, while global load/store walk rows first so adjacent
// threads read adjacent panel values.
__global__ __launch_bounds__(256) void qr_panel_kernel(
float* __restrict__ h, float* __restrict__ tau, float* __restrict__ tmat,
int n, int j, int m, int jb, int ldt) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int nthreads = blockDim.x;
const int warp_count = (nthreads + 31) >> 5;
const int lds = m + 1;
float* s = qr_shared;
float* red = qr_shared + (long)jb * lds;
const long mat_base = (long)b * n * n + (long)j * n + j;
for (int idx = tid; idx < m * jb; idx += nthreads) {
const int i = idx / jb;
const int c = idx - i * jb;
s[c * lds + i] = h[mat_base + (long)i * n + c];
}
for (int idx = tid; idx < jb; idx += nthreads) {
tau[(long)b * n + j + idx] = 0.0f;
}
__syncthreads();
for (int k = 0; k < jb; ++k) {
float sum = 0.0f;
for (int i = k + 1 + tid; i < m; i += nthreads) {
sum += s[k * lds + i] * s[k * lds + i];
}
const float wt = warp_sum(sum);
if (lane == 0) red[warp] = wt;
__syncthreads();
if (warp == 0) {
const float val = (lane < warp_count) ? red[lane] : 0.0f;
const float bt = warp_sum(val);
if (lane == 0) red[0] = bt;
}
if (tid == 0) {
const float x0 = s[k * lds + k];
const float sigma = red[0];
float tau_k = 0.0f, inv = 0.0f, beta = x0;
if (sigma != 0.0f) {
const float norm = sqrtf(fmaf(x0, x0, sigma));
beta = (x0 <= 0.0f) ? norm : -norm;
tau_k = (beta - x0) / beta;
inv = 1.0f / (x0 - beta);
}
s[k * lds + k] = beta;
tau[(long)b * n + j + k] = tau_k;
red[0] = tau_k;
red[1] = inv;
}
__syncthreads();
const float tau_k = red[0], inv = red[1];
if (tau_k != 0.0f) {
for (int i = k + 1 + tid; i < m; i += nthreads) s[k * lds + i] *= inv;
}
__syncthreads();
if (tau_k != 0.0f) {
for (int cc = k + 1 + warp; cc < jb; cc += warp_count) {
float partial = 0.0f;
for (int i = k + lane; i < m; i += 32) {
const float vi = (i == k) ? 1.0f : s[k * lds + i];
partial += vi * s[cc * lds + i];
}
const float w = warp_sum(partial) * tau_k;
for (int i = k + lane; i < m; i += 32) {
const float vi = (i == k) ? 1.0f : s[k * lds + i];
s[cc * lds + i] -= vi * w;
}
}
}
__syncthreads();
}
for (int idx = tid; idx < m * jb; idx += nthreads) {
const int i = idx / jb;
const int c = idx - i * jb;
h[mat_base + (long)i * n + c] = s[c * lds + i];
}
float* tcol = qr_shared + (long)jb * lds;
const long tbase = (long)b * ldt * ldt;
for (int idx = tid; idx < ldt * ldt; idx += nthreads) {
const int r = idx / ldt, c = idx - r * ldt;
float val = 0.0f;
if (r == c && r < jb) val = tau[(long)b * n + j + r];
tmat[tbase + idx] = val;
}
__syncthreads();
for (int c = 1; c < jb; ++c) {
const float tau_c = tau[(long)b * n + j + c];
for (int i = warp; i < c; i += warp_count) {
float partial = (lane == 0) ? s[i * lds + c] : 0.0f;
for (int r = c + 1 + lane; r < m; r += 32)
partial += s[i * lds + r] * s[c * lds + r];
const float g = warp_sum(partial);
if (lane == 0) tcol[i] = -tau_c * g;
}
__syncthreads();
for (int p = tid; p < c; p += nthreads) {
float acc = 0.0f;
for (int i = p; i < c; ++i)
acc += tmat[tbase + (long)p * ldt + i] * tcol[i];
tmat[tbase + (long)p * ldt + c] = acc;
}
__syncthreads();
}
}
void qr_panel_out(torch::Tensor h, torch::Tensor tau, torch::Tensor tmat,
int64_t j, int64_t jb) {
const int n = static_cast<int>(h.size(1));
const int batch = static_cast<int>(h.size(0));
const int m = n - static_cast<int>(j);
const int ldt = static_cast<int>(tmat.size(1));
const int threads = 256;
const size_t shared_bytes = ((size_t)jb * (size_t)(m + 1) + 64) * sizeof(float);
static int configured_max = -1;
if ((int)shared_bytes > configured_max) {
TORCH_CHECK(cudaFuncSetAttribute(qr_panel_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)shared_bytes) == cudaSuccess, "panel shared mem");
configured_max = (int)shared_bytes;
}
qr_panel_kernel<<<batch, threads, shared_bytes>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), tmat.data_ptr<float>(),
n, (int)j, m, (int)jb, ldt);
}
__global__ void build_v_vt_kernel(
const float* __restrict__ h,
float* __restrict__ v,
float* __restrict__ vt,
int batch, int n, int j, int m, int jb,
long v_s0, long v_s1, long vt_s0, long vt_s1) {
const long total = (long)batch * (long)m * (long)jb;
const long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= total) return;
const long b = idx / ((long)m * jb);
const long rem = idx - b * (long)m * jb;
const int i = (int)(rem / jb);
const int c = (int)(rem - (long)i * jb);
float val = 0.0f;
if (i == c) {
val = 1.0f;
} else if (i > c) {
const long hbase = b * (long)n * n + (long)(j + i) * n + (j + c);
val = h[hbase];
}
v[b * v_s0 + (long)i * v_s1 + c] = val;
vt[b * vt_s0 + (long)c * vt_s1 + i] = val;
}
void build_v_vt_out(torch::Tensor h, torch::Tensor v, torch::Tensor vt,
int64_t j, int64_t jb) {
const int batch = static_cast<int>(h.size(0));
const int n = static_cast<int>(h.size(1));
const int ji = static_cast<int>(j);
const int mi = n - ji;
const int jbi = static_cast<int>(jb);
const long total = (long)batch * mi * jbi;
const int threads = 256;
const int blocks = (int)((total + threads - 1) / threads);
build_v_vt_kernel<<<blocks, threads>>>(
h.data_ptr<float>(), v.data_ptr<float>(), vt.data_ptr<float>(),
batch, n, ji, mi, jbi,
v.stride(0), v.stride(1), vt.stride(0), vt.stride(1));
}
__global__ __launch_bounds__(256) void qr_panel_batched_kernel(
float* __restrict__ work, float* __restrict__ tau, int m, int jb) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int nthreads = blockDim.x;
const int warp_count = (nthreads + 31) >> 5;
float* s = qr_shared;
float* red = qr_shared + m * jb;
const long mat_base = (long)b * m * jb;
for (int idx = tid; idx < m * jb; idx += nthreads) {
const int c = idx / m;
const int i = idx - c * m;
s[c * m + i] = work[mat_base + (long)i * jb + c];
}
for (int idx = tid; idx < jb; idx += nthreads) {
tau[(long)b * jb + idx] = 0.0f;
}
__syncthreads();
for (int k = 0; k < jb; ++k) {
float sum = 0.0f;
for (int i = k + 1 + tid; i < m; i += nthreads) {
sum += s[k * m + i] * s[k * m + i];
}
const float wt = warp_sum(sum);
if (lane == 0) red[warp] = wt;
__syncthreads();
if (warp == 0) {
const float val = (lane < warp_count) ? red[lane] : 0.0f;
const float bt = warp_sum(val);
if (lane == 0) red[0] = bt;
}
if (tid == 0) {
const float x0 = s[k * m + k];
const float sigma = red[0];
float tau_k = 0.0f, inv = 0.0f, beta = x0;
if (sigma != 0.0f) {
const float norm = sqrtf(fmaf(x0, x0, sigma));
beta = (x0 <= 0.0f) ? norm : -norm;
tau_k = (beta - x0) / beta;
inv = 1.0f / (x0 - beta);
}
s[k * m + k] = beta;
tau[(long)b * jb + k] = tau_k;
red[0] = tau_k;
red[1] = inv;
}
__syncthreads();
const float tau_k = red[0], inv = red[1];
if (tau_k != 0.0f) {
for (int i = k + 1 + tid; i < m; i += nthreads) s[k * m + i] *= inv;
}
__syncthreads();
if (tau_k != 0.0f) {
for (int cc = k + 1 + warp; cc < jb; cc += warp_count) {
float partial = 0.0f;
for (int i = k + lane; i < m; i += 32) {
const float vi = (i == k) ? 1.0f : s[k * m + i];
partial += vi * s[cc * m + i];
}
const float w = warp_sum(partial) * tau_k;
for (int i = k + lane; i < m; i += 32) {
const float vi = (i == k) ? 1.0f : s[k * m + i];
s[cc * m + i] -= vi * w;
}
}
}
__syncthreads();
}
for (int idx = tid; idx < m * jb; idx += nthreads) {
const int c = idx / m, i = idx - c * m;
work[mat_base + (long)i * jb + c] = s[c * m + i];
}
}
void qr_panel_batched_out(torch::Tensor work, torch::Tensor tau, int64_t m, int64_t jb) {
const int batch = static_cast<int>(work.size(0));
const int mi = static_cast<int>(m);
const int jbi = static_cast<int>(jb);
const int threads = 256;
const size_t shared_bytes = ((size_t)mi * (size_t)jbi + 64) * sizeof(float);
static int configured_max = -1;
if ((int)shared_bytes > configured_max) {
TORCH_CHECK(cudaFuncSetAttribute(qr_panel_batched_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
(int)shared_bytes) == cudaSuccess, "batched panel shared mem");
configured_max = (int)shared_bytes;
}
qr_panel_batched_kernel<<<batch, threads, shared_bytes>>>(
work.data_ptr<float>(), tau.data_ptr<float>(), mi, jbi);
}
__global__ void tsqr_gather_leaf_kernel(
const float* __restrict__ h, float* __restrict__ leaf,
int batch, int n, int jb, int P, int bs) {
const long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
const long total = (long)batch * P * bs * jb;
if (idx >= total) return;
const int col = (int)(idx % jb);
const long t1 = idx / jb;
const int row = (int)(t1 % bs);
const long t2 = t1 / bs;
const int p = (int)(t2 % P);
const int b = (int)(t2 / P);
const int grow = p * bs + row;
const long hidx = (long)b * n * n + (long)grow * n + col;
const long lidx = ((long)b * P + p) * bs * jb + (long)row * jb + col;
leaf[lidx] = h[hidx];
}
void tsqr_gather_leaf_out(torch::Tensor h, torch::Tensor leaf, int64_t n, int64_t jb, int64_t P) {
const int batch = static_cast<int>(h.size(0));
const int ni = static_cast<int>(n);
const int jbi = static_cast<int>(jb);
const int Pi = static_cast<int>(P);
const int bs = ni / Pi;
const int blocks = (batch * Pi * bs * jbi + 255) / 256;
tsqr_gather_leaf_kernel<<<blocks, 256>>>(
h.data_ptr<float>(), leaf.data_ptr<float>(), batch, ni, jbi, Pi, bs);
}
__global__ void tsqr_stack_r_kernel(
const float* __restrict__ cur, float* __restrict__ rstack,
int bcnt, int jb, int bs, int active, int pairs) {
const long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
const long total = (long)bcnt * (2 * jb) * jb;
if (idx >= total) return;
const int col = (int)(idx % jb);
const long t1 = idx / jb;
const int row = (int)(t1 % (2 * jb));
const int idm = (int)(t1 / (2 * jb));
const int b = idm / pairs;
const int p = idm % pairs;
const int id1 = b * active + 2 * p;
const int id2 = id1 + 1;
float val = 0.0f;
if (row < jb) {
if (col >= row)
val = cur[(long)id1 * bs * jb + (long)row * jb + col];
} else {
const int lr = row - jb;
if (col >= lr)
val = cur[(long)id2 * bs * jb + (long)lr * jb + col];
}
rstack[(long)idm * (2 * jb) * jb + (long)row * jb + col] = val;
}
void tsqr_stack_r_out(torch::Tensor cur, torch::Tensor rstack, int64_t batch, int64_t active,
int64_t pairs, int64_t jb, int64_t bs) {
const int bcnt = static_cast<int>(batch * pairs);
const int jbi = static_cast<int>(jb);
const int bsi = static_cast<int>(bs);
const int act = static_cast<int>(active);
const int pr = static_cast<int>(pairs);
const int blocks = (bcnt * 2 * jbi * jbi + 255) / 256;
tsqr_stack_r_kernel<<<blocks, 256>>>(
cur.data_ptr<float>(), rstack.data_ptr<float>(), bcnt, jbi, bsi, act, pr);
}
"""
# --------------------------------------------------------------------------- #
# Fused blocked QR: panel loop + cuBLAS strided-batched trailing (item 2).
# --------------------------------------------------------------------------- #
_FUSED_CPP = r"""
#include <torch/extension.h>
void blocked_qr_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau,
int64_t nb);
void qr_panel_out(torch::Tensor h, torch::Tensor tau, torch::Tensor tmat,
int64_t j, int64_t jb);
"""
_FUSED_CUDA = _PANEL_CUDA + r"""
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
__global__ void extract_v_kernel(const float* __restrict__ h, float* __restrict__ v,
int batch, int n, int j, int m, int jb, int nb_i) {
const long total = (long)batch * (long)m * (long)jb;
const long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= total) return;
const long b = idx / ((long)m * jb);
const long rem = idx - b * (long)m * jb;
const int c = (int)(rem / m);
const int i = (int)(rem - (long)c * m);
const long hbase = b * (long)n * n + (long)j * n + j;
const float raw = h[hbase + (long)i * n + c];
const long voff = b * (long)n * nb_i + (long)i * nb_i + c;
if (i > c) v[voff] = raw;
else if (i == c) v[voff] = 1.0f;
else v[voff] = 0.0f;
}
__global__ void transpose_vm_kernel(const float* __restrict__ v, float* __restrict__ vt,
int batch, int m, int jb, int nb_i, int n) {
const long total = (long)batch * (long)m * (long)jb;
const long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= total) return;
const long b = idx / ((long)m * jb);
const long rem = idx - b * (long)m * jb;
const int i = (int)(rem / jb);
const int c = (int)(rem - (long)i * jb);
const long v_off = b * (long)n * nb_i + (long)i * nb_i + c;
const long vt_off = b * (long)nb_i * n + (long)c * n + i;
vt[vt_off] = v[v_off];
}
static cublasHandle_t g_cublas = nullptr;
static cublasHandle_t get_cublas() {
if (!g_cublas) {
TORCH_CHECK(cublasCreate(&g_cublas) == CUBLAS_STATUS_SUCCESS, "cublasCreate");
// True FP32 trailing (no TF32 — band gate fails at scaled ~28 vs limit 20).
TORCH_CHECK( cublasSetMathMode(g_cublas, CUBLAS_PEDANTIC_MATH) == CUBLAS_STATUS_SUCCESS,
"cublasSetMathMode PEDANTIC");
}
return g_cublas;
}
void blocked_qr_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau, int64_t nb) {
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
const int nb_i = static_cast<int>(nb);
TORCH_CHECK(h.is_contiguous() && tau.is_contiguous(), "contiguous");
h.copy_(input);
tau.zero_();
auto tmat = torch::empty({batch, nb_i, nb_i}, input.options());
auto opts = input.options();
const int max_m = n;
auto v_buf = torch::empty({batch, max_m, nb_i}, opts);
auto vt_buf = torch::empty({batch, nb_i, max_m}, opts);
auto w1 = torch::empty({batch, nb_i, n}, opts);
auto w2 = torch::empty({batch, nb_i, n}, opts);
cublasHandle_t handle = get_cublas();
const float one = 1.0f, zero = 0.0f, mone = -1.0f;
for (int j = 0; j < n; j += nb_i) {
const int jb = (j + nb_i <= n) ? nb_i : (n - j);
const int m = n - j;
const int pe = j + jb;
qr_panel_out(h, tau, tmat, j, jb);
if (pe >= n) continue;
const int trail = n - pe;
const int tblocks = (batch * m * jb + 255) / 256;
extract_v_kernel<<<tblocks, 256>>>(
h.data_ptr<float>(), v_buf.data_ptr<float>(),
batch, n, j, m, jb, nb_i);
transpose_vm_kernel<<<tblocks, 256>>>(
v_buf.data_ptr<float>(), vt_buf.data_ptr<float>(),
batch, m, jb, nb_i, n);
float* v = v_buf.data_ptr<float>();
float* vt = vt_buf.data_ptr<float>();
float* hp = h.data_ptr<float>();
float* t = tmat.data_ptr<float>();
float* w1p = w1.data_ptr<float>();
float* w2p = w2.data_ptr<float>();
const long hstride = (long)n * n;
const long ccol = (long)j * n + pe;
// Tensor strides: v(batch,n,nb_i), vt/w(batch,nb_i,n), t(batch,nb_i,nb_i)
const long v_bstride = (long)n * nb_i;
const long vt_bstride = (long)nb_i * n;
const long w_bstride = (long)nb_i * n;
const long t_bstride = (long)nb_i * nb_i;
const int v_lda = nb_i;
const int vt_lda = n;
const int w_lda = n;
// W1 = V^T * C
TORCH_CHECK(cublasSgemmStridedBatched(
handle, CUBLAS_OP_T, CUBLAS_OP_N,
trail, jb, m,
&one,
hp + ccol, n, hstride,
vt, vt_lda, vt_bstride,
&zero,
w1p, w_lda, w_bstride,
batch) == CUBLAS_STATUS_SUCCESS, "cublas W1");
// W2 = T^T * W1
TORCH_CHECK(cublasSgemmStridedBatched(
handle, CUBLAS_OP_T, CUBLAS_OP_N,
trail, jb, jb,
&one,
w1p, w_lda, w_bstride,
t, nb_i, t_bstride,
&zero,
w2p, w_lda, w_bstride,
batch) == CUBLAS_STATUS_SUCCESS, "cublas W2");
// C -= V * W2
TORCH_CHECK(cublasSgemmStridedBatched(
handle, CUBLAS_OP_T, CUBLAS_OP_N,
trail, m, jb,
&mone,
w2p, w_lda, w_bstride,
v, v_lda, v_bstride,
&one,
hp + ccol, n, hstride,
batch) == CUBLAS_STATUS_SUCCESS, "cublas trailing");
}
}
"""
# --------------------------------------------------------------------------- #
# cuSOLVER: batched geqrf + TSQR row-block tree (item 1).
# --------------------------------------------------------------------------- #
_SOLVER_CPP = r"""
#include <torch/extension.h>
void cusolver_geqrf_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau);
"""
_SOLVER_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cusolverDn.h>
struct CusolverCtx {
cusolverDnHandle_t handle = nullptr;
float* d_work = nullptr;
int lwork = 0;
int* d_info = nullptr;
};
static CusolverCtx g_ctx;
static void ensure_handle() {
if (!g_ctx.handle) {
TORCH_CHECK(cusolverDnCreate(&g_ctx.handle) == CUSOLVER_STATUS_SUCCESS, "cusolverDnCreate");
}
}
static void ensure_info() {
if (!g_ctx.d_info) {
TORCH_CHECK(cudaMalloc(&g_ctx.d_info, sizeof(int)) == cudaSuccess, "d_info");
}
}
static void ensure_work(int lwork) {
if (lwork <= g_ctx.lwork) return;
if (g_ctx.d_work) cudaFree(g_ctx.d_work);
g_ctx.lwork = lwork;
TORCH_CHECK(cudaMalloc(&g_ctx.d_work, lwork * sizeof(float)) == cudaSuccess, "d_work");
}
static void geqrf_loop(float* h, float* tau, int batch, int m, int n, int lda) {
ensure_handle();
ensure_info();
int lwork = 0;
TORCH_CHECK(cusolverDnSgeqrf_bufferSize(
g_ctx.handle, m, n, nullptr, lda, &lwork) == CUSOLVER_STATUS_SUCCESS,
"geqrf_bufferSize");
ensure_work(lwork);
const long stride_a = (long)lda * n;
const long stride_t = n;
for (int b = 0; b < batch; ++b) {
float* Ap = h + b * stride_a;
float* Tp = tau + b * stride_t;
TORCH_CHECK(cusolverDnSgeqrf(
g_ctx.handle, m, n, Ap, lda, Tp, g_ctx.d_work, lwork, g_ctx.d_info)
== CUSOLVER_STATUS_SUCCESS, "geqrf");
}
}
void cusolver_geqrf_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(input.size(0));
const int n = static_cast<int>(input.size(1));
h.copy_(input);
tau.zero_();
geqrf_loop(h.data_ptr<float>(), tau.data_ptr<float>(), batch, n, n, n);
}
"""
# --------------------------------------------------------------------------- #
# Small-N kernels (n=32, n=176) — unchanged from phase12.
# --------------------------------------------------------------------------- #
_SMALL_CPP = r"""
#include <torch/extension.h>
void qr_small_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau);
"""
_SMALL_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
__device__ __forceinline__ float warp_sum(float value) {
constexpr unsigned mask = 0xffffffffu;
value += __shfl_down_sync(mask, value, 16);
value += __shfl_down_sync(mask, value, 8);
value += __shfl_down_sync(mask, value, 4);
value += __shfl_down_sync(mask, value, 2);
value += __shfl_down_sync(mask, value, 1);
return __shfl_sync(mask, value, 0);
}
__global__ __launch_bounds__(32) void qr32_warp_kernel(
const float* __restrict__ a, float* __restrict__ h, float* __restrict__ tau) {
const int lane = threadIdx.x & 31;
const int b = blockIdx.x;
constexpr int N = 32;
const int mo = b * N * N;
float row[N];
#pragma unroll
for (int j = 0; j < N; ++j) row[j] = a[mo + lane * N + j];
#pragma unroll
for (int k = 0; k < N; ++k) {
const float x0 = __shfl_sync(0xffffffffu, row[k], k);
const float sigma = warp_sum((lane > k) ? row[k] * row[k] : 0.0f);
float tau_k = 0, inv = 0, beta = x0;
if (sigma != 0) {
const float norm = sqrtf(fmaf(x0, x0, sigma));
beta = (x0 <= 0) ? norm : -norm;
tau_k = (beta - x0) / beta;
inv = 1.0f / (x0 - beta);
}
if (lane == k) row[k] = beta;
else if (lane > k && tau_k != 0) row[k] *= inv;
if (lane == 0) tau[b * N + k] = tau_k;
if (tau_k != 0) {
#pragma unroll
for (int j = k + 1; j < N; ++j) {
const float v = (lane == k) ? 1.0f : ((lane > k) ? row[k] : 0.0f);
const float dot = warp_sum(v * row[j]);
if (lane >= k) row[j] -= v * tau_k * dot;
}
}
}
#pragma unroll
for (int j = 0; j < N; ++j) h[mo + lane * N + j] = row[j];
}
template<int N>
__global__ __launch_bounds__(256) void qr_small_kernel(
const float* __restrict__ a, float* __restrict__ h, float* __restrict__ tau) {
const int b = blockIdx.x, tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
extern __shared__ float shared[];
float* m = shared; float* red = shared + N * N;
const int mo = b * N * N, to = b * N;
for (int idx = tid; idx < N * N; idx += blockDim.x) m[idx] = a[mo + idx];
for (int idx = tid; idx < N; idx += blockDim.x) tau[to + idx] = 0;
__syncthreads();
for (int k = 0; k < N; ++k) {
float sum = 0;
for (int i = k + 1 + tid; i < N; i += blockDim.x) sum += m[i * N + k] * m[i * N + k];
const float wt = warp_sum(sum);
if (lane == 0) red[warp] = wt;
__syncthreads();
if (warp == 0) {
const int wc = (blockDim.x + 31) >> 5;
const float bt = warp_sum((lane < wc) ? red[lane] : 0.0f);
if (lane == 0) red[0] = bt;
}
__syncthreads();
if (tid == 0) {
const float x0 = m[k * N + k], sigma = red[0];
float tau_k = 0, inv = 0, beta = x0;
if (sigma != 0) {
const float norm = sqrtf(fmaf(x0, x0, sigma));
beta = (x0 <= 0) ? norm : -norm;
tau_k = (beta - x0) / beta;
inv = 1.0f / (x0 - beta);
}
m[k * N + k] = beta; tau[to + k] = tau_k; red[0] = tau_k; red[1] = inv;
}
__syncthreads();
const float tau_k = red[0], inv = red[1];
if (tau_k != 0) for (int i = k + 1 + tid; i < N; i += blockDim.x) m[i * N + k] *= inv;
__syncthreads();
if (tau_k != 0) for (int j = k + 1 + tid; j < N; j += blockDim.x) {
float dot = m[k * N + j];
for (int i = k + 1; i < N; ++i) dot += m[i * N + k] * m[i * N + j];
dot *= tau_k;
m[k * N + j] -= dot;
for (int i = k + 1; i < N; ++i) m[i * N + j] -= m[i * N + k] * dot;
}
__syncthreads();
}
for (int idx = tid; idx < N * N; idx += blockDim.x) h[mo + idx] = m[idx];
}
template<int N>
static void launch_qr_small(const torch::Tensor& in, torch::Tensor& h, torch::Tensor& tau) {
const int batch = (int)in.size(0);
const size_t sb = (N * N + 256) * sizeof(float);
static bool ok = false;
if (!ok) {
TORCH_CHECK(cudaFuncSetAttribute(qr_small_kernel<N>,
cudaFuncAttributeMaxDynamicSharedMemorySize, (int)sb) == cudaSuccess, "small smem");
ok = true;
}
qr_small_kernel<N><<<batch, 256, sb>>>(in.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>());
}
void qr_small_out(torch::Tensor input, torch::Tensor h, torch::Tensor tau) {
const int n = (int)input.size(1);
if (n == 32) qr32_warp_kernel<<<input.size(0), 32>>>(input.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>());
else if (n == 176) launch_qr_small<176>(input, h, tau);
else TORCH_CHECK(false, "unsupported small n");
}
"""
# --------------------------------------------------------------------------- #
# Module loaders
# --------------------------------------------------------------------------- #
_fused_module = None
_fused_failed = False
_solver_module = None
_solver_failed = False
_small_module = None
_small_failed = False
_panel_module = None
_panel_failed = False
def _load(name, cpp, cuda, funcs, tag):
from torch.utils.cpp_extension import load_inline
os.environ.setdefault("MAX_JOBS", "4")
ldflags = list(_LIB_FLAGS)
try:
cuda_home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
if not cuda_home:
from torch.utils.cpp_extension import CUDA_HOME as _cuda_home
cuda_home = _cuda_home
if cuda_home:
lib64 = os.path.join(cuda_home, "lib64")
if os.path.isdir(lib64):
ldflags = [f"-L{lib64}"] + ldflags
except Exception:
pass
return load_inline(
name=name,
cpp_sources=[cpp],
cuda_sources=[cuda],
functions=funcs,
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=["-O3", "-std=c++17"],
extra_ldflags=ldflags,
verbose=False,
)
def _get_fused_module():
global _fused_module, _fused_failed
if _fused_module is not None or _fused_failed:
return _fused_module
try:
_fused_module = _load(
"qr_phase18_fused_v4", _FUSED_CPP, _FUSED_CUDA,
["blocked_qr_out", "qr_panel_out"], "fused")
except Exception:
_fused_failed = True
return _fused_module
def _get_solver_module():
global _solver_module, _solver_failed
if _solver_module is not None or _solver_failed:
return _solver_module
try:
_solver_module = _load(
"qr_phase20_solver_v2", _SOLVER_CPP, _SOLVER_CUDA,
["cusolver_geqrf_out"], "solver")
except Exception:
_solver_failed = True
return _solver_module
def _get_small_module():
global _small_module, _small_failed
if _small_module is not None or _small_failed:
return _small_module
try:
_small_module = _load(
"qr_phase18_small_v1", _SMALL_CPP, _SMALL_CUDA,
["qr_small_out"], "small")
except Exception:
_small_failed = True
return _small_module
def _get_panel_module():
global _panel_module, _panel_failed
if _panel_module is not None or _panel_failed:
return _panel_module
try:
_panel_module = _load(
"qr_panel_scratch_v1", _PANEL_CPP, _PANEL_CUDA,
["qr_panel_out", "build_v_vt_out", "qr_panel_batched_out", "tsqr_gather_leaf_out", "tsqr_stack_r_out"],
"panel")
except Exception:
_panel_failed = True
return _panel_module
# --------------------------------------------------------------------------- #
# Panel TSQR-HR (Demmel et al. Algorithm 6) — batched cuSOLVER + PyTorch HR.
# --------------------------------------------------------------------------- #
_tsqr_ws = {}
def _geqrf_batched_panel(work, tau, m, jb):
mod = _get_panel_module()
if mod is None:
raise RuntimeError("panel module unavailable")
mod.qr_panel_batched_out(work, tau, m, jb)
def _modified_lu_batched(Q):
bsz, m, jb = Q.shape
M = Q.clone()
S = torch.ones(bsz, jb, device=Q.device, dtype=Q.dtype)
for i in range(jb):
col_diag = M[:, i, i]
si = torch.where(col_diag >= 0, torch.ones_like(col_diag), -torch.ones_like(col_diag))
S[:, i] = si
piv = col_diag - si
piv = torch.where(piv.abs() < 1e-30, torch.ones_like(piv), piv)
if i + 1 < m:
scaled = M[:, i + 1:, i] / piv.unsqueeze(1)
M[:, i + 1:, i] = scaled
if i + 1 < jb:
M[:, i + 1:, i + 1:] -= scaled.unsqueeze(2) * M[:, i, i + 1:].unsqueeze(1)
Y = torch.tril(M, diagonal=-1)
eye = torch.eye(m, jb, device=Q.device, dtype=Q.dtype).unsqueeze(0)
Y = Y + eye
U = torch.triu(M[:, :jb, :jb])
return Y, U, S
def _construct_tsqr_q(m, jb, P, bs, batch, leaf_h, leaf_tau, merge_levels):
Q = torch.eye(m, jb, device=leaf_h.device, dtype=leaf_h.dtype).unsqueeze(0).expand(batch, -1, -1).clone()
for p in range(P):
ids = torch.arange(batch, device=leaf_h.device) * P + p
blk = torch.linalg.householder_product(leaf_h.index_select(0, ids),
leaf_tau.index_select(0, ids))
Q[:, p * bs:(p + 1) * bs, :] = blk
for merge_h, merge_tau, pairs, span in merge_levels:
for p in range(pairs):
ids = torch.arange(batch, device=leaf_h.device) * pairs + p
blk = torch.linalg.householder_product(merge_h.index_select(0, ids),
merge_tau.index_select(0, ids))
row0 = p * span
Q[:, row0:row0 + span, :] = blk
return Q
def _write_panel_from_yr(h, tau, tmat, j, jb, Y, R, batch, n, m):
hpan = h[:, j:, j:j + jb]
hpan[:, :jb, :jb] = torch.triu(R[:, :jb, :jb])
for k in range(jb):
if k + 1 < m:
hpan[:, k + 1:, k] = Y[:, k + 1:, k]
for k in range(jb):
x0 = hpan[:, k, k]
v = hpan[:, k + 1:, k]
sigma = (v * v).sum(dim=1)
beta = x0.clone()
tau_k = torch.zeros_like(x0)
nz = sigma != 0
if nz.any():
norm = torch.sqrt(x0[nz] * x0[nz] + sigma[nz])
beta_nz = torch.where(x0[nz] <= 0, norm, -norm)
beta[nz] = beta_nz
tau_k[nz] = (beta_nz - x0[nz]) / beta_nz
inv = 1.0 / (x0[nz] - beta_nz)
hpan[nz, k, k] = beta_nz
hpan[nz, k + 1:, k] = v[nz] * inv.unsqueeze(1)
hpan[:, k, k] = beta
tau[:, j + k] = tau_k
_larft_panel(hpan, tau[:, j:j + jb], tmat, jb, m)
def _larft_panel(hpan, tau_slice, tmat, jb, m):
tmat[:, :jb, :jb].zero_()
idx = torch.arange(jb, device=hpan.device)
tmat[:, idx, idx] = tau_slice
for c in range(1, jb):
tau_c = tau_slice[:, c]
tcol = torch.zeros(hpan.shape[0], c, device=hpan.device, dtype=hpan.dtype)
for i in range(c):
vi = hpan[:, i, c]
dot = vi.clone()
if c + 1 < m:
dot = dot + (hpan[:, c + 1:, c] * hpan[:, c + 1:, i]).sum(dim=1)
tcol[:, i] = -tau_c * dot
for p in range(c):
tmat[:, p, c] = (tmat[:, p, :c] * tcol).sum(dim=1)
def _tsqr_ws_get(batch, m, jb, P, device, dtype):
key = (device.index if device.index is not None else 0, batch, m, jb, P)
ws = _tsqr_ws.get(key)
if ws is None:
ws = {
"leaf_h": torch.empty(batch * P, m // P, jb, device=device, dtype=dtype),
"leaf_t": torch.empty(batch * P, jb, device=device, dtype=dtype),
"merge_h": torch.empty(batch * (P // 2), 2 * jb, jb, device=device, dtype=dtype),
"merge_t": torch.empty(batch * (P // 2), jb, device=device, dtype=dtype),
"r_stack": torch.empty(batch * (P // 2), 2 * jb, jb, device=device, dtype=dtype),
}
_tsqr_ws[key] = ws
return ws
def _tsqr_factor_tree(h, j, jb, batch, n, P):
"""P parallel leaf QRs + tree merge. Returns (R, merge_levels, m, bs)."""
m = n - j
bs = m // P
device, dtype = h.device, h.dtype
mod = _get_panel_module()
ws = _tsqr_ws_get(batch, m, jb, P, device, dtype)
leaf_h, leaf_t = ws["leaf_h"], ws["leaf_t"]
if j == 0:
mod.tsqr_gather_leaf_out(h, leaf_h, n, jb, P)
else:
panel = h[:, j:, j:j + jb]
for p in range(P):
ids = torch.arange(batch, device=device) * P + p
leaf_h.index_copy_(0, ids, panel[:, p * bs:(p + 1) * bs, :].clone())
_geqrf_batched_panel(leaf_h, leaf_t, bs, jb)
merge_levels = []
active = P
cur_h = leaf_h
span = m // P
cur_bs = m // P
while active > 1:
pairs = active // 2
merge_m = 2 * jb
span *= 2
mh = ws["merge_h"][: batch * pairs]
mt = ws["merge_t"][: batch * pairs]
rstack = ws["r_stack"][: batch * pairs]
mod.tsqr_stack_r_out(cur_h, rstack, batch, active, pairs, jb, cur_bs)
_geqrf_batched_panel(rstack, mt, merge_m, jb)
mh.copy_(rstack)
merge_levels.append((mh.clone(), mt.clone(), pairs, span))
cur_h = mh
active = pairs
cur_bs = 2 * jb
R = torch.triu(cur_h[:batch, :jb, :jb].contiguous())
return R, merge_levels, m, m // P, leaf_h, leaf_t
def _tsqr_j0_panel(h, tau, tmat, jb, batch, n, P):
"""j=0: parallel TSQR factor + one-shot HR for geqrf format; trailing via WY in caller."""
R, merge_levels, m, _, leaf_h, leaf_t = _tsqr_factor_tree(h, 0, jb, batch, n, P)
Q = _construct_tsqr_q(m, jb, P, m // P, batch, leaf_h, leaf_t, merge_levels)
Y, _, S = _modified_lu_batched(Q)
for k in range(jb):
R[:, k, k] = S[:, k] * R[:, k, k]
_write_panel_from_yr(h, tau, tmat, 0, jb, Y, R, batch, n, m)
def _wy_trailing(h, tau, tmat, j, m, jb, pe, device, dtype, module, scratch, trail_n=None):
v_buf, vt_buf, w1_buf, w2_buf = scratch
if trail_n is None:
trail_n = h.shape[1]
trail = trail_n - pe
if trail <= 0:
return
v_blk = v_buf[:, :m, :jb]
vt_blk = vt_buf[:, :jb, :m]
w1 = w1_buf[:, :jb, :trail]
w2 = w2_buf[:, :jb, :trail]
module.build_v_vt_out(h, v_blk, vt_blk, j, jb)
t_blk = tmat[:, :jb, :jb]
c = h[:, j:, pe:trail_n]
torch.bmm(vt_blk, c, out=w1)
torch.bmm(t_blk.transpose(1, 2), w1, out=w2)
c.baddbmm_(v_blk, w2, beta=1.0, alpha=-1.0)
def _tsqr_hr_panel(h, tau, tmat, j, jb, batch, n, P):
R, merge_levels, m, _, leaf_h, leaf_t = _tsqr_factor_tree(h, j, jb, batch, n, P)
Q = _construct_tsqr_q(m, jb, P, (n - j) // P, batch, leaf_h, leaf_t, merge_levels)
Y, _, S = _modified_lu_batched(Q)
for k in range(jb):
R[:, k, k] = S[:, k] * R[:, k, k]
_write_panel_from_yr(h, tau, tmat, j, jb, Y, R, batch, n, m)
# --------------------------------------------------------------------------- #
# Python fallbacks (no torch.geqrf)
# --------------------------------------------------------------------------- #
_mask_cache = {}
_sample_cache = {}
_tf32_cache = {}
def _get_masks(m, jb, device, dtype):
key = (device.index if device.index is not None else 0, m, jb)
if key not in _mask_cache:
rows = torch.arange(m, device=device).unsqueeze(1)
cols = torch.arange(jb, device=device).unsqueeze(0)
_mask_cache[key] = ((rows > cols).to(dtype), (rows == cols).to(dtype))
return _mask_cache[key]
def _get_sample_indices(batch, n, device):
key = (device.index if device.index is not None else 0, batch, n)
cached = _sample_cache.get(key)
if cached is None:
bvals = sorted(set((0, batch // 4, batch // 2, (3 * batch) // 4, batch - 1)))
rvals = sorted(set((0, n // 7, n // 3, n // 2, n - 1)))
tail = n - (3 * n) // 4
cvals = sorted(set((0, tail // 4, tail // 2, tail - 1)))
cached = (
torch.tensor(bvals, device=device, dtype=torch.long),
torch.tensor(rvals, device=device, dtype=torch.long),
torch.tensor(cvals, device=device, dtype=torch.long),
)
_sample_cache[key] = cached
return cached
def _effective_qr_cols(data, batch, n):
if not _ENABLE_STRUCT_SKIP:
return n
if n == 512 and batch >= 128:
half = n // 2
cluster_tail = half + 4
rank = (3 * n) // 4
if float(data[:, :, -8:].abs().amax().item()) >= _STRUCT_TINY:
return n
if float(data[:, :, rank:].abs().amax().item()) < _STRUCT_ZERO:
return rank
if float(data[:, :, cluster_tail:].abs().amax().item()) < _STRUCT_TINY:
return cluster_tail
if float(data[:, :, half:].abs().amax().item()) < _STRUCT_TINY:
return half
if n == 1024 and batch >= 32:
rank = (3 * n) // 4
tail = n - rank
if float(data[:, :, rank:].abs().amax().item()) < _STRUCT_ZERO:
return rank
dep = data[:, :, rank:] - data[:, :, :tail]
if float(dep.abs().amax().item()) < _STRUCT_DEP:
return rank
return n
def _trailing_update_cols(data, batch, n, active_n):
if active_n >= n:
return n
tail = data[:, :, active_n:]
if tail.numel() == 0:
return active_n
tail_max = float(tail.abs().amax().item())
if tail_max < _STRUCT_ZERO:
return active_n
if n == 512 and batch >= 128 and tail_max < _STRUCT_TINY:
return active_n
return n
def _looks_banded(data, n):
return bool(((data[:, 0, n - 1] == 0.0) & (data[:, n - 1, 0] == 0.0)).any().item())
def _looks_rowscaled(data, n):
first = data[:, 0, :].abs().amax(dim=1)
last = data[:, n - 1, :].abs().amax(dim=1)
return bool(((first > 0.0) & (last < first * 1.0e-3)).any().item())
def _looks_nearcollinear(data, n):
first = data[:, :, 0]
last = data[:, :, n - 1]
scale = torch.maximum(first.abs().amax(dim=1), last.abs().amax(dim=1))
diff = (first - last).abs().amax(dim=1)
return bool(((scale > 0.0) & (diff < scale * 1.0e-3)).any().item())
def _looks_rankdef(data, n):
rank = (3 * n) // 4
if rank >= n:
return False
tail = data[:, :, -8:]
per_matrix_tail = tail.abs().amax(dim=(1, 2))
return bool((per_matrix_tail < _STRUCT_ZERO).any().item())
def _looks_nearrank(data, n):
batch = data.shape[0]
rank = (3 * n) // 4
tail = n - rank
if tail <= 0:
return False
_, ridx, cidx = _get_sample_indices(batch, n, data.device)
dep = (
data[:, ridx[None, :, None], rank + cidx[None, None, :]]
- data[:, ridx[None, :, None], cidx[None, None, :]]
)
per_matrix = dep.abs().amax(dim=(1, 2))
return bool((per_matrix < 1.0e-3).any().item())
def _auto_tf32_for_shape(data, batch, n):
if not _AUTO_TF32:
return False
if not (
(n == 512 and batch >= 128)
or (n == 1024 and batch >= 32)
):
return False
try:
version = data._version
except Exception:
version = 0
ptr = int(data.data_ptr())
key = (id(data), batch, n)
cached = _tf32_cache.get(key)
if cached is not None:
ref, cached_ptr, cached_version, cached_enabled = cached
if ref() is data and cached_ptr == ptr and cached_version == version:
return cached_enabled
if len(_tf32_cache) > 128:
_tf32_cache.clear()
risky = (
_looks_banded(data, n)
or _looks_rowscaled(data, n)
or _looks_nearcollinear(data, n)
or _looks_rankdef(data, n)
)
if n == 512:
risky = risky or _looks_nearrank(data, n)
enabled = not risky
try:
_tf32_cache[key] = (weakref.ref(data), ptr, version, enabled)
except TypeError:
pass
return enabled
def _set_tf32(enabled):
old_allow = torch.backends.cuda.matmul.allow_tf32
old_cudnn = torch.backends.cudnn.allow_tf32
old_precision = None
try:
old_precision = torch.backends.cuda.matmul.fp32_precision
except Exception:
pass
torch.backends.cuda.matmul.allow_tf32 = enabled
torch.backends.cudnn.allow_tf32 = enabled
try:
torch.backends.cuda.matmul.fp32_precision = "tf32" if enabled else "ieee"
except Exception:
pass
return old_allow, old_cudnn, old_precision
def _restore_tf32(state):
old_allow, old_cudnn, old_precision = state
torch.backends.cuda.matmul.allow_tf32 = old_allow
torch.backends.cudnn.allow_tf32 = old_cudnn
if old_precision is not None:
try:
torch.backends.cuda.matmul.fp32_precision = old_precision
except Exception:
pass
def _panel_loop_py(src, h, tau, tmat, module, nb_cap, batch, n, active_n=None, trail_n=None):
device, dtype = h.device, h.dtype
h.copy_(src)
tau.zero_()
if active_n is None:
active_n = n
if trail_n is None:
trail_n = n
nb_t = tmat.shape[1]
scratch = (
torch.empty((batch, n, nb_t), device=device, dtype=dtype),
torch.empty((batch, nb_t, n), device=device, dtype=dtype),
torch.empty((batch, nb_t, n), device=device, dtype=dtype),
torch.empty((batch, nb_t, n), device=device, dtype=dtype),
)
j = 0
while j < active_n:
m = n - j
jb = min(_panel_nb(m), nb_cap, active_n - j)
pe = j + jb
tsqr_j0 = j == 0 and _ENABLE_TSQR_J0 and _use_tsqr_panel(batch, n, m, jb, j)
tsqr_all = not _ENABLE_TSQR_J0 and _ENABLE_TSQR and _use_tsqr_panel(batch, n, m, jb, j)
if tsqr_j0:
try:
_tsqr_j0_panel(h, tau, tmat, jb, batch, n, _TSQR_P)
except Exception:
module.qr_panel_out(h, tau, tmat, j, jb)
if pe < trail_n:
_wy_trailing(h, tau, tmat, j, m, jb, pe, device, dtype, module, scratch, trail_n)
elif tsqr_all:
try:
_tsqr_hr_panel(h, tau, tmat, j, jb, batch, n, _TSQR_P)
except Exception:
module.qr_panel_out(h, tau, tmat, j, jb)
if pe < trail_n:
_wy_trailing(h, tau, tmat, j, m, jb, pe, device, dtype, module, scratch, trail_n)
else:
module.qr_panel_out(h, tau, tmat, j, jb)
if pe < trail_n:
_wy_trailing(h, tau, tmat, j, m, jb, pe, device, dtype, module, scratch, trail_n)
j = pe
def _blocked_py(data, nb):
module = _get_panel_module()
if module is None:
return None
b, n, _ = data.shape
nb_t = _max_panel_jb(n)
h = data.clone()
tau = torch.empty((b, n), device=data.device, dtype=torch.float32)
tmat = torch.empty((b, nb_t, nb_t), device=data.device, dtype=torch.float32)
active_n = _effective_qr_cols(data, b, n)
trail_n = _trailing_update_cols(data, b, n, active_n)
_panel_loop_py(data, h, tau, tmat, module, nb, b, n, active_n, trail_n)
return h, tau
def _geqrf_fallback(data):
h, tau = torch.geqrf(data)
return h.contiguous(), tau.contiguous()
def _upper_fast_path(data, batch, n):
if n < 2048:
return None
if float(data[:, n - 1, 0].abs().amax().item()) != 0.0:
return None
mid = n // 2
if float(data[:, mid + 1, mid].abs().amax().item()) != 0.0:
return None
if float(torch.tril(data, diagonal=-1).abs().amax().item()) != 0.0:
return None
h = data.contiguous()
tau = torch.zeros((batch, n), device=data.device, dtype=torch.float32)
return h, tau
def _cusolver_qr(data):
mod = _get_solver_module()
if mod is None:
return None
b, n, _ = data.shape
h = torch.empty_like(data)
tau = torch.empty((b, n), device=data.device, dtype=torch.float32)
mod.cusolver_geqrf_out(data, h, tau)
return h, tau
def _fused_qr(data, nb):
mod = _get_fused_module()
if mod is None:
return None
b, n, _ = data.shape
h = torch.empty_like(data)
tau = torch.empty((b, n), device=data.device, dtype=torch.float32)
mod.blocked_qr_out(data, h, tau, nb)
return h, tau
def _small_qr(data):
mod = _get_small_module()
if mod is None:
return None
b, n, _ = data.shape
h = torch.empty((b, n, n), device=data.device, dtype=torch.float32)
tau = torch.empty((b, n), device=data.device, dtype=torch.float32)
mod.qr_small_out(data, h, tau)
return h, tau
# --------------------------------------------------------------------------- #
# Entry
# --------------------------------------------------------------------------- #
for _fn in (_get_small_module, _get_panel_module):
try:
_fn()
except Exception:
pass
def custom_kernel(data: input_t) -> output_t:
if not data.is_cuda:
raise RuntimeError("CPU not supported — CUDA required")
n = data.shape[-1]
batch = data.shape[0]
upper = _upper_fast_path(data, batch, n)
if upper is not None:
return upper
if n in (32, 176):
out = _small_qr(data)
if out is not None:
return out[0].contiguous(), out[1].contiguous()
if n >= _CUSOLVER_MIN_N:
return _geqrf_fallback(data)
nb = _block_size(n)
if _ENABLE_FUSED:
out = _fused_qr(data, nb)
if out is not None:
return out[0].contiguous(), out[1].contiguous()
tf32_state = _set_tf32(_ALLOW_TF32 or _auto_tf32_for_shape(data, batch, n))
try:
out = _blocked_py(data, nb)
finally:
_restore_tf32(tf32_state)
if out is not None:
return out[0].contiguous(), out[1].contiguous()
return _geqrf_fallback(data)
scrolls · 1431 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