submission 844791
arseni_ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1101 lines, June 9 Researcher Reciprocity License v1.0.
blackwell_qr.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844791?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:c07ee29592c7fe93db8b81af770d3780fcbab7dca3c6ae79167535fe34ed715e
license declaredunknown
license concludedunknown
authorsarseni_ivanov
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
_coop_smem_cache = {}Kernel source
blackwell_qr.py1101 lines
#!POPCORN leaderboard qr_v2
import cutlass
import cutlass.cute as cute
import torch
import cutlass.torch as cutlass_torch
from cutlass.cute.runtime import from_dlpack
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm, vector
from cutlass.cute.typing import (
Float32,
)
_NB = 64 # panel / block width
NMAX = 4096
_ASM_FAST = r"""
{
.reg .f32 %asq, %full, %nfull, %beta, %amb, %sc, %rb, %tau, %zero;
.reg .pred %neg, %active;
mov.f32 %zero, 0f00000000;
fma.rn.f32 %asq, $3, $3, $4; // alpha^2 + sumsq
sqrt.approx.f32 %full, %asq; // ||x|| (SFU)
neg.f32 %nfull, %full;
setp.lt.f32 %neg, $3, %zero; // alpha < 0 ?
selp.f32 %beta, %full, %nfull, %neg; // beta = (alpha<0)? +||x|| : -||x||
sub.f32 %amb, $3, %beta; // alpha - beta
rcp.approx.f32 %sc, %amb; // scale = 1/(alpha-beta) (SFU)
rcp.approx.f32 %rb, %beta; // 1/beta (SFU)
mul.f32 %tau, %amb, %rb;
neg.f32 %tau, %tau; // tau = (beta-alpha)/beta = -(amb/beta)
setp.gt.f32 %active, $4, %zero; // sumsq > 0 ?
selp.f32 $0, %beta, $3, %active; // new_diag = active? beta : alpha
selp.f32 $1, %sc, %zero, %active; // scale = active? sc : 0
selp.f32 $2, %tau, %zero, %active; // tau = active? tau : 0
}
"""
@dsl_user_op
def householder_solve(alpha, sumsq, *, loc=None, ip=None):
f32 = Float32.mlir_type
res = llvm.inline_asm(
llvm.StructType.get_literal([f32, f32, f32]),
[alpha.ir_value(), sumsq.ir_value()],
_ASM_FAST,
"=f,=f,=f,f,f",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc, ip=ip,
)
return tuple(cutlass.Float32(llvm.extractvalue(f32, res, [i], loc=loc, ip=ip)) for i in range(3))
@dsl_user_op
def load4(ptr, loc=None, ip=None):
"""128-bit vectorized load: ptr -> (f0,f1,f2,f3). ptr must be 16B-aligned. The loaded
vector indexes directly, so no extractelement/wrapping is needed."""
v4 = ir.VectorType.get([4], Float32.mlir_type, loc=loc)
vv = cute.arch.load(ptr, v4, loc=loc, ip=ip)
return vv[0], vv[1], vv[2], vv[3]
@dsl_user_op
def store4(ptr, a, b, c, d, loc=None, ip=None):
"""128-bit vectorized store of 4 Float32 to ptr (must be 16B-aligned)."""
v4 = ir.VectorType.get([4], Float32.mlir_type, loc=loc)
vec = vector.from_elements(v4, [a.ir_value(), b.ir_value(), c.ir_value(), d.ir_value()], loc=loc, ip=ip)
cute.arch.store(ptr, vec, loc=loc, ip=ip)
class PanelQRGmem:
@cute.jit
def __call__(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
self.kernel(mH, mtau, mV, k, b, n).launch(
grid=[cute.size(mH, mode=[0]), 1, 1], block=[THREADS, 1, 1]
)
@cute.jit
def warp_reduce(self, val):
for i in range(5):
val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
return val
@cute.kernel
def kernel(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
tidx, _, _ = cute.arch.thread_idx()
bi, _, _ = cute.arch.block_idx()
warp_id = tidx >> 5
lane_id = tidx & 31
m = n - k
smem = cutlass.utils.SmemAllocator()
vsm = smem.allocate_tensor(cutlass.Float32, NMAX)
red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
wsm = smem.allocate_tensor(cutlass.Float32, _NB)
scal = smem.allocate_tensor(cutlass.Float32, 4)
for j in cutlass.range(0, b, 1, unroll=1):
col_len = m - j
base = k + j
ssq = cutlass.Float32(0.0)
r = tidx
if r == 0:
vsm[0] = mH[bi, base, base]
r += THREADS
while r < col_len:
x = mH[bi, base + r, base]
vsm[r] = x
ssq += x * x
r += THREADS
cute.arch.sync_threads()
ssq = self.warp_reduce(ssq)
if lane_id == 0:
red[warp_id] = ssq
cute.arch.sync_threads()
if warp_id == 0:
v0 = red[lane_id] if lane_id < WARPS else cutlass.Float32(0.0)
v0 = self.warp_reduce(v0)
if lane_id == 0:
red[WARPS] = v0
cute.arch.sync_threads()
sumsq = red[WARPS]
if tidx == 0:
alpha = vsm[0]
new_diag, scale, tau_j = householder_solve(alpha, sumsq)
scal[0] = new_diag
scal[1] = tau_j
scal[2] = scale
mtau[bi, base] = tau_j
cute.arch.sync_threads()
new_diag = scal[0]
tau_j = scal[1]
scale = scal[2]
r = tidx
if r == 0:
vsm[0] = cutlass.Float32(1.0)
mH[bi, base, base] = new_diag
mV[bi, j, j] = cutlass.Float32(1.0)
r += THREADS
while r < col_len:
vv = vsm[r] * scale
vsm[r] = vv
mH[bi, base + r, base] = vv
mV[bi, j + r, j] = vv
r += THREADS
r = tidx
while r < j:
mV[bi, r, j] = cutlass.Float32(0.0)
r += THREADS
cute.arch.sync_threads()
ncols = b - 1 - j
cc = tidx
while cc < ncols:
acc = cutlass.Float32(0.0)
r = 0
while r < col_len:
acc += vsm[r] * mH[bi, base + r, base + 1 + cc]
r += 1
wsm[cc] = acc
cc += THREADS
cute.arch.sync_threads()
r = tidx
while r < col_len:
vr = vsm[r]
row = base + r
cc = 0
while cc < ncols:
mH[bi, row, base + 1 + cc] -= tau_j * wsm[cc] * vr
cc += 1
r += THREADS
cute.arch.sync_threads()
THREADS = 256
WARPS = THREADS // 32
WPCTA = 8 # warps per CTA for the tiny T-from-G kernel
_PF_MAX = 57000 # max panel-cache floats (B200 ~232KB smem)
class PanelQRv2:
def __init__(self, panel_floats, nb, use_tpc=True, wide_red=False, vec=False, fast_update=False, fuse_update=False, use_regblock=False):
self.PF = panel_floats
self.NB = nb # actual panel width -> right-size wacc / wred / wsm / scales
self.use_tpc = use_tpc
self.wide_red = wide_red # TPC reduce: True -> 1 warp/col (conflict-free, n<=512); False -> 8 thr/col
self.vec = vec # 128-bit vectorized gmem staging/writeback (high-occupancy shapes
# only; at low batch the latency-bound panel sees vector-op overhead)
self.fast_update = fast_update # no-B3 per-warp column trailing update for small n where it wins
self.fuse_update = fuse_update # fused reduction+update: no wsm round-trip, no B3, wval in register
self.use_regblock = use_regblock # register-block 4 columns in fused tpc32 (1 row-pass vs 4)
@cute.jit
def __call__(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
self.kernel(mH, mtau, mV, k, b, n).launch(
grid=[cute.size(mH, mode=[0]), 1, 1], block=[THREADS, 1, 1]
)
@cute.jit
def warp_reduce(self, val):
for i in range(5):
val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
return val
@cute.jit
def _reduce_tpc8(self, sp, wsm, j, m, ncols, sbp, scale, tidx):
g = tidx // 8
gl = tidx % 8
cc = g
acc = cutlass.Float32(0.0)
if cc < ncols:
r = j + 1 + gl
while r < m:
acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
r += 8
for o in (1, 2, 4):
acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
if cc < ncols:
if gl == 0:
wsm[cc] = sp[j * sbp + j + 1 + cc] + scale * acc
@cute.jit
def _reduce_tpc32(self, sp, wsm, j, m, ncols, sbp, scale, tidx):
g = tidx // 32
gl = tidx % 32
cc = g
while cc < ncols:
acc = cutlass.Float32(0.0)
r = j + 1 + gl
while r < m:
acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
r += 32
for i in range(5): # butterfly offsets 1,2,4,8,16
acc += cute.arch.shuffle_sync_bfly(acc, offset=1 << i)
if gl == 0:
wsm[cc] = sp[j * sbp + j + 1 + cc] + scale * acc
cc += WARPS
@cute.jit
def _reduce_update_tpc32(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
g = tidx // 32
gl = tidx % 32
cc = g
tsc = tau_j * scale
while cc < ncols:
acc = cutlass.Float32(0.0)
r = j + 1 + gl
while r < m:
acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
r += 32
for i in range(5):
acc += cute.arch.shuffle_sync_bfly(acc, offset=1 << i)
wval = sp[j * sbp + j + 1 + cc] + scale * acc
if gl == 0:
sp[j * sbp + j + 1 + cc] -= tau_j * wval
r = j + 1 + gl
while r < m:
tvr = tsc * sp[r * sbp + j]
sp[r * sbp + j + 1 + cc] -= tvr * wval
r += 32
cc += WARPS
@cute.jit
def _reduce_update_tpc32_regblock(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
g = tidx // 32
gl = tidx % 32
tsc = tau_j * scale
cc = g
while cc + 24 < ncols:
cc0 = cc
cc1 = cc + 8
cc2 = cc + 16
cc3 = cc + 24
# --- accumulate 4 columns in one row pass ---
acc0 = cutlass.Float32(0.0)
acc1 = cutlass.Float32(0.0)
acc2 = cutlass.Float32(0.0)
acc3 = cutlass.Float32(0.0)
r = j + 1 + gl
while r < m:
vv = sp[r * sbp + j]
acc0 += vv * sp[r * sbp + j + 1 + cc0]
acc1 += vv * sp[r * sbp + j + 1 + cc1]
acc2 += vv * sp[r * sbp + j + 1 + cc2]
acc3 += vv * sp[r * sbp + j + 1 + cc3]
r += 32
for i in range(5):
acc0 += cute.arch.shuffle_sync_bfly(acc0, offset=1 << i)
acc1 += cute.arch.shuffle_sync_bfly(acc1, offset=1 << i)
acc2 += cute.arch.shuffle_sync_bfly(acc2, offset=1 << i)
acc3 += cute.arch.shuffle_sync_bfly(acc3, offset=1 << i)
wval0 = sp[j * sbp + j + 1 + cc0] + scale * acc0
wval1 = sp[j * sbp + j + 1 + cc1] + scale * acc1
wval2 = sp[j * sbp + j + 1 + cc2] + scale * acc2
wval3 = sp[j * sbp + j + 1 + cc3] + scale * acc3
if gl == 0:
sp[j * sbp + j + 1 + cc0] -= tau_j * wval0
sp[j * sbp + j + 1 + cc1] -= tau_j * wval1
sp[j * sbp + j + 1 + cc2] -= tau_j * wval2
sp[j * sbp + j + 1 + cc3] -= tau_j * wval3
r = j + 1 + gl
while r < m:
tvr = tsc * sp[r * sbp + j]
sp[r * sbp + j + 1 + cc0] -= tvr * wval0
sp[r * sbp + j + 1 + cc1] -= tvr * wval1
sp[r * sbp + j + 1 + cc2] -= tvr * wval2
sp[r * sbp + j + 1 + cc3] -= tvr * wval3
r += 32
cc += 32
# Tail: remaining 0-3 columns, single-column fallback
while cc < ncols:
acc = cutlass.Float32(0.0)
r = j + 1 + gl
while r < m:
acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
r += 32
for o in (1, 2, 4, 8, 16):
acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
wval = sp[j * sbp + j + 1 + cc] + scale * acc
if gl == 0:
sp[j * sbp + j + 1 + cc] -= tau_j * wval
r = j + 1 + gl
while r < m:
tvr = tsc * sp[r * sbp + j]
sp[r * sbp + j + 1 + cc] -= tvr * wval
r += 32
cc += 8
@cute.jit
def _reduce_update_tpc8(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
g = tidx // 8
gl = tidx % 8
cc = g
tsc = tau_j * scale
acc = cutlass.Float32(0.0)
if cc < ncols:
r = j + 1 + gl
while r < m:
acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
r += 8
for o in (1, 2, 4):
acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
if cc < ncols:
wval = sp[j * sbp + j + 1 + cc] + scale * acc
if gl == 0:
sp[j * sbp + j + 1 + cc] -= tau_j * wval
r = j + 1 + gl
while r < m:
tvr = tsc * sp[r * sbp + j]
sp[r * sbp + j + 1 + cc] -= tvr * wval
r += 8
@cute.jit
def _reduce_wr(self, sp, wsm, wred, j, m, ncols, sbp, scale, tidx, warp_id, lane_id):
wacc = cute.make_rmem_tensor(self.NB, cutlass.Float32)
r = tidx
while r < m:
if r >= j:
vr = cutlass.Float32(1.0) if r == j else scale * sp[r * sbp + j] # scale inline
base = r * sbp + j + 1
cc = 0
while cc < ncols:
wacc[cc] += vr * sp[base + cc]
cc += 1
r += THREADS
cc = 0
while cc < ncols:
wv = self.warp_reduce(wacc[cc])
if lane_id == 0:
wred[warp_id * self.NB + cc] = wv
cc += 1
if warp_id == 0:
cc = lane_id
while cc < ncols:
acc = cutlass.Float32(0.0)
w = 0
while w < WARPS:
acc += wred[w * self.NB + cc]
w += 1
wsm[cc] = acc
cc += 32
@cute.kernel
def kernel(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
tidx, _, _ = cute.arch.thread_idx()
bi, _, _ = cute.arch.block_idx()
warp_id = tidx >> 5
lane_id = tidx & 31
m = n - k
sbp = b + 1
smem = cutlass.utils.SmemAllocator()
sp = smem.allocate_tensor(cutlass.Float32, self.PF)
red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
wsm = smem.allocate_tensor(cutlass.Float32, self.NB)
scales = smem.allocate_tensor(cutlass.Float32, self.NB) # per-column Householder scale
if cutlass.const_expr(not self.use_tpc):
wred = smem.allocate_tensor(cutlass.Float32, WARPS * self.NB)
tot = m * b # used by the writeback loop at the end
b4 = b >> 2
rem = b - (b4 << 2) # leftover cols (0..3) when b not a multiple of 4
# ---- stage the m x b panel into smem ----
if cutlass.const_expr(self.vec):
idx = tidx
tot4 = m * b4
while idx < tot4:
r = idx // b4
c = (idx - r * b4) << 2
f0, f1, f2, f3 = load4(cute.domain_offset((bi, k + r, k + c), mH).iterator)
base = r * sbp + c
sp[base] = f0
sp[base + 1] = f1
sp[base + 2] = f2
sp[base + 3] = f3
idx += THREADS
if rem > 0:
c0 = b4 << 2
idx = tidx
totr = m * rem
while idx < totr:
r = idx // rem
c = c0 + (idx - r * rem)
sp[r * sbp + c] = mH[bi, k + r, k + c]
idx += THREADS
else:
idx = tidx
while idx < tot:
r = idx // b
c = idx - r * b
sp[r * sbp + c] = mH[bi, k + r, k + c]
idx += THREADS
cute.arch.sync_threads()
for j in cutlass.range(0, b, 1, unroll=1):
# ---- column norm (single barrier) -------------------------
alpha = sp[j * sbp + j]
ssq = cutlass.Float32(0.0)
r = tidx
while r < m:
if r > j:
x = sp[r * sbp + j]
ssq += x * x
r += THREADS
ssq = self.warp_reduce(ssq)
if lane_id == 0:
red[warp_id] = ssq
cute.arch.sync_threads() # B1
sumsq = cutlass.Float32(0.0)
w = 0
while w < WARPS:
sumsq += red[w]
w += 1
new_diag, scale, tau_j = householder_solve(alpha, sumsq)
if tidx == 0:
scales[j] = scale
sp[j * sbp + j] = new_diag
mtau[bi, k + j] = tau_j
# ---- w = v^T C (reduction) and trailing update (fused) -----------------
ncols = b - 1 - j
if cutlass.const_expr(self.fuse_update):
if cutlass.const_expr(self.use_regblock):
self._reduce_update_tpc32_regblock(sp, j, m, ncols, sbp, scale, tau_j, tidx)
elif cutlass.const_expr(self.wide_red):
self._reduce_update_tpc32(sp, j, m, ncols, sbp, scale, tau_j, tidx)
else:
self._reduce_update_tpc8(sp, j, m, ncols, sbp, scale, tau_j, tidx)
elif cutlass.const_expr(self.use_tpc):
if cutlass.const_expr(self.wide_red):
self._reduce_tpc32(sp, wsm, j, m, ncols, sbp, scale, tidx)
else:
self._reduce_tpc8(sp, wsm, j, m, ncols, sbp, scale, tidx)
if cutlass.const_expr(self.fast_update):
g = warp_id
gl = lane_id
tsc = tau_j * scale
col = g
while col < ncols:
wval = wsm[col]
if gl == 0:
sp[j * sbp + j + 1 + col] -= tau_j * wval
r = j + 1 + gl
while r < m:
tvr = tsc * sp[r * sbp + j]
sp[r * sbp + j + 1 + col] -= tvr * wval
r += 32
col += WARPS
else:
cute.arch.sync_threads() # B3
tsc = tau_j * scale
r = tidx
while r < m:
if r >= j:
tvr = tau_j if r == j else tsc * sp[r * sbp + j]
base = r * sbp + j + 1
cc = 0
while cc + 3 < ncols:
w0 = wsm[cc]; w1 = wsm[cc+1]; w2 = wsm[cc+2]; w3 = wsm[cc+3]
sp[base + cc] -= tvr * w0
sp[base + cc + 1] -= tvr * w1
sp[base + cc + 2] -= tvr * w2
sp[base + cc + 3] -= tvr * w3
cc += 4
while cc < ncols:
sp[base + cc] -= tvr * wsm[cc]
cc += 1
r += THREADS
else:
self._reduce_wr(sp, wsm, wred, j, m, ncols, sbp, scale, tidx, warp_id, lane_id)
cute.arch.sync_threads() # B3
tsc = tau_j * scale
r = tidx
while r < m:
if r >= j:
tvr = tau_j if r == j else tsc * sp[r * sbp + j]
base = r * sbp + j + 1
cc = 0
while cc + 3 < ncols:
w0 = wsm[cc]; w1 = wsm[cc+1]; w2 = wsm[cc+2]; w3 = wsm[cc+3]
sp[base + cc] -= tvr * w0
sp[base + cc + 1] -= tvr * w1
sp[base + cc + 2] -= tvr * w2
sp[base + cc + 3] -= tvr * w3
cc += 4
while cc < ncols:
sp[base + cc] -= tvr * wsm[cc]
cc += 1
r += THREADS
cute.arch.sync_threads() # B4
# ---- writeback: apply the deferred Householder scale to the stored v's ----
if cutlass.const_expr(self.vec):
nbr = m - b
if nbr > 0:
tot4 = nbr * b4
idx = tidx
while idx < tot4:
rr = idx // b4
r = b + rr
c = (idx - rr * b4) << 2
base = r * sbp + c
v0 = sp[base] * scales[c]
v1 = sp[base + 1] * scales[c + 1]
v2 = sp[base + 2] * scales[c + 2]
v3 = sp[base + 3] * scales[c + 3]
store4(cute.domain_offset((bi, k + r, k + c), mH).iterator, v0, v1, v2, v3)
store4(cute.domain_offset((bi, r, c), mV).iterator, v0, v1, v2, v3)
idx += THREADS
if rem > 0:
c0 = b4 << 2
idx = tidx
totr = nbr * rem
while idx < totr:
rr = idx // rem
r = b + rr
c = c0 + (idx - rr * rem)
val = sp[r * sbp + c] * scales[c]
mH[bi, k + r, k + c] = val
mV[bi, r, c] = val
idx += THREADS
# top b x b block (r < b): diagonal / above-diagonal -> scalar branchy path
idx = tidx
while idx < b * b:
r = idx // b
c = idx - r * b
val = sp[r * sbp + c]
if r > c:
mH[bi, k + r, k + c] = val * scales[c]
mV[bi, r, c] = val * scales[c]
elif r == c:
mH[bi, k + r, k + c] = val # beta (R diagonal)
mV[bi, r, c] = cutlass.Float32(1.0)
else:
mH[bi, k + r, k + c] = val # R (above diagonal)
mV[bi, r, c] = cutlass.Float32(0.0)
idx += THREADS
else:
idx = tidx
while idx < tot:
r = idx // b
c = idx - r * b
val = sp[r * sbp + c]
if r > c:
val = val * scales[c]
mH[bi, k + r, k + c] = val
mV[bi, r, c] = val
elif r == c:
mH[bi, k + r, k + c] = val # beta (R diagonal)
mV[bi, r, c] = cutlass.Float32(1.0)
else:
mH[bi, k + r, k + c] = val # R (above diagonal)
mV[bi, r, c] = cutlass.Float32(0.0)
idx += THREADS
class PanelQRCoopSmem:
def __init__(self, cpm, pf, nb):
self.CPM = cpm
self.PF = pf
self.NB = nb
@cute.jit
def __call__(self, mH, mtau, mV, mpart, mbar, k: cutlass.Int32, b: cutlass.Int32,
n: cutlass.Int32, base: cutlass.Int32):
batch = cute.size(mH, mode=[0])
grid = batch * self.CPM
self.kernel(mH, mtau, mV, mpart, mbar, k, b, n, base, batch).launch(
grid=[grid, 1, 1], block=[THREADS, 1, 1], cooperative=True
)
@cute.jit
def warp_reduce(self, val):
for i in range(5):
val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
return val
@cute.jit
def cta_reduce(self, val, red, warp_id, lane_id):
val = self.warp_reduce(val)
if lane_id == 0:
red[warp_id] = val
cute.arch.sync_threads()
if warp_id == 0:
v0 = red[lane_id] if lane_id < WARPS else cutlass.Float32(0.0)
v0 = self.warp_reduce(v0)
if lane_id == 0:
red[WARPS] = v0
cute.arch.sync_threads()
return red[WARPS]
@cute.jit
def gbar(self, mbar, slot, total, tidx):
cute.arch.sync_threads()
cute.arch.fence_acq_rel_gpu()
if tidx == 0:
cute.arch.atomic_add(mbar.iterator + slot, cutlass.Int32(1), sem="release", scope="gpu")
done = cutlass.Int32(0)
while done == 0:
# acquire LOAD (not atom.add 0): a read doesn't serialize on the L2 atomic unit.
cur = cute.arch.load(mbar.iterator + slot, cutlass.Int32, sem="acquire", scope="gpu")
if cur >= total:
done = cutlass.Int32(1)
cute.arch.sync_threads()
cute.arch.fence_acq_rel_gpu()
@cute.kernel
def kernel(self, mH, mtau, mV, mpart, mbar, k: cutlass.Int32, b: cutlass.Int32,
n: cutlass.Int32, base: cutlass.Int32, batch: cutlass.Int32):
tidx, _, _ = cute.arch.thread_idx()
g, _, _ = cute.arch.block_idx()
bi = g // self.CPM
sub = g - bi * self.CPM
warp_id = tidx >> 5
lane_id = tidx & 31
m = n - k
sbp = b + 1
total = batch * self.CPM
rpc = (m + self.CPM - 1) // self.CPM
r0 = sub * rpc
rend = r0 + rpc
if rend > m:
rend = m
nloc = rend - r0
if nloc < 0:
nloc = cutlass.Int32(0)
smem = cutlass.utils.SmemAllocator()
sp = smem.allocate_tensor(cutlass.Float32, self.PF)
red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
wred = smem.allocate_tensor(cutlass.Float32, WARPS * self.NB)
wsm = smem.allocate_tensor(cutlass.Float32, self.NB)
tot = nloc * b
idx = tidx
while idx < tot:
rr = idx // b
c = idx - rr * b
sp[rr * sbp + c] = mH[bi, k + r0 + rr, k + c]
idx += THREADS
cute.arch.sync_threads()
for j in cutlass.range(0, b, 1, unroll=1):
owner = j // rpc
ncols = b - 1 - j
# ---- local norm AND w_raw (off-diagonal, unscaled v_raw) ----
ssq = cutlass.Float32(0.0)
wacc = cute.make_fragment(self.NB, cutlass.Float32)
for ci in range(self.NB):
wacc[ci] = cutlass.Float32(0.0)
rr = tidx
while rr < nloc:
gr = r0 + rr
if gr > j:
v_raw = sp[rr * sbp + j]
ssq += v_raw * v_raw
bcol = rr * sbp + j + 1
cc = 0
while cc < ncols:
wacc[cc] += v_raw * sp[bcol + cc]
cc += 1
rr += THREADS
ssq = self.cta_reduce(ssq, red, warp_id, lane_id)
# w_raw warp-reduce + cross-warp combine
for ci in range(self.NB):
wv = self.warp_reduce(wacc[ci])
if lane_id == 0:
wred[warp_id * self.NB + ci] = wv
cute.arch.sync_threads()
if warp_id == 0:
ci = lane_id
while ci < self.NB:
acc = cutlass.Float32(0.0)
w = 0
while w < WARPS:
acc += wred[w * self.NB + ci]
w += 1
if ci < ncols:
mpart[bi, sub, 1 + ci] = acc
ci += 32
cute.arch.sync_threads()
# ---- write to mpart and single gbar ----
if tidx == 0:
mpart[bi, sub, 0] = ssq
if sub == owner:
mpart[bi, sub, self.NB + 1] = sp[(j - r0) * sbp + j]
# store C[j, j+1+cc] for later w combination
ci = 0
while ci < ncols:
mpart[bi, sub, self.NB + 2 + ci] = sp[(j - r0) * sbp + j + 1 + ci]
ci += 1
self.gbar(mbar, base + j, total, tidx) # single barrier (was 2)
# ---- after barrier: combine ssq + w_raw, compute tau/scale/w, update ----
sumsq = cutlass.Float32(0.0)
s = 0
while s < self.CPM:
sumsq += mpart[bi, s, 0]
s += 1
alpha = mpart[bi, owner, self.NB + 1]
new_diag, scale, tau_j = householder_solve(alpha, sumsq)
cc = tidx
while cc < ncols:
w_raw_global = cutlass.Float32(0.0)
s = 0
while s < self.CPM:
w_raw_global += mpart[bi, s, 1 + cc]
s += 1
c_diag = mpart[bi, owner, self.NB + 2 + cc]
wsm[cc] = c_diag + scale * w_raw_global
cc += THREADS
cute.arch.sync_threads()
# ---- scale v in-place ----
rr = tidx
while rr < nloc:
gr = r0 + rr
if gr > j:
sp[rr * sbp + j] = sp[rr * sbp + j] * scale
elif gr == j:
sp[rr * sbp + j] = new_diag
rr += THREADS
if sub == 0 and tidx == 0:
mtau[bi, k + j] = tau_j
cute.arch.sync_threads()
# ---- trailing update ----
rr = tidx
while rr < nloc:
gr = r0 + rr
if gr >= j:
vr = cutlass.Float32(1.0) if gr == j else sp[rr * sbp + j]
tvr = tau_j * vr
bcol = rr * sbp + j + 1
cc = 0
while cc < ncols:
sp[bcol + cc] -= tvr * wsm[cc]
cc += 1
rr += THREADS
cute.arch.sync_threads()
idx = tidx
while idx < tot:
rr = idx // b
c = idx - rr * b
gr = r0 + rr
val = sp[rr * sbp + c]
if gr > c:
mH[bi, k + gr, k + c] = val
mV[bi, gr, c] = val
elif gr == c:
mH[bi, k + gr, k + c] = val
mV[bi, gr, c] = cutlass.Float32(1.0)
else:
mH[bi, k + gr, k + c] = val
mV[bi, gr, c] = cutlass.Float32(0.0)
idx += THREADS
class TFromG:
def __init__(self, b):
self.B = b
@cute.jit
def __call__(self, mG, mtau, mT, k: cutlass.Int32, b: cutlass.Int32):
batch = cute.size(mG, mode=[0])
grid = (batch + WPCTA - 1) // WPCTA
self.kernel(mG, mtau, mT, k, b, batch).launch(grid=[grid, 1, 1], block=[WPCTA * 32, 1, 1])
@cute.kernel
def kernel(self, mG, mtau, mT,
k: cutlass.Int32,
b: cutlass.Int32,
batch: cutlass.Int32):
B = self.B
sbp = B + 1
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
warp = tidx >> 5
lane = tidx & 31
bi = bidx * WPCTA + warp
smem = cutlass.utils.SmemAllocator()
gsm = smem.allocate_tensor(cutlass.Float32, WPCTA * B * sbp)
tsm = smem.allocate_tensor(cutlass.Float32, WPCTA * B * sbp)
taus = smem.allocate_tensor(cutlass.Float32, WPCTA * B)
if bi < batch:
base = warp * B * sbp
tau_base = warp * B
#
# Stage tau once
#
if lane < b:
taus[tau_base + lane] = mtau[bi, k + lane]
#
# Vectorized stage of G
#
b4 = b >> 2
rem = b - (b4 << 2)
idx = lane
tot4 = b * b4
while idx < tot4:
r = idx // b4
c = (idx - r * b4) << 2
f0, f1, f2, f3 = load4(
cute.domain_offset((bi, r, c), mG).iterator
)
row = base + r * sbp + c
gsm[row] = f0
gsm[row + 1] = f1
gsm[row + 2] = f2
gsm[row + 3] = f3
tsm[row] = cutlass.Float32(0.0)
tsm[row + 1] = cutlass.Float32(0.0)
tsm[row + 2] = cutlass.Float32(0.0)
tsm[row + 3] = cutlass.Float32(0.0)
idx += 32
if rem > 0:
c0 = b4 << 2
idx = lane
totr = b * rem
while idx < totr:
r = idx // rem
c = c0 + (idx - r * rem)
row = base + r * sbp + c
gsm[row] = mG[bi, r, c]
tsm[row] = cutlass.Float32(0.0)
idx += 32
if lane == 0:
tsm[base] = taus[tau_base]
cute.arch.sync_warp()
#
# WY recurrence
#
i = 1
while i < b:
tau_i = taus[tau_base + i]
r = lane
while r < i:
grow = base + r * sbp
trow = grow
tau_r = taus[tau_base + r]
acc = tau_r * gsm[grow + i]
c = r + 1
while c + 3 < i:
acc += tsm[trow + c] * gsm[base + (c ) * sbp + i]
acc += tsm[trow + c + 1] * gsm[base + (c + 1) * sbp + i]
acc += tsm[trow + c + 2] * gsm[base + (c + 2) * sbp + i]
acc += tsm[trow + c + 3] * gsm[base + (c + 3) * sbp + i]
c += 4
while c < i:
acc += tsm[trow + c] * gsm[base + c * sbp + i]
c += 1
tsm[trow + i] = -tau_i * acc
r += 32
if lane == 0:
tsm[base + i * sbp + i] = tau_i
i += 1
cute.arch.sync_warp()
#
# Vectorized writeback
#
idx = lane
while idx < tot4:
r = idx // b4
c = (idx - r * b4) << 2
row = base + r * sbp + c
store4(
cute.domain_offset((bi, r, c), mT).iterator,
tsm[row],
tsm[row + 1],
tsm[row + 2],
tsm[row + 3],
)
idx += 32
if rem > 0:
c0 = b4 << 2
idx = lane
totr = b * rem
while idx < totr:
r = idx // rem
c = c0 + (idx - r * rem)
mT[bi, r, c] = tsm[base + r * sbp + c]
idx += 32
def _make_cute(t):
"""Wrap a (M, K, L) / (M, N, L) batch-last torch tensor as an fp32 cute tensor."""
ct = from_dlpack(t, assumed_align=16)
ct.element_type = cutlass.Float32
ct = ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
return ct
def _make_cute_i(t):
ct = from_dlpack(t, assumed_align=16)
ct.element_type = cutlass.Int32
ct = ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
return ct
_panel_cache = {}
def _panel(cH, ctau, cV, k, b, n, pf, nb, use_tpc=True, wide_red=False, vec=False, fast_update=False, fuse_update=False, use_regblock=False):
ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n)
key = (pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock)
fn = _panel_cache.get(key)
if fn is None:
fn = cute.compile(PanelQRv2(pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock), cH, ctau, cV, ck, cb, cn)
_panel_cache[key] = fn
fn(cH, ctau, cV, ck, cb, cn)
_coop_smem_cache = {}
def _panel_coop_smem(H, tau, V, part, bar, cpm, pf, nb, k, b, n, base):
ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n); cbase = cutlass.Int32(base)
key = (cpm, pf, nb)
fn = _coop_smem_cache.get(key)
if fn is None:
fn = cute.compile(PanelQRCoopSmem(cpm, pf, nb), _make_cute(H), _make_cute(tau), _make_cute(V),
_make_cute(part), _make_cute_i(bar), ck, cb, cn, cbase)
_coop_smem_cache[key] = fn
fn(_make_cute(H), _make_cute(tau), _make_cute(V), _make_cute(part), _make_cute_i(bar),
ck, cb, cn, cbase)
_COOP_NB = 32
def _coop_config(n):
if n not in (2048, 4096):
return None
rpc = 128 if n == 2048 else 256
nb = _COOP_NB
cpm = max(1, (n + rpc - 1) // rpc)
if rpc * (nb + 1) > _PF_MAX or rpc < nb:
return None
return cpm, nb
_tfromg_cache = {}
def _tfromg(cG, ctau, cT, k, b):
ck = cutlass.Int32(k); cb = cutlass.Int32(b)
fn = _tfromg_cache.get(b) # G is now b×b (varying layout) -> key per b; b -> smem size
if fn is None:
fn = cute.compile(TFromG(b), cG, ctau, cT, ck, cb)
_tfromg_cache[b] = fn
fn(cG, ctau, cT, ck, cb)
_panel_gmem_cache = {}
def _panel_gmem(mH, mtau, mV, k, b, n):
ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n)
key = (mH.shape[0], n)
fn = _panel_gmem_cache.get(key)
if fn is None:
fn = cute.compile(
PanelQRGmem(), _make_cute(mH), _make_cute(mtau), _make_cute(mV), ck, cb, cn
)
_panel_gmem_cache[key] = fn
fn(_make_cute(mH), _make_cute(mtau), _make_cute(mV), ck, cb, cn)
def _choose_nb(n):
for nb in (32, 16, 8):
if nb <= _NB and n * (nb + 1) <= _PF_MAX:
return nb
return 4
def custom_kernel(data: input_t) -> output_t:
A = data
batch, n, _ = A.shape
#---Set up branches and flags for various shapes---
# Large low-batch shapes need different CTA split
coop_cfg = _coop_config(n)
# Always use SMEM when size allows
cached = (n * (16 + 1) <= _PF_MAX)
# Small panels do differently depending on amount of threads participating in reduction
use_tpc = n <= 2048
wide_red = (n <= 512) or (n == 2048)
# Large batches benefit from vectorized loads, small don't somehow
vec = batch >= 48
# Use SMEM to do local rank-1 updates without sync with other warps, works for small shapes
fast_update = n <= 352
# Fused reduction+update: wval in RMEM
fuse_update = wide_red
use_regblock = fuse_update # all fused-tpc32 shapes use register blocking
if coop_cfg is not None:
coop_cpm, nb = coop_cfg
coop_pf = ((n + coop_cpm - 1) // coop_cpm) * (nb + 1)
elif cached:
nb = _choose_nb(n)
pf = n * (nb + 1)
else:
nb = 32
pf = 0
#---Prep tensors---
H = A.clone()
tau = torch.empty(batch, n, device=A.device, dtype=torch.float32) # panel writes all of tau
Vbuf = torch.zeros(batch, n, _NB, device=A.device, dtype=torch.float32)
Tbuf = torch.zeros(batch, _NB, _NB, device=A.device, dtype=torch.float32)
if coop_cfg is not None:
coop_part = torch.zeros(batch, coop_cpm, 2 * _NB + 2, device=A.device, dtype=torch.float32)
coop_bar = torch.zeros(n + 1, device=A.device, dtype=torch.int32)
torch.backends.cuda.matmul.allow_tf32 = False
#---Compute loop---
k = 0
while k < n:
b = min(nb, n - k)
m = n - k
if coop_cfg is not None:
_panel_coop_smem(H, tau, Vbuf, coop_part, coop_bar, coop_cpm, coop_pf, nb,
k, b, n, k)
elif cached:
_panel(H, tau, Vbuf, k, b, n, pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock)
else:
_panel_gmem(H, tau, Vbuf, k, b, n)
if k + b < n:
V = Vbuf[:, :m, :b]
G = torch.matmul(V.transpose(-1, -2), V) # (l,b,b) FP32
_tfromg(G, tau, Tbuf[:, :b, :b], k, b)
C0 = H[:, k:, k + b:]
T = Tbuf[:, :b, :b]
torch.backends.cuda.matmul.allow_tf32 = True
W = torch.matmul(V.transpose(-1, -2), C0) # (l,b,rest) — W only needs TF32x1
torch.backends.cuda.matmul.allow_tf32 = False
Z = torch.matmul(T.transpose(-1, -2), W) # (l,b,rest) FP32
C0.baddbmm_(V, Z, beta=1, alpha=-1)
k += b #Move to next panel
torch.backends.cuda.matmul.allow_tf32 = True
return H, tau
scrolls · 1101 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