submission 808705
trxonphoenix · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1523 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-808705?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:7be935bf0b5e65cec3d01deebec0cb01b05ee7ae19a3f94fcadd24cb32d7aea9
license declaredunknown
license concludedunknown
authorstrxonphoenix
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float a[MAX_N * MAX_N];Kernel source
submission.py1523 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = """
void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int custom_small);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cmath>
#include <stdexcept>
#include <string>
namespace {
constexpr int MAX_N = 32;
constexpr int PANEL_THREADS = 512;
inline void check_cuda(cudaError_t status, const char* where) {
if (status != cudaSuccess) {
throw std::runtime_error(std::string(where) + ": " + cudaGetErrorString(status));
}
}
inline void check_blas(cublasStatus_t status, const char* where) {
if (status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(std::string(where) + ": cuBLAS status " + std::to_string(status));
}
}
inline void check_solver(cusolverStatus_t status, const char* where) {
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error(std::string(where) + ": cuSOLVER status " + std::to_string(status));
}
}
cublasHandle_t get_blas_handle() {
static cublasHandle_t handle = nullptr;
static bool initialized = false;
if (!initialized) {
check_blas(cublasCreate(&handle), "cublasCreate");
#ifdef CUBLAS_TF32_TENSOR_OP_MATH
check_blas(cublasSetMathMode(handle, CUBLAS_TF32_TENSOR_OP_MATH), "cublasSetMathMode");
#endif
initialized = true;
}
return handle;
}
template <int QR_THREADS>
__global__ void qr32_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int batch,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
__shared__ float a[MAX_N * MAX_N];
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
const float* src = a_in + static_cast<long long>(b) * n * n;
float* dst = h_out + static_cast<long long>(b) * n * n;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
for (int idx = tid; idx < n * n; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[i * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const float alpha = a[k * n + k];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[k * n + k] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[i * n + k] *= scale_s;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[k * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[i * n + k], a[i * n + j], dot);
}
const float w = tau * dot;
a[k * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
a[i * n + j] = fmaf(-a[i * n + k], w, a[i * n + j]);
}
}
__syncthreads();
}
for (int idx = tid; idx < n * n; idx += blockDim.x) {
dst[idx] = a[idx];
}
}
template <int QR_THREADS>
__global__ void qr_shared_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src = a_in + static_cast<long long>(b) * matrix_size;
float* dst = h_out + static_cast<long long>(b) * matrix_size;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
extern __shared__ float a[];
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[static_cast<long long>(i) * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(k) * n + k;
const float alpha = a[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[diag] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[static_cast<long long>(i) * n + k] *= scale;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[static_cast<long long>(k) * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[static_cast<long long>(i) * n + k],
a[static_cast<long long>(i) * n + j],
dot);
}
const float w = tau * dot;
a[static_cast<long long>(k) * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
const long long idx = static_cast<long long>(i) * n + j;
a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
}
}
__syncthreads();
}
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
dst[idx] = a[idx];
}
}
template <int QR_THREADS>
__global__ void qr_global_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src = a_in + static_cast<long long>(b) * matrix_size;
float* a = h_out + static_cast<long long>(b) * matrix_size;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[static_cast<long long>(i) * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(k) * n + k;
const float alpha = a[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[diag] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[static_cast<long long>(i) * n + k] *= scale_s;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[static_cast<long long>(k) * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[static_cast<long long>(i) * n + k],
a[static_cast<long long>(i) * n + j],
dot);
}
const float w = tau * dot;
a[static_cast<long long>(k) * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
const long long idx = static_cast<long long>(i) * n + j;
a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
}
}
__syncthreads();
}
}
__global__ void row_to_col_major_kernel(const float* __restrict__ row,
float* __restrict__ col,
int n,
long long total) {
const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx >= total) {
return;
}
const long long matrix_size = static_cast<long long>(n) * n;
const long long b = idx / matrix_size;
const long long rem = idx - b * matrix_size;
const int i = static_cast<int>(rem / n);
const int j = static_cast<int>(rem - static_cast<long long>(i) * n);
col[b * matrix_size + i + static_cast<long long>(j) * n] =
row[b * matrix_size + static_cast<long long>(i) * n + j];
}
__global__ void col_to_row_major_kernel(const float* __restrict__ col,
float* __restrict__ row,
int n,
long long total) {
const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx >= total) {
return;
}
const long long matrix_size = static_cast<long long>(n) * n;
const long long b = idx / matrix_size;
const long long rem = idx - b * matrix_size;
const int i = static_cast<int>(rem / n);
const int j = static_cast<int>(rem - static_cast<long long>(i) * n);
row[b * matrix_size + static_cast<long long>(i) * n + j] =
col[b * matrix_size + i + static_cast<long long>(j) * n];
}
__global__ void transpose_tiled_kernel(const float* __restrict__ src,
float* __restrict__ dst,
int n) {
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
__shared__ float tile[TILE_DIM][TILE_DIM + 1];
const int b = blockIdx.z;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src_b = src + static_cast<long long>(b) * matrix_size;
float* dst_b = dst + static_cast<long long>(b) * matrix_size;
int x = blockIdx.x * TILE_DIM + threadIdx.x;
int y = blockIdx.y * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
tile[threadIdx.y + j][threadIdx.x] =
src_b[static_cast<long long>(y + j) * n + x];
}
}
__syncthreads();
x = blockIdx.y * TILE_DIM + threadIdx.x;
y = blockIdx.x * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
dst_b[static_cast<long long>(y + j) * n + x] =
tile[threadIdx.x][threadIdx.y + j];
}
}
}
__global__ void col_to_row_major_restore_v_kernel(const float* __restrict__ col,
float* __restrict__ row,
const float* __restrict__ saved,
int n,
int nb,
long long total) {
const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx >= total) {
return;
}
const long long matrix_size = static_cast<long long>(n) * n;
const long long b = idx / matrix_size;
const long long rem = idx - b * matrix_size;
const int i = static_cast<int>(rem / n);
const int j = static_cast<int>(rem - static_cast<long long>(i) * n);
float value = col[b * matrix_size + i + static_cast<long long>(j) * n];
const int panel_k = (j / nb) * nb;
const int r = i - panel_k;
const int s = j - panel_k;
if (r >= 0 && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb;
value = saved[b * save_stride + static_cast<long long>(panel_k) * nb +
r + static_cast<long long>(s) * nb];
}
row[b * matrix_size + static_cast<long long>(i) * n + j] = value;
}
__global__ void transpose_tiled_restore_v_kernel(const float* __restrict__ col,
float* __restrict__ row,
const float* __restrict__ saved,
int n,
int nb) {
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
__shared__ float tile[TILE_DIM][TILE_DIM + 1];
const int b = blockIdx.z;
const long long matrix_size = static_cast<long long>(n) * n;
const float* col_b = col + static_cast<long long>(b) * matrix_size;
float* row_b = row + static_cast<long long>(b) * matrix_size;
int x = blockIdx.x * TILE_DIM + threadIdx.x;
int y = blockIdx.y * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
tile[threadIdx.y + j][threadIdx.x] =
col_b[static_cast<long long>(y + j) * n + x];
}
}
__syncthreads();
x = blockIdx.y * TILE_DIM + threadIdx.x;
y = blockIdx.x * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
const int row_idx = y + j;
const int col_idx = x;
if (row_idx < n && col_idx < n) {
float value = tile[threadIdx.x][threadIdx.y + j];
const int panel_k = (col_idx / nb) * nb;
const int r = row_idx - panel_k;
const int s = col_idx - panel_k;
if (r >= 0 && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb;
value = saved[static_cast<long long>(b) * save_stride +
static_cast<long long>(panel_k) * nb +
r + static_cast<long long>(s) * nb];
}
row_b[static_cast<long long>(row_idx) * n + col_idx] = value;
}
}
}
__global__ void build_explicit_v_kernel(const float* __restrict__ col,
float* __restrict__ v,
int n,
int k,
int m,
int ib,
long long stride_v,
long long total) {
const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx >= total) {
return;
}
const long long per_batch = static_cast<long long>(m) * ib;
const long long b = idx / per_batch;
const long long rem = idx - b * per_batch;
const int r = static_cast<int>(rem % m);
const int s = static_cast<int>(rem / m);
const long long matrix_size = static_cast<long long>(n) * n;
const float* mat = col + b * matrix_size;
float value = 0.0f;
if (r == s) {
value = 1.0f;
} else if (r > s) {
value = mat[static_cast<long long>(k + r) + static_cast<long long>(k + s) * n];
}
v[b * stride_v + r + static_cast<long long>(s) * m] = value;
}
__global__ void prepare_panel_v_inplace_kernel(float* __restrict__ col,
float* __restrict__ saved,
int n,
int k,
int ib,
long long stride_small) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
float* sb = saved + static_cast<long long>(b) * stride_small;
for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
const int r = idx % ib;
const int s = idx / ib;
const long long mat_idx = static_cast<long long>(k + r) +
static_cast<long long>(k + s) * n;
const float original = mat[mat_idx];
if (r <= s) {
sb[idx] = original;
}
if (r < s) {
mat[mat_idx] = 0.0f;
} else if (r == s) {
mat[mat_idx] = 1.0f;
}
}
}
__global__ void restore_panel_r_kernel(float* __restrict__ col,
const float* __restrict__ saved,
int n,
int k,
int ib,
long long stride_small) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
const float* sb = saved + static_cast<long long>(b) * stride_small;
for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
const int r = idx % ib;
const int s = idx / ib;
const long long mat_idx = static_cast<long long>(k + r) +
static_cast<long long>(k + s) * n;
if (r <= s) {
mat[mat_idx] = sb[idx];
}
}
}
__global__ void build_t_from_gram_kernel(const float* __restrict__ gram,
float* __restrict__ t,
const float* __restrict__ tau,
int n,
int k,
int ib,
long long stride_small) {
const int b = blockIdx.x;
const float* g = gram + static_cast<long long>(b) * stride_small;
float* tb = t + static_cast<long long>(b) * stride_small;
const float* tau_b = tau + static_cast<long long>(b) * n + k;
for (int idx = 0; idx < ib * ib; ++idx) {
tb[idx] = 0.0f;
}
float work[64];
for (int i = 0; i < ib; ++i) {
const float tau_i = tau_b[i];
if (tau_i == 0.0f) {
tb[i + static_cast<long long>(i) * ib] = 0.0f;
continue;
}
for (int j = 0; j < i; ++j) {
work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
}
for (int row = 0; row < i; ++row) {
float sum = 0.0f;
for (int col = row; col < i; ++col) {
sum = fmaf(tb[row + static_cast<long long>(col) * ib], work[col], sum);
}
tb[row + static_cast<long long>(i) * ib] = sum;
}
tb[i + static_cast<long long>(i) * ib] = tau_i;
}
}
__global__ void build_gram_small_kernel(const float* __restrict__ v,
float* __restrict__ gram,
int m,
int ib,
long long stride_v,
long long stride_small) {
const int b = blockIdx.x;
const int row = blockIdx.y;
const int col = blockIdx.z;
const int tid = threadIdx.x;
__shared__ float red[PANEL_THREADS];
const float* vb = v + static_cast<long long>(b) * stride_v;
float local = 0.0f;
for (int r = tid; r < m; r += blockDim.x) {
local = fmaf(vb[r + static_cast<long long>(row) * m],
vb[r + static_cast<long long>(col) * m],
local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
gram[static_cast<long long>(b) * stride_small + row + static_cast<long long>(col) * ib] = red[0];
}
}
__global__ void build_gram_strided_small_kernel(const float* __restrict__ v,
float* __restrict__ gram,
int m,
int ib,
int ldv,
long long stride_v,
long long stride_small) {
const int b = blockIdx.x;
const int row = blockIdx.y;
const int col = blockIdx.z;
const int tid = threadIdx.x;
__shared__ float red[PANEL_THREADS];
const float* vb = v + static_cast<long long>(b) * stride_v;
float local = 0.0f;
for (int r = tid; r < m; r += blockDim.x) {
local = fmaf(vb[r + static_cast<long long>(row) * ldv],
vb[r + static_cast<long long>(col) * ldv],
local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
gram[static_cast<long long>(b) * stride_small + row + static_cast<long long>(col) * ib] = red[0];
}
}
__global__ void panel_factor_colmajor_kernel(float* __restrict__ col,
float* __restrict__ tau,
int n,
int k0,
int ib) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
float* tau_b = tau + static_cast<long long>(b) * n;
__shared__ float red[PANEL_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
__shared__ float w_s;
for (int s = 0; s < ib; ++s) {
const int k = k0 + s;
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = mat[i + static_cast<long long>(k) * n];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(k) + static_cast<long long>(k) * n;
const float alpha = mat[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_b[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau_val = (beta - alpha) / beta;
tau_s = tau_val;
scale_s = 1.0f / (alpha - beta);
mat[diag] = beta;
tau_b[k] = tau_val;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
mat[i + static_cast<long long>(k) * n] *= scale;
}
}
__syncthreads();
for (int jj = s + 1; jj < ib; ++jj) {
const int j = k0 + jj;
float dot_local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
dot_local = fmaf(mat[i + static_cast<long long>(k) * n],
mat[i + static_cast<long long>(j) * n],
dot_local);
}
red[tid] = dot_local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const float dot = mat[static_cast<long long>(k) + static_cast<long long>(j) * n] + red[0];
const float w = tau_s * dot;
w_s = w;
mat[static_cast<long long>(k) + static_cast<long long>(j) * n] -= w;
}
__syncthreads();
const float w = w_s;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const long long idx = static_cast<long long>(i) + static_cast<long long>(j) * n;
mat[idx] = fmaf(-mat[i + static_cast<long long>(k) * n], w, mat[idx]);
}
__syncthreads();
}
}
}
__global__ void panel_factor_colmajor_shared_kernel(float* __restrict__ col,
float* __restrict__ tau,
float* __restrict__ saved,
int n,
int k0,
int ib,
int nb_storage,
int make_explicit_v) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int m = n - k0;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
float* tau_b = tau + static_cast<long long>(b) * n;
extern __shared__ float panel[];
__shared__ float red[PANEL_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
__shared__ float w_s;
const int panel_elems = m * ib;
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
panel[idx] = mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n];
}
__syncthreads();
for (int s = 0; s < ib; ++s) {
float local = 0.0f;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
const float x = panel[r + static_cast<long long>(s) * m];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(s) + static_cast<long long>(s) * m;
const float alpha = panel[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_b[k0 + s] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau_val = (beta - alpha) / beta;
tau_s = tau_val;
scale_s = 1.0f / (alpha - beta);
panel[diag] = beta;
tau_b[k0 + s] = tau_val;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
panel[r + static_cast<long long>(s) * m] *= scale;
}
}
__syncthreads();
for (int jj = s + 1; jj < ib; ++jj) {
float dot_local = 0.0f;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
dot_local = fmaf(panel[r + static_cast<long long>(s) * m],
panel[r + static_cast<long long>(jj) * m],
dot_local);
}
red[tid] = dot_local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const float dot = panel[s + static_cast<long long>(jj) * m] + red[0];
const float w = tau_s * dot;
w_s = w;
panel[s + static_cast<long long>(jj) * m] -= w;
}
__syncthreads();
const float w = w_s;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
const long long idx = static_cast<long long>(r) + static_cast<long long>(jj) * m;
panel[idx] = fmaf(-panel[r + static_cast<long long>(s) * m], w, panel[idx]);
}
__syncthreads();
}
}
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
float value = panel[idx];
if (make_explicit_v && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb_storage;
saved[static_cast<long long>(b) * save_stride +
static_cast<long long>(k0) * nb_storage +
r + static_cast<long long>(s) * nb_storage] = value;
value = (r == s) ? 1.0f : 0.0f;
}
mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n] = value;
}
}
__global__ void panel_factor_colmajor_shared_warpcols_kernel(float* __restrict__ col,
float* __restrict__ tau,
float* __restrict__ saved,
int n,
int k0,
int ib,
int nb_storage,
int make_explicit_v) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int warp_count = blockDim.x >> 5;
const int m = n - k0;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
float* tau_b = tau + static_cast<long long>(b) * n;
extern __shared__ float panel[];
__shared__ float red[PANEL_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
const int panel_elems = m * ib;
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
panel[idx] = mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n];
}
__syncthreads();
for (int s = 0; s < ib; ++s) {
float local = 0.0f;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
const float x = panel[r + static_cast<long long>(s) * m];
local = fmaf(x, x, local);
}
unsigned mask = 0xffffffffu;
for (int offset = 16; offset > 0; offset >>= 1) {
local += __shfl_down_sync(mask, local, offset);
}
if (lane == 0) {
red[warp] = local;
}
__syncthreads();
if (warp == 0) {
float warp_sum = (lane < warp_count) ? red[lane] : 0.0f;
for (int offset = 16; offset > 0; offset >>= 1) {
warp_sum += __shfl_down_sync(mask, warp_sum, offset);
}
if (lane == 0) {
red[0] = warp_sum;
}
}
__syncthreads();
if (tid == 0) {
const long long diag = static_cast<long long>(s) + static_cast<long long>(s) * m;
const float alpha = panel[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_b[k0 + s] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau_val = (beta - alpha) / beta;
tau_s = tau_val;
scale_s = 1.0f / (alpha - beta);
panel[diag] = beta;
tau_b[k0 + s] = tau_val;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
panel[r + static_cast<long long>(s) * m] *= scale;
}
}
__syncthreads();
const int jj = s + 1 + warp;
if (jj < ib) {
float dot = 0.0f;
for (int r = s + 1 + lane; r < m; r += 32) {
dot = fmaf(panel[r + static_cast<long long>(s) * m],
panel[r + static_cast<long long>(jj) * m],
dot);
}
for (int offset = 16; offset > 0; offset >>= 1) {
dot += __shfl_down_sync(mask, dot, offset);
}
float w = 0.0f;
if (lane == 0) {
const float full_dot = panel[s + static_cast<long long>(jj) * m] + dot;
w = tau_s * full_dot;
panel[s + static_cast<long long>(jj) * m] -= w;
}
w = __shfl_sync(mask, w, 0);
for (int r = s + 1 + lane; r < m; r += 32) {
const long long idx = static_cast<long long>(r) + static_cast<long long>(jj) * m;
panel[idx] = fmaf(-panel[r + static_cast<long long>(s) * m], w, panel[idx]);
}
}
__syncthreads();
}
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
float value = panel[idx];
if (make_explicit_v && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb_storage;
saved[static_cast<long long>(b) * save_stride +
static_cast<long long>(k0) * nb_storage +
r + static_cast<long long>(s) * nb_storage] = value;
value = (r == s) ? 1.0f : 0.0f;
}
mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n] = value;
}
}
__global__ void apply_t_transpose_small_kernel(const float* __restrict__ t,
const float* __restrict__ w,
float* __restrict__ w2,
int ib,
int trailing,
long long stride_small,
long long stride_w,
long long total) {
const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx >= total) {
return;
}
const long long per_batch = static_cast<long long>(ib) * trailing;
const long long b = idx / per_batch;
const long long rem = idx - b * per_batch;
const int p = static_cast<int>(rem % ib);
const int j = static_cast<int>(rem / ib);
const float* tb = t + b * stride_small;
const float* wb = w + b * stride_w;
float* w2b = w2 + b * stride_w;
float sum = 0.0f;
for (int q = 0; q < ib; ++q) {
sum = fmaf(tb[q + static_cast<long long>(p) * ib],
wb[q + static_cast<long long>(j) * ib],
sum);
}
w2b[p + static_cast<long long>(j) * ib] = sum;
}
__global__ void build_t_apply_transpose_fused_kernel(const float* __restrict__ gram,
const float* __restrict__ tau,
const float* __restrict__ w,
float* __restrict__ w2,
int n,
int k,
int ib,
int trailing,
long long stride_small,
long long stride_w) {
constexpr int COL_TILE = 64;
const int b = blockIdx.x;
const int tile_start = blockIdx.y * COL_TILE;
const int tid = threadIdx.x;
const int cols_left = trailing - tile_start;
const int cols = (cols_left < COL_TILE) ? cols_left : COL_TILE;
if (cols <= 0) {
return;
}
__shared__ float ts[64 * 64];
if (tid < ib * ib) {
ts[tid] = 0.0f;
}
__syncthreads();
if (tid == 0) {
const float* g = gram + static_cast<long long>(b) * stride_small;
const float* tau_b = tau + static_cast<long long>(b) * n + k;
float work[64];
for (int i = 0; i < ib; ++i) {
const float tau_i = tau_b[i];
if (tau_i == 0.0f) {
ts[i + static_cast<long long>(i) * ib] = 0.0f;
continue;
}
for (int j = 0; j < i; ++j) {
work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
}
for (int row = 0; row < i; ++row) {
float sum = 0.0f;
for (int col = row; col < i; ++col) {
sum = fmaf(ts[row + static_cast<long long>(col) * ib], work[col], sum);
}
ts[row + static_cast<long long>(i) * ib] = sum;
}
ts[i + static_cast<long long>(i) * ib] = tau_i;
}
}
__syncthreads();
const float* wb = w + static_cast<long long>(b) * stride_w;
float* w2b = w2 + static_cast<long long>(b) * stride_w;
const int tile_outputs = ib * cols;
for (int idx = tid; idx < tile_outputs; idx += blockDim.x) {
const int p = idx % ib;
const int j = tile_start + idx / ib;
float sum = 0.0f;
for (int q = 0; q < ib; ++q) {
sum = fmaf(ts[q + static_cast<long long>(p) * ib],
wb[q + static_cast<long long>(j) * ib],
sum);
}
w2b[p + static_cast<long long>(j) * ib] = sum;
}
}
} // namespace
void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr32_kernel<128><<<batch, 128>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
batch,
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
const int shmem = n * n * static_cast<int>(sizeof(float));
check_cuda(
cudaFuncSetAttribute(
qr_shared_kernel<256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shmem),
"cudaFuncSetAttribute qr_shared_kernel");
qr_shared_kernel<256><<<batch, 256, shmem>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr_global_kernel<256><<<batch, 256>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr_global_kernel<512><<<batch, 512>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int custom_small) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
const long long matrix_size = static_cast<long long>(n) * n;
const long long total = static_cast<long long>(batch) * matrix_size;
if (nb <= 0 || nb > 64) {
throw std::runtime_error("unsupported blocked-WY panel width");
}
auto col = torch::empty_like(a);
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
const dim3 transpose_block(TILE_DIM, BLOCK_ROWS);
const dim3 transpose_grid((n + TILE_DIM - 1) / TILE_DIM,
(n + TILE_DIM - 1) / TILE_DIM,
batch);
transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
a.data_ptr<float>(),
col.data_ptr<float>(),
n);
check_cuda(cudaGetLastError(), "transpose_tiled_kernel row_to_col");
cublasHandle_t blas = get_blas_handle();
const int max_panel_shmem = n * nb * static_cast<int>(sizeof(float));
check_cuda(
cudaFuncSetAttribute(
panel_factor_colmajor_shared_warpcols_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shmem),
"cudaFuncSetAttribute panel_factor_colmajor_shared_warpcols_kernel");
auto opts = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
const bool delayed_inplace_v_mode = (custom_small == 8 || custom_small == 9 || custom_small == 10);
const bool inplace_v_mode = (custom_small == 6 || custom_small == 7 || delayed_inplace_v_mode);
torch::Tensor v;
if (!inplace_v_mode) {
v = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
}
auto gram = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
auto t = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
auto w = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
auto w2 = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
torch::Tensor saved;
if (delayed_inplace_v_mode) {
saved = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
} else if (inplace_v_mode) {
saved = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
}
float* col_ptr = col.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
const float one = 1.0f;
const float zero = 0.0f;
const float minus_one = -1.0f;
for (int k = 0; k < n; k += nb) {
const int m = n - k;
const int ib = (m < nb) ? m : nb;
const int trailing = n - k - ib;
const bool delayed_inplace_v = (custom_small == 8 || custom_small == 9 || custom_small == 10);
const bool inplace_v = (custom_small == 6 || custom_small == 7 || delayed_inplace_v);
const int panel_shmem = (n - k) * ib * static_cast<int>(sizeof(float));
panel_factor_colmajor_shared_warpcols_kernel<<<batch, PANEL_THREADS, panel_shmem>>>(
col_ptr,
tau_ptr,
delayed_inplace_v ? saved.data_ptr<float>() : nullptr,
n,
k,
ib,
nb,
delayed_inplace_v ? 1 : 0);
check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel");
const long long stride_v = static_cast<long long>(m) * ib;
const long long stride_small = static_cast<long long>(ib) * ib;
float* v_ptr = nullptr;
int ldv = m;
long long stride_v_use = stride_v;
if (delayed_inplace_v) {
v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
ldv = n;
stride_v_use = matrix_size;
} else if (inplace_v) {
prepare_panel_v_inplace_kernel<<<batch, 256>>>(
col_ptr,
saved.data_ptr<float>(),
n,
k,
ib,
stride_small);
check_cuda(cudaGetLastError(), "prepare_panel_v_inplace_kernel");
v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
ldv = n;
stride_v_use = matrix_size;
} else {
v_ptr = v.data_ptr<float>();
const long long v_total = static_cast<long long>(batch) * stride_v;
const int v_blocks = static_cast<int>((v_total + 255) / 256);
build_explicit_v_kernel<<<v_blocks, 256>>>(
col_ptr,
v.data_ptr<float>(),
n,
k,
m,
ib,
stride_v,
v_total);
check_cuda(cudaGetLastError(), "build_explicit_v_kernel");
}
if (custom_small == 1 || custom_small == 2) {
const dim3 gram_grid(batch, ib, ib);
build_gram_small_kernel<<<gram_grid, PANEL_THREADS>>>(
v.data_ptr<float>(),
gram.data_ptr<float>(),
m,
ib,
stride_v,
stride_small);
check_cuda(cudaGetLastError(), "build_gram_small_kernel");
} else if (custom_small == 10) {
const dim3 gram_grid(batch, ib, ib);
build_gram_strided_small_kernel<<<gram_grid, PANEL_THREADS>>>(
v_ptr,
gram.data_ptr<float>(),
m,
ib,
ldv,
stride_v_use,
stride_small);
check_cuda(cudaGetLastError(), "build_gram_strided_small_kernel");
} else {
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
ib,
m,
&one,
v_ptr,
ldv,
stride_v_use,
v_ptr,
ldv,
stride_v_use,
&zero,
gram.data_ptr<float>(),
ib,
stride_small,
batch),
"cublasSgemmStridedBatched gram");
}
const bool fused_t_apply = (custom_small == 2 || custom_small == 5 || custom_small == 9 || custom_small == 10);
if (!fused_t_apply) {
build_t_from_gram_kernel<<<batch, 1>>>(
gram.data_ptr<float>(),
t.data_ptr<float>(),
tau_ptr,
n,
k,
ib,
stride_small);
check_cuda(cudaGetLastError(), "build_t_from_gram_kernel");
}
if (trailing > 0) {
const long long stride_w = static_cast<long long>(ib) * trailing;
float* c_panel = col_ptr + static_cast<long long>(k) +
static_cast<long long>(k + ib) * n;
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
trailing,
m,
&one,
v_ptr,
ldv,
stride_v_use,
c_panel,
n,
matrix_size,
&zero,
w.data_ptr<float>(),
ib,
stride_w,
batch),
"cublasSgemmStridedBatched vt_c");
if (fused_t_apply) {
constexpr int fused_cols = 64;
const dim3 fused_grid(batch, (trailing + fused_cols - 1) / fused_cols);
build_t_apply_transpose_fused_kernel<<<fused_grid, 256>>>(
gram.data_ptr<float>(),
tau_ptr,
w.data_ptr<float>(),
w2.data_ptr<float>(),
n,
k,
ib,
trailing,
stride_small,
stride_w);
check_cuda(cudaGetLastError(), "build_t_apply_transpose_fused_kernel");
} else if (custom_small == 1) {
const long long tw_total = static_cast<long long>(batch) * stride_w;
const int tw_blocks = static_cast<int>((tw_total + 255) / 256);
apply_t_transpose_small_kernel<<<tw_blocks, 256>>>(
t.data_ptr<float>(),
w.data_ptr<float>(),
w2.data_ptr<float>(),
ib,
trailing,
stride_small,
stride_w,
tw_total);
check_cuda(cudaGetLastError(), "apply_t_transpose_small_kernel");
} else {
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
trailing,
ib,
&one,
t.data_ptr<float>(),
ib,
stride_small,
w.data_ptr<float>(),
ib,
stride_w,
&zero,
w2.data_ptr<float>(),
ib,
stride_w,
batch),
"cublasSgemmStridedBatched t_w");
}
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_N,
CUBLAS_OP_N,
m,
trailing,
ib,
&minus_one,
v_ptr,
ldv,
stride_v_use,
w2.data_ptr<float>(),
ib,
stride_w,
&one,
c_panel,
n,
matrix_size,
batch),
"cublasSgemmStridedBatched update");
}
if (inplace_v && !delayed_inplace_v) {
restore_panel_r_kernel<<<batch, 256>>>(
col_ptr,
saved.data_ptr<float>(),
n,
k,
ib,
stride_small);
check_cuda(cudaGetLastError(), "restore_panel_r_kernel");
}
}
if (delayed_inplace_v_mode) {
transpose_tiled_restore_v_kernel<<<transpose_grid, transpose_block>>>(
col.data_ptr<float>(),
h.data_ptr<float>(),
saved.data_ptr<float>(),
n,
nb);
check_cuda(cudaGetLastError(), "transpose_tiled_restore_v_kernel");
} else {
transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
col.data_ptr<float>(),
h.data_ptr<float>(),
n);
check_cuda(cudaGetLastError(), "transpose_tiled_kernel col_to_row");
}
}
"""
_module = load_inline(
name="qr_compact_householder_b200_v99_n352_nb16_retry",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["qr32", "qr_shared_176", "qr_global_256", "qr_global_512", "qr_blocked_wy"],
verbose=False,
extra_cuda_cflags=["-O3", "--use_fast_math"],
extra_ldflags=["-lcublas"],
)
def custom_kernel(data: input_t) -> output_t:
if data.is_cuda and data.dtype == torch.float32 and data.is_contiguous():
batch, n, m = data.shape
if m == n and n <= 32:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr32(data, h, tau)
return h, tau
if m == n and n == 176:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_shared_176(data, h, tau)
return h, tau
if m == n and n <= 176:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_global_256(data, h, tau)
return h, tau
if m == n and n <= 512 and not (n == 352 or (n == 512 and batch >= 128)):
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_global_512(data, h, tau)
return h, tau
if m == n and n == 512 and batch >= 128:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_blocked_wy(data, h, tau, 16, 8)
return h, tau
if m == n and n == 352:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_blocked_wy(data, h, tau, 16, 9)
return h, tau
if m == n and n == 1024:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_blocked_wy(data, h, tau, 16, 9)
return h, tau
if m == n and n == 2048:
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
_module.qr_blocked_wy(data, h, tau, 16, 10)
return h, tau
return torch.geqrf(data)
scrolls · 1523 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