submission 844731
sankalp1999 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 5211 lines, June 9 Researcher Reciprocity License v1.0.
submission_yui.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844731?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:1c9581795051ccfbb4ef040799271c522aa8ebd2571ef7e66be468187c2842e1
license declaredunknown
license concludedunknown
authorssankalp1999
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
tmp = tl.dot(t_left, cross, input_precision="tf32", out_dtype=tl.float32)num-warps = 8
num_warps=8 if rows >= 64 else 4,split-k
def _larfb16_reduce_apply_splitk_kernel(Kernel source
submission_yui.py5211 lines
import torch
from task import input_t, output_t
try:
import triton
import triton.language as tl
_HAS_TRITON = True
except Exception:
triton = None
tl = None
_HAS_TRITON = False
# Use ONLY the new matmul-precision API; mixing it with the legacy allow_tf32
# flag makes get_float32_matmul_precision() raise on B200.
# "medium" is bf16 matmul (~3.9e-3 rel err) and fails the qr_v2 per-matrix gate
# on n512 mixed batches. "high" is single-pass TF32 (~5e-4): same tensor-core
# speed, ~8x more accurate, clears the tight n512 tol with margin.
torch.set_float32_matmul_precision("high")
_BLOCKED_CASES = {
(40, 176): 16,
(40, 352): 16,
}
_GRAPHS = {}
def _graphed(name: str, fn, data: torch.Tensor) -> output_t:
if not data.is_cuda:
return fn(data)
key = (name, tuple(data.shape), data.dtype, data.device.index)
entry = _GRAPHS.get(key)
if entry is None:
static_in = data.clone()
for _ in range(3):
fn(static_in)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
static_in.copy_(data)
with torch.cuda.graph(graph):
out_h, out_tau = fn(static_in)
entry = (graph, static_in, out_h, out_tau)
_GRAPHS[key] = entry
graph, static_in, out_h, out_tau = entry
static_in.copy_(data)
graph.replay()
return out_h.clone(), out_tau.clone()
def _graphed_inplace_input(name: str, fn, data: torch.Tensor) -> output_t:
if not data.is_cuda:
return fn(data, False)
key = (name, tuple(data.shape), data.dtype, data.device.index)
entry = _GRAPHS.get(key)
if entry is None:
static_in = data.clone()
for _ in range(3):
static_in.copy_(data)
fn(static_in, True)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
static_in.copy_(data)
with torch.cuda.graph(graph):
out_h, out_tau = fn(static_in, True)
entry = (graph, static_in, out_h, out_tau)
_GRAPHS[key] = entry
graph, static_in, out_h, out_tau = entry
static_in.copy_(data)
graph.replay()
return out_h.clone(), out_tau.clone()
def _graphed_inplace_input_h(name: str, fn, data: torch.Tensor) -> output_t:
if not data.is_cuda:
return fn(data, False)
key = (name, tuple(data.shape), data.dtype, data.device.index)
entry = _GRAPHS.get(key)
if entry is None:
static_in = data.clone()
for _ in range(3):
static_in.copy_(data)
fn(static_in, True)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
static_in.copy_(data)
with torch.cuda.graph(graph):
out_h, out_tau = fn(static_in, True)
entry = (graph, static_in, out_h, out_tau)
_GRAPHS[key] = entry
graph, static_in, out_h, out_tau = entry
static_in.copy_(data)
graph.replay()
return out_h, out_tau.clone()
def _next_pow2(x: int) -> int:
return 1 << (x - 1).bit_length()
def _with_matmul_precision(precision: str, fn):
old_precision = torch.get_float32_matmul_precision()
torch.set_float32_matmul_precision(precision)
try:
return fn()
finally:
torch.set_float32_matmul_precision(old_precision)
def _tf32_hi(x: torch.Tensor) -> torch.Tensor:
bits = x.contiguous().view(torch.int32)
return torch.bitwise_and(bits, -8192).view(torch.float32)
def _bmm_3xtf32(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
def run():
a_hi = _tf32_hi(a)
b_hi = _tf32_hi(b)
a_lo = a.contiguous() - a_hi
b_lo = b.contiguous() - b_hi
out = torch.bmm(a_hi, b_hi)
out.add_(torch.bmm(a_hi, b_lo))
out.add_(torch.bmm(a_lo, b_hi))
return out
return _with_matmul_precision("high", run)
def _bmm_fp32(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
return _with_matmul_precision("highest", lambda: torch.bmm(a, b))
def _baddbmm_fp32_(target: torch.Tensor, a: torch.Tensor, b: torch.Tensor) -> None:
_with_matmul_precision(
"highest",
lambda: torch.baddbmm(target, a, b, beta=1.0, alpha=-1.0, out=target),
)
if _HAS_TRITON:
@triton.jit
def _sum_pair(a0, a1, b0, b1):
return a0 + b0, a1 + b1
@triton.jit
def _full_qr_read_write_kernel(
data,
h,
tau,
data_s0: tl.constexpr,
data_s1: tl.constexpr,
data_s2: tl.constexpr,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
n: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, n)
cols = tl.arange(0, n)
in_ptrs = data + batch_id * data_s0 + rows[:, None] * data_s1 + cols[None, :] * data_s2
out_ptrs = h + batch_id * h_s0 + rows[:, None] * h_s1 + cols[None, :] * h_s2
a = tl.load(in_ptrs).to(tl.float32)
for j in tl.static_range(0, n):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == j, col_j, 0.0), axis=0)
tail_norm2 = tl.sum(tl.where(rows > j, col_j * col_j, 0.0), axis=0)
use_reflector = tail_norm2 > 0.0
norm = tl.sqrt(alpha * alpha + tail_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where(rows > j, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(rows == j, 1.0, tl.where(rows > j, packed_col, 0.0))
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where(cols[None, :] > j, a - update, a)
tl.store(tau + batch_id * tau_s0 + j * tau_s1, tau_j)
tl.store(out_ptrs, a)
@triton.jit
def _tail_qr_inplace_kernel(
h,
tau,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k: tl.constexpr,
rows_count: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, rows_count)
cols = tl.arange(0, rows_count)
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs).to(tl.float32)
for j in tl.static_range(0, rows_count):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha_terms = tl.where(rows == j, col_j, 0.0)
norm_terms = tl.where(rows >= j, col_j * col_j, 0.0)
alpha, full_norm2 = tl.reduce(
(alpha_terms, norm_terms),
axis=0,
combine_fn=_sum_pair,
)
alpha_sq = alpha * alpha
use_reflector = full_norm2 > alpha_sq
norm = tl.sqrt(full_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where(rows > j, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(rows == j, 1.0, tl.where(rows > j, packed_col, 0.0))
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where(cols[None, :] > j, a - update, a)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a)
@triton.jit
def _panel16_qr_kernel(
h,
tau,
v_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
k,
rows_count,
block_m: tl.constexpr,
emit_v: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == j, col_j, 0.0), axis=0)
tail_norm2 = tl.sum(
tl.where((rows > j) & row_mask, col_j * col_j, 0.0),
axis=0,
)
use_reflector = tail_norm2 > 0.0
norm = tl.sqrt(alpha * alpha + tail_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where(
(cols[None, :] > j) & row_mask[:, None],
a - update,
a,
)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a, mask=row_mask[:, None])
if emit_v:
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
@triton.jit
def _panel16_qr_fixed_rows_kernel(
h,
tau,
v_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
k,
rows_count: tl.constexpr,
block_m: tl.constexpr,
emit_v: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == j, col_j, 0.0), axis=0)
tail_norm2 = tl.sum(
tl.where((rows > j) & row_mask, col_j * col_j, 0.0),
axis=0,
)
use_reflector = tail_norm2 > 0.0
norm = tl.sqrt(alpha * alpha + tail_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where(
(cols[None, :] > j) & row_mask[:, None],
a - update,
a,
)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a, mask=row_mask[:, None])
if emit_v:
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
@triton.jit
def _panel16_qr_fixed_rows_paired_norm_kernel(
h,
tau,
v_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
k,
rows_count: tl.constexpr,
block_m: tl.constexpr,
emit_v: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha_terms = tl.where(rows == j, col_j, 0.0)
norm_terms = tl.where((rows >= j) & row_mask, col_j * col_j, 0.0)
alpha, full_norm2 = tl.reduce(
(alpha_terms, norm_terms),
axis=0,
combine_fn=_sum_pair,
)
alpha_sq = alpha * alpha
use_reflector = full_norm2 > alpha_sq
norm = tl.sqrt(full_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where(
(cols[None, :] > j) & row_mask[:, None],
a - update,
a,
)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a, mask=row_mask[:, None])
if emit_v:
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
@triton.jit
def _panel16_qr_update_next16_fixed_rows_paired_norm_kernel(
h,
tau,
v_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
k,
rows_count: tl.constexpr,
block_m: tl.constexpr,
emit_v: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
panel_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
next_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(panel_ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
nxt = tl.load(next_ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha_terms = tl.where(rows == j, col_j, 0.0)
norm_terms = tl.where((rows >= j) & row_mask, col_j * col_j, 0.0)
alpha, full_norm2 = tl.reduce(
(alpha_terms, norm_terms),
axis=0,
combine_fn=_sum_pair,
)
alpha_sq = alpha * alpha
use_reflector = full_norm2 > alpha_sq
norm = tl.sqrt(full_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
panel_dots = tl.sum(v[:, None] * a, axis=0)
panel_update = tau_j * v[:, None] * panel_dots[None, :]
a = tl.where(
(cols[None, :] > j) & row_mask[:, None],
a - panel_update,
a,
)
next_dots = tl.sum(v[:, None] * nxt, axis=0)
next_update = tau_j * v[:, None] * next_dots[None, :]
nxt = tl.where(row_mask[:, None], nxt - next_update, nxt)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(panel_ptrs, a, mask=row_mask[:, None])
tl.store(next_ptrs, nxt, mask=row_mask[:, None])
if emit_v:
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
@triton.jit
def _larft16_kernel(
gram,
tau,
t,
gram_s0: tl.constexpr,
gram_s1: tl.constexpr,
gram_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
width: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, width)
rr = idx[:, None]
cc = idx[None, :]
g = tl.load(
gram + batch_id * gram_s0 + rr * gram_s1 + cc * gram_s2
).to(tl.float32)
t_mat = tl.zeros((width, width), tl.float32)
for j in tl.static_range(0, width):
tau_j = tl.load(tau + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
gram_col = tl.sum(tl.where(cc == j, g, 0.0), axis=1)
col = tl.where(idx < j, -tau_j * gram_col, 0.0)
new_col = tl.sum(t_mat * col[None, :], axis=1)
t_mat = tl.where((cc == j) & (rr < j), new_col[:, None], t_mat)
t_mat = tl.where((rr == j) & (cc == j), tau_j, t_mat)
tl.store(t + batch_id * t_s0 + rr * t_s1 + cc * t_s2, t_mat)
@triton.jit
def _larft64_blocked2_from_gram_kernel(
gram,
tau,
t,
gram_s0: tl.constexpr,
gram_s1: tl.constexpr,
gram_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 32)
rr = idx[:, None]
cc = idx[None, :]
g_left = tl.load(
gram + batch_id * gram_s0 + rr * gram_s1 + cc * gram_s2
).to(tl.float32)
t_left = tl.zeros((32, 32), tl.float32)
for j in tl.static_range(0, 32):
tau_j = tl.load(tau + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
gram_col = tl.sum(tl.where(cc == j, g_left, 0.0), axis=1)
col = tl.where(idx < j, -tau_j * gram_col, 0.0)
new_col = tl.sum(t_left * col[None, :], axis=1)
t_left = tl.where((cc == j) & (rr < j), new_col[:, None], t_left)
t_left = tl.where((rr == j) & (cc == j), tau_j, t_left)
g_right = tl.load(
gram + batch_id * gram_s0 + (rr + 32) * gram_s1 + (cc + 32) * gram_s2
).to(tl.float32)
t_right = tl.zeros((32, 32), tl.float32)
for j in tl.static_range(0, 32):
tau_j = tl.load(tau + batch_id * tau_s0 + (j + 32) * tau_s1).to(tl.float32)
gram_col = tl.sum(tl.where(cc == j, g_right, 0.0), axis=1)
col = tl.where(idx < j, -tau_j * gram_col, 0.0)
new_col = tl.sum(t_right * col[None, :], axis=1)
t_right = tl.where((cc == j) & (rr < j), new_col[:, None], t_right)
t_right = tl.where((rr == j) & (cc == j), tau_j, t_right)
cross = tl.load(
gram + batch_id * gram_s0 + rr * gram_s1 + (cc + 32) * gram_s2
).to(tl.float32)
tmp = tl.dot(t_left, cross, input_precision="tf32", out_dtype=tl.float32)
top_right = -tl.dot(tmp, t_right, input_precision="tf32", out_dtype=tl.float32)
base = t + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, t_left)
tl.store(base + rr * t_s1 + (cc + 32) * t_s2, top_right)
tl.store(base + (rr + 32) * t_s1 + cc * t_s2, tl.zeros((32, 32), tl.float32))
tl.store(base + (rr + 32) * t_s1 + (cc + 32) * t_s2, t_right)
@triton.jit
def _chol64_upper_noinfo_kernel(
gram,
r,
gram_s0: tl.constexpr,
gram_s1: tl.constexpr,
gram_s2: tl.constexpr,
r_s0: tl.constexpr,
r_s1: tl.constexpr,
r_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
a = tl.load(
gram + batch_id * gram_s0 + rr * gram_s1 + cc * gram_s2
).to(tl.float32)
for j in tl.static_range(0, 64):
col_j = tl.sum(tl.where(cc == j, a, 0.0), axis=1)
pivot = tl.sum(tl.where(idx == j, col_j, 0.0), axis=0)
diag = tl.sqrt(pivot)
l_col = col_j / diag
update = l_col[:, None] * l_col[None, :]
a = tl.where((rr > j) & (cc > j), a - update, a)
a = tl.where((rr >= j) & (cc == j), l_col[:, None], a)
r_vals = tl.where(cc >= rr, tl.trans(a), 0.0)
tl.store(r + batch_id * r_s0 + rr * r_s1 + cc * r_s2, r_vals)
@triton.jit
def _panel32_qr_tail_gram_kernel(
h,
tau,
v_out,
gram_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
gram_s0: tl.constexpr,
gram_s1: tl.constexpr,
gram_s2: tl.constexpr,
k,
rows_count,
block_m: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 32)
row_mask = rows < rows_count
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 32):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == j, col_j, 0.0), axis=0)
tail_norm2 = tl.sum(
tl.where((rows > j) & row_mask, col_j * col_j, 0.0),
axis=0,
)
use_reflector = tail_norm2 > 0.0
norm = tl.sqrt(alpha * alpha + tail_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where((cols[None, :] > j) & row_mask[:, None], a - update, a)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a, mask=row_mask[:, None])
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
gram = tl.dot(tl.trans(dense_v), dense_v, input_precision="tf32", out_dtype=tl.float32)
rr = cols[:, None]
cc = cols[None, :]
tl.store(gram_out + batch_id * gram_s0 + rr * gram_s1 + cc * gram_s2, gram)
@triton.jit
def _panel32_qr_tail_gram_fixed_rows_kernel(
h,
tau,
v_out,
gram_out,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
gram_s0: tl.constexpr,
gram_s1: tl.constexpr,
gram_s2: tl.constexpr,
k,
rows_count: tl.constexpr,
block_m: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 32)
row_mask = rows < rows_count
ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + cols)[None, :] * h_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 32):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha_terms = tl.where(rows == j, col_j, 0.0)
norm_terms = tl.where((rows >= j) & row_mask, col_j * col_j, 0.0)
alpha, full_norm2 = tl.reduce(
(alpha_terms, norm_terms),
axis=0,
combine_fn=_sum_pair,
)
alpha_sq = alpha * alpha
use_reflector = full_norm2 > alpha_sq
norm = tl.sqrt(full_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where((cols[None, :] > j) & row_mask[:, None], a - update, a)
tl.store(tau + batch_id * tau_s0 + (k + j) * tau_s1, tau_j)
tl.store(ptrs, a, mask=row_mask[:, None])
v_ptrs = v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2
dense_v = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where((rows[:, None] > cols[None, :]) & row_mask[:, None], a, 0.0),
)
tl.store(v_ptrs, dense_v, mask=row_mask[:, None])
gram = tl.dot(tl.trans(dense_v), dense_v, input_precision="tf32", out_dtype=tl.float32)
rr = cols[:, None]
cc = cols[None, :]
tl.store(gram_out + batch_id * gram_s0 + rr * gram_s1 + cc * gram_s2, gram)
@triton.jit
def _lu64_no_pivot_kernel(
m,
lu,
m_s0: tl.constexpr,
m_s1: tl.constexpr,
m_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
a = tl.load(m + batch_id * m_s0 + rr * m_s1 + cc * m_s2).to(tl.float32)
for j in tl.static_range(0, 64):
row_j = tl.sum(tl.where(rr == j, a, 0.0), axis=0)
col_j = tl.sum(tl.where(cc == j, a, 0.0), axis=1)
pivot = tl.sum(tl.where(idx == j, row_j, 0.0), axis=0)
mult = col_j / pivot
update = mult[:, None] * row_j[None, :]
a = tl.where((rr > j) & (cc == j), mult[:, None], a)
a = tl.where((rr > j) & (cc > j), a - update, a)
tl.store(lu + batch_id * lu_s0 + rr * lu_s1 + cc * lu_s2, a)
@triton.jit
def _orhr64_pack_panel_v_kernel(
panel,
m,
lu_top,
r,
signs,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
m_s0: tl.constexpr,
m_s1: tl.constexpr,
m_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
r_s0: tl.constexpr,
r_s1: tl.constexpr,
r_s2: tl.constexpr,
signs_s0: tl.constexpr,
signs_s1: tl.constexpr,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
row_mask = rows < rows_count
m_vals = tl.load(
m + batch_id * m_s0 + rows[:, None] * m_s1 + cols[None, :] * m_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
lu_vals = tl.load(
lu_top + batch_id * lu_s0 + rows[:, None] * lu_s1 + cols[None, :] * lu_s2,
mask=(rows[:, None] < 64) & row_mask[:, None],
other=0.0,
).to(tl.float32)
r_vals = tl.load(
r + batch_id * r_s0 + rows[:, None] * r_s1 + cols[None, :] * r_s2,
mask=(rows[:, None] < 64) & row_mask[:, None],
other=0.0,
).to(tl.float32)
sign_vals = tl.load(
signs + batch_id * signs_s0 + rows * signs_s1,
mask=(rows < 64) & row_mask,
other=0.0,
).to(tl.float32)
top = rows[:, None] < 64
lower = rows[:, None] > cols[None, :]
upper = cols[None, :] >= rows[:, None]
top_packed = tl.where(lower, lu_vals, 0.0) + tl.where(upper, sign_vals[:, None] * r_vals, 0.0)
packed = tl.where(top, top_packed, m_vals)
v_vals = tl.where(
rows[:, None] == cols[None, :],
1.0,
tl.where(lower, packed, 0.0),
)
tl.store(
panel + batch_id * panel_s0 + rows[:, None] * panel_s1 + cols[None, :] * panel_s2,
packed,
mask=row_mask[:, None],
)
tl.store(
m + batch_id * m_s0 + rows[:, None] * m_s1 + cols[None, :] * m_s2,
v_vals,
mask=row_mask[:, None],
)
@triton.jit
def _orhr64_top_lu_pack_v_kernel(
panel,
q,
r,
lu_top,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
q_s0,
q_s1: tl.constexpr,
q_s2: tl.constexpr,
r_s0: tl.constexpr,
r_s1: tl.constexpr,
r_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
q_vals = tl.load(
q + batch_id * q_s0 + rr * q_s1 + cc * q_s2,
).to(tl.float32)
q_diag = tl.load(
q + batch_id * q_s0 + idx * q_s1 + idx * q_s2,
).to(tl.float32)
signs = tl.where(q_diag > 0.0, -1.0, 1.0)
a = tl.where(rr == cc, q_vals - signs[:, None], q_vals)
for j in tl.static_range(0, 64):
row_j = tl.sum(tl.where(rr == j, a, 0.0), axis=0)
col_j = tl.sum(tl.where(cc == j, a, 0.0), axis=1)
pivot = tl.sum(tl.where(idx == j, row_j, 0.0), axis=0)
mult = col_j / pivot
update = mult[:, None] * row_j[None, :]
a = tl.where((rr > j) & (cc == j), mult[:, None], a)
a = tl.where((rr > j) & (cc > j), a - update, a)
tl.store(
lu_top + batch_id * lu_s0 + rr * lu_s1 + cc * lu_s2,
a,
)
r_vals = tl.load(
r + batch_id * r_s0 + rr * r_s1 + cc * r_s2,
).to(tl.float32)
lower = rr > cc
upper = cc >= rr
packed = tl.where(lower, a, 0.0) + tl.where(upper, signs[:, None] * r_vals, 0.0)
v_vals = tl.where(rr == cc, 1.0, tl.where(lower, packed, 0.0))
tl.store(
panel + batch_id * panel_s0 + rr * panel_s1 + cc * panel_s2,
packed,
)
tl.store(
q + batch_id * q_s0 + rr * q_s1 + cc * q_s2,
v_vals,
)
@triton.jit
def _orhr64_tail_solve_pack_v_kernel(
panel,
q,
lu_top,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
q_s0,
q_s1: tl.constexpr,
q_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = 64 + row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
row_mask = rows < rows_count
x = tl.load(
q + batch_id * q_s0 + rows[:, None] * q_s1 + cols[None, :] * q_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 64):
u_col = tl.load(
lu_top + batch_id * lu_s0 + cols * lu_s1 + j * lu_s2,
).to(tl.float32)
pivot = tl.load(
lu_top + batch_id * lu_s0 + j * lu_s1 + j * lu_s2,
).to(tl.float32)
acc = tl.sum(
tl.where(cols[None, :] == j, x, 0.0),
axis=1,
)
prev = tl.sum(
tl.where(cols[None, :] < j, x * u_col[None, :], 0.0),
axis=1,
)
solved = (acc - prev) / pivot
x = tl.where(cols[None, :] == j, solved[:, None], x)
tl.store(
panel + batch_id * panel_s0 + rows[:, None] * panel_s1 + cols[None, :] * panel_s2,
x,
mask=row_mask[:, None],
)
tl.store(
q + batch_id * q_s0 + rows[:, None] * q_s1 + cols[None, :] * q_s2,
x,
mask=row_mask[:, None],
)
@triton.jit
def _orhr64_top_lu_pack_v_cols_qt_kernel(
panel,
q_t,
v_out,
r,
lu_cols,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
qt_s0,
qt_s1,
qt_s2,
v_s0,
v_s1: tl.constexpr,
v_s2: tl.constexpr,
r_s0: tl.constexpr,
r_s1: tl.constexpr,
r_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
q_vals = tl.load(
q_t + batch_id * qt_s0 + cc * qt_s1 + rr * qt_s2,
).to(tl.float32)
q_diag = tl.load(
q_t + batch_id * qt_s0 + idx * qt_s1 + idx * qt_s2,
).to(tl.float32)
signs = tl.where(q_diag > 0.0, -1.0, 1.0)
a = tl.where(rr == cc, q_vals - signs[:, None], q_vals)
for j in tl.static_range(0, 64):
row_j = tl.sum(tl.where(rr == j, a, 0.0), axis=0)
col_j = tl.sum(tl.where(cc == j, a, 0.0), axis=1)
pivot = tl.sum(tl.where(idx == j, row_j, 0.0), axis=0)
mult = col_j / pivot
update = mult[:, None] * row_j[None, :]
a = tl.where((rr > j) & (cc == j), mult[:, None], a)
a = tl.where((rr > j) & (cc > j), a - update, a)
tl.store(
lu_cols + batch_id * lu_s0 + cc * lu_s1 + rr * lu_s2,
a,
)
r_vals = tl.load(
r + batch_id * r_s0 + rr * r_s1 + cc * r_s2,
).to(tl.float32)
lower = rr > cc
upper = cc >= rr
packed = tl.where(lower, a, 0.0) + tl.where(upper, signs[:, None] * r_vals, 0.0)
v_vals = tl.where(rr == cc, 1.0, tl.where(lower, packed, 0.0))
tl.store(
panel + batch_id * panel_s0 + rr * panel_s1 + cc * panel_s2,
packed,
)
tl.store(
v_out + batch_id * v_s0 + rr * v_s1 + cc * v_s2,
v_vals,
)
@triton.jit
def _orhr64_tail_solve_pack_v_cols_qt_kernel(
panel,
q_t,
v_out,
lu_cols,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
qt_s0,
qt_s1,
qt_s2,
v_s0,
v_s1: tl.constexpr,
v_s2: tl.constexpr,
lu_s0: tl.constexpr,
lu_s1: tl.constexpr,
lu_s2: tl.constexpr,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = 64 + row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
row_mask = rows < rows_count
x = tl.load(
q_t + batch_id * qt_s0 + cols[None, :] * qt_s1 + rows[:, None] * qt_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 64):
u_col = tl.load(
lu_cols + batch_id * lu_s0 + j * lu_s1 + cols * lu_s2,
).to(tl.float32)
pivot = tl.load(
lu_cols + batch_id * lu_s0 + j * lu_s1 + j * lu_s2,
).to(tl.float32)
acc = tl.sum(
tl.where(cols[None, :] == j, x, 0.0),
axis=1,
)
prev = tl.sum(
tl.where(cols[None, :] < j, x * u_col[None, :], 0.0),
axis=1,
)
solved = (acc - prev) / pivot
x = tl.where(cols[None, :] == j, solved[:, None], x)
tl.store(
panel + batch_id * panel_s0 + rows[:, None] * panel_s1 + cols[None, :] * panel_s2,
x,
mask=row_mask[:, None],
)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
x,
mask=row_mask[:, None],
)
@triton.jit
def _larfb16_update_kernel(
h,
v,
t,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
k,
rows_count,
cols_count,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
panel_cols = tl.arange(0, 16)
row_mask = rows < rows_count
col_mask = cols < cols_count
v_tile = tl.load(
v
+ batch_id * v_s0
+ rows[:, None] * v_s1
+ panel_cols[None, :] * v_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a_tile = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
t_tile = tl.load(
t
+ batch_id * t_s0
+ panel_cols[:, None] * t_s1
+ panel_cols[None, :] * t_s2
).to(tl.float32)
w = tl.dot(tl.trans(v_tile), a_tile, input_precision="ieee", out_dtype=tl.float32)
w = tl.dot(tl.trans(t_tile), w, input_precision="ieee", out_dtype=tl.float32)
update = tl.dot(v_tile, w, input_precision="ieee", out_dtype=tl.float32)
tl.store(a_ptrs, a_tile - update, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _larfb16_update_x3_kernel(
h,
v,
t,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
k,
rows_count,
cols_count,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
panel_cols = tl.arange(0, 16)
row_mask = rows < rows_count
col_mask = cols < cols_count
v_tile = tl.load(
v
+ batch_id * v_s0
+ rows[:, None] * v_s1
+ panel_cols[None, :] * v_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a_tile = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
t_tile = tl.load(
t
+ batch_id * t_s0
+ panel_cols[:, None] * t_s1
+ panel_cols[None, :] * t_s2
).to(tl.float32)
w = tl.dot(tl.trans(v_tile), a_tile, input_precision="tf32x3", out_dtype=tl.float32)
w = tl.dot(tl.trans(t_tile), w, input_precision="ieee", out_dtype=tl.float32)
update = tl.dot(v_tile, w, input_precision="tf32x3", out_dtype=tl.float32)
tl.store(a_ptrs, a_tile - update, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _split16_local_direct_forward_kernel(
h,
v,
tau_panel,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k,
rows_count,
block_m: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(a_ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
v_col = tl.load(
v + batch_id * v_s0 + rows * v_s1 + j * v_s2,
mask=row_mask,
other=0.0,
).to(tl.float32)
tau_j = tl.load(tau_panel + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
dots = tl.sum(v_col[:, None] * a, axis=0)
a = tl.where(row_mask[:, None], a - tau_j * v_col[:, None] * dots[None, :], a)
tl.store(a_ptrs, a, mask=row_mask[:, None])
@triton.jit
def _split16_trailing_direct_forward_kernel(
h,
v,
tau_panel,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k,
rows_count,
cols_count,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
row_mask = rows < rows_count
col_mask = cols < cols_count
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
v_col = tl.load(
v + batch_id * v_s0 + rows * v_s1 + j * v_s2,
mask=row_mask,
other=0.0,
).to(tl.float32)
tau_j = tl.load(tau_panel + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
dots = tl.sum(v_col[:, None] * a, axis=0)
a = tl.where(row_mask[:, None], a - tau_j * v_col[:, None] * dots[None, :], a)
tl.store(a_ptrs, a, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _split16_local_direct_forward_fixed_rows_kernel(
h,
v,
tau_panel,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k,
rows_count: tl.constexpr,
block_m: tl.constexpr,
):
batch_id = tl.program_id(0)
rows = tl.arange(0, block_m)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(a_ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
v_col = tl.load(
v + batch_id * v_s0 + rows * v_s1 + j * v_s2,
mask=row_mask,
other=0.0,
).to(tl.float32)
tau_j = tl.load(tau_panel + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
dots = tl.sum(v_col[:, None] * a, axis=0)
a = tl.where(row_mask[:, None], a - tau_j * v_col[:, None] * dots[None, :], a)
tl.store(a_ptrs, a, mask=row_mask[:, None])
@triton.jit
def _split16_trailing_direct_forward_fixed_rows_kernel(
h,
v,
tau_panel,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k,
rows_count: tl.constexpr,
cols_count,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
row_mask = rows < rows_count
col_mask = cols < cols_count
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
v_col = tl.load(
v + batch_id * v_s0 + rows * v_s1 + j * v_s2,
mask=row_mask,
other=0.0,
).to(tl.float32)
tau_j = tl.load(tau_panel + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
dots = tl.sum(v_col[:, None] * a, axis=0)
a = tl.where(row_mask[:, None], a - tau_j * v_col[:, None] * dots[None, :], a)
tl.store(a_ptrs, a, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _split16_trailing_direct_forward_fixed_k_rows_kernel(
h,
v,
tau_panel,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
k: tl.constexpr,
rows_count: tl.constexpr,
cols_count,
block_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
rows = tl.arange(0, block_m)
cols = col_block * block_n + tl.arange(0, block_n)
row_mask = rows < rows_count
col_mask = cols < cols_count
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
for j in tl.static_range(0, 16):
v_col = tl.load(
v + batch_id * v_s0 + rows * v_s1 + j * v_s2,
mask=row_mask,
other=0.0,
).to(tl.float32)
tau_j = tl.load(tau_panel + batch_id * tau_s0 + j * tau_s1).to(tl.float32)
dots = tl.sum(v_col[:, None] * a, axis=0)
a = tl.where(row_mask[:, None], a - tau_j * v_col[:, None] * dots[None, :], a)
tl.store(a_ptrs, a, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _larfb16_partial_w_kernel(
h,
v,
partial,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
p_s0,
p_s1,
p_s2,
p_s3,
p_s4,
k,
rows_count,
cols_count,
chunk_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
row_part = tl.program_id(2)
rows = row_part * chunk_m + tl.arange(0, chunk_m)
cols = col_block * block_n + tl.arange(0, block_n)
panel_cols = tl.arange(0, 16)
row_mask = rows < rows_count
col_mask = cols < cols_count
v_tile = tl.load(
v
+ batch_id * v_s0
+ rows[:, None] * v_s1
+ panel_cols[None, :] * v_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
a_tile = tl.load(
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
w = tl.dot(tl.trans(v_tile), a_tile, input_precision="ieee", out_dtype=tl.float32)
tl.store(
partial
+ batch_id * p_s0
+ col_block * p_s1
+ row_part * p_s2
+ panel_cols[:, None] * p_s3
+ tl.arange(0, block_n)[None, :] * p_s4,
w,
)
@triton.jit
def _larfb16_reduce_apply_splitk_kernel(
h,
v,
t,
partial,
h_s0: tl.constexpr,
h_s1: tl.constexpr,
h_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
p_s0,
p_s1,
p_s2,
p_s3,
p_s4,
k,
rows_count,
cols_count,
row_parts: tl.constexpr,
chunk_m: tl.constexpr,
block_n: tl.constexpr,
):
batch_id = tl.program_id(0)
col_block = tl.program_id(1)
row_part = tl.program_id(2)
rows = row_part * chunk_m + tl.arange(0, chunk_m)
cols = col_block * block_n + tl.arange(0, block_n)
panel_cols = tl.arange(0, 16)
row_mask = rows < rows_count
col_mask = cols < cols_count
w = tl.zeros((16, block_n), dtype=tl.float32)
for part in tl.static_range(0, row_parts):
w += tl.load(
partial
+ batch_id * p_s0
+ col_block * p_s1
+ part * p_s2
+ panel_cols[:, None] * p_s3
+ tl.arange(0, block_n)[None, :] * p_s4
).to(tl.float32)
t_tile = tl.load(
t
+ batch_id * t_s0
+ panel_cols[:, None] * t_s1
+ panel_cols[None, :] * t_s2
).to(tl.float32)
w = tl.dot(tl.trans(t_tile), w, input_precision="ieee", out_dtype=tl.float32)
v_tile = tl.load(
v
+ batch_id * v_s0
+ rows[:, None] * v_s1
+ panel_cols[None, :] * v_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
a_ptrs = (
h
+ batch_id * h_s0
+ (k + rows)[:, None] * h_s1
+ (k + 16 + cols)[None, :] * h_s2
)
a_tile = tl.load(
a_ptrs,
mask=row_mask[:, None] & col_mask[None, :],
other=0.0,
).to(tl.float32)
update = tl.dot(v_tile, w, input_precision="ieee", out_dtype=tl.float32)
tl.store(a_ptrs, a_tile - update, mask=row_mask[:, None] & col_mask[None, :])
@triton.jit
def _assemble_v64_from32_kernel(
v1,
v2,
v_out,
v1_s0,
v1_s1,
v1_s2,
v2_s0,
v2_s1,
v2_s2,
v_s0,
v_s1,
v_s2,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
first = cols < 32
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v1 + batch_id * v1_s0 + rows[:, None] * v1_s1 + cols[None, :] * v1_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
v2_rows = rows - 32
v2_cols = cols - 32
val2 = tl.load(
v2 + batch_id * v2_s0 + v2_rows[:, None] * v2_s1 + v2_cols[None, :] * v2_s2,
mask=(rows[:, None] >= 32) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_v64_from32_fixed_rows_kernel(
v1,
v2,
v_out,
v1_s0,
v1_s1,
v1_s2,
v2_s0,
v2_s1,
v2_s2,
v_s0,
v_s1,
v_s2,
rows_count: tl.constexpr,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
first = cols < 32
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v1 + batch_id * v1_s0 + rows[:, None] * v1_s1 + cols[None, :] * v1_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
v2_rows = rows - 32
v2_cols = cols - 32
val2 = tl.load(
v2 + batch_id * v2_s0 + v2_rows[:, None] * v2_s1 + v2_cols[None, :] * v2_s2,
mask=(rows[:, None] >= 32) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_vtau32_from16_kernel(
v1,
v2,
tau1,
tau2,
v_out,
tau_out,
v1_s0,
v1_s1,
v1_s2,
v2_s0,
v2_s1,
v2_s2,
tau1_s0: tl.constexpr,
tau1_s1: tl.constexpr,
tau2_s0: tl.constexpr,
tau2_s1: tl.constexpr,
v_s0,
v_s1,
v_s2,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 32)
first = cols < 16
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v1 + batch_id * v1_s0 + rows[:, None] * v1_s1 + cols[None, :] * v1_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
v2_rows = rows - 16
v2_cols = cols - 16
val2 = tl.load(
v2 + batch_id * v2_s0 + v2_rows[:, None] * v2_s1 + v2_cols[None, :] * v2_s2,
mask=(rows[:, None] >= 16) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
tau_cols = tl.arange(0, 32)
tau_first = tau_cols < 16
tau_val1 = tl.load(
tau1 + batch_id * tau1_s0 + tau_cols * tau1_s1,
mask=tau_first,
other=0.0,
)
tau_val2 = tl.load(
tau2 + batch_id * tau2_s0 + (tau_cols - 16) * tau2_s1,
mask=~tau_first,
other=0.0,
)
tau_val = tl.where(tau_first, tau_val1, tau_val2)
tl.store(tau_out + batch_id * tau_s0 + tau_cols * tau_s1, tau_val)
@triton.jit
def _assemble_v128_from64_kernel(
v_left,
v_right,
v_out,
vl_s0,
vl_s1,
vl_s2,
vr_s0,
vr_s1,
vr_s2,
v_s0,
v_s1,
v_s2,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 128)
first = cols < 64
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v_left + batch_id * vl_s0 + rows[:, None] * vl_s1 + cols[None, :] * vl_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
right_rows = rows - 64
right_cols = cols - 64
val2 = tl.load(
v_right + batch_id * vr_s0 + right_rows[:, None] * vr_s1 + right_cols[None, :] * vr_s2,
mask=(rows[:, None] >= 64) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_t128_from64_kernel(
t_left,
t_right,
cross,
t_out,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
cross_s0: tl.constexpr,
cross_s1: tl.constexpr,
cross_s2: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
top_left = tl.load(t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2)
top_right = tl.load(cross + batch_id * cross_s0 + rr * cross_s1 + cc * cross_s2)
bottom_right = tl.load(t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, top_left)
tl.store(base + rr * t_s1 + (cc + 64) * t_s2, top_right)
tl.store(base + (rr + 64) * t_s1 + cc * t_s2, tl.zeros((64, 64), tl.float32))
tl.store(base + (rr + 64) * t_s1 + (cc + 64) * t_s2, bottom_right)
@triton.jit
def _assemble_v256_from128_kernel(
v_left,
v_right,
v_out,
vl_s0,
vl_s1,
vl_s2,
vr_s0,
vr_s1,
vr_s2,
v_s0,
v_s1,
v_s2,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 256)
first = cols < 128
second = ~first
row_mask = rows < rows_count
left_vals = tl.load(
v_left + batch_id * vl_s0 + rows[:, None] * vl_s1 + cols[None, :] * vl_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
right_rows = rows - 128
right_cols = cols - 128
right_vals = tl.load(
v_right + batch_id * vr_s0 + right_rows[:, None] * vr_s1 + right_cols[None, :] * vr_s2,
mask=(rows[:, None] >= 128) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], left_vals, right_vals)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_t256_from128_kernel(
t_left,
t_right,
cross,
t_out,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
cross_s0: tl.constexpr,
cross_s1: tl.constexpr,
cross_s2: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
block_c: tl.constexpr,
):
batch_id = tl.program_id(0)
task = tl.program_id(1)
rows = tl.arange(0, 128)
cols = tl.arange(0, block_c)
rr = rows[:, None]
cc = task * block_c + cols[None, :]
col_mask = cc < 128
left = tl.load(
t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2,
mask=col_mask,
other=0.0,
)
top_right = tl.load(
cross + batch_id * cross_s0 + rr * cross_s1 + cc * cross_s2,
mask=col_mask,
other=0.0,
)
bottom_right = tl.load(
t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2,
mask=col_mask,
other=0.0,
)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left, mask=col_mask)
tl.store(base + rr * t_s1 + (cc + 128) * t_s2, top_right, mask=col_mask)
tl.store(base + (rr + 128) * t_s1 + cc * t_s2, tl.zeros((128, block_c), tl.float32), mask=col_mask)
tl.store(base + (rr + 128) * t_s1 + (cc + 128) * t_s2, bottom_right, mask=col_mask)
@triton.jit
def _assemble_v512_from256_kernel(
v_left,
v_right,
v_out,
vl_s0,
vl_s1,
vl_s2,
vr_s0,
vr_s1,
vr_s2,
v_s0,
v_s1,
v_s2,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 512)
first = cols < 256
second = ~first
row_mask = rows < rows_count
left_vals = tl.load(
v_left + batch_id * vl_s0 + rows[:, None] * vl_s1 + cols[None, :] * vl_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
right_rows = rows - 256
right_cols = cols - 256
right_vals = tl.load(
v_right + batch_id * vr_s0 + right_rows[:, None] * vr_s1 + right_cols[None, :] * vr_s2,
mask=(rows[:, None] >= 256) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], left_vals, right_vals)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_t512_from256_kernel(
t_left,
t_right,
cross,
t_out,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
cross_s0: tl.constexpr,
cross_s1: tl.constexpr,
cross_s2: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
block_c: tl.constexpr,
):
batch_id = tl.program_id(0)
task = tl.program_id(1)
rows = tl.arange(0, 256)
cols = tl.arange(0, block_c)
rr = rows[:, None]
cc = task * block_c + cols[None, :]
col_mask = cc < 256
left = tl.load(
t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2,
mask=col_mask,
other=0.0,
)
top_right = tl.load(
cross + batch_id * cross_s0 + rr * cross_s1 + cc * cross_s2,
mask=col_mask,
other=0.0,
)
bottom_right = tl.load(
t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2,
mask=col_mask,
other=0.0,
)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left, mask=col_mask)
tl.store(base + rr * t_s1 + (cc + 256) * t_s2, top_right, mask=col_mask)
tl.store(base + (rr + 256) * t_s1 + cc * t_s2, tl.zeros((256, block_c), tl.float32), mask=col_mask)
tl.store(base + (rr + 256) * t_s1 + (cc + 256) * t_s2, bottom_right, mask=col_mask)
@triton.jit
def _compose_t128_from_cross64_kernel(
t_left,
t_right,
cross0,
t_out,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
c_s0: tl.constexpr,
c_s1: tl.constexpr,
c_s2: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 64)
rr = idx[:, None]
cc = idx[None, :]
left = tl.load(t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2)
right = tl.load(t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2)
raw_cross = tl.load(cross0 + batch_id * c_s0 + rr * c_s1 + cc * c_s2)
tmp = tl.dot(left, raw_cross, input_precision="tf32x3", out_dtype=tl.float32)
top_right = -tl.dot(tmp, right, input_precision="tf32x3", out_dtype=tl.float32)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left)
tl.store(base + rr * t_s1 + (cc + 64) * t_s2, top_right)
tl.store(base + (rr + 64) * t_s1 + cc * t_s2, tl.zeros((64, 64), tl.float32))
tl.store(base + (rr + 64) * t_s1 + (cc + 64) * t_s2, right)
@triton.jit
def _compose_t64_from_cross32_kernel(
t_left,
t_right,
cross0,
t_out,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
c_s0: tl.constexpr,
c_s1: tl.constexpr,
c_s2: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 32)
rr = idx[:, None]
cc = idx[None, :]
left = tl.load(t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2)
right = tl.load(t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2)
raw_cross = tl.load(cross0 + batch_id * c_s0 + rr * c_s1 + cc * c_s2)
tmp = tl.dot(left, raw_cross, input_precision="tf32x3", out_dtype=tl.float32)
top_right = -tl.dot(tmp, right, input_precision="tf32x3", out_dtype=tl.float32)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left)
tl.store(base + rr * t_s1 + (cc + 32) * t_s2, top_right)
tl.store(base + (rr + 32) * t_s1 + cc * t_s2, tl.zeros((32, 32), tl.float32))
tl.store(base + (rr + 32) * t_s1 + (cc + 32) * t_s2, right)
@triton.jit
def _assemble_v64_compose_t64_from32_kernel(
v1,
v2,
t_left,
t_right,
cross0,
v_out,
t_out,
v1_s0,
v1_s1,
v1_s2,
v2_s0,
v2_s1,
v2_s2,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
c_s0: tl.constexpr,
c_s1: tl.constexpr,
c_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
t_s0,
t_s1,
t_s2,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
task = tl.program_id(1)
if task == 0:
idx = tl.arange(0, 32)
rr = idx[:, None]
cc = idx[None, :]
left = tl.load(t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2)
right = tl.load(t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2)
raw_cross = tl.load(cross0 + batch_id * c_s0 + rr * c_s1 + cc * c_s2)
tmp = tl.dot(left, raw_cross, input_precision="tf32x3", out_dtype=tl.float32)
top_right = -tl.dot(tmp, right, input_precision="tf32x3", out_dtype=tl.float32)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left)
tl.store(base + rr * t_s1 + (cc + 32) * t_s2, top_right)
tl.store(base + (rr + 32) * t_s1 + cc * t_s2, tl.zeros((32, 32), tl.float32))
tl.store(base + (rr + 32) * t_s1 + (cc + 32) * t_s2, right)
else:
row_block = task - 1
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
first = cols < 32
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v1 + batch_id * v1_s0 + rows[:, None] * v1_s1 + cols[None, :] * v1_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
v2_rows = rows - 32
v2_cols = cols - 32
val2 = tl.load(
v2 + batch_id * v2_s0 + v2_rows[:, None] * v2_s1 + v2_cols[None, :] * v2_s2,
mask=(rows[:, None] >= 32) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _assemble_v64_compose_t64_direct_cross32_kernel(
v1,
v2,
t_left,
t_right,
v_out,
t_out,
v1_s0,
v1_s1,
v1_s2,
v2_s0,
v2_s1,
v2_s2,
tl_s0: tl.constexpr,
tl_s1: tl.constexpr,
tl_s2: tl.constexpr,
tr_s0: tl.constexpr,
tr_s1: tl.constexpr,
tr_s2: tl.constexpr,
v_s0,
v_s1,
v_s2,
t_s0,
t_s1,
t_s2,
rows_count,
cross_k: tl.constexpr,
block_r: tl.constexpr,
block_k: tl.constexpr,
):
batch_id = tl.program_id(0)
task = tl.program_id(1)
if task == 0:
idx = tl.arange(0, 32)
kk = tl.arange(0, block_k)
rr = idx[:, None]
cc = idx[None, :]
cross = tl.zeros((32, 32), tl.float32)
for start in tl.range(0, cross_k, block_k):
k_offsets = start + kk
k_mask = k_offsets < cross_k
left_tail = tl.load(
v1
+ batch_id * v1_s0
+ (32 + k_offsets)[None, :] * v1_s1
+ idx[:, None] * v1_s2,
mask=k_mask[None, :],
other=0.0,
).to(tl.float32)
right_tail = tl.load(
v2
+ batch_id * v2_s0
+ k_offsets[:, None] * v2_s1
+ idx[None, :] * v2_s2,
mask=k_mask[:, None],
other=0.0,
).to(tl.float32)
cross += tl.dot(left_tail, right_tail, input_precision="tf32", out_dtype=tl.float32)
left = tl.load(t_left + batch_id * tl_s0 + rr * tl_s1 + cc * tl_s2)
right = tl.load(t_right + batch_id * tr_s0 + rr * tr_s1 + cc * tr_s2)
tmp = tl.dot(left, cross, input_precision="tf32x3", out_dtype=tl.float32)
top_right = -tl.dot(tmp, right, input_precision="tf32x3", out_dtype=tl.float32)
base = t_out + batch_id * t_s0
tl.store(base + rr * t_s1 + cc * t_s2, left)
tl.store(base + rr * t_s1 + (cc + 32) * t_s2, top_right)
tl.store(base + (rr + 32) * t_s1 + cc * t_s2, tl.zeros((32, 32), tl.float32))
tl.store(base + (rr + 32) * t_s1 + (cc + 32) * t_s2, right)
else:
row_block = task - 1
rows = row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 64)
first = cols < 32
second = ~first
row_mask = rows < rows_count
val1 = tl.load(
v1 + batch_id * v1_s0 + rows[:, None] * v1_s1 + cols[None, :] * v1_s2,
mask=row_mask[:, None] & first[None, :],
other=0.0,
)
v2_rows = rows - 32
v2_cols = cols - 32
val2 = tl.load(
v2 + batch_id * v2_s0 + v2_rows[:, None] * v2_s1 + v2_cols[None, :] * v2_s2,
mask=(rows[:, None] >= 32) & row_mask[:, None] & second[None, :],
other=0.0,
)
out = tl.where(first[None, :], val1, val2)
tl.store(
v_out + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
@triton.jit
def _house_r16_chunks_kernel(
panel,
out_r,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
out_s0: tl.constexpr,
out_s1: tl.constexpr,
out_s2: tl.constexpr,
out_s3: tl.constexpr,
rows_count: tl.constexpr,
CHUNK_ROWS: tl.constexpr,
BLOCK_M: tl.constexpr,
POS_DIAG: tl.constexpr,
):
batch_id = tl.program_id(0)
chunk_id = tl.program_id(1)
rows = tl.arange(0, BLOCK_M)
cols = tl.arange(0, 16)
chunk_base = chunk_id * CHUNK_ROWS
chunk_n = tl.minimum(CHUNK_ROWS, rows_count - chunk_base)
row_mask = rows < chunk_n
ptrs = (
panel
+ batch_id * panel_s0
+ (chunk_base + rows)[:, None] * panel_s1
+ cols[None, :] * panel_s2
)
a = tl.load(ptrs, mask=row_mask[:, None], other=0.0).to(tl.float32)
for j in tl.static_range(0, 16):
col_j = tl.sum(tl.where(cols[None, :] == j, a, 0.0), axis=1)
alpha = tl.sum(tl.where(rows == j, col_j, 0.0), axis=0)
tail_norm2 = tl.sum(
tl.where((rows > j) & row_mask, col_j * col_j, 0.0),
axis=0,
)
use_reflector = tail_norm2 > 0.0
norm = tl.sqrt(alpha * alpha + tail_norm2)
sign = tl.where(alpha >= 0.0, 1.0, -1.0)
beta = tl.where(use_reflector, -sign * norm, alpha)
denom = tl.where(use_reflector, alpha - beta, 1.0)
tau_j = tl.where(use_reflector, (beta - alpha) / beta, 0.0)
packed_col = tl.where(
rows == j,
beta,
tl.where((rows > j) & row_mask, col_j / denom, col_j),
)
a = tl.where(cols[None, :] == j, packed_col[:, None], a)
v = tl.where(
rows == j,
1.0,
tl.where((rows > j) & row_mask, packed_col, 0.0),
)
dots = tl.sum(v[:, None] * a, axis=0)
update = tau_j * v[:, None] * dots[None, :]
a = tl.where((cols[None, :] > j) & row_mask[:, None], a - update, a)
out_ptrs = (
out_r
+ batch_id * out_s0
+ chunk_id * out_s1
+ rows[:, None] * out_s2
+ cols[None, :] * out_s3
)
r_vals = tl.where(cols[None, :] >= rows[:, None], a, 0.0)
if POS_DIAG:
diag = tl.sum(tl.where(rows[:, None] == cols[None, :], r_vals, 0.0), axis=1)
sign = tl.where(diag < 0.0, -1.0, 1.0)
r_vals = r_vals * sign[:, None]
tl.store(out_ptrs, r_vals, mask=(rows[:, None] < 16) & (cols[None, :] < 16))
@triton.jit
def _orhr16_folded_top_setup_direct_t_kernel(
panel,
r,
v,
n_out,
tau_out,
t_out,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
r_s0: tl.constexpr,
r_s1: tl.constexpr,
r_s2: tl.constexpr,
v_s0,
v_s1: tl.constexpr,
v_s2: tl.constexpr,
n_s0: tl.constexpr,
n_s1: tl.constexpr,
n_s2: tl.constexpr,
tau_s0: tl.constexpr,
tau_s1: tl.constexpr,
t_s0: tl.constexpr,
t_s1: tl.constexpr,
t_s2: tl.constexpr,
):
batch_id = tl.program_id(0)
idx = tl.arange(0, 16)
rr = idx[:, None]
cc = idx[None, :]
p = tl.load(
panel + batch_id * panel_s0 + rr * panel_s1 + cc * panel_s2,
).to(tl.float32)
r_vals = tl.load(
r + batch_id * r_s0 + rr * r_s1 + cc * r_s2,
).to(tl.float32)
q_top = p
for j in tl.static_range(0, 16):
r_col_j = tl.sum(tl.where(cc == j, r_vals, 0.0), axis=1)
rhs_j = tl.sum(tl.where(cc == j, q_top, 0.0), axis=1)
prev = tl.sum(
tl.where(idx[None, :] < j, q_top * r_col_j[None, :], 0.0),
axis=1,
)
pivot = tl.sum(tl.where(idx == j, r_col_j, 0.0), axis=0)
solved = (rhs_j - prev) / pivot
q_top = tl.where(cc == j, solved[:, None], q_top)
q_diag = tl.sum(tl.where(rr == cc, q_top, 0.0), axis=1)
signs = tl.where(q_diag > 0.0, -1.0, 1.0)
a = tl.where(rr == cc, q_top - signs[:, None], q_top)
for j in tl.static_range(0, 16):
row_j = tl.sum(tl.where(rr == j, a, 0.0), axis=0)
col_j = tl.sum(tl.where(cc == j, a, 0.0), axis=1)
pivot = tl.sum(tl.where(idx == j, row_j, 0.0), axis=0)
mult = col_j / pivot
update = mult[:, None] * row_j[None, :]
a = tl.where((rr > j) & (cc == j), mult[:, None], a)
a = tl.where((rr > j) & (cc > j), a - update, a)
lower = rr > cc
upper = cc >= rr
packed = tl.where(lower, a, 0.0) + tl.where(upper, signs[:, None] * r_vals, 0.0)
v_top = tl.where(rr == cc, 1.0, tl.where(lower, a, 0.0))
tl.store(panel + batch_id * panel_s0 + rr * panel_s1 + cc * panel_s2, packed)
tl.store(v + batch_id * v_s0 + rr * v_s1 + cc * v_s2, v_top)
u_vals = tl.where(upper, a, 0.0)
l_inv_t = tl.zeros((16, 16), tl.float32)
for ii in tl.static_range(0, 16):
i = 15 - ii
l_col_i = tl.sum(tl.where(cc == i, a, 0.0), axis=1)
prev = tl.sum(
tl.where(idx[:, None] > i, l_col_i[:, None] * l_inv_t, 0.0),
axis=0,
)
rhs = tl.where(idx == i, 1.0, 0.0)
solved = rhs - prev
l_inv_t = tl.where(rr == i, solved[None, :], l_inv_t)
u_signed = u_vals * signs[None, :]
t_vals = -tl.dot(u_signed, l_inv_t, input_precision="ieee")
t_vals = tl.where(upper, t_vals, 0.0)
u_diag = tl.sum(tl.where(rr == cc, u_vals, 0.0), axis=1)
tau_vals = -signs * u_diag
tl.store(tau_out + batch_id * tau_s0 + idx * tau_s1, tau_vals)
tl.store(t_out + batch_id * t_s0 + rr * t_s1 + cc * t_s2, t_vals)
u_inv = tl.zeros((16, 16), tl.float32)
for ii in tl.static_range(0, 16):
i = 15 - ii
u_row = tl.sum(tl.where(rr == i, u_vals, 0.0), axis=0)
prev = tl.sum(
tl.where(idx[:, None] > i, u_row[:, None] * u_inv, 0.0),
axis=0,
)
rhs = tl.where(idx == i, 1.0, 0.0)
pivot = tl.sum(tl.where(idx == i, u_row, 0.0), axis=0)
solved = (rhs - prev) / pivot
u_inv = tl.where(rr == i, solved[None, :], u_inv)
n_mat = tl.zeros((16, 16), tl.float32)
for ii in tl.static_range(0, 16):
i = 15 - ii
r_row = tl.sum(tl.where(rr == i, r_vals, 0.0), axis=0)
rhs = tl.sum(tl.where(rr == i, u_inv, 0.0), axis=0)
prev = tl.sum(
tl.where(idx[:, None] > i, r_row[:, None] * n_mat, 0.0),
axis=0,
)
pivot = tl.sum(tl.where(idx == i, r_row, 0.0), axis=0)
solved = (rhs - prev) / pivot
n_mat = tl.where(rr == i, solved[None, :], n_mat)
tl.store(n_out + batch_id * n_s0 + rr * n_s1 + cc * n_s2, n_mat)
@triton.jit
def _orhr16_folded_tail_matmul_pack_v_kernel(
panel,
v,
n_mat,
panel_s0: tl.constexpr,
panel_s1: tl.constexpr,
panel_s2: tl.constexpr,
v_s0,
v_s1: tl.constexpr,
v_s2: tl.constexpr,
n_s0: tl.constexpr,
n_s1: tl.constexpr,
n_s2: tl.constexpr,
rows_count,
block_r: tl.constexpr,
):
batch_id = tl.program_id(0)
row_block = tl.program_id(1)
rows = 16 + row_block * block_r + tl.arange(0, block_r)
cols = tl.arange(0, 16)
row_mask = rows < rows_count
p_vals = tl.load(
panel + batch_id * panel_s0 + rows[:, None] * panel_s1 + cols[None, :] * panel_s2,
mask=row_mask[:, None],
other=0.0,
).to(tl.float32)
n_vals = tl.load(
n_mat + batch_id * n_s0 + cols[:, None] * n_s1 + cols[None, :] * n_s2,
).to(tl.float32)
out = tl.dot(p_vals, n_vals, input_precision="ieee")
tl.store(
panel + batch_id * panel_s0 + rows[:, None] * panel_s1 + cols[None, :] * panel_s2,
out,
mask=row_mask[:, None],
)
tl.store(
v + batch_id * v_s0 + rows[:, None] * v_s1 + cols[None, :] * v_s2,
out,
mask=row_mask[:, None],
)
def _full_qr(data: torch.Tensor) -> output_t:
batch, n, _ = data.shape
h = torch.empty_like(data)
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
_full_qr_read_write_kernel[(batch,)](
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),
n,
)
return h, tau
def _finish_tail_qr_inplace(h: torch.Tensor, tau: torch.Tensor, k: int) -> None:
rows = h.shape[1] - k
_tail_qr_inplace_kernel[(h.shape[0],)](
h,
tau,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
k,
rows,
num_warps=8 if rows >= 64 else 4,
)
def _gram_fp32(v: torch.Tensor) -> torch.Tensor:
# The LARFT Gram uses single-pass TF32 here. Earlier full-bf16/medium Gram
# variants failed tolerance, but the active "high" setting has passed.
old_prec = torch.get_float32_matmul_precision()
torch.set_float32_matmul_precision("high")
gram = torch.bmm(v.transpose(1, 2), v)
torch.set_float32_matmul_precision(old_prec)
return gram
def _larft_forward(v: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
batch, _, width = v.shape
gram = _gram_fp32(v)
t = torch.zeros((batch, width, width), device=v.device, dtype=v.dtype)
t.diagonal(dim1=-2, dim2=-1).copy_(tau)
for j in range(1, width):
col = -tau[:, j][:, None, None] * gram[:, :j, j : j + 1]
t[:, :j, j : j + 1] = torch.bmm(t[:, :j, :j], col)
return t
def _larft_forward_from_gram(gram: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
batch, width, _ = gram.shape
t = torch.zeros((batch, width, width), device=gram.device, dtype=gram.dtype)
t.diagonal(dim1=-2, dim2=-1).copy_(tau)
for j in range(1, width):
col = -tau[:, j][:, None, None] * gram[:, :j, j : j + 1]
t[:, :j, j : j + 1] = torch.bmm(t[:, :j, :j], col)
return t
def _larft_triton(v: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
width = v.shape[2]
if not _HAS_TRITON or width not in (16, 32, 64):
return _larft_forward(v, tau)
gram = _gram_fp32(v)
batch = v.shape[0]
t = torch.empty((batch, width, width), device=v.device, dtype=v.dtype)
_larft16_kernel[(batch,)](
gram,
tau,
t,
gram.stride(0),
gram.stride(1),
gram.stride(2),
tau.stride(0),
tau.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
width,
)
return t
def _larft_triton_current_gram(v: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
width = v.shape[2]
if not _HAS_TRITON or width not in (16, 32, 64):
return _larft_forward(v, tau)
gram = torch.bmm(v.transpose(1, 2), v)
batch = v.shape[0]
t = torch.empty((batch, width, width), device=v.device, dtype=v.dtype)
_larft16_kernel[(batch,)](
gram,
tau,
t,
gram.stride(0),
gram.stride(1),
gram.stride(2),
tau.stride(0),
tau.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
width,
)
return t
def _larft_triton_from_gram(gram: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
width = tau.shape[1]
if not _HAS_TRITON or width not in (16, 32, 64):
return _larft_forward_from_gram(gram, tau)
batch = tau.shape[0]
t = torch.empty((batch, width, width), device=tau.device, dtype=tau.dtype)
_larft16_kernel[(batch,)](
gram,
tau,
t,
gram.stride(0),
gram.stride(1),
gram.stride(2),
tau.stride(0),
tau.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
width,
)
return t
def _larft64_blocked2_from_gram(gram: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
batch = tau.shape[0]
t = torch.empty((batch, 64, 64), device=tau.device, dtype=tau.dtype)
_larft64_blocked2_from_gram_kernel[(batch,)](
gram,
tau,
t,
gram.stride(0),
gram.stride(1),
gram.stride(2),
tau.stride(0),
tau.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
num_warps=8,
)
return t
def _chol64_upper_noinfo(gram: torch.Tensor) -> torch.Tensor:
if not _HAS_TRITON or gram.shape[-1] != 64:
chol = torch.linalg.cholesky_ex(gram, check_errors=False)[0]
return chol.transpose(1, 2).contiguous()
r = torch.empty_like(gram)
_chol64_upper_noinfo_kernel[(gram.shape[0],)](
gram,
r,
gram.stride(0),
gram.stride(1),
gram.stride(2),
r.stride(0),
r.stride(1),
r.stride(2),
num_warps=8,
)
return r
def _larft_triton_high_gram(v: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
width = v.shape[2]
if not _HAS_TRITON or width not in (16, 32, 64):
return _larft_forward(v, tau)
torch.set_float32_matmul_precision("high")
gram = torch.bmm(v.transpose(1, 2), v)
torch.set_float32_matmul_precision("medium")
batch = v.shape[0]
t = torch.empty((batch, width, width), device=v.device, dtype=v.dtype)
_larft16_kernel[(batch,)](
gram,
tau,
t,
gram.stride(0),
gram.stride(1),
gram.stride(2),
tau.stride(0),
tau.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
width,
)
return t
def _larft16(v: torch.Tensor, tau: torch.Tensor) -> torch.Tensor:
return _larft_triton(v, tau)
def _panel_reflectors(h_panel: torch.Tensor) -> torch.Tensor:
v = torch.tril(h_panel, diagonal=-1)
v.diagonal(dim1=-2, dim2=-1).fill_(1.0)
return v
def _larfb16_update(h: torch.Tensor, v: torch.Tensor, t: torch.Tensor, k: int) -> None:
rows = h.shape[1] - k
cols = h.shape[2] - (k + 16)
if cols <= 0:
return
block_m = _next_pow2(rows)
block_n = 16
_larfb16_update_kernel[(h.shape[0], triton.cdiv(cols, block_n))](
h,
v,
t,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
k,
rows,
cols,
block_m,
block_n,
num_warps=4,
)
def _larfb16_update_splitk352(h: torch.Tensor, v: torch.Tensor, t: torch.Tensor, k: int) -> None:
rows = h.shape[1] - k
cols = h.shape[2] - (k + 16)
if cols <= 0:
return
if rows < 176:
_larfb16_update(h, v, t, k)
return
chunk_m = 32
block_n = 32
row_parts = triton.cdiv(rows, chunk_m)
col_blocks = triton.cdiv(cols, block_n)
partial = torch.empty(
(h.shape[0], col_blocks, row_parts, 16, block_n),
device=h.device,
dtype=torch.float32,
)
grid = (h.shape[0], col_blocks, row_parts)
_larfb16_partial_w_kernel[grid](
h,
v,
partial,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
partial.stride(0),
partial.stride(1),
partial.stride(2),
partial.stride(3),
partial.stride(4),
k,
rows,
cols,
chunk_m,
block_n,
num_warps=4,
)
_larfb16_reduce_apply_splitk_kernel[grid](
h,
v,
t,
partial,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
partial.stride(0),
partial.stride(1),
partial.stride(2),
partial.stride(3),
partial.stride(4),
k,
rows,
cols,
row_parts,
chunk_m,
block_n,
num_warps=4,
)
def _larfb16_update_x3(h: torch.Tensor, v: torch.Tensor, t: torch.Tensor, k: int) -> None:
rows = h.shape[1] - k
cols = h.shape[2] - (k + 16)
if cols <= 0:
return
block_m = _next_pow2(rows)
block_n = 32
_larfb16_update_x3_kernel[(h.shape[0], triton.cdiv(cols, block_n))](
h,
v,
t,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
k,
rows,
cols,
block_m,
block_n,
num_warps=4,
)
def _factor_panel(h: torch.Tensor, tau: torch.Tensor, k: int, width: int, emit_v: bool = True):
if _HAS_TRITON and h.shape[1] in (176, 352, 512, 1024, 2048) and width == 16:
rows = h.shape[1] - k
block_m = _next_pow2(rows)
v = torch.empty((h.shape[0], rows, width), device=h.device, dtype=h.dtype) if emit_v else torch.empty((h.shape[0], 1, width), device=h.device, dtype=h.dtype)
panel_warps = (
(4 if rows <= 320 else 8)
if h.shape[1] == 512
else (
(16 if rows <= 976 else 32)
if h.shape[1] == 1024
else (
8
if h.shape[1] in (176, 352) and block_m >= 256
else (
(
32
if k < 640 and block_m >= 2048
else (16 if block_m >= 1024 else (8 if block_m >= 512 else 4))
)
if h.shape[1] == 2048
else (8 if block_m >= 1024 else 4)
)
)
)
)
panel16_kernel = (
_panel16_qr_fixed_rows_paired_norm_kernel
if h.shape[1] in (512, 1024, 2048)
else (
_panel16_qr_fixed_rows_kernel
if h.shape[1] in (176, 352)
else _panel16_qr_kernel
)
)
panel16_kernel[(h.shape[0],)](
h,
tau,
v,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
v.stride(0),
v.stride(1),
v.stride(2),
k,
rows,
block_m,
emit_v,
num_warps=panel_warps,
)
return h[:, k:, k : k + width], tau[:, k : k + width], (v if emit_v else None)
h_panel, tau_panel = torch.geqrf(h[:, k:, k : k + width])
h[:, k:, k : k + width] = h_panel
tau[:, k : k + width] = tau_panel
return h_panel, tau_panel, None
def _factor_panel16_update_next16(
h: torch.Tensor,
tau: torch.Tensor,
k: int,
emit_v: bool = True,
):
if not (_HAS_TRITON and h.shape[1] == 1024 and k + 32 <= h.shape[2]):
h_panel, tau_panel, v = _factor_panel(h, tau, k, 16, emit_v=emit_v)
if v is None:
v = _panel_reflectors(h_panel)
_apply_split16_local_direct_forward(h, v, tau_panel, k, num_warps=16)
return h_panel, tau_panel, (v if emit_v else None)
rows = h.shape[1] - k
block_m = _next_pow2(rows)
v = (
torch.empty((h.shape[0], rows, 16), device=h.device, dtype=h.dtype)
if emit_v
else torch.empty((h.shape[0], 1, 16), device=h.device, dtype=h.dtype)
)
_panel16_qr_update_next16_fixed_rows_paired_norm_kernel[(h.shape[0],)](
h,
tau,
v,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
v.stride(0),
v.stride(1),
v.stride(2),
k,
rows,
block_m,
emit_v,
num_warps=16,
)
return h[:, k:, k : k + 16], tau[:, k : k + 16], (v if emit_v else None)
def _factor_panel16_n1024_tail_warps(
h: torch.Tensor,
tau: torch.Tensor,
k: int,
emit_v: bool = True,
panel_warps: int = 4,
):
if not (_HAS_TRITON and h.shape[1] == 1024):
return _factor_panel(h, tau, k, 16, emit_v=emit_v)
rows = h.shape[1] - k
block_m = _next_pow2(rows)
v = (
torch.empty((h.shape[0], rows, 16), device=h.device, dtype=h.dtype)
if emit_v
else torch.empty((h.shape[0], 1, 16), device=h.device, dtype=h.dtype)
)
_panel16_qr_fixed_rows_paired_norm_kernel[(h.shape[0],)](
h,
tau,
v,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
v.stride(0),
v.stride(1),
v.stride(2),
k,
rows,
block_m,
emit_v,
num_warps=panel_warps,
)
return h[:, k:, k : k + 16], tau[:, k : k + 16], (v if emit_v else None)
def _rowsplit_tsqr16_r(panel: torch.Tensor, chunk_rows: int = 256) -> torch.Tensor:
batch, rows, width = panel.shape
if not _HAS_TRITON or width != 16:
raise RuntimeError("rowsplit TSQR16 requires Triton and width 16")
n_chunks = triton.cdiv(rows, chunk_rows)
block_m = _next_pow2(chunk_rows)
r_chunks = torch.empty((batch, n_chunks, 16, 16), device=panel.device, dtype=panel.dtype)
_house_r16_chunks_kernel[(batch, n_chunks)](
panel,
r_chunks,
panel.stride(0),
panel.stride(1),
panel.stride(2),
r_chunks.stride(0),
r_chunks.stride(1),
r_chunks.stride(2),
r_chunks.stride(3),
rows,
chunk_rows,
block_m,
False,
num_warps=8 if block_m >= 512 else 4,
)
stacked = r_chunks.reshape(batch, n_chunks * 16, 16).contiguous()
red_rows = n_chunks * 16
r_slot = torch.empty((batch, 1, 16, 16), device=panel.device, dtype=panel.dtype)
_house_r16_chunks_kernel[(batch, 1)](
stacked,
r_slot,
stacked.stride(0),
stacked.stride(1),
stacked.stride(2),
r_slot.stride(0),
r_slot.stride(1),
r_slot.stride(2),
r_slot.stride(3),
red_rows,
red_rows,
_next_pow2(red_rows),
True,
num_warps=4,
)
r = r_slot[:, 0].contiguous()
return r
def _orhr16_folded_direct_t(panel: torch.Tensor, r: torch.Tensor, tail_block: int = 64):
batch, rows, width = panel.shape
if not _HAS_TRITON or width != 16:
raise RuntimeError("ORHR16 direct-T bridge requires Triton and width 16")
v = torch.empty_like(panel)
n_mat = torch.empty((batch, 16, 16), device=panel.device, dtype=panel.dtype)
tau_panel = torch.empty((batch, 16), device=panel.device, dtype=panel.dtype)
t = torch.empty((batch, 16, 16), device=panel.device, dtype=panel.dtype)
_orhr16_folded_top_setup_direct_t_kernel[(batch,)](
panel,
r,
v,
n_mat,
tau_panel,
t,
panel.stride(0),
panel.stride(1),
panel.stride(2),
r.stride(0),
r.stride(1),
r.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
n_mat.stride(0),
n_mat.stride(1),
n_mat.stride(2),
tau_panel.stride(0),
tau_panel.stride(1),
t.stride(0),
t.stride(1),
t.stride(2),
num_warps=1,
)
if rows > 16:
_orhr16_folded_tail_matmul_pack_v_kernel[(batch, triton.cdiv(rows - 16, tail_block))](
panel,
v,
n_mat,
panel.stride(0),
panel.stride(1),
panel.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
n_mat.stride(0),
n_mat.stride(1),
n_mat.stride(2),
rows,
tail_block,
num_warps=4 if tail_block >= 64 else 1,
)
return v, tau_panel, t
def _factor_panel_rowsplit_tsqr16_direct_t(
h: torch.Tensor,
tau: torch.Tensor,
k: int,
chunk_rows: int = 256,
tail_block: int = 64,
):
panel = h[:, k:, k : k + 16]
r = _rowsplit_tsqr16_r(panel, chunk_rows=chunk_rows)
v, tau_panel, t = _orhr16_folded_direct_t(panel, r, tail_block=tail_block)
tau[:, k : k + 16] = tau_panel
return panel, tau_panel, v, t
def _factor_superpanel32_tail_gram(
h: torch.Tensor,
tau: torch.Tensor,
k: int,
panel_warps: int = 16,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
rows = h.shape[1] - k
block_m = _next_pow2(rows)
v = torch.empty((h.shape[0], rows, 32), device=h.device, dtype=h.dtype)
gram = torch.empty((h.shape[0], 32, 32), device=h.device, dtype=h.dtype)
panel32_tail_kernel = (
_panel32_qr_tail_gram_fixed_rows_kernel
if h.shape[1] in (512, 1024)
else _panel32_qr_tail_gram_kernel
)
panel32_tail_kernel[(h.shape[0],)](
h,
tau,
v,
gram,
h.stride(0),
h.stride(1),
h.stride(2),
tau.stride(0),
tau.stride(1),
v.stride(0),
v.stride(1),
v.stride(2),
gram.stride(0),
gram.stride(1),
gram.stride(2),
k,
rows,
block_m,
num_warps=panel_warps,
)
return tau[:, k : k + 32], v, gram
_SPLIT16_DIRECT_FIXED_ROW_SHAPES = (176, 352, 512, 1024, 2048)
def _apply_split16_local_direct_forward(
h: torch.Tensor,
v: torch.Tensor,
tau_panel: torch.Tensor,
k: int,
num_warps: int = 4,
) -> None:
if not _HAS_TRITON:
t = _larft_triton_current_gram(v, tau_panel)
_apply_wy_update_tfp32_baddbmm(h[:, k:, k + 16 : k + 32], v, t)
return
rows = h.shape[1] - k
if rows <= 0:
return
block_m = _next_pow2(rows)
direct_kernel = (
_split16_local_direct_forward_fixed_rows_kernel
if h.shape[1] in _SPLIT16_DIRECT_FIXED_ROW_SHAPES
else _split16_local_direct_forward_kernel
)
direct_kernel[(h.shape[0],)](
h,
v,
tau_panel,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
tau_panel.stride(0),
tau_panel.stride(1),
k,
rows,
block_m,
num_warps=num_warps,
)
def _apply_split16_trailing_direct_forward(
h: torch.Tensor,
v: torch.Tensor,
tau_panel: torch.Tensor,
k: int,
block_n: int = 16,
num_warps: int = 4,
fixed_k: bool = False,
) -> None:
if not _HAS_TRITON:
t = _larft_triton(v, tau_panel)
_larfb16_update(h, v, t, k)
return
rows = h.shape[1] - k
cols = h.shape[2] - (k + 16)
if rows <= 0 or cols <= 0:
return
block_m = _next_pow2(rows)
direct_kernel = (
_split16_trailing_direct_forward_fixed_k_rows_kernel
if fixed_k
else (
_split16_trailing_direct_forward_fixed_rows_kernel
if h.shape[1] in _SPLIT16_DIRECT_FIXED_ROW_SHAPES
else _split16_trailing_direct_forward_kernel
)
)
direct_kernel[(h.shape[0], triton.cdiv(cols, block_n))](
h,
v,
tau_panel,
h.stride(0),
h.stride(1),
h.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
tau_panel.stride(0),
tau_panel.stride(1),
k,
rows,
cols,
block_m,
block_n,
num_warps=num_warps,
)
def _factor_superpanel32_split16(h: torch.Tensor, tau: torch.Tensor, k: int):
h_panel1, tau1, v1 = _factor_panel(h, tau, k, 16)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
mid = k + 16
end = k + 32
local = h[:, k:, mid:end]
if local.shape[2] > 0:
_apply_split16_local_direct_forward(h, v1, tau1, k)
h_panel2, tau2, v2 = _factor_panel(h, tau, mid, 16)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
return _assemble_vtau32_from16(v1, v2, tau1, tau2)
def _factor_superpanel32_split16_fastlocal(h: torch.Tensor, tau: torch.Tensor, k: int):
rows = h.shape[1] - k
use_fused_next16 = _HAS_TRITON and h.shape[1] == 1024 and rows >= 768 and k + 32 <= h.shape[2]
if use_fused_next16:
h_panel1, tau1, v1 = _factor_panel16_update_next16(h, tau, k)
else:
h_panel1, tau1, v1 = _factor_panel(h, tau, k, 16)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
mid = k + 16
end = k + 32
local = h[:, k:, mid:end]
if local.shape[2] > 0 and not use_fused_next16:
_apply_split16_local_direct_forward(h, v1, tau1, k, num_warps=16)
h_panel2, tau2, v2 = _factor_panel(h, tau, mid, 16)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
return _assemble_vtau32_from16(v1, v2, tau1, tau2)
def _assemble_v64_from32(
v1: torch.Tensor,
v2: torch.Tensor,
fixed_rows: bool = False,
) -> torch.Tensor:
batch, rows, _ = v1.shape
v = torch.empty((batch, rows, 64), device=v1.device, dtype=v1.dtype)
if not _HAS_TRITON:
v[:, :, :32] = v1
v[:, :32, 32:] = 0.0
v[:, 32:, 32:] = v2
return v
block_r = 32
assemble_kernel = (
_assemble_v64_from32_fixed_rows_kernel
if fixed_rows
else _assemble_v64_from32_kernel
)
assemble_kernel[(batch, triton.cdiv(rows, block_r))](
v1,
v2,
v,
v1.stride(0),
v1.stride(1),
v1.stride(2),
v2.stride(0),
v2.stride(1),
v2.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
rows,
block_r,
num_warps=4,
)
return v
def _assemble_v64_compose_t64_from32(
v1: torch.Tensor,
v2: torch.Tensor,
t_left: torch.Tensor,
t_right: torch.Tensor,
cross0: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
batch, rows, _ = v1.shape
v = torch.empty((batch, rows, 64), device=v1.device, dtype=v1.dtype)
t = torch.empty((batch, 64, 64), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
v[:, :, :32] = v1
v[:, :32, 32:] = 0.0
v[:, 32:, 32:] = v2
cross = -torch.bmm(torch.bmm(t_left, cross0), t_right)
t[:, :32, :32] = t_left
t[:, :32, 32:] = cross
t[:, 32:, :32] = 0.0
t[:, 32:, 32:] = t_right
return v, t
block_r = 128 if rows >= 448 else (64 if rows >= 320 else 32)
_assemble_v64_compose_t64_from32_kernel[(batch, 1 + triton.cdiv(rows, block_r))](
v1,
v2,
t_left,
t_right,
cross0,
v,
t,
v1.stride(0),
v1.stride(1),
v1.stride(2),
v2.stride(0),
v2.stride(1),
v2.stride(2),
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross0.stride(0),
cross0.stride(1),
cross0.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
rows,
block_r,
num_warps=4,
)
return v, t
def _assemble_v64_compose_t64_direct_cross32(
v1: torch.Tensor,
v2: torch.Tensor,
t_left: torch.Tensor,
t_right: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
batch, rows, _ = v1.shape
if not _HAS_TRITON:
cross0 = torch.bmm(v1[:, 32:, :].transpose(1, 2), v2)
return _assemble_v64_compose_t64_from32(v1, v2, t_left, t_right, cross0)
v = torch.empty((batch, rows, 64), device=v1.device, dtype=v1.dtype)
t = torch.empty((batch, 64, 64), device=t_left.device, dtype=t_left.dtype)
block_r = 128 if rows >= 448 else (64 if rows >= 320 else 32)
_assemble_v64_compose_t64_direct_cross32_kernel[
(batch, 1 + triton.cdiv(rows, block_r))
](
v1,
v2,
t_left,
t_right,
v,
t,
v1.stride(0),
v1.stride(1),
v1.stride(2),
v2.stride(0),
v2.stride(1),
v2.stride(2),
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
rows,
v2.shape[1],
block_r,
128,
num_warps=4,
)
return v, t
def _assemble_v64_tail_gram_compose_t64_from32(
v1: torch.Tensor,
v2: torch.Tensor,
t_left: torch.Tensor,
tau2: torch.Tensor,
full_n: int | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
v = _assemble_v64_from32(v1, v2, fixed_rows=full_n == 512)
gram_tail = torch.bmm(v[:, 32:, :].transpose(1, 2), v[:, 32:, :])
t_right = _larft_triton_from_gram(gram_tail[:, 32:, 32:], tau2)
t = _compose_t64_from_cross32(t_left, t_right, gram_tail[:, :32, 32:])
return v, t
def _assemble_vtau32_from16(
v1: torch.Tensor,
v2: torch.Tensor,
tau1: torch.Tensor,
tau2: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
batch, rows, _ = v1.shape
v = torch.empty((batch, rows, 32), device=v1.device, dtype=v1.dtype)
tau = torch.empty((batch, 32), device=v1.device, dtype=v1.dtype)
if not _HAS_TRITON:
v[:, :, :16] = v1
v[:, :16, 16:] = 0.0
v[:, 16:, 16:] = v2
tau[:, :16] = tau1
tau[:, 16:] = tau2
return v, tau
block_r = 32
_assemble_vtau32_from16_kernel[(batch, triton.cdiv(rows, block_r))](
v1,
v2,
tau1,
tau2,
v,
tau,
v1.stride(0),
v1.stride(1),
v1.stride(2),
v2.stride(0),
v2.stride(1),
v2.stride(2),
tau1.stride(0),
tau1.stride(1),
tau2.stride(0),
tau2.stride(1),
v.stride(0),
v.stride(1),
v.stride(2),
tau.stride(0),
tau.stride(1),
rows,
block_r,
num_warps=4,
)
return v, tau
def _assemble_v128_from64(v_left: torch.Tensor, v_right: torch.Tensor) -> torch.Tensor:
batch, rows, _ = v_left.shape
v = torch.empty((batch, rows, 128), device=v_left.device, dtype=v_left.dtype)
if not _HAS_TRITON:
v[:, :, :64] = v_left
v[:, :64, 64:] = 0.0
v[:, 64:, 64:] = v_right
return v
block_r = 16
_assemble_v128_from64_kernel[(batch, triton.cdiv(rows, block_r))](
v_left,
v_right,
v,
v_left.stride(0),
v_left.stride(1),
v_left.stride(2),
v_right.stride(0),
v_right.stride(1),
v_right.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
rows,
block_r,
num_warps=4,
)
return v
def _assemble_t128_from64(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross: torch.Tensor,
) -> torch.Tensor:
batch = t_left.shape[0]
t = torch.empty((batch, 128, 128), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
t[:, :64, :64] = t_left
t[:, :64, 64:] = cross
t[:, 64:, :64] = 0.0
t[:, 64:, 64:] = t_right
return t
_assemble_t128_from64_kernel[(batch,)](
t_left,
t_right,
cross,
t,
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross.stride(0),
cross.stride(1),
cross.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
num_warps=8,
)
return t
def _assemble_v256_from128(v_left: torch.Tensor, v_right: torch.Tensor) -> torch.Tensor:
batch, rows, _ = v_left.shape
v = torch.empty((batch, rows, 256), device=v_left.device, dtype=v_left.dtype)
if not _HAS_TRITON:
v[:, :, :128] = v_left
v[:, :128, 128:] = 0.0
v[:, 128:, 128:] = v_right
return v
block_r = 8
_assemble_v256_from128_kernel[(batch, triton.cdiv(rows, block_r))](
v_left,
v_right,
v,
v_left.stride(0),
v_left.stride(1),
v_left.stride(2),
v_right.stride(0),
v_right.stride(1),
v_right.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
rows,
block_r,
num_warps=8,
)
return v
def _assemble_t256_from128(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross: torch.Tensor,
) -> torch.Tensor:
batch = t_left.shape[0]
t = torch.empty((batch, 256, 256), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
t[:, :128, :128] = t_left
t[:, :128, 128:] = cross
t[:, 128:, :128] = 0.0
t[:, 128:, 128:] = t_right
return t
block_c = 32
_assemble_t256_from128_kernel[(batch, triton.cdiv(128, block_c))](
t_left,
t_right,
cross,
t,
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross.stride(0),
cross.stride(1),
cross.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
block_c,
num_warps=8,
)
return t
def _compose_t256_from_cross128(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross0: torch.Tensor,
) -> torch.Tensor:
cross = -torch.bmm(torch.bmm(t_left, cross0), t_right)
return _assemble_t256_from128(t_left, t_right, cross)
def _assemble_v512_from256(v_left: torch.Tensor, v_right: torch.Tensor) -> torch.Tensor:
batch, rows, _ = v_left.shape
v = torch.empty((batch, rows, 512), device=v_left.device, dtype=v_left.dtype)
if not _HAS_TRITON:
v[:, :, :256] = v_left
v[:, :256, 256:] = 0.0
v[:, 256:, 256:] = v_right
return v
block_r = 4
_assemble_v512_from256_kernel[(batch, triton.cdiv(rows, block_r))](
v_left,
v_right,
v,
v_left.stride(0),
v_left.stride(1),
v_left.stride(2),
v_right.stride(0),
v_right.stride(1),
v_right.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
rows,
block_r,
num_warps=8,
)
return v
def _assemble_t512_from256(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross: torch.Tensor,
) -> torch.Tensor:
batch = t_left.shape[0]
t = torch.empty((batch, 512, 512), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
t[:, :256, :256] = t_left
t[:, :256, 256:] = cross
t[:, 256:, :256] = 0.0
t[:, 256:, 256:] = t_right
return t
block_c = 16
_assemble_t512_from256_kernel[(batch, triton.cdiv(256, block_c))](
t_left,
t_right,
cross,
t,
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross.stride(0),
cross.stride(1),
cross.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
block_c,
num_warps=8,
)
return t
def _compose_t512_from_cross256(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross0: torch.Tensor,
) -> torch.Tensor:
cross = -torch.bmm(torch.bmm(t_left, cross0), t_right)
return _assemble_t512_from256(t_left, t_right, cross)
def _compose_t128_from_cross64(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross0: torch.Tensor,
) -> torch.Tensor:
batch = t_left.shape[0]
t = torch.empty((batch, 128, 128), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
cross = -torch.bmm(torch.bmm(t_left, cross0), t_right)
return _assemble_t128_from64(t_left, t_right, cross)
_compose_t128_from_cross64_kernel[(batch,)](
t_left,
t_right,
cross0,
t,
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross0.stride(0),
cross0.stride(1),
cross0.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
num_warps=8,
)
return t
def _compose_t64_from_cross32(
t_left: torch.Tensor,
t_right: torch.Tensor,
cross0: torch.Tensor,
) -> torch.Tensor:
batch = t_left.shape[0]
t = torch.empty((batch, 64, 64), device=t_left.device, dtype=t_left.dtype)
if not _HAS_TRITON:
cross = -torch.bmm(torch.bmm(t_left, cross0), t_right)
t[:, :32, :32] = t_left
t[:, :32, 32:] = cross
t[:, 32:, :32] = 0.0
t[:, 32:, 32:] = t_right
return t
_compose_t64_from_cross32_kernel[(batch,)](
t_left,
t_right,
cross0,
t,
t_left.stride(0),
t_left.stride(1),
t_left.stride(2),
t_right.stride(0),
t_right.stride(1),
t_right.stride(2),
cross0.stride(0),
cross0.stride(1),
cross0.stride(2),
t.stride(0),
t.stride(1),
t.stride(2),
num_warps=4,
)
return t
def _tensor_geqrf(data: torch.Tensor) -> output_t:
geqrf = getattr(data, "geqrf", None)
if geqrf is not None:
return geqrf()
return torch.geqrf(data)
def _lu64_no_pivot(m_top: torch.Tensor) -> torch.Tensor:
if not _HAS_TRITON:
lu_top, _, _ = torch.linalg.lu_factor_ex(m_top, pivot=False, check_errors=False)
return lu_top
lu_top = torch.empty_like(m_top)
_lu64_no_pivot_kernel[(m_top.shape[0],)](
m_top,
lu_top,
m_top.stride(0),
m_top.stride(1),
m_top.stride(2),
lu_top.stride(0),
lu_top.stride(1),
lu_top.stride(2),
num_warps=8,
)
return lu_top
def _orhr64_pack_panel_v_(
panel: torch.Tensor,
m: torch.Tensor,
lu_top: torch.Tensor,
r: torch.Tensor,
signs: torch.Tensor,
) -> torch.Tensor:
if not _HAS_TRITON:
width = 64
packed = m
packed[:, :width, :].copy_(torch.tril(lu_top, diagonal=-1))
packed[:, :width, :].add_(torch.triu(signs[:, :, None] * r))
panel.copy_(packed)
v = packed
v.tril_(diagonal=-1)
v.diagonal(dim1=1, dim2=2).fill_(1.0)
return v
block_r = 16
rows = m.shape[1]
_orhr64_pack_panel_v_kernel[(m.shape[0], (rows + block_r - 1) // block_r)](
panel,
m,
lu_top,
r,
signs,
panel.stride(0),
panel.stride(1),
panel.stride(2),
m.stride(0),
m.stride(1),
m.stride(2),
lu_top.stride(0),
lu_top.stride(1),
lu_top.stride(2),
r.stride(0),
r.stride(1),
r.stride(2),
signs.stride(0),
signs.stride(1),
rows,
block_r,
num_warps=4,
)
return m
def _orhr64_reconstruct_panel_v_triton_(
panel: torch.Tensor,
q: torch.Tensor,
r: torch.Tensor,
) -> torch.Tensor:
if not _HAS_TRITON:
width = 64
diag_view = q.diagonal(dim1=1, dim2=2)
signs = torch.where(diag_view > 0, -1.0, 1.0)
m = q
m.diagonal(dim1=1, dim2=2).sub_(signs)
lu_top = _lu64_no_pivot(m[:, :width, :])
if m.shape[1] > width:
u = torch.triu(lu_top)
lower_tail_t = torch.linalg.solve_triangular(
u.transpose(1, 2),
m[:, width:, :].transpose(1, 2).contiguous(),
upper=False,
left=True,
)
m[:, width:, :] = lower_tail_t.transpose(1, 2)
return _orhr64_pack_panel_v_(panel, m, lu_top, r, signs)
batch, rows, width = q.shape
if width != 64:
diag_view = q.diagonal(dim1=1, dim2=2)
signs = torch.where(diag_view > 0, -1.0, 1.0)
m = q
m.diagonal(dim1=1, dim2=2).sub_(signs)
lu_top = _lu64_no_pivot(m[:, :width, :])
if rows > width:
u = torch.triu(lu_top)
lower_tail_t = torch.linalg.solve_triangular(
u.transpose(1, 2),
m[:, width:, :].transpose(1, 2).contiguous(),
upper=False,
left=True,
)
m[:, width:, :] = lower_tail_t.transpose(1, 2)
return _orhr64_pack_panel_v_(panel, m, lu_top, r, signs)
lu_top = torch.empty((batch, 64, 64), device=q.device, dtype=q.dtype)
_orhr64_top_lu_pack_v_kernel[(batch,)](
panel,
q,
r,
lu_top,
panel.stride(0),
panel.stride(1),
panel.stride(2),
q.stride(0),
q.stride(1),
q.stride(2),
r.stride(0),
r.stride(1),
r.stride(2),
lu_top.stride(0),
lu_top.stride(1),
lu_top.stride(2),
num_warps=4,
)
if rows > 64:
block_r = 16
_orhr64_tail_solve_pack_v_kernel[(batch, (rows - 64 + block_r - 1) // block_r)](
panel,
q,
lu_top,
panel.stride(0),
panel.stride(1),
panel.stride(2),
q.stride(0),
q.stride(1),
q.stride(2),
lu_top.stride(0),
lu_top.stride(1),
lu_top.stride(2),
rows,
block_r,
num_warps=4,
)
return q
def _orhr64_reconstruct_panel_v_cols_from_qt_triton_(
panel: torch.Tensor,
q_t: torch.Tensor,
r: torch.Tensor,
) -> torch.Tensor:
batch, width, rows = q_t.shape
if not _HAS_TRITON:
q = q_t.transpose(1, 2).contiguous()
return _orhr64_reconstruct_panel_v_triton_(panel, q, r)
if width != 64:
q = q_t.transpose(1, 2).contiguous()
return _orhr64_reconstruct_panel_v_triton_(panel, q, r)
v_out = torch.empty_like(panel)
lu_cols = torch.empty((batch, 64, 64), device=q_t.device, dtype=q_t.dtype)
_orhr64_top_lu_pack_v_cols_qt_kernel[(batch,)](
panel,
q_t,
v_out,
r,
lu_cols,
panel.stride(0),
panel.stride(1),
panel.stride(2),
q_t.stride(0),
q_t.stride(1),
q_t.stride(2),
v_out.stride(0),
v_out.stride(1),
v_out.stride(2),
r.stride(0),
r.stride(1),
r.stride(2),
lu_cols.stride(0),
lu_cols.stride(1),
lu_cols.stride(2),
num_warps=4,
)
if rows > 64:
block_r = 16
_orhr64_tail_solve_pack_v_cols_qt_kernel[(batch, (rows - 64 + block_r - 1) // block_r)](
panel,
q_t,
v_out,
lu_cols,
panel.stride(0),
panel.stride(1),
panel.stride(2),
q_t.stride(0),
q_t.stride(1),
q_t.stride(2),
v_out.stride(0),
v_out.stride(1),
v_out.stride(2),
lu_cols.stride(0),
lu_cols.stride(1),
lu_cols.stride(2),
rows,
block_r,
num_warps=4,
)
return v_out
def _factor_panel_cholesky_orhr_v_trace_cols_qt(panel: torch.Tensor) -> torch.Tensor:
batch, rows, width = panel.shape
gram = torch.bmm(panel.transpose(1, 2), panel)
r = _chol64_upper_noinfo(gram)
q_t = torch.linalg.solve_triangular(
r.transpose(1, 2),
panel.transpose(1, 2),
upper=False,
left=True,
)
return _orhr64_reconstruct_panel_v_cols_from_qt_triton_(panel, q_t, r)
def _factor_panel_cholesky_orhr_v_trace_cols_medium_gram_qt(panel: torch.Tensor) -> torch.Tensor:
old_precision = torch.get_float32_matmul_precision()
torch.set_float32_matmul_precision("medium")
gram = torch.bmm(panel.transpose(1, 2), panel)
torch.set_float32_matmul_precision(old_precision)
r = _chol64_upper_noinfo(gram)
q_t = torch.linalg.solve_triangular(
r.transpose(1, 2),
panel.transpose(1, 2),
upper=False,
left=True,
)
return _orhr64_reconstruct_panel_v_cols_from_qt_triton_(panel, q_t, r)
def _larft_triton_tau_from_gram(v: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
width = v.shape[2]
if not _HAS_TRITON or width not in (16, 32, 64):
tau = 2.0 / torch.sum(v * v, dim=1)
return tau, _larft_forward(v, tau)
gram = _gram_fp32(v)
diag = gram.diagonal(dim1=1, dim2=2)
tau = (2.0 / diag).contiguous()
if width == 64 and v.shape[0] == 2:
return tau, _larft64_blocked2_from_gram(gram, tau)
return tau, _larft_triton_from_gram(gram, tau)
def _factor_cholesky_orhr256_taugram_4096_group_packed(
h: torch.Tensor,
tau: torch.Tensor,
k: int,
update_end: int | None = None,
):
_, n, _ = h.shape
block = 64
pair = 128
group = 256
mid1 = k + block
end1 = min(k + pair, n)
end2 = min(k + group, n)
panel1 = h[:, k:, k:mid1]
if panel1.shape[1] <= 2048:
v1 = _factor_panel_cholesky_orhr_v_trace_cols_medium_gram_qt(panel1)
else:
v1 = _factor_panel_cholesky_orhr_v_trace_cols_qt(panel1)
tau1, t1 = _larft_triton_tau_from_gram(v1)
tau[:, k:mid1] = tau1
if mid1 >= n:
return None, None, n
local = h[:, k:, mid1:end1]
if local.shape[2] > 0:
_apply_wy_update_medium_tf32_t_baddbmm(local, v1, t1)
panel2 = h[:, mid1:, mid1:end1]
if panel2.shape[1] <= 2048:
v2 = _factor_panel_cholesky_orhr_v_trace_cols_medium_gram_qt(panel2)
else:
v2 = _factor_panel_cholesky_orhr_v_trace_cols_qt(panel2)
tau2, t2 = _larft_triton_tau_from_gram(v2)
tau[:, mid1:end1] = tau2
if end1 >= n:
return None, None, end1
cross_raw = torch.bmm(v1[:, block:, :].transpose(1, 2), v2)
v128_left = _assemble_v128_from64(v1, v2)
if n - end1 <= 512:
t128_left = _compose_t128_from_cross64(t1, t2, cross_raw)
else:
cross = -torch.bmm(torch.bmm(t1, cross_raw), t2)
t128_left = _assemble_t128_from64(t1, t2, cross)
local_next = h[:, k:, end1:end2]
if local_next.shape[2] > 0:
_apply_wy_update_medium_tf32_t_baddbmm(local_next, v128_left, t128_left)
mid2 = end1 + block
panel3 = h[:, end1:, end1:mid2]
if panel3.shape[1] <= 2048:
v3 = _factor_panel_cholesky_orhr_v_trace_cols_medium_gram_qt(panel3)
else:
v3 = _factor_panel_cholesky_orhr_v_trace_cols_qt(panel3)
tau3, t3 = _larft_triton_tau_from_gram(v3)
tau[:, end1:mid2] = tau3
if mid2 >= n:
return None, None, end2
local = h[:, end1:, mid2:end2]
if local.shape[2] > 0:
_apply_wy_update_medium_tf32_t_baddbmm(local, v3, t3)
panel4 = h[:, mid2:, mid2:end2]
if panel4.shape[1] <= 2048:
v4 = _factor_panel_cholesky_orhr_v_trace_cols_medium_gram_qt(panel4)
else:
v4 = _factor_panel_cholesky_orhr_v_trace_cols_qt(panel4)
tau4, t4 = _larft_triton_tau_from_gram(v4)
tau[:, mid2:end2] = tau4
if end2 >= n:
return None, None, end2
cross_raw = torch.bmm(v3[:, block:, :].transpose(1, 2), v4)
v128_right = _assemble_v128_from64(v3, v4)
if n - end2 <= 512:
t128_right = _compose_t128_from_cross64(t3, t4, cross_raw)
else:
cross = -torch.bmm(torch.bmm(t3, cross_raw), t4)
t128_right = _assemble_t128_from64(t3, t4, cross)
cross_raw = torch.bmm(v128_left[:, pair:, :].transpose(1, 2), v128_right)
v256 = _assemble_v256_from128(v128_left, v128_right)
t256 = _compose_t256_from_cross128(t128_left, t128_right, cross_raw)
target_end = n if update_end is None else min(update_end, n)
if target_end > end2:
_apply_wy_update_medium_tf32_t_baddbmm(h[:, k:, end2:target_end], v256, t256)
return v256, t256, end2
def _blocked_cholesky_orhr512_taugram_4096_cols_packed_early(
data: torch.Tensor,
inplace_input: bool = False,
) -> output_t:
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("high")
group = 256
early_limit = 1536
try:
k = 0
while k < n:
if k < early_limit and k + 2 * group < n:
pair_end = k + 2 * group
v256_left, t256_left, end_left = _factor_cholesky_orhr256_taugram_4096_group_packed(
h, tau, k, update_end=pair_end
)
if v256_left is None or t256_left is None:
k = end_left
continue
v256_right, t256_right, end_right = _factor_cholesky_orhr256_taugram_4096_group_packed(
h, tau, end_left, update_end=pair_end
)
if v256_right is None or t256_right is None:
k = end_right
continue
if end_right < n:
cross_raw = torch.bmm(v256_left[:, group:, :].transpose(1, 2), v256_right)
v512 = _assemble_v512_from256(v256_left, v256_right)
t512 = _compose_t512_from_cross256(t256_left, t256_right, cross_raw)
_apply_wy_update_medium_tf32_t_baddbmm(h[:, k:, end_right:], v512, t512)
k = end_right
else:
_, _, end = _factor_cholesky_orhr256_taugram_4096_group_packed(
h, tau, k, update_end=None
)
if end <= k:
break
k = end
return h, tau
except Exception:
if inplace_input:
raise
return _tensor_geqrf(data)
def _apply_wy_update(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
w = torch.bmm(v.transpose(1, 2), target)
w = torch.bmm(t.transpose(1, 2), w)
target.sub_(torch.bmm(v, w))
def _apply_wy_update_x3(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
w = _bmm_3xtf32(v.transpose(1, 2), target)
w = _bmm_fp32(t.transpose(1, 2), w)
target.sub_(_bmm_3xtf32(v, w))
def _apply_wy_update_k0_split(
target: torch.Tensor,
v: torch.Tensor,
t: torch.Tensor,
exact_cols: int,
) -> None:
w = torch.bmm(v.transpose(1, 2), target)
w = _bmm_fp32(t.transpose(1, 2), w)
exact_cols = min(exact_cols, target.shape[2])
if exact_cols > 0:
_baddbmm_fp32_(target[:, :, :exact_cols], v, w[:, :, :exact_cols])
if exact_cols < target.shape[2]:
target_fast = target[:, :, exact_cols:]
w_fast = w[:, :, exact_cols:]
if target_fast.shape[2] >= 128:
torch.baddbmm(target_fast, v, w_fast, beta=1.0, alpha=-1.0, out=target_fast)
else:
target_fast.sub_(torch.bmm(v, w_fast))
def _apply_wy_update_tfp32_baddbmm(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
w = torch.bmm(v.transpose(1, 2), target)
w = _bmm_fp32(t.transpose(1, 2), w)
torch.baddbmm(target, v, w, beta=1.0, alpha=-1.0, out=target)
def _apply_wy_update_tsplit_baddbmm(
target: torch.Tensor,
v: torch.Tensor,
t: torch.Tensor,
exact_cols: int,
) -> None:
w = torch.bmm(v.transpose(1, 2), target)
exact_cols = min(max(exact_cols, 0), target.shape[2])
tt = t.transpose(1, 2)
if exact_cols <= 0:
w_fast = torch.bmm(tt, w)
torch.baddbmm(target, v, w_fast, beta=1.0, alpha=-1.0, out=target)
return
if exact_cols >= target.shape[2]:
w_exact = _bmm_fp32(tt, w)
torch.baddbmm(target, v, w_exact, beta=1.0, alpha=-1.0, out=target)
return
w_exact = _bmm_fp32(tt, w[:, :, :exact_cols])
w_fast = torch.bmm(tt, w[:, :, exact_cols:])
target_exact = target[:, :, :exact_cols]
torch.baddbmm(target_exact, v, w_exact, beta=1.0, alpha=-1.0, out=target_exact)
target_fast = target[:, :, exact_cols:]
torch.baddbmm(target_fast, v, w_fast, beta=1.0, alpha=-1.0, out=target_fast)
def _apply_wy_update_high_all_baddbmm(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
old_precision = torch.get_float32_matmul_precision()
torch.set_float32_matmul_precision("high")
try:
_apply_wy_update_baddbmm(target, v, t)
finally:
torch.set_float32_matmul_precision(old_precision)
def _apply_wy_update_medium_tf32_t_baddbmm(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
old_precision = torch.get_float32_matmul_precision()
torch.set_float32_matmul_precision("medium")
w = torch.bmm(v.transpose(1, 2), target)
torch.set_float32_matmul_precision("high")
w = torch.bmm(t.transpose(1, 2), w)
torch.set_float32_matmul_precision("medium")
torch.baddbmm(target, v, w, beta=1.0, alpha=-1.0, out=target)
torch.set_float32_matmul_precision(old_precision)
def _apply_wy_update_baddbmm(target: torch.Tensor, v: torch.Tensor, t: torch.Tensor) -> None:
w = torch.bmm(v.transpose(1, 2), target)
w = torch.bmm(t.transpose(1, 2), w)
torch.baddbmm(target, v, w, beta=1.0, alpha=-1.0, out=target)
def _blocked_wy_qr_group2_512(data: torch.Tensor) -> output_t:
batch, n, _ = data.shape
h = data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("high")
block = 16
group = 32
for k in range(0, n, group):
if n - k <= 512:
for kk in range(k, n, group):
rows_tail = n - kk
panel_warps_tail = 4 if rows_tail <= 64 else (8 if rows_tail <= 256 else 16)
tau_tail, v_tail, gram_tail = _factor_superpanel32_tail_gram(
h, tau, kk, panel_warps_tail
)
end_tail = kk + group
if end_tail >= n:
continue
t_tail = _larft_triton_from_gram(gram_tail, tau_tail)
_apply_wy_update(h[:, kk:, end_tail:], v_tail, t_tail)
return h, tau
h_panel1, tau_panel1, v1 = _factor_panel(h, tau, k, block)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
mid = k + block
end = min(k + group, n)
if mid >= n:
break
t1 = _larft16(v1, tau_panel1)
_apply_wy_update_x3(h[:, k:, mid:end], v1, t1)
h_panel2, tau_panel2, v2 = _factor_panel(h, tau, mid, block)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
if end >= n:
continue
rows = n - k
v = torch.empty((batch, rows, group), device=data.device, dtype=data.dtype)
v[:, :, :block] = v1
v[:, :block, block:] = 0.0
v[:, block:, block:] = v2
tau_group = torch.cat((tau_panel1, tau_panel2), dim=1)
t_group = _larft_triton(v, tau_group)
_apply_wy_update_x3(h[:, k:, end:], v, t_group)
return h, tau
def _blocked_wy_qr_group64_512_superpanel_split16(
data: torch.Tensor,
inplace_input: bool = False,
) -> output_t:
if not _HAS_TRITON:
return _blocked_wy_qr_group2_512(data)
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("high")
block = 32
group = 64
loop_start = 0
if n == 512:
k0 = 0
k1 = k0 + block
k2 = k0 + group
k3 = k2 + block
k4 = k2 + group
v1, tau1 = _factor_superpanel32_split16(h, tau, k0)
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_k0_split(h[:, k0:, k1:k2], v1, t1, exact_cols=block)
v2, tau2 = _factor_superpanel32_split16(h, tau, k1)
v64a, t64a = _assemble_v64_tail_gram_compose_t64_from32(
v1,
v2,
t1,
tau2,
full_n=n,
)
_apply_wy_update_k0_split(h[:, k0:, k2:k4], v64a, t64a, exact_cols=0)
v3, tau3 = _factor_superpanel32_split16(h, tau, k2)
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_tfp32_baddbmm(h[:, k2:, k3:k4], v3, t3)
v4, tau4 = _factor_superpanel32_split16(h, tau, k3)
v64b, t64b = _assemble_v64_tail_gram_compose_t64_from32(
v3,
v4,
t3,
tau4,
full_n=n,
)
cross = torch.bmm(v64a[:, group:, :].transpose(1, 2), v64b)
v128 = _assemble_v128_from64(v64a, v64b)
t128 = _compose_t128_from_cross64(t64a, t64b, cross)
_apply_wy_update_k0_split(h[:, k0:, k4:], v128, t128, exact_cols=0)
k5 = k4
k6 = k5 + block
k7 = k5 + group
k8 = k7 + block
k9 = k7 + group
v5, tau5 = _factor_superpanel32_split16(h, tau, k5)
t5 = _larft_triton_current_gram(v5, tau5)
_apply_wy_update_tfp32_baddbmm(h[:, k5:, k6:k7], v5, t5)
v6, tau6 = _factor_superpanel32_split16(h, tau, k6)
v64c, t64c = _assemble_v64_tail_gram_compose_t64_from32(
v5,
v6,
t5,
tau6,
full_n=n,
)
_apply_wy_update_tfp32_baddbmm(h[:, k5:, k7:k9], v64c, t64c)
v7, tau7 = _factor_superpanel32_split16(h, tau, k7)
t7 = _larft_triton_current_gram(v7, tau7)
_apply_wy_update_tfp32_baddbmm(h[:, k7:, k8:k9], v7, t7)
v8, tau8 = _factor_superpanel32_split16(h, tau, k8)
v64d, t64d = _assemble_v64_tail_gram_compose_t64_from32(
v7,
v8,
t7,
tau8,
full_n=n,
)
cross_second = torch.bmm(v64c[:, group:, :].transpose(1, 2), v64d)
v128_second = _assemble_v128_from64(v64c, v64d)
t128_second = _compose_t128_from_cross64(t64c, t64d, cross_second)
_apply_wy_update_tsplit_baddbmm(h[:, k5:, k9:], v128_second, t128_second, exact_cols=128)
loop_start = k9
for k in range(loop_start, n, group):
if n - k <= 64:
for kk in range(k, n, block):
rows_tail = n - kk
panel_warps_tail = 4 if rows_tail <= 64 else (8 if rows_tail <= 256 else 16)
tau_tail, v_tail, gram_tail = _factor_superpanel32_tail_gram(
h, tau, kk, panel_warps_tail
)
end_tail = kk + block
if end_tail >= n:
continue
t_tail = _larft_triton_from_gram(gram_tail, tau_tail)
far_tail = h[:, kk:, end_tail:]
_apply_wy_update_tfp32_baddbmm(far_tail, v_tail, t_tail)
return h, tau
v1, tau1 = _factor_superpanel32_split16(h, tau, k)
mid = k + block
end = k + group
if mid >= n:
break
local = h[:, k:, mid:min(end, n)]
if local.shape[2] > 0:
t1 = _larft_triton_current_gram(v1, tau1)
if k == 0:
_apply_wy_update_k0_split(local, v1, t1, exact_cols=local.shape[2])
else:
_apply_wy_update_tfp32_baddbmm(local, v1, t1)
v2, tau2 = _factor_superpanel32_split16(h, tau, mid)
if end >= n:
continue
v, t_group = _assemble_v64_tail_gram_compose_t64_from32(v1, v2, t1, tau2, full_n=n)
far = h[:, k:, end:]
if k == 0:
_apply_wy_update_k0_split(far, v, t_group, exact_cols=0)
else:
if far.shape[2] >= 128:
_apply_wy_update_baddbmm(far, v, t_group)
else:
_apply_wy_update_tfp32_baddbmm(far, v, t_group)
return h, tau
def _blocked_wy_qr_group128_1024_superpanel(
data: torch.Tensor,
inplace_input: bool = False,
) -> output_t:
if not _HAS_TRITON:
return _blocked_wy_qr(data, 16, trailing_tf32=True)
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("high")
block = 32
half_group = 64
group = 128
loop_start = 0
if n == 1024:
k = 0
v1, tau1 = _factor_superpanel32_split16_fastlocal(h, tau, k)
k2 = k + block
k3 = k + half_group
k4 = k3 + block
end = k + group
local1 = h[:, k:, k2:k3]
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_baddbmm(local1, v1, t1)
v2, tau2 = _factor_superpanel32_split16_fastlocal(h, tau, k2)
t2 = _larft_triton_current_gram(v2, tau2)
v64a, t64a = _assemble_v64_compose_t64_direct_cross32(v1, v2, t1, t2)
local2 = h[:, k:, k3:end]
_apply_wy_update_baddbmm(local2, v64a, t64a)
v3, tau3 = _factor_superpanel32_split16_fastlocal(h, tau, k3)
local3 = h[:, k3:, k4:end]
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_baddbmm(local3, v3, t3)
v4, tau4 = _factor_superpanel32_split16_fastlocal(h, tau, k4)
t4 = _larft_triton_current_gram(v4, tau4)
v64b, t64b = _assemble_v64_compose_t64_direct_cross32(v3, v4, t3, t4)
cross = torch.bmm(v64a[:, half_group:, :].transpose(1, 2), v64b)
v128_left = _assemble_v128_from64(v64a, v64b)
t128_left = _compose_t128_from_cross64(t64a, t64b, cross)
second_start = end
second_end = second_start + group
_apply_wy_update_baddbmm(h[:, k:, second_start:second_end], v128_left, t128_left)
k = second_start
v1, tau1 = _factor_superpanel32_split16_fastlocal(h, tau, k)
k2 = k + block
k3 = k + half_group
k4 = k3 + block
end = k + group
local1 = h[:, k:, k2:k3]
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_baddbmm(local1, v1, t1)
v2, tau2 = _factor_superpanel32_split16_fastlocal(h, tau, k2)
t2 = _larft_triton_current_gram(v2, tau2)
v64a, t64a = _assemble_v64_compose_t64_direct_cross32(v1, v2, t1, t2)
local2 = h[:, k:, k3:end]
_apply_wy_update_baddbmm(local2, v64a, t64a)
v3, tau3 = _factor_superpanel32_split16_fastlocal(h, tau, k3)
local3 = h[:, k3:, k4:end]
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_baddbmm(local3, v3, t3)
v4, tau4 = _factor_superpanel32_split16_fastlocal(h, tau, k4)
t4 = _larft_triton_current_gram(v4, tau4)
v64b, t64b = _assemble_v64_compose_t64_direct_cross32(v3, v4, t3, t4)
cross = torch.bmm(v64a[:, half_group:, :].transpose(1, 2), v64b)
v128_right = _assemble_v128_from64(v64a, v64b)
t128_right = _compose_t128_from_cross64(t64a, t64b, cross)
cross_wide = torch.bmm(v128_left[:, group:, :].transpose(1, 2), v128_right)
v256 = _assemble_v256_from128(v128_left, v128_right)
t256 = _compose_t256_from_cross128(t128_left, t128_right, cross_wide)
_apply_wy_update_baddbmm(h[:, 0:, second_end:], v256, t256)
pair_start = second_end
k = pair_start
v1, tau1 = _factor_superpanel32_split16_fastlocal(h, tau, k)
k2 = k + block
k3 = k + half_group
k4 = k3 + block
end = k + group
local1 = h[:, k:, k2:k3]
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_baddbmm(local1, v1, t1)
v2, tau2 = _factor_superpanel32_split16_fastlocal(h, tau, k2)
t2 = _larft_triton_current_gram(v2, tau2)
v64a, t64a = _assemble_v64_compose_t64_direct_cross32(v1, v2, t1, t2)
local2 = h[:, k:, k3:end]
_apply_wy_update_baddbmm(local2, v64a, t64a)
v3, tau3 = _factor_superpanel32_split16_fastlocal(h, tau, k3)
local3 = h[:, k3:, k4:end]
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_baddbmm(local3, v3, t3)
v4, tau4 = _factor_superpanel32_split16_fastlocal(h, tau, k4)
t4 = _larft_triton_current_gram(v4, tau4)
v64b, t64b = _assemble_v64_compose_t64_direct_cross32(v3, v4, t3, t4)
cross = torch.bmm(v64a[:, half_group:, :].transpose(1, 2), v64b)
v128_left = _assemble_v128_from64(v64a, v64b)
t128_left = _compose_t128_from_cross64(t64a, t64b, cross)
second_start = end
second_end = second_start + group
_apply_wy_update_baddbmm(h[:, k:, second_start:second_end], v128_left, t128_left)
k = second_start
v1, tau1 = _factor_superpanel32_split16_fastlocal(h, tau, k)
k2 = k + block
k3 = k + half_group
k4 = k3 + block
end = k + group
local1 = h[:, k:, k2:k3]
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_baddbmm(local1, v1, t1)
v2, tau2 = _factor_superpanel32_split16_fastlocal(h, tau, k2)
t2 = _larft_triton_current_gram(v2, tau2)
v64a, t64a = _assemble_v64_compose_t64_direct_cross32(v1, v2, t1, t2)
local2 = h[:, k:, k3:end]
_apply_wy_update_baddbmm(local2, v64a, t64a)
v3, tau3 = _factor_superpanel32_split16_fastlocal(h, tau, k3)
local3 = h[:, k3:, k4:end]
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_baddbmm(local3, v3, t3)
v4, tau4 = _factor_superpanel32_split16_fastlocal(h, tau, k4)
t4 = _larft_triton_current_gram(v4, tau4)
v64b, t64b = _assemble_v64_compose_t64_direct_cross32(v3, v4, t3, t4)
cross = torch.bmm(v64a[:, half_group:, :].transpose(1, 2), v64b)
v128_right = _assemble_v128_from64(v64a, v64b)
t128_right = _compose_t128_from_cross64(t64a, t64b, cross)
cross_wide = torch.bmm(v128_left[:, group:, :].transpose(1, 2), v128_right)
v256 = _assemble_v256_from128(v128_left, v128_right)
t256 = _compose_t256_from_cross128(t128_left, t128_right, cross_wide)
_apply_wy_update_baddbmm(h[:, pair_start:, second_end:], v256, t256)
loop_start = second_end
for k in range(loop_start, n, group):
if n - k <= (256 if n == 1024 else 512):
tail_block = 16
for kk in range(k, n, tail_block):
if n == 1024 and n - kk <= 32:
_finish_tail_qr_inplace(h, tau, kk)
break
panel_warps_tail = 4
if kk + tail_block >= n:
_factor_panel16_n1024_tail_warps(
h, tau, kk, emit_v=False, panel_warps=panel_warps_tail
)
break
h_panel, tau_panel, v = _factor_panel16_n1024_tail_warps(
h, tau, kk, emit_v=True, panel_warps=panel_warps_tail
)
if v is None:
v = _panel_reflectors(h_panel)
_apply_split16_trailing_direct_forward(
h, v, tau_panel, kk, block_n=32, num_warps=4, fixed_k=True
)
return h, tau
v1, tau1 = _factor_superpanel32_split16_fastlocal(h, tau, k)
k2 = k + block
k3 = k + half_group
k4 = k3 + block
end = min(k + group, n)
if k2 >= n:
break
local1 = h[:, k:, k2:min(k3, n)]
if local1.shape[2] > 0:
t1 = _larft_triton_current_gram(v1, tau1)
_apply_wy_update_baddbmm(local1, v1, t1)
rows2 = n - k2
panel_warps2 = 32 if rows2 > 512 else 16
v2, tau2 = _factor_superpanel32_split16_fastlocal(h, tau, k2)
t2 = _larft_triton_current_gram(v2, tau2)
v64a, t64a = _assemble_v64_compose_t64_direct_cross32(v1, v2, t1, t2)
if k3 >= n:
continue
local2 = h[:, k:, k3:end]
if local2.shape[2] > 0:
_apply_wy_update_baddbmm(local2, v64a, t64a)
rows3 = n - k3
panel_warps3 = 32 if rows3 > 512 else 16
v3, tau3 = _factor_superpanel32_split16_fastlocal(h, tau, k3)
if k4 >= n:
continue
local3 = h[:, k3:, k4:end]
if local3.shape[2] > 0:
t3 = _larft_triton_current_gram(v3, tau3)
_apply_wy_update_baddbmm(local3, v3, t3)
rows4 = n - k4
panel_warps4 = 32 if rows4 > 512 else 16
v4, tau4 = _factor_superpanel32_split16_fastlocal(h, tau, k4)
if end >= n:
continue
t4 = _larft_triton_current_gram(v4, tau4)
v64b, t64b = _assemble_v64_compose_t64_direct_cross32(v3, v4, t3, t4)
cross = torch.bmm(v64a[:, half_group:, :].transpose(1, 2), v64b)
v128 = _assemble_v128_from64(v64a, v64b)
t128 = _compose_t128_from_cross64(t64a, t64b, cross)
_apply_wy_update_baddbmm(h[:, k:, end:], v128, t128)
return h, tau
def _blocked_wy_qr_group2_2048(data: torch.Tensor, inplace_input: bool = False) -> output_t:
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("medium")
block = 16
group = 32
k = 0
while k < n:
if k < 512 and k + 64 <= n:
k2 = k + block
k3 = k + group
k4 = k3 + block
end64 = k + 64
if _HAS_TRITON and k < 256:
h_panel1, tau_panel1, v1, t1 = _factor_panel_rowsplit_tsqr16_direct_t(
h, tau, k, chunk_rows=256, tail_block=64
)
else:
h_panel1, tau_panel1, v1 = _factor_panel(h, tau, k, block)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
t1 = _larft_triton_high_gram(v1, tau_panel1)
local1 = h[:, k:, k2:k3]
if local1.shape[2] > 0:
_apply_wy_update_baddbmm(local1, v1, t1)
h_panel2, tau_panel2, v2 = _factor_panel(h, tau, k2, block)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
v32a, tau32a = _assemble_vtau32_from16(v1, v2, tau_panel1, tau_panel2)
t32a = _larft_triton_high_gram(v32a, tau32a)
local2 = h[:, k:, k3:end64]
if local2.shape[2] > 0:
_apply_wy_update_baddbmm(local2, v32a, t32a)
if _HAS_TRITON and k < 256:
h_panel3, tau_panel3, v3, t3 = _factor_panel_rowsplit_tsqr16_direct_t(
h, tau, k3, chunk_rows=256, tail_block=64
)
else:
h_panel3, tau_panel3, v3 = _factor_panel(h, tau, k3, block)
if v3 is None:
v3 = _panel_reflectors(h_panel3)
t3 = _larft_triton_high_gram(v3, tau_panel3)
local3 = h[:, k3:, k4:end64]
if local3.shape[2] > 0:
_apply_wy_update_baddbmm(local3, v3, t3)
h_panel4, tau_panel4, v4 = _factor_panel(h, tau, k4, block)
if v4 is None:
v4 = _panel_reflectors(h_panel4)
v32b, tau32b = _assemble_vtau32_from16(v3, v4, tau_panel3, tau_panel4)
v = _assemble_v64_from32(v32a, v32b)
t32b = _larft_triton_high_gram(v32b, tau32b)
cross64 = _with_matmul_precision(
"high",
lambda: torch.bmm(v32a[:, group:, :].transpose(1, 2), v32b),
)
t_group = _compose_t64_from_cross32(t32a, t32b, cross64)
far = h[:, k:, end64:]
if far.shape[2] > 0:
if n - end64 >= 512:
_apply_wy_update_baddbmm(far, v, t_group)
else:
_apply_wy_update(far, v, t_group)
k += 64
continue
rows_tail = n - k
if rows_tail <= 704 and rows_tail >= group:
h_panel1, tau_panel1, v1 = _factor_panel(h, tau, k, block)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
mid = k + block
end = min(k + group, n)
if mid >= n:
break
_apply_split16_trailing_direct_forward(
h, v1, tau_panel1, k, block_n=16, num_warps=4, fixed_k=True
)
if end >= n:
_factor_panel(h, tau, mid, block, emit_v=False)
k += group
continue
h_panel2, tau_panel2, v2 = _factor_panel(h, tau, mid, block)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
_apply_split16_trailing_direct_forward(
h, v2, tau_panel2, mid, block_n=16, num_warps=4, fixed_k=True
)
k += group
continue
h_panel1, tau_panel1, v1 = _factor_panel(h, tau, k, block)
if v1 is None:
v1 = _panel_reflectors(h_panel1)
mid = k + block
end = min(k + group, n)
if mid >= n:
break
t1 = _larft_triton_high_gram(v1, tau_panel1)
local = h[:, k:, mid:end]
_apply_wy_update_baddbmm(local, v1, t1)
if end >= n:
_factor_panel(h, tau, mid, block, emit_v=False)
k += group
continue
h_panel2, tau_panel2, v2 = _factor_panel(h, tau, mid, block)
if v2 is None:
v2 = _panel_reflectors(h_panel2)
v, tau_group = _assemble_vtau32_from16(v1, v2, tau_panel1, tau_panel2)
t_group = _larft_triton_high_gram(v, tau_group)
far = h[:, k:, end:]
if n - end >= 512:
_apply_wy_update_baddbmm(far, v, t_group)
else:
_apply_wy_update(far, v, t_group)
k += group
return h, tau
def _blocked_wy_qr(
data: torch.Tensor,
block: int,
trailing_tf32: bool = True,
split_tf32_trailing: bool = False,
inplace_input: bool = False,
) -> output_t:
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
# "high" = single-pass TF32 trailing (~5e-4); "highest" = full fp32. n512
# mixed needs TF32 (not bf16) to clear the per-matrix gate; n176/n352 use
# full fp32 because their tighter tol can't absorb even TF32 error.
if batch == 8 and n == 2048:
torch.set_float32_matmul_precision("medium")
else:
torch.set_float32_matmul_precision("high" if trailing_tf32 else "highest")
for k in range(0, n, block):
width = min(block, n - k)
h_panel, tau_panel, v = _factor_panel(h, tau, k, width)
if k + width < n:
if v is None:
v = _panel_reflectors(h_panel)
t = _larft16(v, tau_panel) if width == 16 else _larft_forward(v, tau_panel)
if _HAS_TRITON and batch == 40 and width == 16 and n in (176, 352):
if n == 352:
_larfb16_update_splitk352(h, v, t, k)
else:
_larfb16_update(h, v, t, k)
elif _HAS_TRITON and batch == 640 and n == 512 and width == 16 and split_tf32_trailing:
_larfb16_update_x3(h, v, t, k)
else:
trailing = h[:, k:, k + width :]
if split_tf32_trailing:
w = _bmm_3xtf32(v.transpose(1, 2), trailing)
w = _bmm_fp32(t.transpose(1, 2), w)
update = _bmm_3xtf32(v, w)
else:
w = torch.bmm(v.transpose(1, 2), trailing)
w = torch.bmm(t.transpose(1, 2), w)
update = torch.bmm(v, w)
h[:, k:, k + width :] = trailing - update
return h, tau
def _blocked_wy_qr_n352_pair32_high_all(
data: torch.Tensor,
inplace_input: bool = False,
) -> output_t:
if not _HAS_TRITON:
return _blocked_wy_qr(data, 16, trailing_tf32=False, inplace_input=inplace_input)
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("highest")
for k in range(0, n, 32):
v, tau_group = _factor_superpanel32_split16(h, tau, k)
end = k + 32
if end >= n:
continue
t_group = _larft_triton(v, tau_group)
_apply_wy_update_high_all_baddbmm(h[:, k:, end:], v, t_group)
return h, tau
def _blocked_wy_qr_n176_direct_trailing(
data: torch.Tensor,
inplace_input: bool = False,
) -> output_t:
if not _HAS_TRITON:
return _blocked_wy_qr(data, 16, trailing_tf32=False, inplace_input=inplace_input)
batch, n, _ = data.shape
h = data if inplace_input else data.clone()
tau = torch.empty((batch, n), device=data.device, dtype=data.dtype)
torch.set_float32_matmul_precision("highest")
for k in range(0, n, 16):
width = min(16, n - k)
if k + width >= n:
_finish_tail_qr_inplace(h, tau, k)
continue
h_panel, tau_panel, v = _factor_panel(h, tau, k, width)
if v is None:
v = _panel_reflectors(h_panel)
_apply_split16_trailing_direct_forward(h, v, tau_panel, k, block_n=16, num_warps=4)
return h, tau
def custom_kernel(data: input_t) -> output_t:
batch, n, _ = data.shape
if _HAS_TRITON and batch == 20 and n == 32:
return _full_qr(data)
if batch == 640 and n == 512:
return _graphed_inplace_input_h(
"n512_stay_a_while_refrain",
_blocked_wy_qr_group64_512_superpanel_split16,
data,
)
if batch == 60 and n == 1024:
return _graphed_inplace_input_h("sakura", _blocked_wy_qr_group128_1024_superpanel, data)
if batch == 8 and n == 2048:
return _graphed_inplace_input("hana", _blocked_wy_qr_group2_2048, data)
if batch == 2 and n == 4096:
return _graphed_inplace_input(
"n4096_tf32_street_refrain",
_blocked_cholesky_orhr512_taugram_4096_cols_packed_early,
data,
)
block = _BLOCKED_CASES.get((batch, n))
if block is not None:
if batch == 40 and n == 176:
return _graphed_inplace_input(
"n176_yui",
_blocked_wy_qr_n176_direct_trailing,
data,
)
if batch == 40 and n == 352:
return _graphed_inplace_input(
"n352_faint_signal_refrain",
_blocked_wy_qr_n352_pair32_high_all,
data,
)
return _graphed(
f"blocked_{n}",
lambda x: _blocked_wy_qr(x, block, trailing_tf32=False),
data,
)
return torch.geqrf(data)
scrolls · 5211 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