submission 843487
Mertt dönmez · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1828 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-843487?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:919297efe174e78047f82b58f785cbce6289484bac63ad43f51fc9b8a240940b
license declaredunknown
license concludedunknown
authorsMertt dönmez
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float a[MAX_N * MAX_N];Kernel source
submission.py1828 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
from pathlib import Path
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = """
void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int mode);
void qr_n512_custom(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_n512_human(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cmath>
#include <stdexcept>
#include <string>
namespace {
constexpr int MAX_N = 32;
constexpr int PANEL_THREADS = 512;
constexpr int MODE_CUBLAS_GRAM_CUBLAS_T = 8; // n512 large batch
constexpr int MODE_CUBLAS_GRAM_FUSED_T = 9; // n352 / n1024
constexpr int MODE_CUSTOM_GRAM_FUSED_T = 10; // n2048
constexpr int MODE_N512_CUSTOM_UPDATE = 11; // n512 scaffold: custom C -= V @ W2
inline void check_cuda(cudaError_t status, const char* where) {
if (status != cudaSuccess) {
throw std::runtime_error(std::string(where) + ": " + cudaGetErrorString(status));
}
}
inline void check_blas(cublasStatus_t status, const char* where) {
if (status != CUBLAS_STATUS_SUCCESS) {
throw std::runtime_error(std::string(where) + ": cuBLAS status " + std::to_string(status));
}
}
inline void check_solver(cusolverStatus_t status, const char* where) {
if (status != CUSOLVER_STATUS_SUCCESS) {
throw std::runtime_error(std::string(where) + ": cuSOLVER status " + std::to_string(status));
}
}
cublasHandle_t get_blas_handle() {
static cublasHandle_t handle = nullptr;
static bool initialized = false;
if (!initialized) {
check_blas(cublasCreate(&handle), "cublasCreate");
#ifdef CUBLAS_TF32_TENSOR_OP_MATH
check_blas(cublasSetMathMode(handle, CUBLAS_TF32_TENSOR_OP_MATH), "cublasSetMathMode");
#endif
initialized = true;
}
return handle;
}
template <int QR_THREADS>
__global__ void qr32_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int batch,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
__shared__ float a[MAX_N * MAX_N];
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
const float* src = a_in + static_cast<long long>(b) * n * n;
float* dst = h_out + static_cast<long long>(b) * n * n;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
for (int idx = tid; idx < n * n; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[i * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const float alpha = a[k * n + k];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[k * n + k] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[i * n + k] *= scale_s;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[k * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[i * n + k], a[i * n + j], dot);
}
const float w = tau * dot;
a[k * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
a[i * n + j] = fmaf(-a[i * n + k], w, a[i * n + j]);
}
}
__syncthreads();
}
for (int idx = tid; idx < n * n; idx += blockDim.x) {
dst[idx] = a[idx];
}
}
// n176 fits in B200 dynamic shared memory, so the whole matrix can be factored
// inside one CTA without repeated global-memory panel traffic.
template <int QR_THREADS>
__global__ void qr_shared_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src = a_in + static_cast<long long>(b) * matrix_size;
float* dst = h_out + static_cast<long long>(b) * matrix_size;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
extern __shared__ float a[];
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[static_cast<long long>(i) * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(k) * n + k;
const float alpha = a[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[diag] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[static_cast<long long>(i) * n + k] *= scale;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[static_cast<long long>(k) * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[static_cast<long long>(i) * n + k],
a[static_cast<long long>(i) * n + j],
dot);
}
const float w = tau * dot;
a[static_cast<long long>(k) * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
const long long idx = static_cast<long long>(i) * n + j;
a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
}
}
__syncthreads();
}
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
dst[idx] = a[idx];
}
}
// Simple one-CTA Householder QR for mid-size shapes where launch count and
// batch-level parallelism beat cuSOLVER overhead.
template <int QR_THREADS>
__global__ void qr_global_kernel(const float* __restrict__ a_in,
float* __restrict__ h_out,
float* __restrict__ tau_out,
int n) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src = a_in + static_cast<long long>(b) * matrix_size;
float* a = h_out + static_cast<long long>(b) * matrix_size;
float* tau_dst = tau_out + static_cast<long long>(b) * n;
__shared__ float red[QR_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
a[idx] = src[idx];
}
__syncthreads();
for (int k = 0; k < n; ++k) {
float local = 0.0f;
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
const float x = a[static_cast<long long>(i) * n + k];
local = fmaf(x, x, local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
const long long diag = static_cast<long long>(k) * n + k;
const float alpha = a[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_dst[k] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau = (beta - alpha) / beta;
tau_s = tau;
scale_s = 1.0f / (alpha - beta);
a[diag] = beta;
tau_dst[k] = tau;
}
}
__syncthreads();
if (tau_s != 0.0f) {
for (int i = k + 1 + tid; i < n; i += blockDim.x) {
a[static_cast<long long>(i) * n + k] *= scale_s;
}
}
__syncthreads();
const float tau = tau_s;
for (int j = k + 1 + tid; j < n; j += blockDim.x) {
float dot = a[static_cast<long long>(k) * n + j];
for (int i = k + 1; i < n; ++i) {
dot = fmaf(a[static_cast<long long>(i) * n + k],
a[static_cast<long long>(i) * n + j],
dot);
}
const float w = tau * dot;
a[static_cast<long long>(k) * n + j] -= w;
for (int i = k + 1; i < n; ++i) {
const long long idx = static_cast<long long>(i) * n + j;
a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
}
}
__syncthreads();
}
}
// Tiled row-major <-> column-major transform. The blocked-WY path stores its
// working matrix in column-major order but must return row-major compact H.
__global__ void transpose_tiled_kernel(const float* __restrict__ src,
float* __restrict__ dst,
int n) {
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
__shared__ float tile[TILE_DIM][TILE_DIM + 1];
const int b = blockIdx.z;
const long long matrix_size = static_cast<long long>(n) * n;
const float* src_b = src + static_cast<long long>(b) * matrix_size;
float* dst_b = dst + static_cast<long long>(b) * matrix_size;
int x = blockIdx.x * TILE_DIM + threadIdx.x;
int y = blockIdx.y * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
tile[threadIdx.y + j][threadIdx.x] =
src_b[static_cast<long long>(y + j) * n + x];
}
}
__syncthreads();
x = blockIdx.y * TILE_DIM + threadIdx.x;
y = blockIdx.x * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
dst_b[static_cast<long long>(y + j) * n + x] =
tile[threadIdx.x][threadIdx.y + j];
}
}
}
// Final transpose plus delayed restore of the panel R blocks that were
// temporarily overwritten while using the panel storage as explicit V.
__global__ void transpose_tiled_restore_v_kernel(const float* __restrict__ col,
float* __restrict__ row,
const float* __restrict__ saved,
int n,
int nb) {
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
__shared__ float tile[TILE_DIM][TILE_DIM + 1];
const int b = blockIdx.z;
const long long matrix_size = static_cast<long long>(n) * n;
const float* col_b = col + static_cast<long long>(b) * matrix_size;
float* row_b = row + static_cast<long long>(b) * matrix_size;
int x = blockIdx.x * TILE_DIM + threadIdx.x;
int y = blockIdx.y * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
if (x < n && y + j < n) {
tile[threadIdx.y + j][threadIdx.x] =
col_b[static_cast<long long>(y + j) * n + x];
}
}
__syncthreads();
x = blockIdx.y * TILE_DIM + threadIdx.x;
y = blockIdx.x * TILE_DIM + threadIdx.y;
for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
const int row_idx = y + j;
const int col_idx = x;
if (row_idx < n && col_idx < n) {
float value = tile[threadIdx.x][threadIdx.y + j];
const int panel_k = (col_idx / nb) * nb;
const int r = row_idx - panel_k;
const int s = col_idx - panel_k;
if (r >= 0 && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb;
value = saved[static_cast<long long>(b) * save_stride +
static_cast<long long>(panel_k) * nb +
r + static_cast<long long>(s) * nb];
}
row_b[static_cast<long long>(row_idx) * n + col_idx] = value;
}
}
}
// Build the triangular WY T factor from G = V.T @ V and tau. This path is kept
// for n512, where cuBLAS handles the following T.T @ W better than the fused
// custom kernel.
__global__ void build_t_from_gram_kernel(const float* __restrict__ gram,
float* __restrict__ t,
const float* __restrict__ tau,
int n,
int k,
int ib,
long long stride_small) {
const int b = blockIdx.x;
const float* g = gram + static_cast<long long>(b) * stride_small;
float* tb = t + static_cast<long long>(b) * stride_small;
const float* tau_b = tau + static_cast<long long>(b) * n + k;
for (int idx = 0; idx < ib * ib; ++idx) {
tb[idx] = 0.0f;
}
float work[64];
for (int i = 0; i < ib; ++i) {
const float tau_i = tau_b[i];
if (tau_i == 0.0f) {
tb[i + static_cast<long long>(i) * ib] = 0.0f;
continue;
}
for (int j = 0; j < i; ++j) {
work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
}
for (int row = 0; row < i; ++row) {
float sum = 0.0f;
for (int col = row; col < i; ++col) {
sum = fmaf(tb[row + static_cast<long long>(col) * ib], work[col], sum);
}
tb[row + static_cast<long long>(i) * ib] = sum;
}
tb[i + static_cast<long long>(i) * ib] = tau_i;
}
}
// n2048 has a small batch; a custom dot kernel wins over cuBLAS for the tiny
// per-panel Gram matrix.
__global__ void build_gram_strided_small_kernel(const float* __restrict__ v,
float* __restrict__ gram,
int m,
int ib,
int ldv,
long long stride_v,
long long stride_small) {
const int b = blockIdx.x;
const int row = blockIdx.y;
const int col = blockIdx.z;
const int tid = threadIdx.x;
__shared__ float red[PANEL_THREADS];
const float* vb = v + static_cast<long long>(b) * stride_v;
float local = 0.0f;
for (int r = tid; r < m; r += blockDim.x) {
local = fmaf(vb[r + static_cast<long long>(row) * ldv],
vb[r + static_cast<long long>(col) * ldv],
local);
}
red[tid] = local;
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
red[tid] += red[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
gram[static_cast<long long>(b) * stride_small + row + static_cast<long long>(col) * ib] = red[0];
}
}
// Shared-memory panel factorization. One CTA owns one matrix in the batch; each
// warp updates one target panel column after a reflector is formed.
__global__ void panel_factor_colmajor_shared_warpcols_kernel(float* __restrict__ col,
float* __restrict__ tau,
float* __restrict__ saved,
int n,
int k0,
int ib,
int nb_storage,
int make_explicit_v) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int warp_count = blockDim.x >> 5;
const int m = n - k0;
const long long matrix_size = static_cast<long long>(n) * n;
float* mat = col + static_cast<long long>(b) * matrix_size;
float* tau_b = tau + static_cast<long long>(b) * n;
extern __shared__ float panel[];
__shared__ float red[PANEL_THREADS];
__shared__ float tau_s;
__shared__ float scale_s;
const int panel_elems = m * ib;
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
panel[idx] = mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n];
}
__syncthreads();
for (int s = 0; s < ib; ++s) {
float local = 0.0f;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
const float x = panel[r + static_cast<long long>(s) * m];
local = fmaf(x, x, local);
}
unsigned mask = 0xffffffffu;
for (int offset = 16; offset > 0; offset >>= 1) {
local += __shfl_down_sync(mask, local, offset);
}
if (lane == 0) {
red[warp] = local;
}
__syncthreads();
if (warp == 0) {
float warp_sum = (lane < warp_count) ? red[lane] : 0.0f;
for (int offset = 16; offset > 0; offset >>= 1) {
warp_sum += __shfl_down_sync(mask, warp_sum, offset);
}
if (lane == 0) {
red[0] = warp_sum;
}
}
__syncthreads();
if (tid == 0) {
const long long diag = static_cast<long long>(s) + static_cast<long long>(s) * m;
const float alpha = panel[diag];
const float xnorm2 = red[0];
if (xnorm2 == 0.0f) {
tau_s = 0.0f;
scale_s = 0.0f;
tau_b[k0 + s] = 0.0f;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
const float beta = (alpha >= 0.0f) ? -norm : norm;
const float tau_val = (beta - alpha) / beta;
tau_s = tau_val;
scale_s = 1.0f / (alpha - beta);
panel[diag] = beta;
tau_b[k0 + s] = tau_val;
}
}
__syncthreads();
if (tau_s != 0.0f) {
const float scale = scale_s;
for (int r = s + 1 + tid; r < m; r += blockDim.x) {
panel[r + static_cast<long long>(s) * m] *= scale;
}
}
__syncthreads();
const int jj = s + 1 + warp;
if (jj < ib) {
float dot = 0.0f;
for (int r = s + 1 + lane; r < m; r += 32) {
dot = fmaf(panel[r + static_cast<long long>(s) * m],
panel[r + static_cast<long long>(jj) * m],
dot);
}
for (int offset = 16; offset > 0; offset >>= 1) {
dot += __shfl_down_sync(mask, dot, offset);
}
float w = 0.0f;
if (lane == 0) {
const float full_dot = panel[s + static_cast<long long>(jj) * m] + dot;
w = tau_s * full_dot;
panel[s + static_cast<long long>(jj) * m] -= w;
}
w = __shfl_sync(mask, w, 0);
for (int r = s + 1 + lane; r < m; r += 32) {
const long long idx = static_cast<long long>(r) + static_cast<long long>(jj) * m;
panel[idx] = fmaf(-panel[r + static_cast<long long>(s) * m], w, panel[idx]);
}
}
__syncthreads();
}
for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
const int r = idx % m;
const int s = idx / m;
float value = panel[idx];
if (make_explicit_v && r <= s) {
const long long save_stride = static_cast<long long>(n) * nb_storage;
saved[static_cast<long long>(b) * save_stride +
static_cast<long long>(k0) * nb_storage +
r + static_cast<long long>(s) * nb_storage] = value;
value = (r == s) ? 1.0f : 0.0f;
}
mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n] = value;
}
}
// Fused small operation for modes where building T and then launching cuBLAS for
// T.T @ W costs more than doing both in one custom kernel.
__global__ void build_t_apply_transpose_fused_kernel(const float* __restrict__ gram,
const float* __restrict__ tau,
const float* __restrict__ w,
float* __restrict__ w2,
int n,
int k,
int ib,
int trailing,
long long stride_small,
long long stride_w) {
constexpr int COL_TILE = 64;
const int b = blockIdx.x;
const int tile_start = blockIdx.y * COL_TILE;
const int tid = threadIdx.x;
const int cols_left = trailing - tile_start;
const int cols = (cols_left < COL_TILE) ? cols_left : COL_TILE;
if (cols <= 0) {
return;
}
__shared__ float ts[64 * 64];
if (tid < ib * ib) {
ts[tid] = 0.0f;
}
__syncthreads();
if (tid == 0) {
const float* g = gram + static_cast<long long>(b) * stride_small;
const float* tau_b = tau + static_cast<long long>(b) * n + k;
float work[64];
for (int i = 0; i < ib; ++i) {
const float tau_i = tau_b[i];
if (tau_i == 0.0f) {
ts[i + static_cast<long long>(i) * ib] = 0.0f;
continue;
}
for (int j = 0; j < i; ++j) {
work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
}
for (int row = 0; row < i; ++row) {
float sum = 0.0f;
for (int col = row; col < i; ++col) {
sum = fmaf(ts[row + static_cast<long long>(col) * ib], work[col], sum);
}
ts[row + static_cast<long long>(i) * ib] = sum;
}
ts[i + static_cast<long long>(i) * ib] = tau_i;
}
}
__syncthreads();
const float* wb = w + static_cast<long long>(b) * stride_w;
float* w2b = w2 + static_cast<long long>(b) * stride_w;
const int tile_outputs = ib * cols;
for (int idx = tid; idx < tile_outputs; idx += blockDim.x) {
const int p = idx % ib;
const int j = tile_start + idx / ib;
float sum = 0.0f;
for (int q = 0; q < ib; ++q) {
sum = fmaf(ts[q + static_cast<long long>(p) * ib],
wb[q + static_cast<long long>(j) * ib],
sum);
}
w2b[p + static_cast<long long>(j) * ib] = sum;
}
}
// Human-editable replacement for the n512 cuBLAS W = V.T @ C call.
//
// cuBLAS sees W = V.T @ C as a strided-batched column-major SGEMM:
// V : m x ib, lda = n, stride = n*n
// C : m x trailing, ldc = n, stride = n*n
// W : ib x trailing, ldw = ib, stride = ib*trailing
//
// The production library uses a tiled SIMT GEMM. This starter version computes
// one tile of output columns per CTA and loops over all ib panel rows inside the
// CTA. It is intentionally regular and boring, so it can be rewritten by hand.
__global__ void n512_vt_c_tile_kernel(const float* __restrict__ v,
const float* __restrict__ c,
float* __restrict__ w,
int n,
int m,
int ib,
int trailing,
long long matrix_stride,
long long w_stride) {
constexpr int COL_TILE = 8;
constexpr int THREADS = 256;
__shared__ float red[COL_TILE][THREADS];
const int b = blockIdx.x;
const int tile_col = blockIdx.y * COL_TILE;
const int tid = threadIdx.x;
const float* vb = v + static_cast<long long>(b) * matrix_stride;
const float* cb = c + static_cast<long long>(b) * matrix_stride;
float* wb = w + static_cast<long long>(b) * w_stride;
for (int p = 0; p < ib; ++p) {
#pragma unroll
for (int tc = 0; tc < COL_TILE; ++tc) {
const int col = tile_col + tc;
float sum = 0.0f;
if (col < trailing) {
for (int r = tid; r < m; r += blockDim.x) {
sum = fmaf(vb[r + static_cast<long long>(p) * n],
cb[r + static_cast<long long>(col) * n],
sum);
}
}
red[tc][tid] = sum;
}
__syncthreads();
for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
if (tid < stride) {
#pragma unroll
for (int tc = 0; tc < COL_TILE; ++tc) {
red[tc][tid] += red[tc][tid + stride];
}
}
__syncthreads();
}
if (tid == 0) {
#pragma unroll
for (int tc = 0; tc < COL_TILE; ++tc) {
const int col = tile_col + tc;
if (col < trailing) {
wb[p + static_cast<long long>(col) * ib] = red[tc][0];
}
}
}
__syncthreads();
}
}
// Human-editable replacement for the n512 cuBLAS SGEMM update:
//
// C = C - V @ W2
//
// cuBLAS sees this as a strided-batched column-major SGEMM:
// A = V : m x ib, lda = n, stride = n*n
// B = W2 : ib x trailing, ldb = ib, stride = ib*trailing
// C = C : m x trailing, ldc = n, stride = n*n
//
// The library implementation is a tiled SIMT SGEMM: load a tile of V and W2,
// compute many C elements per CTA, and write back beta*C + alpha*A*B. This
// version is intentionally simple instead: one thread computes one C element
// and the small K dimension (ib <= 16 here) is unrolled. It is not trying to be
// cuBLAS-fast yet; it is a compact, correct place to start rewriting.
__global__ void n512_update_trailing_elementwise_kernel(const float* __restrict__ v,
const float* __restrict__ w2,
float* __restrict__ c,
int n,
int m,
int ib,
int trailing,
long long matrix_stride,
long long w_stride) {
const int col = blockIdx.x * blockDim.x + threadIdx.x;
const int row = blockIdx.y * blockDim.y + threadIdx.y;
const int b = blockIdx.z;
if (row >= m || col >= trailing) {
return;
}
const float* vb = v + static_cast<long long>(b) * matrix_stride;
const float* w2b = w2 + static_cast<long long>(b) * w_stride;
float* cb = c + static_cast<long long>(b) * matrix_stride;
float acc = 0.0f;
#pragma unroll
for (int q = 0; q < 16; ++q) {
if (q < ib) {
acc = fmaf(vb[row + static_cast<long long>(q) * n],
w2b[q + static_cast<long long>(col) * ib],
acc);
}
}
cb[row + static_cast<long long>(col) * n] -= acc;
}
} // namespace
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr32_kernel<128><<<batch, 128>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
batch,
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
const int shmem = n * n * static_cast<int>(sizeof(float));
cudaError_t err = cudaFuncSetAttribute(
qr_shared_kernel<256>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
shmem);
if (err != cudaSuccess) {
cudaGetLastError(); // Clear the error
qr_global_256(a, h, tau);
return;
}
qr_shared_kernel<256><<<batch, 256, shmem>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t launch_err = cudaGetLastError();
if (launch_err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(launch_err));
}
}
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr_global_kernel<256><<<batch, 256>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
qr_global_kernel<512><<<batch, 512>>>(
a.data_ptr<float>(),
h.data_ptr<float>(),
tau.data_ptr<float>(),
n);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int mode) {
const int batch = static_cast<int>(a.size(0));
const int n = static_cast<int>(a.size(1));
const long long matrix_size = static_cast<long long>(n) * n;
int device_id = 0;
check_cuda(cudaGetDevice(&device_id), "cudaGetDevice");
cudaDeviceProp prop;
check_cuda(cudaGetDeviceProperties(&prop, device_id), "cudaGetDeviceProperties");
int max_shmem = prop.sharedMemPerBlockOptin;
if (max_shmem <= 0) {
max_shmem = prop.sharedMemPerBlock;
}
while (nb > 1 && n * nb * static_cast<int>(sizeof(float)) > max_shmem) {
nb /= 2;
}
if (nb <= 0 || nb > 64) {
throw std::runtime_error("unsupported blocked-WY panel width");
}
if (mode != MODE_CUBLAS_GRAM_CUBLAS_T &&
mode != MODE_CUBLAS_GRAM_FUSED_T &&
mode != MODE_CUSTOM_GRAM_FUSED_T &&
mode != MODE_N512_CUSTOM_UPDATE) {
throw std::runtime_error("unsupported blocked-WY mode");
}
// Work in column-major layout so panel columns and cuBLAS operands are natural.
auto col = torch::empty_like(a);
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
const dim3 transpose_block(TILE_DIM, BLOCK_ROWS);
const dim3 transpose_grid((n + TILE_DIM - 1) / TILE_DIM,
(n + TILE_DIM - 1) / TILE_DIM,
batch);
transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
a.data_ptr<float>(),
col.data_ptr<float>(),
n);
check_cuda(cudaGetLastError(), "transpose_tiled_kernel row_to_col");
cublasHandle_t blas = get_blas_handle();
const int max_panel_shmem = n * nb * static_cast<int>(sizeof(float));
check_cuda(
cudaFuncSetAttribute(
panel_factor_colmajor_shared_warpcols_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shmem),
"cudaFuncSetAttribute panel_factor_colmajor_shared_warpcols_kernel");
auto opts = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
auto gram = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
auto t = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
auto w = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
auto w2 = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
// Panel storage is temporarily overwritten with explicit V. Save the upper
// triangular R blocks here and restore them during the final transpose.
auto saved = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
float* col_ptr = col.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
const float one = 1.0f;
const float zero = 0.0f;
const float minus_one = -1.0f;
for (int k = 0; k < n; k += nb) {
const int m = n - k;
const int ib = (m < nb) ? m : nb;
const int trailing = n - k - ib;
// Factor A[k:n, k:k+ib] in shared memory. The kernel leaves the panel
// as explicit V in-place and writes compact Householder tau.
const int panel_shmem = (n - k) * ib * static_cast<int>(sizeof(float));
panel_factor_colmajor_shared_warpcols_kernel<<<batch, PANEL_THREADS, panel_shmem>>>(
col_ptr,
tau_ptr,
saved.data_ptr<float>(),
n,
k,
ib,
nb,
1);
check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel");
const long long stride_small = static_cast<long long>(ib) * ib;
float* v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
const int ldv = n;
const long long stride_v_use = matrix_size;
// Gram is tiny (ib x ib) but repeated per panel and per batch. n2048
// prefers the custom strided dot kernel; n352/n512/n1024 use cuBLAS.
if (mode == MODE_CUSTOM_GRAM_FUSED_T) {
const dim3 gram_grid(batch, ib, ib);
build_gram_strided_small_kernel<<<gram_grid, PANEL_THREADS>>>(
v_ptr,
gram.data_ptr<float>(),
m,
ib,
ldv,
stride_v_use,
stride_small);
check_cuda(cudaGetLastError(), "build_gram_strided_small_kernel");
} else {
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
ib,
m,
&one,
v_ptr,
ldv,
stride_v_use,
v_ptr,
ldv,
stride_v_use,
&zero,
gram.data_ptr<float>(),
ib,
stride_small,
batch),
"cublasSgemmStridedBatched gram");
}
const bool fused_t_apply = (mode == MODE_CUBLAS_GRAM_FUSED_T ||
mode == MODE_CUSTOM_GRAM_FUSED_T);
if (!fused_t_apply) {
build_t_from_gram_kernel<<<batch, 1>>>(
gram.data_ptr<float>(),
t.data_ptr<float>(),
tau_ptr,
n,
k,
ib,
stride_small);
check_cuda(cudaGetLastError(), "build_t_from_gram_kernel");
}
if (trailing > 0) {
const long long stride_w = static_cast<long long>(ib) * trailing;
float* c_panel = col_ptr + static_cast<long long>(k) +
static_cast<long long>(k + ib) * n;
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
trailing,
m,
&one,
v_ptr,
ldv,
stride_v_use,
c_panel,
n,
matrix_size,
&zero,
w.data_ptr<float>(),
ib,
stride_w,
batch),
"cublasSgemmStridedBatched vt_c");
if (fused_t_apply) {
constexpr int fused_cols = 64;
const dim3 fused_grid(batch, (trailing + fused_cols - 1) / fused_cols);
build_t_apply_transpose_fused_kernel<<<fused_grid, 256>>>(
gram.data_ptr<float>(),
tau_ptr,
w.data_ptr<float>(),
w2.data_ptr<float>(),
n,
k,
ib,
trailing,
stride_small,
stride_w);
check_cuda(cudaGetLastError(), "build_t_apply_transpose_fused_kernel");
} else {
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_T,
CUBLAS_OP_N,
ib,
trailing,
ib,
&one,
t.data_ptr<float>(),
ib,
stride_small,
w.data_ptr<float>(),
ib,
stride_w,
&zero,
w2.data_ptr<float>(),
ib,
stride_w,
batch),
"cublasSgemmStridedBatched t_w");
}
if (mode == MODE_N512_CUSTOM_UPDATE) {
const dim3 update_block(16, 16);
const dim3 update_grid((trailing + update_block.x - 1) / update_block.x,
(m + update_block.y - 1) / update_block.y,
batch);
n512_update_trailing_elementwise_kernel<<<update_grid, update_block>>>(
v_ptr,
w2.data_ptr<float>(),
c_panel,
n,
m,
ib,
trailing,
matrix_size,
stride_w);
check_cuda(cudaGetLastError(), "n512_update_trailing_elementwise_kernel");
} else {
check_blas(
cublasSgemmStridedBatched(
blas,
CUBLAS_OP_N,
CUBLAS_OP_N,
m,
trailing,
ib,
&minus_one,
v_ptr,
ldv,
stride_v_use,
w2.data_ptr<float>(),
ib,
stride_w,
&one,
c_panel,
n,
matrix_size,
batch),
"cublasSgemmStridedBatched update");
}
}
}
transpose_tiled_restore_v_kernel<<<transpose_grid, transpose_block>>>(
col.data_ptr<float>(),
h.data_ptr<float>(),
saved.data_ptr<float>(),
n,
nb);
check_cuda(cudaGetLastError(), "transpose_tiled_restore_v_kernel");
}
void qr_n512_custom(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
const int batch = static_cast<int>(a.size(0));
constexpr int n = 512;
constexpr int nb = 16;
constexpr long long matrix_size = static_cast<long long>(n) * n;
// Same layout and panel QR as the cuBLAS-backed n512 path. The difference is
// that every trailing-matrix operation below is now an editable CUDA kernel.
auto col = torch::empty_like(a);
constexpr int TILE_DIM = 32;
constexpr int BLOCK_ROWS = 8;
const dim3 transpose_block(TILE_DIM, BLOCK_ROWS);
const dim3 transpose_grid((n + TILE_DIM - 1) / TILE_DIM,
(n + TILE_DIM - 1) / TILE_DIM,
batch);
transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
a.data_ptr<float>(),
col.data_ptr<float>(),
n);
check_cuda(cudaGetLastError(), "transpose_tiled_kernel row_to_col n512_custom");
const int max_panel_shmem = n * nb * static_cast<int>(sizeof(float));
check_cuda(
cudaFuncSetAttribute(
panel_factor_colmajor_shared_warpcols_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_panel_shmem),
"cudaFuncSetAttribute panel_factor_colmajor_shared_warpcols_kernel n512_custom");
auto opts = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
auto gram = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
auto w = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
auto w2 = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
auto saved = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
float* col_ptr = col.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
for (int k = 0; k < n; k += nb) {
const int m = n - k;
const int ib = (m < nb) ? m : nb;
const int trailing = n - k - ib;
const long long stride_small = static_cast<long long>(ib) * ib;
const int panel_shmem = m * ib * static_cast<int>(sizeof(float));
panel_factor_colmajor_shared_warpcols_kernel<<<batch, PANEL_THREADS, panel_shmem>>>(
col_ptr,
tau_ptr,
saved.data_ptr<float>(),
n,
k,
ib,
nb,
1);
check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel n512_custom");
float* v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
const dim3 gram_grid(batch, ib, ib);
build_gram_strided_small_kernel<<<gram_grid, PANEL_THREADS>>>(
v_ptr,
gram.data_ptr<float>(),
m,
ib,
n,
matrix_size,
stride_small);
check_cuda(cudaGetLastError(), "build_gram_strided_small_kernel n512_custom");
if (trailing > 0) {
const long long stride_w = static_cast<long long>(ib) * trailing;
float* c_panel = col_ptr + static_cast<long long>(k) +
static_cast<long long>(k + ib) * n;
constexpr int vt_cols = 8;
const dim3 vt_grid(batch, (trailing + vt_cols - 1) / vt_cols);
n512_vt_c_tile_kernel<<<vt_grid, 256>>>(
v_ptr,
c_panel,
w.data_ptr<float>(),
n,
m,
ib,
trailing,
matrix_size,
stride_w);
check_cuda(cudaGetLastError(), "n512_vt_c_tile_kernel");
constexpr int fused_cols = 64;
const dim3 fused_grid(batch, (trailing + fused_cols - 1) / fused_cols);
build_t_apply_transpose_fused_kernel<<<fused_grid, 256>>>(
gram.data_ptr<float>(),
tau_ptr,
w.data_ptr<float>(),
w2.data_ptr<float>(),
n,
k,
ib,
trailing,
stride_small,
stride_w);
check_cuda(cudaGetLastError(), "build_t_apply_transpose_fused_kernel n512_custom");
const dim3 update_block(16, 16);
const dim3 update_grid((trailing + update_block.x - 1) / update_block.x,
(m + update_block.y - 1) / update_block.y,
batch);
n512_update_trailing_elementwise_kernel<<<update_grid, update_block>>>(
v_ptr,
w2.data_ptr<float>(),
c_panel,
n,
m,
ib,
trailing,
matrix_size,
stride_w);
check_cuda(cudaGetLastError(), "n512_update_trailing_elementwise_kernel");
}
}
transpose_tiled_restore_v_kernel<<<transpose_grid, transpose_block>>>(
col.data_ptr<float>(),
h.data_ptr<float>(),
saved.data_ptr<float>(),
n,
nb);
check_cuda(cudaGetLastError(), "transpose_tiled_restore_v_kernel n512_custom");
}
"""
# Popcorn uploads only submission.py. During local development we compile the
# neighboring human.cu directly; this embedded copy keeps the uploaded file
# self-contained. Run studies/embed_human_cuda.py after editing human.cu.
# BEGIN HUMAN_CUDA_EMBEDDED
HUMAN_CUDA_EMBEDDED = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <stdexcept>
#define N 512
#define THREADS_PER_COL 32
#define BLOCK_SIZE 32
#define APPLY_THREADS 256
#define COLS_PER_TRAILING_BLOCK 32
#ifndef HUMAN_N512_IMPLEMENTED
#define HUMAN_N512_IMPLEMENTED 1
#endif
//=============================================================================
// Kernel 1: Panel Factorization
// Her panel sutunu icin: norm -> alpha/beta/tau/scale -> v scale -> panel-ici update
// Grid: <<<batch, 512>>>
//=============================================================================
__global__ void __launch_bounds__(512, 3) kernel_factor_panel(float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ T_workspace,
int k, int ib, int batch)
{
if (blockIdx.x >= batch) return;
int tid = threadIdx.x;
int base = blockIdx.x * N * N;
__shared__ float warp_sums[16]; // 512 / 32
__shared__ float v_col[N];
__shared__ float alpha;
__shared__ float beta;
__shared__ float scale;
__shared__ float tau_col;
__shared__ float T_shared[BLOCK_SIZE][BLOCK_SIZE + 1];
__shared__ float w_shared[BLOCK_SIZE];
// Initialize T_shared to 0
if (tid < BLOCK_SIZE * BLOCK_SIZE) {
int r = tid % BLOCK_SIZE;
int c = tid / BLOCK_SIZE;
T_shared[r][c] = 0.0f;
}
for (int j = 0; j < ib; ++j) {
int col = k + j;
v_col[tid] = h[base + col * N + tid];
//__syncthreads(); gerek var mi
// Norm reduction: once her warp kendi toplamini hesaplar.. i think this approach
// is not much necessary
float norm_sq = 0.0f;
if (tid >= col) {
float val = v_col[tid];
norm_sq = val * val;
}
for (int offset = 16; offset > 0; offset >>= 1) {
norm_sq += __shfl_down_sync(0xffffffff, norm_sq, offset);
}
int norm_lane = tid & 31;
int norm_warp = tid >> 5;
// 16 warp'in sonuclarini shared memory'ye yaz
if (norm_lane == 0) {
warp_sums[norm_warp] = norm_sq;
}
__syncthreads();
// ilk warp, 16 ara sonucu toplar
if (norm_warp == 0) {
norm_sq = norm_lane < 16 ? warp_sums[norm_lane] : 0.0f;
for (int offset = 16; offset > 0; offset >>= 1) {
norm_sq += __shfl_down_sync(0xffffffff, norm_sq, offset);
}
if (norm_lane == 0) {
warp_sums[0] = norm_sq;
}
}
// Compute scalars (thread 0 only)
if (tid == 0) {
alpha = v_col[col];
float norm = sqrtf(warp_sums[0]);
if (norm < 1e-20f) {
beta = alpha;
tau[blockIdx.x * N + col] = 0.0f;
tau_col = 0.0f;
scale = 0.0f;
} else {
beta = (alpha >= 0.0f) ? -norm : norm;
float tau_val = (beta - alpha) / beta;
tau[blockIdx.x * N + col] = tau_val;
tau_col = tau_val;
scale = 1.0f / (alpha - beta);
}
h[base + col * N + col] = beta;
v_col[col] = 1.0f;
}
__syncthreads();
// Scale Householder vector and write back
if (tid > col) {
v_col[tid] *= scale;
}
if (tid > col) {
//burada globale yazmak yerine başka biryerde yazilabilrimi
h[base + col * N + tid] = v_col[tid];
}
__syncthreads();
// Update remaining columns in the current panel
int local_tid = tid % THREADS_PER_COL;
int col_group = tid / THREADS_PER_COL;
int cols_per_block = blockDim.x / THREADS_PER_COL;
for (int panel_col = col + 1; panel_col < k + ib; panel_col += cols_per_block) {
int j_panel = panel_col + col_group;
float sum = 0.0f;
if (j_panel < k + ib) {
int row = (col & ~31) + local_tid;
// Yalnızca ilk, kısmi warp parçasında kontrol gerekli
if (row >= col) {
sum = fmaf(
v_col[row],
h[base + j_panel * N + row],
sum
);
}
// Bundan sonraki bütün parçalar tam ve hizalı
for (row += 32; row < N; row += 32) {
sum = fmaf(
v_col[row],
h[base + j_panel * N + row],
sum
);
}
}
// Warp reduction
for (int offset = THREADS_PER_COL / 2; offset > 0; offset /= 2) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
float dot = __shfl_sync(0xffffffff, sum, 0);
if (j_panel < k + ib) {
dot *= tau_col;
int row = (col & ~31) + local_tid;
// The first partial warp includes the diagonal. v_col[col] is
// already 1, so the same FMA also performs C[col] -= dot.
if (row >= col) {
h[base + j_panel * N + row] =
fmaf(-v_col[row], dot,
h[base + j_panel * N + row]);
}
// All subsequent warp accesses start on a 128-byte boundary.
for (row += THREADS_PER_COL; row < N;
row += THREADS_PER_COL) {
h[base + j_panel * N + row] =
fmaf(-v_col[row], dot,
h[base + j_panel * N + row]);
}
}
}
//__syncthreads();
// Compute w_shared[p] = -tau_col * dot(v_p, v_j) for p < j in parallel
int lane_id = tid % 16;
int group_id = tid / 16;
if (group_id < j) {
int p = group_id;
float sum = 0.0f;
int row = (col & ~15) + lane_id;
// The first partial half-warp includes v_j's unit diagonal.
if (row >= col) {
float val_p = h[base + (k + p) * N + row];
sum = fmaf(val_p, v_col[row], sum);
}
// Remaining half-warp accesses are aligned to 64-byte boundaries.
for (row += 16; row < N; row += 16) {
float val_p = h[base + (k + p) * N + row];
sum = fmaf(val_p, v_col[row], sum);
}
unsigned int mask = (tid % 32 < 16) ? 0x0000ffff : 0xffff0000;
for (int offset = 8; offset > 0; offset /= 2) {
sum += __shfl_down_sync(mask, sum, offset, 16);
}
if (lane_id == 0) {
w_shared[p] = -tau_col * sum;
}
}
__syncthreads();
// Compute T_shared[i][j]
if (tid < j) {
float sum_T = 0.0f;
for (int p = tid; p < j; ++p) {
sum_T = fmaf(T_shared[tid][p], w_shared[p], sum_T);
}
T_shared[tid][j] = sum_T;
}
if (tid == j) {
T_shared[j][j] = tau_col;
}
__syncthreads();
}
// Write T_shared to global workspace
for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
int r = idx % ib;
int c = idx / ib;
T_workspace[blockIdx.x * BLOCK_SIZE * BLOCK_SIZE + c * BLOCK_SIZE + r] = T_shared[r][c];
}
}
//=============================================================================
// Kernel 2: Load T + Apply to Trailing Matrix
// Grid: <<<batch * num_col_blocks, 512>>>
//=============================================================================
__global__ void __launch_bounds__(APPLY_THREADS) kernel_build_T_apply(float* __restrict__ h,
const float* __restrict__ tau,
const float* __restrict__ T_workspace,
int k, int ib, int batch,
int num_col_blocks)
{
int batch_id = blockIdx.x / num_col_blocks;
int col_block_id = blockIdx.x % num_col_blocks;
if (batch_id >= batch) return;
int tid = threadIdx.x;
int base = batch_id * N * N;
int trailingStart = k + ib;
__shared__ float T[BLOCK_SIZE][BLOCK_SIZE + 1];
constexpr int TILE_ROWS = 64;
__shared__ float V_shared[TILE_ROWS][BLOCK_SIZE + 1];
// Load T from global workspace
for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
int r = idx % ib;
int c = idx / ib;
T[r][c] = T_workspace[batch_id * BLOCK_SIZE * BLOCK_SIZE + c * BLOCK_SIZE + r];
}
// == Step 2-4: Apply block reflector to this block's trailing columns ==
int col_start = trailingStart + col_block_id * COLS_PER_TRAILING_BLOCK;
int col_end = col_start + COLS_PER_TRAILING_BLOCK;
if (col_end > N) col_end = N;
if (col_start >= N) return;
int col_group = tid / 32; // 0..7 (8 warps)
int local_tid = tid % 32; // lane within warp
// Warp processes up to 4 columns in parallel (unrolled)
int c0 = col_start + col_group + 0 * 8;
int c1 = col_start + col_group + 1 * 8;
int c2 = col_start + col_group + 2 * 8;
int c3 = col_start + col_group + 3 * 8;
float w_val0 = 0.0f;
float w_val1 = 0.0f;
float w_val2 = 0.0f;
float w_val3 = 0.0f;
// Pass 1: Dot Product (Compute W = V^T * C)
for (int tile_row_start = k; tile_row_start < N; tile_row_start += TILE_ROWS) {
__syncthreads();
for (int load_idx = tid; load_idx < TILE_ROWS * ib; load_idx += APPLY_THREADS) {
int r = load_idx % TILE_ROWS;
int c_V = load_idx / TILE_ROWS;
int global_row = tile_row_start + r;
float val = 0.0f;
if (global_row < N && c_V < ib) {
if (global_row == k + c_V) {
val = 1.0f;
} else if (global_row > k + c_V) {
val = h[base + (k + c_V) * N + global_row];
}
}
V_shared[r][c_V] = val;
}
__syncthreads();
if (local_tid < ib) {
int limit = min(TILE_ROWS, N - tile_row_start);
for (int r = 0; r < limit; ++r) {
int global_row = tile_row_start + r;
float v_val = V_shared[r][local_tid];
if (c0 < col_end) w_val0 = fmaf(v_val, h[base + c0 * N + global_row], w_val0);
if (c1 < col_end) w_val1 = fmaf(v_val, h[base + c1 * N + global_row], w_val1);
if (c2 < col_end) w_val2 = fmaf(v_val, h[base + c2 * N + global_row], w_val2);
if (c3 < col_end) w_val3 = fmaf(v_val, h[base + c3 * N + global_row], w_val3);
}
}
}
__syncthreads();
// Pass 2: Compute Y = T * W
float y_val0 = 0.0f;
float y_val1 = 0.0f;
float y_val2 = 0.0f;
float y_val3 = 0.0f;
for (int p = 0; p < ib; ++p) {
float wp0 = __shfl_sync(0xffffffff, w_val0, p);
float wp1 = __shfl_sync(0xffffffff, w_val1, p);
float wp2 = __shfl_sync(0xffffffff, w_val2, p);
float wp3 = __shfl_sync(0xffffffff, w_val3, p);
if (local_tid < ib && p <= local_tid) {
y_val0 = fmaf(T[p][local_tid], wp0, y_val0);
y_val1 = fmaf(T[p][local_tid], wp1, y_val1);
y_val2 = fmaf(T[p][local_tid], wp2, y_val2);
y_val3 = fmaf(T[p][local_tid], wp3, y_val3);
}
}
// Pass 3: Apply Update (C -= V * Y)
for (int tile_row_start = k; tile_row_start < N; tile_row_start += TILE_ROWS) {
__syncthreads();
for (int load_idx = tid; load_idx < TILE_ROWS * ib; load_idx += APPLY_THREADS) {
int r = load_idx % TILE_ROWS;
int c_V = load_idx / TILE_ROWS;
int global_row = tile_row_start + r;
float val = 0.0f;
if (global_row < N && c_V < ib) {
if (global_row == k + c_V) {
val = 1.0f;
} else if (global_row > k + c_V) {
val = h[base + (k + c_V) * N + global_row];
}
}
V_shared[r][c_V] = val;
}
__syncthreads();
int limit = min(TILE_ROWS, N - tile_row_start);
for (int r = local_tid; r < limit; r += 32) {
int global_row = tile_row_start + r;
float sum0 = 0.0f;
float sum1 = 0.0f;
float sum2 = 0.0f;
float sum3 = 0.0f;
for (int i = 0; i < ib; ++i) {
float v_val = V_shared[r][i];
float yi0 = __shfl_sync(0xffffffff, y_val0, i);
float yi1 = __shfl_sync(0xffffffff, y_val1, i);
float yi2 = __shfl_sync(0xffffffff, y_val2, i);
float yi3 = __shfl_sync(0xffffffff, y_val3, i);
sum0 = fmaf(v_val, yi0, sum0);
sum1 = fmaf(v_val, yi1, sum1);
sum2 = fmaf(v_val, yi2, sum2);
sum3 = fmaf(v_val, yi3, sum3);
}
if (c0 < col_end) h[base + c0 * N + global_row] -= sum0;
if (c1 < col_end) h[base + c1 * N + global_row] -= sum1;
if (c2 < col_end) h[base + c2 * N + global_row] -= sum2;
if (c3 < col_end) h[base + c3 * N + global_row] -= sum3;
}
}
}
//=============================================================================
// Host Function
//=============================================================================
void qr_n512_human(torch::Tensor a, torch::Tensor h, torch::Tensor tau)
{
TORCH_CHECK(a.is_cuda() && h.is_cuda() && tau.is_cuda(),
"human n512 expects CUDA tensors");
TORCH_CHECK(a.scalar_type() == torch::kFloat32 &&
h.scalar_type() == torch::kFloat32 &&
tau.scalar_type() == torch::kFloat32,
"human n512 expects FP32 tensors");
TORCH_CHECK(a.dim() == 3 && a.size(1) == N && a.size(2) == N,
"human n512 input must have shape [batch, 512, 512]");
TORCH_CHECK(h.sizes() == a.sizes(),
"human n512 H must match input shape");
TORCH_CHECK(tau.dim() == 2 &&
tau.size(0) == a.size(0) &&
tau.size(1) == N,
"human n512 tau must have shape [batch, 512]");
TORCH_CHECK(a.is_contiguous() && h.is_contiguous() && tau.is_contiguous(),
"human n512 expects contiguous tensors");
#if HUMAN_N512_IMPLEMENTED
constexpr int threads = 512;
int batch = static_cast<int>(a.size(0));
// Transpose input
torch::Tensor a_T = a.transpose(1, 2).contiguous();
torch::Tensor h_T = torch::empty_like(a_T);
// Copy a_T -> h_T (replaces the in-kernel copy)
h_T.copy_(a_T);
float* h_ptr = h_T.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
// Allocate workspace for T matrix
auto options = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
torch::Tensor work_T = torch::empty({batch, BLOCK_SIZE, BLOCK_SIZE}, options);
float* work_T_ptr = work_T.data_ptr<float>();
for (int k = 0; k < N; k += BLOCK_SIZE) {
int ib = (N - k < BLOCK_SIZE) ? (N - k) : BLOCK_SIZE;
int trailingCols = N - k - ib;
// Kernel 1: Panel factorization
kernel_factor_panel<<<batch, threads>>>(h_ptr, tau_ptr, work_T_ptr, k, ib, batch);
// Kernel 2: Build T + apply to trailing matrix (multi-block per batch)
if (trailingCols > 0) {
int num_col_blocks = (trailingCols + COLS_PER_TRAILING_BLOCK - 1)
/ COLS_PER_TRAILING_BLOCK;
kernel_build_T_apply<<<batch * num_col_blocks, APPLY_THREADS>>>(
h_ptr, tau_ptr, work_T_ptr, k, ib, batch, num_col_blocks);
}
}
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
// Transpose back
h.copy_(h_T.transpose(1, 2));
#else
constexpr int MODE_CUBLAS_GRAM_CUBLAS_T = 8;
qr_blocked_wy(a, h, tau, 16, MODE_CUBLAS_GRAM_CUBLAS_T);
#endif
}
"""
# END HUMAN_CUDA_EMBEDDED
def _load_human_cuda_source() -> str:
try:
human_path = Path(__file__).with_name("human_adaptive.cu")
except NameError:
return HUMAN_CUDA_EMBEDDED
if human_path.is_file():
return human_path.read_text(encoding="utf-8")
return HUMAN_CUDA_EMBEDDED
HUMAN_CUDA_SRC = _load_human_cuda_source()
import os
_module = load_inline(
name="qr_compact_householder_b200_v100_human_n512",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC, HUMAN_CUDA_SRC],
functions=[
"qr32",
"qr_shared_176",
"qr_global_256",
"qr_global_512",
"qr_blocked_wy",
"qr_n512_custom",
"qr_n512_human",
],
verbose=False,
extra_cuda_cflags=["-O3", "--use_fast_math"],
extra_ldflags=["cublas.lib"] if os.name == "nt" else ["-lcublas"],
)
MODE_CUBLAS_GRAM_CUBLAS_T = 8
MODE_CUBLAS_GRAM_FUSED_T = 9
MODE_CUSTOM_GRAM_FUSED_T = 10
MODE_N512_CUSTOM_UPDATE = 11
def _empty_output(data: input_t, batch: int, n: int) -> output_t:
return (
torch.empty_like(data),
torch.empty((batch, n), device=data.device, dtype=torch.float32),
)
def custom_kernel(data: input_t) -> output_t:
if data.is_cuda and data.dtype == torch.float32 and data.is_contiguous():
batch, n, m = data.shape
if m == n and n <= 32:
h, tau = _empty_output(data, batch, n)
_module.qr32(data, h, tau)
return h, tau
if m == n and n == 176:
h, tau = _empty_output(data, batch, n)
_module.qr_shared_176(data, h, tau)
return h, tau
if m == n and n <= 176:
h, tau = _empty_output(data, batch, n)
_module.qr_global_256(data, h, tau)
return h, tau
if m == n and n <= 512 and not (n == 352 or (n == 512 and batch >= 128)):
h, tau = _empty_output(data, batch, n)
_module.qr_global_512(data, h, tau)
return h, tau
if m == n and n == 512 and batch >= 128:
# human.cu owns this route. Its disabled stub safely falls back to
# the old cuBLAS blocked-WY implementation.
h, tau = _empty_output(data, batch, n)
_module.qr_n512_human(data, h, tau)
return h, tau
if m == n and n == 352:
h, tau = _empty_output(data, batch, n)
_module.qr_blocked_wy(data, h, tau, 16, MODE_CUBLAS_GRAM_FUSED_T)
return h, tau
if m == n and n == 1024:
h, tau = _empty_output(data, batch, n)
_module.qr_blocked_wy(data, h, tau, 16, MODE_CUBLAS_GRAM_FUSED_T)
return h, tau
if m == n and n == 2048:
h, tau = _empty_output(data, batch, n)
_module.qr_blocked_wy(data, h, tau, 16, MODE_CUSTOM_GRAM_FUSED_T)
return h, tau
return torch.geqrf(data)
scrolls · 1828 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