submission 838628
Dortamac · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 730 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-838628?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:f0355594d23bb1fc453e16920c8c12cd5f765a41cff249320b92b195dd4b9ec8
license declaredunknown
license concludedunknown
authorsDortamac
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ double scratch[THREADS];Kernel source
submission.py730 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_PANEL = 8
_PANEL_32 = 32
_PANEL_352 = 8
_PANEL_2048 = 4
_PANEL_4096 = 4
_SUPERPANEL_512 = 32
_SUPERPANEL_NS = ()
_PANEL_T_NS = (176,)
_USE_PRECOMPUTE_U_NS = ()
_BLOCKED_NS = (32, 176, 352, 512, 1024, 2048)
_NATIVE_PANEL_NS = (32, 176, 352, 512, 1024, 2048)
CPP_SRC = r"""
#include <torch/extension.h>
void geqrf_32_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_176_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_352_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_512_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_1024_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_2048_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_4096_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width);
void geqrf_176_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_352_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_512_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_1024_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void geqrf_2048_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width);
void pack_v_panel(torch::Tensor h, torch::Tensor v, int64_t k, int64_t width);
void make_t_panel(torch::Tensor v, torch::Tensor tau, torch::Tensor t, int64_t width);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <stdexcept>
#include <string>
namespace {
constexpr int PANEL = 64;
constexpr int MAX_PANEL = 128;
constexpr int THREADS = 256;
inline void check_cuda(cudaError_t status, const char* what) {
if (status != cudaSuccess) {
throw std::runtime_error(std::string(what) + ": " + cudaGetErrorString(status));
}
}
__device__ double shfl_down_double(double value, int offset) {
int2 words = *reinterpret_cast<int2*>(&value);
words.x = __shfl_down_sync(0xffffffffu, words.x, offset);
words.y = __shfl_down_sync(0xffffffffu, words.y, offset);
return *reinterpret_cast<double*>(&words);
}
__device__ double warp_sum(double value) {
for (int offset = 16; offset > 0; offset >>= 1) {
value += shfl_down_double(value, offset);
}
return value;
}
__device__ double block_sum(double value, double* scratch) {
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
value = warp_sum(value);
if (lane == 0) {
scratch[warp] = value;
}
__syncthreads();
value = 0.0;
const int num_warps = (blockDim.x + 31) >> 5;
if (warp == 0 && lane < num_warps) {
value = scratch[lane];
}
value = warp_sum(value);
if (tid == 0) {
scratch[0] = value;
}
__syncthreads();
return scratch[0];
}
template <int N, bool BUILD_T>
__global__ void geqrf_panel_kernel(float* __restrict__ h,
float* __restrict__ tau,
float* __restrict__ t_out,
int batch,
int k,
int width) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
if (b >= batch) {
return;
}
__shared__ double scratch[THREADS];
__shared__ float beta_s;
__shared__ float tau_s;
__shared__ float inv_s;
__shared__ float update_s;
float* a = h + static_cast<long long>(b) * N * N;
float* t = tau + static_cast<long long>(b) * N;
for (int jj = 0; jj < width; ++jj) {
const int col = k + jj;
double tail_local = 0.0;
for (int row = col + 1 + tid; row < N; row += blockDim.x) {
const float x = a[row * N + col];
tail_local += static_cast<double>(x) * static_cast<double>(x);
}
const double tail_norm_sq = block_sum(tail_local, scratch);
if (tid == 0) {
const float alpha = a[col * N + col];
if (tail_norm_sq == 0.0) {
beta_s = alpha;
tau_s = 0.0f;
inv_s = 0.0f;
} else {
const double norm = sqrt(static_cast<double>(alpha) * static_cast<double>(alpha) + tail_norm_sq);
const double beta = (alpha >= 0.0f) ? -norm : norm;
beta_s = static_cast<float>(beta);
tau_s = static_cast<float>((beta - static_cast<double>(alpha)) / beta);
inv_s = static_cast<float>(1.0 / (static_cast<double>(alpha) - beta));
}
t[col] = tau_s;
}
__syncthreads();
if (tau_s != 0.0f) {
for (int row = col + 1 + tid; row < N; row += blockDim.x) {
a[row * N + col] *= inv_s;
}
}
__syncthreads();
if (tau_s != 0.0f) {
for (int j2 = jj + 1; j2 < width; ++j2) {
const int update_col = k + j2;
double dot_local = (tid == 0) ? static_cast<double>(a[col * N + update_col]) : 0.0;
for (int row = col + 1 + tid; row < N; row += blockDim.x) {
dot_local += static_cast<double>(a[row * N + col]) *
static_cast<double>(a[row * N + update_col]);
}
const double dot = block_sum(dot_local, scratch);
if (tid == 0) {
update_s = tau_s * static_cast<float>(dot);
a[col * N + update_col] -= update_s;
}
__syncthreads();
for (int row = col + 1 + tid; row < N; row += blockDim.x) {
a[row * N + update_col] -= a[row * N + col] * update_s;
}
}
__syncthreads();
}
if (tid == 0) {
a[col * N + col] = beta_s;
}
__syncthreads();
}
if constexpr (BUILD_T) {
__shared__ float t_col[MAX_PANEL];
__shared__ float tau_j_s;
float* tb = t_out + static_cast<long long>(b) * width * width;
const int rows = N - k;
for (int idx = tid; idx < width * width; idx += blockDim.x) {
tb[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < width; ++j) {
if (tid == 0) {
tau_j_s = t[k + j];
}
__syncthreads();
if (j > 0 && tau_j_s != 0.0f) {
for (int i = 0; i < j; ++i) {
double local = 0.0;
for (int row = j + tid; row < rows; row += blockDim.x) {
float vi = 0.0f;
if (row == i) {
vi = 1.0f;
} else if (row > i) {
vi = a[(k + row) * N + (k + i)];
}
float vj = 0.0f;
if (row == j) {
vj = 1.0f;
} else if (row > j) {
vj = a[(k + row) * N + (k + j)];
}
local += static_cast<double>(vi) * static_cast<double>(vj);
}
const double dot = block_sum(local, scratch);
if (tid == 0) {
t_col[i] = -tau_j_s * static_cast<float>(dot);
}
}
__syncthreads();
for (int i = tid; i < j; i += blockDim.x) {
float acc = 0.0f;
for (int q = 0; q < j; ++q) {
acc += tb[i * width + q] * t_col[q];
}
tb[i * width + j] = acc;
}
__syncthreads();
}
if (tid == 0) {
tb[j * width + j] = tau_j_s;
}
__syncthreads();
}
}
}
__global__ void make_t_panel_kernel(const float* __restrict__ v,
const float* __restrict__ tau,
float* __restrict__ t,
int batch,
int rows,
int width,
long long tau_s0,
long long tau_s1) {
const int b = blockIdx.x;
const int tid = threadIdx.x;
if (b >= batch) {
return;
}
__shared__ double scratch[THREADS];
__shared__ float col[MAX_PANEL];
__shared__ float tau_j_s;
const float* vb = v + static_cast<long long>(b) * rows * width;
const float* taub = tau + static_cast<long long>(b) * tau_s0;
float* tb = t + static_cast<long long>(b) * width * width;
for (int idx = tid; idx < width * width; idx += blockDim.x) {
tb[idx] = 0.0f;
}
__syncthreads();
for (int j = 0; j < width; ++j) {
if (tid == 0) {
tau_j_s = taub[static_cast<long long>(j) * tau_s1];
}
__syncthreads();
if (j > 0 && tau_j_s != 0.0f) {
for (int i = 0; i < j; ++i) {
double local = 0.0;
for (int row = j + tid; row < rows; row += blockDim.x) {
local += static_cast<double>(vb[row * width + i]) *
static_cast<double>(vb[row * width + j]);
}
const double dot = block_sum(local, scratch);
if (tid == 0) {
col[i] = -tau_j_s * static_cast<float>(dot);
}
}
__syncthreads();
for (int i = tid; i < j; i += blockDim.x) {
float acc = 0.0f;
for (int q = 0; q < j; ++q) {
acc += tb[i * width + q] * col[q];
}
tb[i * width + j] = acc;
}
__syncthreads();
}
if (tid == 0) {
tb[j * width + j] = tau_j_s;
}
__syncthreads();
}
}
__global__ void pack_v_panel_kernel(const float* __restrict__ h,
float* __restrict__ v,
int batch,
int n,
int rows,
int k,
int width) {
const int b = blockIdx.z;
const int row = blockIdx.y * blockDim.y + threadIdx.y;
const int col = blockIdx.x * blockDim.x + threadIdx.x;
if (b >= batch || row >= rows || col >= width) {
return;
}
float value = 0.0f;
if (row == col) {
value = 1.0f;
} else if (row > col) {
value = h[(static_cast<long long>(b) * n + (k + row)) * n + (k + col)];
}
v[(static_cast<long long>(b) * rows + row) * width + col] = value;
}
} // namespace
template <int N>
void geqrf_panel(torch::Tensor h, torch::Tensor tau, torch::Tensor t_panel, int64_t k, int64_t width, const char* name) {
TORCH_CHECK(h.is_cuda() && tau.is_cuda(), name, " expects CUDA tensors");
TORCH_CHECK(h.scalar_type() == torch::kFloat32, "h must be float32");
TORCH_CHECK(tau.scalar_type() == torch::kFloat32, "tau must be float32");
TORCH_CHECK(h.is_contiguous(), "h must be contiguous");
TORCH_CHECK(tau.is_contiguous(), "tau must be contiguous");
TORCH_CHECK(h.dim() == 3 && h.size(1) == N && h.size(2) == N, "h has wrong shape");
TORCH_CHECK(tau.dim() == 2 && tau.size(1) == N, "tau has wrong shape");
float* t_ptr = nullptr;
if (t_panel.defined()) {
TORCH_CHECK(t_panel.is_cuda(), name, " t expects CUDA tensor");
TORCH_CHECK(t_panel.scalar_type() == torch::kFloat32, "t must be float32");
TORCH_CHECK(t_panel.is_contiguous(), "t must be contiguous");
TORCH_CHECK(t_panel.dim() == 3, "t must be batch x width x width");
t_ptr = t_panel.data_ptr<float>();
}
const int kk = static_cast<int>(k);
const int ww = static_cast<int>(width);
TORCH_CHECK(ww > 0 && ww <= MAX_PANEL, "invalid panel width");
TORCH_CHECK(kk >= 0 && kk + ww <= N, "invalid panel offset");
const int batch = static_cast<int>(h.size(0));
if (t_panel.defined()) {
TORCH_CHECK(t_panel.size(0) == batch && t_panel.size(1) == ww && t_panel.size(2) == ww, "t has wrong shape");
}
if (batch == 0) {
return;
}
if (t_panel.defined()) {
geqrf_panel_kernel<N, true><<<batch, THREADS>>>(h.data_ptr<float>(), tau.data_ptr<float>(), t_ptr, batch, kk, ww);
} else {
geqrf_panel_kernel<N, false><<<batch, THREADS>>>(h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, batch, kk, ww);
}
check_cuda(cudaGetLastError(), name);
}
void geqrf_32_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<32>(h, tau, torch::Tensor(), k, width, "geqrf_32_panel");
}
void geqrf_176_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<176>(h, tau, torch::Tensor(), k, width, "geqrf_176_panel");
}
void geqrf_352_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<352>(h, tau, torch::Tensor(), k, width, "geqrf_352_panel");
}
void geqrf_512_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<512>(h, tau, torch::Tensor(), k, width, "geqrf_512_panel");
}
void geqrf_1024_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<1024>(h, tau, torch::Tensor(), k, width, "geqrf_1024_panel");
}
void geqrf_2048_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<2048>(h, tau, torch::Tensor(), k, width, "geqrf_2048_panel");
}
void geqrf_4096_panel(torch::Tensor h, torch::Tensor tau, int64_t k, int64_t width) {
geqrf_panel<4096>(h, tau, torch::Tensor(), k, width, "geqrf_4096_panel");
}
void geqrf_176_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
geqrf_panel<176>(h, tau, t, k, width, "geqrf_176_panel_t");
}
void geqrf_352_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
geqrf_panel<352>(h, tau, t, k, width, "geqrf_352_panel_t");
}
void geqrf_512_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
geqrf_panel<512>(h, tau, t, k, width, "geqrf_512_panel_t");
}
void geqrf_1024_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
geqrf_panel<1024>(h, tau, t, k, width, "geqrf_1024_panel_t");
}
void geqrf_2048_panel_t(torch::Tensor h, torch::Tensor tau, torch::Tensor t, int64_t k, int64_t width) {
geqrf_panel<2048>(h, tau, t, k, width, "geqrf_2048_panel_t");
}
void pack_v_panel(torch::Tensor h, torch::Tensor v, int64_t k, int64_t width) {
TORCH_CHECK(h.is_cuda() && v.is_cuda(), "pack_v_panel expects CUDA tensors");
TORCH_CHECK(h.scalar_type() == torch::kFloat32, "h must be float32");
TORCH_CHECK(v.scalar_type() == torch::kFloat32, "v must be float32");
TORCH_CHECK(h.is_contiguous(), "h must be contiguous");
TORCH_CHECK(v.is_contiguous(), "v must be contiguous");
TORCH_CHECK(h.dim() == 3 && h.size(1) == h.size(2), "h must be batch x n x n");
TORCH_CHECK(v.dim() == 3, "v must be batch x rows x width");
const int batch = static_cast<int>(h.size(0));
const int n = static_cast<int>(h.size(1));
const int rows = static_cast<int>(v.size(1));
const int kk = static_cast<int>(k);
const int ww = static_cast<int>(width);
TORCH_CHECK(v.size(0) == batch, "v batch mismatch");
TORCH_CHECK(ww > 0 && ww <= MAX_PANEL && v.size(2) == ww, "invalid pack_v_panel width");
TORCH_CHECK(kk >= 0 && kk + ww <= n && rows == n - kk, "invalid pack_v_panel shape");
if (batch == 0) {
return;
}
const dim3 block(16, 16, 1);
const dim3 grid((ww + block.x - 1) / block.x, (rows + block.y - 1) / block.y, batch);
pack_v_panel_kernel<<<grid, block>>>(h.data_ptr<float>(), v.data_ptr<float>(), batch, n, rows, kk, ww);
check_cuda(cudaGetLastError(), "pack_v_panel");
}
void make_t_panel(torch::Tensor v, torch::Tensor tau, torch::Tensor t, int64_t width) {
TORCH_CHECK(v.is_cuda() && tau.is_cuda() && t.is_cuda(), "make_t_panel expects CUDA tensors");
TORCH_CHECK(v.scalar_type() == torch::kFloat32, "v must be float32");
TORCH_CHECK(tau.scalar_type() == torch::kFloat32, "tau must be float32");
TORCH_CHECK(t.scalar_type() == torch::kFloat32, "t must be float32");
TORCH_CHECK(v.is_contiguous(), "v must be contiguous");
TORCH_CHECK(t.is_contiguous(), "t must be contiguous");
TORCH_CHECK(v.dim() == 3, "v must be batch x rows x width");
TORCH_CHECK(tau.dim() == 2, "tau must be batch x width");
TORCH_CHECK(t.dim() == 3, "t must be batch x width x width");
const int batch = static_cast<int>(v.size(0));
const int rows = static_cast<int>(v.size(1));
const int ww = static_cast<int>(width);
TORCH_CHECK(ww > 0 && ww <= MAX_PANEL, "invalid make_t_panel width");
TORCH_CHECK(v.size(2) == ww && tau.size(1) == ww && t.size(1) == ww && t.size(2) == ww, "make_t_panel shape mismatch");
if (batch == 0) {
return;
}
make_t_panel_kernel<<<batch, THREADS>>>(
v.data_ptr<float>(),
tau.data_ptr<float>(),
t.data_ptr<float>(),
batch,
rows,
ww,
static_cast<long long>(tau.stride(0)),
static_cast<long long>(tau.stride(1)));
check_cuda(cudaGetLastError(), "make_t_panel");
}
"""
if torch.cuda.is_available():
_native_module = load_inline(
name="qr_v2_native_panel_p8_n32_n4096_v1",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=[
"geqrf_32_panel",
"geqrf_176_panel",
"geqrf_352_panel",
"geqrf_512_panel",
"geqrf_1024_panel",
"geqrf_2048_panel",
"geqrf_4096_panel",
"geqrf_176_panel_t",
"geqrf_352_panel_t",
"geqrf_512_panel_t",
"geqrf_1024_panel_t",
"geqrf_2048_panel_t",
"pack_v_panel",
"make_t_panel",
],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3"],
verbose=False,
)
else:
_native_module = None
def _make_v_inplace(panel_h: torch.Tensor, width: int) -> torch.Tensor:
v = panel_h[:, :, :width]
v.tril_(-1)
diag = torch.arange(width, device=panel_h.device)
v[:, diag, diag] = 1.0
return v
def _make_t(
v: torch.Tensor,
tau: torch.Tensor,
width: int,
use_native_t: bool,
) -> torch.Tensor:
batch = v.shape[0]
if use_native_t and width <= 128 and v.is_cuda and _native_module is not None:
t = torch.empty((batch, width, width), device=v.device, dtype=v.dtype)
_native_module.make_t_panel(v, tau, t, width)
return t
t = torch.zeros((batch, width, width), device=v.device, dtype=v.dtype)
for j in range(width):
tau_j = tau[:, j]
if j > 0:
col = -tau_j[:, None] * torch.bmm(
v[:, j:, :j].transpose(1, 2),
v[:, j:, j : j + 1],
).squeeze(-1)
t[:, :j, j] = torch.bmm(t[:, :j, :j], col.unsqueeze(-1)).squeeze(-1)
t[:, j, j] = tau_j
return t
def _panel_for_n(n: int) -> int:
if n == 32:
return _PANEL_32
if n == 2048:
return _PANEL_2048
if n == 4096:
return _PANEL_4096
if n == 352:
return _PANEL_352
return _PANEL
def _apply_trailing_update(
trailing: torch.Tensor,
v: torch.Tensor,
t: torch.Tensor,
precompute_u: bool,
) -> None:
work = torch.bmm(v.transpose(1, 2), trailing)
if precompute_u:
u = torch.bmm(v, t.transpose(1, 2))
torch.baddbmm(trailing, u, work, beta=1.0, alpha=-1.0, out=trailing)
else:
work = torch.bmm(t.transpose(1, 2), work)
torch.baddbmm(trailing, v, work, beta=1.0, alpha=-1.0, out=trailing)
def _native_geqrf_panel(
h: torch.Tensor,
tau: torch.Tensor,
n: int,
k: int,
width: int,
t: torch.Tensor | None = None,
) -> None:
if n == 32:
_native_module.geqrf_32_panel(h, tau, k, width)
elif n == 176:
if t is not None:
_native_module.geqrf_176_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_176_panel(h, tau, k, width)
elif n == 352:
if t is not None:
_native_module.geqrf_352_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_352_panel(h, tau, k, width)
elif n == 512:
if t is not None:
_native_module.geqrf_512_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_512_panel(h, tau, k, width)
elif n == 1024:
if t is not None:
_native_module.geqrf_1024_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_1024_panel(h, tau, k, width)
elif n == 2048:
if t is not None:
_native_module.geqrf_2048_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_2048_panel(h, tau, k, width)
def _blocked_qr_superpanel_512(data: torch.Tensor) -> output_t:
h = data.clone()
batch, n, _ = h.shape
tau = torch.empty((batch, n), device=h.device, dtype=h.dtype)
panel = _PANEL
superpanel = _SUPERPANEL_512
precompute_u = n in _USE_PRECOMPUTE_U_NS
for k in range(0, n, superpanel):
block_width = min(superpanel, n - k)
block_end = k + block_width
for kk in range(k, block_end, panel):
width = min(panel, block_end - kk)
next_col = kk + width
_native_geqrf_panel(h, tau, n, kk, width)
if next_col >= block_end:
continue
v = torch.empty((batch, n - kk, width), device=h.device, dtype=h.dtype)
_native_module.pack_v_panel(h, v, kk, width)
panel_tau = tau[:, kk : kk + width]
t = _make_t(v, panel_tau, width, True)
local_trailing = h[:, kk:, next_col:block_end]
_apply_trailing_update(local_trailing, v, t, precompute_u)
if block_end >= n:
continue
v_block = torch.empty((batch, n - k, block_width), device=h.device, dtype=h.dtype)
_native_module.pack_v_panel(h, v_block, k, block_width)
t_block = _make_t(v_block, tau[:, k:block_end], block_width, True)
trailing = h[:, k:, block_end:]
_apply_trailing_update(trailing, v_block, t_block, precompute_u)
return h, tau
def _blocked_qr(data: torch.Tensor, panel: int) -> output_t:
h = data.clone()
batch, n, _ = h.shape
tau = torch.empty((batch, n), device=h.device, dtype=h.dtype)
panel = _panel_for_n(n)
use_native_panel = n in _NATIVE_PANEL_NS and _native_module is not None
use_native_t = use_native_panel
precompute_u = n in _USE_PRECOMPUTE_U_NS
for k in range(0, n, panel):
width = min(panel, n - k)
next_col = k + width
has_trailing = next_col < n
use_panel_t = use_native_panel and n in _PANEL_T_NS and has_trailing
t = None
if use_native_panel:
if use_panel_t:
t = torch.empty((batch, width, width), device=h.device, dtype=h.dtype)
if n == 32:
_native_module.geqrf_32_panel(h, tau, k, width)
elif n == 176:
if use_panel_t:
_native_module.geqrf_176_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_176_panel(h, tau, k, width)
elif n == 352:
if use_panel_t:
_native_module.geqrf_352_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_352_panel(h, tau, k, width)
elif n == 512:
if use_panel_t:
_native_module.geqrf_512_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_512_panel(h, tau, k, width)
elif n == 1024:
if use_panel_t:
_native_module.geqrf_1024_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_1024_panel(h, tau, k, width)
elif n == 2048:
if use_panel_t:
_native_module.geqrf_2048_panel_t(h, tau, t, k, width)
else:
_native_module.geqrf_2048_panel(h, tau, k, width)
elif n == 4096:
_native_module.geqrf_4096_panel(h, tau, k, width)
panel_tau = tau[:, k : k + width]
panel_h = None
else:
panel_h, panel_tau = torch.geqrf(h[:, k:, k : k + width].contiguous())
h[:, k:, k : k + width] = panel_h
tau[:, k : k + width] = panel_tau
if not has_trailing:
continue
if use_native_panel:
v = torch.empty((batch, n - k, width), device=h.device, dtype=h.dtype)
_native_module.pack_v_panel(h, v, k, width)
if t is None:
t = _make_t(v, panel_tau, width, use_native_t)
else:
v = _make_v_inplace(panel_h, width)
t = _make_t(v, panel_tau, width, use_native_t)
trailing = h[:, k:, next_col:]
_apply_trailing_update(trailing, v, t, precompute_u)
return h, tau
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] in _BLOCKED_NS
):
if data.shape[-1] in _SUPERPANEL_NS and _native_module is not None:
return _blocked_qr_superpanel_512(
data.contiguous() if not data.is_contiguous() else data,
)
return _blocked_qr(
data.contiguous() if not data.is_contiguous() else data,
_panel_for_n(data.shape[-1]),
)
return torch.geqrf(data)
scrolls · 730 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