submission 833276
ptxv · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3814 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833276?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:c4ca7060bfabce7b27eb0dd0efc9de7d2d68b8504af41dc91dd6bc37dccdc125
license declaredunknown
license concludedunknown
authorsptxv
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ unsigned int warp_maxima[2 * Warps];vector-width = float4
__device__ __forceinline__ float4 ptx_ld_global_v4_f32(const float* ptr) {Kernel source
submission.py3814 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
import tempfile
from pathlib import Path
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
_CPP_SRC = r"""
#include <torch/extension.h>
std::vector<torch::Tensor> qr_small_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_small_prefix_cuda(torch::Tensor a, int64_t factor_cols);
std::vector<torch::Tensor> qr_cholqr_hr512_cuda(torch::Tensor a, torch::Tensor r, torch::Tensor info);
int64_t detect_tiny_suffix_512_cuda(torch::Tensor a);
int64_t detect_upper_512_cuda(torch::Tensor a);
int64_t detect_upper_1024_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_2048_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_4096_cuda(torch::Tensor a);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("qr_small", &qr_small_cuda, "batched Householder QR");
m.def("qr_small_prefix", &qr_small_prefix_cuda, "batched prefix Householder QR");
m.def("qr_cholqr_hr512", &qr_cholqr_hr512_cuda, "n512 CholeskyQR direct Householder reconstruction");
m.def("detect_tiny_suffix_512", &detect_tiny_suffix_512_cuda, "certified n512 tiny-suffix detector");
m.def("detect_upper_512", &detect_upper_512_cuda, "certified n512 upper-triangular detector");
m.def("detect_upper_1024", &detect_upper_1024_cuda, "certified n1024 upper-triangular detector");
m.def("qr_2048", &qr_2048_cuda, "batched n=2048 Householder QR");
m.def("qr_4096", &qr_4096_cuda, "batched n=4096 Householder QR");
}
"""
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cublasLt.h>
#include <cublas_v2.h>
#include <cuda_runtime.h>
namespace {
constexpr int kN = 32;
constexpr int kLD = 33;
constexpr int kThreads = 32;
constexpr int kN176 = 176;
constexpr int kLD176Resident = 177;
constexpr int kPanel176 = 8;
constexpr int kPanelThreads176 = 256;
constexpr int kUpdateThreads176 = 128;
constexpr int kN352 = 352;
constexpr int kN512 = 512;
constexpr int kN1024 = 1024;
constexpr int kN2048 = 2048;
constexpr int kN4096 = 4096;
constexpr int kPanel352 = 8;
constexpr int kPanelThreads352 = 256;
constexpr int kUpdateThreads352 = 128;
constexpr int kTileUpdate352 = 32;
constexpr int kPanel512 = 16;
constexpr int kPanelThreads512 = 128;
constexpr int kPanel1024 = 16;
constexpr int kPanelThreads1024 = 512;
constexpr int kPanel2048 = 16;
constexpr int kPanelThreads2048 = 512;
constexpr int kPanel4096 = 8;
constexpr int kPanelThreads4096 = 256;
constexpr int kPanel4096Late = 16;
constexpr int kPanelThreads4096Late = 512;
// B200 / compute capability 10.0 permits a 16-column panel up to 3584 active
// rows within the large dynamic shared-memory limit.
constexpr int kPanel4096LateMaxRows = 3584;
inline void cublas_leaf0_gram_from_macro(
float* gram,
long long gram_stride0,
const float* v_macro,
long long v_stride0,
int macro_cols,
int active_rows,
int batch,
bool allow_tf32) {
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
const float alpha = 1.0f;
const float beta = 0.0f;
const cublasComputeType_t compute_type = allow_tf32
? CUBLAS_COMPUTE_32F_FAST_TF32
: CUBLAS_COMPUTE_32F;
const cublasGemmAlgo_t algo = allow_tf32
? CUBLAS_GEMM_DEFAULT_TENSOR_OP
: CUBLAS_GEMM_DEFAULT;
cublasStatus_t status = cublasGemmStridedBatchedEx(
handle,
CUBLAS_OP_N,
CUBLAS_OP_T,
kPanel512,
kPanel512,
active_rows,
&alpha,
v_macro,
CUDA_R_32F,
macro_cols,
v_stride0,
v_macro,
CUDA_R_32F,
macro_cols,
v_stride0,
&beta,
gram,
CUDA_R_32F,
kPanel512,
gram_stride0,
batch,
compute_type,
algo);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "leaf0 Gram cuBLAS call failed");
}
inline void cublas_leaf1_cross_gram_from_macro(
float* s_macro,
long long s_stride0,
const float* v_macro,
long long v_stride0,
int macro_cols,
int active_rows,
int batch,
bool allow_tf32) {
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
const float alpha = 1.0f;
const float beta = 0.0f;
const cublasComputeType_t compute_type = allow_tf32
? CUBLAS_COMPUTE_32F_FAST_TF32
: CUBLAS_COMPUTE_32F;
const cublasGemmAlgo_t algo = allow_tf32
? CUBLAS_GEMM_DEFAULT_TENSOR_OP
: CUBLAS_GEMM_DEFAULT;
cublasStatus_t status = cublasGemmStridedBatchedEx(
handle,
CUBLAS_OP_N,
CUBLAS_OP_T,
kPanel512,
macro_cols,
active_rows,
&alpha,
v_macro + kPanel512,
CUDA_R_32F,
macro_cols,
v_stride0,
v_macro,
CUDA_R_32F,
macro_cols,
v_stride0,
&beta,
s_macro,
CUDA_R_32F,
kPanel512,
s_stride0,
batch,
compute_type,
algo);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "leaf1 cross Gram cuBLAS call failed");
}
inline void cublas_w_from_vt_c(
float* w,
const float* c_tail,
const float* v,
long long v_stride0,
int n,
int panel_cols,
int v_ld_cols,
int active_rows,
int trailing_cols,
int batch,
bool allow_tf32) {
cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
const float alpha = 1.0f;
const float beta = 0.0f;
const cublasComputeType_t compute_type = allow_tf32
? CUBLAS_COMPUTE_32F_FAST_TF32
: CUBLAS_COMPUTE_32F;
const cublasGemmAlgo_t algo = allow_tf32
? CUBLAS_GEMM_DEFAULT_TENSOR_OP
: CUBLAS_GEMM_DEFAULT;
cublasStatus_t status = cublasGemmStridedBatchedEx(
handle,
CUBLAS_OP_N,
CUBLAS_OP_T,
trailing_cols,
panel_cols,
active_rows,
&alpha,
c_tail,
CUDA_R_32F,
n,
static_cast<long long>(n) * n,
v,
CUDA_R_32F,
v_ld_cols,
v_stride0,
&beta,
w,
CUDA_R_32F,
trailing_cols,
static_cast<long long>(panel_cols) * trailing_cols,
batch,
compute_type,
algo);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "W=V^T C cuBLAS call failed");
}
inline void cublaslt_tail_update_out_of_place(
float* d_tail,
const float* c_tail,
const float* v_macro,
long long v_stride0,
const float* z,
int n,
int macro_cols,
int active_rows,
int trailing_cols,
int batch,
bool allow_tf32,
void* workspace,
size_t workspace_bytes) {
cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
cublasLtMatmulDesc_t op_desc = nullptr;
cublasLtMatrixLayout_t a_desc = nullptr;
cublasLtMatrixLayout_t b_desc = nullptr;
cublasLtMatrixLayout_t c_desc = nullptr;
cublasLtMatrixLayout_t d_desc = nullptr;
cublasLtMatmulPreference_t pref = nullptr;
const cublasComputeType_t compute_type = allow_tf32
? CUBLAS_COMPUTE_32F_FAST_TF32
: CUBLAS_COMPUTE_32F;
TORCH_CHECK(
cublasLtMatmulDescCreate(&op_desc, compute_type, CUDA_R_32F) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt desc create failed");
cublasOperation_t trans = CUBLAS_OP_N;
TORCH_CHECK(
cublasLtMatmulDescSetAttribute(
op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &trans, sizeof(trans)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt transa set failed");
TORCH_CHECK(
cublasLtMatmulDescSetAttribute(
op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &trans, sizeof(trans)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt transb set failed");
TORCH_CHECK(
cublasLtMatrixLayoutCreate(&a_desc, CUDA_R_32F, trailing_cols, macro_cols, trailing_cols) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt A layout create failed");
TORCH_CHECK(
cublasLtMatrixLayoutCreate(&b_desc, CUDA_R_32F, macro_cols, active_rows, macro_cols) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt B layout create failed");
TORCH_CHECK(
cublasLtMatrixLayoutCreate(&c_desc, CUDA_R_32F, trailing_cols, active_rows, n) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt C layout create failed");
TORCH_CHECK(
cublasLtMatrixLayoutCreate(&d_desc, CUDA_R_32F, trailing_cols, active_rows, n) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt D layout create failed");
const int batch_count = batch;
const int64_t a_stride = static_cast<int64_t>(macro_cols) * trailing_cols;
const int64_t b_stride = static_cast<int64_t>(v_stride0);
const int64_t cd_stride = static_cast<int64_t>(n) * n;
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
a_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt A batch set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
b_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt B batch set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
c_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt C batch set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
d_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt D batch set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
a_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &a_stride, sizeof(a_stride)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt A stride set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
b_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &b_stride, sizeof(b_stride)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt B stride set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
c_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &cd_stride, sizeof(cd_stride)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt C stride set failed");
TORCH_CHECK(
cublasLtMatrixLayoutSetAttribute(
d_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &cd_stride, sizeof(cd_stride)) ==
CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt D stride set failed");
TORCH_CHECK(
cublasLtMatmulPreferenceCreate(&pref) == CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt preference create failed");
TORCH_CHECK(
cublasLtMatmulPreferenceSetAttribute(
pref,
CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
&workspace_bytes,
sizeof(workspace_bytes)) == CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt workspace set failed");
cublasLtMatmulHeuristicResult_t heuristic;
int returned = 0;
TORCH_CHECK(
cublasLtMatmulAlgoGetHeuristic(
handle,
op_desc,
a_desc,
b_desc,
c_desc,
d_desc,
pref,
1,
&heuristic,
&returned) == CUBLAS_STATUS_SUCCESS,
"tail update cuBLASLt heuristic query failed");
TORCH_CHECK(returned > 0, "tail update cuBLASLt returned no heuristic");
const float alpha = -1.0f;
const float beta = 1.0f;
const cublasStatus_t status = cublasLtMatmul(
handle,
op_desc,
&alpha,
z,
a_desc,
v_macro,
b_desc,
&beta,
c_tail,
c_desc,
d_tail,
d_desc,
&heuristic.algo,
workspace,
workspace_bytes,
nullptr);
cublasLtMatmulPreferenceDestroy(pref);
cublasLtMatrixLayoutDestroy(d_desc);
cublasLtMatrixLayoutDestroy(c_desc);
cublasLtMatrixLayoutDestroy(b_desc);
cublasLtMatrixLayoutDestroy(a_desc);
cublasLtMatmulDescDestroy(op_desc);
TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "tail update cuBLASLt matmul failed");
}
__device__ __forceinline__ float warp_sum(float value) {
value += __shfl_down_sync(0xffffffff, value, 16);
value += __shfl_down_sync(0xffffffff, value, 8);
value += __shfl_down_sync(0xffffffff, value, 4);
value += __shfl_down_sync(0xffffffff, value, 2);
value += __shfl_down_sync(0xffffffff, value, 1);
return value;
}
__device__ __forceinline__ float4 ptx_ld_global_v4_f32(const float* ptr) {
float4 value;
asm volatile(
"ld.global.v4.f32 {%0, %1, %2, %3}, [%4];"
: "=f"(value.x), "=f"(value.y), "=f"(value.z), "=f"(value.w)
: "l"(ptr));
return value;
}
__device__ __forceinline__ void ptx_st_global_v4_f32(float* ptr, float4 value) {
asm volatile(
"st.global.v4.f32 [%0], {%1, %2, %3, %4};"
:
: "l"(ptr), "f"(value.x), "f"(value.y), "f"(value.z), "f"(value.w));
}
template <int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_matrix_kernel(
const float* __restrict__ a,
float* __restrict__ h,
long long total) {
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < total;
idx += static_cast<long long>(gridDim.x) * Threads) {
h[idx] = a[idx];
}
}
template <int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_matrix_v4_kernel(
const float* __restrict__ a,
float* __restrict__ h,
long long total) {
const long long vec_total = total >> 2;
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < vec_total;
idx += static_cast<long long>(gridDim.x) * Threads) {
const long long base = idx << 2;
const float4 vals = ptx_ld_global_v4_f32(a + base);
ptx_st_global_v4_f32(h + base, vals);
}
}
template <int N, int Cols, int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_first_cols_v4_kernel(
const float* __restrict__ a,
float* __restrict__ h,
int batch) {
constexpr int VecCols = Cols / 4;
const long long total = static_cast<long long>(batch) * N * VecCols;
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < total;
idx += static_cast<long long>(gridDim.x) * Threads) {
const int vec_col = static_cast<int>(idx % VecCols);
const long long row_tmp = idx / VecCols;
const int row = static_cast<int>(row_tmp % N);
const int b = static_cast<int>(row_tmp / N);
const long long base =
(static_cast<long long>(b) * N + row) * N + vec_col * 4;
const float4 vals = ptx_ld_global_v4_f32(a + base);
ptx_st_global_v4_f32(h + base, vals);
}
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void zero_suffix_cols_v4_kernel(
float* __restrict__ h,
int start_col,
int batch) {
const int vec_cols = (N - start_col) / 4;
const long long total = static_cast<long long>(batch) * N * vec_cols;
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < total;
idx += static_cast<long long>(gridDim.x) * Threads) {
const int vec_col = static_cast<int>(idx % vec_cols);
const long long row_tmp = idx / vec_cols;
const int row = static_cast<int>(row_tmp % N);
const int b = static_cast<int>(row_tmp / N);
float* dst =
h + (static_cast<long long>(b) * N + row) * N + start_col + vec_col * 4;
ptx_st_global_v4_f32(dst, make_float4(0.0f, 0.0f, 0.0f, 0.0f));
}
}
__device__ __forceinline__ unsigned int warp_max_u32(unsigned int value) {
value = max(value, __shfl_down_sync(0xffffffff, value, 16));
value = max(value, __shfl_down_sync(0xffffffff, value, 8));
value = max(value, __shfl_down_sync(0xffffffff, value, 4));
value = max(value, __shfl_down_sync(0xffffffff, value, 2));
value = max(value, __shfl_down_sync(0xffffffff, value, 1));
return value;
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void suffix_sample_reject_kernel(
const float* __restrict__ a,
int* __restrict__ reject) {
constexpr int Warps = Threads / 32;
__shared__ unsigned int warp_maxima[2 * Warps];
const int b = blockIdx.x;
const float* a_b = a + static_cast<long long>(b) * N * N;
const int tid = threadIdx.x;
// This stage can only reject. It never authorizes a shortcut.
const int row0 = (tid * 73) & (N - 1);
const int col0 = tid & 15;
const int row1 = (tid * 37) & (N - 1);
const int col1 = (3 * N) / 4 + ((tid * 29) & (N / 4 - 1));
unsigned int prefix = __float_as_uint(fabsf(a_b[row0 * N + col0]));
unsigned int tail = __float_as_uint(fabsf(a_b[row1 * N + col1]));
prefix = warp_max_u32(prefix);
tail = warp_max_u32(tail);
const int lane = tid & 31;
const int warp = tid >> 5;
if (lane == 0) {
warp_maxima[warp] = prefix;
warp_maxima[Warps + warp] = tail;
}
__syncthreads();
if (warp == 0) {
unsigned int p = (lane < Warps) ? warp_maxima[lane] : 0;
unsigned int t = (lane < Warps) ? warp_maxima[Warps + lane] : 0;
p = warp_max_u32(p);
t = warp_max_u32(t);
if (lane == 0) {
const float pv = __uint_as_float(p);
const float tv = __uint_as_float(t);
if (pv > 0.0f && tv > 1.0e-3f * pv) {
atomicExch(reject, 1);
}
}
}
}
template <int N, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void suffix_factor_cols_kernel(
const float* __restrict__ a,
int* __restrict__ factors,
int k0,
int k1,
int k2) {
constexpr int Warps = Threads / 32;
constexpr int RowVecs = N / 4;
constexpr int MatrixVecs = N * RowVecs;
__shared__ unsigned int warp_maxima[4 * Warps];
const int b = blockIdx.x;
const float* a_b = a + static_cast<long long>(b) * N * N;
const int tid = threadIdx.x;
unsigned int local0 = 0;
unsigned int local1 = 0;
unsigned int local2 = 0;
unsigned int local3 = 0;
for (int vec = tid; vec < MatrixVecs; vec += Threads) {
const int col4 = (vec % RowVecs) * 4;
const float4 vals = ptx_ld_global_v4_f32(a_b + static_cast<long long>(vec) * 4);
unsigned int v = __float_as_uint(fabsf(vals.x));
v = max(v, __float_as_uint(fabsf(vals.y)));
v = max(v, __float_as_uint(fabsf(vals.z)));
v = max(v, __float_as_uint(fabsf(vals.w)));
local0 = max(local0, v);
if (col4 >= k0) local1 = max(local1, v);
if (col4 >= k1) local2 = max(local2, v);
if (col4 >= k2) local3 = max(local3, v);
}
local0 = warp_max_u32(local0);
local1 = warp_max_u32(local1);
local2 = warp_max_u32(local2);
local3 = warp_max_u32(local3);
const int lane = tid & 31;
const int warp = tid >> 5;
if (lane == 0) {
warp_maxima[0 * Warps + warp] = local0;
warp_maxima[1 * Warps + warp] = local1;
warp_maxima[2 * Warps + warp] = local2;
warp_maxima[3 * Warps + warp] = local3;
}
__syncthreads();
if (warp == 0) {
unsigned int block0 = (lane < Warps) ? warp_maxima[0 * Warps + lane] : 0;
unsigned int block1 = (lane < Warps) ? warp_maxima[1 * Warps + lane] : 0;
unsigned int block2 = (lane < Warps) ? warp_maxima[2 * Warps + lane] : 0;
unsigned int block3 = (lane < Warps) ? warp_maxima[3 * Warps + lane] : 0;
block0 = warp_max_u32(block0);
block1 = warp_max_u32(block1);
block2 = warp_max_u32(block2);
block3 = warp_max_u32(block3);
if (lane == 0) {
constexpr float eps32 = 1.1920928955078125e-7f;
constexpr float route_budget = 6.0f;
const float all = __uint_as_float(block0);
const float tail0 = __uint_as_float(block1);
const float tail1 = __uint_as_float(block2);
const float tail2 = __uint_as_float(block3);
int matrix_cols = 0;
if (all == 0.0f) {
matrix_cols = Panel;
} else {
const float limit = route_budget * eps32 * all;
if (tail0 <= limit) matrix_cols = k0;
else if (tail1 <= limit) matrix_cols = k1;
else if (tail2 <= limit) matrix_cols = k2;
}
factors[b] = matrix_cols;
}
}
}
template <int Threads>
__global__ __launch_bounds__(Threads, 1) void reduce_factor_cols_kernel(
const int* __restrict__ factors,
int* __restrict__ result,
int batch) {
constexpr int Warps = Threads / 32;
__shared__ int warp_values[2 * Warps];
const int tid = threadIdx.x;
int local_max = 0;
int local_bad = 0;
for (int idx = tid; idx < batch; idx += Threads) {
const int value = factors[idx];
if (value == 0) {
local_bad = 1;
} else {
local_max = max(local_max, value);
}
}
local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 16));
local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 8));
local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 4));
local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 2));
local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 1));
local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 16));
local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 8));
local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 4));
local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 2));
local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 1));
const int lane = tid & 31;
const int warp = tid >> 5;
if (lane == 0) {
warp_values[warp] = local_max;
warp_values[Warps + warp] = local_bad;
}
__syncthreads();
if (warp == 0) {
int block_max = (lane < Warps) ? warp_values[lane] : 0;
int block_bad = (lane < Warps) ? warp_values[Warps + lane] : 0;
block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 16));
block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 8));
block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 4));
block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 2));
block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 1));
block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 16));
block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 8));
block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 4));
block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 2));
block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 1));
if (lane == 0) {
result[0] = (block_bad != 0) ? 0 : block_max;
}
}
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_prefix_zero_suffix_v4_kernel(
const float* __restrict__ a,
float* __restrict__ h,
long long total,
int factor_cols) {
const long long vec_total = total >> 2;
constexpr int RowVecs = N / 4;
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < vec_total;
idx += static_cast<long long>(gridDim.x) * Threads) {
const int col4 = static_cast<int>(idx % RowVecs) * 4;
const long long base = idx << 2;
const float4 vals = (col4 < factor_cols)
? ptx_ld_global_v4_f32(a + base)
: make_float4(0.0f, 0.0f, 0.0f, 0.0f);
ptx_st_global_v4_f32(h + base, vals);
}
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void zero_tau_suffix_kernel(
float* __restrict__ tau,
int factor_cols,
int batch) {
const int suffix = N - factor_cols;
const long long total = static_cast<long long>(batch) * suffix;
for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
idx < total;
idx += static_cast<long long>(gridDim.x) * Threads) {
const int b = static_cast<int>(idx / suffix);
const int j = factor_cols + static_cast<int>(idx - static_cast<long long>(b) * suffix);
tau[static_cast<long long>(b) * N + j] = 0.0f;
}
}
template <int Threads>
__device__ __forceinline__ float block_sum_thread0(float value, float* work) {
constexpr int Warps = Threads / 32;
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
value = warp_sum(value);
if (lane == 0) {
work[warp] = value;
}
__syncthreads();
float total = 0.0f;
if (threadIdx.x < 32) {
total = (threadIdx.x < Warps) ? work[threadIdx.x] : 0.0f;
total = warp_sum(total);
}
return total;
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 1) void cholqr_hr_lu_kernel(
const float* __restrict__ r,
const int* __restrict__ info,
float* __restrict__ h,
float* __restrict__ tau,
int* __restrict__ ok) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
constexpr long long MatrixElems = static_cast<long long>(N) * N;
const float* r_b = r + static_cast<long long>(b) * MatrixElems;
float* h_b = h + static_cast<long long>(b) * MatrixElems;
float* tau_b = tau + static_cast<long long>(b) * N;
__shared__ float reduce[32];
__shared__ float row_sign;
__shared__ float pivot;
__shared__ float diag_scale;
__shared__ int bad;
if (tid == 0) {
bad = (info[b] != 0);
float min_diag = 3.4028234663852886e38f;
float max_diag = 0.0f;
if (!bad) {
for (int j = 0; j < N; ++j) {
const float d = fabsf(r_b[j * N + j]);
if (!(isfinite(d) && d > 0.0f)) {
bad = 1;
break;
}
min_diag = fminf(min_diag, d);
max_diag = fmaxf(max_diag, d);
}
}
diag_scale = fmaxf(max_diag, 1.0e-30f);
if (!bad && min_diag < 1.0e-7f * diag_scale) {
bad = 1;
}
}
__syncthreads();
if (bad) {
for (int j = tid; j < N; j += Threads) {
tau_b[j] = 0.0f;
}
if (tid == 0) ok[b] = 0;
return;
}
for (int k = 0; k < N; ++k) {
if (tid == 0) {
const float x = h_b[k * N + k];
if (isfinite(x)) {
row_sign = (x >= 0.0f) ? -1.0f : 1.0f;
} else {
bad = 1;
row_sign = -1.0f;
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
for (int j = k + tid; j < N; j += Threads) {
const float shifted = h_b[k * N + j] - row_sign * r_b[k * N + j];
h_b[k * N + j] = shifted;
if (!isfinite(shifted)) {
atomicExch(&bad, 1);
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
if (tid == 0) {
pivot = h_b[k * N + k];
const float pivot_floor = fmaxf(1.0e-20f, 1.0e-7f * diag_scale);
if (!(isfinite(pivot) && fabsf(pivot) > pivot_floor)) {
bad = 1;
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
const float inv_pivot = 1.0f / pivot;
float local_norm = 0.0f;
for (int i = k + 1 + tid; i < N; i += Threads) {
const float v = h_b[i * N + k] * inv_pivot;
h_b[i * N + k] = v;
local_norm = fmaf(v, v, local_norm);
if (!isfinite(v)) {
atomicExch(&bad, 1);
}
}
const float tail_norm = block_sum_thread0<Threads>(local_norm, reduce);
if (tid == 0) {
const float tau_k = 2.0f / (1.0f + tail_norm);
tau_b[k] = tau_k;
if (!(isfinite(tau_k) && tau_k >= 0.0f && tau_k <= 2.0f)) {
bad = 1;
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
const int trailing = N - k - 1;
const int update_total = trailing * trailing;
for (int idx = tid; idx < update_total; idx += Threads) {
const int rel_i = idx / trailing;
const int rel_j = idx - rel_i * trailing;
const int i = k + 1 + rel_i;
const int j = k + 1 + rel_j;
const float updated = fmaf(-h_b[i * N + k], h_b[k * N + j], h_b[i * N + j]);
h_b[i * N + j] = updated;
if (!isfinite(updated)) {
atomicExch(&bad, 1);
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
for (int j = k + tid; j < N; j += Threads) {
const float out = row_sign * r_b[k * N + j];
h_b[k * N + j] = out;
if (!isfinite(out)) {
atomicExch(&bad, 1);
}
}
__syncthreads();
if (bad) {
if (tid == 0) ok[b] = 0;
return;
}
}
if (tid == 0) ok[b] = 1;
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void upper_max_classifier_kernel(
const float* __restrict__ a,
unsigned int* __restrict__ maxima) {
constexpr int Warps = Threads / 32;
__shared__ unsigned int warp_maxima[2 * Warps];
const int b = blockIdx.y;
constexpr int MatrixElems = N * N;
const float* a_b = a + static_cast<long long>(b) * MatrixElems;
unsigned int local_all = 0;
unsigned int local_lower = 0;
for (int idx = blockIdx.x * Threads + threadIdx.x;
idx < MatrixElems;
idx += gridDim.x * Threads) {
const unsigned int bits = __float_as_uint(fabsf(a_b[idx]));
local_all = max(local_all, bits);
const int row = idx / N;
const int col = idx - row * N;
if (row > col) {
local_lower = max(local_lower, bits);
}
}
local_all = warp_max_u32(local_all);
local_lower = warp_max_u32(local_lower);
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
if (lane == 0) {
warp_maxima[warp] = local_all;
warp_maxima[Warps + warp] = local_lower;
}
__syncthreads();
if (warp == 0) {
unsigned int block_all = (lane < Warps) ? warp_maxima[lane] : 0;
unsigned int block_lower = (lane < Warps) ? warp_maxima[Warps + lane] : 0;
block_all = warp_max_u32(block_all);
block_lower = warp_max_u32(block_lower);
if (lane == 0) {
unsigned int* out = maxima + static_cast<long long>(b) * 2;
atomicMax(out + 0, block_all);
atomicMax(out + 1, block_lower);
}
}
}
template <int N>
bool all_matrices_upper_certified(torch::Tensor a) {
constexpr int threads = 256;
constexpr int blocks_per_matrix = (N >= 4096) ? 512 : 128;
constexpr float eps32 = 1.1920928955078125e-7f;
constexpr float route_budget = 2.0f;
const int batch = static_cast<int>(a.size(0));
auto maxima = torch::empty({a.size(0), 2}, a.options().dtype(at::kInt));
C10_CUDA_CHECK(cudaMemset(
maxima.data_ptr<int>(), 0, static_cast<size_t>(batch) * 2 * sizeof(int)));
upper_max_classifier_kernel<N, threads>
<<<dim3(blocks_per_matrix, batch), threads, 0>>>(
a.data_ptr<float>(),
reinterpret_cast<unsigned int*>(maxima.data_ptr<int>()));
std::vector<unsigned int> host(static_cast<size_t>(batch) * 2);
C10_CUDA_CHECK(cudaMemcpy(
host.data(),
maxima.data_ptr<int>(),
host.size() * sizeof(unsigned int),
cudaMemcpyDeviceToHost));
union BitsFloat { unsigned int u; float f; };
for (int b = 0; b < batch; ++b) {
BitsFloat all{host[static_cast<size_t>(b) * 2 + 0]};
BitsFloat lower{host[static_cast<size_t>(b) * 2 + 1]};
if (all.f != 0.0f && lower.f > route_budget * eps32 * all.f) {
return false;
}
}
return true;
}
__global__ __launch_bounds__(kThreads, 16) void qr32_kernel(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const float* a_b = a + static_cast<long long>(b) * kN * kN;
float* h_b = h + static_cast<long long>(b) * kN * kN;
float* tau_b = tau + static_cast<long long>(b) * kN;
__shared__ float s[kN * kLD];
constexpr int RowVecs = kN / 4;
for (int idx = tid; idx < kN * RowVecs; idx += kThreads) {
const int row = idx / RowVecs;
const int col = (idx - row * RowVecs) * 4;
const float4 vals = ptx_ld_global_v4_f32(a_b + row * kN + col);
s[row * kLD + col + 0] = vals.x;
s[row * kLD + col + 1] = vals.y;
s[row * kLD + col + 2] = vals.z;
s[row * kLD + col + 3] = vals.w;
}
__syncwarp();
#pragma unroll 32
for (int k = 0; k < kN; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < kN; i += kThreads) {
const float x = s[i * kLD + k];
local = fmaf(x, x, local);
}
const float sigma = warp_sum(local);
float tau_k = 0.0f;
float inv = 0.0f;
if (tid == 0) {
const float alpha = s[k * kLD + k];
if (sigma == 0.0f) {
tau_b[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, sigma));
const float beta = (alpha < 0.0f) ? norm : -norm;
tau_k = (beta - alpha) / beta;
inv = 1.0f / (alpha - beta);
tau_b[k] = tau_k;
s[k * kLD + k] = beta;
}
}
tau_k = __shfl_sync(0xffffffff, tau_k, 0);
inv = __shfl_sync(0xffffffff, inv, 0);
for (int i = k + 1 + tid; i < kN; i += kThreads) {
s[i * kLD + k] *= inv;
}
__syncwarp();
for (int j = k + 1 + tid; j < kN; j += kThreads) {
float dot = s[k * kLD + j];
#pragma unroll 4
for (int i = k + 1; i < kN; ++i) {
dot = fmaf(s[i * kLD + k], s[i * kLD + j], dot);
}
dot *= tau_k;
s[k * kLD + j] -= dot;
#pragma unroll 4
for (int i = k + 1; i < kN; ++i) {
s[i * kLD + j] = fmaf(-s[i * kLD + k], dot, s[i * kLD + j]);
}
}
__syncwarp();
}
for (int idx = tid; idx < kN * RowVecs; idx += kThreads) {
const int row = idx / RowVecs;
const int col = (idx - row * RowVecs) * 4;
ptx_st_global_v4_f32(
h_b + row * kN + col,
make_float4(
s[row * kLD + col + 0],
s[row * kLD + col + 1],
s[row * kLD + col + 2],
s[row * kLD + col + 3]));
}
}
template <int Threads>
__global__ __launch_bounds__(Threads, 1) void qr176_resident_kernel(
const float* __restrict__ a,
float* __restrict__ h,
float* __restrict__ tau) {
constexpr int Warps = Threads / 32;
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const float* a_b = a + static_cast<long long>(b) * kN176 * kN176;
float* h_b = h + static_cast<long long>(b) * kN176 * kN176;
float* tau_b = tau + static_cast<long long>(b) * kN176;
extern __shared__ float s[];
__shared__ float reduce[Warps];
__shared__ float params[2];
for (int idx = tid; idx < kN176 * kN176; idx += Threads) {
const int row = idx / kN176;
const int col = idx - row * kN176;
s[row * kLD176Resident + col] = a_b[idx];
}
__syncthreads();
for (int k = 0; k < kN176; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < kN176; i += Threads) {
const float x = s[i * kLD176Resident + k];
local = fmaf(x, x, local);
}
const float sigma = block_sum_thread0<Threads>(local, reduce);
if (tid == 0) {
const float alpha = s[k * kLD176Resident + k];
if (sigma == 0.0f) {
tau_b[k] = 0.0f;
params[0] = 0.0f;
params[1] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, sigma));
const float beta = (alpha < 0.0f) ? norm : -norm;
const float tau_k = (beta - alpha) / beta;
tau_b[k] = tau_k;
s[k * kLD176Resident + k] = beta;
params[0] = tau_k;
params[1] = 1.0f / (alpha - beta);
}
}
__syncthreads();
const float tau_k = params[0];
const float inv = params[1];
for (int i = k + 1 + tid; i < kN176; i += Threads) {
s[i * kLD176Resident + k] *= inv;
}
__syncthreads();
for (int j_base = k + 1; j_base < kN176; j_base += Warps) {
const int j = j_base + warp;
float dot = 0.0f;
if (j < kN176) {
dot = (lane == 0) ? s[k * kLD176Resident + j] : 0.0f;
for (int i = k + 1 + lane; i < kN176; i += 32) {
dot = fmaf(
s[i * kLD176Resident + k],
s[i * kLD176Resident + j],
dot);
}
}
dot = warp_sum(dot);
dot = __shfl_sync(0xffffffff, dot, 0) * tau_k;
if (j < kN176) {
if (lane == 0) {
s[k * kLD176Resident + j] -= dot;
}
for (int i = k + 1 + lane; i < kN176; i += 32) {
const int offset = i * kLD176Resident + j;
s[offset] = fmaf(-s[i * kLD176Resident + k], dot, s[offset]);
}
}
}
__syncthreads();
}
for (int idx = tid; idx < kN176 * kN176; idx += Threads) {
const int row = idx / kN176;
const int col = idx - row * kN176;
h_b[idx] = s[row * kLD176Resident + col];
}
}
template <
int N,
int Panel,
int Threads,
bool PackV,
bool DynamicStride = false,
bool BuildT = true,
bool FusedT = true,
bool WriteMacroV = false>
__global__ __launch_bounds__(Threads, 1) void qr_panel_cached_kernel(
float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ t_scratch,
float* __restrict__ v_pack,
long long v_stride0,
int panel_start,
float* __restrict__ v_macro = nullptr,
long long v_macro_stride0 = 0,
int macro_cols = 0,
int macro_row_offset = 0,
int macro_col_offset = 0) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_stride = static_cast<long long>(N) * N;
float* h_b = h + static_cast<long long>(b) * matrix_stride;
float* tau_b = tau + static_cast<long long>(b) * N;
float* t_b = nullptr;
if constexpr (BuildT) {
t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
}
float* v_b = nullptr;
if constexpr (PackV) {
v_b = v_pack + static_cast<long long>(b) * v_stride0;
}
float* v_macro_b = nullptr;
if constexpr (WriteMacroV) {
v_macro_b = v_macro + static_cast<long long>(b) * v_macro_stride0;
}
extern __shared__ float smem[];
float* panel = smem;
const int active_rows = N - panel_start;
const int panel_stride = DynamicStride ? (active_rows + 1) : (N + 1);
float* work = smem + panel_stride * Panel;
constexpr int Warps = Threads / 32;
constexpr int PanelVec = Panel / 4;
for (int idx = tid; idx < active_rows * PanelVec; idx += Threads) {
const int rel = idx / PanelVec;
const int t = (idx - rel * PanelVec) * 4;
const float4 vals = ptx_ld_global_v4_f32(
h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t);
panel[(t + 0) * panel_stride + rel] = vals.x;
panel[(t + 1) * panel_stride + rel] = vals.y;
panel[(t + 2) * panel_stride + rel] = vals.z;
panel[(t + 3) * panel_stride + rel] = vals.w;
}
__syncthreads();
if constexpr (BuildT) {
for (int idx = tid; idx < Panel * Panel; idx += Threads) {
t_b[idx] = 0.0f;
}
}
__syncthreads();
#pragma unroll
for (int col = 0; col < Panel; ++col) {
const int k = panel_start + col;
float local = 0.0f;
for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
const float x = panel[col * panel_stride + rel];
local = fmaf(x, x, local);
}
const float sigma = block_sum_thread0<Threads>(local, work);
if (tid == 0) {
const float alpha = panel[col * panel_stride + col];
if (sigma == 0.0f) {
tau_b[k] = 0.0f;
work[0] = 0.0f;
work[1] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, sigma));
const float beta = (alpha < 0.0f) ? norm : -norm;
const float tau_k = (beta - alpha) / beta;
tau_b[k] = tau_k;
panel[col * panel_stride + col] = beta;
work[0] = tau_k;
work[1] = 1.0f / (alpha - beta);
}
}
__syncthreads();
const float tau_k = work[0];
const float inv = work[1];
if constexpr (BuildT && FusedT) {
const int lane = tid & 31;
const int warp = tid >> 5;
float t_dots[Panel];
#pragma unroll
for (int prev = 0; prev < Panel; ++prev) {
t_dots[prev] = 0.0f;
}
for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
const float v_cur = panel[col * panel_stride + rel] * inv;
panel[col * panel_stride + rel] = v_cur;
#pragma unroll
for (int prev = 0; prev < Panel; ++prev) {
if (prev < col) {
t_dots[prev] = fmaf(
panel[prev * panel_stride + rel],
v_cur,
t_dots[prev]);
}
}
}
if (tid == 0) {
#pragma unroll
for (int prev = 0; prev < Panel; ++prev) {
if (prev < col) {
t_dots[prev] += panel[prev * panel_stride + col];
}
}
}
#pragma unroll
for (int prev = 0; prev < Panel; ++prev) {
const float partial = (prev < col) ? warp_sum(t_dots[prev]) : 0.0f;
if (lane == 0) {
work[warp * Panel + prev] = partial;
}
}
__syncthreads();
if (tid < Panel && tid < col) {
float dot = 0.0f;
#pragma unroll
for (int w = 0; w < Warps; ++w) {
dot += work[w * Panel + tid];
}
work[Warps * Panel + tid] = -tau_k * dot;
}
__syncthreads();
if (tid < col) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < col) {
accum = fmaf(
t_b[tid * Panel + inner],
work[Warps * Panel + inner],
accum);
}
}
t_b[tid * Panel + col] = accum;
}
if (tid == col) {
t_b[col * Panel + col] = tau_k;
}
} else {
for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
panel[col * panel_stride + rel] *= inv;
}
}
__syncthreads();
const int lane = tid & 31;
const int warp = tid >> 5;
for (int j_base = col + 1; j_base < Panel; j_base += Warps) {
const int j = j_base + warp;
float local_dot = 0.0f;
if (j < Panel) {
local_dot = (lane == 0) ? panel[j * panel_stride + col] : 0.0f;
for (int rel = col + 1 + lane; rel < active_rows; rel += 32) {
local_dot = fmaf(
panel[col * panel_stride + rel],
panel[j * panel_stride + rel],
local_dot);
}
}
float dot = warp_sum(local_dot);
dot = __shfl_sync(0xffffffff, dot, 0);
if (j < Panel) {
dot *= tau_k;
if (lane == 0) {
panel[j * panel_stride + col] -= dot;
}
for (int rel = col + 1 + lane; rel < active_rows; rel += 32) {
const int offset = j * panel_stride + rel;
panel[offset] = fmaf(
-panel[col * panel_stride + rel],
dot,
panel[offset]);
}
}
}
__syncthreads();
}
if constexpr (BuildT && !FusedT) {
#pragma unroll
for (int j = 0; j < Panel; ++j) {
const float tau_j = tau_b[panel_start + j];
#pragma unroll
for (int i = 0; i < Panel; ++i) {
if (i < j) {
float local = 0.0f;
for (int rel = j + 1 + tid; rel < active_rows; rel += Threads) {
local = fmaf(
panel[i * panel_stride + rel],
panel[j * panel_stride + rel],
local);
}
if (tid == 0) {
local += panel[i * panel_stride + j];
}
const float dot = block_sum_thread0<Threads>(local, work);
if (tid == 0) {
work[Warps + i] = -tau_j * dot;
}
__syncthreads();
}
}
if (tid < j) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < j) {
accum = fmaf(
t_b[tid * Panel + inner],
work[Warps + inner],
accum);
}
}
t_b[tid * Panel + j] = accum;
}
if (tid == j) {
t_b[j * Panel + j] = tau_j;
}
__syncthreads();
}
}
if constexpr (WriteMacroV) {
for (int idx = tid; idx < macro_row_offset * PanelVec; idx += Threads) {
const int rel = idx / PanelVec;
const int t = (idx - rel * PanelVec) * 4;
ptx_st_global_v4_f32(
v_macro_b + static_cast<long long>(rel) * macro_cols + macro_col_offset + t,
make_float4(0.0f, 0.0f, 0.0f, 0.0f));
}
}
for (int idx = tid; idx < active_rows * PanelVec; idx += Threads) {
const int rel = idx / PanelVec;
const int t = (idx - rel * PanelVec) * 4;
const float h0 = panel[(t + 0) * panel_stride + rel];
const float h1 = panel[(t + 1) * panel_stride + rel];
const float h2 = panel[(t + 2) * panel_stride + rel];
const float h3 = panel[(t + 3) * panel_stride + rel];
ptx_st_global_v4_f32(
h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t,
make_float4(h0, h1, h2, h3));
if constexpr (PackV) {
const float v0 = (rel == t + 0) ? 1.0f : ((rel > t + 0) ? h0 : 0.0f);
const float v1 = (rel == t + 1) ? 1.0f : ((rel > t + 1) ? h1 : 0.0f);
const float v2 = (rel == t + 2) ? 1.0f : ((rel > t + 2) ? h2 : 0.0f);
const float v3 = (rel == t + 3) ? 1.0f : ((rel > t + 3) ? h3 : 0.0f);
const float4 v_vals = make_float4(v0, v1, v2, v3);
ptx_st_global_v4_f32(
v_b + static_cast<long long>(rel) * Panel + t,
v_vals);
if constexpr (WriteMacroV) {
ptx_st_global_v4_f32(
v_macro_b +
static_cast<long long>(macro_row_offset + rel) * macro_cols +
macro_col_offset + t,
v_vals);
}
} else if constexpr (WriteMacroV) {
const float v0 = (rel == t + 0) ? 1.0f : ((rel > t + 0) ? h0 : 0.0f);
const float v1 = (rel == t + 1) ? 1.0f : ((rel > t + 1) ? h1 : 0.0f);
const float v2 = (rel == t + 2) ? 1.0f : ((rel > t + 2) ? h2 : 0.0f);
const float v3 = (rel == t + 3) ? 1.0f : ((rel > t + 3) ? h3 : 0.0f);
ptx_st_global_v4_f32(
v_macro_b +
static_cast<long long>(macro_row_offset + rel) * macro_cols +
macro_col_offset + t,
make_float4(v0, v1, v2, v3));
}
}
}
template <int N, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 1) void rebuild_t_from_panel_kernel(
const float* __restrict__ h,
const float* __restrict__ tau,
float* __restrict__ t_scratch,
int panel_start) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
constexpr int Warps = Threads / 32;
const int active_rows = N - panel_start;
const long long matrix_stride = static_cast<long long>(N) * N;
const float* h_b = h + static_cast<long long>(b) * matrix_stride;
const float* tau_b = tau + static_cast<long long>(b) * N;
float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
__shared__ float reduce[Warps];
__shared__ float y[Panel];
for (int idx = tid; idx < Panel * Panel; idx += Threads) {
t_b[idx] = 0.0f;
}
__syncthreads();
#pragma unroll
for (int j = 0; j < Panel; ++j) {
const float tau_j = tau_b[panel_start + j];
#pragma unroll
for (int i = 0; i < Panel; ++i) {
if (i < j) {
float local = 0.0f;
for (int rel = j + 1 + tid; rel < active_rows; rel += Threads) {
const float vi = h_b[
static_cast<long long>(panel_start + rel) * N + panel_start + i];
const float vj = h_b[
static_cast<long long>(panel_start + rel) * N + panel_start + j];
local = fmaf(vi, vj, local);
}
if (tid == 0) {
local += h_b[
static_cast<long long>(panel_start + j) * N + panel_start + i];
}
const float dot = block_sum_thread0<Threads>(local, reduce);
if (tid == 0) {
y[i] = -tau_j * dot;
}
__syncthreads();
}
}
if (tid < j) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < j) {
accum = fmaf(t_b[tid * Panel + inner], y[inner], accum);
}
}
t_b[tid * Panel + j] = accum;
}
if (tid == j) {
t_b[j * Panel + j] = tau_j;
}
__syncthreads();
}
}
template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void build_t_from_gram_kernel(
float* __restrict__ t_scratch,
const float* __restrict__ tau,
int n,
int panel_start) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
const float* tau_b = tau + static_cast<long long>(b) * n;
__shared__ float y[Panel];
#pragma unroll
for (int j = 0; j < Panel; ++j) {
const float tau_j = tau_b[panel_start + j];
if (tid < j) {
y[tid] = -tau_j * t_b[tid * Panel + j];
}
__syncthreads();
if (tid < j) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < j && tid <= inner) {
accum = fmaf(t_b[tid * Panel + inner], y[inner], accum);
}
}
t_b[tid * Panel + j] = accum;
}
if (tid == j) {
t_b[j * Panel + j] = tau_j;
}
__syncthreads();
}
}
template <int N, int Panel, int Threads, bool PackedV>
__global__ __launch_bounds__(Threads, 1) void qr_panel_wy_update_kernel(
float* __restrict__ h,
const float* __restrict__ t_scratch,
int panel_start,
const float* __restrict__ v_pack,
long long v_stride0) {
const int tile = blockIdx.x;
const int b = blockIdx.y;
const int tid = threadIdx.x;
const int panel_end = (panel_start + Panel < N) ? (panel_start + Panel) : N;
const int panel_cols = panel_end - panel_start;
const int active_rows = N - panel_start;
const int j = panel_end + tile * Threads + tid;
const long long matrix_stride = static_cast<long long>(N) * N;
float* h_b = h + static_cast<long long>(b) * matrix_stride;
const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
const float* v_b = nullptr;
if constexpr (PackedV) {
v_b = v_pack + static_cast<long long>(b) * v_stride0;
}
extern __shared__ float v_panel[];
if constexpr (!PackedV && Panel == 8) {
for (int rel_load = tid; rel_load < active_rows; rel_load += Threads) {
const int row = panel_start + rel_load;
const float4 lo = ptx_ld_global_v4_f32(
h_b + static_cast<long long>(row) * N + panel_start);
const float4 hi = ptx_ld_global_v4_f32(
h_b + static_cast<long long>(row) * N + panel_start + 4);
float* dst = v_panel + rel_load * Panel;
dst[0] = (rel_load == 0) ? 1.0f : ((rel_load > 0) ? lo.x : 0.0f);
dst[1] = (rel_load == 1) ? 1.0f : ((rel_load > 1) ? lo.y : 0.0f);
dst[2] = (rel_load == 2) ? 1.0f : ((rel_load > 2) ? lo.z : 0.0f);
dst[3] = (rel_load == 3) ? 1.0f : ((rel_load > 3) ? lo.w : 0.0f);
dst[4] = (rel_load == 4) ? 1.0f : ((rel_load > 4) ? hi.x : 0.0f);
dst[5] = (rel_load == 5) ? 1.0f : ((rel_load > 5) ? hi.y : 0.0f);
dst[6] = (rel_load == 6) ? 1.0f : ((rel_load > 6) ? hi.z : 0.0f);
dst[7] = (rel_load == 7) ? 1.0f : ((rel_load > 7) ? hi.w : 0.0f);
}
} else {
constexpr int RowThreads = Threads / Panel;
const int t_load = tid - (tid / Panel) * Panel;
for (int rel_load = tid / Panel; rel_load < active_rows; rel_load += RowThreads) {
const int row = panel_start + rel_load;
const int k = panel_start + t_load;
float value = 0.0f;
if constexpr (PackedV) {
value = v_b[static_cast<long long>(rel_load) * Panel + t_load];
} else if (t_load < panel_cols) {
if (rel_load == t_load) {
value = 1.0f;
} else if (rel_load > t_load) {
value = h_b[static_cast<long long>(row) * N + k];
}
}
v_panel[rel_load * Panel + t_load] = value;
}
}
__syncthreads();
if (j >= N) {
return;
}
float w0 = 0.0f;
float w1 = 0.0f;
float w2 = 0.0f;
float w3 = 0.0f;
float w4 = 0.0f;
float w5 = 0.0f;
float w6 = 0.0f;
float w7 = 0.0f;
long long offset = static_cast<long long>(panel_start) * N + j;
int rel = 0;
#pragma unroll 24
for (int row = panel_start; row < N; ++row, ++rel, offset += N) {
const float c = h_b[offset];
const int v_offset = rel * Panel;
w0 = fmaf(v_panel[v_offset], c, w0);
w1 = fmaf(v_panel[v_offset + 1], c, w1);
w2 = fmaf(v_panel[v_offset + 2], c, w2);
w3 = fmaf(v_panel[v_offset + 3], c, w3);
w4 = fmaf(v_panel[v_offset + 4], c, w4);
w5 = fmaf(v_panel[v_offset + 5], c, w5);
w6 = fmaf(v_panel[v_offset + 6], c, w6);
w7 = fmaf(v_panel[v_offset + 7], c, w7);
}
const float z0 = t_b[0] * w0;
const float z1 = fmaf(t_b[1], w0, t_b[Panel + 1] * w1);
const float z2 = fmaf(t_b[2], w0, fmaf(t_b[Panel + 2], w1, t_b[2 * Panel + 2] * w2));
const float z3 = fmaf(t_b[3], w0, fmaf(t_b[Panel + 3], w1, fmaf(t_b[2 * Panel + 3], w2, t_b[3 * Panel + 3] * w3)));
const float z4 = fmaf(t_b[4], w0, fmaf(t_b[Panel + 4], w1, fmaf(t_b[2 * Panel + 4], w2, fmaf(t_b[3 * Panel + 4], w3, t_b[4 * Panel + 4] * w4))));
const float z5 = fmaf(t_b[5], w0, fmaf(t_b[Panel + 5], w1, fmaf(t_b[2 * Panel + 5], w2, fmaf(t_b[3 * Panel + 5], w3, fmaf(t_b[4 * Panel + 5], w4, t_b[5 * Panel + 5] * w5)))));
const float z6 = fmaf(t_b[6], w0, fmaf(t_b[Panel + 6], w1, fmaf(t_b[2 * Panel + 6], w2, fmaf(t_b[3 * Panel + 6], w3, fmaf(t_b[4 * Panel + 6], w4, fmaf(t_b[5 * Panel + 6], w5, t_b[6 * Panel + 6] * w6))))));
const float z7 = fmaf(t_b[7], w0, fmaf(t_b[Panel + 7], w1, fmaf(t_b[2 * Panel + 7], w2, fmaf(t_b[3 * Panel + 7], w3, fmaf(t_b[4 * Panel + 7], w4, fmaf(t_b[5 * Panel + 7], w5, fmaf(t_b[6 * Panel + 7], w6, t_b[7 * Panel + 7] * w7)))))));
offset = static_cast<long long>(panel_start) * N + j;
rel = 0;
#pragma unroll 24
for (int row = panel_start; row < N; ++row, ++rel, offset += N) {
const int v_offset = rel * Panel;
float delta = v_panel[v_offset] * z0;
delta = fmaf(v_panel[v_offset + 1], z1, delta);
delta = fmaf(v_panel[v_offset + 2], z2, delta);
delta = fmaf(v_panel[v_offset + 3], z3, delta);
delta = fmaf(v_panel[v_offset + 4], z4, delta);
delta = fmaf(v_panel[v_offset + 5], z5, delta);
delta = fmaf(v_panel[v_offset + 6], z6, delta);
delta = fmaf(v_panel[v_offset + 7], z7, delta);
h_b[offset] -= delta;
}
}
template <int N, int Panel, int TileCols, int Threads>
__global__ __launch_bounds__(Threads, 1) void qr_panel_tile_update_kernel(
float* __restrict__ h,
const float* __restrict__ t_scratch,
int panel_start) {
const int tile = blockIdx.x;
const int b = blockIdx.y;
const int tid = threadIdx.x;
const int panel_end = panel_start + Panel;
const int active_rows = N - panel_start;
const int tile_start = panel_end + tile * TileCols;
if (tile_start >= N) {
return;
}
const int tile_cols = min(TileCols, N - tile_start);
const long long matrix_stride = static_cast<long long>(N) * N;
float* h_b = h + static_cast<long long>(b) * matrix_stride;
const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
extern __shared__ float smem[];
float* c_tile = smem;
float* v_panel = c_tile + active_rows * TileCols;
float* w_tile = v_panel + active_rows * Panel;
float* z_tile = w_tile + Panel * TileCols;
for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
const int rel = idx / TileCols;
const int col = idx - rel * TileCols;
float value = 0.0f;
if (col < tile_cols) {
value = h_b[
static_cast<long long>(panel_start + rel) * N +
tile_start + col];
}
c_tile[idx] = value;
}
for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
const int rel = idx / Panel;
const int t = idx - rel * Panel;
float value = 0.0f;
if (rel == t) {
value = 1.0f;
} else if (rel > t) {
value = h_b[
static_cast<long long>(panel_start + rel) * N +
panel_start + t];
}
v_panel[idx] = value;
}
__syncthreads();
if (tid < TileCols) {
const int col = tid;
float w0 = 0.0f;
float w1 = 0.0f;
float w2 = 0.0f;
float w3 = 0.0f;
float w4 = 0.0f;
float w5 = 0.0f;
float w6 = 0.0f;
float w7 = 0.0f;
if (col < tile_cols) {
#pragma unroll 24
for (int rel = 0; rel < active_rows; ++rel) {
const float c = c_tile[rel * TileCols + col];
const float* v = v_panel + rel * Panel;
w0 = fmaf(v[0], c, w0);
w1 = fmaf(v[1], c, w1);
w2 = fmaf(v[2], c, w2);
w3 = fmaf(v[3], c, w3);
w4 = fmaf(v[4], c, w4);
w5 = fmaf(v[5], c, w5);
w6 = fmaf(v[6], c, w6);
w7 = fmaf(v[7], c, w7);
}
}
w_tile[0 * TileCols + col] = w0;
w_tile[1 * TileCols + col] = w1;
w_tile[2 * TileCols + col] = w2;
w_tile[3 * TileCols + col] = w3;
w_tile[4 * TileCols + col] = w4;
w_tile[5 * TileCols + col] = w5;
w_tile[6 * TileCols + col] = w6;
w_tile[7 * TileCols + col] = w7;
z_tile[0 * TileCols + col] = t_b[0] * w0;
z_tile[1 * TileCols + col] = fmaf(t_b[1], w0, t_b[Panel + 1] * w1);
z_tile[2 * TileCols + col] =
fmaf(t_b[2], w0, fmaf(t_b[Panel + 2], w1, t_b[2 * Panel + 2] * w2));
z_tile[3 * TileCols + col] =
fmaf(t_b[3], w0, fmaf(t_b[Panel + 3], w1, fmaf(t_b[2 * Panel + 3], w2, t_b[3 * Panel + 3] * w3)));
z_tile[4 * TileCols + col] =
fmaf(t_b[4], w0, fmaf(t_b[Panel + 4], w1, fmaf(t_b[2 * Panel + 4], w2, fmaf(t_b[3 * Panel + 4], w3, t_b[4 * Panel + 4] * w4))));
z_tile[5 * TileCols + col] =
fmaf(t_b[5], w0, fmaf(t_b[Panel + 5], w1, fmaf(t_b[2 * Panel + 5], w2, fmaf(t_b[3 * Panel + 5], w3, fmaf(t_b[4 * Panel + 5], w4, t_b[5 * Panel + 5] * w5)))));
z_tile[6 * TileCols + col] =
fmaf(t_b[6], w0, fmaf(t_b[Panel + 6], w1, fmaf(t_b[2 * Panel + 6], w2, fmaf(t_b[3 * Panel + 6], w3, fmaf(t_b[4 * Panel + 6], w4, fmaf(t_b[5 * Panel + 6], w5, t_b[6 * Panel + 6] * w6))))));
z_tile[7 * TileCols + col] =
fmaf(t_b[7], w0, fmaf(t_b[Panel + 7], w1, fmaf(t_b[2 * Panel + 7], w2, fmaf(t_b[3 * Panel + 7], w3, fmaf(t_b[4 * Panel + 7], w4, fmaf(t_b[5 * Panel + 7], w5, fmaf(t_b[6 * Panel + 7], w6, t_b[7 * Panel + 7] * w7)))))));
}
__syncthreads();
for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
const int rel = idx / TileCols;
const int col = idx - rel * TileCols;
if (col < tile_cols) {
const float* v = v_panel + rel * Panel;
const float* z = z_tile + col;
float delta = v[0] * z[0 * TileCols];
delta = fmaf(v[1], z[1 * TileCols], delta);
delta = fmaf(v[2], z[2 * TileCols], delta);
delta = fmaf(v[3], z[3 * TileCols], delta);
delta = fmaf(v[4], z[4 * TileCols], delta);
delta = fmaf(v[5], z[5 * TileCols], delta);
delta = fmaf(v[6], z[6 * TileCols], delta);
delta = fmaf(v[7], z[7 * TileCols], delta);
c_tile[idx] -= delta;
}
}
__syncthreads();
for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
const int rel = idx / TileCols;
const int col = idx - rel * TileCols;
if (col < tile_cols) {
h_b[
static_cast<long long>(panel_start + rel) * N +
tile_start + col] = c_tile[idx];
}
}
}
template <int N, int Panel, int Cols, int Threads>
__global__ __launch_bounds__(Threads, 1) void apply_panel_to_next_cols_kernel(
float* __restrict__ h,
const float* __restrict__ t_scratch,
int panel_start) {
const int target_col = blockIdx.x;
const int b = blockIdx.y;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
constexpr int Warps = Threads / 32;
const int active_rows = N - panel_start;
const int j = panel_start + Panel + target_col;
const long long matrix_stride = static_cast<long long>(N) * N;
float* h_b = h + static_cast<long long>(b) * matrix_stride;
const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
extern __shared__ float smem[];
float* v_panel = smem;
float* partials = smem + Panel * active_rows;
float* z = partials + Warps * Panel;
for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
const int rel = idx / Panel;
const int t = idx - rel * Panel;
float value = 0.0f;
if (rel == t) {
value = 1.0f;
} else if (rel > t) {
value = h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t];
}
v_panel[t * active_rows + rel] = value;
}
__syncthreads();
if (j >= N || target_col >= Cols) {
return;
}
float local[Panel];
#pragma unroll
for (int t = 0; t < Panel; ++t) {
local[t] = 0.0f;
}
for (int rel = tid; rel < active_rows; rel += Threads) {
const float c = h_b[static_cast<long long>(panel_start + rel) * N + j];
#pragma unroll
for (int t = 0; t < Panel; ++t) {
local[t] = fmaf(v_panel[t * active_rows + rel], c, local[t]);
}
}
#pragma unroll
for (int t = 0; t < Panel; ++t) {
const float sum = warp_sum(local[t]);
if (lane == 0) {
partials[warp * Panel + t] = sum;
}
}
__syncthreads();
if (tid < Panel) {
float w = 0.0f;
#pragma unroll
for (int r = 0; r < Warps; ++r) {
w += partials[r * Panel + tid];
}
partials[Warps * Panel + tid] = w;
}
__syncthreads();
if (tid < Panel) {
float value = 0.0f;
#pragma unroll
for (int r = 0; r < Panel; ++r) {
value = fmaf(t_b[r * Panel + tid], partials[Warps * Panel + r], value);
}
z[tid] = value;
}
__syncthreads();
for (int rel = tid; rel < active_rows; rel += Threads) {
float delta = 0.0f;
#pragma unroll
for (int t = 0; t < Panel; ++t) {
delta = fmaf(v_panel[t * active_rows + rel], z[t], delta);
}
h_b[static_cast<long long>(panel_start + rel) * N + j] -= delta;
}
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 1) void build_t16_from_two_t8_singleblock_kernel(
const float* __restrict__ h,
const float* __restrict__ t_first,
const float* __restrict__ t_second,
float* __restrict__ t_out,
float* __restrict__ v_pack,
long long v_stride,
int panel_start) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
constexpr int Half = 8;
constexpr int Super = 16;
constexpr int Warps = Threads / 32;
const int active_rows = N - panel_start;
const long long matrix_stride = static_cast<long long>(N) * N;
const float* h_b = h + static_cast<long long>(b) * matrix_stride;
const float* t1 = t_first + static_cast<long long>(b) * Half * Half;
const float* t2 = t_second + static_cast<long long>(b) * Half * Half;
float* tout = t_out + static_cast<long long>(b) * Super * Super;
float* v_b = v_pack + static_cast<long long>(b) * v_stride;
constexpr int SuperVec = Super / 4;
for (int idx = tid; idx < active_rows * SuperVec; idx += Threads) {
const int rel = idx / SuperVec;
const int t = (idx - rel * SuperVec) * 4;
float4 values;
if (rel > t + 3) {
values = ptx_ld_global_v4_f32(
h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t);
} else {
values.x = (rel == t) ? 1.0f :
((rel > t) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t] : 0.0f);
values.y = (rel == t + 1) ? 1.0f :
((rel > t + 1) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 1] : 0.0f);
values.z = (rel == t + 2) ? 1.0f :
((rel > t + 2) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 2] : 0.0f);
values.w = (rel == t + 3) ? 1.0f :
((rel > t + 3) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 3] : 0.0f);
}
ptx_st_global_v4_f32(v_b + static_cast<long long>(rel) * Super + t, values);
}
__syncthreads();
__shared__ float s_shared[Half * Half];
__shared__ float middle[Half * Half];
for (int group = warp; group < Half * 2; group += Warps) {
const int i = group >> 1;
const int j_base = (group & 1) << 2;
float local0 = 0.0f;
float local1 = 0.0f;
float local2 = 0.0f;
float local3 = 0.0f;
for (int rel = lane; rel < active_rows; rel += 32) {
const float* v_row = v_b + static_cast<long long>(rel) * Super;
const float v1 = v_row[i];
const float4 vals = ptx_ld_global_v4_f32(v_row + Half + j_base);
const float v20 = vals.x;
const float v21 = vals.y;
const float v22 = vals.z;
const float v23 = vals.w;
local0 = fmaf(v1, v20, local0);
local1 = fmaf(v1, v21, local1);
local2 = fmaf(v1, v22, local2);
local3 = fmaf(v1, v23, local3);
}
const float total0 = warp_sum(local0);
const float total1 = warp_sum(local1);
const float total2 = warp_sum(local2);
const float total3 = warp_sum(local3);
if (lane == 0) {
s_shared[i * Half + j_base] = total0;
s_shared[i * Half + j_base + 1] = total1;
s_shared[i * Half + j_base + 2] = total2;
s_shared[i * Half + j_base + 3] = total3;
}
}
__syncthreads();
for (int idx = tid; idx < Super * Super; idx += Threads) {
tout[idx] = 0.0f;
}
__syncthreads();
for (int idx = tid; idx < Half * Half; idx += Threads) {
const int row = idx / Half;
const int col = idx - row * Half;
tout[row * Super + col] = t1[idx];
tout[(Half + row) * Super + Half + col] = t2[idx];
}
for (int idx = tid; idx < Half * Half; idx += Threads) {
const int row = idx / Half;
const int col = idx - row * Half;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < Half; ++k) {
value = fmaf(t1[row * Half + k], s_shared[k * Half + col], value);
}
middle[idx] = value;
}
__syncthreads();
for (int idx = tid; idx < Half * Half; idx += Threads) {
const int row = idx / Half;
const int col = idx - row * Half;
float value = 0.0f;
#pragma unroll
for (int k = 0; k < Half; ++k) {
value = fmaf(middle[row * Half + k], t2[k * Half + col], value);
}
tout[row * Super + Half + col] = -value;
}
}
template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void apply_t_transpose_kernel(
const float* __restrict__ t_scratch,
const float* __restrict__ w,
float* __restrict__ z,
int trailing_cols) {
const int b = blockIdx.y;
const int col = blockIdx.x * Threads + threadIdx.x;
const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols;
float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
__shared__ float t_shared[Panel * Panel];
for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
t_shared[idx] = t_b[idx];
}
__syncthreads();
if (col >= trailing_cols) {
return;
}
float wv[Panel];
#pragma unroll
for (int i = 0; i < Panel; ++i) {
wv[i] = w_b[i * trailing_cols + col];
}
#pragma unroll
for (int row = 0; row < Panel; ++row) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner <= row) {
accum = fmaf(t_shared[inner * Panel + row], wv[inner], accum);
}
}
z_b[row * trailing_cols + col] = accum;
}
}
template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void solve_inverse_wy_from_gram_kernel(
const float* __restrict__ gram_scratch,
const float* __restrict__ tau,
const float* __restrict__ w,
float* __restrict__ z,
int n,
int panel_start,
int trailing_cols) {
const int b = blockIdx.y;
const int col = blockIdx.x * Threads + threadIdx.x;
const float* g_b = gram_scratch + static_cast<long long>(b) * Panel * Panel;
const float* tau_b = tau + static_cast<long long>(b) * n + panel_start;
const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols;
float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
__shared__ float g_shared[Panel * Panel];
__shared__ float tau_shared[Panel];
for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
g_shared[idx] = g_b[idx];
}
for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
tau_shared[idx] = tau_b[idx];
}
__syncthreads();
if (col >= trailing_cols) {
return;
}
float zv[Panel];
#pragma unroll
for (int row = 0; row < Panel; ++row) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < row) {
accum = fmaf(g_shared[inner * Panel + row], zv[inner], accum);
}
}
const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum);
zv[row] = zi;
z_b[row * trailing_cols + col] = zi;
}
}
template <int Panel, int Block, int Threads>
__global__ __launch_bounds__(Threads, 2) void solve_inverse_wy_from_gram_blocked_kernel(
const float* __restrict__ gram_scratch,
const float* __restrict__ tau,
const float* __restrict__ w,
float* __restrict__ z,
int n,
int panel_start,
int trailing_cols) {
const int b = blockIdx.y;
const int col = blockIdx.x * Threads + threadIdx.x;
const float* g_b = gram_scratch + static_cast<long long>(b) * Panel * Panel;
const float* tau_b = tau + static_cast<long long>(b) * n + panel_start;
const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols;
float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
__shared__ float g_shared[Panel * Panel];
__shared__ float tau_shared[Panel];
__shared__ float z_shared[Panel * Threads];
for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
g_shared[idx] = g_b[idx];
}
for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
tau_shared[idx] = tau_b[idx];
}
__syncthreads();
#pragma unroll
for (int block_start = 0; block_start < Panel; block_start += Block) {
float zv[Block];
#pragma unroll
for (int r = 0; r < Block; ++r) {
zv[r] = 0.0f;
}
if (col < trailing_cols) {
#pragma unroll
for (int r = 0; r < Block; ++r) {
const int row = block_start + r;
float accum = 0.0f;
for (int inner = 0; inner < Panel; ++inner) {
if (inner < block_start) {
accum = fmaf(
g_shared[inner * Panel + row],
z_shared[inner * Threads + threadIdx.x],
accum);
}
}
#pragma unroll
for (int inner = 0; inner < Block; ++inner) {
if (inner < r) {
accum = fmaf(
g_shared[(block_start + inner) * Panel + row],
zv[inner],
accum);
}
}
const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum);
zv[r] = zi;
z_shared[row * Threads + threadIdx.x] = zi;
z_b[row * trailing_cols + col] = zi;
}
}
__syncthreads();
}
}
template <int Panel, int Block, int Threads>
__global__ __launch_bounds__(Threads, 2) void solve_inverse_wy_from_split_gram_blocked_kernel(
const float* __restrict__ g11_scratch,
long long g11_stride0,
const float* __restrict__ s_scratch,
const float* __restrict__ tau,
const float* __restrict__ w,
float* __restrict__ z,
int n,
int panel_start,
int trailing_cols) {
const int b = blockIdx.y;
const int col = blockIdx.x * Threads + threadIdx.x;
const float* g11_b = g11_scratch + static_cast<long long>(b) * g11_stride0;
const float* s_b = s_scratch + static_cast<long long>(b) * Panel * Block;
const float* tau_b = tau + static_cast<long long>(b) * n + panel_start;
const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols;
float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
__shared__ float g11_shared[Block * Block];
__shared__ float s_shared[Panel * Block];
__shared__ float tau_shared[Panel];
__shared__ float z_shared[Panel * Threads];
for (int idx = threadIdx.x; idx < Block * Block; idx += Threads) {
g11_shared[idx] = g11_b[idx];
}
for (int idx = threadIdx.x; idx < Panel * Block; idx += Threads) {
s_shared[idx] = s_b[idx];
}
for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
tau_shared[idx] = tau_b[idx];
}
__syncthreads();
#pragma unroll
for (int block_start = 0; block_start < Panel; block_start += Block) {
float zv[Block];
#pragma unroll
for (int r = 0; r < Block; ++r) {
zv[r] = 0.0f;
}
if (col < trailing_cols) {
#pragma unroll
for (int r = 0; r < Block; ++r) {
const int row = block_start + r;
float accum = 0.0f;
for (int inner = 0; inner < Panel; ++inner) {
if (inner < block_start) {
float gij = 0.0f;
if (row < Block) {
gij = g11_shared[inner * Block + row];
} else if (inner < Block) {
gij = s_shared[inner * Block + (row - Block)];
} else {
gij = s_shared[row * Block + (inner - Block)];
}
accum = fmaf(
gij,
z_shared[inner * Threads + threadIdx.x],
accum);
}
}
#pragma unroll
for (int inner = 0; inner < Block; ++inner) {
if (inner < r) {
const int gram_inner = block_start + inner;
float gij = 0.0f;
if (row < Block) {
gij = g11_shared[gram_inner * Block + row];
} else if (gram_inner < Block) {
gij = s_shared[gram_inner * Block + (row - Block)];
} else {
gij = s_shared[row * Block + (gram_inner - Block)];
}
accum = fmaf(gij, zv[inner], accum);
}
}
const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum);
zv[r] = zi;
z_shared[row * Threads + threadIdx.x] = zi;
z_b[row * trailing_cols + col] = zi;
}
}
__syncthreads();
}
}
template <int N, int MacroCols, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 2) void apply_prev_leaves_to_next_leaf_kernel(
float* __restrict__ h,
const float* __restrict__ v_macro,
long long v_stride0,
const float* __restrict__ leaf_grams,
const float* __restrict__ tau,
int macro_start,
int prev_leaves,
int target_offset) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int active_rows = N - macro_start;
float* h_b = h + static_cast<long long>(b) * N * N;
const float* v_b = v_macro + static_cast<long long>(b) * v_stride0;
constexpr int Leaves = MacroCols / Panel;
const float* gram_b = leaf_grams + static_cast<long long>(b) * Leaves * Panel * Panel;
const float* tau_b = tau + static_cast<long long>(b) * N + macro_start;
extern __shared__ float smem[];
float* c_tile = smem;
float* w_tile = c_tile + active_rows * Panel;
float* z_tile = w_tile + Panel * Panel;
float* g_tile = z_tile + Panel * Panel;
float* tau_tile = g_tile + Panel * Panel;
for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
const int row = idx / Panel;
const int col = idx - row * Panel;
c_tile[idx] = h_b[
static_cast<long long>(macro_start + row) * N +
macro_start + target_offset + col];
}
__syncthreads();
for (int leaf = 0; leaf < prev_leaves; ++leaf) {
for (int idx = tid; idx < Panel * Panel; idx += Threads) {
g_tile[idx] = gram_b[leaf * Panel * Panel + idx];
}
for (int idx = tid; idx < Panel; idx += Threads) {
tau_tile[idx] = tau_b[leaf * Panel + idx];
}
__syncthreads();
for (int idx = tid; idx < Panel * Panel; idx += Threads) {
const int row = idx / Panel;
const int col = idx - row * Panel;
float accum = 0.0f;
for (int rel = 0; rel < active_rows; ++rel) {
accum = fmaf(
v_b[static_cast<long long>(rel) * MacroCols + leaf * Panel + row],
c_tile[rel * Panel + col],
accum);
}
w_tile[idx] = accum;
}
__syncthreads();
if (tid < Panel) {
const int col = tid;
float zv[Panel];
#pragma unroll
for (int row = 0; row < Panel; ++row) {
float accum = 0.0f;
#pragma unroll
for (int inner = 0; inner < Panel; ++inner) {
if (inner < row) {
accum = fmaf(g_tile[inner * Panel + row], zv[inner], accum);
}
}
const float zi = tau_tile[row] * (w_tile[row * Panel + col] - accum);
zv[row] = zi;
z_tile[row * Panel + col] = zi;
}
}
__syncthreads();
for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
const int rel = idx / Panel;
const int col = idx - rel * Panel;
float value = c_tile[idx];
#pragma unroll
for (int k = 0; k < Panel; ++k) {
value = fmaf(
-v_b[static_cast<long long>(rel) * MacroCols + leaf * Panel + k],
z_tile[k * Panel + col],
value);
}
c_tile[idx] = value;
}
__syncthreads();
}
for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
const int row = idx / Panel;
const int col = idx - row * Panel;
h_b[
static_cast<long long>(macro_start + row) * N +
macro_start + target_offset + col] = c_tile[idx];
}
}
} // namespace
int64_t detect_upper_512_cuda(torch::Tensor a) {
TORCH_CHECK(a.is_cuda(), "a must be CUDA");
TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
TORCH_CHECK(a.dim() == 3 && a.size(1) == kN512 && a.size(2) == kN512,
"upper detector expects [batch,512,512]");
TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
return all_matrices_upper_certified<kN512>(a) ? 1 : 0;
}
int64_t detect_upper_1024_cuda(torch::Tensor a) {
TORCH_CHECK(a.is_cuda(), "a must be CUDA");
TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
TORCH_CHECK(a.dim() == 3 && a.size(1) == kN1024 && a.size(2) == kN1024,
"upper detector expects [batch,1024,1024]");
TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
return all_matrices_upper_certified<kN1024>(a) ? 1 : 0;
}
std::vector<torch::Tensor> qr_small_cuda(torch::Tensor a) {
const int n = static_cast<int>(a.size(1));
const int batch = static_cast<int>(a.size(0));
auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
if (n == kN) {
qr32_kernel<<<batch, kThreads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>());
} else if (n == kN176) {
constexpr int threads = 1024;
constexpr int shared_bytes = kN176 * kLD176Resident * sizeof(float);
static bool attrs_set = false;
if (!attrs_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr176_resident_kernel<threads>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr176_resident_kernel<threads>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
attrs_set = true;
}
qr176_resident_kernel<threads><<<batch, threads, shared_bytes>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>());
} else if (n == kN352) {
constexpr int panel_shared_bytes = ((kN352 + 1) * kPanel352 + kPanelThreads352) * sizeof(float);
constexpr int update_shared_bytes =
(kN352 * kTileUpdate352 + kN352 * kPanel352 + 2 * kPanel352 * kTileUpdate352) *
static_cast<int>(sizeof(float));
constexpr int copy_threads = 256;
constexpr int copy_blocks = 1024;
static bool update_attrs_set = false;
if (!update_attrs_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_tile_update_kernel<
kN352,
kPanel352,
kTileUpdate352,
kPanelThreads352>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
update_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_tile_update_kernel<
kN352,
kPanel352,
kTileUpdate352,
kPanelThreads352>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
update_attrs_set = true;
}
auto t_scratch = torch::empty(
{a.size(0), kPanel352, kPanel352},
a.options());
copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
static_cast<long long>(batch) * kN352 * kN352);
for (int panel_start = 0; panel_start < kN352; panel_start += kPanel352) {
qr_panel_cached_kernel<kN352, kPanel352, kPanelThreads352, false>
<<<batch, kPanelThreads352, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
nullptr,
0,
panel_start);
const int panel_end = panel_start + kPanel352;
if (panel_end < kN352) {
const int trailing_cols = kN352 - panel_end;
const int tiles = (trailing_cols + kTileUpdate352 - 1) / kTileUpdate352;
qr_panel_tile_update_kernel<
kN352,
kPanel352,
kTileUpdate352,
kPanelThreads352>
<<<dim3(tiles, batch), kPanelThreads352, update_shared_bytes>>>(
h.data_ptr<float>(),
t_scratch.data_ptr<float>(),
panel_start);
}
}
} else if (n == kN512) {
const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS();
constexpr int kMacro512 = 2 * kPanel512;
constexpr int kMacro64_512 = 4 * kPanel512;
constexpr int kMacroLeaves64_512 = kMacro64_512 / kPanel512;
constexpr int panel_shared_bytes = ((kN512 + 1) * kPanel512 + kPanelThreads512) * sizeof(float);
constexpr int prep_threads = 256;
constexpr int prep_shared_bytes =
(kN512 * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) * sizeof(float);
auto leaf_grams = torch::empty(
{a.size(0), kMacroLeaves64_512, kPanel512, kPanel512},
a.options());
auto g_local = torch::empty({a.size(0), kMacro512, kMacro512}, a.options());
auto g_macro = torch::empty({a.size(0), kMacro64_512, kMacro64_512}, a.options());
auto w_local_workspace = torch::empty({a.size(0), kMacro512, kMacro512}, a.options());
auto w_workspace = torch::empty({a.size(0), kMacro64_512, kN512}, a.options());
constexpr size_t lt_workspace_bytes = 4 * 1024 * 1024;
auto lt_workspace = torch::empty(
{static_cast<long long>(lt_workspace_bytes)},
a.options().dtype(at::kByte));
constexpr int copy_threads = 256;
constexpr int copy_blocks = 4096;
copy_first_cols_v4_kernel<kN512, kMacro64_512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
batch);
for (int macro_start = 0; macro_start < kN512; macro_start += kMacro64_512) {
const int active_rows = kN512 - macro_start;
const int second_start = macro_start + kMacro512;
const int macro_end = macro_start + kMacro64_512;
const int prep_active_shared_bytes =
(active_rows * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) *
static_cast<int>(sizeof(float));
auto v_macro = torch::empty({a.size(0), active_rows, kMacro64_512}, a.options());
for (int leaf = 0; leaf < 2; ++leaf) {
const int leaf_offset = leaf * kPanel512;
const int panel_start = macro_start + leaf_offset;
if (leaf == 0) {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_512,
leaf_offset,
leaf_offset);
auto g_leaf = leaf_grams.select(1, leaf);
cublas_leaf0_gram_from_macro(
g_leaf.data_ptr<float>(),
g_leaf.stride(0),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_512,
active_rows,
batch,
true);
apply_prev_leaves_to_next_leaf_kernel<
kN512,
kMacro64_512,
kPanel512,
prep_threads>
<<<batch, prep_threads, prep_active_shared_bytes>>>(
h.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
leaf_grams.data_ptr<float>(),
tau.data_ptr<float>(),
macro_start,
leaf + 1,
leaf_offset + kPanel512);
} else {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_512,
leaf_offset,
leaf_offset);
}
}
auto v_first = v_macro.as_strided(
{a.size(0), active_rows, kMacro512},
{v_macro.stride(0), kMacro64_512, 1});
auto c_next = h.slice(1, macro_start, kN512).slice(2, second_start, macro_end);
at::bmm_out(g_local, v_first.transpose(1, 2), v_first);
at::bmm_out(w_local_workspace, v_first.transpose(1, 2), c_next);
constexpr int local_apply_threads = 128;
solve_inverse_wy_from_gram_blocked_kernel<
kMacro512,
kPanel512,
local_apply_threads>
<<<dim3(1, batch), local_apply_threads, 0>>>(
g_local.data_ptr<float>(),
tau.data_ptr<float>(),
w_local_workspace.data_ptr<float>(),
w_local_workspace.data_ptr<float>(),
kN512,
macro_start,
kMacro512);
c_next.baddbmm_(v_first, w_local_workspace, 1.0, -1.0);
const int second_active_rows = kN512 - second_start;
const int second_prep_active_shared_bytes =
(second_active_rows * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) *
static_cast<int>(sizeof(float));
for (int leaf = 0; leaf < 2; ++leaf) {
const int leaf_offset = leaf * kPanel512;
const int panel_start = second_start + leaf_offset;
const int macro_row_offset = kMacro512 + leaf_offset;
const int macro_col_offset = kMacro512 + leaf_offset;
float* v_second_base =
v_macro.data_ptr<float>() +
static_cast<long long>(kMacro512) * kMacro64_512 +
kMacro512;
float* g_second_base =
leaf_grams.data_ptr<float>() +
static_cast<long long>(2) * kPanel512 * kPanel512;
if (leaf == 0) {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_512,
macro_row_offset,
macro_col_offset);
auto g_leaf = leaf_grams.select(1, 2);
cublas_leaf0_gram_from_macro(
g_leaf.data_ptr<float>(),
g_leaf.stride(0),
v_second_base,
v_macro.stride(0),
kMacro64_512,
second_active_rows,
batch,
true);
apply_prev_leaves_to_next_leaf_kernel<
kN512,
kMacro64_512,
kPanel512,
prep_threads>
<<<batch, prep_threads, second_prep_active_shared_bytes>>>(
h.data_ptr<float>(),
v_second_base,
v_macro.stride(0),
g_second_base,
tau.data_ptr<float>(),
second_start,
leaf + 1,
leaf_offset + kPanel512);
} else {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_512,
macro_row_offset,
macro_col_offset);
}
}
if (macro_end < kN512) {
const int trailing_cols = kN512 - macro_end;
at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
auto w = w_workspace.as_strided(
{a.size(0), kMacro64_512, trailing_cols},
{kMacro64_512 * trailing_cols, trailing_cols, 1});
auto z = w;
auto c = h.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
if (macro_start == 0) {
auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
at::bmm_out(w, v_macro.transpose(1, 2), c_in);
} else {
at::bmm_out(w, v_macro.transpose(1, 2), c);
}
constexpr int apply_t_threads = 96;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_gram_blocked_kernel<kMacro64_512, kPanel512, apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
g_macro.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
kN512,
macro_start,
trailing_cols);
if (macro_start < 32) {
auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
cublaslt_tail_update_out_of_place(
c.data_ptr<float>(),
c_in.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
z.data_ptr<float>(),
kN512,
kMacro64_512,
active_rows,
trailing_cols,
batch,
false,
lt_workspace.data_ptr(),
lt_workspace_bytes);
} else {
c.baddbmm_(v_macro, z, 1.0, -1.0);
}
}
}
at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
} else if (n == kN1024) {
const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS();
constexpr int panel_work1024 = ((kPanelThreads1024 / 32) + 1) * kPanel1024;
constexpr int panel_shared_bytes = ((kN1024 + 1) * kPanel1024 + panel_work1024) * sizeof(float);
constexpr int kMacro1024 = 2 * kPanel1024;
constexpr int kMacro64_1024 = 4 * kPanel1024;
constexpr int kMacroLeaves64_1024 = kMacro64_1024 / kPanel1024;
constexpr int prep_threads = 256;
constexpr int prep_shared_bytes =
(kN1024 * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) * sizeof(float);
auto leaf_grams = torch::empty(
{a.size(0), kMacroLeaves64_1024, kPanel1024, kPanel1024},
a.options());
auto g_local = torch::empty({a.size(0), kMacro1024, kMacro1024}, a.options());
auto g_macro = torch::empty({a.size(0), kMacro64_1024, kMacro64_1024}, a.options());
auto w_local_workspace = torch::empty({a.size(0), kMacro1024, kMacro1024}, a.options());
auto w_workspace = torch::empty({a.size(0), kMacro64_1024, kN1024}, a.options());
constexpr size_t lt_workspace_bytes = 4 * 1024 * 1024;
auto lt_workspace = torch::empty(
{static_cast<long long>(lt_workspace_bytes)},
a.options().dtype(at::kByte));
constexpr int copy_threads = 256;
constexpr int copy_blocks = 4096;
static bool panel_attrs_set = false;
if (!panel_attrs_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro1024,
kPanel1024,
prep_threads>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
prep_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro1024,
kPanel1024,
prep_threads>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro64_1024,
kPanel1024,
prep_threads>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
prep_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro64_1024,
kPanel1024,
prep_threads>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
panel_attrs_set = true;
}
copy_first_cols_v4_kernel<kN1024, kMacro64_1024, copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
batch);
for (int macro_start = 0; macro_start < kN1024; macro_start += kMacro64_1024) {
const int active_rows = kN1024 - macro_start;
const int second_start = macro_start + kMacro1024;
const int macro_end = macro_start + kMacro64_1024;
const int prep_active_shared_bytes =
(active_rows * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) *
static_cast<int>(sizeof(float));
auto v_macro = torch::empty({a.size(0), active_rows, kMacro64_1024}, a.options());
for (int leaf = 0; leaf < 2; ++leaf) {
const int leaf_offset = leaf * kPanel1024;
const int panel_start = macro_start + leaf_offset;
if (leaf == 0) {
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
<<<batch, kPanelThreads1024, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_1024,
leaf_offset,
leaf_offset);
auto g_leaf = leaf_grams.select(1, leaf);
cublas_leaf0_gram_from_macro(
g_leaf.data_ptr<float>(),
g_leaf.stride(0),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_1024,
active_rows,
batch,
true);
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro64_1024,
kPanel1024,
prep_threads>
<<<batch, prep_threads, prep_active_shared_bytes>>>(
h.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
leaf_grams.data_ptr<float>(),
tau.data_ptr<float>(),
macro_start,
leaf + 1,
leaf_offset + kPanel1024);
} else {
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
<<<batch, kPanelThreads1024, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_1024,
leaf_offset,
leaf_offset);
}
}
auto v_first = v_macro.as_strided(
{a.size(0), active_rows, kMacro1024},
{v_macro.stride(0), kMacro64_1024, 1});
auto c_next = h.slice(1, macro_start, kN1024).slice(2, second_start, macro_end);
at::bmm_out(g_local, v_first.transpose(1, 2), v_first);
at::bmm_out(w_local_workspace, v_first.transpose(1, 2), c_next);
constexpr int local_apply_threads = 128;
solve_inverse_wy_from_gram_blocked_kernel<
kMacro1024,
kPanel1024,
local_apply_threads>
<<<dim3(1, batch), local_apply_threads, 0>>>(
g_local.data_ptr<float>(),
tau.data_ptr<float>(),
w_local_workspace.data_ptr<float>(),
w_local_workspace.data_ptr<float>(),
kN1024,
macro_start,
kMacro1024);
c_next.baddbmm_(v_first, w_local_workspace, 1.0, -1.0);
const int second_active_rows = kN1024 - second_start;
const int second_prep_active_shared_bytes =
(second_active_rows * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) *
static_cast<int>(sizeof(float));
for (int leaf = 0; leaf < 2; ++leaf) {
const int leaf_offset = leaf * kPanel1024;
const int panel_start = second_start + leaf_offset;
const int macro_row_offset = kMacro1024 + leaf_offset;
const int macro_col_offset = kMacro1024 + leaf_offset;
float* v_second_base =
v_macro.data_ptr<float>() +
static_cast<long long>(kMacro1024) * kMacro64_1024 +
kMacro1024;
float* g_second_base =
leaf_grams.data_ptr<float>() +
static_cast<long long>(2) * kPanel1024 * kPanel1024;
if (leaf == 0) {
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
<<<batch, kPanelThreads1024, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_1024,
macro_row_offset,
macro_col_offset);
auto g_leaf = leaf_grams.select(1, 2);
cublas_leaf0_gram_from_macro(
g_leaf.data_ptr<float>(),
g_leaf.stride(0),
v_second_base,
v_macro.stride(0),
kMacro64_1024,
second_active_rows,
batch,
true);
apply_prev_leaves_to_next_leaf_kernel<
kN1024,
kMacro64_1024,
kPanel1024,
prep_threads>
<<<batch, prep_threads, second_prep_active_shared_bytes>>>(
h.data_ptr<float>(),
v_second_base,
v_macro.stride(0),
g_second_base,
tau.data_ptr<float>(),
second_start,
leaf + 1,
leaf_offset + kPanel1024);
} else {
qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
<<<batch, kPanelThreads1024, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacro64_1024,
macro_row_offset,
macro_col_offset);
}
}
if (macro_end < kN1024) {
const int trailing_cols = kN1024 - macro_end;
at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
auto w = w_workspace.as_strided(
{a.size(0), kMacro64_1024, trailing_cols},
{kMacro64_1024 * trailing_cols, trailing_cols, 1});
auto c = h.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
if (macro_start == 0) {
auto c_in = a.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c_in.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kN1024,
kMacro64_1024,
kMacro64_1024,
active_rows,
trailing_cols,
batch,
true);
} else {
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kN1024,
kMacro64_1024,
kMacro64_1024,
active_rows,
trailing_cols,
batch,
true);
}
constexpr int apply_t_threads = 96;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_gram_blocked_kernel<
kMacro64_1024,
kPanel1024,
apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
g_macro.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
w.data_ptr<float>(),
kN1024,
macro_start,
trailing_cols);
if (macro_start == 0) {
auto c_in = a.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
cublaslt_tail_update_out_of_place(
c.data_ptr<float>(),
c_in.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
w.data_ptr<float>(),
kN1024,
kMacro64_1024,
active_rows,
trailing_cols,
batch,
true,
lt_workspace.data_ptr(),
lt_workspace_bytes);
} else {
c.baddbmm_(v_macro, w, 1.0, -1.0);
}
}
}
at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
}
return {h, tau};
}
int64_t detect_tiny_suffix_512_cuda(torch::Tensor a) {
TORCH_CHECK(a.is_cuda(), "a must be CUDA");
TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
TORCH_CHECK(a.dim() == 3 && a.size(1) == kN512 && a.size(2) == kN512,
"detector expects [batch,512,512]");
constexpr int k0 = kN512 / 2;
constexpr int k1 = kN512 / 2 + 2 * kPanel512;
constexpr int k2 = (3 * kN512) / 4;
constexpr int threads = 256;
const int batch = static_cast<int>(a.size(0));
// Cheap deterministic rejector. A rejection only selects the safe path;
// acceptance still requires the complete per-matrix bound below.
auto reject = torch::empty({1}, a.options().dtype(at::kInt));
C10_CUDA_CHECK(cudaMemset(reject.data_ptr<int>(), 0, sizeof(int)));
suffix_sample_reject_kernel<kN512, 32><<<batch, 32, 0>>>(
a.data_ptr<float>(), reject.data_ptr<int>());
int host_reject = 0;
C10_CUDA_CHECK(cudaMemcpy(
&host_reject, reject.data_ptr<int>(), sizeof(int), cudaMemcpyDeviceToHost));
if (host_reject != 0) return 0;
auto factors = torch::empty({batch}, a.options().dtype(at::kInt));
suffix_factor_cols_kernel<kN512, kPanel512, threads>
<<<batch, threads, 0>>>(
a.data_ptr<float>(),
factors.data_ptr<int>(),
k0,
k1,
k2);
auto factor_result = torch::empty({1}, a.options().dtype(at::kInt));
reduce_factor_cols_kernel<threads>
<<<1, threads, 0>>>(
factors.data_ptr<int>(),
factor_result.data_ptr<int>(),
batch);
int factor_cols = 0;
C10_CUDA_CHECK(cudaMemcpy(
&factor_cols,
factor_result.data_ptr<int>(),
sizeof(int),
cudaMemcpyDeviceToHost));
return factor_cols;
}
std::vector<torch::Tensor> qr_cholqr_hr512_cuda(
torch::Tensor a,
torch::Tensor r,
torch::Tensor info) {
TORCH_CHECK(a.is_cuda(), "a must be CUDA");
TORCH_CHECK(r.is_cuda(), "r must be CUDA");
TORCH_CHECK(info.is_cuda(), "info must be CUDA");
TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
TORCH_CHECK(r.scalar_type() == at::kFloat, "r must be float32");
TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
TORCH_CHECK(a.dim() == 3 && a.size(1) == kN512 && a.size(2) == kN512,
"CholQR-HR expects a [batch,512,512] input");
TORCH_CHECK(r.dim() == 3 && r.size(0) == a.size(0) && r.size(1) == kN512 && r.size(2) == kN512,
"CholQR-HR expects r [batch,512,512]");
TORCH_CHECK(info.dim() == 1 && info.size(0) == a.size(0),
"CholQR-HR expects info [batch]");
TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
TORCH_CHECK(r.is_contiguous(), "r must be contiguous");
TORCH_CHECK(info.is_contiguous(), "info must be contiguous");
auto h = a.clone();
auto tau = torch::empty({a.size(0), kN512}, a.options());
auto ok = torch::empty({a.size(0)}, a.options().dtype(at::kInt));
constexpr int threads = 256;
const int batch = static_cast<int>(a.size(0));
cholqr_hr_lu_kernel<kN512, threads><<<batch, threads, 0>>>(
r.data_ptr<float>(),
info.data_ptr<int>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
ok.data_ptr<int>());
C10_CUDA_CHECK(cudaGetLastError());
return {h, tau, ok};
}
std::vector<torch::Tensor> qr_small_prefix_cuda(torch::Tensor a, int64_t factor_cols_arg) {
const int n = static_cast<int>(a.size(1));
const int batch = static_cast<int>(a.size(0));
const int factor_cols = static_cast<int>(factor_cols_arg);
if (n != kN512 || factor_cols <= 0 || factor_cols >= n || (factor_cols & 15) != 0) {
return qr_small_cuda(a);
}
auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS();
constexpr int panel_shared_bytes = ((kN512 + 1) * kPanel512 + kPanelThreads512) * sizeof(float);
constexpr int copy_threads = 256;
constexpr int copy_blocks = 4096;
constexpr int kMacroPrefix512 = 2 * kPanel512;
constexpr int kMacroPrefixLeaves512 = kMacroPrefix512 / kPanel512;
constexpr int prep_threads = 256;
constexpr int prep_shared_bytes =
(kN512 * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) * sizeof(float);
if (factor_cols >= kMacroPrefix512) {
copy_first_cols_v4_kernel<kN512, kMacroPrefix512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
batch);
zero_suffix_cols_v4_kernel<kN512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
h.data_ptr<float>(),
factor_cols,
batch);
} else {
copy_prefix_zero_suffix_v4_kernel<kN512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
static_cast<long long>(batch) * kN512 * kN512,
factor_cols);
}
if ((factor_cols % kMacroPrefix512) == 0) {
auto leaf_grams = torch::empty(
{a.size(0), kMacroPrefixLeaves512, kPanel512, kPanel512},
a.options());
auto s_macro = torch::empty({a.size(0), kMacroPrefix512, kPanel512}, a.options());
auto w_workspace = torch::empty({a.size(0), kMacroPrefix512, factor_cols}, a.options());
constexpr size_t lt_workspace_bytes = 4 * 1024 * 1024;
auto lt_workspace = torch::empty(
{static_cast<long long>(lt_workspace_bytes)},
a.options().dtype(at::kByte));
for (int macro_start = 0; macro_start < factor_cols; macro_start += kMacroPrefix512) {
const int active_rows = kN512 - macro_start;
const int macro_end = macro_start + kMacroPrefix512;
const int prep_active_shared_bytes =
(active_rows * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) *
static_cast<int>(sizeof(float));
auto v_macro = torch::empty({a.size(0), active_rows, kMacroPrefix512}, a.options());
for (int leaf = 0; leaf < kMacroPrefixLeaves512; ++leaf) {
const int leaf_offset = leaf * kPanel512;
const int panel_start = macro_start + leaf_offset;
const int leaf_active_rows = kN512 - panel_start;
if (leaf + 1 < kMacroPrefixLeaves512) {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacroPrefix512,
leaf_offset,
leaf_offset);
auto g_leaf = leaf_grams.select(1, leaf);
cublas_leaf0_gram_from_macro(
g_leaf.data_ptr<float>(),
g_leaf.stride(0),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacroPrefix512,
active_rows,
batch,
false);
apply_prev_leaves_to_next_leaf_kernel<
kN512,
kMacroPrefix512,
kPanel512,
prep_threads>
<<<batch, prep_threads, prep_active_shared_bytes>>>(
h.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
leaf_grams.data_ptr<float>(),
tau.data_ptr<float>(),
macro_start,
leaf + 1,
leaf_offset + kPanel512);
} else {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
nullptr,
nullptr,
0,
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacroPrefix512,
leaf_offset,
leaf_offset);
}
}
if (macro_end < factor_cols) {
const int trailing_cols = factor_cols - macro_end;
auto g_first = leaf_grams.select(1, 0);
cublas_leaf1_cross_gram_from_macro(
s_macro.data_ptr<float>(),
s_macro.stride(0),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kMacroPrefix512,
active_rows,
batch,
false);
auto w = w_workspace.as_strided(
{a.size(0), kMacroPrefix512, trailing_cols},
{kMacroPrefix512 * trailing_cols, trailing_cols, 1});
auto z = w;
auto c = h.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
if (macro_start == 0) {
auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
at::bmm_out(w, v_macro.transpose(1, 2), c_in);
} else {
at::bmm_out(w, v_macro.transpose(1, 2), c);
}
constexpr int apply_t_threads = 96;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_split_gram_blocked_kernel<
kMacroPrefix512,
kPanel512,
apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
g_first.data_ptr<float>(),
g_first.stride(0),
s_macro.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
kN512,
macro_start,
trailing_cols);
if (macro_start == 0) {
auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
cublaslt_tail_update_out_of_place(
c.data_ptr<float>(),
c_in.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
z.data_ptr<float>(),
kN512,
kMacroPrefix512,
active_rows,
trailing_cols,
batch,
old_allow_tf32,
lt_workspace.data_ptr(),
lt_workspace_bytes);
} else {
c.baddbmm_(v_macro, z, 1.0, -1.0);
}
}
}
} else {
auto t_scratch = torch::empty({a.size(0), kPanel512, kPanel512}, a.options());
auto w_workspace = torch::empty({a.size(0), kPanel512, factor_cols}, a.options());
for (int panel_start = 0; panel_start < factor_cols; panel_start += kPanel512) {
const int panel_end = panel_start + kPanel512;
if (panel_end < factor_cols) {
const int active_rows = kN512 - panel_start;
const int trailing_cols = factor_cols - panel_end;
// Keep V compact. A max-stride slab measurably regressed the batched GEMMs.
auto v = torch::empty({a.size(0), active_rows, kPanel512}, a.options());
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, true, false, false, false>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
v.data_ptr<float>(),
v.stride(0),
panel_start);
at::globalContext().setAllowTF32CuBLAS(false);
at::bmm_out(t_scratch, v.transpose(1, 2), v);
at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
auto w = w_workspace.as_strided(
{a.size(0), kPanel512, trailing_cols},
{kPanel512 * trailing_cols, trailing_cols, 1});
auto z = w;
auto c = h.slice(1, panel_start, kN512).slice(2, panel_end, factor_cols);
at::bmm_out(w, v.transpose(1, 2), c);
constexpr int apply_t_threads = 128;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_gram_kernel<kPanel512, apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
t_scratch.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
kN512,
panel_start,
trailing_cols);
c.baddbmm_(v, z, 1.0, -1.0);
} else {
qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false>
<<<batch, kPanelThreads512, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
nullptr,
0,
panel_start);
}
}
}
constexpr int zero_threads = 256;
const long long tau_total = static_cast<long long>(batch) * (kN512 - factor_cols);
int zero_blocks = static_cast<int>((tau_total + zero_threads - 1) / zero_threads);
if (zero_blocks < 1) zero_blocks = 1;
if (zero_blocks > 1024) zero_blocks = 1024;
zero_tau_suffix_kernel<kN512, zero_threads><<<zero_blocks, zero_threads, 0>>>(
tau.data_ptr<float>(), factor_cols, batch);
at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
return {h, tau};
}
std::vector<torch::Tensor> qr_2048_cuda(torch::Tensor a) {
const int batch = static_cast<int>(a.size(0));
auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
constexpr int macro_cols = 4 * kPanel2048;
auto t_scratch = torch::empty({a.size(0), kPanel2048, kPanel2048}, a.options());
auto g_macro = torch::empty({a.size(0), macro_cols, macro_cols}, a.options());
auto w_workspace = torch::empty({a.size(0), macro_cols, kN2048}, a.options());
auto z_workspace = torch::empty({a.size(0), kPanel2048, kN2048}, a.options());
constexpr int max_panel_shared_bytes = ((kN2048 + 1) * kPanel2048 + kPanelThreads2048) * sizeof(float);
constexpr int copy_threads = 256;
constexpr int copy_blocks = 1024;
static bool panel_attrs_set = false;
if (!panel_attrs_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
panel_attrs_set = true;
}
copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
static_cast<long long>(batch) * kN2048 * kN2048);
for (int macro_start = 0; macro_start < kN2048; macro_start += macro_cols) {
const int active_rows = kN2048 - macro_start;
const int macro_end = (macro_start + macro_cols < kN2048)
? macro_start + macro_cols
: kN2048;
const int panel_shared_bytes =
((active_rows + 1) * kPanel2048 + kPanelThreads2048) * sizeof(float);
auto v_macro = torch::empty({a.size(0), active_rows, macro_cols}, a.options());
for (int panel_start = macro_start; panel_start < macro_end; panel_start += kPanel2048) {
const int panel_end = panel_start + kPanel2048;
const int leaf_active_rows = kN2048 - panel_start;
const int macro_offset = panel_start - macro_start;
auto v_leaf = torch::empty({a.size(0), leaf_active_rows, kPanel2048}, a.options());
qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>
<<<batch, kPanelThreads2048, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
v_leaf.data_ptr<float>(),
v_leaf.stride(0),
panel_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
macro_cols,
macro_offset,
macro_offset);
if (panel_end < macro_end) {
const int local_cols = macro_end - panel_end;
auto w = w_workspace.as_strided(
{a.size(0), kPanel2048, local_cols},
{kPanel2048 * local_cols, local_cols, 1});
auto z = z_workspace.as_strided(
{a.size(0), kPanel2048, local_cols},
{kPanel2048 * local_cols, local_cols, 1});
auto c = h.slice(1, panel_start, kN2048).slice(2, panel_end, macro_end);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v_leaf.data_ptr<float>(),
v_leaf.stride(0),
kN2048,
kPanel2048,
kPanel2048,
leaf_active_rows,
local_cols,
batch,
true);
constexpr int local_apply_t_threads = 64;
const int local_apply_t_blocks =
(local_cols + local_apply_t_threads - 1) / local_apply_t_threads;
apply_t_transpose_kernel<kPanel2048, local_apply_t_threads>
<<<dim3(local_apply_t_blocks, batch), local_apply_t_threads, 0>>>(
t_scratch.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
local_cols);
c.baddbmm_(v_leaf, z, 1.0, -1.0);
}
}
if (macro_end < kN2048) {
const int trailing_cols = kN2048 - macro_end;
at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
auto w = w_workspace.as_strided(
{a.size(0), macro_cols, trailing_cols},
{macro_cols * trailing_cols, trailing_cols, 1});
auto c = h.slice(1, macro_start, kN2048).slice(2, macro_end, kN2048);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kN2048,
macro_cols,
macro_cols,
active_rows,
trailing_cols,
batch,
true);
constexpr int apply_t_threads = 64;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_gram_blocked_kernel<macro_cols, kPanel2048, apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
g_macro.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
w.data_ptr<float>(),
kN2048,
macro_start,
trailing_cols);
c.baddbmm_(v_macro, w, 1.0, -1.0);
}
}
return {h, tau};
}
std::vector<torch::Tensor> qr_4096_cuda(torch::Tensor a) {
const int batch = static_cast<int>(a.size(0));
auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
constexpr int super_panel = 2 * kPanel4096;
constexpr int early_macro = 8 * kPanel4096;
auto t_first = torch::empty({a.size(0), kPanel4096, kPanel4096}, a.options());
auto t_second = torch::empty({a.size(0), kPanel4096, kPanel4096}, a.options());
auto t_scratch = torch::empty({a.size(0), super_panel, super_panel}, a.options());
auto g_early = torch::empty({a.size(0), early_macro, early_macro}, a.options());
auto w_workspace = torch::empty({a.size(0), early_macro, kN4096}, a.options());
auto z_workspace = torch::empty({a.size(0), super_panel, kN4096}, a.options());
constexpr int panel_shared_bytes = ((kN4096 + 1) * kPanel4096 + kPanelThreads4096) * sizeof(float);
constexpr int apply_threads = 512;
constexpr int apply_shared_bytes =
(kN4096 * kPanel4096 + (apply_threads / 32) * kPanel4096 + kPanel4096) * sizeof(float);
constexpr int build_threads = 512;
constexpr int copy_threads = 256;
constexpr int copy_blocks = 8192;
static int late_max_rows = kPanel4096LateMaxRows;
static int late_panel_shared_bytes =
((kPanel4096LateMaxRows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
static bool panel_attrs_set = false;
if (!panel_attrs_set) {
int device = 0;
int max_shared_bytes = late_panel_shared_bytes;
C10_CUDA_CHECK(cudaGetDevice(&device));
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&max_shared_bytes,
cudaDevAttrMaxSharedMemoryPerBlockOptin,
device));
int device_rows =
((max_shared_bytes / static_cast<int>(sizeof(float))) - kPanelThreads4096Late) /
kPanel4096Late - 1;
if (device_rows > kPanel4096LateMaxRows) {
device_rows = kPanel4096LateMaxRows;
}
if (device_rows < 0) {
device_rows = 0;
}
late_max_rows = (device_rows / super_panel) * super_panel;
late_panel_shared_bytes =
((late_max_rows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
late_panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
late_panel_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
panel_attrs_set = true;
}
static bool apply_attrs_set = false;
if (!apply_attrs_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_panel_to_next_cols_kernel<kN4096, kPanel4096, kPanel4096, apply_threads>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
apply_shared_bytes));
C10_CUDA_CHECK(cudaFuncSetAttribute(
apply_panel_to_next_cols_kernel<kN4096, kPanel4096, kPanel4096, apply_threads>,
cudaFuncAttributePreferredSharedMemoryCarveout,
100));
apply_attrs_set = true;
}
copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
static_cast<long long>(batch) * kN4096 * kN4096);
for (int panel_start = 0; panel_start < kN4096;) {
const int active_rows = kN4096 - panel_start;
int panel_end = (panel_start + super_panel < kN4096)
? panel_start + super_panel
: kN4096;
if (active_rows <= late_max_rows) {
const int direct_shared_bytes =
((active_rows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
if (panel_end >= kN4096) {
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>
<<<batch, kPanelThreads4096Late, direct_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
nullptr,
0,
panel_start);
break;
}
const int trailing_cols = kN4096 - panel_end;
auto v = torch::empty({a.size(0), active_rows, super_panel}, a.options());
qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>
<<<batch, kPanelThreads4096Late, direct_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_scratch.data_ptr<float>(),
v.data_ptr<float>(),
v.stride(0),
panel_start);
auto w = w_workspace.as_strided(
{a.size(0), super_panel, trailing_cols},
{super_panel * trailing_cols, trailing_cols, 1});
auto z = z_workspace.as_strided(
{a.size(0), super_panel, trailing_cols},
{super_panel * trailing_cols, trailing_cols, 1});
auto c = h.slice(1, panel_start, kN4096).slice(2, panel_end, kN4096);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v.data_ptr<float>(),
v.stride(0),
kN4096,
super_panel,
super_panel,
active_rows,
trailing_cols,
batch,
true);
constexpr int apply_t_threads = 128;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
apply_t_transpose_kernel<super_panel, apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
t_scratch.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
trailing_cols);
c.baddbmm_(v, z, 1.0, -1.0);
panel_start += super_panel;
continue;
}
panel_end = (panel_start + early_macro < kN4096)
? panel_start + early_macro
: kN4096;
auto v_macro = torch::empty({a.size(0), active_rows, early_macro}, a.options());
for (int leaf_start = panel_start; leaf_start < panel_end; leaf_start += kPanel4096) {
const int leaf_end = leaf_start + kPanel4096;
const int leaf_active_rows = kN4096 - leaf_start;
const int macro_offset = leaf_start - panel_start;
auto v_leaf = torch::empty({a.size(0), leaf_active_rows, kPanel4096}, a.options());
qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>
<<<batch, kPanelThreads4096, panel_shared_bytes>>>(
h.data_ptr<float>(),
tau.data_ptr<float>(),
t_first.data_ptr<float>(),
v_leaf.data_ptr<float>(),
v_leaf.stride(0),
leaf_start,
v_macro.data_ptr<float>(),
v_macro.stride(0),
early_macro,
macro_offset,
macro_offset);
if (leaf_end < panel_end) {
const int local_cols = panel_end - leaf_end;
auto w = w_workspace.as_strided(
{a.size(0), kPanel4096, local_cols},
{kPanel4096 * local_cols, local_cols, 1});
auto z = z_workspace.as_strided(
{a.size(0), kPanel4096, local_cols},
{kPanel4096 * local_cols, local_cols, 1});
auto c = h.slice(1, leaf_start, kN4096).slice(2, leaf_end, panel_end);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v_leaf.data_ptr<float>(),
v_leaf.stride(0),
kN4096,
kPanel4096,
kPanel4096,
leaf_active_rows,
local_cols,
batch,
true);
constexpr int local_apply_t_threads = 64;
const int local_apply_t_blocks =
(local_cols + local_apply_t_threads - 1) / local_apply_t_threads;
apply_t_transpose_kernel<kPanel4096, local_apply_t_threads>
<<<dim3(local_apply_t_blocks, batch), local_apply_t_threads, 0>>>(
t_first.data_ptr<float>(),
w.data_ptr<float>(),
z.data_ptr<float>(),
local_cols);
c.baddbmm_(v_leaf, z, 1.0, -1.0);
}
}
if (panel_end < kN4096) {
const int trailing_cols = kN4096 - panel_end;
at::bmm_out(g_early, v_macro.transpose(1, 2), v_macro);
auto w = w_workspace.as_strided(
{a.size(0), early_macro, trailing_cols},
{early_macro * trailing_cols, trailing_cols, 1});
auto c = h.slice(1, panel_start, kN4096).slice(2, panel_end, kN4096);
cublas_w_from_vt_c(
w.data_ptr<float>(),
c.data_ptr<float>(),
v_macro.data_ptr<float>(),
v_macro.stride(0),
kN4096,
early_macro,
early_macro,
active_rows,
trailing_cols,
batch,
true);
constexpr int apply_t_threads = 64;
const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
solve_inverse_wy_from_gram_blocked_kernel<
early_macro,
kPanel4096,
apply_t_threads>
<<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
g_early.data_ptr<float>(),
tau.data_ptr<float>(),
w.data_ptr<float>(),
w.data_ptr<float>(),
kN4096,
panel_start,
trailing_cols);
c.baddbmm_(v_macro, w, 1.0, -1.0);
}
panel_start += early_macro;
}
return {h, tau};
}
"""
_EXT = None
_EXT_FAILED = False
def _load_ext():
global _EXT, _EXT_FAILED
if _EXT is None and not _EXT_FAILED:
try:
from torch.utils.cpp_extension import load
source_dir = Path(tempfile.gettempdir()) / "qr_householder_ext_n512_all_tf32_gram_v12_macro32"
source_dir.mkdir(parents=True, exist_ok=True)
cuda_path = source_dir / "qr_householder_all.cu"
cuda_path.write_text(_CPP_SRC + "\n" + _CUDA_SRC)
_EXT = load(
name="qr_householder_ext_n512_all_tf32_gram_v12_macro32",
sources=[str(cuda_path)],
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
],
verbose=False,
)
except Exception as exc:
_EXT_FAILED = True
raise RuntimeError("qr_householder extension build failed") from exc
return _EXT
def _cholqr_hr512_experiment(contiguous: torch.Tensor, ext) -> output_t:
old_allow_tf32 = torch.backends.cuda.matmul.allow_tf32
torch.backends.cuda.matmul.allow_tf32 = False
try:
gram = torch.bmm(contiguous.transpose(1, 2), contiguous)
r, info = torch.linalg.cholesky_ex(gram, upper=True, check_errors=False)
finally:
torch.backends.cuda.matmul.allow_tf32 = old_allow_tf32
cand_h, cand_tau, ok = ext.qr_cholqr_hr512(
contiguous,
r.contiguous(),
info.contiguous(),
)
ok_mask = ok.to(dtype=torch.bool)
if bool(ok_mask.all()):
return cand_h, cand_tau
safe_h, safe_tau = ext.qr_small(contiguous)
h = torch.where(ok_mask.view(-1, 1, 1), cand_h, safe_h)
tau = torch.where(ok_mask.view(-1, 1), cand_tau, safe_tau)
return h, tau
def _small_qr(data: torch.Tensor) -> output_t:
ext = _load_ext()
if ext is None:
return torch.ops.aten.geqrf.default(data)
contiguous = data if data.is_contiguous() else data.contiguous()
if contiguous.shape[-1] == 512 and contiguous.shape[0] <= 32:
return torch.ops.aten.geqrf.default(contiguous)
if contiguous.shape[-1] == 512:
factor_cols = int(ext.detect_tiny_suffix_512(contiguous))
if 0 < factor_cols < 512:
h, tau = ext.qr_small_prefix(contiguous, factor_cols)
return h, tau
h, tau = ext.qr_small(contiguous)
return h, tau
def _large_qr(data: torch.Tensor) -> output_t:
ext = _load_ext()
if ext is not None:
contiguous = data if data.is_contiguous() else data.contiguous()
if contiguous.shape[-1] == 2048:
h, tau = ext.qr_2048(contiguous)
return h, tau
if contiguous.shape[-1] == 4096:
h, tau = ext.qr_4096(contiguous)
return h, tau
return torch.ops.aten.geqrf.default(data)
def custom_kernel(data: input_t) -> output_t:
if (
data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[-1] == data.shape[-2]
and data.shape[-1] in (32, 176, 352, 512, 1024)
):
return _small_qr(data)
if (
data.is_cuda
and data.dtype == torch.float32
and data.dim() == 3
and data.shape[-1] == data.shape[-2]
and data.shape[-1] in (2048, 4096)
):
return _large_qr(data)
return torch.ops.aten.geqrf.default(data)
scrolls · 3814 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