submission 839047
benfattori · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1317 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-839047?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:289c318462d714efcd7afc2139ff155eb57422913b30c48994ed4f33aaec8f24
license declaredunknown
license concludedunknown
authorsbenfattori
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
triton.Config(fused-epilogue
DO_EPILOGUE: tl.constexpr,mma
acc = tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)num-warps = 4
num_warps=4,split-k
def _prune_bmm_t_split_k_configs(configs, named_args, **_kwargs):stages = 3
num_stages=3,Kernel source
submission.py1317 lines
import torch
from task import input_t, output_t
import triton.language as tl
import triton
_BMM_ACC_AUTOTUNE_CONFIGS = [
triton.Config(
{"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 256},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 128},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 256},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 32},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 64},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 128},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256},
num_warps=8,
num_stages=3,
),
]
_BMM_T_AUTOTUNE_CONFIGS = [
triton.Config(
{"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 128},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 256},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 128},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 256},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 32},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 64},
num_warps=4,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 128},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 256},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 32},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 64},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 128},
num_warps=8,
num_stages=3,
),
triton.Config(
{"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 256},
num_warps=8,
num_stages=3,
),
]
def _prune_bmm_acc_configs(configs, named_args, **_kwargs):
dim_m = named_args["dim_m"]
dim_n = named_args["dim_n"]
pruned = []
for config in configs:
block_m = config.kwargs["BLOCK_SIZE_M"]
block_n = config.kwargs["BLOCK_SIZE_N"]
if dim_m <= 16 and block_m != 16:
continue
if dim_m > 16 and block_m < 32:
continue
if dim_n <= 16 and block_n > 16:
continue
if dim_n <= 32 and block_n > 32:
continue
if dim_n <= 64 and block_n > 64:
continue
pruned.append(config)
return pruned or configs[:1]
def _prune_bmm_t_configs(configs, named_args, **_kwargs):
dim_k = named_args["dim_k"]
dim_n = named_args["dim_n"]
pruned = []
for config in configs:
block_k = config.kwargs["BLOCK_SIZE_K"]
block_n = config.kwargs["BLOCK_SIZE_N"]
if dim_k <= 16 and block_k != 16:
continue
if dim_k <= 32 and block_k > 32:
continue
if dim_k <= 64 and block_k > 64:
continue
if dim_k >= 128 and block_k < 32:
continue
if dim_n <= 16 and block_n > 16:
continue
if dim_n <= 32 and block_n > 32:
continue
if dim_n <= 64 and block_n > 64:
continue
pruned.append(config)
return pruned or configs[:1]
def _prune_bmm_t_split_k_configs(configs, named_args, **_kwargs):
dim_k = named_args["tune_split_k"]
dim_n = named_args["dim_n"]
pruned = []
for config in configs:
block_k = config.kwargs["BLOCK_SIZE_K"]
block_n = config.kwargs["BLOCK_SIZE_N"]
if dim_k <= 16 and block_k != 16:
continue
if dim_k <= 32 and block_k > 32:
continue
if dim_k <= 64 and block_k > 64:
continue
if dim_k >= 128 and block_k < 32:
continue
if dim_n <= 16 and block_n > 16:
continue
if dim_n <= 32 and block_n > 32:
continue
if dim_n <= 64 and block_n > 64:
continue
pruned.append(config)
return pruned or configs[:1]
def _shape_bucket(value: int) -> int:
return triton.next_power_of_2(value)
@triton.autotune(
configs=_BMM_ACC_AUTOTUNE_CONFIGS,
key=["op_kind", "tune_m", "tune_k", "tune_n"],
prune_configs_by={"early_config_prune": _prune_bmm_acc_configs},
restore_value=["c_out_ptr"],
warmup=5,
rep=15,
)
@triton.jit
def _bmm_tf32x3_kernel(
a_ptr,
b_ptr,
c_in_ptr,
c_out_ptr,
bounds_ptr,
panel_end,
a_b_stride,
a_m_stride,
a_k_stride,
b_b_stride,
b_k_stride,
b_n_stride,
c_in_b_stride,
c_in_m_stride,
c_in_n_stride,
c_out_b_stride,
c_out_m_stride,
c_out_n_stride,
dim_m,
dim_k,
dim_n,
op_kind,
tune_m,
tune_k,
tune_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
specialize_constexpr: tl.constexpr,
):
pid = tl.program_id(0)
b_pid = tl.program_id(1)
active_n = dim_n
if specialize_constexpr:
bounds_ptr += b_pid
bound = tl.load(bounds_ptr)
active_n = bound - panel_end
if active_n <= 0:
return
num_pid_n = tl.cdiv(dim_n, BLOCK_SIZE_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
n_tile_start = pid_n * BLOCK_SIZE_N
if specialize_constexpr:
if n_tile_start >= active_n:
return
a_base = a_ptr + b_pid * a_b_stride
b_base = b_ptr + b_pid * b_b_stride
c_in_base = c_in_ptr + b_pid * c_in_b_stride
c_out_base = c_out_ptr + b_pid * c_out_b_stride
a_block_ptr = tl.make_block_ptr(
base=a_base,
shape=(dim_m, dim_k),
strides=(a_m_stride, a_k_stride),
offsets=(pid_m * BLOCK_SIZE_M, 0),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_K),
order=(1, 0),
)
b_block_ptr = tl.make_block_ptr(
base=b_base,
shape=(dim_k, dim_n),
strides=(b_k_stride, b_n_stride),
offsets=(0, n_tile_start),
block_shape=(BLOCK_SIZE_K, BLOCK_SIZE_N),
order=(1, 0),
)
a = tl.load(a_block_ptr, boundary_check=(0, 1))
b = tl.load(b_block_ptr, boundary_check=(0, 1))
acc = tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)
c_in_block_ptr = tl.make_block_ptr(
base=c_in_base,
shape=(dim_m, dim_n),
strides=(c_in_m_stride, c_in_n_stride),
offsets=(pid_m * BLOCK_SIZE_M, n_tile_start),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
order=(1, 0),
)
c = tl.load(c_in_block_ptr, boundary_check=(0, 1))
acc = c - acc
c_out_block_ptr = tl.make_block_ptr(
base=c_out_base,
shape=(dim_m, dim_n),
strides=(c_out_m_stride, c_out_n_stride),
offsets=(pid_m * BLOCK_SIZE_M, n_tile_start),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
order=(1, 0),
)
tl.store(c_out_block_ptr, acc, boundary_check=(0, 1))
@triton.autotune(
configs=_BMM_T_AUTOTUNE_CONFIGS,
key=["tune_m", "tune_k", "tune_n"],
prune_configs_by={"early_config_prune": _prune_bmm_t_configs},
warmup=5,
rep=15,
)
@triton.jit
def bmm_then_t_tf32x3_kernel(
a_ptr,
b_ptr,
t_ptr,
c_ptr,
bounds_ptr,
panel_end,
a_b_stride,
a_m_stride,
a_k_stride,
b_b_stride,
b_k_stride,
b_n_stride,
t_b_stride,
t_m_stride,
t_k_stride,
c_b_stride,
c_m_stride,
c_n_stride,
dim_m,
dim_k,
dim_n,
tune_m,
tune_k,
tune_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
specialize_constexpr: tl.constexpr,
):
pid = tl.program_id(0)
b_pid = tl.program_id(1)
active_n = dim_n
if specialize_constexpr:
bounds_ptr += b_pid
bound = tl.load(bounds_ptr)
active_n = bound - panel_end
if active_n <= 0:
return
# dim_m is the panel width. One program handles all panel rows for one
# output-column tile so the following T multiply has the full small block.
pid_m = 0 # noqa
pid_n = pid
n_tile_start = pid_n * BLOCK_SIZE_N
if specialize_constexpr:
if n_tile_start >= active_n:
return
a_base = a_ptr + b_pid * a_b_stride
b_base = b_ptr + b_pid * b_b_stride
t_base = t_ptr + b_pid * t_b_stride
c_base = c_ptr + b_pid * c_b_stride
a_block_ptr = tl.make_block_ptr(
base=a_base,
shape=(dim_m, dim_k),
strides=(a_m_stride, a_k_stride),
offsets=(0, 0),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_K),
order=(1, 0),
)
b_block_ptr = tl.make_block_ptr(
base=b_base,
shape=(dim_k, dim_n),
strides=(b_k_stride, b_n_stride),
offsets=(0, n_tile_start),
block_shape=(BLOCK_SIZE_K, BLOCK_SIZE_N),
order=(1, 0),
)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(0, tl.cdiv(dim_k, BLOCK_SIZE_K)):
a = tl.load(a_block_ptr, boundary_check=(0, 1))
b = tl.load(b_block_ptr, boundary_check=(0, 1))
acc += tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)
a_block_ptr = tl.advance(a_block_ptr, (0, BLOCK_SIZE_K))
b_block_ptr = tl.advance(b_block_ptr, (BLOCK_SIZE_K, 0))
t_block_ptr = tl.make_block_ptr(
base=t_base,
shape=(dim_m, dim_m),
strides=(t_m_stride, t_k_stride),
offsets=(0, 0),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_M),
order=(1, 0),
)
t = tl.load(t_block_ptr, boundary_check=(0, 1))
acc = tl.dot(t, acc, input_precision="tf32x3", out_dtype=tl.float32)
c_block_ptr = tl.make_block_ptr(
base=c_base,
shape=(dim_m, dim_n),
strides=(c_m_stride, c_n_stride),
offsets=(0, n_tile_start),
block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
order=(1, 0),
)
tl.store(c_block_ptr, acc, boundary_check=(0, 1))
@triton.autotune(
configs=_BMM_T_AUTOTUNE_CONFIGS,
key=["tune_m", "tune_split_k", "tune_n", "split_k_slices"],
prune_configs_by={"early_config_prune": _prune_bmm_t_split_k_configs},
restore_value=["c_ptr"],
warmup=5,
rep=15,
)
@triton.jit
def bmm_then_t_split_k_tf32x3_kernel(
a_ptr,
b_ptr,
t_ptr,
c_ptr,
bounds_ptr,
panel_end,
a_b_stride,
a_m_stride,
a_k_stride,
b_b_stride,
b_k_stride,
b_n_stride,
t_b_stride,
t_m_stride,
t_k_stride,
c_b_stride,
c_m_stride,
c_n_stride,
dim_m,
dim_k,
dim_n,
split_k_slices,
tune_m,
tune_split_k,
tune_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
specialize_constexpr: tl.constexpr,
):
pid_n = tl.program_id(0)
b_pid = tl.program_id(1)
split_pid = tl.program_id(2)
n_tile_start = pid_n * BLOCK_SIZE_N
active_n = dim_n
if specialize_constexpr:
bounds_ptr += b_pid
bound = tl.load(bounds_ptr)
active_n = bound - panel_end
if active_n <= 0:
return
if n_tile_start >= active_n:
return
split_size = tl.cdiv(dim_k, split_k_slices)
split_start = split_pid * split_size
split_end = min(split_start + split_size, dim_k)
a_base = a_ptr + b_pid * a_b_stride
b_base = b_ptr + b_pid * b_b_stride
t_base = t_ptr + b_pid * t_b_stride
c_base = c_ptr + b_pid * c_b_stride
offs_m = tl.arange(0, BLOCK_SIZE_M)
offs_k = tl.arange(0, BLOCK_SIZE_K)
offs_n = n_tile_start + tl.arange(0, BLOCK_SIZE_N)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_offset in range(0, tl.cdiv(split_size, BLOCK_SIZE_K)):
k_idxs = split_start + k_offset * BLOCK_SIZE_K + offs_k
a = tl.load(
a_base + offs_m[:, None] * a_m_stride + k_idxs[None, :] * a_k_stride,
mask=(offs_m[:, None] < dim_m) & (k_idxs[None, :] < split_end),
other=0.0,
)
b = tl.load(
b_base + k_idxs[:, None] * b_k_stride + offs_n[None, :] * b_n_stride,
mask=(k_idxs[:, None] < split_end) & (offs_n[None, :] < active_n),
other=0.0,
)
acc += tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)
offs_t = tl.arange(0, BLOCK_SIZE_M)
t = tl.load(
t_base + offs_m[:, None] * t_m_stride + offs_t[None, :] * t_k_stride,
mask=(offs_m[:, None] < dim_m) & (offs_t[None, :] < dim_m),
other=0.0,
)
acc = tl.dot(t, acc, input_precision="tf32x3", out_dtype=tl.float32)
tl.atomic_add(
c_base + offs_m[:, None] * c_m_stride + offs_n[None, :] * c_n_stride,
acc,
sem="relaxed",
mask=(offs_m[:, None] < dim_m) & (offs_n[None, :] < active_n),
)
def bmm_then_t_tf32x3(
v: torch.Tensor,
c: torch.Tensor,
t: torch.Tensor,
local_n_bounds: torch.Tensor,
panel_end: int,
specialize_constexpr: bool,
split_k: bool = False,
split_k_slices: int = 4,
) -> torch.Tensor:
B, K, M = v.shape
N = c.shape[2]
if split_k and split_k_slices > 1:
actual_split_k = min(split_k_slices, K)
split_k_size = triton.cdiv(K, actual_split_k)
out = torch.empty((B, M, N), device=v.device, dtype=v.dtype)
out.zero_()
grid = lambda META: ( # noqa
triton.cdiv(N, META["BLOCK_SIZE_N"]),
B,
actual_split_k,
)
bmm_then_t_split_k_tf32x3_kernel[grid](
v,
c,
t,
out,
local_n_bounds,
panel_end,
v.stride(0),
v.stride(2),
v.stride(1),
c.stride(0),
c.stride(1),
c.stride(2),
t.stride(0),
t.stride(2),
t.stride(1),
out.stride(0),
out.stride(1),
out.stride(2),
M,
K,
N,
actual_split_k,
_shape_bucket(M),
_shape_bucket(split_k_size),
_shape_bucket(N),
BLOCK_SIZE_M=_shape_bucket(M),
specialize_constexpr=specialize_constexpr,
)
return out
out = torch.empty((B, M, N), device=v.device, dtype=v.dtype)
grid = lambda META: ( # noqa
triton.cdiv(N, META["BLOCK_SIZE_N"]),
B,
)
bmm_then_t_tf32x3_kernel[grid](
v,
c,
t,
out,
local_n_bounds,
panel_end,
v.stride(0),
v.stride(2),
v.stride(1),
c.stride(0),
c.stride(1),
c.stride(2),
t.stride(0),
t.stride(2),
t.stride(1),
out.stride(0),
out.stride(1),
out.stride(2),
M,
K,
N,
_shape_bucket(M),
_shape_bucket(K),
_shape_bucket(N),
BLOCK_SIZE_M=_shape_bucket(M),
specialize_constexpr=specialize_constexpr,
)
return out
def baddbmm_sub_tf32x3(
c: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
local_bounds: torch.Tensor,
panel_end: int,
specialize_constexpr: bool,
c_source: torch.Tensor | None = None,
) -> None:
if c_source is None:
c_source = c
B, M, K = a.shape
N = b.shape[2]
grid = lambda META: ( # noqa
triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),
B,
)
_bmm_tf32x3_kernel[grid](
a,
b,
c_source,
c,
local_bounds,
panel_end,
a.stride(0),
a.stride(1),
a.stride(2),
b.stride(0),
b.stride(1),
b.stride(2),
c_source.stride(0),
c_source.stride(1),
c_source.stride(2),
c.stride(0),
c.stride(1),
c.stride(2),
M,
K,
N,
1,
_shape_bucket(M),
_shape_bucket(K),
_shape_bucket(N),
BLOCK_SIZE_K=_shape_bucket(K),
specialize_constexpr=specialize_constexpr,
)
@triton.jit
def tiled_qr_naive_kernel(
panel_src_ptr,
H_ptr,
tau_ptr,
V_ptr,
T_ptr,
bounds_ptr,
panel_src_b_stride,
panel_src_r_stride,
panel_src_c_stride,
H_b_stride,
H_r_stride,
H_c_stride,
V_b_stride,
V_r_stride,
V_c_stride,
T_b_stride,
T_r_stride,
T_c_stride,
tau_b_stride,
tau_r_stride,
col_start,
row_start,
n,
valid_cols,
TILE_SIZE: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
DO_EPILOGUE: tl.constexpr,
specialize_constexpr: tl.constexpr,
):
b_pid = tl.program_id(0)
if specialize_constexpr:
bounds_ptr += b_pid
bound = tl.load(bounds_ptr)
if col_start >= bound:
return
# offset to the start of the loop
panel_src_ptr += (
b_pid * panel_src_b_stride
+ row_start * panel_src_r_stride
+ col_start * panel_src_c_stride
)
H_ptr += b_pid * H_b_stride + row_start * H_r_stride + col_start * H_c_stride
tau_ptr += b_pid * tau_b_stride + col_start * tau_r_stride
V_ptr += b_pid * V_b_stride
T_ptr += b_pid * T_b_stride
row_offsets = tl.arange(0, BLOCK_SIZE)
col_offsets = tl.arange(0, TILE_SIZE)
rows = row_offsets[:, None]
cols = col_offsets[None, :]
m = n - col_start
tau_smem = tl.zeros((TILE_SIZE,), dtype=tl.float32)
mask = (rows < m) & (cols < valid_cols)
H_panel = tl.load(
panel_src_ptr + rows * panel_src_r_stride + cols * panel_src_c_stride,
mask=mask,
other=0.0,
) # [BLOCK_SIZE, TILE_SIZE]
for j in range(TILE_SIZE):
H_col = tl.sum(tl.where(cols == j, H_panel, 0.0), axis=1)
alpha = tl.sum(tl.where(row_offsets == j, H_col, 0.0))
tail = tl.where(row_offsets > j, H_col, 0.0)
tail_norm = tl.sum(tail * tail)
active = tail_norm != 0
x_norm = tl.sqrt(alpha * alpha + tail_norm)
beta_new = tl.where(alpha >= 0, -x_norm, x_norm)
beta = tl.where(active, beta_new, alpha)
denom = alpha - beta_new
denom_safe = tl.where(active, denom, 1.0)
tau_col = tl.where(active, (beta_new - alpha) / beta_new, 0.0)
scale = tl.where(active, 1.0 / denom_safe, 0.0)
col_value_tail = tail * scale
col_value = tl.where(row_offsets > j, col_value_tail, H_col)
col_value = tl.where(row_offsets == j, beta, col_value)
H_panel = tl.where(cols == j, col_value[:, None], H_panel)
tau_smem = tl.where(col_offsets == j, tau_col, tau_smem)
v = tl.where(row_offsets == j, 1.0, col_value_tail)
vtC = tl.sum(tl.where(cols > j, v[:, None] * H_panel, 0.0), axis=0)
H_applied = H_panel - tau_col * v[:, None] * vtC[None, :]
H_panel = tl.where(cols > j, H_applied, H_panel)
tl.store(H_ptr + rows * H_r_stride + cols * H_c_stride, mask=mask, value=H_panel)
store_offs = tl.arange(0, TILE_SIZE)
tl.store(
tau_ptr + store_offs * tau_r_stride, tau_smem, mask=store_offs < valid_cols
)
V_panel = tl.where(
rows == cols,
1.0,
tl.where(rows > cols, H_panel, 0.0),
)
tl.store(
V_ptr + rows * V_r_stride + cols * V_c_stride,
mask=mask,
value=V_panel,
)
if DO_EPILOGUE:
T_panel = tl.zeros([TILE_SIZE, TILE_SIZE], dtype=tl.float32)
T_rows = col_offsets[:, None]
T_cols = col_offsets[None, :]
T_panel = tl.where(T_rows == T_cols, tau_smem, T_panel)
for j in range(1, TILE_SIZE):
Vi = tl.sum(tl.where(cols == j, V_panel, 0.0), axis=1)
tau_col = tl.sum(tl.where(col_offsets == j, tau_smem, 0.0))
dots = tl.sum(V_panel * Vi[:, None], axis=0)
z = tl.where(col_offsets < j, -tau_col * dots, 0.0)
t = tl.sum(T_panel * z[None, :], axis=1)
T_panel = tl.where(
(T_rows < j) & (T_cols == j),
t[:, None],
T_panel,
)
tl.store(
T_ptr + T_rows * T_r_stride + T_cols * T_c_stride,
mask=(T_rows < valid_cols) & (T_cols < valid_cols),
value=T_panel,
)
# this is only for the n = 32 case
@triton.jit
def tiled_qr_naive_kernel_smol_h(
panel_src_ptr,
H_ptr,
tau_ptr,
panel_src_b_stride,
panel_src_r_stride,
panel_src_c_stride,
H_b_stride,
H_r_stride,
H_c_stride,
tau_b_stride,
tau_r_stride,
col_start,
row_start,
n,
valid_cols,
TILE_SIZE: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
b_pid = tl.program_id(0)
n = n.to(tl.int32)
col_start = col_start.to(tl.int32)
valid_cols = valid_cols.to(tl.int32)
# offset to the start of the loop
panel_src_ptr += (
b_pid * panel_src_b_stride
+ row_start * panel_src_r_stride
+ col_start * panel_src_c_stride
)
H_ptr += b_pid * H_b_stride + row_start * H_r_stride + col_start * H_c_stride
tau_ptr += b_pid * tau_b_stride + col_start * tau_r_stride
row_offsets = tl.arange(0, BLOCK_SIZE).to(tl.int32)
col_offsets = tl.arange(0, TILE_SIZE).to(tl.int32)
rows = row_offsets[:, None]
cols = col_offsets[None, :]
m = n - col_start
tau_smem = tl.zeros((TILE_SIZE,), dtype=tl.float32)
mask = (rows < m) & (cols < valid_cols)
H_panel = tl.load(
panel_src_ptr + rows * panel_src_r_stride + cols * panel_src_c_stride,
mask=mask,
other=0.0,
) # [BLOCK_SIZE, TILE_SIZE]
for j in range(TILE_SIZE):
H_col = tl.sum(tl.where(cols == j, H_panel, 0.0), axis=1)
alpha = tl.sum(tl.where(row_offsets == j, H_col, 0.0))
tail = tl.where(row_offsets > j, H_col, 0.0)
tail_norm = tl.sum(tail * tail)
active = tail_norm != 0
x_norm = tl.sqrt(alpha * alpha + tail_norm)
beta_new = tl.where(alpha >= 0, -x_norm, x_norm)
beta = tl.where(active, beta_new, alpha)
denom = alpha - beta_new
denom_safe = tl.where(active, denom, 1.0)
tau_col = tl.where(active, (beta_new - alpha) / beta_new, 0.0)
scale = tl.where(active, 1.0 / denom_safe, 0.0)
col_value_tail = tail * scale
col_value = tl.where(row_offsets > j, col_value_tail, H_col)
col_value = tl.where(row_offsets == j, beta, col_value)
H_panel = tl.where(cols == j, col_value[:, None], H_panel)
tau_smem = tl.where(col_offsets == j, tau_col, tau_smem)
v = tl.where(row_offsets == j, 1.0, col_value_tail)
vtC = tl.sum(tl.where(cols > j, v[:, None] * H_panel, 0.0), axis=0)
H_applied = H_panel - tau_col * v[:, None] * vtC[None, :]
H_panel = tl.where(cols > j, H_applied, H_panel)
tl.store(H_ptr + rows * H_r_stride + cols * H_c_stride, mask=mask, value=H_panel)
store_offs = tl.arange(0, TILE_SIZE).to(tl.int32)
tl.store(
tau_ptr + store_offs * tau_r_stride, tau_smem, mask=store_offs < valid_cols
)
def launch_tiled_qr_panel(
H: torch.Tensor, # [b,n,n]
tau: torch.Tensor, # [b,n]
V: torch.Tensor, # [b, n - col_start, tile_size]
T: torch.Tensor, # [b,tile_size,tile_size]
local_bounds: torch.Tensor, # [b]
col_start: int,
specialize_constexpr: bool,
tile_size: int = 16,
panel_source: torch.Tensor | None = None,
):
B, n, _ = H.shape
grid = (B,)
if panel_source is None:
panel_source = H
row_start = col_start
actual_block_size = n - col_start
block_size = triton.next_power_of_2(actual_block_size)
valid_cols = min(tile_size, n - col_start)
if n == 32:
tiled_qr_naive_kernel_smol_h[grid](
panel_source,
H,
tau,
panel_source.stride(0),
panel_source.stride(1),
panel_source.stride(2),
H.stride(0),
H.stride(1),
H.stride(2),
tau.stride(0),
tau.stride(1),
col_start,
row_start,
n,
valid_cols,
TILE_SIZE=tile_size, # type: ignore
BLOCK_SIZE=block_size, # type: ignore
num_warps=1,
)
return
# best config a quick scan
if block_size <= 128:
num_warps = 1
elif block_size <= 256:
num_warps = 2
elif block_size < 1024:
num_warps = 4
else:
num_warps = 8
# fmt: off
tiled_qr_naive_kernel[grid](
panel_source, H, tau, V, T, local_bounds,
panel_source.stride(0), panel_source.stride(1), panel_source.stride(2),
H.stride(0), H.stride(1), H.stride(2),
V.stride(0), V.stride(1), V.stride(2),
T.stride(0), T.stride(1), T.stride(2),
tau.stride(0), tau.stride(1),
col_start, row_start, n, valid_cols,
TILE_SIZE=tile_size, # type: ignore
BLOCK_SIZE=block_size, # type: ignore
DO_EPILOGUE=((col_start + valid_cols) < n), # type: ignore
specialize_constexpr = specialize_constexpr, # type: ignore
num_warps = num_warps, # type: ignore
)
# fmt: on
@triton.jit
def check_structure_kernel(
A_ptr,
local_bounds_ptr,
structure_meta_ptr,
A_stride_b,
A_stride_r,
A_stride_c,
local_stride_b,
BLOCK_SIZE_ROWS: tl.constexpr,
BLOCK_SIZE_COLS: tl.constexpr,
ROWS: tl.constexpr,
RANK_OFFSET: tl.constexpr,
HEAD_BOUND: tl.constexpr,
CLUSTER_BOUND: tl.constexpr,
CLUSTERED_BOUND: tl.constexpr,
):
pid = tl.program_id(0)
A_ptr += A_stride_b * pid
local_bounds_ptr += local_stride_b * pid
offs_rows = tl.arange(0, BLOCK_SIZE_ROWS)[:, None]
offs_cols = tl.arange(0, BLOCK_SIZE_COLS)[None, :]
head_max = tl.full((), 0.0, tl.float32)
cluster_tail_max = tl.full((), 0.0, tl.float32)
rank_tail_max = tl.full((), 0.0, tl.float32)
for row_start in range(0, ROWS, BLOCK_SIZE_ROWS):
rows = row_start + offs_rows
row_mask = rows < ROWS
head_cols = offs_cols
head_mask = row_mask & (head_cols < HEAD_BOUND)
head = tl.load(
A_ptr + rows * A_stride_r + head_cols * A_stride_c,
mask=head_mask,
other=0.0,
)
head_abs = tl.where(head_mask, tl.abs(head), 0.0)
head_max = tl.maximum(head_max, tl.max(tl.max(head_abs, axis=0), axis=0))
tail_cols = CLUSTER_BOUND + offs_cols
tail_mask = row_mask & (tail_cols < ROWS)
tail = tl.load(
A_ptr + rows * A_stride_r + tail_cols * A_stride_c,
mask=tail_mask,
other=0.0,
)
tail_abs = tl.where(tail_mask, tl.abs(tail), 0.0)
cluster_tail_max = tl.maximum(
cluster_tail_max, tl.max(tl.max(tail_abs, axis=0), axis=0)
)
rank_tail_abs = tl.where(tail_cols >= RANK_OFFSET, tail_abs, 0.0)
rank_tail_max = tl.maximum(
rank_tail_max, tl.max(tl.max(rank_tail_abs, axis=0), axis=0)
)
bound = tl.where(rank_tail_max == 0.0, RANK_OFFSET, ROWS)
clustered = cluster_tail_max <= tl.maximum(head_max, 1.0e-30) * 1.0e-5
bound = tl.where(clustered, tl.minimum(bound, CLUSTERED_BOUND), bound)
tl.store(local_bounds_ptr, bound)
tl.atomic_max(structure_meta_ptr, bound, sem="relaxed")
tl.atomic_max(structure_meta_ptr + 1, tl.where(bound < ROWS, 1, 0), sem="relaxed")
@triton.jit
def cleanup_structure_tail_kernel(
H_ptr,
tau_ptr,
local_bounds_ptr,
H_b_stride,
H_r_stride,
H_c_stride,
tau_b_stride,
tau_r_stride,
local_stride_b,
n,
BLOCK_SIZE_ROWS: tl.constexpr,
BLOCK_SIZE_COLS: tl.constexpr,
):
b_pid = tl.program_id(0)
row_pid = tl.program_id(1)
col_pid = tl.program_id(2)
bound = tl.load(local_bounds_ptr + b_pid * local_stride_b)
col_start = col_pid * BLOCK_SIZE_COLS
if col_start + BLOCK_SIZE_COLS <= bound:
return
rows = row_pid * BLOCK_SIZE_ROWS + tl.arange(0, BLOCK_SIZE_ROWS)
cols = col_start + tl.arange(0, BLOCK_SIZE_COLS)
mask = (rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] >= bound)
tl.store(
H_ptr
+ b_pid * H_b_stride
+ rows[:, None] * H_r_stride
+ cols[None, :] * H_c_stride,
0.0,
mask=mask,
)
if row_pid == 0:
tau_mask = (cols < n) & (cols >= bound)
tl.store(
tau_ptr + b_pid * tau_b_stride + cols * tau_r_stride,
0.0,
mask=tau_mask,
)
def geqrf_blocked_batched(
A: torch.Tensor,
tile_size: int = 16,
split_k: bool = False,
split_k_slices: int = 4,
exploit_structure: bool = True,
):
if A.ndim != 3 or A.shape[-1] != A.shape[-2]:
raise ValueError("Expected A with shape (batch, n, n)")
if tile_size <= 0:
raise ValueError("tile_size must be positive")
H = A.new_empty(A.shape)
B, n, _ = H.shape
tau = H.new_empty(B, n)
# super early exit on n = 32 case, we can save some allocations and writes (V & T)
if n == 32:
launch_tiled_qr_panel(
H,
tau,
None, # type: ignore
None, # type: ignore
None, # type: ignore
0,
False,
tile_size,
panel_source=A,
)
return H, tau
if n >= 2048:
tile_size = 16
if n > 2048:
split_k_slices = 8
T = H.new_empty(B, tile_size, tile_size)
if n in {512}:
split_k = False
# these are the cases we scan for rankdef, clustered, or mixed
cases_to_specialize = {512}
if exploit_structure and n in cases_to_specialize:
local_n_bounds = A.new_empty((B,), dtype=torch.long)
else:
local_n_bounds = A.new_empty((1,), dtype=torch.long)
maybe_local_bounds = False
structure_meta = None
if exploit_structure and n in cases_to_specialize:
rank = (3 * n) // 4
cluster_bound = n // 2 + 2
head_bound = n // 2 - 2
clustered_bound = min(n, triton.cdiv(cluster_bound, tile_size) * tile_size)
block_size_cols = triton.next_power_of_2(n - cluster_bound)
structure_meta = A.new_zeros((2,), dtype=torch.long)
grid = (B,)
check_structure_kernel[grid](
A,
local_n_bounds,
structure_meta,
A.stride(0),
A.stride(1),
A.stride(2),
local_n_bounds.stride(0),
BLOCK_SIZE_ROWS=64, # type: ignore
BLOCK_SIZE_COLS=block_size_cols, # type: ignore
ROWS=n, # type: ignore
RANK_OFFSET=rank, # type: ignore
HEAD_BOUND=head_bound, # type: ignore
CLUSTER_BOUND=cluster_bound, # type: ignore
CLUSTERED_BOUND=clustered_bound, # type: ignore
num_warps=8, # type: ignore
)
maybe_local_bounds = True
max_n_bound = int(structure_meta[0].item()) if maybe_local_bounds else n
specialize_constexpr = maybe_local_bounds
for k in range(0, max_n_bound, tile_size):
ib = min(tile_size, n - k)
panel_end = k + ib
is_first_panel = k == 0
V = H.new_empty(B, n - k, ib)
launch_tiled_qr_panel(
H,
tau,
V,
T,
local_n_bounds,
k,
specialize_constexpr,
tile_size,
panel_source=A if is_first_panel else H,
)
if panel_end < n:
C = H[:, k:, panel_end:n]
C_source = A[:, k:, panel_end:n] if is_first_panel else C
split_k_panel = False
if split_k:
trailing_n = max_n_bound - panel_end
block_n_est = 128
num_w_ctas = B * triton.cdiv(trailing_n, block_n_est)
split_k_panel = split_k and (num_w_ctas < 300) and (trailing_n >= 512)
W = bmm_then_t_tf32x3(
V,
C_source,
T,
local_n_bounds,
panel_end,
specialize_constexpr,
split_k=split_k_panel,
split_k_slices=split_k_slices,
)
baddbmm_sub_tf32x3(
C,
V,
W,
local_n_bounds,
panel_end,
specialize_constexpr,
c_source=C_source,
)
if maybe_local_bounds and int(structure_meta[1].item()) != 0:
cleanup_grid = (
B,
triton.cdiv(n, 16),
triton.cdiv(n, 32),
)
cleanup_structure_tail_kernel[cleanup_grid](
H,
tau,
local_n_bounds,
H.stride(0),
H.stride(1),
H.stride(2),
tau.stride(0),
tau.stride(1),
local_n_bounds.stride(0),
n,
BLOCK_SIZE_ROWS=16, # type: ignore
BLOCK_SIZE_COLS=32, # type: ignore
num_warps=4, # type: ignore
)
return H, tau
def custom_kernel(data: input_t) -> output_t:
return geqrf_blocked_batched(
data,
tile_size=32,
split_k=True,
split_k_slices=4,
exploit_structure=True,
)
scrolls · 1317 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