submission 843441
Subho Ghosh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1084 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-843441?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:03592dacd34bf7c56a34dbcc32d074636444a25d4fac2110147a97101d6e2c3f
license declaredunknown
license concludedunknown
authorsSubho Ghosh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
def _panel_cluster_launch(mbarrier
"st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"shared-memory
def _set_block_rank(smem_ptr, peer, *, loc=None, ip=None):Kernel source
submission.py1084 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import operator
import torch
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32
from cutlass._mlir.dialects import llvm
from cutlass.cute.runtime import from_dlpack
from cutlass.cutlass_dsl import T, dsl_user_op
from task import input_t, output_t
_NB = 32
_CN = 64
_TPB = 512
_PANEL_NW = _TPB // 32
_GEQR2_TPB = 256
_GEQR2_TDIM = 16
_PC_TPB = 512
_PC_NW = _PC_TPB // 32
_PC_C = 8
def _t2c(t, align=32):
return from_dlpack(t, assumed_align=align)
@cute.jit
def _block_sum(val, red, warp, lane):
NW = cutlass.const_expr(cute.size(red))
v = cute.arch.warp_reduction(val, operator.add)
if lane == 0:
red[warp] = v
cute.arch.barrier()
total = cutlass.Float32(0.0)
for w in cutlass.range_constexpr(NW):
total = total + red[w]
cute.arch.barrier()
return total
@dsl_user_op
def _set_block_rank(smem_ptr, peer, *, loc=None, ip=None):
p = smem_ptr.toint(loc=loc, ip=ip).ir_value()
return Int32(
llvm.inline_asm(
T.i32(), [p, peer.ir_value()],
"mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r",
has_side_effects=False, is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def _load_remote(smem_ptr, peer, *, loc=None, ip=None):
addr = _set_block_rank(smem_ptr, peer, loc=loc, ip=ip).ir_value()
return Float32(
llvm.inline_asm(
T.f32(), [addr],
"ld.shared::cluster.f32 $0, [$1];", "=f,r",
has_side_effects=True, is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def _store_remote_v4(
v0: Float32,
v1: Float32,
v2: Float32,
v3: Float32,
smem_ptr,
mbar_ptr,
peer,
*,
loc=None,
ip=None,
):
dst = _set_block_rank(smem_ptr, peer, loc=loc, ip=ip).ir_value()
bar = _set_block_rank(mbar_ptr, peer, loc=loc, ip=ip).ir_value()
llvm.inline_asm(
None,
[
dst,
bar,
v0.ir_value(loc=loc, ip=ip),
v1.ir_value(loc=loc, ip=ip),
v2.ir_value(loc=loc, ip=ip),
v3.ir_value(loc=loc, ip=ip),
],
"{\n\t"
".reg .v4 .f32 abcd;\n\t"
"mov.f32 abcd.x, $2;\n\t"
"mov.f32 abcd.y, $3;\n\t"
"mov.f32 abcd.z, $4;\n\t"
"mov.f32 abcd.w, $5;\n\t"
"st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"
"}\n",
"r,r,f,f,f,f",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@cute.jit
def _cluster_vsum(redbuf, total, cnt, rank, tidx, C: cutlass.Constexpr):
cute.arch.barrier()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
for v in cutlass.range(tidx, cnt, _PC_TPB):
s = redbuf[v]
for p in cutlass.range_constexpr(C):
if p != rank:
s = s + _load_remote(redbuf.iterator + v, Int32(p))
total[v] = s
cute.arch.barrier()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
@cute.jit
def _cluster_vsum_db(redbuf, total, buf, first, rank, tidx, C: cutlass.Constexpr):
cute.arch.barrier()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
for q in cutlass.range(tidx, 2 * _NB, _PC_TPB):
v = buf + q
s = redbuf[v]
for p in cutlass.range_constexpr(C):
if p != rank:
s = s + _load_remote(redbuf.iterator + v, Int32(p))
total[v] = s
cute.arch.barrier()
@cute.jit
def _cluster_vsum_async(
redbuf, total, recv, mbars, buf, jj, rank, tidx,
C: cutlass.Constexpr, TPB: cutlass.Constexpr,
):
cute.arch.barrier()
slot = jj & 1
if tidx < 16:
q = 4 * tidx
dst = recv.iterator + (slot * C + rank) * (2 * _NB) + q
for peer in cutlass.range_constexpr(C):
_store_remote_v4(
redbuf[buf + q],
redbuf[buf + q + 1],
redbuf[buf + q + 2],
redbuf[buf + q + 3],
dst,
mbars + slot,
Int32(peer),
)
cute.arch.mbarrier_wait(mbars + slot, (jj // 2) & 1)
for q in cutlass.range(tidx, 2 * _NB, TPB):
s = cutlass.Float32(0.0)
for peer in cutlass.range_constexpr(C):
s = s + recv[slot, peer, q]
total[buf + q] = s
cute.arch.barrier()
if tidx == 0 and jj + 2 < _NB:
cute.arch.mbarrier_arrive_and_expect_tx(
mbars + slot, C * (2 * _NB) * 4
)
@cute.kernel
def _geqr2_small_kernel(
mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor,
tiled_copy: cute.TiledCopy, rows: cutlass.Constexpr,
cols: cutlass.Constexpr, TPB: cutlass.Constexpr,
NW: cutlass.Constexpr,
):
tidx, _, _ = cute.arch.thread_idx()
b, _, _ = cute.arch.block_idx()
warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
lane = cute.arch.lane_idx()
smem = cutlass.utils.SmemAllocator()
sH = smem.allocate_tensor(
cutlass.Float32,
cute.make_layout((rows, cols), stride=(cols + 1, 1)),
byte_alignment=16,
)
red = smem.allocate_tensor(cutlass.Float32, cute.make_layout(NW), byte_alignment=16)
gA = mA[b, None, None]
gH = mH[b, None, None]
tc = tiled_copy.get_slice(tidx)
if cutlass.const_expr(rows == cols):
g_src = tc.partition_S(gA)
frag = cute.make_rmem_tensor_like(g_src)
cute.copy(tiled_copy, g_src, frag)
cute.copy(tiled_copy, frag, tc.partition_D(sH))
else:
for idx in cutlass.range(tidx, rows * cols, TPB):
sH[idx // cols, idx % cols] = gA[idx // cols, idx % cols]
cute.arch.barrier()
for j in cutlass.range(0, cols):
m = rows - j
p = cols - j
sT = cute.domain_offset((j, j), sH)
local = cutlass.Float32(0.0)
for i in cutlass.range(1 + tidx, m, TPB):
x = sT[i, 0]
local = local + x * x
xnorm2 = _block_sum(local, red, warp, lane)
alpha = sT[0, 0]
nrm = cute.math.sqrt(alpha * alpha + xnorm2)
beta = nrm
if alpha >= 0.0:
beta = -nrm
tau_j = cutlass.Float32(0.0)
scale = cutlass.Float32(0.0)
beta_diag = alpha
if xnorm2 > 0.0:
tau_j = (beta - alpha) / beta
scale = 1.0 / (alpha - beta)
beta_diag = beta
if tidx == 0:
sT[0, 0] = beta_diag
mTau[b, j] = tau_j
for i in cutlass.range(1 + tidx, m, TPB):
sT[i, 0] = sT[i, 0] * scale
cute.arch.barrier()
for c in cutlass.range(1 + warp, p, NW):
dot = cutlass.Float32(0.0)
for i in cutlass.range(1 + lane, m, 32):
dot = dot + sT[i, 0] * sT[i, c]
dot = cute.arch.warp_reduction(dot, operator.add)
tw = tau_j * (dot + sT[0, c])
if lane == 0:
sT[0, c] = sT[0, c] - tw
for i in cutlass.range(1 + lane, m, 32):
sT[i, c] = sT[i, c] - sT[i, 0] * tw
cute.arch.barrier()
if cutlass.const_expr(rows == cols):
s_src = tc.partition_S(sH)
frag2 = cute.make_rmem_tensor_like(s_src)
cute.copy(tiled_copy, s_src, frag2)
cute.copy(tiled_copy, frag2, tc.partition_D(gH))
else:
for idx in cutlass.range(tidx, rows * cols, TPB):
gH[idx // cols, idx % cols] = sH[idx // cols, idx % cols]
@cute.jit
def _geqr2_launch(mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor):
rows = mA.shape[1]
cols = mA.shape[2]
tpb = 1024 if rows == 128 and cols == 128 else (512 if rows >= 64 else 256)
TR = 32 if tpb == 1024 else _GEQR2_TDIM
TC = 32 if tpb >= 512 else _GEQR2_TDIM
VR = cutlass.const_expr(rows // TR)
VC = cutlass.const_expr(cols // TC)
thr = cute.make_ordered_layout((TR, TC), order=(1, 0))
val = cute.make_ordered_layout((VR, VC), order=(1, 0))
atom = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(), cutlass.Float32, num_bits_per_copy=32
)
tiled_copy = cute.make_tiled_copy_tv(atom, thr, val)
_geqr2_small_kernel(mA, mH, mTau, tiled_copy, rows, cols, tpb, tpb // 32).launch(
grid=[mA.shape[0], 1, 1], block=[tpb, 1, 1]
)
_geqr2_cache: dict = {}
_tail_cache: dict = {}
def geqr2_small(A: torch.Tensor):
batch, n, _ = A.shape
H = torch.empty_like(A)
tau = torch.empty((batch, n), device=A.device, dtype=torch.float32)
key = (batch, n)
if key not in _geqr2_cache:
_geqr2_cache[key] = cute.compile(
_geqr2_launch, _t2c(A, 16), _t2c(H, 16), _t2c(tau, 16)
)
_geqr2_cache[key](_t2c(A, 16), _t2c(H, 16), _t2c(tau, 16))
return H, tau
def _geqr2_tail(H: torch.Tensor, tau: torch.Tensor, start: int, stop: int):
tail = H[:, start:, start:stop]
tail_tau = tau[:, start:stop]
key = (H.shape[0], H.shape[1], start, stop)
mTail = _t2c(tail, 16)
mTau = _t2c(tail_tau, 16)
if key not in _tail_cache:
_tail_cache[key] = cute.compile(_geqr2_launch, mTail, mTail, mTau)
_tail_cache[key](mTail, mTail, mTau)
@cute.kernel
def _panel_resident_kernel(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mV: cute.Tensor,
k: cutlass.Int32, n: cutlass.Constexpr,
form_t: cutlass.Constexpr, TPB: cutlass.Constexpr,
NW: cutlass.Constexpr,
):
tidx, _, _ = cute.arch.thread_idx()
b, _, _ = cute.arch.block_idx()
k = cute.assume(k, divby=_NB)
warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
lane = cute.arch.lane_idx()
m = n - k
gP = cute.domain_offset((k, k), mH[b, None, None])
gT = mT[b, k // _NB, None, None]
smem = cutlass.utils.SmemAllocator()
sP = smem.allocate_tensor(cutlass.Float32, cute.make_layout((n, _NB), stride=(_NB + 1, 1)), byte_alignment=16)
sT = smem.allocate_tensor(cutlass.Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
sS = smem.allocate_tensor(cutlass.Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
red = smem.allocate_tensor(cutlass.Float32, cute.make_layout(NW), byte_alignment=16)
s_tau = smem.allocate_tensor(cutlass.Float32, cute.make_layout(_NB), byte_alignment=16)
s_sc = smem.allocate_tensor(cutlass.Float32, cute.make_layout(2), byte_alignment=16)
for idx in cutlass.range(tidx, m * _NB, TPB):
sP[idx // _NB, idx % _NB] = gP[idx // _NB, idx % _NB]
for idx in cutlass.range(tidx, _NB * _NB, TPB):
sT[idx // _NB, idx % _NB] = 0.0
cute.arch.barrier()
for jj in cutlass.range(0, _NB, 1, unroll=1):
local = cutlass.Float32(0.0)
for i in cutlass.range(jj + 1 + tidx, m, TPB):
x = sP[i, jj]
local = local + x * x
partial = cute.arch.warp_reduction(local, operator.add)
if lane == 0:
red[warp] = partial
cute.arch.barrier()
if tidx == 0:
xnorm2 = cutlass.Float32(0.0)
for w in cutlass.range_constexpr(NW):
xnorm2 = xnorm2 + red[w]
alpha = sP[jj, jj]
nrm = cute.math.sqrt(alpha * alpha + xnorm2)
beta = nrm
if alpha >= 0.0:
beta = -nrm
tau_j = cutlass.Float32(0.0)
scale = cutlass.Float32(0.0)
beta_diag = alpha
if xnorm2 > 0.0:
tau_j = (beta - alpha) / beta
scale = 1.0 / (alpha - beta)
beta_diag = beta
sP[jj, jj] = beta_diag
mTau[b, k + jj] = tau_j
s_tau[jj] = tau_j
s_sc[0] = tau_j
s_sc[1] = scale
cute.arch.barrier()
tau_j = s_sc[0]
scale = s_sc[1]
for i in cutlass.range(jj + 1 + tidx, m, TPB):
sP[i, jj] = sP[i, jj] * scale
cute.arch.barrier()
for c in cutlass.range(jj + 1 + warp, _NB, NW):
dot = cutlass.Float32(0.0)
for i in cutlass.range(jj + 1 + lane, m, 32):
dot = dot + sP[i, jj] * sP[i, c]
dot = cute.arch.warp_reduction(dot, operator.add)
tw = tau_j * (dot + sP[jj, c])
if lane == 0:
sP[jj, c] = sP[jj, c] - tw
for i in cutlass.range(jj + 1 + lane, m, 32):
sP[i, c] = sP[i, c] - sP[i, jj] * tw
cute.arch.barrier()
if cutlass.const_expr(form_t):
for idx in cutlass.range(tidx, _NB * _NB, TPB):
l = idx // _NB
jc = idx % _NB
if l < jc:
s = sP[jc, l]
for r in cutlass.range(jc + 1, m, 1):
s = s + sP[r, l] * sP[r, jc]
sS[l, jc] = s
cute.arch.barrier()
for i in cutlass.range_constexpr(_NB):
if tidx < _NB:
l = tidx
if l == i:
sT[i, i] = s_tau[i]
elif l < i:
acc = cutlass.Float32(0.0)
for p in cutlass.range(l, i, 1):
acc = acc + sT[l, p] * sS[p, i]
sT[l, i] = -s_tau[i] * acc
cute.arch.barrier()
for idx in cutlass.range(tidx, m * _NB, TPB):
r = idx // _NB
c = idx % _NB
value = sP[r, c]
logical = value
if r < _NB:
if r < c:
logical = cutlass.Float32(0.0)
elif r == c:
logical = cutlass.Float32(1.0)
gP[r, c] = value
mV[b, r, c] = logical
if cutlass.const_expr(form_t):
for idx in cutlass.range(tidx, _NB * _NB, TPB):
gT[idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]
@cute.jit
def _panel_resident_launch(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mV: cute.Tensor,
k: cutlass.Int32,
):
tpb = 1024 if mH.shape[1] == 352 else 512
_panel_resident_kernel(mH, mTau, mT, mV, k, mH.shape[1], True, tpb, tpb // 32).launch(
grid=[mH.shape[0], 1, 1], block=[tpb, 1, 1]
)
@cute.jit
def _panel_resident_no_t_launch(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mV: cute.Tensor,
k: cutlass.Int32,
):
tpb = 1024 if mH.shape[1] == 352 else 512
_panel_resident_kernel(mH, mTau, mT, mV, k, mH.shape[1], False, tpb, tpb // 32).launch(
grid=[mH.shape[0], 1, 1], block=[tpb, 1, 1]
)
@cute.kernel
def _larft_from_gram_kernel(
mGram: cute.Tensor,
mTau: cute.Tensor,
mT: cute.Tensor,
k: cutlass.Int32,
):
tidx, _, _ = cute.arch.thread_idx()
b, _, _ = cute.arch.block_idx()
smem = cutlass.utils.SmemAllocator()
sT = smem.allocate_tensor(
cutlass.Float32,
cute.make_layout((_NB, _NB), stride=(_NB, 1)),
byte_alignment=128,
)
for idx in cutlass.range(tidx, _NB * _NB, 32):
sT[idx // _NB, idx % _NB] = 0.0
cute.arch.barrier()
for col in cutlass.range_constexpr(_NB):
row = tidx
if row == col:
sT[row, col] = mTau[b, k + col]
elif row < col:
value = cutlass.Float32(0.0)
for p in cutlass.range(row, col, 1):
value = value + sT[row, p] * mGram[b, p, col]
sT[row, col] = -mTau[b, k + col] * value
cute.arch.sync_warp()
for idx in cutlass.range(tidx, _NB * _NB, 32):
mT[b, k // _NB, idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]
@cute.jit
def _larft_from_gram_launch(
mGram: cute.Tensor,
mTau: cute.Tensor,
mT: cute.Tensor,
k: cutlass.Int32,
):
_larft_from_gram_kernel(mGram, mTau, mT, k).launch(
grid=[mGram.shape[0], 1, 1], block=[32, 1, 1]
)
@cute.kernel
def _diag_prepare_kernel(
mH: cute.Tensor,
mR: cute.Tensor,
k: cutlass.Int32,
):
tidx, _, _ = cute.arch.thread_idx()
b, _, _ = cute.arch.block_idx()
p = k // _NB
for idx in cutlass.range(tidx, _NB * _NB, _TPB):
i = idx // _NB
j = idx % _NB
value = mH[b, k + i, k + j]
mR[b, p, i, j] = value
if i < j:
value = cutlass.Float32(0.0)
elif i == j:
value = cutlass.Float32(1.0)
mH[b, k + i, k + j] = value
@cute.jit
def _diag_prepare_launch(
mH: cute.Tensor,
mR: cute.Tensor,
k: cutlass.Int32,
):
_diag_prepare_kernel(mH, mR, k).launch(
grid=[mH.shape[0], 1, 1], block=[_TPB, 1, 1]
)
@cute.kernel
def _diag_restore_kernel(
mH: cute.Tensor,
mR: cute.Tensor,
):
tidx, _, _ = cute.arch.thread_idx()
p, b, _ = cute.arch.block_idx()
k = p * _NB
for idx in cutlass.range(tidx, _NB * _NB, _TPB):
i = idx // _NB
j = idx % _NB
mH[b, k + i, k + j] = mR[b, p, i, j]
@cute.jit
def _diag_restore_launch(
mH: cute.Tensor,
mR: cute.Tensor,
):
_diag_restore_kernel(mH, mR).launch(
grid=[mR.shape[1], mH.shape[0], 1], block=[_TPB, 1, 1]
)
@cute.kernel
def _panel_cluster_kernel(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mR: cute.Tensor,
k: cutlass.Int32, n: cutlass.Constexpr, C: cutlass.Constexpr,
rows_cap: cutlass.Constexpr,
form_t: cutlass.Constexpr, TPB: cutlass.Constexpr,
NW: cutlass.Constexpr,
):
tidx, _, _ = cute.arch.thread_idx()
rank = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster())
_, b, _ = cute.arch.block_idx()
k = cute.assume(k, divby=_NB)
warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
lane = cute.arch.lane_idx()
m = n - k
S = cute.ceil_div(m, C)
row0 = rank * S
avail = m - row0
rows = avail if avail < S else S
rows = rows if rows > 0 else 0
gP = cute.domain_offset((k, k), mH[b, None, None])
gT = mT[b, k // _NB, None, None]
smem = cutlass.utils.SmemAllocator()
sP = smem.allocate_tensor(Float32, cute.make_layout((rows_cap, _NB), stride=(_NB + 1, 1)), byte_alignment=16)
sT = smem.allocate_tensor(Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
redbuf = smem.allocate_tensor(Float32, cute.make_layout(2 * _NB * _NB), byte_alignment=16)
total = smem.allocate_tensor(Float32, cute.make_layout(2 * _NB * _NB), byte_alignment=16)
recv = smem.allocate_tensor(
Float32,
cute.make_layout(
(2, C, 2 * _NB),
stride=(C * 2 * _NB, 2 * _NB, 1),
),
byte_alignment=16,
)
mbars = smem.allocate_array(cutlass.Int64, num_elems=2)
red = smem.allocate_tensor(Float32, cute.make_layout(NW), byte_alignment=16)
s_tau = smem.allocate_tensor(Float32, cute.make_layout(_NB), byte_alignment=16)
s_sc = smem.allocate_tensor(Float32, cute.make_layout(2), byte_alignment=16)
for idx in cutlass.range(tidx, rows * _NB, TPB):
sP[idx // _NB, idx % _NB] = gP[row0 + idx // _NB, idx % _NB]
for idx in cutlass.range(tidx, _NB * _NB, TPB):
sT[idx // _NB, idx % _NB] = 0.0
if tidx < 2:
cute.arch.mbarrier_init(mbars + tidx, 1)
cute.arch.mbarrier_arrive_and_expect_tx(
mbars + tidx, C * (2 * _NB) * 4
)
cute.arch.mbarrier_init_fence()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
cute.arch.barrier()
for jj in cutlass.range(0, _NB, 1, unroll=1):
rbuf = (jj & 1) * _NB * _NB
jloc = jj - row0
owns = (jloc >= 0) and (jloc < rows)
lo = jj + 1 - row0
lo = lo if lo > 0 else 0
for c in cutlass.range(jj + warp, _NB, NW):
dot = cutlass.Float32(0.0)
for i in cutlass.range(lo + lane, rows, 32):
dot = dot + sP[i, jj] * sP[i, c]
dot = cute.arch.warp_reduction(dot, operator.add)
if lane == 0:
redbuf[rbuf + c] = dot
redbuf[rbuf + _NB + c] = sP[jloc, c] if owns else cutlass.Float32(0.0)
_cluster_vsum_async(redbuf, total, recv, mbars, rbuf, jj, rank, tidx, C, TPB)
xnorm2 = total[rbuf + jj]
alpha = total[rbuf + _NB + jj]
nrm = cute.math.sqrt(alpha * alpha + xnorm2)
beta = nrm
if alpha >= 0.0:
beta = -nrm
tau_j = cutlass.Float32(0.0)
scale = cutlass.Float32(0.0)
if xnorm2 > 0.0:
tau_j = (beta - alpha) / beta
scale = 1.0 / (alpha - beta)
if tidx == 0:
s_tau[jj] = tau_j
s_sc[0] = tau_j
s_sc[1] = scale
if owns:
sP[jloc, jj] = beta if xnorm2 > 0.0 else alpha
mTau[b, k + jj] = tau_j
cute.arch.barrier()
tau_j = s_sc[0]
scale = s_sc[1]
for i in cutlass.range(lo + tidx, rows, TPB):
sP[i, jj] = sP[i, jj] * scale
cute.arch.barrier()
for c in cutlass.range(jj + 1 + warp, _NB, NW):
tw = tau_j * (total[rbuf + _NB + c] + scale * total[rbuf + c])
if owns and lane == 0:
sP[jloc, c] = sP[jloc, c] - tw
for i in cutlass.range(lo + lane, rows, 32):
sP[i, c] = sP[i, c] - sP[i, jj] * tw
cute.arch.barrier()
if cutlass.const_expr(form_t):
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
for idx in cutlass.range(tidx, _NB * _NB, TPB):
redbuf[idx] = 0.0
cute.arch.barrier()
for idx in cutlass.range(tidx, _NB * _NB, TPB):
l = idx // _NB
jc = idx % _NB
if l < jc:
s = cutlass.Float32(0.0)
jc_loc = jc - row0
if (jc_loc >= 0) and (jc_loc < rows):
s = s + sP[jc_loc, l]
lo2 = jc + 1 - row0
lo2 = lo2 if lo2 > 0 else 0
for r in cutlass.range(lo2, rows, 1):
s = s + sP[r, l] * sP[r, jc]
redbuf[idx] = s
cute.arch.barrier()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
if rank == 0:
for v in cutlass.range(tidx, _NB * _NB, TPB):
s = redbuf[v]
for p in cutlass.range_constexpr(C):
if p != 0:
s = s + _load_remote(redbuf.iterator + v, Int32(p))
total[v] = s
cute.arch.barrier()
for i in cutlass.range_constexpr(_NB):
if tidx < _NB:
l = tidx
if l == i:
sT[i, i] = s_tau[i]
elif l < i:
acc = cutlass.Float32(0.0)
for p in cutlass.range(l, i, 1):
acc = acc + sT[l, p] * total[p * _NB + i]
sT[l, i] = -s_tau[i] * acc
cute.arch.sync_warp()
cute.arch.cluster_arrive_relaxed()
cute.arch.cluster_wait()
for idx in cutlass.range(tidx, rows * _NB, TPB):
r = idx // _NB
c = idx % _NB
gr = row0 + r
value = sP[r, c]
if gr < _NB:
logical = value
if gr < c:
logical = cutlass.Float32(0.0)
elif gr == c:
logical = cutlass.Float32(1.0)
mR[b, gr, c] = logical
else:
mR[b, gr, c] = value
gP[gr, c] = value
if rank == 0 and cutlass.const_expr(form_t):
for idx in cutlass.range(tidx, _NB * _NB, TPB):
gT[idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]
@cute.jit
def _panel_cluster_launch(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mR: cute.Tensor,
k: cutlass.Int32, C: cutlass.Constexpr,
rows_cap: cutlass.Constexpr,
):
n = mH.shape[1]
batch = mH.shape[0]
tpb = 1024 if n >= 2048 else 512
_panel_cluster_kernel(mH, mTau, mT, mR, k, n, C, rows_cap, True, tpb, tpb // 32).launch(
grid=[C, batch, 1], block=[tpb, 1, 1], cluster=(C, 1, 1)
)
@cute.jit
def _panel_cluster_no_t_launch(
mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
mR: cute.Tensor,
k: cutlass.Int32, C: cutlass.Constexpr,
rows_cap: cutlass.Constexpr,
):
n = mH.shape[1]
batch = mH.shape[0]
tpb = 1024 if n >= 2048 else 512
_panel_cluster_kernel(mH, mTau, mT, mR, k, n, C, rows_cap, False, tpb, tpb // 32).launch(
grid=[C, batch, 1], block=[tpb, 1, 1], cluster=(C, 1, 1)
)
def _alloc_vg(batch, n):
return torch.empty((batch, n, _NB), device="cuda", dtype=torch.float32)
def _materialize_vg(H, Vg, k, strict_lower=None, eye=None):
n = H.shape[1]
m = n - k
if strict_lower is None:
ii = torch.arange(_NB, device=H.device)
strict_lower = (ii[:, None] > ii[None, :]).to(torch.float32)
if eye is None:
eye = torch.eye(_NB, device=H.device, dtype=torch.float32)
Vg[:, :_NB, :] = H[:, k : k + _NB, k : k + _NB] * strict_lower + eye
if m > _NB:
Vg[:, _NB:m, :] = H[:, k + _NB : n, k : k + _NB]
_panel_cache: dict = {}
_larft_cache: dict = {}
_diag_cache: dict = {}
_buf_cache: dict = {}
def _effective_factor_size(data: torch.Tensor):
batch, n, _ = data.shape
if n == 512 and batch >= 16:
edge_by_matrix = torch.maximum(
data[:, :8, -1].abs().amax(dim=1),
data[:, -8:, -1].abs().amax(dim=1),
)
edge_min_t, edge_max_t = torch.aminmax(edge_by_matrix)
edge_min = edge_min_t.item()
edge_max = edge_max_t.item()
if edge_max == 0.0:
return 384, "rankdef"
if edge_max < 1.0e-4:
first_col_offdiag = data[:, 1:8, 0].abs().amax().item()
if first_col_offdiag > 1.0e-4:
return 256, "clustered"
if batch >= 640:
if edge_min > 1.0e-4:
return n, "dense512"
if edge_min == 0.0:
return n, "mixed512"
if n == 1024 and batch >= 4:
tail_error = (data[:, 255, -1] - data[:, 255, 255]).abs().amax().item()
if tail_error < 2.0e-4:
return 768, "nearrank"
return n, None
def _nbo(n):
if n <= 352:
return _NB
if n == 512:
return _NB
return _NB
def _cluster_schedule(n):
if n == 1024:
return ((0, 4, 256),)
if n == 2048:
return ((0, 8, 256),)
if n == 4096:
return ((0, 16, 256),)
return ((0, 8, (n + 7) // 8),)
def _get_buf(batch, n, NBO, device):
key = (batch, n, NBO)
if key not in _buf_cache:
slL = torch.tril(torch.ones(NBO, NBO, device=device), -1)
eyeL = torch.eye(NBO, device=device, dtype=torch.float32)
_buf_cache[key] = (
torch.empty(batch, n, NBO, device=device, dtype=torch.float32),
torch.empty(batch, NBO, NBO, device=device, dtype=torch.float32),
slL, eyeL,
)
return _buf_cache[key]
def _blocked_qr(data: torch.Tensor):
batch, n, _ = data.shape
factor_size, structure = _effective_factor_size(data)
update_size = factor_size if structure is not None else n
tail_start = (
64 if n == 192 and factor_size == n
# The batched 64x64 tail kernel is nondeterministic at the n=512,
# batch=640 scored shapes. Finish those with the resident panel path.
else n - 64 if n >= 352 and n != 512 and factor_size == n
else n
)
if n >= 1024:
use_tf32 = True
elif n == 512:
use_tf32 = False
elif n >= 512:
use_tf32 = True
else:
use_tf32 = False
torch.backends.cuda.matmul.allow_tf32 = use_tf32
tf32_state = use_tf32
H = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
T = torch.empty((batch, (n + _NB - 1) // _NB, _NB, _NB), device=data.device, dtype=torch.float32)
Vpanel = _alloc_vg(batch, n)
mH, mTau, mT = _t2c(H), _t2c(tau), _t2c(T)
mR = _t2c(Vpanel)
panel_schedule = None
if n >= 1024:
panel_schedule = []
for start, C, rows_cap in _cluster_schedule(n):
full_key = (batch, n, "cluster", C, rows_cap, True)
no_t_key = (batch, n, "cluster", C, rows_cap, False)
if full_key not in _panel_cache:
_panel_cache[full_key] = cute.compile(
_panel_cluster_launch, mH, mTau, mT, mR,
cutlass.Int32(0), C, rows_cap,
)
if no_t_key not in _panel_cache:
_panel_cache[no_t_key] = cute.compile(
_panel_cluster_no_t_launch, mH, mTau, mT, mR,
cutlass.Int32(0), C, rows_cap,
)
panel_schedule.append(
(start, _panel_cache[full_key], _panel_cache[no_t_key])
)
panel = panel_schedule[0][1]
last_panel = panel_schedule[0][2]
else:
key = (batch, n, "no_t" if n == 512 else "full")
if key not in _panel_cache:
launch = _panel_resident_no_t_launch if n == 512 else _panel_resident_launch
_panel_cache[key] = cute.compile(
launch, mH, mTau, mT, mR, cutlass.Int32(0)
)
panel = _panel_cache[key]
if n == 512:
last_panel = panel
else:
last_key = (batch, n, "last_no_t")
if last_key not in _panel_cache:
_panel_cache[last_key] = cute.compile(
_panel_resident_no_t_launch,
mH, mTau, mT, mR, cutlass.Int32(0),
)
last_panel = _panel_cache[last_key]
NBO = _nbo(n)
if NBO == _NB:
Vg = Vpanel
Gram = (
torch.empty((batch, _NB, _NB), device=data.device, dtype=torch.float32)
if n == 512
else None
)
if n == 512:
mGram = _t2c(Gram)
lkey = (batch, n)
if lkey not in _larft_cache:
_larft_cache[lkey] = cute.compile(
_larft_from_gram_launch,
mGram,
mTau,
mT,
cutlass.Int32(0),
)
larft = _larft_cache[lkey]
Wbuf = torch.empty(
(batch, _NB, n), device=data.device, dtype=torch.float32
)
Auxbuf = torch.empty(
(batch, _NB, n) if n == 512 else (batch, n, _NB),
device=data.device,
dtype=torch.float32,
)
for k in range(0, n, _NB):
if k >= tail_start:
break
if k >= factor_size:
break
if panel_schedule is not None:
for start, full_panel, no_t_panel in panel_schedule:
if k >= start:
panel = full_panel
last_panel = no_t_panel
active_panel = last_panel if k + _NB >= factor_size else panel
active_panel(mH, mTau, mT, mR, cutlass.Int32(k))
if k + _NB >= factor_size:
break
mm = n - k
V = Vpanel[:, :mm, :]
A22 = H[:, k:n, k + _NB : update_size]
if n == 512:
full_tf32 = batch >= 640 and (
k >= 32
or structure in ("dense512", "rankdef", "clustered")
)
projection_only_tf32 = structure == "mixed512" and k == 0
if full_tf32 or projection_only_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
tf32_state = full_tf32
cols = update_size - k - _NB
W = Wbuf[:, :, :cols]
gram_tf32 = structure in (
"dense512", "rankdef", "clustered"
)
if gram_tf32 and not tf32_state:
torch.backends.cuda.matmul.allow_tf32 = True
tf32_state = True
elif not gram_tf32 and tf32_state:
torch.backends.cuda.matmul.allow_tf32 = False
tf32_state = False
torch.bmm(V.transpose(1, 2), V, out=Gram)
larft(mGram, mTau, mT, cutlass.Int32(k))
if full_tf32 or projection_only_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
tf32_state = full_tf32
torch.bmm(V.transpose(1, 2), A22, out=W)
transform_tf32 = full_tf32 and structure in (
"dense512", "rankdef", "clustered"
)
if (tf32_state or projection_only_tf32) and not transform_tf32:
torch.backends.cuda.matmul.allow_tf32 = False
W2 = Auxbuf[:, :, :cols]
torch.bmm(T[:, k // _NB].transpose(1, 2), W, out=W2)
if tf32_state and not transform_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
torch.baddbmm(A22, V, W2, beta=1.0, alpha=-1.0, out=A22)
elif n == 352:
W = torch.bmm(V.transpose(1, 2), A22)
W2 = torch.bmm(T[:, k // _NB].transpose(1, 2), W)
torch.baddbmm(A22, V, W2, beta=1.0, alpha=-1.0, out=A22)
else:
cols = update_size - k - _NB
W = Wbuf[:, :, :cols]
YT = Auxbuf[:, :mm, :]
torch.bmm(V, T[:, k // _NB].transpose(1, 2), out=YT)
torch.bmm(V.transpose(1, 2), A22, out=W)
torch.baddbmm(A22, YT, W, beta=1.0, alpha=-1.0, out=A22)
if tail_start < factor_size:
_geqr2_tail(H, tau, tail_start, factor_size)
if factor_size < n:
tau[:, factor_size:].zero_()
if structure == "clustered":
H[:, :, factor_size:].zero_()
elif structure == "nearrank":
H[:, :, factor_size:].zero_()
H[:, :256, factor_size:].copy_(torch.triu(H[:, :256, :256]))
return H, tau
ii = torch.arange(_NB, device=data.device)
slS = (ii[:, None] > ii[None, :]).to(torch.float32)
eyeS = torch.eye(_NB, device=data.device, dtype=torch.float32)
Vbuf, Tbuf, slL, eyeL = _get_buf(batch, n, NBO, data.device)
for kb in range(0, n, NBO):
if kb >= factor_size:
break
bw = min(NBO, n - kb)
nsub = (bw + _NB - 1) // _NB
for s in range(nsub):
ki = kb + s * _NB
panel(mH, mTau, mT, mR, cutlass.Int32(ki))
ie = kb + bw
if ki + _NB < ie:
if n == 512 and batch >= 640:
desired_tf32 = kb >= 2 * NBO
if desired_tf32 != tf32_state:
torch.backends.cuda.matmul.allow_tf32 = desired_tf32
tf32_state = desired_tf32
mm = n - ki
V = Vbuf[:, :mm, :_NB]
V[:, :_NB, :].copy_(H[:, ki:ki + _NB, ki:ki + _NB])
V[:, :_NB, :].mul_(slS).add_(eyeS)
if mm > _NB:
V[:, _NB:mm, :].copy_(H[:, ki + _NB:n, ki:ki + _NB])
Ain = H[:, ki:n, ki + _NB : ie]
YT = torch.bmm(V, T[:, ki // _NB].transpose(1, 2))
W = torch.bmm(V.transpose(1, 2), Ain)
torch.baddbmm(Ain, YT, W, beta=1.0, alpha=-1.0, out=Ain)
if kb + bw >= n:
break
if n == 512 and batch >= 640:
desired_tf32 = kb >= 2 * NBO
if desired_tf32 != tf32_state:
torch.backends.cuda.matmul.allow_tf32 = desired_tf32
tf32_state = desired_tf32
mm = n - kb
Vblk = Vbuf[:, :mm, :bw]
Vblk[:, :bw, :].copy_(H[:, kb:kb + bw, kb:kb + bw])
Vblk[:, :bw, :].mul_(slL[:bw, :bw]).add_(eyeL[:bw, :bw])
if mm > bw:
Vblk[:, bw:mm, :].copy_(H[:, kb + bw : n, kb:kb + bw])
Sb = torch.bmm(Vblk.transpose(1, 2), Vblk)
Tb = Tbuf[:, :bw, :bw]
W = min(_NB, bw)
Tb[:, :W, :W] = T[:, kb // _NB][:, :W, :W]
while W < bw:
w = min(_NB, bw - W)
o = W
Tb[:, o:o + w, :o].zero_()
Tb[:, o:o + w, o:o + w] = T[:, (kb + o) // _NB][:, :w, :w]
Cross = Sb[:, :W, o:o + w]
Tb[:, :W, o:o + w] = -torch.bmm(torch.bmm(Tb[:, :W, :W], Cross), Tb[:, o:o + w, o:o + w])
W += w
Af = H[:, kb:n, kb + bw : n]
YTb = torch.bmm(Vblk, Tb.transpose(1, 2))
W = torch.bmm(Vblk.transpose(1, 2), Af)
torch.baddbmm(Af, YTb, W, beta=1.0, alpha=-1.0, out=Af)
if factor_size < n:
tau[:, factor_size:].zero_()
return H, tau
def _geqr176_padded(data: torch.Tensor):
padded = torch.nn.functional.pad(data, (0, 16, 0, 16))
H, tau = _blocked_qr(padded)
return H[:, :176, :176], tau[:, :176]
def custom_kernel(data: input_t) -> output_t:
n = data.shape[-1]
if n == 176:
return _geqr176_padded(data)
if n <= 176 and n % _GEQR2_TDIM == 0:
return geqr2_small(data)
return _blocked_qr(data)
scrolls · 1084 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