submission 834880
moogician · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 438 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-834880?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:46c90429a0d55dd1c5936858e331a10e334fed045061426ca2b32ce3c9175649
license declaredunknown
license concludedunknown
authorsmoogician
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
G = tl.dot(tl.trans(Vc), Vc,Kernel source
submission.py438 lines
import torch
from task import input_t, output_t
# Trailing-update GEMMs (V^T C, T^T Wt, V Wt) dominate the n=512/1024 FLOPs.
# TF32 tensor cores run them far faster; whether it stays inside the per-matrix
# factor-residual gate depends on n (the gate scales with n, so larger n has more
# slack). We toggle the global flag per call rather than globally.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
try:
import triton
import triton.language as tl
@triton.jit
def _panel_kernel(H_ptr, tau_ptr, V_ptr, T_ptr, n, m, w, k,
sb, srow, scol, stb, vb, vrow, vcol, tb, trow, tcol,
BLOCK_M: tl.constexpr, BLOCK_W: tl.constexpr,
DOT_TF32: tl.constexpr, NO_VT: tl.constexpr,
LOG2W: tl.constexpr, TPREC: tl.constexpr = "ieee"):
b = tl.program_id(0)
rm = tl.arange(0, BLOCK_M)
rw = tl.arange(0, BLOCK_W)
rmask = rm < m
wmask = rw < w
base = H_ptr + b * sb + k * srow + k * scol
ptrs = base + rm[:, None] * srow + rw[None, :] * scol
mask2d = rmask[:, None] & wmask[None, :]
P = tl.load(ptrs, mask=mask2d, other=0.0)
tau_acc = tl.zeros([BLOCK_W], dtype=tl.float32)
for c in range(0, BLOCK_W):
if c < w:
colc = tl.sum(tl.where((rw == c)[None, :], P, 0.0), axis=1) # [BLOCK_M]
ge = rm >= c
# fuse the two M-dim reductions (alpha=colc[c], nf2=sum_{i>=c} colc[i]^2)
# into one reduction over a [BLOCK_M, 2] stack to cut sequential
# warp-reduction latency on the per-column critical path.
two = tl.arange(0, 2)
stk = tl.where(two[None, :] == 0,
tl.where(rm == c, colc, 0.0)[:, None],
tl.where(ge, colc * colc, 0.0)[:, None]) # [BLOCK_M, 2]
red = tl.sum(stk, axis=0) # [2]
alpha = tl.sum(tl.where(two == 0, red, 0.0))
nf2 = tl.sum(tl.where(two == 1, red, 0.0))
nf = tl.sqrt(nf2)
an = tl.abs(alpha)
pos = nf > 0.0
sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
tauv = tl.where(pos, 1.0 + an / tl.where(pos, nf, 1.0), 0.0)
inv = tl.where(pos, sgn / (nf + an), 0.0)
beta = tl.where(pos, -sgn * nf, alpha)
v = tl.where(rm == c, 1.0, tl.where(rm > c, colc * inv, 0.0))
wvec = tl.sum(v[:, None] * P, axis=0) # [BLOCK_W]
upd = (tauv * v)[:, None] * wvec[None, :]
P = tl.where((rw > c)[None, :], P - upd, P)
newcol = tl.where(rm < c, colc, tl.where(rm == c, beta, colc * inv))
P = tl.where((rw == c)[None, :], newcol[:, None], P)
tau_acc = tl.where(rw == c, tauv, tau_acc)
tl.store(ptrs, P, mask=mask2d)
tl.store(tau_ptr + b * stb + k + rw, tau_acc, mask=wmask)
if NO_VT:
return
# clean unit-lower V (1 on diagonal, v below, 0 above) for the trailing update
Vc = tl.where(rm[:, None] > rw[None, :], P,
tl.where(rm[:, None] == rw[None, :], 1.0, 0.0))
vptrs = V_ptr + b * vb + rm[:, None] * vrow + rw[None, :] * vcol
tl.store(vptrs, Vc, mask=mask2d)
# block reflector T (upper-tri): Q = I - V T V^T, via LARFT recurrence on G = V^T V
G = tl.dot(tl.trans(Vc), Vc,
input_precision=("tf32" if DOT_TF32 else TPREC)) # [BLOCK_W, BLOCK_W]
# Block reflector T = (I - N)^{-1} diag(tau), N = -diag(tau) striu(G)
# (strictly upper, nilpotent: N^BLOCK_W = 0). The 32-step sequential LARFT
# recurrence is latency-bound on the warp-reduction hardware; instead build
# (I-N)^{-1} = sum_{i<32} N^i = prod_{k} (I + N^{2^k}) with ~8 tl.dot's, which
# run on the MMA units (off the shuffle-reduction critical path). fp32 dots
# keep T as accurate as the exact recurrence. (BLOCK_W==32 -> 5 factors.)
eye = rw[:, None] == rw[None, :]
Imat = tl.where(eye, 1.0, 0.0)
N = tl.where(rw[:, None] < rw[None, :], -tau_acc[:, None] * G, 0.0)
# (I-N)^{-1} = sum_i N^i = prod_k (I + N^{2^k}); N is BLOCK_W-nilpotent, so
# LOG2W-1 doubling stages (covering up to N^{BLOCK_W-1}) are exact.
acc = Imat + N
Np = N
for _ in range(LOG2W - 1):
Np = tl.dot(Np, Np, input_precision=TPREC)
acc = tl.dot(acc, Imat + Np, input_precision=TPREC)
T = acc * tau_acc[None, :] # (I-N)^{-1} diag(tau)
tptrs = T_ptr + b * tb + rw[:, None] * trow + rw[None, :] * tcol
tl.store(tptrs, T, mask=wmask[:, None] & wmask[None, :])
_HAS_TRITON = True
except Exception:
_HAS_TRITON = False
def _next_pow2(x):
p = 1
while p < x:
p <<= 1
return p
@torch.jit.script
def _refl(colf: torch.Tensor, alpha: torch.Tensor):
nf = torch.sqrt((colf * colf).sum(dim=1)) # full-column norm
an = torch.abs(alpha)
mask = nf > 0
safen = torch.where(mask, nf, torch.ones_like(nf))
tauj = torch.where(mask, 1.0 + an / safen, torch.zeros_like(nf))
inv = torch.where(mask, torch.copysign(1.0 / (nf + an), alpha), torch.zeros_like(nf))
newdiag = torch.where(mask, -torch.copysign(nf, alpha), alpha)
return inv, tauj, newdiag
def _house_qr_blocked(A, nb=64):
"""Batched blocked Householder QR producing geqrf-compact (H, tau)."""
B, n, _ = A.shape
dev, dt = A.device, A.dtype
H = A.clone()
tau = torch.zeros(B, n, device=dev, dtype=dt)
zeros = torch.zeros(B, device=dev, dtype=dt)
for k in range(0, n, nb):
kb = min(nb, n - k)
jend = k + kb
# ---- unblocked panel factorization (columns k..jend) ----
for jj in range(kb):
j = k + jj
colf = H[:, j:, j] # (B, m) view, full column
alpha = H[:, j, j] # (B,)
inv, tauj, newdiag = _refl(colf, alpha)
tau[:, j] = tauj
if j + 1 < n:
H[:, j + 1:, j].mul_(inv[:, None]) # scale v below diagonal
# apply reflector within panel to cols j+1..jend-1
if j + 1 < jend:
colf[:, 0] = 1.0 # v[0]=1 (overwrites alpha)
sub = H[:, j:, j + 1:jend] # (B, m, w) view of H
w = torch.bmm(colf[:, None, :], sub).squeeze(1) # (B, w) = v^T C
w.mul_(tauj[:, None])
torch.baddbmm(sub, colf[:, :, None], w[:, None, :],
beta=1.0, alpha=-1.0, out=sub) # C -= v w
H[:, j, j] = newdiag # write R diagonal
# ---- build block reflector V, T (compact WY) ----
# T satisfies T^{-1} = diag(1/tau) + striu(V^T V), i.e.
# (I + diag(tau) striu(G)) T = diag(tau) -> one triangular solve.
Vp = H[:, k:, k:jend] # (B, m, kb)
V = torch.tril(Vp, -1)
di = torch.arange(kb, device=dev)
V[:, di, di] = 1.0
G = torch.bmm(V.transpose(1, 2), V) # (B, kb, kb) Gram
taup = tau[:, k:jend] # (B, kb)
M = torch.triu(G, 1) * taup[:, :, None] # diag(tau) @ striu(G)
M[:, di, di] = 1.0 # unit upper triangular
RHS = torch.diag_embed(taup)
T = torch.linalg.solve_triangular(M, RHS, upper=True, unitriangular=True)
# ---- trailing update: C -= V (T^T (V^T C)) ----
if jend < n:
C = H[:, k:, jend:]
Wt = torch.bmm(V.transpose(1, 2), C)
Wt = torch.bmm(T.transpose(1, 2), Wt)
torch.baddbmm(C, V, Wt, beta=1.0, alpha=-1.0, out=C)
return H, tau
def _panel_warps(BLOCK_M, B, nwarps):
# Only the BLOCK_M==1024 panels (n=1024 early panels) truly need 8 warps;
# at 4 they are thread-starved (256 rows/warp).
# The mid tiers (BLOCK_M==512/256) DIVERGE by shape, keyed on the nwarps hint
# (n=1024 path passes nwarps=8, n=512 path passes 4):
# - n=1024 LATER panels (m<=512, nwarps>=8): need 4 -- at 2 they are
# thread-starved (CATASTROPHIC 6.8ms, the 512-tile reduction can't hide).
# - n=512 EARLY panels (the DOMINANT ones, narrow nb=16 tile, nwarps=4):
# reduce FASTER at 2 warps under heavy cross-call overlap (nslots=10) --
# more concurrent CTAs prefer fewer warps/CTA (fewer cross-warp combines).
# Re-swept @ nslots=10: 512=2,256=2 -> 4.84->4.80 (the old 4 was tuned at
# nslots=6, before every-iter-its-own-buffer max overlap shifted the optimum).
if BLOCK_M >= 1024:
return nwarps
elif BLOCK_M >= 256:
return 4 if nwarps >= 8 else 2
elif BLOCK_M >= 64:
return 2 if B >= 256 else 4
return 2
def _house_qr_triton_agg(A, nb=32, ag=2, tf32=False, nwarps=8, dot_tf32=None,
inplace=False, wide_tf32=None, tprec="ieee"):
"""Blocked Householder QR, panel factored by the fused Triton kernel, but the
WIDE trailing update is AGGREGATED across `ag` consecutive nb-panels into one
rank-(ag*nb) update. The panel kernel stays at the fast nb=32 width; only the
memory-bound wide trailing is widened, halving (for ag=2) the number of full
passes over the trailing matrix C -- the dominant n=512 cost."""
B, n, _ = A.shape
dev, dt = A.device, A.dtype
if dot_tf32 is None:
dot_tf32 = tf32
H = A if inplace else A.clone()
tau = torch.empty(B, n, device=dev, dtype=dt)
sb, srow, scol = H.stride()
stb = tau.stride(0)
BLOCK_W = nb
LOG2W = max(1, _next_pow2(nb).bit_length() - 1)
GW = ag * nb # combined block width
# combined-group V (unit lower-trapezoidal, GW wide) and combined T (GW x GW)
Vg = torch.zeros(B, n, GW, device=dev, dtype=dt)
Tg = torch.zeros(B, GW, GW, device=dev, dtype=dt)
Wt1 = torch.empty(B, GW, n, device=dev, dtype=dt)
Wt2 = torch.empty(B, GW, n, device=dev, dtype=dt)
Gc = torch.empty(B, GW, nb, device=dev, dtype=dt) # cross Gram / cross-T scratch
Cc = torch.empty(B, GW, nb, device=dev, dtype=dt)
torch.backends.cuda.matmul.allow_tf32 = tf32
for kk in range(0, n, GW):
gend = min(kk + GW, n)
mg = n - kk # rows in this group's frame
gw = gend - kk # actual group width (<=GW)
wide = gend < n # a wide trailing update follows
subs = list(range(kk, gend, nb))
# NB: the block-upper-triangular corners of Vg (rows above each sub-panel's
# diagonal start) stay zero from the initial torch.zeros -- they are never
# written by any kernel/bmm, and the initial zeros re-runs on each graph
# replay -- so no per-group re-zeroing is needed (combined V stays unit
# lower-trapezoidal; reads always slice to the valid mg rows).
# ---- factor sub-panels; narrow-update later sub-panels; accumulate combined T ----
for i, k in enumerate(subs):
kb = min(nb, n - k)
jend = k + kb
m = n - k
p = i * nb # combined width before this sub-panel
BLOCK_M = _next_pow2(m)
pw = _panel_warps(BLOCK_M, B, nwarps)
Vslice = Vg[:, k - kk:, p:p + kb]
Tslice = Tg[:, p:p + kb, p:p + kb]
vbi, vrowi, vcoli = Vslice.stride()
tbi, trowi, tcoli = Tslice.stride()
no_vt = (jend >= n) # only the final panel skips V/T
_panel_kernel[(B,)](H, tau, Vslice, Tslice, n, m, kb, k,
sb, srow, scol, stb, vbi, vrowi, vcoli,
tbi, trowi, tcoli, BLOCK_M=BLOCK_M, BLOCK_W=BLOCK_W,
DOT_TF32=dot_tf32, NO_VT=no_vt, LOG2W=LOG2W,
TPREC=tprec, num_warps=pw)
# accumulate the combined-WY cross block for the wide update:
# T[:p, p:p+kb] = -T[:p,:p] @ (V[:,:p]^T @ V[:,p:p+kb]) @ T_i
if wide and p > 0:
Vprev = Vg[:, :mg, :p]
Vi = Vg[:, :mg, p:p + kb]
Tprev = Tg[:, :p, :p]
Ti = Tg[:, p:p + kb, p:p + kb]
gram = torch.bmm(Vprev.transpose(1, 2), Vi, out=Gc[:, :p, :kb])
cross = torch.bmm(Tprev, gram, out=Cc[:, :p, :kb])
cross = torch.bmm(cross, Ti, out=Gc[:, :p, :kb]).neg_()
Tg[:, :p, p:p + kb] = cross
# narrow update: apply this block to the REMAINING sub-panels in the group
if jend < gend:
V = Vg[:, k - kk:k - kk + m, p:p + kb]
T = Tslice
ncols = gend - jend
C = H[:, k:, jend:gend]
w1 = Wt1[:, :kb, :ncols]
w2 = Wt2[:, :kb, :ncols]
torch.bmm(V.transpose(1, 2), C, out=w1)
torch.bmm(T.transpose(1, 2), w1, out=w2)
torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
# ---- aggregated WIDE trailing update over cols [gend:n] (rank gw) ----
if wide:
V = Vg[:, :mg, :gw]
T = Tg[:, :gw, :gw]
cols = n - gend
C = H[:, kk:, gend:]
w1 = Wt1[:, :gw, :cols]
w2 = Wt2[:, :gw, :cols]
if wide_tf32 is None:
torch.bmm(V.transpose(1, 2), C, out=w1)
torch.bmm(T.transpose(1, 2), w1, out=w2)
torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
else:
# PER-GEMM precision split (n=512): the two BIG wide-trailing GEMMs
# (V^T C and C-=V w2, each m*gw*cols MACs) run in TF32 tensor cores
# (~2x the CUDA-core fp32 path), while the tiny step2 (T^T w1) and
# all narrow/cross-T work stay fp32. Only ONE of the two big GEMMs'
# worth of tf32 error lands per pass and the aggregated path makes
# only gw/nb-fewer passes, so mixed-n512 resid 1.12e-3 < gate
# 1.22e-3 holds (full-tf32 was 1.65e-3, FAIL). Other n512 cases have
# ~1000x slack. fp32-only schedules left the trailing CUDA-core-bound.
torch.backends.cuda.matmul.allow_tf32 = wide_tf32
torch.bmm(V.transpose(1, 2), C, out=w1)
torch.backends.cuda.matmul.allow_tf32 = tf32
torch.bmm(T.transpose(1, 2), w1, out=w2)
torch.backends.cuda.matmul.allow_tf32 = wide_tf32
torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
torch.backends.cuda.matmul.allow_tf32 = tf32
torch.backends.cuda.matmul.allow_tf32 = False
return H, tau
def _house_qr_triton(A, nb=32, tf32=False, nwarps=8, dot_tf32=None, inplace=False):
"""Blocked Householder QR with the panel factored by a fused Triton kernel."""
B, n, _ = A.shape
dev, dt = A.device, A.dtype
if dot_tf32 is None:
dot_tf32 = tf32
# inplace=True factors A directly (caller owns a scratch buffer); avoids the
# clone so a captured CUDA graph needs only ONE input copy per replay.
H = A if inplace else A.clone()
# every tau entry is written by the panel kernel (kb per panel, covering 0..n),
# so an uninitialized buffer is safe and skips the zero-init launch.
tau = torch.empty(B, n, device=dev, dtype=dt)
sb, srow, scol = H.stride()
stb = tau.stride(0)
BLOCK_W = nb
LOG2W = max(1, _next_pow2(nb).bit_length() - 1) # log2(nb) for the T recurrence
Vbuf = torch.empty(B, n, BLOCK_W, device=dev, dtype=dt)
Tbuf = torch.empty(B, BLOCK_W, BLOCK_W, device=dev, dtype=dt)
vb, vrow, vcol = Vbuf.stride()
tbs, trow, tcol = Tbuf.stride()
# reusable trailing-update scratch (V^T C and T^T(V^T C)); allocating once and
# slicing per panel keeps the CUDA-graph capture pool from bloating (one buffer
# instead of one fresh allocation per panel -> graphable even for n=512/1024).
Wt1 = torch.empty(B, BLOCK_W, n, device=dev, dtype=dt)
Wt2 = torch.empty(B, BLOCK_W, n, device=dev, dtype=dt)
torch.backends.cuda.matmul.allow_tf32 = tf32
for k in range(0, n, nb):
kb = min(nb, n - k)
jend = k + kb
m = n - k
BLOCK_M = _next_pow2(m)
# Per-panel warp count: big early panels (BLOCK_M>=256) need all 8 warps
# (fewer starves them -- memory: global nwarps=4 catastrophic), but the
# small later panels (BLOCK_M<256) are oversubscribed at 8 warps (256
# threads for <256 rows) -> fewer warps cuts cross-warp reduction overhead.
# Optimal warp count falls with BLOCK_M (small tiles oversubscribe at 8
# warps) but also depends on OCCUPANCY: a large batch (B>=256, e.g. n=512
# b=640) saturates the GPU so its small tail panels tolerate just 2 warps,
# while a small batch (n=176/352 b=40) needs more warps to hide latency.
if BLOCK_M >= 512:
pw = nwarps
elif BLOCK_M >= 256:
pw = 4
elif BLOCK_M >= 64:
pw = 2 if B >= 256 else 4
else:
pw = 2
_panel_kernel[(B,)](H, tau, Vbuf, Tbuf, n, m, kb, k, sb, srow, scol, stb,
vb, vrow, vcol, tbs, trow, tcol,
BLOCK_M=BLOCK_M, BLOCK_W=BLOCK_W, DOT_TF32=dot_tf32,
NO_VT=(jend >= n), LOG2W=LOG2W, num_warps=pw)
V = Vbuf[:, :m, :kb] # clean unit-lower, from kernel
T = Tbuf[:, :kb, :kb] # block reflector, from kernel
if jend < n:
C = H[:, k:, jend:]
cols = n - jend
w1 = Wt1[:, :kb, :cols]
w2 = Wt2[:, :kb, :cols]
# fp32 trailing for n<=512 (tight gate); TF32 tensor cores for n>=1024
# where the looser gate absorbs the rounding.
torch.bmm(V.transpose(1, 2), C, out=w1)
torch.bmm(T.transpose(1, 2), w1, out=w2)
torch.baddbmm(C, V, w2, beta=1.0, alpha=-1.0, out=C)
torch.backends.cuda.matmul.allow_tf32 = False
return H, tau
# Small shapes (n<=352) are launch-overhead-bound (~4 launches/panel, several
# panels). Their working set is tiny, so a captured CUDA graph collapses all the
# per-panel launches into one replay without the cache-thrash that killed graphs
# for n>=512 (large per-panel trailing temporaries -> capture-pool bloat).
_GRAPH_CACHE = {}
def _graphed_triton(data, nb, tf32, dot_tf32, nwarps=8, agg=1, wide_tf32=None,
tprec="ieee", nslots=1):
# Direct factorization on the default execution queue. CUDA graphs use capture
# (an internal side queue), which the leaderboard disallows, so the graph
# capture/replay path is removed. We clone the input since the factorization
# runs in place.
buf = data.clone()
if agg >= 2:
return _house_qr_triton_agg(buf, nb, ag=agg, tf32=tf32,
dot_tf32=dot_tf32, nwarps=nwarps,
inplace=True, wide_tf32=wide_tf32,
tprec=tprec)
return _house_qr_triton(buf, nb, tf32=tf32, dot_tf32=dot_tf32,
nwarps=nwarps, inplace=True)
def custom_kernel(data: input_t) -> output_t:
B, n, _ = data.shape
if _HAS_TRITON and 32 <= n <= 1024 and B >= 8:
try:
# n>=1024's factor-residual gate (20*n*eps) is loose enough to absorb
# TF32 trailing updates; n<=512 stays fp32 to keep ill-conditioned
# (mixed) batches inside the tighter gate.
# n=512: nb=16 (NOT 32). The within-panel unblocked wvec work is
# O(n*nb*m) -- it scales with the panel width -- so a NARROWER panel
# halves the panel-kernel's dominant cross-warp v^T C reduction. The
# memory's old "nb=16 is slower" finding was with PER-PANEL trailing
# (which then makes 32 memory-bound passes over C); the aggregation path
# DECOUPLES trailing rank from panel width (ag=4 keeps the rank-64 wide
# trailing = 8 passes), so nb=16 wins ONLY in the agg context. nb must
# be >=16 (tl.dot MMA minimum for the in-kernel G-dot). n=1024 is
# occupancy-bound (60 progs, BLOCK_M=1024), NOT within-panel-bound, so
# narrowing nb there just adds non-overlapped passes -> stays nb=32.
nb = 16 if n == 512 else 32
# n=512 stays fp32: a TF32 precision SCHEDULE was tried (TF32 early/late
# windows) and the mixed-batch error stays ~1.6e-3 > 1.22e-3 gate no matter
# which panels are TF32 -- a single extreme matrix (nearcollinear/rowscaled)
# in the mixed batch is ruined by any TF32 trailing. fp32-locked confirmed.
# All other Triton shapes (32/176/352 dense-only, 1024 looser gate) tolerate TF32.
# TF32 trailing updates overran the leaderboard's correctness tolerance
# on ill-conditioned public-test matrices (the locally-tuned gates were
# looser than the grader's), so the Triton path runs fully fp32.
use_tf32 = False
# n=512: aggregate the memory-bound wide trailing update into rank-64
# (nb=16 * ag=4) -- keeps trailing passes low while the narrow panel cuts
# within-panel work. n=1024: rank-128 (nb=32 * ag=4).
agg = 4 if n in (512, 1024) else 1
# n=512: keep panel + step2 fp32 (tight mixed gate) but run the two big
# wide-trailing GEMMs in TF32 -- per-GEMM split passes (1.12e-3<1.22e-3)
# where full-TF32 fails, recovering the tensor-core trailing speedup.
wide_tf32 = None
# nwarps=4 for the BLOCK_M==512 panels (faster cross-warp reduction than
# 8): wins for n=512 (nb=16 narrow tile) AND n=352 (plain-path BLOCK_M=512
# early panels). ONLY n=1024 needs 8 (its BLOCK_M==1024 panels are
# thread-starved at 4 -> 53ms). n=176/32 have no BLOCK_M>=512 panel so the
# value is irrelevant for them. (The memory's "nwarps=4 catastrophic" was
# for BLOCK_M==1024 / the wider nb=32 tile.)
nwarps = 8 if n == 1024 else 4
return _graphed_triton(data, nb, tf32=use_tf32, dot_tf32=use_tf32,
agg=agg, wide_tf32=wide_tf32)
except Exception:
pass
# Batched blocked Householder wins only when batch is large and n moderate;
# otherwise cuSOLVER per-matrix (geqrf) is faster.
if 352 < n <= 1024 and B >= 32:
return _house_qr_blocked(data, 192 if n > 512 else 96)
# Large-n small-batch and everything else: plain batched cuSOLVER geqrf on the
# default execution queue.
return torch.geqrf(data)
scrolls · 438 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