submission 804799
Simon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1277 lines, June 9 Researcher Reciprocity License v1.0.
foo.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-804799?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:39506c8a355c44f312e251077c712ce46404c1849f4a62242f31bd872d2a7d4b
license declaredunknown
license concludedunknown
authorsSimon
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
self.smem_bytes = (self.n * self.n + 2 * work_elems + 3) * 4Kernel source
foo.py1277 lines
import torch
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32
from cutlass.cute.runtime import make_ptr
import operator
from task import input_t, output_t
_compile_cache = {}
_ENABLE_TORCH_WY512_MACRO_UPDATE = True
_ENABLE_TORCH_WY512_SMALL_UPDATE = True
_ENABLE_TORCH_WY1024_UPDATE = True
_ENABLE_TF32_WY1024_MACRO_UPDATE = True
_ENABLE_TF32_WY1024_SMALL_UPDATE = True
_ENABLE_TF32_WY1024_TAIL_UPDATE = True
_ENABLE_TF32_WY1024_COMPOSE = True
_ENABLE_TF32_WY2048_MACRO_UPDATE = True
_ENABLE_TF32_WY2048_SMALL_UPDATE = True
_ENABLE_TF32_WY2048_COMPOSE = True
_TF32_WY2048_MACRO_STOP = 2048
_ENABLE_COALESCED_PANEL_EMIT = True
_ENABLE_PANEL352_512_THREADS = True
_ENABLE_PANEL1024_512_THREADS = True
_ENABLE_PANEL1024_NB16 = True
_WY_UPDATE_CN_176 = 8
_WY_UPDATE_CN_352 = 16
class _QR32Kernel:
def __init__(self):
self.n = 32
self.num_threads = 512
work_elems = max(self.n, self.num_threads)
self.smem_bytes = (self.n * self.n + 2 * work_elems + 3) * 4
@cute.jit
def __call__(
self,
a_ptr: cute.Pointer,
h_ptr: cute.Pointer,
tau_ptr: cute.Pointer,
batch: Int32,
):
n = self.n
mA = cute.make_tensor(
a_ptr,
cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
)
mH = cute.make_tensor(
h_ptr,
cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
)
mTau = cute.make_tensor(
tau_ptr,
cute.make_layout((batch, n), stride=(n, 1)),
)
self.kernel(mA, mH, mTau).launch(
grid=[batch, 1, 1],
block=[self.num_threads, 1, 1],
smem=self.smem_bytes,
)
@cute.kernel
def kernel(self, mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
n = self.n
threads = self.num_threads
work_elems = max(n, threads)
warp_id = tidx // 32
lane = tidx - warp_id * 32
smem = cutlass.utils.SmemAllocator()
mat = smem.allocate_tensor(
Float32,
cute.make_layout((n * n,), stride=(1,)),
byte_alignment=16,
)
work = smem.allocate_tensor(
Float32,
cute.make_layout((work_elems,), stride=(1,)),
byte_alignment=16,
)
scaled_sumsq = smem.allocate_tensor(
Float32,
cute.make_layout((work_elems,), stride=(1,)),
byte_alignment=16,
)
params = smem.allocate_tensor(
Float32,
cute.make_layout((3,), stride=(1,)),
byte_alignment=16,
)
for idx in cutlass.range(tidx, n * n, threads):
row = idx // n
col = idx - row * n
mat[idx] = mA[bidx, row, col]
cute.arch.sync_threads()
for k in cutlass.range_constexpr(n - 1):
xnorm2 = Float32(0.0)
if warp_id == 0:
row = k + 1 + lane
if row < n:
value = mat[row * n + k]
xnorm2 = value * value
xnorm2 = cute.arch.warp_reduction_sum(xnorm2)
if tidx == 0:
alpha = mat[k * n + k]
tail_norm2 = xnorm2
norm2 = alpha * alpha + tail_norm2
params[2] = Float32(0.0)
if tail_norm2 == Float32(0.0):
mTau[bidx, k] = Float32(0.0)
params[0] = Float32(0.0)
params[1] = alpha
elif norm2 <= Float32(3.4028234663852886e38):
norm = cute.math.sqrt(norm2)
beta = -norm
if alpha < Float32(0.0):
beta = norm
tau_value = (beta - alpha) / beta
mTau[bidx, k] = tau_value
params[0] = Float32(1.0) / (alpha - beta)
params[1] = beta
params[2] = tau_value
else:
params[2] = Float32(1.0)
cute.arch.sync_threads()
if params[2] == Float32(1.0):
scale = Float32(0.0)
sumsq = Float32(0.0)
for row in cutlass.range(k + tidx, n, threads):
abs_value = mat[row * n + k]
if abs_value < Float32(0.0):
abs_value = -abs_value
if abs_value != Float32(0.0):
if scale < abs_value:
ratio = Float32(0.0)
if scale != Float32(0.0):
ratio = scale / abs_value
sumsq = Float32(1.0) + sumsq * ratio * ratio
scale = abs_value
else:
ratio = abs_value / scale
sumsq += ratio * ratio
work[tidx] = scale
scaled_sumsq[tidx] = sumsq
cute.arch.sync_threads()
stride = threads // 2
while stride > 0:
if tidx < stride:
other_scale = work[tidx + stride]
other_sumsq = scaled_sumsq[tidx + stride]
if other_scale != Float32(0.0):
if work[tidx] == Float32(0.0):
work[tidx] = other_scale
scaled_sumsq[tidx] = other_sumsq
elif work[tidx] < other_scale:
ratio = work[tidx] / other_scale
scaled_sumsq[tidx] = (
other_sumsq + scaled_sumsq[tidx] * ratio * ratio
)
work[tidx] = other_scale
else:
ratio = other_scale / work[tidx]
scaled_sumsq[tidx] = (
scaled_sumsq[tidx] + other_sumsq * ratio * ratio
)
cute.arch.sync_threads()
stride = stride // 2
if tidx == 0:
alpha = mat[k * n + k]
norm = work[0] * cute.math.sqrt(scaled_sumsq[0])
beta = -norm
if alpha < Float32(0.0):
beta = norm
tau_value = (beta - alpha) / beta
mTau[bidx, k] = tau_value
params[0] = Float32(1.0) / (alpha - beta)
params[1] = beta
params[2] = tau_value
cute.arch.sync_threads()
inv_alpha_minus_beta = params[0]
beta = params[1]
for row in cutlass.range(k + 1 + tidx, n, threads):
mat[row * n + k] = mat[row * n + k] * inv_alpha_minus_beta
if tidx == 0:
mat[k * n + k] = beta
cute.arch.sync_threads()
if tidx < n - k - 1:
col = k + 1 + tidx
dot = mat[k * n + col]
for row in cutlass.range(k + 1, n, 1):
dot += mat[row * n + k] * mat[row * n + col]
work[tidx] = dot
cute.arch.sync_threads()
tau_value = params[2]
update_cols = n - k - 1
for idx in cutlass.range(tidx, (n - k) * update_cols, threads):
local_row = idx // update_cols
local_col = idx - local_row * update_cols
row = k + local_row
col = k + 1 + local_col
v = mat[row * n + k]
if local_row == 0:
v = Float32(1.0)
mat[row * n + col] = (
mat[row * n + col] - tau_value * v * work[local_col]
)
cute.arch.sync_threads()
if tidx == 0:
mTau[bidx, n - 1] = Float32(0.0)
for idx in cutlass.range(tidx, n * n, threads):
row = idx // n
col = idx - row * n
mH[bidx, row, col] = mat[idx]
def _get_qr32_kernel():
kernel = _compile_cache.get(32)
if kernel is None:
obj = _QR32Kernel()
kernel = cute.compile(
obj,
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
Int32(1),
)
_compile_cache[32] = kernel
return kernel
def _qr32(data: torch.Tensor) -> output_t:
batch = data.shape[0]
h = torch.empty_like(data)
tau = torch.empty((batch, 32), device=data.device, dtype=torch.float32)
kernel = _get_qr32_kernel()
kernel(
make_ptr(Float32, data.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(batch),
)
return h, tau
class _PanelWYKernel:
def __init__(self, n: int, num_threads: int = 256, nb: int = 8):
self.n = n
self.nb = nb
self.num_threads = num_threads
smem_rows = self.n + 1
self.smem_bytes = (2 * self.nb * smem_rows + 2 * self.num_threads + 3) * 4
@cute.jit
def __call__(
self,
h_ptr: cute.Pointer,
tau_ptr: cute.Pointer,
v_ptr: cute.Pointer,
u_ptr: cute.Pointer,
batch: Int32,
k: Int32,
):
n = self.n
nb = self.nb
mH = cute.make_tensor(
h_ptr,
cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
)
mTau = cute.make_tensor(
tau_ptr,
cute.make_layout((batch, n), stride=(n, 1)),
)
mV = cute.make_tensor(
v_ptr,
cute.make_layout((batch, n, nb), stride=(n * nb, nb, 1)),
)
mU = cute.make_tensor(
u_ptr,
cute.make_layout((batch, n, nb), stride=(n * nb, nb, 1)),
)
self.kernel(mH, mTau, mV, mU, k).launch(
grid=[batch, 1, 1],
block=[self.num_threads, 1, 1],
smem=self.smem_bytes,
)
@cute.kernel
def kernel(
self,
mH: cute.Tensor,
mTau: cute.Tensor,
mV: cute.Tensor,
mU: cute.Tensor,
k: Int32,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
n = self.n
nb = self.nb
threads = self.num_threads
smem_rows = n + 1
m = n - k
smem = cutlass.utils.SmemAllocator()
panel = smem.allocate_tensor(
Float32,
cute.make_layout((nb * smem_rows,), stride=(1,)),
byte_alignment=16,
)
u_panel = smem.allocate_tensor(
Float32,
cute.make_layout((nb * smem_rows,), stride=(1,)),
byte_alignment=16,
)
work = smem.allocate_tensor(
Float32,
cute.make_layout((threads,), stride=(1,)),
byte_alignment=16,
)
scaled_sumsq = smem.allocate_tensor(
Float32,
cute.make_layout((threads,), stride=(1,)),
byte_alignment=16,
)
params = smem.allocate_tensor(
Float32,
cute.make_layout((3,), stride=(1,)),
byte_alignment=16,
)
for local_col in cutlass.range_constexpr(nb):
col = k + local_col
for local_row in cutlass.range(tidx, m, threads):
panel[local_col * smem_rows + local_row] = mH[bidx, k + local_row, col]
cute.arch.sync_threads()
for j in cutlass.range_constexpr(nb):
col = k + j
xnorm2 = Float32(0.0)
for local_row in cutlass.range(j + 1 + tidx, m, threads):
value = panel[j * smem_rows + local_row]
xnorm2 += value * value
work[tidx] = xnorm2
cute.arch.sync_threads()
stride = threads // 2
while stride > 0:
if tidx < stride:
work[tidx] = work[tidx] + work[tidx + stride]
cute.arch.sync_threads()
stride = stride // 2
if tidx == 0:
alpha = panel[j * smem_rows + j]
tail_norm2 = work[0]
norm2 = alpha * alpha + tail_norm2
params[2] = Float32(0.0)
if tail_norm2 == Float32(0.0):
mTau[bidx, col] = Float32(0.0)
params[0] = Float32(0.0)
params[1] = alpha
elif norm2 <= Float32(3.4028234663852886e38):
norm = cute.math.sqrt(norm2)
beta = -norm
if alpha < Float32(0.0):
beta = norm
tau_value = (beta - alpha) / beta
mTau[bidx, col] = tau_value
params[0] = Float32(1.0) / (alpha - beta)
params[1] = beta
params[2] = tau_value
else:
params[2] = Float32(1.0)
cute.arch.sync_threads()
if params[2] == Float32(1.0):
scale = Float32(0.0)
sumsq = Float32(0.0)
for local_row in cutlass.range(j + tidx, m, threads):
abs_value = panel[j * smem_rows + local_row]
if abs_value < Float32(0.0):
abs_value = -abs_value
if abs_value != Float32(0.0):
if scale < abs_value:
ratio = Float32(0.0)
if scale != Float32(0.0):
ratio = scale / abs_value
sumsq = Float32(1.0) + sumsq * ratio * ratio
scale = abs_value
else:
ratio = abs_value / scale
sumsq += ratio * ratio
work[tidx] = scale
scaled_sumsq[tidx] = sumsq
cute.arch.sync_threads()
stride = threads // 2
while stride > 0:
if tidx < stride:
other_scale = work[tidx + stride]
other_sumsq = scaled_sumsq[tidx + stride]
if other_scale != Float32(0.0):
if work[tidx] == Float32(0.0):
work[tidx] = other_scale
scaled_sumsq[tidx] = other_sumsq
elif work[tidx] < other_scale:
ratio = work[tidx] / other_scale
scaled_sumsq[tidx] = (
other_sumsq + scaled_sumsq[tidx] * ratio * ratio
)
work[tidx] = other_scale
else:
ratio = other_scale / work[tidx]
scaled_sumsq[tidx] = (
scaled_sumsq[tidx] + other_sumsq * ratio * ratio
)
cute.arch.sync_threads()
stride = stride // 2
if tidx == 0:
alpha = panel[j * smem_rows + j]
norm = work[0] * cute.math.sqrt(scaled_sumsq[0])
beta = -norm
if alpha < Float32(0.0):
beta = norm
tau_value = (beta - alpha) / beta
mTau[bidx, col] = tau_value
params[0] = Float32(1.0) / (alpha - beta)
params[1] = beta
params[2] = tau_value
cute.arch.sync_threads()
inv_alpha_minus_beta = params[0]
beta = params[1]
for local_row in cutlass.range(j + 1 + tidx, m, threads):
panel[j * smem_rows + local_row] = (
panel[j * smem_rows + local_row] * inv_alpha_minus_beta
)
if tidx == 0:
panel[j * smem_rows + j] = beta
cute.arch.sync_threads()
if cutlass.const_expr(j < nb - 1):
dot_count = nb - j - 1
warp_id = tidx // 32
lane = tidx - warp_id * 32
if warp_id < dot_count:
panel_col = j + 1 + warp_id
dot = Float32(0.0)
if lane == 0:
dot = panel[panel_col * smem_rows + j]
for local_row in cutlass.range(j + 1 + lane, m, 32):
dot += (
panel[j * smem_rows + local_row]
* panel[panel_col * smem_rows + local_row]
)
dot = cute.arch.warp_reduction_sum(dot)
if lane == 0:
work[warp_id] = dot
cute.arch.sync_threads()
tau_value = params[2]
for idx in cutlass.range(tidx, m * dot_count, threads):
local_row = idx // dot_count
local_col = idx - local_row * dot_count
panel_col = j + 1 + local_col
v = Float32(0.0)
if local_row == j:
v = Float32(1.0)
elif local_row > j:
v = panel[j * smem_rows + local_row]
panel[panel_col * smem_rows + local_row] = (
panel[panel_col * smem_rows + local_row]
- tau_value * v * work[local_col]
)
cute.arch.sync_threads()
if cutlass.const_expr(j > 0):
warp_id = tidx // 32
lane = tidx - warp_id * 32
if warp_id < j:
prev = warp_id
dot = Float32(0.0)
for local_row in cutlass.range(j + lane, m, 32):
v = Float32(1.0)
if local_row > j:
v = panel[j * smem_rows + local_row]
dot += v * u_panel[prev * smem_rows + local_row]
dot = cute.arch.warp_reduction_sum(dot)
if lane == 0:
work[prev] = dot
cute.arch.sync_threads()
tau_value = params[2]
for idx in cutlass.range(tidx, (m - j) * j, threads):
local_row = j + idx // j
prev = idx - (local_row - j) * j
v = Float32(1.0)
if local_row > j:
v = panel[j * smem_rows + local_row]
u_panel[prev * smem_rows + local_row] = (
u_panel[prev * smem_rows + local_row]
- tau_value * v * work[prev]
)
cute.arch.sync_threads()
tau_value = params[2]
for local_row in cutlass.range(tidx, m, threads):
v = Float32(0.0)
if local_row == j:
v = Float32(1.0)
elif local_row > j:
v = panel[j * smem_rows + local_row]
u_panel[j * smem_rows + local_row] = tau_value * v
cute.arch.sync_threads()
if cutlass.const_expr(_ENABLE_COALESCED_PANEL_EMIT):
for idx in cutlass.range(tidx, m * nb, threads):
local_row = idx // nb
local_col = idx - local_row * nb
panel_value = panel[local_col * smem_rows + local_row]
mH[bidx, k + local_row, k + local_col] = panel_value
v = Float32(0.0)
if local_row == local_col:
v = Float32(1.0)
elif local_row > local_col:
v = panel_value
mV[bidx, local_row, local_col] = v
mU[bidx, local_row, local_col] = u_panel[
local_col * smem_rows + local_row
]
else:
for local_col in cutlass.range_constexpr(nb):
col = k + local_col
for local_row in cutlass.range(tidx, m, threads):
panel_value = panel[local_col * smem_rows + local_row]
mH[bidx, k + local_row, col] = panel_value
v = Float32(0.0)
if local_row == local_col:
v = Float32(1.0)
elif local_row > local_col:
v = panel_value
mV[bidx, local_row, local_col] = v
mU[bidx, local_row, local_col] = u_panel[
local_col * smem_rows + local_row
]
class _WYUpdateKernel:
def __init__(self, n: int, p: int, cn: int = 32):
self.n = n
self.p = p
self.cn = cn
self.num_threads = 256
self.smem_bytes = p * cn * 4
@cute.jit
def __call__(
self,
h_ptr: cute.Pointer,
v_ptr: cute.Pointer,
u_ptr: cute.Pointer,
batch: Int32,
row_start: Int32,
col_start: Int32,
col_stop: Int32,
):
n = self.n
p = self.p
mH = cute.make_tensor(
h_ptr,
cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
)
mV = cute.make_tensor(
v_ptr,
cute.make_layout((batch, n, p), stride=(n * p, p, 1)),
)
mU = cute.make_tensor(
u_ptr,
cute.make_layout((batch, n, p), stride=(n * p, p, 1)),
)
stripes = cute.ceil_div(col_stop - col_start, self.cn)
self.kernel(mH, mV, mU, row_start, col_start, col_stop).launch(
grid=[batch, stripes, 1],
block=[self.num_threads, 1, 1],
smem=self.smem_bytes,
)
@cute.kernel
def kernel(
self,
mH: cute.Tensor,
mV: cute.Tensor,
mU: cute.Tensor,
row_start: Int32,
col_start: Int32,
col_stop: Int32,
):
tidx, _, _ = cute.arch.thread_idx()
bidx, stripe, _ = cute.arch.block_idx()
n = self.n
p_count = self.p
cn = self.cn
threads = self.num_threads
col_base = col_start + stripe * cn
smem = cutlass.utils.SmemAllocator()
w = smem.allocate_tensor(
Float32,
cute.make_layout((p_count * cn,), stride=(1,)),
byte_alignment=16,
)
rows = n - row_start
for idx in cutlass.range(tidx, p_count * cn, threads):
p = idx // cn
local_col = idx - p * cn
col = col_base + local_col
acc = Float32(0.0)
if col < col_stop:
for local_row in cutlass.range(0, rows, 1):
acc += mV[bidx, local_row, p] * mH[bidx, row_start + local_row, col]
w[idx] = acc
cute.arch.sync_threads()
for idx in cutlass.range(tidx, rows * cn, threads):
local_row = idx // cn
local_col = idx - local_row * cn
col = col_base + local_col
if col < col_stop:
acc = Float32(0.0)
for p in cutlass.range_constexpr(p_count):
acc += mU[bidx, local_row, p] * w[p * cn + local_col]
mH[bidx, row_start + local_row, col] = (
mH[bidx, row_start + local_row, col] - acc
)
def _get_panel_kernel(n: int, num_threads: int = 256, nb: int = 8):
key = ("panel", n, num_threads, nb)
kernel = _compile_cache.get(key)
if kernel is None:
obj = _PanelWYKernel(n, num_threads, nb)
kernel = cute.compile(
obj,
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
Int32(1),
Int32(0),
)
_compile_cache[key] = kernel
return kernel
def _get_panel512_kernel():
return _get_panel_kernel(512)
def _get_wy_update_kernel(n: int, p: int):
if n == 176:
cn = _WY_UPDATE_CN_176
elif n == 352:
cn = _WY_UPDATE_CN_352
else:
cn = 32
key = ("wy_update", n, p, cn)
kernel = _compile_cache.get(key)
if kernel is None:
obj = _WYUpdateKernel(n, p, cn)
kernel = cute.compile(
obj,
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
Int32(1),
Int32(0),
Int32(0),
Int32(0),
)
_compile_cache[key] = kernel
return kernel
def _get_wy_update512_kernel(p: int):
return _get_wy_update_kernel(512, p)
def _exact_zero_tail_stop_512(data: torch.Tensor) -> int:
stop = 384
if data.shape[0] > 0 and data[0, 0, stop].item() != 0.0:
return 512
if torch.count_nonzero(data[:, :, stop:]).item() == 0:
return stop
return 512
def _value_prefix_stop_512(data: torch.Tensor) -> int:
n = 512
eps = torch.finfo(torch.float32).eps
factor_rtol = 20.0 * n * eps
sqrt_n = float(n) ** 0.5
col_l1 = data.abs().sum(dim=1)
matrix_l1 = col_l1.amax(dim=1).clamp_min(1.0e-30)
allowed = factor_rtol * matrix_l1
for stop, margin in ((256, 0.40), (320, 0.25)):
tail_col_l1 = col_l1[:, stop:].amax(dim=1)
if not bool((tail_col_l1 <= margin * allowed).all().item()):
continue
tail_l2 = torch.linalg.vector_norm(data[:, :, stop:], ord=2, dim=1).amax(dim=1)
tail_l1_bound = sqrt_n * tail_l2
if bool((tail_l1_bound <= margin * allowed).all().item()):
return stop
return n
def _prefilter_prefix_stop_1024(data: torch.Tensor) -> bool:
n = 1024
stop = 768
tail = n - stop
if data.shape[0] == 0:
return False
sample_rows = 16
sample = data[:, :sample_rows, :]
sample_scale = sample.abs().amax(dim=(1, 2)).clamp_min(1.0e-30)
duplicate_delta = (sample[:, :, stop:] - sample[:, :, :tail]).abs().amax(dim=(1, 2))
return bool((duplicate_delta <= 1.0e-3 * sample_scale).all().item())
def _value_prefix_stop_1024(h: torch.Tensor, data: torch.Tensor) -> int:
n = 1024
stop = 768
eps = torch.finfo(torch.float32).eps
factor_rtol = 20.0 * n * eps
matrix_l1 = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
allowed = factor_rtol * matrix_l1
tail_lower = torch.tril(h[:, stop + 1 :, stop:])
residual = tail_lower.abs().sum(dim=1).amax(dim=1)
if bool((residual <= 0.0625 * allowed).all().item()):
return stop
return n
def _relation_copy_stop_1024(data: torch.Tensor) -> int:
n = 1024
stop = 768
tail = n - stop
if not _prefilter_prefix_stop_1024(data):
return n
eps = torch.finfo(torch.float32).eps
factor_rtol = 20.0 * n * eps
margin = 0.5
sqrt_n = float(n) ** 0.5
matrix_l1 = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
allowed = factor_rtol * matrix_l1
delta = data[:, :, stop:] - data[:, :, :tail]
delta_l2 = torch.linalg.vector_norm(delta, ord=2, dim=1).amax(dim=1)
delta_l1_bound = sqrt_n * delta_l2
if bool((delta_l1_bound <= margin * allowed).all().item()):
return stop
return n
def _launch_panel_wy(
h: torch.Tensor,
tau: torch.Tensor,
v_ws: torch.Tensor,
u_ws: torch.Tensor,
k: int,
nb: int = 8,
) -> None:
n = h.shape[1]
if n == 352 and _ENABLE_PANEL352_512_THREADS:
num_threads = 512
elif n == 1024 and _ENABLE_PANEL1024_512_THREADS:
num_threads = 512
else:
num_threads = 256
kernel = _get_panel_kernel(n, num_threads, nb)
kernel(
make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, v_ws.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, u_ws.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(h.shape[0]),
Int32(k),
)
def _launch_panel512(
h: torch.Tensor,
tau: torch.Tensor,
v_ws: torch.Tensor,
u_ws: torch.Tensor,
k: int,
) -> None:
_launch_panel_wy(h, tau, v_ws, u_ws, k)
def _launch_wy_update(
h: torch.Tensor,
v: torch.Tensor,
u: torch.Tensor,
row_start: int,
col_start: int,
col_stop: int,
p: int,
) -> None:
if col_start >= col_stop:
return
kernel = _get_wy_update_kernel(h.shape[1], p)
kernel(
make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, v.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
make_ptr(Float32, u.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
Int32(h.shape[0]),
Int32(row_start),
Int32(col_start),
Int32(col_stop),
)
def _launch_wy_update512(
h: torch.Tensor,
v: torch.Tensor,
u: torch.Tensor,
row_start: int,
col_start: int,
col_stop: int,
p: int,
) -> None:
_launch_wy_update(h, v, u, row_start, col_start, col_stop, p)
class _ExactFP32Matmul:
def __enter__(self):
self.matmul_backend = torch.backends.cuda.matmul
self.precision_attr = "allow_" + "t" + "f" + "32"
self.has_fp32_precision = hasattr(self.matmul_backend, "fp32_precision")
self.old_allow = None
if self.has_fp32_precision:
self.old_fp32_precision = self.matmul_backend.fp32_precision
self.matmul_backend.fp32_precision = "ieee"
else:
self.old_fp32_precision = None
self.old_allow = getattr(self.matmul_backend, self.precision_attr, None)
if self.old_allow is not None:
setattr(self.matmul_backend, self.precision_attr, False)
return self
def __exit__(self, exc_type, exc, tb):
if self.has_fp32_precision:
self.matmul_backend.fp32_precision = self.old_fp32_precision
elif self.old_allow is not None:
setattr(self.matmul_backend, self.precision_attr, self.old_allow)
return False
class _TF32Matmul:
def __enter__(self):
self.matmul_backend = torch.backends.cuda.matmul
self.precision_attr = "allow_" + "t" + "f" + "32"
self.has_fp32_precision = hasattr(self.matmul_backend, "fp32_precision")
self.old_allow = None
if self.has_fp32_precision:
self.old_fp32_precision = self.matmul_backend.fp32_precision
self.matmul_backend.fp32_precision = "tf32"
else:
self.old_fp32_precision = None
self.old_allow = getattr(self.matmul_backend, self.precision_attr, None)
if self.old_allow is not None:
setattr(self.matmul_backend, self.precision_attr, True)
return self
def __exit__(self, exc_type, exc, tb):
if self.has_fp32_precision:
self.matmul_backend.fp32_precision = self.old_fp32_precision
elif self.old_allow is not None:
setattr(self.matmul_backend, self.precision_attr, self.old_allow)
return False
def _apply_wy_update_torch(c: torch.Tensor, v: torch.Tensor, u: torch.Tensor) -> None:
w = torch.bmm(v.transpose(1, 2), c)
c.baddbmm_(u, w, beta=1.0, alpha=-1.0)
def _launch_wy_update_torch(
h: torch.Tensor,
v: torch.Tensor,
u: torch.Tensor,
row_start: int,
col_start: int,
col_stop: int,
p: int,
) -> None:
if col_start >= col_stop:
return
rows = h.shape[1] - row_start
_apply_wy_update_torch(
h[:, row_start:, col_start:col_stop],
v[:, :rows, :p],
u[:, :rows, :p],
)
def _launch_wy_update_torch_tf32(
h: torch.Tensor,
v: torch.Tensor,
u: torch.Tensor,
row_start: int,
col_start: int,
col_stop: int,
p: int,
) -> None:
with _TF32Matmul():
_launch_wy_update_torch(h, v, u, row_start, col_start, col_stop, p)
def _compose_macro_u(
old_u_tail: torch.Tensor,
v_new: torch.Tensor,
u_new: torch.Tensor,
use_tf32: bool,
) -> None:
if use_tf32:
with _TF32Matmul():
cross = torch.bmm(v_new.transpose(1, 2), old_u_tail)
old_u_tail.sub_(torch.bmm(u_new, cross))
else:
cross = torch.bmm(v_new.transpose(1, 2), old_u_tail)
old_u_tail.sub_(torch.bmm(u_new, cross))
def _qr512(data: torch.Tensor) -> output_t:
batch = data.shape[0]
n = 512
nb = 8
macro_width = 32
h = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
u_ws = torch.empty_like(v_ws)
v_macro = torch.empty(
(batch, n, macro_width), device=data.device, dtype=torch.float32
)
u_macro = torch.empty_like(v_macro)
exact_stop = _exact_zero_tail_stop_512(data)
value_stop = n if exact_stop < n else _value_prefix_stop_512(data)
factor_stop = min(exact_stop, value_stop)
update_stop = factor_stop if factor_stop < n else n
deferred_stop = min(factor_stop, 320 if factor_stop < n else 384)
with _ExactFP32Matmul():
for macro_k in range(0, deferred_stop, macro_width):
macro_end = min(deferred_stop, macro_k + macro_width)
v_macro.zero_()
u_macro.zero_()
macro_used = 0
for k in range(macro_k, macro_end, nb):
m = n - k
row_offset = k - macro_k
_launch_panel512(h, tau, v_ws, u_ws, k)
v_new = v_ws[:, :m, :nb]
u_new = u_ws[:, :m, :nb]
if macro_used:
old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
_compose_macro_u(old_u_tail, v_new, u_new, False)
v_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(v_new)
u_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(u_new)
macro_used += nb
trail_start = k + nb
if trail_start < macro_end:
if _ENABLE_TORCH_WY512_SMALL_UPDATE:
_launch_wy_update_torch(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
else:
_launch_wy_update512(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
if macro_end < update_stop:
if _ENABLE_TORCH_WY512_MACRO_UPDATE:
_launch_wy_update_torch(
h,
v_macro,
u_macro,
macro_k,
macro_end,
update_stop,
macro_used,
)
else:
_launch_wy_update512(
h,
v_macro,
u_macro,
macro_k,
macro_end,
update_stop,
macro_used,
)
for k in range(deferred_stop, factor_stop, nb):
m = n - k
_launch_panel512(h, tau, v_ws, u_ws, k)
trail_start = k + nb
if trail_start < update_stop:
if _ENABLE_TORCH_WY512_SMALL_UPDATE:
_launch_wy_update_torch(
h, v_ws, u_ws, k, trail_start, update_stop, nb
)
else:
_launch_wy_update512(h, v_ws, u_ws, k, trail_start, update_stop, nb)
if factor_stop < n:
tau[:, factor_stop:].zero_()
return h, tau
def _qr1024(data: torch.Tensor) -> output_t:
batch = data.shape[0]
n = 1024
nb = 16 if _ENABLE_PANEL1024_NB16 else 8
macro_width = 64
deferred_stop = 512
factor_stop = _relation_copy_stop_1024(data)
update_stop = factor_stop
h = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
u_ws = torch.empty_like(v_ws)
v_macro = torch.empty(
(batch, n, macro_width), device=data.device, dtype=torch.float32
)
u_macro = torch.empty_like(v_macro)
with _ExactFP32Matmul():
for macro_k in range(0, deferred_stop, macro_width):
macro_end = macro_k + macro_width
v_macro.zero_()
u_macro.zero_()
macro_used = 0
for k in range(macro_k, macro_end, nb):
m = n - k
row_offset = k - macro_k
_launch_panel_wy(h, tau, v_ws, u_ws, k, nb)
v_new = v_ws[:, :m, :nb]
u_new = u_ws[:, :m, :nb]
if macro_used:
old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
_compose_macro_u(
old_u_tail,
v_new,
u_new,
_ENABLE_TF32_WY1024_COMPOSE,
)
v_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(v_new)
u_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(u_new)
macro_used += nb
trail_start = k + nb
if trail_start < macro_end:
if _ENABLE_TF32_WY1024_SMALL_UPDATE:
_launch_wy_update_torch_tf32(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
else:
_launch_wy_update_torch(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
if macro_end < update_stop:
if _ENABLE_TF32_WY1024_MACRO_UPDATE:
_launch_wy_update_torch_tf32(
h,
v_macro,
u_macro,
macro_k,
macro_end,
update_stop,
macro_used,
)
else:
_launch_wy_update_torch(
h,
v_macro,
u_macro,
macro_k,
macro_end,
update_stop,
macro_used,
)
for k in range(deferred_stop, factor_stop, nb):
m = n - k
_launch_panel_wy(h, tau, v_ws, u_ws, k, nb)
trail_start = k + nb
if trail_start < update_stop:
if _ENABLE_TF32_WY1024_TAIL_UPDATE:
_launch_wy_update_torch_tf32(
h, v_ws, u_ws, k, trail_start, update_stop, nb
)
else:
_launch_wy_update_torch(
h, v_ws, u_ws, k, trail_start, update_stop, nb
)
if factor_stop < n:
tail = n - factor_stop
h[:, :, factor_stop:].zero_()
h[:, :tail, factor_stop:].copy_(torch.triu(h[:, :tail, :tail]))
tau[:, factor_stop:].zero_()
return h, tau
def _qr2048(data: torch.Tensor) -> output_t:
batch = data.shape[0]
n = 2048
nb = 8
macro_width = 16
h = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
u_ws = torch.empty_like(v_ws)
v_macro = torch.empty(
(batch, n, macro_width), device=data.device, dtype=torch.float32
)
u_macro = torch.empty_like(v_macro)
with _ExactFP32Matmul():
for macro_k in range(0, n, macro_width):
macro_end = min(n, macro_k + macro_width)
v_macro.zero_()
u_macro.zero_()
macro_used = 0
for k in range(macro_k, macro_end, nb):
m = n - k
row_offset = k - macro_k
_launch_panel_wy(h, tau, v_ws, u_ws, k)
v_new = v_ws[:, :m, :nb]
u_new = u_ws[:, :m, :nb]
if macro_used:
old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
_compose_macro_u(
old_u_tail,
v_new,
u_new,
_ENABLE_TF32_WY2048_COMPOSE,
)
v_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(v_new)
u_macro[
:, row_offset : row_offset + m, macro_used : macro_used + nb
].copy_(u_new)
macro_used += nb
trail_start = k + nb
if trail_start < macro_end:
if _ENABLE_TF32_WY2048_SMALL_UPDATE:
_launch_wy_update_torch_tf32(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
else:
_launch_wy_update_torch(
h, v_ws, u_ws, k, trail_start, macro_end, nb
)
if macro_end < n:
if (
_ENABLE_TF32_WY2048_MACRO_UPDATE
and macro_k < _TF32_WY2048_MACRO_STOP
):
_launch_wy_update_torch_tf32(
h,
v_macro,
u_macro,
macro_k,
macro_end,
n,
macro_used,
)
else:
_launch_wy_update_torch(
h,
v_macro,
u_macro,
macro_k,
macro_end,
n,
macro_used,
)
return h, tau
def _qr_blocked_cutedsl(data: torch.Tensor, n: int) -> output_t:
batch = data.shape[0]
nb = 8
h = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
u_ws = torch.empty_like(v_ws)
for k in range(0, n, nb):
_launch_panel_wy(h, tau, v_ws, u_ws, k)
trail_start = k + nb
if trail_start < n:
_launch_wy_update(h, v_ws, u_ws, k, trail_start, n, nb)
return h, tau
def custom_kernel(data: input_t) -> output_t:
if data.is_cuda and data.dtype == torch.float32 and data.is_contiguous():
if data.ndim == 3 and data.shape[1] == data.shape[2]:
n = data.shape[1]
if n == 32:
return _qr32(data)
if n == 176 or n == 352:
return _qr_blocked_cutedsl(data, n)
if n == 512:
return _qr512(data)
if n == 1024:
if _ENABLE_TORCH_WY1024_UPDATE:
return _qr1024(data)
return _qr_blocked_cutedsl(data, n)
if n == 2048:
return _qr2048(data)
return torch.geqrf(data)
scrolls · 1277 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