Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
1.80ms
#25 of 515
2026-06-29

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.

mmatmp = tl.dot(t_left, cross, input_precision="tf32", out_dtype=tl.float32)
num-warps = 8num_warps=8 if rows >= 64 else 4,
split-kdef _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