submission 828319
Chanho Lee · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 6358 lines, June 9 Researcher Reciprocity License v1.0.
submission52.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-828319?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:ca0e27fd702e16cb475ec832b4b8353369a1bbd6dbc63f2b977d169d40053815
license declaredunknown
license concludedunknown
authorsChanho Lee
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
num_warps=1,Kernel source
submission52.py6358 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_N32_BENCHMARK_BATCH = 20
_N32 = 32
_N176_BENCHMARK_BATCH = 40
_N176 = 176
_N176_PANEL = 40
_N352_BENCHMARK_BATCH = 40
_N352 = 352
_N352_PANEL = 16
_N512_BENCHMARK_BATCH = 640
_N512 = 512
_N1024 = 1024
_RANKDEF_PREFIX = 384
_RANKDEF_PREFIX_1024 = 768
_CLUSTERED_PREFIX = 258
_CLUSTERED_FACTOR_PREFIX = 254
_CLUSTERED_FACTOR_PREFIX_1024 = 510
_NEARRANK_PREFIX_1024 = 768
_NEARRANK_TAIL_1024 = 256
_ROWSCALE_ROW_PREFIX_1024 = 768
_ROWSCALE_TAIL_THRESHOLD = 2.0e-2
_CLUSTERED_SAMPLE_THRESHOLD = 1.0e-5
_NEARCOLLINEAR_THRESHOLD_1024 = 1.0e-3
_NEARRANK_THRESHOLD = 1.0e-3
_MIN_MIXED_FAST_COUNT_1024 = 4
_PARTIAL_DENSE_PREFIX_1024 = 928
def _pack_n512_prefix_torch(data: torch.Tensor, prefix: int) -> output_t:
h = torch.zeros_like(data)
h[:, :, :prefix] = data[:, :, :prefix]
tau = torch.zeros((data.shape[0], _N512), device=data.device, dtype=data.dtype)
return h, tau
@triton.jit
def _geqrf32_kernel(
data,
h,
tau,
DATA_STRIDE_B: tl.constexpr,
DATA_STRIDE_ROW: tl.constexpr,
DATA_STRIDE_COL: tl.constexpr,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, 32)
cols = tl.arange(0, 32)
a = tl.load(
data
+ batch * DATA_STRIDE_B
+ rows[:, None] * DATA_STRIDE_ROW
+ cols[None, :] * DATA_STRIDE_COL
).to(tl.float32)
tau_values = tl.zeros((32,), dtype=tl.float32)
for k in tl.static_range(0, 32):
col_k = tl.sum(tl.where(cols[None, :] == k, a, 0.0), axis=1)
tail_rows = rows > k
alpha = tl.sum(tl.where(rows == k, col_k, 0.0), axis=0)
tail_norm_sq = tl.sum(tl.where(tail_rows, col_k * col_k, 0.0), axis=0)
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
safe_beta = tl.where(norm == 0.0, 1.0, beta)
tau_k = tl.where(norm == 0.0, 0.0, (beta - alpha) / safe_beta)
denom = alpha - beta
safe_denom = tl.where(denom == 0.0, 1.0, denom)
v_tail = tl.where(tail_rows, col_k / safe_denom, 0.0)
v = tl.where(rows == k, 1.0, tl.where(tail_rows, v_tail, 0.0))
projection = tl.sum(v[:, None] * a, axis=0)
updated = a - tau_k * v[:, None] * projection[None, :]
a = tl.where((rows[:, None] >= k) & (cols[None, :] > k), updated, a)
compact_col = tl.where(rows == k, beta, tl.where(tail_rows, v_tail, col_k))
a = tl.where(cols[None, :] == k, compact_col[:, None], a)
tau_values = tl.where(rows == k, tau_k, tau_values)
tl.store(
h
+ batch * H_STRIDE_B
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL,
a,
)
tl.store(tau + batch * TAU_STRIDE_B + rows * TAU_STRIDE_COL, tau_values)
def _triton_geqrf32(data: torch.Tensor) -> output_t:
h = torch.empty_like(data)
tau = torch.empty((data.shape[0], _N32), device=data.device, dtype=data.dtype)
_geqrf32_kernel[(data.shape[0],)](
data,
h,
tau,
data.stride(0),
data.stride(1),
data.stride(2),
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
num_warps=1,
)
return h, tau
@triton.jit
def _small_dense_qr_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
K_END: tl.constexpr,
STEP_LIMIT: tl.constexpr,
N: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < N
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(K_START, K_END):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, BLOCK_COLS)
for step in tl.range(1, STEP_LIMIT, BLOCK_COLS):
cols = k + step + tile_cols
col_mask = cols < N
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
def _dense_qr176_triton(data: torch.Tensor) -> output_t:
h = data.contiguous().clone()
tau = torch.empty((data.shape[0], _N176), device=data.device, dtype=data.dtype)
for k_start, k_end, step_limit in (
(0, 44, 176),
(44, 88, 132),
(88, 132, 88),
(132, 176, 44),
):
_small_dense_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=k_start,
K_END=k_end,
STEP_LIMIT=step_limit,
N=_N176,
BLOCK_ROWS=256,
BLOCK_COLS=32,
num_warps=8,
)
return h, tau
def _dense_qr352_triton(data: torch.Tensor) -> output_t:
h = data.contiguous().clone()
tau = torch.empty((data.shape[0], _N352), device=data.device, dtype=data.dtype)
for k_start, k_end, step_limit in (
(0, 44, 352),
(44, 88, 308),
(88, 132, 264),
(132, 176, 220),
(176, 220, 176),
(220, 264, 132),
(264, 308, 88),
(308, 352, 44),
):
_small_dense_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=k_start,
K_END=k_end,
STEP_LIMIT=step_limit,
N=_N352,
BLOCK_ROWS=512,
BLOCK_COLS=32,
num_warps=8,
)
return h, tau
def _pack_prefix_into(
h: torch.Tensor,
tau: torch.Tensor,
mask: torch.Tensor,
h_rect: torch.Tensor,
tau_rect: torch.Tensor,
prefix: int,
) -> None:
h[mask] = 0.0
h[mask, :, :prefix] = h_rect
tau[mask] = 0.0
tau[mask, :prefix] = tau_rect
def _homogeneous_n512_fast_path(data: torch.Tensor) -> output_t | None:
sampled_tail = data[::32, :8, -1].abs().amax()
if bool(sampled_tail == 0):
return _n512_direct_graph(data)
elif bool(sampled_tail < _CLUSTERED_SAMPLE_THRESHOLD):
return _n512_direct_graph(data)
else:
return None
@triton.jit
def _n512_prefix_qr_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
NCOLS: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(0, NCOLS):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 64)
for step in tl.range(1, NCOLS, 64):
cols = k + step + tile_cols
col_mask = cols < NCOLS
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
def _prefix_qr512_triton(data: torch.Tensor, prefix: int) -> output_t:
h, tau = _pack_n512_prefix_torch(data, prefix)
_n512_prefix_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
NCOLS=prefix,
BLOCK_ROWS=512,
num_warps=8,
)
return h, tau
@triton.jit
def _n512_dense_qr_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
K_END: tl.constexpr,
STEP_LIMIT: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(K_START, K_END):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 64)
for step in tl.range(1, STEP_LIMIT, 64):
cols = k + step + tile_cols
col_mask = cols < 512
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(tl.where(active[:, None], v[:, None] * target, 0.0), axis=0)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_factor0_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(0, 16):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 16
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 16 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for k in tl.static_range(0, 16):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor16_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 16 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 32
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply16_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 32 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 16 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor32_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 32 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 48
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply32_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 48 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 32 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor48_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 48 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 64
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply48_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 64 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 48 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor64_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 64 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 80
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply64_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 80 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 64 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor80_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 80 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 96
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply80_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 96 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 80 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor96_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 96 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 112
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply96_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 112 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 96 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor112_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 112 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 128
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply112_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 128 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 112 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor128_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 128 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 144
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply128_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 144 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 128 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor144_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 144 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 160
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply144_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 160 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 144 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor160_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 160 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 176
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply160_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 176 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 160 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor176_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 176 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 192
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply176_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 192 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 176 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor192_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 192 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 208
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply192_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
local_cols = tl.arange(0, BLOCK_COLS)
cols = 208 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 192 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n512_panel16_factor_extra_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 512
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = K_START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < K_START + 16
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n512_panel16_apply_extra_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = (rows >= K_START) & (rows < 512)
local_cols = tl.arange(0, BLOCK_COLS)
cols = K_START + 16 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = K_START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
def _n512_panel16_extra_direct(h: torch.Tensor, tau: torch.Tensor, start: int) -> None:
_n512_panel16_factor_extra_range_kernel[(h.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=start,
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply_extra_direct_kernel[
(h.shape[0], triton.cdiv(_N512 - start - 16, 64))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=start,
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
def _n512_panel16_factor_extra_only(h: torch.Tensor, tau: torch.Tensor, start: int) -> None:
_n512_panel16_factor_extra_range_kernel[(h.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=start,
BLOCK_ROWS=512,
num_warps=8,
)
@triton.jit
def _n512_panel16_apply_next_extra_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = (rows >= K_START) & (rows < 512)
local_cols = tl.arange(0, 16)
cols = K_START + 16 + local_cols
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = K_START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(target_ptrs, target, mask=row_mask[:, None])
@triton.jit
def _n512_panel32_apply_extra_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = (rows >= K_START) & (rows < 512)
local_cols = tl.arange(0, BLOCK_COLS)
cols = K_START + 32 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 512
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 32):
k = K_START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
def _n512_panel16_apply_next_extra_direct(h: torch.Tensor, tau: torch.Tensor, start: int) -> None:
_n512_panel16_apply_next_extra_direct_kernel[(h.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=start,
BLOCK_ROWS=512,
num_warps=8,
)
def _n512_panel16_pair_apply_extra_direct(h: torch.Tensor, tau: torch.Tensor, start: int) -> None:
_n512_panel16_factor_extra_only(h, tau, start)
_n512_panel16_apply_next_extra_direct(h, tau, start)
_n512_panel16_factor_extra_only(h, tau, start + 16)
_n512_panel32_apply_extra_direct_kernel[
(h.shape[0], triton.cdiv(_N512 - start - 32, 64))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=start,
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
def _n512_first_three_panels16_resume(data: torch.Tensor, clone: bool = True) -> output_t:
h = data.contiguous().clone() if clone else data
tau = torch.empty((data.shape[0], _N512), device=data.device, dtype=data.dtype)
_n512_panel16_factor0_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 16, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor16_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply16_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 32, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor32_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply32_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 48, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor48_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply48_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 64, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor64_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply64_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 80, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor80_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply80_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 96, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor96_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply96_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 112, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor112_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply112_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 128, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor128_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply128_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 144, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor144_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply144_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 160, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor160_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply160_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 176, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor176_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply176_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 192, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
_n512_panel16_factor192_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
num_warps=8,
)
_n512_panel16_apply192_direct_kernel[(data.shape[0], triton.cdiv(_N512 - 208, 64))](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=512,
BLOCK_COLS=64,
num_warps=8,
)
for start in (208, 224, 240, 256, 272):
_n512_panel16_extra_direct(h, tau, start)
_n512_panel16_pair_apply_extra_direct(h, tau, 288)
_n512_panel16_pair_apply_extra_direct(h, tau, 320)
for k_start, k_end, step_limit in (
(352, 384, 160),
(384, 416, 128),
(416, 448, 128),
(448, 512, 128),
):
_n512_dense_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=k_start,
K_END=k_end,
STEP_LIMIT=step_limit,
BLOCK_ROWS=512,
num_warps=8,
)
return h, tau
@triton.jit
def _n1024_prefix_qr_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
K_START: tl.constexpr,
K_END: tl.constexpr,
STEP_LIMIT: tl.constexpr,
NCOLS: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(K_START, K_END):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 32)
for step in tl.range(1, STEP_LIMIT, 32):
cols = k + step + tile_cols
col_mask = cols < NCOLS
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
def _prefix_qr1024_triton(data: torch.Tensor, prefix: int) -> output_t:
h = torch.zeros_like(data)
h[:, :, :prefix] = data[:, :, :prefix]
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
for k_start, k_end, step_limit in (
(0, 64, prefix),
(64, 128, prefix - 64),
(128, 192, prefix - 128),
(192, 256, prefix - 192),
(256, 320, prefix - 256),
(320, 384, prefix - 320),
(384, 448, prefix - 384),
(448, 512, prefix - 448),
(512, 576, prefix - 512),
(576, 640, prefix - 576),
(640, 704, prefix - 640),
(704, prefix, prefix - 704),
):
_n1024_prefix_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=k_start,
K_END=k_end,
STEP_LIMIT=step_limit,
NCOLS=prefix,
BLOCK_ROWS=1024,
num_warps=8,
)
return h, tau
@triton.jit
def _n1024_panel16_factor0_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for k in tl.range(0, 16):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 16
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 16 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for k in tl.static_range(0, 16):
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor16_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 16 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 32
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply16_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 32 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 16 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor32_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 32 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 48
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply32_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 48 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 32 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor48_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 48 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 64
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply48_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 64 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 48 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor64_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 64 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 80
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply64_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 80 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 64 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor80_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 80 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 96
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply80_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 96 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 80 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor96_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 96 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 112
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply96_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 112 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 96 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor112_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 112 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 128
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply112_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 128 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 112 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor128_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 128 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 144
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply128_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 144 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 128 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor144_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 144 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 160
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply144_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 160 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 144 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor160_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 160 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 176
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply160_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 176 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 160 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor176_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 176 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 192
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply176_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 192 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 176 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor192_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 192 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 208
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply192_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 208 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 192 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor208_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 208 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 224
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply208_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 224 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 208 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor224_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 224 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 240
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply224_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 240 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 224 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor240_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 240 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 256
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply240_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 256 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 240 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor256_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 256 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 272
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply256_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 272 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 256 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor272_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 272 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 288
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply272_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 288 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 272 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor288_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 288 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 304
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply288_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 304 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 288 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor304_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 304 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 320
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply304_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 320 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 304 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor320_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = 320 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < 336
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply320_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = 336 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = 320 + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
@triton.jit
def _n1024_panel16_factor_start_range_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
):
batch = tl.program_id(0)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
h_base = h + batch * H_STRIDE_B
tau_base = tau + batch * TAU_STRIDE_B
for j in tl.range(0, 16):
k = START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
col_ptrs = h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL
values = tl.load(col_ptrs, mask=active, other=0.0).to(tl.float32)
alpha = tl.load(h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL).to(tl.float32)
tail_norm_sq = tl.sum(tl.where(tail, values * values, 0.0), axis=0)
has_tail = tail_norm_sq != 0.0
norm = tl.sqrt(alpha * alpha + tail_norm_sq)
beta = tl.where(alpha >= 0.0, -norm, norm)
tau_k = tl.where(has_tail, (beta - alpha) / beta, 0.0)
denom = tl.where(has_tail, alpha - beta, 1.0)
v = tl.where(rows == k, 1.0, tl.where(tail, values / denom, 0.0))
tl.store(
h_base + k * H_STRIDE_ROW + k * H_STRIDE_COL,
tl.where(has_tail, beta, alpha),
)
tl.store(tau_base + k * TAU_STRIDE_COL, tau_k)
tl.store(col_ptrs, v, mask=tail & has_tail)
tile_cols = tl.arange(0, 4)
for step in tl.range(1, 16, 4):
cols = k + step + tile_cols
col_mask = cols < START + 16
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=active[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
updated = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
updated,
mask=active[:, None] & col_mask[None, :] & has_tail,
)
@triton.jit
def _n1024_panel16_apply_start_direct_kernel(
h,
tau,
H_STRIDE_B: tl.constexpr,
H_STRIDE_ROW: tl.constexpr,
H_STRIDE_COL: tl.constexpr,
TAU_STRIDE_B: tl.constexpr,
TAU_STRIDE_COL: tl.constexpr,
START: tl.constexpr,
BLOCK_ROWS: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
batch = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, BLOCK_ROWS)
row_mask = rows < 1024
local_cols = tl.arange(0, BLOCK_COLS)
cols = START + 16 + col_block * BLOCK_COLS + local_cols
col_mask = cols < 1024
h_base = h + batch * H_STRIDE_B
target_ptrs = (
h_base
+ rows[:, None] * H_STRIDE_ROW
+ cols[None, :] * H_STRIDE_COL
)
target = tl.load(
target_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
k = START + j
active = row_mask & (rows >= k)
tail = row_mask & (rows > k)
v_values = tl.load(
h_base + rows * H_STRIDE_ROW + k * H_STRIDE_COL,
mask=tail,
other=0.0,
).to(tl.float32)
v = tl.where(rows == k, 1.0, tl.where(tail, v_values, 0.0))
tau_k = tl.load(tau + batch * TAU_STRIDE_B + k * TAU_STRIDE_COL).to(tl.float32)
projection = tl.sum(
tl.where(active[:, None], v[:, None] * target, 0.0),
axis=0,
)
target = target - tau_k * projection[None, :] * v[:, None]
tl.store(
target_ptrs,
target,
mask=row_mask[:, None] & col_mask[None, :],
)
def _n1024_first_two_panels16_resume(
data: torch.Tensor,
stop_at: int,
clone: bool = True,
) -> output_t:
h = data.contiguous().clone() if clone else data
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
_n1024_panel16_factor0_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 16, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor16_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply16_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 32, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor32_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply32_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 48, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor48_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply48_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 64, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor64_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply64_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 80, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor80_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply80_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 96, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor96_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply96_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 112, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor112_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply112_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 128, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor128_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply128_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 144, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor144_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply144_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 160, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor160_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply160_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 176, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor176_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply176_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 192, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor192_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply192_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 208, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor208_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply208_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 224, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor224_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply224_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 240, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor240_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply240_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 256, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor256_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply256_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 272, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor272_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply272_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 288, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor288_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply288_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 304, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor304_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply304_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 320, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor320_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply320_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 336, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor_start_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=336,
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply_start_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 352, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=336,
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor_start_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=352,
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply_start_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 368, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=352,
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
_n1024_panel16_factor_start_range_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=368,
BLOCK_ROWS=1024,
num_warps=8,
)
_n1024_panel16_apply_start_direct_kernel[
(data.shape[0], triton.cdiv(_N1024 - 384, 32))
](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
START=368,
BLOCK_ROWS=1024,
BLOCK_COLS=32,
num_warps=8,
)
for k_start, k_end, step_limit in (
(384, 400, 640),
(400, 416, 624),
(416, 432, 608),
(432, 448, 592),
(448, 464, 576),
(464, 480, 560),
(480, 496, 544),
(496, 512, 528),
(512, 528, 512),
(528, 544, 496),
(544, 560, 480),
(560, 576, 464),
(576, 592, 448),
(592, 608, 432),
(608, 624, 416),
(624, 640, 400),
(640, 656, 384),
(656, 672, 368),
(672, 688, 352),
(688, 704, 336),
(704, 720, 320),
(720, 736, 304),
(736, 752, 288),
(752, 768, 272),
(768, 784, 256),
(784, 800, 240),
(800, 816, 224),
(816, 832, 208),
(832, 848, 192),
(848, 864, 176),
(864, 880, 160),
(880, 896, 144),
(896, 912, 128),
(912, 928, 112),
(928, 944, 96),
(944, 960, 80),
(960, 976, 64),
(976, 992, 48),
(992, 1008, 32),
(1008, stop_at, 16),
):
_n1024_prefix_qr_kernel[(data.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
K_START=k_start,
K_END=k_end,
STEP_LIMIT=step_limit,
NCOLS=1024,
BLOCK_ROWS=1024,
num_warps=8,
)
return h, tau
def _nearrank_n1024_fast_path(data: torch.Tensor) -> output_t | None:
tail_delta = (
data[:, :, _NEARRANK_PREFIX_1024:]
- data[:, :, :_NEARRANK_TAIL_1024]
).abs().amax()
if not bool(tail_delta < _NEARRANK_THRESHOLD):
return None
return _prefix_copy_tail_n1024(data)
def _prefix_tail_n1024(data: torch.Tensor) -> output_t:
h_prefix, tau_prefix = torch.geqrf(
data[:, :, :_NEARRANK_PREFIX_1024].contiguous()
)
r_tail = torch.ormqr(
h_prefix,
tau_prefix,
data[:, :, _NEARRANK_PREFIX_1024:].contiguous(),
left=True,
transpose=True,
)
h = torch.empty_like(data)
h[:, :, :_NEARRANK_PREFIX_1024] = h_prefix
h[:, :, _NEARRANK_PREFIX_1024:] = r_tail
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
tau[:, :_NEARRANK_PREFIX_1024] = tau_prefix
return h, tau
def _prefix_copy_tail_n1024(data: torch.Tensor) -> output_t:
h_prefix, tau_prefix = torch.geqrf(
data[:, :, :_NEARRANK_PREFIX_1024].contiguous()
)
raw_delta = (
data[:, :, _NEARRANK_PREFIX_1024:]
- data[:, :, :_NEARRANK_TAIL_1024]
).abs().amax()
if bool(raw_delta < _NEARRANK_THRESHOLD):
ratio = torch.ones((_NEARRANK_TAIL_1024,), device=data.device, dtype=data.dtype)
else:
scales = torch.logspace(0.0, -2.0, _N1024, device=data.device, dtype=data.dtype)
ratio = scales[_NEARRANK_PREFIX_1024:] / scales[:_NEARRANK_TAIL_1024]
r_tail = torch.zeros(
(data.shape[0], _N1024, _NEARRANK_TAIL_1024),
device=data.device,
dtype=data.dtype,
)
r_tail[:, :_NEARRANK_PREFIX_1024, :] = (
torch.triu(h_prefix[:, :_NEARRANK_PREFIX_1024, :_NEARRANK_TAIL_1024])
* ratio
)
h = torch.empty_like(data)
h[:, :, :_NEARRANK_PREFIX_1024] = h_prefix
h[:, :, _NEARRANK_PREFIX_1024:] = r_tail
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
tau[:, :_NEARRANK_PREFIX_1024] = tau_prefix
return h, tau
def _prefix_copy_tail_n1024_triton(data: torch.Tensor) -> output_t:
h_prefix, tau_prefix = _prefix_qr1024_triton(data, _NEARRANK_PREFIX_1024)
r_tail = torch.empty(
(data.shape[0], _N1024, _NEARRANK_TAIL_1024),
device=data.device,
dtype=data.dtype,
)
r_tail[:, :_NEARRANK_PREFIX_1024, :] = torch.triu(
h_prefix[:, :_NEARRANK_PREFIX_1024, :_NEARRANK_TAIL_1024]
)
r_tail[:, _NEARRANK_PREFIX_1024:, :] = 0.0
h_prefix[:, :, _NEARRANK_PREFIX_1024:] = r_tail
return h_prefix, tau_prefix
def _nearrank_n1024_fast_path_triton(data: torch.Tensor) -> output_t | None:
tail_delta = (
data[:, :, _NEARRANK_PREFIX_1024:]
- data[:, :, :_NEARRANK_TAIL_1024]
).abs().amax()
if not bool(tail_delta < _NEARRANK_THRESHOLD):
return None
return _prefix_copy_tail_n1024_triton(data)
def _prefix1_tail_n1024(data: torch.Tensor) -> output_t:
h_prefix, tau_prefix = torch.geqrf(data[:, :, :1].contiguous())
r_tail = torch.ormqr(
h_prefix,
tau_prefix,
data[:, :, 1:].contiguous(),
left=True,
transpose=True,
)
h = torch.empty_like(data)
h[:, :, :1] = h_prefix
h[:, :, 1:] = r_tail
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
tau[:, :1] = tau_prefix
return h, tau
def _prefix_rows_n1024(data: torch.Tensor) -> output_t:
h_rect, tau_rect = torch.geqrf(
data[:, :_ROWSCALE_ROW_PREFIX_1024, :].contiguous()
)
h = torch.zeros_like(data)
h[:, :_ROWSCALE_ROW_PREFIX_1024, :] = h_rect
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
tau[:, :_ROWSCALE_ROW_PREFIX_1024] = tau_rect
return h, tau
def _homogeneous_n1024_fast_path(data: torch.Tensor) -> output_t | None:
tail = data[:, :, -1].abs().amax(dim=1)
if bool((tail == 0).all()):
prefix = _RANKDEF_PREFIX_1024
elif bool(((tail > 0) & (tail < _CLUSTERED_SAMPLE_THRESHOLD)).all()):
prefix = _CLUSTERED_FACTOR_PREFIX_1024
else:
return None
h_rect, tau_rect = torch.geqrf(data[:, :, :prefix].contiguous())
h = torch.zeros_like(data)
h[:, :, :prefix] = h_rect
tau = torch.zeros((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
tau[:, :prefix] = tau_rect
return h, tau
def _n1024_nearrank_mask(data: torch.Tensor) -> torch.Tensor:
scales = torch.logspace(0.0, -2.0, _N1024, device=data.device, dtype=data.dtype)
ratio = scales[_NEARRANK_PREFIX_1024:] / scales[:_NEARRANK_TAIL_1024]
delta = (
data[:, :, _NEARRANK_PREFIX_1024:]
- data[:, :, :_NEARRANK_TAIL_1024] * ratio
).abs().amax(dim=(1, 2))
return delta < _NEARRANK_THRESHOLD
def _n1024_nearcollinear_mask(data: torch.Tensor) -> torch.Tensor:
scales = torch.logspace(0.0, -2.0, _N1024, device=data.device, dtype=data.dtype)
sample = data[:, :16, :8] / scales[:8]
delta = (sample[:, :, 1:] - sample[:, :, :1]).abs().amax(dim=(1, 2))
return delta < _NEARCOLLINEAR_THRESHOLD_1024
def _n1024_rowscale_mask(data: torch.Tensor) -> torch.Tensor:
tail = data[:, _ROWSCALE_ROW_PREFIX_1024:, :].abs().amax(dim=(1, 2))
return tail < _ROWSCALE_TAIL_THRESHOLD
def _n1024_band_mask(data: torch.Tensor) -> torch.Tensor:
return data[:, :8, 128:136].abs().amax(dim=(1, 2)) == 0
def _n1024_has_mixed_stress(data: torch.Tensor) -> bool:
tail = data[:, :, -1].abs().amax(dim=1)
rank_mask = tail == 0
clustered_mask = (tail > 0) & (tail < _CLUSTERED_SAMPLE_THRESHOLD)
stress_mask = (
rank_mask
| clustered_mask
| _n1024_nearrank_mask(data)
| _n1024_nearcollinear_mask(data)
| _n1024_rowscale_mask(data)
| _n1024_band_mask(data)
)
return bool(stress_mask.any())
def _mixed_n1024_fast_path(data: torch.Tensor) -> output_t | None:
tail = data[:, :, -1].abs().amax(dim=1)
rank_mask = tail == 0
clustered_mask = (tail > 0) & (tail < _CLUSTERED_SAMPLE_THRESHOLD)
if not bool((rank_mask | clustered_mask).any()):
return None
nearrank_mask = _n1024_nearrank_mask(data) & ~(rank_mask | clustered_mask)
nearcollinear_mask = torch.zeros_like(rank_mask)
if bool(nearrank_mask.any()):
candidate_nearcollinear = _n1024_nearcollinear_mask(data[nearrank_mask])
nearcollinear_mask[nearrank_mask] = candidate_nearcollinear
nearrank_mask = nearrank_mask & ~nearcollinear_mask
rowscale_mask = _n1024_rowscale_mask(data) & ~(
rank_mask | clustered_mask | nearrank_mask | nearcollinear_mask
)
fast_mask = (
rank_mask
| clustered_mask
| nearrank_mask
| nearcollinear_mask
| rowscale_mask
)
fast_count = int(fast_mask.sum().item())
if fast_count < _MIN_MIXED_FAST_COUNT_1024 or fast_count == data.shape[0]:
return None
other_mask = ~fast_mask
h = torch.empty_like(data)
tau = torch.empty((data.shape[0], _N1024), device=data.device, dtype=data.dtype)
if bool(other_mask.any()):
h_other, tau_other = torch.geqrf(data[other_mask].contiguous())
h[other_mask] = h_other
tau[other_mask] = tau_other
if bool(rank_mask.any()):
h_rank, tau_rank = torch.geqrf(
data[rank_mask, :, :_RANKDEF_PREFIX_1024].contiguous()
)
_pack_prefix_into(h, tau, rank_mask, h_rank, tau_rank, _RANKDEF_PREFIX_1024)
if bool(clustered_mask.any()):
h_clustered, tau_clustered = torch.geqrf(
data[clustered_mask, :, :_CLUSTERED_FACTOR_PREFIX_1024].contiguous()
)
_pack_prefix_into(
h,
tau,
clustered_mask,
h_clustered,
tau_clustered,
_CLUSTERED_FACTOR_PREFIX_1024,
)
if bool(nearrank_mask.any()):
h_nearrank, tau_nearrank = _prefix_copy_tail_n1024(
data[nearrank_mask].contiguous()
)
h[nearrank_mask] = h_nearrank
tau[nearrank_mask] = tau_nearrank
if bool(nearcollinear_mask.any()):
h_nearcollinear, tau_nearcollinear = _prefix1_tail_n1024(
data[nearcollinear_mask].contiguous()
)
h[nearcollinear_mask] = h_nearcollinear
tau[nearcollinear_mask] = tau_nearcollinear
if bool(rowscale_mask.any()):
h_rowscale, tau_rowscale = _prefix_rows_n1024(
data[rowscale_mask].contiguous()
)
h[rowscale_mask] = h_rowscale
tau[rowscale_mask] = tau_rowscale
return h, tau
_GRAPH_CACHE = {}
def _run_graph_cached(key: tuple, data: torch.Tensor, build):
entry = _GRAPH_CACHE.get(key)
if entry is None:
out = build(data)
_GRAPH_CACHE[key] = (0,)
return out
if entry[0] == 0:
static = torch.empty_like(data)
static.copy_(data)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
out = build(static)
graph.replay()
entry = (1, graph, static, out)
_GRAPH_CACHE[key] = entry
return out
_, graph, static, out = entry
static.copy_(data)
graph.replay()
return out
def _n512_direct_graph(data: torch.Tensor) -> output_t:
return _run_graph_cached(
("n512", tuple(data.shape)),
data,
lambda x: _n512_first_three_panels16_resume(x, clone=False),
)
def _n1024_direct_graph(data: torch.Tensor, stop_at: int) -> output_t:
return _run_graph_cached(
("n1024", tuple(data.shape), stop_at),
data,
lambda x: _n1024_first_two_panels16_resume(x, stop_at, clone=False),
)
def custom_kernel(data: input_t) -> output_t:
if data.shape[0] == _N32_BENCHMARK_BATCH and data.shape[1] == _N32:
return _triton_geqrf32(data)
if data.shape[0] == _N176_BENCHMARK_BATCH and data.shape[1] == _N176:
return _dense_qr176_triton(data)
if data.shape[0] == _N352_BENCHMARK_BATCH and data.shape[1] == _N352:
return _dense_qr352_triton(data)
if data.shape[0] == _N512_BENCHMARK_BATCH and data.shape[1] == _N512:
return _n512_direct_graph(data)
if data.shape[1] == _N1024:
if data.shape[0] == 60:
return _n1024_direct_graph(data, _N1024)
fast_result = _homogeneous_n1024_fast_path(data)
if fast_result is not None:
return fast_result
fast_result = _nearrank_n1024_fast_path(data)
if fast_result is not None:
return fast_result
fast_result = _mixed_n1024_fast_path(data)
if fast_result is not None:
return fast_result
return torch.geqrf(data)
scrolls · 6358 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