submission 798644
switchtovim · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 984 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-798644?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:cfd33e46ce59a58098e9f6e2805dc6d35dbb3ac842decf1d1d9e9c566fe1918c
license declaredunknown
license concludedunknown
authorsswitchtovim
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):shared-memory
extern __shared__ float s[]; // m x bb, column-major: s[c*m + r], m = n-ktile-m = 128
def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):Kernel source
submission.py984 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
try:
from task import input_t, output_t
except ImportError:
input_t = torch.Tensor
output_t = tuple[torch.Tensor, torch.Tensor]
# ----------------------------------------------------------------------------
# Shared-memory-resident panel factorization (CUDA, load_inline). One thread
# block per matrix factors the m x bb panel H[k:n, k:k+bb] entirely in shared
# memory: the panel is read from global ONCE (column-major) rather than re-read
# per column as the Triton panel does. Householder vectors / beta / tau written
# back in place. Used where the panel fits B200 smem (n=512: 512x64x4 = 128KB).
# Each trailing-column rank-1 update is handled by one warp (shuffle-reduced dot)
# so all 256 threads stay busy through the dominant in-panel apply.
# ----------------------------------------------------------------------------
_PANEL_CUDA = r"""
#include <ATen/cuda/CUDAContext.h>
// Single-block smem-resident panel. In one launch it (1) factors the m x bb panel,
// (2) writes the in-place H result (beta on diag, v below), (3) emits the dense
// unit-lower V (B,n,bcap) and (4) the bb x bb WY T-factor (B,bcap,bcap) — the latter
// two were previously separate Python ops (_extract_V + _form_T = bmm + triangular
// solve). T via closed form T = inv(diag(1/tau) + striu(V^T V, 1)) with V resident.
__global__ void qr_panel_k(float* __restrict__ H, float* __restrict__ tau,
float* __restrict__ Vd, float* __restrict__ Tout,
int n, int k, int bb, int bcap, int do_wy) {
extern __shared__ float s[]; // m x bb, column-major: s[c*m + r], m = n-k
int m = n - k;
int t = threadIdx.x, nt = blockDim.x;
int lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
float* Gm = s + (size_t)m * bb; // bb x bb scratch (G then T-inverse)
__shared__ float red[32]; // up to 32 warps (nt<=1024, hw threads/block cap)
__shared__ float sh_tau, sh_vscale;
__shared__ float taus[64]; // bb <= 64
float* Hb = H + (size_t)blockIdx.x * n * n;
for (int idx = t; idx < m * bb; idx += nt) {
int c = idx / m, r = idx - c * m;
s[idx] = Hb[(size_t)(k + r) * n + (k + c)];
}
__syncthreads();
for (int c = 0; c < bb; c++) {
float* sc = s + (size_t)c * m;
float loc = 0.f;
for (int r = c + 1 + t; r < m; r += nt) { float v = sc[r]; loc += v * v; }
for (int o = 16; o > 0; o >>= 1) loc += __shfl_down_sync(0xffffffff, loc, o);
if (lane == 0) red[warp] = loc;
__syncthreads();
if (t == 0) {
float sigma = 0.f;
for (int i = 0; i < nwarps; i++) sigma += red[i];
float alpha = sc[c], norm = sqrtf(alpha * alpha + sigma);
float beta = (alpha >= 0.f) ? -norm : norm;
bool nz = sigma > 0.f;
sh_tau = nz ? (beta - alpha) / beta : 0.f;
sh_vscale = nz ? 1.f / (alpha - beta) : 0.f;
sc[c] = nz ? beta : alpha;
if (do_wy) taus[c] = sh_tau; // taus[64] only read on do_wy path; guard lets bb>64 (fused whole-matrix factor)
tau[(size_t)blockIdx.x * n + (k + c)] = sh_tau;
}
__syncthreads();
float tj = sh_tau, vs = sh_vscale;
if (tj != 0.f) {
for (int r = c + 1 + t; r < m; r += nt) sc[r] *= vs;
__syncthreads();
for (int cc = c + 1 + warp; cc < bb; cc += nwarps) {
float* scc = s + (size_t)cc * m;
float w = 0.f;
for (int r = c + 1 + lane; r < m; r += 32) w += sc[r] * scc[r];
for (int o = 16; o > 0; o >>= 1) w += __shfl_down_sync(0xffffffff, w, o);
w = __shfl_sync(0xffffffff, w, 0);
w = tj * (w + scc[c]);
if (lane == 0) scc[c] -= w;
for (int r = c + 1 + lane; r < m; r += 32) scc[r] -= sc[r] * w;
}
}
__syncthreads();
}
// write H in place (+ dense unit-lower V if WY factors requested)
float* Vb = Vd + (size_t)blockIdx.x * n * bcap;
for (int idx = t; idx < m * bb; idx += nt) {
int c = idx / m, r = idx - c * m;
float val = s[idx];
Hb[(size_t)(k + r) * n + (k + c)] = val;
if (do_wy) Vb[(size_t)r * bcap + c] = (r < c) ? 0.f : (r == c) ? 1.f : val;
}
if (!do_wy) return;
// G[i][j] = (V^T V)[i][j] for i<j (one warp per pair, shfl-reduced dot)
for (int idx = warp; idx < bb * bb; idx += nwarps) {
int i = idx / bb, j = idx - i * bb;
if (i < j) {
float* si = s + (size_t)i * m;
float* sj = s + (size_t)j * m;
float acc = (lane == 0) ? si[j] : 0.f; // r==j term: v_j[j]=1
for (int r = j + 1 + lane; r < m; r += 32) acc += si[r] * sj[r];
for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffff, acc, o);
if (lane == 0) Gm[(size_t)i * bb + j] = acc;
}
}
__syncthreads();
// T = inv(M), M upper-tri with M[i][i]=1/tau_i, M[i][j>i]=G[i][j]. Solve M X = I.
// Columns of X are independent -> one thread per column (smem back-substitution),
// O(bb^2) depth instead of an O(bb^3) single-thread serial tail.
float* Tsm = Gm + (size_t)bb * bb;
for (int j = t; j < bb; j += nt) {
Tsm[(size_t)j * bb + j] = taus[j]; // 1/M[j][j] = tau_j (0 -> ~identity)
for (int i = j - 1; i >= 0; i--) {
float sum = 0.f;
for (int l = i + 1; l <= j; l++)
sum += Gm[(size_t)i * bb + l] * Tsm[(size_t)l * bb + j];
Tsm[(size_t)i * bb + j] = -taus[i] * sum;
}
for (int i = j + 1; i < bb; i++) Tsm[(size_t)i * bb + j] = 0.f;
}
__syncthreads();
float* Tb = Tout + (size_t)blockIdx.x * bcap * bcap;
for (int idx = t; idx < bb * bb; idx += nt) {
int i = idx / bb, j = idx - i * bb;
Tb[(size_t)i * bcap + j] = Tsm[(size_t)i * bb + j];
}
}
// ---- Multi-block cooperative panel (tiny-batch huge-n). G blocks per matrix each
// own a BM-row slice resident in smem; cross-block reductions via global scratch +
// a hand-rolled device barrier (same structure as the Triton mb path, + residency).
__device__ __forceinline__ void grid_bar(int* counter, volatile int* sense_arr,
int cb, int G, int* msh, int t) {
__syncthreads();
if (t == 0) {
int ms = *msh ^ 1; *msh = ms;
__threadfence();
int old = atomicAdd(&counter[cb], 1);
if (old == G - 1) {
atomicExch(&counter[cb], 0);
__threadfence();
atomicExch((int*)&sense_arr[cb], ms);
} else {
while (sense_arr[cb] != ms) { }
}
}
__syncthreads();
}
__global__ void qr_panel_mb_k(float* __restrict__ H, float* __restrict__ tau,
float* __restrict__ Vd,
int* counter, int* sense_arr, float* sigp,
float* alphasc, float* wp,
int n, int k, int bb, int G, int BM, int bcap, int do_wy) {
extern __shared__ float s[]; // bb x BM, column-major: s[c*BM + rl]
int t = threadIdx.x, nt = blockDim.x;
int lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
int mid = blockIdx.x / G, g = blockIdx.x % G, m = n - k;
float* Hb = H + (size_t)mid * n * n;
int gbase = mid * G + g, go = g * BM;
__shared__ int msh;
__shared__ float red[32], sh_alpha, sh_sigma;
if (t == 0) msh = 0;
for (int idx = t; idx < bb * BM; idx += nt) { // load slice into smem (once)
int c = idx / BM, rl = idx - c * BM, pr = go + rl;
s[idx] = (pr < m) ? Hb[(size_t)(k + pr) * n + (k + c)] : 0.f;
}
__syncthreads();
int BG = gridDim.x; // B*G; double-buffer stride
for (int c = 0; c < bb; c++) {
int g_piv = c / BM, p = c & 1; // ping-pong scratch by column parity
float* sigp_p = sigp + p * BG; // -> removes the 3rd grid barrier
float* alp_p = alphasc + p * (BG / G);
float* wp_p = wp + (size_t)p * BG * bb;
float loc = 0.f; // partial sigma over local rows pr>c
for (int rl = t; rl < BM; rl += nt) {
int pr = go + rl;
if (pr > c && pr < m) { float v = s[c * BM + rl]; loc += v * v; }
}
for (int o = 16; o > 0; o >>= 1) loc += __shfl_down_sync(0xffffffff, loc, o);
if (lane == 0) red[warp] = loc;
__syncthreads();
if (t == 0) { float ss = 0.f; for (int i = 0; i < nwarps; i++) ss += red[i]; sigp_p[gbase] = ss; }
if (g == g_piv && t == 0) alp_p[mid] = s[c * BM + (c - go)];
grid_bar(counter, sense_arr, mid, G, &msh, t);
if (t == 0) {
float sg = 0.f; for (int i = 0; i < G; i++) sg += sigp_p[mid * G + i];
sh_sigma = sg; sh_alpha = alp_p[mid];
}
__syncthreads();
float sigma = sh_sigma, alpha = sh_alpha;
float norm = sqrtf(alpha * alpha + sigma);
float beta = (alpha >= 0.f) ? -norm : norm;
bool nz = sigma > 0.f;
float tj = nz ? (beta - alpha) / beta : 0.f, vs = nz ? 1.f / (alpha - beta) : 0.f;
if (g == g_piv && t == 0) {
s[c * BM + (c - go)] = nz ? beta : alpha;
tau[(size_t)mid * n + (k + c)] = tj;
}
if (nz) {
for (int rl = t; rl < BM; rl += nt) { int pr = go + rl; if (pr > c && pr < m) s[c * BM + rl] *= vs; }
__syncthreads();
for (int cc = c + 1 + warp; cc < bb; cc += nwarps) { // partial w[cc]
float wsum = 0.f;
for (int rl = lane; rl < BM; rl += 32) {
int pr = go + rl;
if (pr >= c && pr < m) { float v = (pr == c) ? 1.f : s[c * BM + rl]; wsum += v * s[cc * BM + rl]; }
}
for (int o = 16; o > 0; o >>= 1) wsum += __shfl_down_sync(0xffffffff, wsum, o);
if (lane == 0) wp_p[(size_t)gbase * bb + cc] = wsum;
}
}
grid_bar(counter, sense_arr, mid, G, &msh, t);
if (nz) {
for (int cc = c + 1 + warp; cc < bb; cc += nwarps) {
float wv = 0.f; for (int i = 0; i < G; i++) wv += wp_p[(size_t)(mid * G + i) * bb + cc];
wv *= tj;
for (int rl = lane; rl < BM; rl += 32) {
int pr = go + rl;
if (pr >= c && pr < m) { float v = (pr == c) ? 1.f : s[c * BM + rl]; s[cc * BM + rl] -= v * wv; }
}
}
}
__syncthreads(); // block-local only: order this column's smem writes before next col
}
float* Vb = Vd + (size_t)mid * n * bcap;
for (int idx = t; idx < bb * BM; idx += nt) { // write slice back (+ dense V)
int c = idx / BM, rl = idx - c * BM, pr = go + rl;
if (pr < m) {
float val = s[idx];
Hb[(size_t)(k + pr) * n + (k + c)] = val;
if (do_wy) Vb[(size_t)pr * bcap + c] = (pr < c) ? 0.f : (pr == c) ? 1.f : val;
}
}
}
// Standalone WY T-factor from a dense unit-lower V (used by the multi-block panel,
// which can't cheaply reduce V^T V across its row-blocks). One block per matrix:
// G = striu(V^T V), then T = inv(diag(1/tau)+G) by one-thread-per-column back-sub.
// Replaces the Python _form_T (V^T V bmm + triangular solve + triu + diag) launches.
__global__ void form_t_k(const float* __restrict__ Vd, const float* __restrict__ tau,
float* __restrict__ Tout, int n, int k, int bb, int bcap) {
extern __shared__ float sm[]; // Gm[bb*bb] then Tsm[bb*bb]
int m = n - k;
int t = threadIdx.x, nt = blockDim.x, lane = t & 31, warp = t >> 5, nwarps = nt >> 5;
int mid = blockIdx.x;
const float* Vb = Vd + (size_t)mid * n * bcap;
float* Gm = sm;
float* Tsm = sm + (size_t)bb * bb;
__shared__ float taus[64];
for (int i = t; i < bb; i += nt) taus[i] = tau[(size_t)mid * n + (k + i)];
__syncthreads();
for (int idx = warp; idx < bb * bb; idx += nwarps) {
int i = idx / bb, j = idx - i * bb;
if (i < j) {
float acc = 0.f;
for (int r = j + lane; r < m; r += 32)
acc += Vb[(size_t)r * bcap + i] * Vb[(size_t)r * bcap + j];
for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xffffffff, acc, o);
if (lane == 0) Gm[(size_t)i * bb + j] = acc;
}
}
__syncthreads();
for (int j = t; j < bb; j += nt) {
Tsm[(size_t)j * bb + j] = taus[j];
for (int i = j - 1; i >= 0; i--) {
float sum = 0.f;
for (int l = i + 1; l <= j; l++)
sum += Gm[(size_t)i * bb + l] * Tsm[(size_t)l * bb + j];
Tsm[(size_t)i * bb + j] = -taus[i] * sum;
}
for (int i = j + 1; i < bb; i++) Tsm[(size_t)i * bb + j] = 0.f;
}
__syncthreads();
float* Tb = Tout + (size_t)mid * bcap * bcap;
for (int idx = t; idx < bb * bb; idx += nt) {
int i = idx / bb, j = idx - i * bb;
Tb[(size_t)i * bcap + j] = Tsm[(size_t)i * bb + j];
}
}
void qr_panel_mb(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,
torch::Tensor counter, torch::Tensor sense_arr, torch::Tensor sigp,
torch::Tensor alphasc, torch::Tensor wp, int k, int bb, int G,
int BM, int do_wy) {
int B = H.size(0), n = H.size(1), nt = 1024, bcap = Vd.size(2); // 32 warps (red[32]):
size_t smem = (size_t)bb * BM * sizeof(float); // fill idle SMs (B*G
qr_panel_mb_k<<<B * G, nt, smem>>>( // ~16 blocks) + hide barrier spin
H.data_ptr<float>(), tau.data_ptr<float>(), Vd.data_ptr<float>(),
counter.data_ptr<int>(), sense_arr.data_ptr<int>(), sigp.data_ptr<float>(),
alphasc.data_ptr<float>(), wp.data_ptr<float>(), n, k, bb, G, BM, bcap, do_wy);
}
void form_t(torch::Tensor Vd, torch::Tensor tau, torch::Tensor Tout, int k, int bb) {
int B = Vd.size(0), n = Vd.size(1), bcap = Vd.size(2), nt = 256;
size_t smem = 2 * (size_t)bb * bb * sizeof(float);
form_t_k<<<B, nt, smem>>>(
Vd.data_ptr<float>(), tau.data_ptr<float>(), Tout.data_ptr<float>(), n, k, bb, bcap);
}
void qr_panel_init(int max_smem) {
cudaFuncSetAttribute(qr_panel_k,
cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
cudaFuncSetAttribute(qr_panel_mb_k,
cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
}
void qr_panel(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,
torch::Tensor Tout, int k, int bb, int do_wy, int nt) {
int B = H.size(0), n = H.size(1), bcap = Vd.size(2);
int m = n - k; // nt (<=1024, red[32]): caller fills the SM when batch is small,
// 512 when the batch already saturates the GPU (else reduce-tree waste)
size_t smem = ((size_t)m * bb + (do_wy ? 2 * (size_t)bb * bb : 0)) * sizeof(float);
qr_panel_k<<<B, nt, smem>>>(
H.data_ptr<float>(), tau.data_ptr<float>(), Vd.data_ptr<float>(),
Tout.data_ptr<float>(), n, k, bb, bcap, do_wy);
}
"""
_PANEL_CPP = (
"void qr_panel(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,"
" torch::Tensor Tout, int k, int bb, int do_wy, int nt);\n"
"void qr_panel_mb(torch::Tensor H, torch::Tensor tau, torch::Tensor Vd,"
" torch::Tensor counter, torch::Tensor sense_arr, torch::Tensor sigp,"
" torch::Tensor alphasc, torch::Tensor wp, int k, int bb, int G, int BM,"
" int do_wy);\n"
"void form_t(torch::Tensor Vd, torch::Tensor tau, torch::Tensor Tout, int k, int bb);\n"
"void qr_panel_init(int max_smem);"
)
_PANEL = None
try:
_PANEL = load_inline(
name="qr_panel_mod",
cpp_sources=[_PANEL_CPP],
cuda_sources=[_PANEL_CUDA],
functions=["qr_panel", "qr_panel_mb", "form_t", "qr_panel_init"],
verbose=False,
)
_PANEL.qr_panel_init(227000)
except Exception as _e:
_PANEL = None
# ----------------------------------------------------------------------------
# CholeskyQR1 + Householder-reconstruction panel (chol_bs) for the n=4096 wall.
# The two ops that would sync via cuSOLVER (32x32 Cholesky, unpivoted LU) are custom
# one-block-per-matrix kernels; the rest (G=A^TA, Q1=A R^-1, Y2=-Q1bot S^-1) are cuBLAS
# GEMM/trsm. This trades the sequential bb-column Householder critical path (the panel
# wall at batch 2) for parallel GEMMs + a trivial 32x32 factorization.
# ----------------------------------------------------------------------------
_TSQR_CUDA = r"""
#include <ATen/cuda/CUDAContext.h>
// Upper-Cholesky R (R^T R = G) of bb x bb SPD G, one block/matrix. info[b]=0 if PD else
// (first non-PD row + 1) so the caller routes that panel to the direct Householder fallback.
__global__ void chol_up_k(const float* __restrict__ G, float* __restrict__ R,
int* __restrict__ info, int bb, float rel_tol) {
extern __shared__ float sm[]; // Gs[bb*bb] | Rs[bb*bb]
float* Gs = sm; float* Rs = sm + (size_t)bb * bb;
int t = threadIdx.x;
const float* Gb = G + (size_t)blockIdx.x * bb * bb;
for (int idx = t; idx < bb * bb; idx += blockDim.x) { Gs[idx] = Gb[idx]; Rs[idx] = 0.f; }
__shared__ int bad; if (t == 0) bad = 0;
__syncthreads();
for (int i = 0; i < bb; i++) {
if (t >= i && t < bb) {
float s = 0.f;
for (int k = 0; k < i; k++) s += Rs[(size_t)k * bb + i] * Rs[(size_t)k * bb + t];
if (t == i) {
float gii = Gs[(size_t)i * bb + i], d = gii - s;
// relative pivot guard: tiny d/gii => column near-dependent (cond(A^TA) past
// fp32) => flag this matrix bad so the caller redoes it with direct geqrf.
if (d <= rel_tol * gii) { atomicExch(&bad, i + 1); d = (d > 0.f) ? d : 1.f; }
Rs[(size_t)i * bb + i] = sqrtf(d);
} else Rs[(size_t)i * bb + t] = Gs[(size_t)i * bb + t] - s;
}
__syncthreads();
if (t > i && t < bb) Rs[(size_t)i * bb + t] /= Rs[(size_t)i * bb + i];
__syncthreads();
}
float* Rb = R + (size_t)blockIdx.x * bb * bb;
for (int idx = t; idx < bb * bb; idx += blockDim.x) Rb[idx] = Rs[idx];
if (t == 0) info[blockIdx.x] = bad;
}
// Unpivoted LU of M0 (bb x bb), one block/matrix: in place strict-lower = Y1 multipliers,
// upper(incl diag) = S so M0 = Y1 @ S. |pivot|<=tol -> degenerate col (identity reflector):
// S row = e_i, Y1 col below = 0, degen[col]=1.
__global__ void lu_unpiv_k(const float* __restrict__ M0, float* __restrict__ LU,
int* __restrict__ degen, int bb, float tol) {
extern __shared__ float sm[]; // Ws[bb*bb]
float* Ws = sm; int t = threadIdx.x;
const float* Mb = M0 + (size_t)blockIdx.x * bb * bb;
for (int idx = t; idx < bb * bb; idx += blockDim.x) Ws[idx] = Mb[idx];
int* db = degen + (size_t)blockIdx.x * bb;
for (int idx = t; idx < bb; idx += blockDim.x) db[idx] = 0;
__syncthreads();
for (int i = 0; i < bb; i++) {
float piv = Ws[(size_t)i * bb + i];
bool deg = fabsf(piv) <= tol;
if (deg) {
if (t == i) { Ws[(size_t)i * bb + i] = 1.f; db[i] = 1; }
if (t > i && t < bb) { Ws[(size_t)t * bb + i] = 0.f; Ws[(size_t)i * bb + t] = 0.f; }
__syncthreads();
} else {
if (t > i && t < bb) Ws[(size_t)t * bb + i] /= piv;
__syncthreads();
if (t > i && t < bb) {
float u = Ws[(size_t)i * bb + t];
for (int r = i + 1; r < bb; r++) Ws[(size_t)r * bb + t] -= Ws[(size_t)r * bb + i] * u;
}
__syncthreads();
}
}
float* Lb = LU + (size_t)blockIdx.x * bb * bb;
for (int idx = t; idx < bb * bb; idx += blockDim.x) Lb[idx] = Ws[idx];
}
void chol_up(torch::Tensor G, torch::Tensor R, torch::Tensor info, int bb, double rel_tol) {
int B = G.size(0), nt = bb < 32 ? 32 : bb;
chol_up_k<<<B, nt, 2 * (size_t)bb * bb * sizeof(float)>>>(
G.data_ptr<float>(), R.data_ptr<float>(), info.data_ptr<int>(), bb, (float)rel_tol);
}
void lu_unpiv(torch::Tensor M0, torch::Tensor LU, torch::Tensor degen, int bb, double tol) {
int B = M0.size(0), nt = bb < 32 ? 32 : bb;
lu_unpiv_k<<<B, nt, (size_t)bb * bb * sizeof(float)>>>(
M0.data_ptr<float>(), LU.data_ptr<float>(), degen.data_ptr<int>(), bb, (float)tol);
}
void tsqr_init(int max_smem) { // opt-in dynamic smem so wide panels (bb=128: chol 128KB) fit
cudaFuncSetAttribute(chol_up_k, cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
cudaFuncSetAttribute(lu_unpiv_k, cudaFuncAttributeMaxDynamicSharedMemorySize, max_smem);
}
"""
_TSQR_CPP = (
"void chol_up(torch::Tensor G, torch::Tensor R, torch::Tensor info, int bb, double rel_tol);\n"
"void lu_unpiv(torch::Tensor M0, torch::Tensor LU, torch::Tensor degen, int bb, double tol);\n"
"void tsqr_init(int max_smem);"
)
_TSQR = None
try:
_TSQR = load_inline(
name="qr_tsqr_mod",
cpp_sources=[_TSQR_CPP],
cuda_sources=[_TSQR_CUDA],
functions=["chol_up", "lu_unpiv", "tsqr_init"],
verbose=False,
)
_TSQR.tsqr_init(227000)
except Exception as _e:
_TSQR = None
_CUDA_MB_SCRATCH = {}
def _panel_factor_cuda_mb(H, k, bb, tau, BM, Vd, do_wy):
"""Multi-block smem-resident CUDA panel. G blocks/matrix cooperate via grid barrier.
Emits dense unit-lower V into Vd (each block writes its own row-slice) when do_wy."""
B, n, _ = H.shape
Gmax = (n + BM - 1) // BM
key = (B, Gmax, bb, H.device)
sc = _CUDA_MB_SCRATCH.get(key)
if sc is None:
z = lambda *s: torch.zeros(*s, dtype=torch.int32, device=H.device)
f = lambda *s: torch.zeros(*s, dtype=torch.float32, device=H.device)
# sigp/alphasc/wp doubled for column-parity ping-pong (drops the 3rd barrier)
sc = (z(B), z(B), f(2 * B * Gmax), f(2 * B), f(2 * B * Gmax * bb))
_CUDA_MB_SCRATCH[key] = sc
counter, sense, sigp, alphasc, wp = sc
G = (n - k + BM - 1) // BM
_PANEL.qr_panel_mb(
H, tau, Vd, counter, sense, sigp, alphasc, wp, k, bb, G, BM, do_wy
)
# TF32 tensor cores for the WY trailing-update bmms (~30% of runtime). TF32 keeps
# an FP32 accumulator; multiplicands round to 19 bits. Toggled per-call: safe for
# n>=1024 (rtol scales with n) but breaks the tightest n=512 wide-dynamic-range
# cases (band/rowscale lose their small entries to the 10-bit TF32 mantissa).
def _set_tf32(on):
torch.backends.cuda.matmul.allow_tf32 = on
# (Shelved: a 3xTF32 / fp16x3 split recovers FP32 accuracy for the trailing GEMM
# but is a net SLOWDOWN on B200 — the GEMM isn't the bottleneck. See qr_notes.md.)
# ----------------------------------------------------------------------------
# Triton panel factorization: factor columns [k, k+bb) of each matrix in the
# batch in place (LARFG convention), writing Householder vectors below the
# diagonal, beta on the diagonal, and tau. grid = (batch,). Sequential over
# the bb columns (runtime loop) with row tiling of size BM.
# ----------------------------------------------------------------------------
@triton.jit
def _panel_kernel(
Hptr,
tauptr,
n,
k,
bb,
sH0,
sH1,
sH2,
stau0,
stau1,
BB: tl.constexpr,
BM: tl.constexpr,
PADDED: tl.constexpr,
):
# BB = tile width (power of 2) >= bb (actual panel columns). When bb < BB the
# extra columns are masked out (PADDED). Benchmark shapes have bb == BB (no mask).
pid = tl.program_id(0)
Hb = Hptr + pid * sH0
coff = tl.arange(0, BB)
cols = k + coff # global panel column indices
for c in range(0, bb):
j = k + c
# ---- pass 1: alpha, sigma = sum_{i>j} H[i,j]^2 ----
alpha = tl.load(Hb + j * sH1 + j * sH2)
sigma = tl.zeros((), dtype=tl.float32)
for r0 in range(0, n, BM):
rows = r0 + tl.arange(0, BM)
m = (rows > j) & (rows < n)
col = tl.load(Hb + rows * sH1 + j * sH2, mask=m, other=0.0)
sigma += tl.sum(col * col)
nrm = tl.sqrt(alpha * alpha + sigma)
beta = tl.where(alpha >= 0, -nrm, nrm)
nz = sigma > 0.0
tau_c = tl.where(nz, (beta - alpha) / beta, 0.0)
inv = tl.where(nz, 1.0 / (alpha - beta), 0.0)
tl.store(Hb + j * sH1 + j * sH2, tl.where(nz, beta, alpha))
tl.store(tauptr + pid * stau0 + j * stau1, tau_c)
if nz:
# scale v below the diagonal
for r0 in range(0, n, BM):
rows = r0 + tl.arange(0, BM)
m = (rows > j) & (rows < n)
col = tl.load(Hb + rows * sH1 + j * sH2, mask=m, other=0.0)
tl.store(Hb + rows * sH1 + j * sH2, col * inv, mask=m)
tl.debug_barrier()
# ---- apply reflector c to panel cols (vectorized over BB) ----
# w[cc] = tau_c * sum_{rows>=j} v_row * H[row, k+cc]
w = tl.zeros((BB,), dtype=tl.float32)
for r0 in range(0, n, BM):
rows = r0 + tl.arange(0, BM)
mr = (rows > j) & (rows < n)
v = tl.load(Hb + rows * sH1 + j * sH2, mask=mr, other=0.0)
v = tl.where(rows == j, 1.0, v)
mt = (rows[:, None] >= j) & (rows[:, None] < n)
if PADDED:
mt = mt & (coff[None, :] < bb)
ct = tl.load(
Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
)
w += tl.sum(v[:, None] * ct, axis=0)
w = tl.where(coff > c, tau_c * w, 0.0) # only update cols cc>c
for r0 in range(0, n, BM):
rows = r0 + tl.arange(0, BM)
mr = (rows > j) & (rows < n)
v = tl.load(Hb + rows * sH1 + j * sH2, mask=mr, other=0.0)
v = tl.where(rows == j, 1.0, v)
mt = (rows[:, None] >= j) & (rows[:, None] < n)
if PADDED:
mt = mt & (coff[None, :] < bb)
ct = tl.load(
Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
)
ct = ct - v[:, None] * w[None, :]
tl.store(Hb + rows[:, None] * sH1 + cols[None, :] * sH2, ct, mask=mt)
tl.debug_barrier()
@triton.jit
def _bar(counter, sense, cb, G, my_sense):
# Device-scope grid barrier across the G blocks of one matrix. acq_rel/gpu makes
# each block's partial-data writes visible to the others. Returns flipped sense.
my_sense = my_sense ^ 1
old = tl.atomic_add(counter + cb, 1, sem="acq_rel", scope="gpu")
if old == G - 1:
tl.atomic_xchg(counter + cb, 0, sem="relaxed", scope="gpu")
tl.atomic_xchg(sense + cb, my_sense, sem="release", scope="gpu")
else:
while tl.load(sense + cb, volatile=True) != my_sense:
pass
return my_sense
@triton.jit
def _panel_kernel_mb(
Hptr,
tauptr,
n,
k,
bb,
G,
sH0,
sH1,
sH2,
stau0,
stau1,
counter,
sense,
sigp,
wp,
alphasc,
BB: tl.constexpr,
BM: tl.constexpr,
PADDED: tl.constexpr,
):
# Multi-block panel factorization: G blocks cooperate on one matrix via a grid
# barrier, each owning a BM-row slice. Used for small-batch large-n (panel-bound).
pid = tl.program_id(0)
mid = pid // G
g = pid % G
Hb = Hptr + mid * sH0
coff = tl.arange(0, BB)
cols = k + coff
rows = k + g * BM + tl.arange(0, BM) # this block's row slice
rv = rows < n
# local sense starts at 0 for all blocks; global sense/counter start at 0 and return
# to 0 after the panel (4*bb barriers = even), so reused scratch stays clean.
ms = 0
for c in range(0, bb):
j = k + c
# block 0 owns row j -> broadcast alpha = H[j,j] (cross-block read otherwise stale)
if g == 0:
tl.atomic_xchg(
alphasc + mid,
tl.load(Hb + j * sH1 + j * sH2),
sem="release",
scope="gpu",
)
# ---- partial sigma over rows>j ----
below = rv & (rows > j)
col = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
tl.atomic_xchg(
sigp + mid * G + g, tl.sum(col * col), sem="release", scope="gpu"
)
ms = _bar(counter, sense, mid, G, ms)
sigma = tl.zeros((), dtype=tl.float32)
for i in range(0, G):
sigma += tl.atomic_add(sigp + mid * G + i, 0.0, sem="acquire", scope="gpu")
alpha = tl.atomic_add(alphasc + mid, 0.0, sem="acquire", scope="gpu")
nrm = tl.sqrt(alpha * alpha + sigma)
beta = tl.where(alpha >= 0, -nrm, nrm)
nz = sigma > 0.0
tau_c = tl.where(nz, (beta - alpha) / beta, 0.0)
inv = tl.where(nz, 1.0 / (alpha - beta), 0.0)
if g == 0:
tl.store(Hb + j * sH1 + j * sH2, tl.where(nz, beta, alpha))
tl.store(tauptr + mid * stau0 + j * stau1, tau_c)
if nz:
tl.store(Hb + rows * sH1 + j * sH2, col * inv, mask=below)
# ---- partial w[cc] = sum rows>=j v*H[row,cc] (own rows -> no barrier vs scale) ----
v = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
v = tl.where(rows == j, 1.0, v)
mt = (rows[:, None] >= j) & rv[:, None]
if PADDED:
mt = mt & (coff[None, :] < bb)
ct = tl.load(
Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
)
tl.atomic_xchg(
wp + (mid * G + g) * BB + coff,
tl.sum(v[:, None] * ct, axis=0),
sem="release",
scope="gpu",
)
ms = _bar(counter, sense, mid, G, ms)
if nz:
w = tl.zeros((BB,), dtype=tl.float32)
for i in range(0, G):
w += tl.atomic_add(
wp + (mid * G + i) * BB + coff, 0.0, sem="acquire", scope="gpu"
)
w = tl.where(coff > c, tau_c * w, 0.0)
v = tl.load(Hb + rows * sH1 + j * sH2, mask=below, other=0.0)
v = tl.where(rows == j, 1.0, v)
mt = (rows[:, None] >= j) & rv[:, None]
if PADDED:
mt = mt & (coff[None, :] < bb)
ct = tl.load(
Hb + rows[:, None] * sH1 + cols[None, :] * sH2, mask=mt, other=0.0
)
ct = ct - v[:, None] * w[None, :]
tl.store(Hb + rows[:, None] * sH1 + cols[None, :] * sH2, ct, mask=mt)
ms = _bar(counter, sense, mid, G, ms)
_MB_SCRATCH = {}
def _panel_factor_mb(H, k, bb, tau, BM=128, num_warps=4):
B, n, _ = H.shape
BB = _next_pow2(bb)
G = (n - k + BM - 1) // BM # blocks per matrix
key = (B, G, BB, H.device)
sc = _MB_SCRATCH.get(key)
if sc is None:
counter = torch.zeros(B, dtype=torch.int32, device=H.device)
sense = torch.zeros(B, dtype=torch.int32, device=H.device)
sigp = torch.zeros(B * G, dtype=torch.float32, device=H.device)
wp = torch.zeros(B * G * BB, dtype=torch.float32, device=H.device)
alphasc = torch.zeros(B, dtype=torch.float32, device=H.device)
sc = (counter, sense, sigp, wp, alphasc)
_MB_SCRATCH[key] = sc
counter, sense, sigp, wp, alphasc = sc
_panel_kernel_mb[(B * G,)](
H,
tau,
n,
k,
bb,
G,
H.stride(0),
H.stride(1),
H.stride(2),
tau.stride(0),
tau.stride(1),
counter,
sense,
sigp,
wp,
alphasc,
BB=BB,
BM=BM,
PADDED=(BB != bb),
num_warps=num_warps,
)
def _next_pow2(x):
return 1 << (x - 1).bit_length()
def _panel_threads(B, width):
"""Threads/block for the single-block panel kernel (qr_panel_k).
Once the batch already saturates the GPU (>~num_SMs blocks) 16 warps is best — extra
warps only add reduction-tree overhead (n=512 B=640 measurably regresses with more).
Otherwise the SM is under-utilized (smem caps it at 1 block/SM), so fill it with the
max 32 warps = 1024 threads (the hardware threads/block limit; can't go to 64 warps).
`width` is accepted for API symmetry but the 1024 cap makes it moot here."""
if B >= 256:
return 512
return 1024
def _panel_factor(H, k, bb, tau, BM=128, num_warps=4):
B, n, _ = H.shape
BB = _next_pow2(bb) # tile width must be a power of 2
_panel_kernel[(B,)](
H,
tau,
n,
k,
bb,
H.stride(0),
H.stride(1),
H.stride(2),
tau.stride(0),
tau.stride(1),
BB=BB,
BM=BM,
PADDED=(BB != bb),
num_warps=num_warps,
)
def _form_T(V, tau):
"""V: (B,m,b) unit-lower; tau: (B,b) -> T (B,b,b) upper, Q=I-V T V^T.
Closed form (no Python loop): T = inv(diag(1/tau) + striu(V^T V, 1)).
tau==0 (identity reflector) handled via a huge diagonal -> T entry ~0."""
B, m, b = V.shape
M = torch.triu(torch.bmm(V.transpose(1, 2), V), 1) # strictly upper
inv_tau = torch.where(tau != 0, 1.0 / tau, tau.new_full((), 1e30))
M.diagonal(dim1=-2, dim2=-1).copy_(inv_tau)
eye = torch.eye(b, device=V.device, dtype=V.dtype).expand(B, b, b)
return torch.linalg.solve_triangular(M, eye, upper=True)
def _extract_V(H, k, bb):
"""Unit-lower-trapezoidal V (B,m,bb) from factored panel."""
B, n, _ = H.shape
V = H[:, k:, k : k + bb].clone()
ii = torch.arange(bb, device=H.device)
top = V[:, :bb, :]
top.masked_fill_(ii[:, None] < ii[None, :], 0.0)
top[:, ii, ii] = 1.0
return V
def _recon_cuda_panel(H, tau, k, bb, tol=1e-7, rel_tol=1e-3):
"""chol_bs reconstruction of panel H[:,k:,k:k+bb]: writes packed H (R above diag, v's below)
+ tau, returns (V, info) where V is the dense unit-lower (B,m,bb) for the trailing WY update
and info[b]!=0 flags an ill-conditioned chol (caller redoes that matrix via reference geqrf)."""
B, n, _ = H.shape
dev = H.device
m = n - k
panel = H[:, k:, k : k + bb]
G = (panel.transpose(1, 2) @ panel).contiguous() # cuBLAS
R = torch.empty(B, bb, bb, device=dev)
info = torch.empty(B, dtype=torch.int32, device=dev)
_TSQR.chol_up(G, R, info, bb, rel_tol) # R upper, R^T R = G
# Q1 = panel @ R^-1 via R^T Q1^T = panel^T
Q1 = torch.linalg.solve_triangular(
R.transpose(1, 2), panel.transpose(1, 2), upper=False
).transpose(1, 2)
M0 = (torch.eye(bb, device=dev) - Q1[:, :bb, :]).contiguous()
LU = torch.empty(B, bb, bb, device=dev)
degen = torch.empty(B, bb, dtype=torch.int32, device=dev)
_TSQR.lu_unpiv(M0, LU, degen, bb, tol) # M0 = Y1 @ S
Y1 = torch.tril(LU, -1) + torch.eye(bb, device=dev)
if m > bb:
S = torch.triu(LU)
Y2 = torch.linalg.solve_triangular(
S.transpose(1, 2), (-Q1[:, bb:, :]).transpose(1, 2), upper=False
).transpose(1, 2) # Y2 = -Q1bot S^-1
V = torch.cat([Y1, Y2], dim=1)
else:
V = Y1
degb = degen.bool()
if degb.any():
ii = torch.arange(bb, device=dev)
cid = torch.arange(m, device=dev)[None, :, None] == ii[None, None, :]
V = torch.where(
degb[:, None, :].expand(B, m, bb), cid.to(V.dtype).expand(B, m, bb), V
)
tk = 2.0 / (V * V).sum(dim=1)
tk = torch.where(degb, torch.zeros_like(tk), tk)
ar = torch.arange(m, device=dev)
ac = torch.arange(bb, device=dev)
upper = (ar[:, None] <= ac[None, :])[None]
lower = (ar[:, None] > ac[None, :])[None]
Rfull = torch.zeros_like(V)
Rfull[:, :bb, :] = R
H[:, k:, k : k + bb] = torch.where(
upper, Rfull, torch.where(lower, V, torch.zeros_like(V))
)
tau[:, k : k + bb] = tk
return V, info
def _qr_tsqr_cholbs(A, b=32, recon_min_m=64, tol=1e-7, rel_tol=1e-3):
"""Block QR for n>=4096 via the chol_bs panel; near-square trailing panels (m<recon_min_m)
use the direct CUDA Householder panel. If any panel's chol is ill-conditioned (cond(A^TA)
past fp32: upper/rankdef/nearcollinear), the whole matrix is redone with reference geqrf -
a rare correctness-only path; dense (the timed case) never trips it (one sync at the end)."""
B, n, _ = A.shape
H = A.clone()
tau = torch.zeros(B, n, device=A.device, dtype=A.dtype)
nt = _panel_threads(B, b)
_dummy = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
any_bad = torch.zeros((), device=A.device)
for k in range(0, n, b):
bb = min(b, n - k)
if (n - k) >= recon_min_m:
V, info = _recon_cuda_panel(H, tau, k, bb, tol, rel_tol)
any_bad = any_bad + info.sum()
else:
_PANEL.qr_panel(
H, tau, _dummy, _dummy, k, bb, 0, nt
) # direct HH, writes H+tau
V = _extract_V(H, k, bb)
if k + bb < n:
C = H[:, k:, k + bb :]
Tm = _form_T(V, tau[:, k : k + bb])
W1 = torch.bmm(V.transpose(1, 2), C)
W2 = torch.bmm(Tm.transpose(1, 2), W1)
C.sub_(torch.bmm(V, W2))
if any_bad.item() > 0: # ill-conditioned -> geqrf
return torch.geqrf(A)
return H, tau
def custom_kernel(data: input_t) -> output_t:
A = data
B, n, _ = A.shape
# Precision: single-pass TF32 tensor cores ok for n>=1024 (factor rtol scales
# with n, loose enough); n<=512 stays fp32 (single TF32 fails band/rowscale, and
# 3xTF32 recovery is correct but slower since the GEMM isn't the bottleneck).
_set_tf32(n >= 1024)
# n>=4096: the panel wall. chol_bs (CholeskyQR1 + Householder reconstruction) trades the
# sequential bb-column Householder critical path for parallel GEMMs + tiny 32x32 chol/LU.
if (_TSQR is not None) and (_PANEL is not None) and (n >= 4096):
return _qr_tsqr_cholbs(
A, b=64, recon_min_m=128
) # wide panel cuts per-panel launches;
# bb=128 is within noise but needs 128KB smem
H = A.clone() # one copy: checker reads original A for the residual
tau = torch.zeros(B, n, device=A.device, dtype=A.dtype)
# Fully-fused small-n path: whole matrix fits B200 smem (n*n*4 <= ~227KB, n<=238).
# One launch does the entire unblocked Householder factor in smem (panel == full
# matrix, bb=n), so the in-kernel rank-1 trailing apply replaces ALL the serialized
# per-panel relaunches + extract_V/form_T/bmm glue that dominate small/mid-n runtime.
if (_PANEL is not None) and (n * n * 4 <= 227000):
# Fused whole-matrix factor (panel == full matrix): width = n.
_dummy = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
_PANEL.qr_panel(H, tau, _dummy, _dummy, 0, n, 0, _panel_threads(B, n))
return H, tau
b = 64 if n <= 1024 else 32
bm = 256 if n <= 1024 else 512
nw = 4 if n <= 1024 else 8
# smem-resident CUDA panel when the largest panel (m=n) fits B200 smem (~195KB).
# n=512: b=64 -> 128KB. n=1024: need b=32 -> 128KB (b=64 would be 256KB, too big).
# Multi-block cooperative panel for small-batch large-n: one block/matrix starves
# the GPU (n=2048 B=8 -> 8 of 148 SMs; n=4096 B=2 -> 2). G blocks/matrix cooperate
# via grid barriers so the panel fills the GPU. BM trades occupancy vs barrier cost
# (barriers scale with G participants): n>=4096 -> mbm=512 (G=8); n=2048 -> mbm=256
# (G=8 -> 64 blocks). [Re-testing n=2048 mb: the prior "B=8 fills enough" was wrong
# for 148 SMs and predates the current smem-resident CUDA cooperative kernel.]
use_mb = (B <= 8) and (n >= 4096)
if use_mb:
mbm = 256 # n=4096: G=16 (32 blocks) is the sweet spot at nt=1024 (~53ms): smaller
# mbm=128 (G=32) hits the 32-way grid-barrier wall (60ms), larger 512
# (G=8) underfills the SMs (56ms). n=2048 stays single-block — its grid
# barriers cost more than filling 140 idle SMs (mb 25.8 vs sb 17.6ms).
use_cuda_panel = (_PANEL is not None) and (n <= 2048) and not use_mb
if use_cuda_panel and n == 1024:
b = 32 # 1024x32x4 = 128KB fits smem (b=64 would be 256KB)
if use_cuda_panel and n == 2048:
b = 24 # widest that fits ~227KB; low batch -> wide panel wins (b sweep: 24 best)
# CUDA multi-block panel: smem-resident row-slice + grid barrier. b even (barrier
# parity) and bb*BM*4 must fit smem (32*256*4=32KB @2048; 32*512*4=64KB @4096).
use_cuda_mb = (_PANEL is not None) and use_mb
if use_cuda_mb:
b = 32
# The panel kernel emits the WY V/T directly (no _extract_V/_form_T launches) only
# where it beats batched cuBLAS form_T: narrow-panel many-panel cases (n=1024/2048).
# At n<=512 the batch is large and the panel wide (b=64) -> cuBLAS V^TV wins; keep it.
# Single-block panel emits V/T in-kernel where it beats batched cuBLAS form_T
# (narrow-panel many-panel n=1024/2048); n<=512 keeps cuBLAS. The multi-block panel
# (n=4096) emits dense V per row-block + a standalone form_t kernel for T -> kills the
# per-panel Python _extract_V/_form_T launches (was 34% of the 4096 time).
use_emit_sb = use_cuda_panel and (
n >= 1024
) # n=512 (B=640): cuBLAS form_T (tensor)
# beats in-kernel G=VtV (scalar) by far
use_emit_mb = use_cuda_mb
use_emit = use_emit_sb or use_emit_mb
panel_nt = _panel_threads(
B, b
) # single-block panel threads (width = block width b)
if use_emit:
Vd_buf = torch.empty(B, n, b, device=A.device, dtype=A.dtype)
T_buf = torch.empty(B, b, b, device=A.device, dtype=A.dtype)
else:
Vd_buf = torch.empty(1, 1, 1, device=A.device, dtype=A.dtype)
T_buf = Vd_buf
for k in range(0, n, b):
bb = min(b, n - k)
if use_cuda_panel:
_PANEL.qr_panel(
H, tau, Vd_buf, T_buf, k, bb, 1 if use_emit_sb else 0, panel_nt
)
elif use_cuda_mb:
_panel_factor_cuda_mb(
H, k, bb, tau, BM=mbm, Vd=Vd_buf, do_wy=1 if use_emit_mb else 0
)
elif use_mb:
_panel_factor_mb(H, k, bb, tau, BM=mbm, num_warps=4)
else:
_panel_factor(H, k, bb, tau, BM=bm, num_warps=nw)
if k + bb < n:
C = H[:, k:, k + bb :]
if use_emit_sb:
V = Vd_buf[:, : n - k, :bb]
Tm = T_buf[:, :bb, :bb]
elif use_emit_mb:
# emitted V skips extract_V; T via cuBLAS form_T (big-K V^TV, cuBLAS wins)
V = Vd_buf[:, : n - k, :bb]
Tm = _form_T(V, tau[:, k : k + bb])
else:
V = _extract_V(H, k, bb)
Tm = _form_T(V, tau[:, k : k + bb])
W1 = torch.bmm(V.transpose(1, 2), C)
W2 = torch.bmm(Tm.transpose(1, 2), W1)
C.sub_(torch.bmm(V, W2))
return H, tau
scrolls · 984 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