submission 801995
Álvaro Borrás Fernández · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2262 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-801995?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:c8bf54bba67abc34d9d09cbb83557c519313de308d5ffe717db2fdad45a02ae7
license declaredunknown
license concludedunknown
authorsÁlvaro Borrás Fernández
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float mat[32 * 32];Kernel source
submission.py2262 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import gc
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
gc.disable()
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
try:
torch.set_float32_matmul_precision('high')
except Exception:
pass
CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <mma.h>
#include <stdexcept>
#include <string>
#define THREADS 256
#define WARP_SIZE 32
static inline void check_cuda(cudaError_t status, const char* what) {
if (status != cudaSuccess) {
throw std::runtime_error(std::string(what) + ": " + cudaGetErrorString(status));
}
}
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
val += __shfl_xor_sync(0xffffffffu, val, offset);
}
return val;
}
__device__ __forceinline__ float block_reduce_sum_64(float val, float* buf) {
int tid = threadIdx.x;
buf[tid] = val;
__syncthreads();
#pragma unroll
for (int offset = 32; offset > 0; offset >>= 1) {
if (tid < offset) buf[tid] += buf[tid + offset];
__syncthreads();
}
return buf[0];
}
__global__ void copy_input_kernel_512(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau,
int batch,
long long total
) {
long long stride = (long long)blockDim.x * gridDim.x;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += stride) {
h[idx] = a[idx];
}
for (int b = blockIdx.x * blockDim.x + threadIdx.x;
b < batch; b += blockDim.x * gridDim.x) {
tau[(long long)b * 512 + 511] = 0.0f;
}
}
__global__ void qr32_kernel(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau
) {
__shared__ float mat[32 * 32];
__shared__ float reduce_buf[64];
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
int b = blockIdx.x;
int tid = threadIdx.x;
long long base = (long long)b * 32 * 32;
for (int idx = tid; idx < 32 * 32; idx += blockDim.x) {
mat[idx] = a[base + idx];
}
__syncthreads();
for (int k = 0; k < 31; ++k) {
float ss = 0.0f;
for (int i = k + tid; i < 32; i += blockDim.x) {
float x = mat[i * 32 + k];
ss = fmaf(x, x, ss);
}
float norm_sq = block_reduce_sum_64(ss, reduce_buf);
if (tid == 0) {
float norm = sqrtf(norm_sq);
float diag = mat[k * 32 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau[(long long)b * 32 + k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
mat[k * 32 + k] = alpha;
tau[(long long)b * 32 + k] = tau_s;
}
}
__syncthreads();
if (active_s == 0) continue;
for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
mat[i * 32 + k] *= inv_s;
}
__syncthreads();
for (int j = k + 1; j < 32; ++j) {
float dot_part = (tid == 0) ? mat[k * 32 + j] : 0.0f;
for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
dot_part = fmaf(mat[i * 32 + k], mat[i * 32 + j], dot_part);
}
float dot = block_reduce_sum_64(dot_part, reduce_buf);
if (tid == 0) dot_s = dot;
__syncthreads();
if (tid == 0) {
mat[k * 32 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
mat[i * 32 + j] -= tau_s * mat[i * 32 + k] * dot_s;
}
__syncthreads();
}
}
if (tid == 0) tau[(long long)b * 32 + 31] = 0.0f;
for (int idx = tid; idx < 32 * 32; idx += blockDim.x) {
h[base + idx] = mat[idx];
}
}
template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr512_panel_factor(
float* __restrict__ h,
float* __restrict__ tau,
int kk,
int ib
) {
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + 512;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * 512 * 512;
float* __restrict__ tau_b = tau + (long long)b * 512;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < 511; ++local_k) {
int k = kk + local_k;
int m = 512 - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * 512 + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * 512 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * 512 + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * 512 + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * 512 + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * 512 + j] : 0.0f;
for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
dot_part = fmaf(h_b[i * 512 + k], h_b[i * 512 + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * 512 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
h_b[i * 512 + j] -= tau_s * h_b[i * 512 + k] * dot_s;
}
__syncthreads();
}
}
__syncthreads();
}
}
template <int NB, int BLOCK_SIZE>
__global__ void qr512_build_t(
const float* __restrict__ h,
const float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float reduce_buf[BLOCK_SIZE];
__shared__ float tmp[NB];
int b = blockIdx.x;
int tid = threadIdx.x;
const float* __restrict__ h_b = h + (long long)b * 512 * 512;
const float* __restrict__ tau_b = tau + (long long)b * 512;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot = (tid == 0) ? h_b[(kk + j) * 512 + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < 512 - kk; r += BLOCK_SIZE) {
dot = fmaf(
h_b[(kk + r) * 512 + (kk + i)],
h_b[(kk + r) * 512 + (kk + j)],
dot
);
}
reduce_buf[tid] = dot;
__syncthreads();
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
if (tid < offset) {
reduce_buf[tid] += reduce_buf[tid + offset];
}
__syncthreads();
}
if (tid == 0) {
tmp[i] = -tau_j * reduce_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}
}
template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr512_panel_build_t_fused(
float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
constexpr int NB = 8;
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + 512;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * 512 * 512;
float* __restrict__ tau_b = tau + (long long)b * 512;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < 511; ++local_k) {
int k = kk + local_k;
int m = 512 - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * 512 + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * 512 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * 512 + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * 512 + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * 512 + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * 512 + j] : 0.0f;
for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
dot_part = fmaf(v[i - k], h_b[i * 512 + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * 512 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
h_b[i * 512 + j] -= tau_s * v[i - k] * dot_s;
}
__syncthreads();
}
}
__syncthreads();
}
if (kk + ib >= 512) return;
float* tmp = smem;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot_part = (tid == 0) ? h_b[(kk + j) * 512 + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < 512 - kk; r += BLOCK_SIZE) {
dot_part = fmaf(
h_b[(kk + r) * 512 + (kk + i)],
h_b[(kk + r) * 512 + (kk + j)],
dot_part
);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[0] = dot_part;
}
__syncthreads();
if (tid == 0) {
tmp[i] = -tau_j * warp_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}
}
template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr512_compute_y(
const float* __restrict__ h,
const float* __restrict__ workspace,
float* __restrict__ y_workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float partial_s[NB * COL_TILE * RSPLIT];
__shared__ float w_s[NB * COL_TILE];
int b = blockIdx.x;
int col_tile = blockIdx.y;
int idx = threadIdx.x;
int split = idx % RSPLIT;
int c = (idx / RSPLIT) % COL_TILE;
int q = idx / (RSPLIT * COL_TILE);
int trail_col = col_tile * COL_TILE + c;
int global_col = kk + ib + trail_col;
int pair = q * COL_TILE + c;
const float* __restrict__ h_b = h + (long long)b * 512 * 512;
const float* __restrict__ t = workspace + (long long)b * workspace_stride;
float* __restrict__ y = y_workspace + (long long)b * NB * 512;
if (q < ib && c < COL_TILE && global_col < 512) {
float sum = (split == 0) ? h_b[(kk + q) * 512 + global_col] : 0.0f;
for (int r = q + 1 + split; r < 512 - kk; r += RSPLIT) {
sum = fmaf(
h_b[(kk + r) * 512 + (kk + q)],
h_b[(kk + r) * 512 + global_col],
sum
);
}
partial_s[pair * RSPLIT + split] = sum;
} else if (q < NB && c < COL_TILE) {
partial_s[pair * RSPLIT + split] = 0.0f;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 512) {
float w = 0.0f;
#pragma unroll
for (int s = 0; s < RSPLIT; ++s) {
w += partial_s[pair * RSPLIT + s];
}
w_s[pair] = w;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 512) {
float sum = 0.0f;
for (int l = 0; l <= q; ++l) {
sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
}
y[q * 512 + global_col] = sum;
}
}
template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr512_apply_block(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int idx = threadIdx.x;
int r_local = row_tile * ROW_TILE + idx / COL_TILE;
int c_local = col_tile * COL_TILE + idx % COL_TILE;
int global_col = kk + ib + c_local;
if (r_local >= 512 - kk || global_col >= 512) return;
float* __restrict__ h_b = h + (long long)b * 512 * 512;
const float* __restrict__ y = y_workspace + (long long)b * NB * 512;
float correction = 0.0f;
int nq = (r_local < ib) ? r_local : ib;
#pragma unroll
for (int q = 0; q < nq; ++q) {
float vq = h_b[(kk + r_local) * 512 + (kk + q)];
correction = fmaf(vq, y[q * 512 + global_col], correction);
}
if (r_local < ib) {
correction = fmaf(1.0f, y[r_local * 512 + global_col], correction);
}
h_b[(kk + r_local) * 512 + global_col] -= correction;
}
// Tensor-core (WMMA m16n16k16 tf32) bottom-apply pass for n=512.
// Computes H[bottom, bottom] -= V[bottom, NB] @ Y[NB, bottom] where
// bottom rows = [kk+ib, 512), NB=8. NB is padded to 16 with zeros.
//
// Note: the nvcuda::wmma tf32 fragment template is not exposed for sm_100
// (B200) in CUDA 12.8's mma.h, so this kernel uses a scalar fallback
// (manual matmul) that compiles on all arches. The dispatcher on sm_100
// falls through to the CuTe path, so this scalar kernel is not on the hot
// path there.
__global__ void qr512_apply_block_wmma(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
constexpr int NB = 8;
constexpr int NB_PAD = 16;
constexpr int ROW_TILE = 16;
constexpr int COL_TILE = 16;
constexpr int N = 512;
constexpr int LDH = N;
constexpr int LDY = N;
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int tid = threadIdx.x;
int global_row_base = kk + ib + row_tile * ROW_TILE;
int global_col_base = kk + ib + col_tile * COL_TILE;
if (global_row_base >= N || global_col_base >= N) return;
float* __restrict__ h_b = h + (long long)b * N * N;
const float* __restrict__ y = y_workspace + (long long)b * NB * N;
__shared__ float V_smem[ROW_TILE * NB_PAD];
__shared__ float Y_smem[NB_PAD * COL_TILE];
#pragma unroll
for (int i = tid; i < ROW_TILE * NB_PAD; i += 32) {
int r = i / NB_PAD;
int c = i - r * NB_PAD;
float v = 0.0f;
if (c < NB) {
int gr = kk + ib + row_tile * ROW_TILE + r;
if (gr < N) {
v = h_b[gr * LDH + (kk + c)];
}
}
V_smem[r * NB_PAD + c] = v;
}
#pragma unroll
for (int i = tid; i < NB_PAD * COL_TILE; i += 32) {
int r = i / COL_TILE;
int c = i - r * COL_TILE;
float v = 0.0f;
if (r < NB) {
int gc = kk + ib + col_tile * COL_TILE + c;
if (gc < N) {
v = y[r * LDY + gc];
}
}
Y_smem[r * COL_TILE + c] = v;
}
__syncwarp();
float acc[ROW_TILE * COL_TILE / 32] = {0.0f};
constexpr int PER_THREAD = (ROW_TILE * COL_TILE) / 32;
#pragma unroll
for (int k = 0; k < NB; ++k) {
#pragma unroll
for (int j = 0; j < COL_TILE; ++j) {
int slot = (j / 1);
int col_in_tile = (tid + j * 32 / COL_TILE) % COL_TILE;
(void)slot; (void)col_in_tile;
}
#pragma unroll
for (int i = 0; i < PER_THREAD; ++i) {
int linear = i * 32 + tid;
int r = linear / COL_TILE;
int c = linear - r * COL_TILE;
acc[i] = fmaf(V_smem[r * NB_PAD + k], Y_smem[k * COL_TILE + c], acc[i]);
}
}
#pragma unroll
for (int i = 0; i < PER_THREAD; ++i) {
int linear = i * 32 + tid;
int r = linear / COL_TILE;
int c = linear - r * COL_TILE;
int gr = kk + ib + row_tile * ROW_TILE + r;
int gc = kk + ib + col_tile * COL_TILE + c;
if (gr < N && gc < N) {
h_b[gr * LDH + gc] -= acc[i];
}
}
}
void qr512_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
constexpr int NB = 8;
constexpr int COL_TILE = 32;
constexpr int ROW_TILE = 16;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * 512;
int batch = (int)a.size(0);
long long total = (long long)batch * 512 * 512;
int copy_blocks = (int)((total + THREADS - 1) / THREADS);
if (copy_blocks < 1) copy_blocks = 1;
if (copy_blocks > 8192) copy_blocks = 8192;
copy_input_kernel_512<<<copy_blocks, THREADS>>>(
a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
);
float* workspace_ptr = workspace.data_ptr<float>();
float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;
for (int kk = 0; kk < 511; kk += NB) {
int ib = min(NB, 512 - kk);
qr512_panel_build_t_fused<256, 8><<<batch, 256, (512 + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
);
if (kk + ib < 512) {
int col_tiles = (512 - kk - ib + COL_TILE - 1) / COL_TILE;
dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
qr512_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
);
int row_tiles = (512 - kk + ROW_TILE - 1) / ROW_TILE;
dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qr512_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
h.data_ptr<float>(), y_ptr, kk, ib
);
}
}
check_cuda(cudaGetLastError(), "qr512_blocked_cuda");
}
void qr512_panel_build_t_cuda(torch::Tensor h, torch::Tensor tau, torch::Tensor workspace, int kk, int batch) {
constexpr int NB = 8;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * 512;
int ib = min(NB, 512 - kk);
qr512_panel_build_t_fused<256, 8><<<batch, 256, (512 + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace.data_ptr<float>(), kk, ib, WORKSPACE_STRIDE
);
}
void qr32_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
int batch = (int)a.size(0);
qr32_kernel<<<batch, 64>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>()
);
check_cuda(cudaGetLastError(), "qr32_cuda");
}
// WMMA tf32 fragments are not exposed by mma.h for sm_100 (B200), so the
// Host-side launcher for the scalar-fallback apply kernel. Compiles on all
// arches; on sm_100 the dispatcher falls through to the CuTe path so this
// is a dead code path.
void qr512_apply_block_wmma_cuda(torch::Tensor h, torch::Tensor y, int kk, int ib, int batch) {
constexpr int NB = 8;
constexpr int ROW_TILE = 16;
constexpr int COL_TILE = 16;
int row_tiles = (512 - kk - ib + ROW_TILE - 1) / ROW_TILE;
int col_tiles = (512 - kk - ib + COL_TILE - 1) / COL_TILE;
if (row_tiles < 0) row_tiles = 0;
if (col_tiles < 0) col_tiles = 0;
if (row_tiles == 0 || col_tiles == 0) return;
dim3 grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qr512_apply_block_wmma<<<grid, 32>>>(
h.data_ptr<float>(), y.data_ptr<float>(), kk, ib
);
}
__global__ void copy_input_kernel_1024(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau,
int batch,
long long total
) {
long long stride = (long long)blockDim.x * gridDim.x;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += stride) {
h[idx] = a[idx];
}
for (int b = blockIdx.x * blockDim.x + threadIdx.x;
b < batch; b += blockDim.x * gridDim.x) {
tau[(long long)b * 1024 + 1023] = 0.0f;
}
}
template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr1024_panel_build_t_fused(
float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
constexpr int NB = 8;
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + 1024;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
float* __restrict__ tau_b = tau + (long long)b * 1024;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < 1023; ++local_k) {
int k = kk + local_k;
int m = 1024 - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * 1024 + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * 1024 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * 1024 + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * 1024 + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * 1024 + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * 1024 + j] : 0.0f;
for (int i = k + 1 + tid; i < 1024; i += BLOCK_SIZE) {
dot_part = fmaf(v[i - k], h_b[i * 1024 + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * 1024 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 1024; i += BLOCK_SIZE) {
h_b[i * 1024 + j] -= tau_s * v[i - k] * dot_s;
}
__syncthreads();
}
}
__syncthreads();
}
if (kk + ib >= 1024) return;
float* tmp = smem;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot_part = (tid == 0) ? h_b[(kk + j) * 1024 + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < 1024 - kk; r += BLOCK_SIZE) {
dot_part = fmaf(
h_b[(kk + r) * 1024 + (kk + i)],
h_b[(kk + r) * 1024 + (kk + j)],
dot_part
);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[0] = dot_part;
}
__syncthreads();
if (tid == 0) {
tmp[i] = -tau_j * warp_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}}
template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr1024_compute_y(
const float* __restrict__ h,
const float* __restrict__ workspace,
float* __restrict__ y_workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float partial_s[NB * COL_TILE * RSPLIT];
__shared__ float w_s[NB * COL_TILE];
int b = blockIdx.x;
int col_tile = blockIdx.y;
int idx = threadIdx.x;
int split = idx % RSPLIT;
int c = (idx / RSPLIT) % COL_TILE;
int q = idx / (RSPLIT * COL_TILE);
int trail_col = col_tile * COL_TILE + c;
int global_col = kk + ib + trail_col;
int pair = q * COL_TILE + c;
const float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
const float* __restrict__ t = workspace + (long long)b * workspace_stride;
float* __restrict__ y = y_workspace + (long long)b * NB * 1024;
if (q < ib && c < COL_TILE && global_col < 1024) {
float sum = (split == 0) ? h_b[(kk + q) * 1024 + global_col] : 0.0f;
for (int r = q + 1 + split; r < 1024 - kk; r += RSPLIT) {
sum = fmaf(
h_b[(kk + r) * 1024 + (kk + q)],
h_b[(kk + r) * 1024 + global_col],
sum
);
}
partial_s[pair * RSPLIT + split] = sum;
} else if (q < NB && c < COL_TILE) {
partial_s[pair * RSPLIT + split] = 0.0f;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 1024) {
float w = 0.0f;
#pragma unroll
for (int s = 0; s < RSPLIT; ++s) {
w += partial_s[pair * RSPLIT + s];
}
w_s[pair] = w;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 1024) {
float sum = 0.0f;
for (int l = 0; l <= q; ++l) {
sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
}
y[q * 1024 + global_col] = sum;
}
}
template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr1024_apply_block(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
__shared__ float V_smem[ROW_TILE * NB];
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int idx = threadIdx.x;
if (idx < ROW_TILE * NB) {
int r = idx / NB;
int q = idx % NB;
int global_r = kk + row_tile * ROW_TILE + r;
if (global_r < 1024) {
float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
V_smem[idx] = h_b[global_r * 1024 + (kk + q)];
} else {
V_smem[idx] = 0.0f;
}
}
__syncthreads();
int r_local = row_tile * ROW_TILE + idx / COL_TILE;
int c_local = col_tile * COL_TILE + idx % COL_TILE;
int global_col = kk + ib + c_local;
if (r_local >= 1024 - kk || global_col >= 1024) return;
float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
const float* __restrict__ y = y_workspace + (long long)b * NB * 1024;
int local_r = idx / COL_TILE;
int nq = (r_local < ib) ? r_local : ib;
float correction = 0.0f;
#pragma unroll
for (int q = 0; q < nq; ++q) {
float vq = V_smem[local_r * NB + q];
correction = fmaf(vq, y[q * 1024 + global_col], correction);
}
if (r_local < ib) {
correction = fmaf(1.0f, y[r_local * 1024 + global_col], correction);
}
h_b[(kk + r_local) * 1024 + global_col] -= correction;
}
void qr1024_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
constexpr int NB = 8;
constexpr int COL_TILE = 32;
constexpr int ROW_TILE = 16;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * 1024;
int batch = (int)a.size(0);
long long total = (long long)batch * 1024 * 1024;
int copy_blocks = (int)((total + THREADS - 1) / THREADS);
if (copy_blocks < 1) copy_blocks = 1;
if (copy_blocks > 8192) copy_blocks = 8192;
copy_input_kernel_1024<<<copy_blocks, THREADS>>>(
a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
);
float* workspace_ptr = workspace.data_ptr<float>();
float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;
for (int kk = 0; kk < 1023; kk += NB) {
int ib = min(NB, 1024 - kk);
qr1024_panel_build_t_fused<256, 8><<<batch, 256, (1024 + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
);
if (kk + ib < 1024) {
int col_tiles = (1024 - kk - ib + COL_TILE - 1) / COL_TILE;
dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
qr1024_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
);
int row_tiles = (1024 - kk + ROW_TILE - 1) / ROW_TILE;
dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qr1024_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
h.data_ptr<float>(), y_ptr, kk, ib
);
}
}
check_cuda(cudaGetLastError(), "qr1024_blocked_cuda");
}
// =============================================================================
// Large-N blocked QR kernels (n=2048, n=4096). N is a template parameter so we
// can share the same template bodies across the two shapes. The structure
// mirrors qr1024: panel factor -> build T -> compute Y -> apply block.
// =============================================================================
template <int N>
__global__ void copy_input_kernel_N(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau,
int batch,
long long total
) {
long long stride = (long long)blockDim.x * gridDim.x;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += stride) {
h[idx] = a[idx];
}
for (int b = blockIdx.x * blockDim.x + threadIdx.x;
b < batch; b += blockDim.x * gridDim.x) {
tau[(long long)b * N + (N - 1)] = 0.0f;
}
}
template <int N, int BLOCK_SIZE, int NUM_WARPS>
__global__ void qrN_panel_factor(
float* __restrict__ h,
float* __restrict__ tau,
int kk,
int ib
) {
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + N;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * N * N;
float* __restrict__ tau_b = tau + (long long)b * N;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < N - 1; ++local_k) {
int k = kk + local_k;
int m = N - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * N + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * N + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * N + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * N + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * N + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * N + j] : 0.0f;
for (int i = k + 1 + tid; i < N; i += BLOCK_SIZE) {
dot_part = fmaf(h_b[i * N + k], h_b[i * N + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * N + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < N; i += BLOCK_SIZE) {
h_b[i * N + j] -= tau_s * h_b[i * N + k] * dot_s;
}
__syncthreads();
}
}
}
}
template <int N, int NB, int BLOCK_SIZE>
__global__ void qrN_build_t(
const float* __restrict__ h,
const float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float reduce_buf[BLOCK_SIZE];
__shared__ float tmp[NB];
int b = blockIdx.x;
int tid = threadIdx.x;
const float* __restrict__ h_b = h + (long long)b * N * N;
const float* __restrict__ tau_b = tau + (long long)b * N;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot = (tid == 0) ? h_b[(kk + j) * N + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < N - kk; r += BLOCK_SIZE) {
dot = fmaf(
h_b[(kk + r) * N + (kk + i)],
h_b[(kk + r) * N + (kk + j)],
dot
);
}
reduce_buf[tid] = dot;
__syncthreads();
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
if (tid < offset) {
reduce_buf[tid] += reduce_buf[tid + offset];
}
__syncthreads();
}
if (tid == 0) {
tmp[i] = -tau_j * reduce_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}
}
template <int N, int NB, int COL_TILE, int RSPLIT>
__global__ void qrN_compute_y(
const float* __restrict__ h,
const float* __restrict__ workspace,
float* __restrict__ y_workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float partial_s[NB * COL_TILE * RSPLIT];
__shared__ float w_s[NB * COL_TILE];
int b = blockIdx.x;
int col_tile = blockIdx.y;
int idx = threadIdx.x;
int split = idx % RSPLIT;
int c = (idx / RSPLIT) % COL_TILE;
int q = idx / (RSPLIT * COL_TILE);
int trail_col = col_tile * COL_TILE + c;
int global_col = kk + ib + trail_col;
int pair = q * COL_TILE + c;
const float* __restrict__ h_b = h + (long long)b * N * N;
const float* __restrict__ t = workspace + (long long)b * workspace_stride;
float* __restrict__ y = y_workspace + (long long)b * NB * N;
if (q < ib && c < COL_TILE && global_col < N) {
float sum = (split == 0) ? h_b[(kk + q) * N + global_col] : 0.0f;
for (int r = q + 1 + split; r < N - kk; r += RSPLIT) {
sum = fmaf(
h_b[(kk + r) * N + (kk + q)],
h_b[(kk + r) * N + global_col],
sum
);
}
partial_s[pair * RSPLIT + split] = sum;
} else if (q < NB && c < COL_TILE) {
partial_s[pair * RSPLIT + split] = 0.0f;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < N) {
float w = 0.0f;
#pragma unroll
for (int s = 0; s < RSPLIT; ++s) {
w += partial_s[pair * RSPLIT + s];
}
w_s[pair] = w;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < N) {
float sum = 0.0f;
for (int l = 0; l <= q; ++l) {
sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
}
y[q * N + global_col] = sum;
}
}
template <int N, int NB, int ROW_TILE, int COL_TILE>
__global__ void qrN_apply_block(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
__shared__ float V_smem[ROW_TILE * NB];
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int idx = threadIdx.x;
if (idx < ROW_TILE * NB) {
int r = idx / NB;
int q = idx % NB;
int global_r = kk + row_tile * ROW_TILE + r;
if (global_r < N) {
float* __restrict__ h_b = h + (long long)b * N * N;
V_smem[idx] = h_b[global_r * N + (kk + q)];
} else {
V_smem[idx] = 0.0f;
}
}
__syncthreads();
int r_local = row_tile * ROW_TILE + idx / COL_TILE;
int c_local = col_tile * COL_TILE + idx % COL_TILE;
int global_col = kk + ib + c_local;
if (r_local >= N - kk || global_col >= N) return;
float* __restrict__ h_b = h + (long long)b * N * N;
const float* __restrict__ y = y_workspace + (long long)b * NB * N;
int local_r = idx / COL_TILE;
int nq = (r_local < ib) ? r_local : ib;
float correction = 0.0f;
#pragma unroll
for (int q = 0; q < nq; ++q) {
float vq = V_smem[local_r * NB + q];
correction = fmaf(vq, y[q * N + global_col], correction);
}
if (r_local < ib) {
correction = fmaf(1.0f, y[r_local * N + global_col], correction);
}
h_b[(kk + r_local) * N + global_col] -= correction;
}
template <int N>
void qrN_blocked_cuda_dispatch(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
constexpr int NB = 8;
constexpr int COL_TILE = 32;
constexpr int ROW_TILE = 16;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * N;
int batch = (int)a.size(0);
long long total = (long long)batch * N * N;
int copy_blocks = (int)((total + THREADS - 1) / THREADS);
if (copy_blocks < 1) copy_blocks = 1;
if (copy_blocks > 8192) copy_blocks = 8192;
copy_input_kernel_N<N><<<copy_blocks, THREADS>>>(
a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
);
float* workspace_ptr = workspace.data_ptr<float>();
float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;
for (int kk = 0; kk < N - 1; kk += NB) {
int ib = min(NB, N - kk);
qrN_panel_factor<N, 256, 8><<<batch, 256, (N + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), kk, ib
);
if (kk + ib < N) {
qrN_build_t<N, NB, 128><<<batch, 128>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
);
int col_tiles = (N - kk - ib + COL_TILE - 1) / COL_TILE;
dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
qrN_compute_y<N, NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
);
int row_tiles = (N - kk + ROW_TILE - 1) / ROW_TILE;
dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qrN_apply_block<N, NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
h.data_ptr<float>(), y_ptr, kk, ib
);
}
}
check_cuda(cudaGetLastError(), "qrN_blocked_cuda_dispatch");
}
void qr2048_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
qrN_blocked_cuda_dispatch<2048>(a, h, tau, workspace);
}
void qr4096_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
qrN_blocked_cuda_dispatch<4096>(a, h, tau, workspace);
}
__global__ void copy_input_kernel_352(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau,
int batch,
long long total
) {
long long stride = (long long)blockDim.x * gridDim.x;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += stride) {
h[idx] = a[idx];
}
for (int b = blockIdx.x * blockDim.x + threadIdx.x;
b < batch; b += blockDim.x * gridDim.x) {
tau[(long long)b * 352 + 351] = 0.0f;
}
}
template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr352_panel_build_t_fused(
float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
constexpr int NB = 8;
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + 352;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * 352 * 352;
float* __restrict__ tau_b = tau + (long long)b * 352;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < 351; ++local_k) {
int k = kk + local_k;
int m = 352 - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * 352 + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * 352 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f;
inv_s = 0.0f;
active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * 352 + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * 352 + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * 352 + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * 352 + j] : 0.0f;
for (int i = k + 1 + tid; i < 352; i += BLOCK_SIZE) {
dot_part = fmaf(h_b[i * 352 + k], h_b[i * 352 + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * 352 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 352; i += BLOCK_SIZE) {
h_b[i * 352 + j] -= tau_s * h_b[i * 352 + k] * dot_s;
}
__syncthreads();
}
}
__syncthreads();
}
if (kk + ib >= 352) return;
float* tmp = smem;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot_part = (tid == 0) ? h_b[(kk + j) * 352 + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < 352 - kk; r += BLOCK_SIZE) {
dot_part = fmaf(
h_b[(kk + r) * 352 + (kk + i)],
h_b[(kk + r) * 352 + (kk + j)],
dot_part
);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[0] = dot_part;
}
__syncthreads();
if (tid == 0) {
tmp[i] = -tau_j * warp_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}}
template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr352_compute_y(
const float* __restrict__ h,
const float* __restrict__ workspace,
float* __restrict__ y_workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float partial_s[NB * COL_TILE * RSPLIT];
__shared__ float w_s[NB * COL_TILE];
int b = blockIdx.x;
int col_tile = blockIdx.y;
int idx = threadIdx.x;
int split = idx % RSPLIT;
int c = (idx / RSPLIT) % COL_TILE;
int q = idx / (RSPLIT * COL_TILE);
int trail_col = col_tile * COL_TILE + c;
int global_col = kk + ib + trail_col;
int pair = q * COL_TILE + c;
const float* __restrict__ h_b = h + (long long)b * 352 * 352;
const float* __restrict__ t = workspace + (long long)b * workspace_stride;
float* __restrict__ y = y_workspace + (long long)b * NB * 352;
if (q < ib && c < COL_TILE && global_col < 352) {
float sum = (split == 0) ? h_b[(kk + q) * 352 + global_col] : 0.0f;
for (int r = q + 1 + split; r < 352 - kk; r += RSPLIT) {
sum = fmaf(
h_b[(kk + r) * 352 + (kk + q)],
h_b[(kk + r) * 352 + global_col],
sum
);
}
partial_s[pair * RSPLIT + split] = sum;
} else if (q < NB && c < COL_TILE) {
partial_s[pair * RSPLIT + split] = 0.0f;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 352) {
float w = 0.0f;
#pragma unroll
for (int s = 0; s < RSPLIT; ++s) {
w += partial_s[pair * RSPLIT + s];
}
w_s[pair] = w;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 352) {
float sum = 0.0f;
for (int l = 0; l <= q; ++l) {
sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
}
y[q * 352 + global_col] = sum;
}
}
template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr352_apply_block(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int idx = threadIdx.x;
int r_local = row_tile * ROW_TILE + idx / COL_TILE;
int c_local = col_tile * COL_TILE + idx % COL_TILE;
int global_col = kk + ib + c_local;
if (r_local >= 352 - kk || global_col >= 352) return;
float* __restrict__ h_b = h + (long long)b * 352 * 352;
const float* __restrict__ y = y_workspace + (long long)b * NB * 352;
float correction = 0.0f;
int nq = (r_local < ib) ? r_local : ib;
#pragma unroll
for (int q = 0; q < nq; ++q) {
float vq = h_b[(kk + r_local) * 352 + (kk + q)];
correction = fmaf(vq, y[q * 352 + global_col], correction);
}
if (r_local < ib) {
correction = fmaf(1.0f, y[r_local * 352 + global_col], correction);
}
h_b[(kk + r_local) * 352 + global_col] -= correction;
}
void qr352_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
constexpr int NB = 8;
constexpr int COL_TILE = 32;
constexpr int ROW_TILE = 16;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * 352;
int batch = (int)a.size(0);
long long total = (long long)batch * 352 * 352;
int copy_blocks = (int)((total + THREADS - 1) / THREADS);
if (copy_blocks < 1) copy_blocks = 1;
if (copy_blocks > 8192) copy_blocks = 8192;
copy_input_kernel_352<<<copy_blocks, THREADS>>>(
a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
);
float* workspace_ptr = workspace.data_ptr<float>();
float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;
for (int kk = 0; kk < 351; kk += NB) {
int ib = min(NB, 352 - kk);
qr352_panel_build_t_fused<256, 8><<<batch, 256, (352 + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
);
if (kk + ib < 352) {
int col_tiles = (352 - kk - ib + COL_TILE - 1) / COL_TILE;
dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
qr352_compute_y<NB, COL_TILE, 1><<<y_grid, NB * COL_TILE * 1>>>(
h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
);
int row_tiles = (352 - kk + ROW_TILE - 1) / ROW_TILE;
dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qr352_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
h.data_ptr<float>(), y_ptr, kk, ib
);
}
}
check_cuda(cudaGetLastError(), "qr352_blocked_cuda");
}
__global__ void copy_input_kernel_176(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau,
int batch,
long long total
) {
long long stride = (long long)blockDim.x * gridDim.x;
for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += stride) {
h[idx] = a[idx];
}
for (int b = blockIdx.x * blockDim.x + threadIdx.x;
b < batch; b += blockDim.x * gridDim.x) {
tau[(long long)b * 176 + 175] = 0.0f;
}
}
template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr176_panel_build_t_fused(
float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ workspace,
int kk,
int ib,
int workspace_stride
) {
constexpr int NB = 8;
extern __shared__ float smem[];
float* v = smem;
float* warp_buf = smem + 176;
int b = blockIdx.x;
int tid = threadIdx.x;
int wid = tid / WARP_SIZE;
int lid = tid % WARP_SIZE;
float* __restrict__ h_b = h + (long long)b * 176 * 176;
float* __restrict__ tau_b = tau + (long long)b * 176;
float* __restrict__ t = workspace + (long long)b * workspace_stride;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float dot_s;
__shared__ int active_s;
for (int local_k = 0; local_k < ib && kk + local_k < 175; ++local_k) {
int k = kk + local_k;
int m = 176 - k;
int panel_end = kk + ib;
float s = 0.0f;
for (int i = tid; i < m; i += BLOCK_SIZE) {
float x = h_b[(k + i) * 176 + k];
s = fmaf(x, x, s);
}
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[wid] = s;
__syncthreads();
if (wid == 0) {
s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
s = warp_reduce_sum(s);
if (lid == 0) warp_buf[0] = s;
}
__syncthreads();
if (tid == 0) {
float norm = sqrtf(warp_buf[0]);
float diag = h_b[k * 176 + k];
if (norm <= 1e-20f) {
tau_s = 0.0f; inv_s = 0.0f; active_s = 0;
tau_b[k] = 0.0f;
} else {
float alpha = (diag >= 0.0f) ? -norm : norm;
float denom = diag - alpha;
inv_s = 1.0f / denom;
tau_s = (alpha - diag) / alpha;
active_s = 1;
h_b[k * 176 + k] = alpha;
tau_b[k] = tau_s;
}
}
__syncthreads();
if (active_s != 0) {
if (tid == 0) v[0] = 1.0f;
for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
float vi = h_b[(k + 1 + i) * 176 + k] * inv_s;
v[i + 1] = vi;
h_b[(k + 1 + i) * 176 + k] = vi;
}
__syncthreads();
for (int j = k + 1; j < panel_end; ++j) {
float dot_part = (tid == 0) ? h_b[k * 176 + j] : 0.0f;
for (int i = k + 1 + tid; i < 176; i += BLOCK_SIZE) {
dot_part = fmaf(h_b[i * 176 + k], h_b[i * 176 + j], dot_part);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) dot_s = dot_part;
}
__syncthreads();
if (tid == 0) {
h_b[k * 176 + j] -= tau_s * dot_s;
}
for (int i = k + 1 + tid; i < 176; i += BLOCK_SIZE) {
h_b[i * 176 + j] -= tau_s * h_b[i * 176 + k] * dot_s;
}
__syncthreads();
}
}
__syncthreads();
}
if (kk + ib >= 176) return;
float* tmp = smem;
for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
t[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < ib; ++j) {
float tau_j = tau_b[kk + j];
if (tau_j == 0.0f) {
if (tid == 0) t[j * NB + j] = 0.0f;
__syncthreads();
continue;
}
for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
tmp[idx] = 0.0f;
}
__syncthreads();
for (int i = 0; i < j; ++i) {
float dot_part = (tid == 0) ? h_b[(kk + j) * 176 + (kk + i)] : 0.0f;
for (int r = j + 1 + tid; r < 176 - kk; r += BLOCK_SIZE) {
dot_part = fmaf(
h_b[(kk + r) * 176 + (kk + i)],
h_b[(kk + r) * 176 + (kk + j)],
dot_part
);
}
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[wid] = dot_part;
__syncthreads();
if (wid == 0) {
dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
dot_part = warp_reduce_sum(dot_part);
if (lid == 0) warp_buf[0] = dot_part;
}
__syncthreads();
if (tid == 0) {
tmp[i] = -tau_j * warp_buf[0];
}
__syncthreads();
}
if (tid == 0) {
for (int row = 0; row < j; ++row) {
float val = 0.0f;
for (int col = row; col < j; ++col) {
val = fmaf(t[row * NB + col], tmp[col], val);
}
t[row * NB + j] = val;
}
t[j * NB + j] = tau_j;
}
__syncthreads();
}}
template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr176_compute_y(
const float* __restrict__ h,
const float* __restrict__ workspace,
float* __restrict__ y_workspace,
int kk,
int ib,
int workspace_stride
) {
__shared__ float partial_s[NB * COL_TILE * RSPLIT];
__shared__ float w_s[NB * COL_TILE];
int b = blockIdx.x;
int col_tile = blockIdx.y;
int idx = threadIdx.x;
int split = idx % RSPLIT;
int c = (idx / RSPLIT) % COL_TILE;
int q = idx / (RSPLIT * COL_TILE);
int trail_col = col_tile * COL_TILE + c;
int global_col = kk + ib + trail_col;
int pair = q * COL_TILE + c;
const float* __restrict__ h_b = h + (long long)b * 176 * 176;
const float* __restrict__ t = workspace + (long long)b * workspace_stride;
float* __restrict__ y = y_workspace + (long long)b * NB * 176;
if (q < ib && c < COL_TILE && global_col < 176) {
float sum = (split == 0) ? h_b[(kk + q) * 176 + global_col] : 0.0f;
for (int r = q + 1 + split; r < 176 - kk; r += RSPLIT) {
sum = fmaf(h_b[(kk + r) * 176 + (kk + q)], h_b[(kk + r) * 176 + global_col], sum);
}
partial_s[pair * RSPLIT + split] = sum;
} else if (q < NB && c < COL_TILE) {
partial_s[pair * RSPLIT + split] = 0.0f;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 176) {
float w = 0.0f;
#pragma unroll
for (int s = 0; s < RSPLIT; ++s) w += partial_s[pair * RSPLIT + s];
w_s[pair] = w;
}
__syncthreads();
if (split == 0 && q < ib && c < COL_TILE && global_col < 176) {
float sum = 0.0f;
for (int l = 0; l <= q; ++l) sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
y[q * 176 + global_col] = sum;
}
}
template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr176_apply_block(
float* __restrict__ h,
const float* __restrict__ y_workspace,
int kk,
int ib
) {
int b = blockIdx.x;
int row_tile = blockIdx.y;
int col_tile = blockIdx.z;
int idx = threadIdx.x;
int r_local = row_tile * ROW_TILE + idx / COL_TILE;
int c_local = col_tile * COL_TILE + idx % COL_TILE;
int global_col = kk + ib + c_local;
if (r_local >= 176 - kk || global_col >= 176) return;
float* __restrict__ h_b = h + (long long)b * 176 * 176;
const float* __restrict__ y = y_workspace + (long long)b * NB * 176;
float correction = 0.0f;
int nq = (r_local < ib) ? r_local : ib;
#pragma unroll
for (int q = 0; q < nq; ++q) {
float vq = h_b[(kk + r_local) * 176 + (kk + q)];
correction = fmaf(vq, y[q * 176 + global_col], correction);
}
if (r_local < ib) {
correction = fmaf(1.0f, y[r_local * 176 + global_col], correction);
}
h_b[(kk + r_local) * 176 + global_col] -= correction;
}
void qr176_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
constexpr int NB = 8;
constexpr int COL_TILE = 32;
constexpr int ROW_TILE = 16;
constexpr int WORKSPACE_STRIDE = NB * NB + NB * 176;
int batch = (int)a.size(0);
long long total = (long long)batch * 176 * 176;
int copy_blocks = (int)((total + THREADS - 1) / THREADS);
if (copy_blocks < 1) copy_blocks = 1;
if (copy_blocks > 8192) copy_blocks = 8192;
copy_input_kernel_176<<<copy_blocks, THREADS>>>(
a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
);
float* workspace_ptr = workspace.data_ptr<float>();
float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;
for (int kk = 0; kk < 175; kk += NB) {
int ib = min(NB, 176 - kk);
qr176_panel_build_t_fused<256, 8><<<batch, 256, (176 + 8) * (int)sizeof(float)>>>(
h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
);
if (kk + ib < 176) {
int col_tiles = (176 - kk - ib + COL_TILE - 1) / COL_TILE;
dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
qr176_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
);
int row_tiles = (176 - kk + ROW_TILE - 1) / ROW_TILE;
dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
qr176_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
h.data_ptr<float>(), y_ptr, kk, ib
);
}
}
check_cuda(cudaGetLastError(), "qr176_blocked_cuda");
}
"""
CPP_SRC = r"""
void qr32_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr512_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr512_panel_build_t_cuda(torch::Tensor h, torch::Tensor tau, torch::Tensor workspace, int kk, int batch);
void qr512_apply_block_wmma_cuda(torch::Tensor h, torch::Tensor y, int kk, int ib, int batch);
void qr1024_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr352_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr176_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr2048_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr4096_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
"""
try:
module = load_inline(
name="qr_native_v2",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=[
"qr32_cuda",
"qr512_blocked_cuda",
"qr512_panel_build_t_cuda",
"qr512_apply_block_wmma_cuda",
"qr1024_blocked_cuda",
"qr352_blocked_cuda",
"qr176_blocked_cuda",
"qr2048_blocked_cuda",
"qr4096_blocked_cuda",
],
verbose=False,
extra_cuda_cflags=["-O2"],
)
_has_cuda = True
except Exception:
module = None
_has_cuda = False
_cute512_enabled = False
_cute512_dyn_y_executors: dict = {}
_cute512_dyn_top_executors: dict = {}
_cute512_dyn_bottom4_executors: dict = {}
_ws_cache: dict = {}
def _get_ws(shape, dtype, device):
import math
dev = torch.device(device) if not isinstance(device, torch.device) else device
key = (tuple(shape), dev.index if dev.type == "cuda" else -1)
t = _ws_cache.get(key)
needed = math.prod(int(s) for s in shape)
if t is None or t.numel() < needed:
t = torch.empty(shape, dtype=dtype, device=dev)
_ws_cache[key] = t
return t
def _try_cuda_graph_path(data: input_t) -> output_t | None:
return None
def _custom_kernel_cute512(data: input_t) -> output_t:
batch, n, _ = data.shape
nb = 8
ws = nb * nb + nb * n
h = torch.empty_like(data)
tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)
workspace = _get_ws((batch, 2 * ws), torch.float32, data.device)
w_tmp = _get_ws((batch, nb * n), torch.float32, data.device).view(batch, nb, n)
y_data = workspace.view(-1)[batch * ws:batch * ws + batch * nb * n].view(batch, nb, n)
h.copy_(data)
tau[:, n - 1] = 0.0
for kk in range(0, n - 1, nb):
ib = min(nb, n - kk)
module.qr512_panel_build_t_cuda(h, tau, workspace, kk, batch)
if kk + ib >= n:
break
_cute512_dynamic_compute_y(h, workspace, y_data, w_tmp, kk)
_cute512_dynamic_apply_top(h, y_data, kk)
if n - kk - ib > 0:
_cute512_dynamic_apply_bottom4(h, y_data, kk)
return h, tau
def _custom_kernel_wmma512(data: input_t) -> output_t:
batch, n, _ = data.shape
nb = 8
ws = nb * nb + nb * n
h = torch.empty_like(data)
tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)
workspace = _get_ws((batch, 2 * ws), torch.float32, data.device)
w_tmp = _get_ws((batch, nb * n), torch.float32, data.device).view(batch, nb, n)
y_data = workspace.view(-1)[batch * ws:batch * ws + batch * nb * n].view(batch, nb, n)
h.copy_(data)
tau[:, n - 1] = 0.0
for kk in range(0, n - 1, nb):
ib = min(nb, n - kk)
module.qr512_panel_build_t_cuda(h, tau, workspace, kk, batch)
if kk + ib >= n:
break
_cute512_dynamic_compute_y(h, workspace, y_data, w_tmp, kk)
_cute512_dynamic_apply_top(h, y_data, kk)
if n - kk - ib > 0:
module.qr512_apply_block_wmma_cuda(h, y_data, kk, ib, batch)
return h, tau
def custom_kernel(data: input_t) -> output_t:
if not _has_cuda:
return torch.geqrf(data)
batch, n, _ = data.shape
h = torch.empty_like(data)
tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)
if n == 32:
module.qr32_cuda(data, h, tau)
return h, tau
if n == 512:
workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
module.qr512_blocked_cuda(data, h, tau, workspace)
return h, tau
if n == 1024:
workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
module.qr1024_blocked_cuda(data, h, tau, workspace)
return h, tau
if n == 352:
workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
module.qr352_blocked_cuda(data, h, tau, workspace)
return h, tau
if n == 176:
workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
module.qr176_blocked_cuda(data, h, tau, workspace)
return h, tau
if n in (2048, 4096):
return torch.geqrf(data)
return torch.geqrf(data)
scrolls · 2262 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