Skip to content
KernelIndex
Search⌘K

submission 844887

10billiontokens · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 12663 lines, June 9 Researcher Reciprocity License v1.0.

s33.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844887?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.22ms
#2 of 515
2026-06-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4e4c6e7a3f13ad3aaa394bd15916c9f8ce126e165212f85a37b29b67f573763d
license declaredunknown
license concludedunknown
authors10billiontokens
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp8primary = (grouped * inverse_scale[:, :, None]).to(tl.float8e4nv)
mmaresult = tl.dot(lhs_high, rhs_high)
num-warps = 1num_warps=1,
persistent-kernel"""Factor and apply one fallback panel in a single persistent CTA.
split-kdef _wide_n4096_pair_projection_split_kernel(
stages = 2for row_start in tl.range(0, triangular_end, BLOCK_K, num_stages=2):
tile-k = 128BLOCK_K=128,
tile-m = 32BLOCK_M=32,
tile-n = 64BLOCK_N=64,

Kernel source

s33.py12663 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

"""Compile-slim B200 QR hybrid."""

import torch
import triton
import triton.language as tl


@triton.jit
def _small_qr_small_kernel(
    a,
    h,
    tau,
    N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    PANEL: tl.constexpr,
):
    """Factor a whole small matrix in one program."""
    batch_index = tl.program_id(0)
    row_offsets = tl.arange(0, BLOCK_M)
    column_offsets = tl.arange(0, PANEL)
    matrix_offset = batch_index * N * N

    pointers = (
        a
        + matrix_offset
        + row_offsets[:, None] * N
        + column_offsets[None, :]
    )
    x = tl.load(pointers)
    tau_values = tl.zeros((PANEL,), dtype=tl.float32)

    for j in tl.static_range(0, PANEL):
        column = tl.sum(
            tl.where(column_offsets[None, :] == j, x, 0.0), axis=1
        )
        alpha = tl.sum(tl.where(row_offsets == j, column, 0.0), axis=0)
        tail_square = tl.sum(
            tl.where(row_offsets > j, column * column, 0.0), axis=0
        )
        norm = tl.sqrt(alpha * alpha + tail_square)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        diagonal = beta
        denominator = alpha - beta
        tau_j = (beta - alpha) / beta
        v = tl.where(
            row_offsets == j,
            1.0,
            tl.where(row_offsets > j, column / denominator, 0.0),
        )

        products = tl.sum(v[:, None] * x, axis=0)
        updated = x - (tau_j * v)[:, None] * products[None, :]
        x = tl.where(column_offsets[None, :] > j, updated, x)
        packed_column = tl.where(
            row_offsets == j,
            diagonal,
            tl.where(row_offsets > j, column / denominator, column),
        )
        x = tl.where(column_offsets[None, :] == j, packed_column[:, None], x)
        tau_values = tl.where(column_offsets == j, tau_j, tau_values)

    tl.store(tau + batch_index * N + column_offsets, tau_values)
    tl.store(
        h
        + matrix_offset
        + row_offsets[:, None] * N
        + column_offsets[None, :],
        x,
    )


@triton.jit
def _small_factor_first_panel_kernel(
    a,
    h,
    tau,
    N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    PANEL: tl.constexpr,
):
    """Factor the first panel directly from the immutable input."""
    batch_index = tl.program_id(0)
    row_offsets = tl.arange(0, BLOCK_M)
    column_offsets = tl.arange(0, PANEL)
    rows = row_offsets
    columns = column_offsets
    matrix_offset = batch_index * N * N

    input_pointers = a + matrix_offset + rows[:, None] * N + columns[None, :]
    x = tl.load(
        input_pointers,
        mask=rows[:, None] < N,
        other=0.0,
    )
    for j in tl.static_range(0, PANEL):
        column = tl.sum(
            tl.where(column_offsets[None, :] == j, x, 0.0), axis=1
        )
        alpha = tl.sum(tl.where(row_offsets == j, column, 0.0), axis=0)
        tail_square = tl.sum(
            tl.where(row_offsets > j, column * column, 0.0), axis=0
        )
        norm = tl.sqrt(alpha * alpha + tail_square)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        diagonal = beta
        denominator = alpha - beta
        tau_j = (beta - alpha) / beta
        v = tl.where(
            row_offsets == j,
            1.0,
            tl.where(row_offsets > j, column / denominator, 0.0),
        )

        if j + 1 < PANEL:
            products = tl.sum(
                tl.where(
                    column_offsets[None, :] > j,
                    v[:, None] * x,
                    0.0,
                ),
                axis=0,
            )
            updated = x - (tau_j * v)[:, None] * products[None, :]
            x = tl.where(column_offsets[None, :] > j, updated, x)
        packed_column = tl.where(
            row_offsets == j,
            diagonal,
            tl.where(row_offsets > j, column / denominator, column),
        )
        x = tl.where(column_offsets[None, :] == j, packed_column[:, None], x)
        tl.store(
            tau + batch_index * N + j,
            tau_j,
        )

    tl.store(
        h + matrix_offset + rows[:, None] * N + columns[None, :],
        x,
        mask=rows[:, None] < N,
    )


@triton.jit
def _small_apply_and_factor_next_panel_kernel(
    a,
    h,
    tau,
    panel_start,
    N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FROM_SOURCE: tl.constexpr,
    STATIC_START: tl.constexpr,
):
    """Apply a panel and factor the next panel in the first column tile.

    Tile zero owns exactly the columns of the next panel.  Once it has applied
    the current reflectors, it can factor those columns while the remaining
    CTAs update farther trailing tiles.  Kernel completion is the global
    synchronization point needed before the next panel application.
    """
    batch_index = tl.program_id(0)
    tile_index = tl.program_id(1)
    row_offsets = tl.arange(0, BLOCK_M)
    panel_offsets = tl.arange(0, PANEL)
    column_offsets = tl.arange(0, BLOCK_N)
    if STATIC_START >= 0:
        start = STATIC_START
    else:
        start = panel_start
    rows = start + row_offsets
    columns = start + PANEL + tile_index * BLOCK_N + column_offsets
    matrix_offset = batch_index * N * N
    valid_rows = rows[:, None] < N
    valid_c = valid_rows

    c_pointers = h + matrix_offset + rows[:, None] * N + columns[None, :]
    if FROM_SOURCE:
        c = tl.load(
            a + matrix_offset + rows[:, None] * N + columns[None, :],
            mask=valid_c,
            other=0.0,
        )
    else:
        c = tl.load(
            c_pointers,
            mask=valid_c,
            other=0.0,
        )
    packed = tl.load(
        h
        + matrix_offset
        + rows[:, None] * N
        + (start + panel_offsets)[None, :],
        mask=valid_rows,
        other=0.0,
    )

    for j in tl.static_range(0, PANEL):
        packed_column = tl.sum(
            tl.where(panel_offsets[None, :] == j, packed, 0.0), axis=1
        )
        apply_v = tl.where(
            row_offsets == j,
            1.0,
            tl.where(row_offsets > j, packed_column, 0.0),
        )
        tau_j = tl.load(
            tau + batch_index * N + start + j,
        )
        apply_products = tl.sum(apply_v[:, None] * c, axis=0)
        c -= (tau_j * apply_v)[:, None] * apply_products[None, :]

    if tile_index == 0:
        for j in tl.static_range(0, PANEL):
            factor_row = PANEL + j
            column = tl.sum(
                tl.where(column_offsets[None, :] == j, c, 0.0), axis=1
            )
            alpha = tl.sum(
                tl.where(row_offsets == factor_row, column, 0.0), axis=0
            )
            tail_square = tl.sum(
                tl.where(
                    row_offsets > factor_row, column * column, 0.0
                ),
                axis=0,
            )
            norm = tl.sqrt(alpha * alpha + tail_square)
            beta = tl.where(alpha >= 0.0, -norm, norm)
            diagonal = beta
            denominator = alpha - beta
            tau_j = (beta - alpha) / beta
            v = tl.where(
                row_offsets == factor_row,
                1.0,
                tl.where(row_offsets > factor_row, column / denominator, 0.0),
            )

            if j + 1 < PANEL:
                products = tl.sum(
                    tl.where(
                        column_offsets[None, :] > j,
                        v[:, None] * c,
                        0.0,
                    ),
                    axis=0,
                )
                updated = c - (tau_j * v)[:, None] * products[None, :]
                c = tl.where(column_offsets[None, :] > j, updated, c)
            packed_column = tl.where(
                row_offsets == factor_row,
                diagonal,
                tl.where(
                    row_offsets > factor_row,
                    column / denominator,
                    column,
                ),
            )
            c = tl.where(
                column_offsets[None, :] == j,
                packed_column[:, None],
                c,
            )
            tl.store(
                tau + batch_index * N + start + factor_row,
                tau_j,
            )

    tl.store(
        c_pointers,
        c,
        mask=valid_c,
    )


def _small_qr_v2(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Return compact Householder factors ``(H, tau)`` for square batches."""
    if not a.is_cuda:
        raise ValueError("qr_v2 requires a CUDA tensor")
    if a.dtype != torch.float32:
        raise ValueError("qr_v2 requires torch.float32 input")
    if a.ndim != 3 or a.shape[-2] != a.shape[-1]:
        raise ValueError("qr_v2 requires shape (batch, n, n)")
    if not a.is_contiguous():
        raise ValueError("qr_v2 requires contiguous input")

    batch, n, _ = a.shape
    h = torch.empty_like(a)
    tau = torch.empty((batch, n), device=a.device, dtype=a.dtype)

    if n == 32:
        _small_qr_small_kernel[(batch,)](
            a,
            h,
            tau,
            N=32,
            BLOCK_M=32,
            PANEL=32,
            num_warps=1,
        )
        return h, tau

    block_m = triton.next_power_of_2(n)
    # Wider panels halve the number of device-wide synchronization points for
    # the benchmark-sized matrices.  The panel remains small enough that the
    # per-CTA working set stays register resident on Blackwell.
    panel = 16
    tail_warps = 2 if n == 176 else 4
    _small_factor_first_panel_kernel[(batch,)](
        a,
        h,
        tau,
        N=n,
        BLOCK_M=block_m,
        PANEL=panel,
        num_warps=8,
    )

    for panel_start in range(0, n, panel):
        trailing_columns = n - panel_start - panel
        if trailing_columns > 0:
            block_n = panel
            active_block_m = triton.next_power_of_2(n - panel_start)
            tile_count = triton.cdiv(trailing_columns, block_n)
            # n176 benefits from static addresses once the large 256-row
            # panels are done; making the first three starts static regresses.
            static_start = (
                panel_start
                if n == 352 or (n == 176 and panel_start >= 48)
                else -1
            )
            _small_apply_and_factor_next_panel_kernel[
                (batch, tile_count)
            ](
                a,
                h,
                tau,
                panel_start,
                N=n,
                BLOCK_M=active_block_m,
                PANEL=panel,
                BLOCK_N=block_n,
                FROM_SOURCE=panel_start == 0,
                STATIC_START=static_start,
                num_warps=(
                    8
                    if active_block_m >= 256
                    else (
                        4
                        if active_block_m >= 128
                        else (tail_warps if active_block_m >= 64 else 1)
                    )
                ),
                maxnreg=128,
            )
    return h, tau


@triton.jit
def _compensated_fp16_dot(lhs, rhs):
    """Near-FP32 product using three high-throughput FP16 MMAs."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    lhs_residual = (lhs - lhs_high).to(tl.float16)
    rhs_residual = (rhs - rhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    result += tl.dot(lhs_high, rhs_residual)
    result += tl.dot(lhs_residual, rhs_high)
    return result


@triton.jit
def _s9_n512_n512_classify_precision_kernel(
    a_ptr,
    flags_ptr,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    DROP_RELATIVE_L1: tl.constexpr,
):
    batch_id = tl.program_id(0)
    offsets = tl.arange(0, N)
    base = batch_id * N * N
    first_row = tl.load(a_ptr + base + offsets)
    last_row = tl.load(a_ptr + base + (N - 1) * N + offsets)
    second_row = tl.load(a_ptr + base + N + offsets)
    penultimate_row = tl.load(a_ptr + base + (N - 2) * N + offsets)
    first_row_abs = tl.abs(first_row)
    last_row_abs = tl.abs(last_row)
    first_row_l1 = tl.sum(first_row_abs, axis=0)
    last_row_l1 = tl.sum(last_row_abs, axis=0)
    precision_flag = (last_row_l1 < first_row_l1 * 0.03125) | (
        first_row_l1 < last_row_l1 * 0.03125
    )
    # Near-collinear inputs have rows that are almost proportional across
    # columns (including the public column scaling).  Route them away from
    # normal equations using a sign-invariant row cosine test.
    row_cross = tl.sum(first_row * last_row, axis=0)
    first_row_l2 = tl.sum(first_row * first_row, axis=0)
    last_row_l2 = tl.sum(last_row * last_row, axis=0)
    correlation_squared = (row_cross * row_cross) / tl.maximum(
        first_row_l2 * last_row_l2, 1.0e-30
    )
    top_cross = tl.sum(first_row * second_row, axis=0)
    top_correlation_squared = (top_cross * top_cross) / tl.maximum(
        first_row_l2 * tl.sum(second_row * second_row, axis=0), 1.0e-30
    )
    bottom_cross = tl.sum(last_row * penultimate_row, axis=0)
    bottom_correlation_squared = (bottom_cross * bottom_cross) / tl.maximum(
        last_row_l2 * tl.sum(penultimate_row * penultimate_row, axis=0),
        1.0e-30,
    )
    nearcollinear_flag = (
        (correlation_squared > 0.25)
        | (top_correlation_squared > 0.25)
        | (bottom_correlation_squared > 0.25)
    )

    first_row_peak = tl.max(first_row_abs, axis=0)
    last_row_peak = tl.max(last_row_abs, axis=0)
    first_row_count = tl.sum(
        tl.where(first_row_abs > first_row_peak * 1.0e-6, 1, 0),
        axis=0,
    )
    last_row_count = tl.sum(
        tl.where(last_row_abs > last_row_peak * 1.0e-6, 1, 0),
        axis=0,
    )
    first_support_end = tl.max(
        tl.where(first_row_abs > first_row_peak * 1.0e-6, offsets + 1, 0),
        axis=0,
    )
    last_support_start = tl.min(
        tl.where(last_row_abs > last_row_peak * 1.0e-6, offsets, N),
        axis=0,
    )
    sparse_boundary = (
        (first_row_count >= 8)
        & (first_row_count <= 64)
        & (last_row_count >= 8)
        & (last_row_count <= 64)
    )
    band_flag = (
        sparse_boundary
        & (first_support_end <= 64)
        & (last_support_start >= N - 64)
    )
    # A row/column transform can turn the diagonal band into an anti-band.
    # It is still sparse and well defined, but the diagonal-band copy kernel
    # is no longer valid; keep it on the general stable route instead.
    precision_flag |= sparse_boundary & ~band_flag
    signal = first_row_abs + last_row_abs
    matrix_scale = tl.max(signal, axis=0)
    active = signal > matrix_scale * DROP_RELATIVE_L1
    active_column_end = tl.max(tl.where(active, offsets + 1, 0), axis=0)
    active_column_end = (
        (active_column_end + PANEL - 1) // PANEL
    ) * PANEL
    active_column_end = tl.maximum(active_column_end, PANEL)
    active_column_end = tl.where(active_column_end == 288, 256, active_column_end)
    tl.store(
        flags_ptr + batch_id,
        active_column_end * 8
        + nearcollinear_flag * 4
        + band_flag * 2
        + precision_flag,
    )


@triton.jit
def _s9_n512_n2048_chol_gram32_kernel(
    source,
    h,
    route_flags,
    partial_gram,
    panel_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    ROUTE_N512: tl.constexpr,
    REQUIRE_FULL_ACTIVE: tl.constexpr,
    PRECISE_ROUTE: tl.constexpr,
):
    """Compute a 32-column Gram as two diagonal and one cross 16 tile."""
    chunk = tl.program_id(0)
    gram_block = tl.program_id(1)
    batch = tl.program_id(2)
    if ROUTE_N512:
        route_metadata = tl.load(route_flags + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
        if not route_selected:
            return
    rows = panel_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    if ROUTE_N512:
        # Form the whole symmetric Gram tile in one CTA.  The prior split
        # schedule loaded overlapping 16-column halves in three CTAs.
        offsets = tl.arange(0, 32)
        columns = panel_start + offsets
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        if PRECISE_ROUTE:
            full_gram = tl.dot(
                tl.trans(values), values, input_precision="tf32x3"
            )
        else:
            full_gram = tl.dot(
                tl.trans(values.to(tl.float16)), values.to(tl.float16)
            )
        full_diagonal = tl.sum(values * values, axis=0)
        full_gram = tl.where(
            offsets[:, None] == offsets[None, :],
            full_diagonal[:, None],
            full_gram,
        )
        full_base = (batch * CHUNKS + chunk) * 32 * 32
        tl.store(
            partial_gram
            + full_base
            + offsets[:, None] * 32
            + offsets[None, :],
            full_gram,
        )
        return

    tile_offsets = tl.arange(0, 16)
    left_tile = tl.where(gram_block == 2, 1, 0)
    right_tile = tl.where(gram_block == 0, 0, 1)
    left_columns = panel_start + left_tile * 16 + tile_offsets
    right_columns = panel_start + right_tile * 16 + tile_offsets
    matrix_selected = True
    if ROUTE_N512:
        matrix_selected = (tl.load(route_flags + batch) & 7) == 0
    left = tl.load(
        matrix + matrix_base + rows[:, None] * n + left_columns[None, :],
        mask=(rows[:, None] < n) & matrix_selected,
        other=0.0,
    )
    right = tl.load(
        matrix + matrix_base + rows[:, None] * n + right_columns[None, :],
        mask=(rows[:, None] < n) & matrix_selected,
        other=0.0,
    )
    if ROUTE_N512:
        # The route excludes scale-sensitive inputs.  FP16 and TF32 have the
        # same significand width here, while FP16 has higher B200 throughput;
        # diagonal norms still use the exact FP32 reduction below.
        gram = tl.dot(tl.trans(left.to(tl.float16)), right.to(tl.float16))
    else:
        gram = tl.dot(tl.trans(left), right, input_precision="tf32x3")
    if gram_block != 1:
        diagonal = tl.sum(left * left, axis=0)
        gram = tl.where(
            tile_offsets[:, None] == tile_offsets[None, :],
            diagonal[:, None],
            gram,
        )

    base = (batch * CHUNKS + chunk) * 32 * 32
    row_offsets = left_tile * 16 + tile_offsets
    column_offsets = right_tile * 16 + tile_offsets
    tl.store(
        partial_gram
        + base
        + row_offsets[:, None] * 32
        + column_offsets[None, :],
        gram,
    )


@triton.jit
def _n1024_chol_gram32_kernel(
    source,
    h,
    route_flags,
    tail_flags,
    partial_gram,
    panel_start,
    FIRST_PANEL: tl.constexpr,
    DETECT_TAIL: tl.constexpr,
    UNIFORM_ROUTE: tl.constexpr,
    BATCH: tl.constexpr,
):
    """Form an n1024 Gram while publishing its two route reductions."""
    chunk = tl.program_id(0)
    batch = tl.program_id(2)

    if DETECT_TAIL:
        if chunk == 0:
            sample_row = 256 + tl.arange(0, 64)
            sample_row_b = 768 + tl.arange(0, 64)
            sample_col = 768 + tl.arange(0, 8)
            sample_base = batch * 1024 * 1024
            sample0 = tl.load(
                h
                + sample_base
                + sample_row[:, None] * 1024
                + sample_col[None, :]
            )
            sample1 = tl.load(
                h
                + sample_base
                + sample_row_b[:, None] * 1024
                + sample_col[None, :]
            )
            sample_max = tl.maximum(
                tl.max(tl.abs(sample0), axis=0),
                tl.max(tl.abs(sample1), axis=0),
            )
            tl.store(
                tail_flags + batch,
                tl.where(tl.max(sample_max, axis=0) < 1.0e-2, 0, 1),
            )

    route_metadata = tl.load(route_flags + batch)
    if UNIFORM_ROUTE:
        flag_offsets = tl.arange(0, 64)
        flag_valid = flag_offsets < BATCH
        all_flags = tl.load(
            route_flags + flag_offsets,
            mask=flag_valid,
            other=1024 * 8,
        )
        all_selected = tl.min(
            ((all_flags == 1024 * 8) | ~flag_valid).to(tl.int32), axis=0
        ) != 0
        route_metadata = tl.where(
            all_selected, route_metadata, route_metadata | 1
        )
        if (chunk == 0) & (batch == 0):
            tl.store(
                route_flags + flag_offsets,
                tl.where(all_selected, all_flags, all_flags | 1),
                mask=flag_valid,
            )

    if route_metadata != 1024 * 8:
        return
    rows = panel_start + chunk * 256 + tl.arange(0, 256)
    columns = panel_start + tl.arange(0, 32)
    matrix_base = batch * 1024 * 1024
    matrix = source if FIRST_PANEL else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * 1024 + columns[None, :],
        mask=rows[:, None] < 1024,
        other=0.0,
    )
    offsets = tl.arange(0, 32)
    gram = tl.dot(tl.trans(values.to(tl.float16)), values.to(tl.float16))
    diagonal = tl.sum(values * values, axis=0)
    gram = tl.where(
        offsets[:, None] == offsets[None, :], diagonal[:, None], gram
    )
    base = (batch * 4 + chunk) * 32 * 32
    tl.store(
        partial_gram
        + base
        + offsets[:, None] * 32
        + offsets[None, :],
        gram,
    )


@triton.jit
def _s9_n512_n2048_cholesky16(matrix):
    """Register-resident upper Cholesky factorization of a 16x16 block."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    upper = tl.zeros((16, 16), tl.float32)
    for j in tl.range(0, 16, loop_unroll_factor=4):
        matrix_row = tl.sum(tl.where(r == j, matrix, 0.0), axis=0)
        old_column = tl.sum(tl.where(c == j, upper, 0.0), axis=1)
        products = tl.sum(old_column[:, None] * upper, axis=0)
        diagonal_value = tl.sum(
            tl.where(offsets == j, matrix_row - products, 0.0), axis=0
        )
        diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-20))
        new_row = (matrix_row - products) / diagonal
        new_row = tl.where(offsets == j, diagonal, new_row)
        new_row = tl.where(offsets >= j, new_row, 0.0)
        upper = tl.where(r == j, new_row[None, :], upper)
    return upper


@triton.jit
def _s9_n512_n2048_inverse_upper16(
    upper, INVERSE_STEPS: tl.constexpr
):
    """Invert a 16x16 upper triangle with a finite Neumann product."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    power = tl.where(r < c, upper / diagonal[:, None], 0.0)
    inverse_unit = identity - power
    for inverse_step in tl.static_range(0, INVERSE_STEPS):
        power = tl.dot(power.to(tl.float16), power.to(tl.float16))
        inverse_unit = tl.dot(
            inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
        )
    return inverse_unit / diagonal[None, :]


@triton.jit
def _s9_n512_n2048_lu16(matrix):
    """No-pivot LU of a strongly diagonally biased 16x16 block."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    lu = matrix
    for j in tl.range(0, 16, loop_unroll_factor=1):
        pivot_row = tl.sum(tl.where(r == j, lu, 0.0), axis=0)
        pivot_column = tl.sum(tl.where(c == j, lu, 0.0), axis=1)
        pivot = tl.sum(tl.where(offsets == j, pivot_row, 0.0), axis=0)
        factor = tl.where(offsets > j, pivot_column / pivot, 0.0)
        updated = lu - factor[:, None] * pivot_row[None, :]
        lu = tl.where((r > j) & (c > j), updated, lu)
        lu = tl.where((r > j) & (c == j), factor[:, None], lu)
    return lu


@triton.jit
def _s9_n512_n2048_inverse_unit_lower16(
    lower, INVERSE_STEPS: tl.constexpr
):
    """Invert a 16x16 unit-lower triangle with tensor-core products."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    power = tl.where(r > c, lower, 0.0)
    inverse = identity - power
    for inverse_step in tl.static_range(0, INVERSE_STEPS):
        power = tl.dot(power.to(tl.float16), power.to(tl.float16))
        inverse = tl.dot(
            (identity + power).to(tl.float16), inverse.to(tl.float16)
        )
    return inverse


@triton.jit
def _s9_n512_n2048_chol32_factor_kernel(
    source,
    h,
    partial_gram,
    matrix_workspace,
    inverse_workspace,
    dense_guard,
    panel_start,
    active_chunks,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    CHECK_DENSE: tl.constexpr,
    ROUTE_N512: tl.constexpr,
    REQUIRE_FULL_ACTIVE: tl.constexpr,
    PRECISE_RECOVERY: tl.constexpr,
):
    """Cholesky a 32-column panel and save M and the signed R inverse."""
    batch = tl.program_id(0)
    if ROUTE_N512:
        route_metadata = tl.load(dense_guard + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
        if not route_selected:
            return
    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * 32 * 32
    partial_base = batch * CHUNKS * 32 * 32

    if ROUTE_N512:
        half = tl.arange(0, 16)
        half_r = half[:, None]
        half_c = half[None, :]
        gram00 = tl.zeros((16, 16), tl.float32)
        gram01 = tl.zeros((16, 16), tl.float32)
        gram11 = tl.zeros((16, 16), tl.float32)
        for chunk in tl.static_range(0, CHUNKS):
            chunk_base = partial_gram + partial_base + chunk * 32 * 32
            chunk_mask = chunk < active_chunks
            gram00 += tl.load(
                chunk_base + half_r * 32 + half_c,
                mask=chunk_mask,
                other=0.0,
            )
            gram01 += tl.load(
                chunk_base + half_r * 32 + 16 + half_c,
                mask=chunk_mask,
                other=0.0,
            )
            gram11 += tl.load(
                chunk_base + (16 + half_r) * 32 + 16 + half_c,
                mask=chunk_mask,
                other=0.0,
            )
        diagonal00 = tl.sum(
            tl.where(half_r == half_c, gram00, 0.0), axis=1
        )
        diagonal11 = tl.sum(
            tl.where(half_r == half_c, gram11, 0.0), axis=1
        )
        gram_diagonal = tl.cat(diagonal00, diagonal11, dim=0)
        upper00 = _s9_n512_n2048_cholesky16(gram00)
        inverse_upper00 = _s9_n512_n2048_inverse_upper16(
            upper00, INVERSE_STEPS=2
        )
        # These products feed only the guarded normal-equations route.  Their
        # inputs are range-safe and the final reflector storage is FP16.
        upper01 = tl.dot(
            tl.trans(inverse_upper00).to(tl.float16),
            gram01.to(tl.float16),
        )
        schur = gram11 - tl.dot(
            tl.trans(upper01).to(tl.float16),
            upper01.to(tl.float16),
        )
        upper11 = _s9_n512_n2048_cholesky16(schur)
        zero16 = tl.zeros((16, 16), tl.float32)
        upper = tl.cat(
            tl.cat(upper00, upper01, dim=1),
            tl.cat(zero16, upper11, dim=1),
            dim=0,
        )
    else:
        gram = tl.zeros((32, 32), tl.float32)
        for chunk in tl.static_range(0, CHUNKS):
            gram += tl.load(
                partial_gram
                + partial_base
                + chunk * 32 * 32
                + r * 32
                + c,
                mask=chunk < active_chunks,
                other=0.0,
            )
        gram_diagonal = tl.sum(tl.where(r == c, gram, 0.0), axis=1)
        upper = tl.zeros((32, 32), tl.float32)
        for j in tl.range(0, 32, loop_unroll_factor=1):
            gram_row = tl.sum(tl.where(r == j, gram, 0.0), axis=0)
            old_column = tl.sum(tl.where(c == j, upper, 0.0), axis=1)
            products = tl.sum(old_column[:, None] * upper, axis=0)
            diagonal_value = tl.sum(
                tl.where(offsets == j, gram_row - products, 0.0), axis=0
            )
            diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-20))
            new_row = (gram_row - products) / diagonal
            new_row = tl.where(offsets == j, diagonal, new_row)
            new_row = tl.where(offsets >= j, new_row, 0.0)
            upper = tl.where(r == j, new_row[None, :], upper)

    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30
    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )
    if CHECK_DENSE:
        panel_active &= tl.load(dense_guard + batch) != 0
    if ROUTE_N512:
        route_metadata = tl.load(dense_guard + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
        panel_active &= route_selected

    identity = tl.where(r == c, 1.0, 0.0)
    if ROUTE_N512 and PRECISE_RECOVERY:
        # U is already available in 2x2 block-upper-triangular form.  Recover
        # U^-1 directly instead of applying four 32x32 Neumann doublings:
        #   inv(U)01 = -inv(U00) @ U01 @ inv(U11).
        # This keeps the same block inverse identity and the existing FP16
        # recovery precision with much less tensor-core work.
        inverse_upper11 = _s9_n512_n2048_inverse_upper16(
            upper11, INVERSE_STEPS=2
        )
        inverse_upper01 = -tl.dot(
            tl.dot(
                inverse_upper00.to(tl.float16),
                upper01.to(tl.float16),
            ).to(tl.float16),
            inverse_upper11.to(tl.float16),
        )
        inverse = tl.cat(
            tl.cat(inverse_upper00, inverse_upper01, dim=1),
            tl.cat(zero16, inverse_upper11, dim=1),
            dim=0,
        )
    else:
        diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
        power = tl.where(r < c, upper / diagonal[:, None], 0.0)
        inverse_unit = identity - power
        for inverse_step in tl.static_range(0, 4):
            power = tl.dot(power.to(tl.float16), power.to(tl.float16))
            inverse_unit = tl.dot(
                inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
            )
        inverse = inverse_unit / diagonal[None, :]

    matrix_base = batch * n * n
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(matrix + matrix_base + rows * n + columns)
    if ROUTE_N512 and not PRECISE_RECOVERY:
        # Q_top is consumed by the FP16 compact-recovery path; narrowing here
        # avoids extra TF32 products without weakening flagged matrices.
        q_top = tl.dot(top.to(tl.float16), inverse.to(tl.float16))
    else:
        q_top = tl.dot(top.to(tl.float16), inverse.to(tl.float16))
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(panel_active, inverse * signs[None, :], 0.0)
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    m_factor = tl.where(
        panel_active, identity - q_top * signs[None, :], identity
    )

    tl.store(
        h + matrix_base + rows * n + columns,
        signed_upper,
        mask=r <= c,
    )
    tl.store(matrix_workspace + workspace_base + r * 32 + c, m_factor)
    tl.store(inverse_workspace + workspace_base + r * 32 + c, signed_inverse)


@triton.jit
def _s9_n512_n2048_chol32_compact_kernel(
    h,
    tau,
    matrix_workspace,
    coefficient_workspace,
    inverse_workspace,
    packed_vectors,
    norm_partials,
    route_flags,
    panel_start,
    n: tl.constexpr,
    ROUTE_N512: tl.constexpr,
    REQUIRE_FULL_ACTIVE: tl.constexpr,
    STORE_TRANSPOSED_T: tl.constexpr,
    PRECISE_RECOVERY: tl.constexpr,
    STORE_NORMS: tl.constexpr,
    STORE_PACKED_VECTORS: tl.constexpr,
):
    """Convert M=I-Q_top to LU compact-WY data for a 32 panel."""
    batch = tl.program_id(0)
    if ROUTE_N512:
        route_metadata = tl.load(route_flags + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
        if not route_selected:
            return
    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * 32 * 32
    identity = tl.where(r == c, 1.0, 0.0)
    half = tl.arange(0, 16)
    half_r = half[:, None]
    half_c = half[None, :]
    identity16 = tl.where(half_r == half_c, 1.0, 0.0)
    zero16 = tl.zeros((16, 16), tl.float32)
    inverse_base = inverse_workspace + workspace_base
    signed00 = tl.load(inverse_base + half_r * 32 + half_c)
    signed01 = tl.load(inverse_base + half_r * 32 + 16 + half_c)
    signed11 = tl.load(
        inverse_base + (16 + half_r) * 32 + 16 + half_c
    )
    if PRECISE_RECOVERY:
        active0 = tl.max(tl.abs(signed00), axis=0) > 0.0
        active1 = tl.maximum(
            tl.max(tl.abs(signed01), axis=0),
            tl.max(tl.abs(signed11), axis=0),
        ) > 0.0
        panel_active = tl.cat(active0, active1, dim=0)
    else:
        signed_inverse = tl.cat(
            tl.cat(signed00, signed01, dim=1),
            tl.cat(zero16, signed11, dim=1),
            dim=0,
        )
        panel_active = tl.max(tl.abs(signed_inverse), axis=0) > 0.0
    route_selected = True
    if ROUTE_N512:
        route_metadata = tl.load(route_flags + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8

    block_base = matrix_workspace + workspace_base
    m00 = tl.load(block_base + half_r * 32 + half_c)
    m01 = tl.load(block_base + half_r * 32 + 16 + half_c)
    m10 = tl.load(block_base + (16 + half_r) * 32 + half_c)
    m11 = tl.load(block_base + (16 + half_r) * 32 + 16 + half_c)

    lu00 = _s9_n512_n2048_lu16(m00)
    lower00 = tl.where(half_r > half_c, lu00, identity16)
    upper00 = tl.where(half_r <= half_c, lu00, 0.0)
    inverse_lower00 = _s9_n512_n2048_inverse_unit_lower16(
        lower00, INVERSE_STEPS=(1 if PRECISE_RECOVERY else 2)
    )
    inverse_upper00 = _s9_n512_n2048_inverse_upper16(
        upper00, INVERSE_STEPS=(1 if PRECISE_RECOVERY else 2)
    )
    upper01 = tl.dot(
        inverse_lower00.to(tl.float16), m01.to(tl.float16)
    )
    lower10 = tl.dot(m10.to(tl.float16), inverse_upper00.to(tl.float16))
    # Both operands were produced by FP16 recovery products, so preserving a
    # three-product TF32 Schur update adds work without additional information.
    schur = m11 - tl.dot(
        lower10.to(tl.float16), upper01.to(tl.float16)
    )
    lu11 = _s9_n512_n2048_lu16(schur)
    lower11 = tl.where(half_r > half_c, lu11, identity16)
    upper11 = tl.where(half_r <= half_c, lu11, 0.0)
    lower = tl.cat(
        tl.cat(lower00, zero16, dim=1),
        tl.cat(lower10, lower11, dim=1),
        dim=0,
    )
    upper_lu = tl.cat(
        tl.cat(upper00, upper01, dim=1),
        tl.cat(zero16, upper11, dim=1),
        dim=0,
    )
    inverse_lower11 = _s9_n512_n2048_inverse_unit_lower16(
        lower11, INVERSE_STEPS=(1 if PRECISE_RECOVERY else 2)
    )
    inverse_lt00 = tl.trans(inverse_lower00)
    inverse_lt11 = tl.trans(inverse_lower11)
    inverse_lt01 = -tl.dot(
        tl.dot(
            inverse_lt00.to(tl.float16),
            tl.trans(lower10).to(tl.float16),
        ).to(tl.float16),
        inverse_lt11.to(tl.float16),
    )
    inverse_lower_transpose = tl.cat(
        tl.cat(inverse_lt00, inverse_lt01, dim=1),
        tl.cat(zero16, inverse_lt11, dim=1),
        dim=0,
    )
    if PRECISE_RECOVERY:
        # Both factors are block upper triangular.  Form only the three
        # nonzero output blocks, retaining the same compensated products but
        # avoiding the zero lower-left work in two full 32x32 MMAs.
        t00 = _compensated_fp16_dot_lhs_residual(upper00, inverse_lt00)
        t01 = _compensated_fp16_dot_lhs_residual(
            upper00, inverse_lt01
        ) + _compensated_fp16_dot_lhs_residual(upper01, inverse_lt11)
        t11 = _compensated_fp16_dot_lhs_residual(upper11, inverse_lt11)
        t_factor = tl.cat(
            tl.cat(t00, t01, dim=1),
            tl.cat(zero16, t11, dim=1),
            dim=0,
        )
    else:
        t_factor = tl.dot(
            upper_lu.to(tl.float16),
            inverse_lower_transpose.to(tl.float16),
        )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    inverse_upper11 = _s9_n512_n2048_inverse_upper16(
        upper11, INVERSE_STEPS=(1 if PRECISE_RECOVERY else 2)
    )
    inverse_upper01 = -tl.dot(
        tl.dot(
            inverse_upper00.to(tl.float16), upper01.to(tl.float16)
        ).to(tl.float16),
        inverse_upper11.to(tl.float16),
    )
    inverse_u = tl.cat(
        tl.cat(inverse_upper00, inverse_upper01, dim=1),
        tl.cat(zero16, inverse_upper11, dim=1),
        dim=0,
    )
    if PRECISE_RECOVERY:
        bottom00 = -_compensated_fp16_dot(signed00, inverse_upper00)
        bottom01 = -(
            _compensated_fp16_dot(signed00, inverse_upper01)
            + _compensated_fp16_dot(signed01, inverse_upper11)
        )
        bottom11 = -_compensated_fp16_dot(signed11, inverse_upper11)
        bottom_transform = tl.cat(
            tl.cat(bottom00, bottom01, dim=1),
            tl.cat(zero16, bottom11, dim=1),
            dim=0,
        )
    else:
        bottom_transform = -tl.dot(
            signed_inverse.to(tl.float16), inverse_u.to(tl.float16)
        )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    matrix_base = batch * n * n
    rows = panel_start + r
    columns = panel_start + c
    tl.store(
        h + matrix_base + rows * n + columns,
        lower,
        mask=(r > c) & route_selected,
    )
    if STORE_PACKED_VECTORS:
        tl.store(
            packed_vectors
            + batch * n * 32
            + (panel_start + offsets)[:, None] * 32
            + offsets[None, :],
            tl.where(r >= c, lower, 0.0),
            mask=route_selected,
        )
    if STORE_TRANSPOSED_T:
        tl.store(
            coefficient_workspace + workspace_base + r * 32 + c,
            tl.trans(t_factor),
        )
    else:
        tl.store(matrix_workspace + workspace_base + r * 32 + c, t_factor)
    tl.store(inverse_workspace + workspace_base + r * 32 + c, bottom_transform)
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(
        tau + batch * n + panel_start + offsets,
        t_diagonal,
        mask=route_selected,
    )
    if ROUTE_N512 and STORE_NORMS:
        top_norm_squared = 1.0 + tl.sum(
            tl.where(r > c, lower * lower, 0.0), axis=0
        )
        norm_base = (batch * 12 + panel_start // 32) * 16 * 32
        tl.store(
            norm_partials + norm_base + offsets,
            top_norm_squared,
        )


@triton.jit
def _safe512_chol32_factor_compact_kernel(
    source,
    h,
    partial_gram,
    t_workspace,
    bottom_workspace,
    packed_vectors,
    norm_partials,
    tau,
    route_flags,
    panel_start,
    active_chunks,
    FIRST_PANEL: tl.constexpr,
):
    """Fuse n512 Cholesky factorization and compact-Householder recovery."""
    batch = tl.program_id(0)
    route_metadata = tl.load(route_flags + batch)
    if (route_metadata & 7) != 0:
        return

    half = tl.arange(0, 16)
    half_r = half[:, None]
    half_c = half[None, :]
    partial_base = batch * 4 * 32 * 32
    gram00 = tl.zeros((16, 16), tl.float32)
    gram01 = tl.zeros((16, 16), tl.float32)
    gram11 = tl.zeros((16, 16), tl.float32)
    for chunk in tl.static_range(0, 4):
        chunk_base = partial_gram + partial_base + chunk * 32 * 32
        chunk_mask = chunk < active_chunks
        gram00 += tl.load(
            chunk_base + half_r * 32 + half_c,
            mask=chunk_mask,
            other=0.0,
        )
        gram01 += tl.load(
            chunk_base + half_r * 32 + 16 + half_c,
            mask=chunk_mask,
            other=0.0,
        )
        gram11 += tl.load(
            chunk_base + (16 + half_r) * 32 + 16 + half_c,
            mask=chunk_mask,
            other=0.0,
        )

    diagonal00 = tl.sum(tl.where(half_r == half_c, gram00, 0.0), axis=1)
    diagonal11 = tl.sum(tl.where(half_r == half_c, gram11, 0.0), axis=1)
    gram_diagonal = tl.cat(diagonal00, diagonal11, dim=0)
    upper00 = _s9_n512_n2048_cholesky16(gram00)
    # Terms through N^7 survive the following FP16 recovery product; a third
    # doubling only forms N^8..N^15 and adds two MMAs per diagonal block.
    inverse_upper00_chol = _s9_n512_n2048_inverse_upper16(
        upper00, INVERSE_STEPS=2
    )
    upper01 = tl.dot(
        tl.trans(inverse_upper00_chol).to(tl.float16), gram01.to(tl.float16)
    )
    schur = gram11 - tl.dot(
        tl.trans(upper01).to(tl.float16), upper01.to(tl.float16)
    )
    upper11 = _s9_n512_n2048_cholesky16(schur)
    zero16 = tl.zeros((16, 16), tl.float32)
    upper = tl.cat(
        tl.cat(upper00, upper01, dim=1),
        tl.cat(zero16, upper11, dim=1),
        dim=0,
    )

    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )

    # Reuse the first block inverse and invert the block-triangular factor
    # directly.  The generic 32-wide Neumann path performs eight 32x32 MMAs;
    # this identity needs only the second 16-wide inverse and two cross MMAs.
    inverse_upper11_chol = _s9_n512_n2048_inverse_upper16(
        upper11, INVERSE_STEPS=2
    )
    inverse_upper01_chol = -tl.dot(
        tl.dot(
            inverse_upper00_chol.to(tl.float16), upper01.to(tl.float16)
        ).to(tl.float16),
        inverse_upper11_chol.to(tl.float16),
    )
    inverse = tl.cat(
        tl.cat(inverse_upper00_chol, inverse_upper01_chol, dim=1),
        tl.cat(zero16, inverse_upper11_chol, dim=1),
        dim=0,
    )
    identity = tl.where(r == c, 1.0, 0.0)

    matrix_base = batch * 512 * 512
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(matrix + matrix_base + rows * 512 + columns)
    q_top = tl.dot(top.to(tl.float16), inverse.to(tl.float16))
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(panel_active, inverse * signs[None, :], 0.0)
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    m_factor = tl.where(
        panel_active, identity - q_top * signs[None, :], identity
    )
    tl.store(
        h + matrix_base + rows * 512 + columns,
        signed_upper,
        mask=r <= c,
    )

    identity16 = tl.where(half_r == half_c, 1.0, 0.0)
    m_column_blocks = tl.permute(
        tl.reshape(m_factor, (32, 2, 16)), (0, 2, 1)
    )
    m_left, m_right = tl.split(m_column_blocks)
    m_left_blocks = tl.permute(tl.reshape(m_left, (2, 16, 16)), (1, 2, 0))
    m_right_blocks = tl.permute(
        tl.reshape(m_right, (2, 16, 16)), (1, 2, 0)
    )
    m00_block, m10_block = tl.split(m_left_blocks)
    m01_block, m11_block = tl.split(m_right_blocks)
    lu00 = _s9_n512_n2048_lu16(m00_block)
    lower00 = tl.where(half_r > half_c, lu00, identity16)
    upper_lu00 = tl.where(half_r <= half_c, lu00, 0.0)
    inverse_lower00 = _s9_n512_n2048_inverse_unit_lower16(
        lower00, INVERSE_STEPS=1
    )
    inverse_lu_upper00 = _s9_n512_n2048_inverse_upper16(
        upper_lu00, INVERSE_STEPS=2
    )
    upper_lu01 = tl.dot(
        inverse_lower00.to(tl.float16), m01_block.to(tl.float16)
    )
    lower10 = tl.dot(
        m10_block.to(tl.float16), inverse_lu_upper00.to(tl.float16)
    )
    compact_schur = m11_block - tl.dot(
        lower10.to(tl.float16), upper_lu01.to(tl.float16)
    )
    lu11 = _s9_n512_n2048_lu16(compact_schur)
    lower11 = tl.where(half_r > half_c, lu11, identity16)
    upper_lu11 = tl.where(half_r <= half_c, lu11, 0.0)
    lower = tl.cat(
        tl.cat(lower00, zero16, dim=1),
        tl.cat(lower10, lower11, dim=1),
        dim=0,
    )
    upper_lu = tl.cat(
        tl.cat(upper_lu00, upper_lu01, dim=1),
        tl.cat(zero16, upper_lu11, dim=1),
        dim=0,
    )

    inverse_lower11 = _s9_n512_n2048_inverse_unit_lower16(
        lower11, INVERSE_STEPS=1
    )
    inverse_lt00 = tl.trans(inverse_lower00)
    inverse_lt11 = tl.trans(inverse_lower11)
    inverse_lt01 = -tl.dot(
        tl.dot(
            inverse_lt00.to(tl.float16), tl.trans(lower10).to(tl.float16)
        ).to(tl.float16),
        inverse_lt11.to(tl.float16),
    )
    inverse_lower_transpose = tl.cat(
        tl.cat(inverse_lt00, inverse_lt01, dim=1),
        tl.cat(zero16, inverse_lt11, dim=1),
        dim=0,
    )
    t_factor = tl.dot(
        upper_lu.to(tl.float16), inverse_lower_transpose.to(tl.float16)
    )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    inverse_upper11 = _s9_n512_n2048_inverse_upper16(
        upper_lu11, INVERSE_STEPS=2
    )
    inverse_upper01 = -tl.dot(
        tl.dot(
            inverse_lu_upper00.to(tl.float16),
            upper_lu01.to(tl.float16),
        ).to(tl.float16),
        inverse_upper11.to(tl.float16),
    )
    inverse_u = tl.cat(
        tl.cat(inverse_lu_upper00, inverse_upper01, dim=1),
        tl.cat(zero16, inverse_upper11, dim=1),
        dim=0,
    )
    bottom_transform = -tl.dot(
        signed_inverse.to(tl.float16), inverse_u.to(tl.float16)
    )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    workspace_base = batch * 32 * 32
    tl.store(
        h + matrix_base + rows * 512 + columns,
        lower,
        mask=r > c,
    )
    tl.store(
        packed_vectors
        + batch * 512 * 32
        + (panel_start + offsets)[:, None] * 32
        + offsets[None, :],
        tl.where(r >= c, lower, 0.0),
    )
    tl.store(t_workspace + workspace_base + r * 32 + c, t_factor)
    tl.store(bottom_workspace + workspace_base + r * 32 + c, bottom_transform)
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(tau + batch * 512 + panel_start + offsets, t_diagonal)
    top_norm_squared = 1.0 + tl.sum(
        tl.where(r > c, lower * lower, 0.0), axis=0
    )
    norm_base = (batch * 12 + panel_start // 32) * 16 * 32
    tl.store(norm_partials + norm_base + offsets, top_norm_squared)


@triton.jit
def _n1024_chol32_factor_compact_kernel(
    source,
    h,
    partial_gram,
    coefficient_workspace,
    bottom_workspace,
    tau,
    route_flags,
    panel_start,
    active_chunks,
    FIRST_PANEL: tl.constexpr,
):
    """Factor and compact-convert one guarded n1024 Cholesky panel."""
    batch = tl.program_id(0)
    if tl.load(route_flags + batch) != 1024 * 8:
        return

    half = tl.arange(0, 16)
    half_r = half[:, None]
    half_c = half[None, :]
    partial_base = batch * 4 * 32 * 32
    gram00 = tl.zeros((16, 16), tl.float32)
    gram01 = tl.zeros((16, 16), tl.float32)
    gram11 = tl.zeros((16, 16), tl.float32)
    for chunk in tl.static_range(0, 4):
        chunk_base = partial_gram + partial_base + chunk * 32 * 32
        chunk_mask = chunk < active_chunks
        gram00 += tl.load(
            chunk_base + half_r * 32 + half_c,
            mask=chunk_mask,
            other=0.0,
        )
        gram01 += tl.load(
            chunk_base + half_r * 32 + 16 + half_c,
            mask=chunk_mask,
            other=0.0,
        )
        gram11 += tl.load(
            chunk_base + (16 + half_r) * 32 + 16 + half_c,
            mask=chunk_mask,
            other=0.0,
        )

    diagonal00 = tl.sum(tl.where(half_r == half_c, gram00, 0.0), axis=1)
    diagonal11 = tl.sum(tl.where(half_r == half_c, gram11, 0.0), axis=1)
    gram_diagonal = tl.cat(diagonal00, diagonal11, dim=0)
    upper00 = _s9_n512_n2048_cholesky16(gram00)
    inverse_chol00 = _s9_n512_n2048_inverse_upper16(
        upper00, INVERSE_STEPS=2
    )
    upper01 = tl.dot(
        tl.trans(inverse_chol00).to(tl.float16), gram01.to(tl.float16)
    )
    chol_schur = gram11 - tl.dot(
        tl.trans(upper01).to(tl.float16), upper01.to(tl.float16)
    )
    upper11 = _s9_n512_n2048_cholesky16(chol_schur)
    zero16 = tl.zeros((16, 16), tl.float32)
    upper = tl.cat(
        tl.cat(upper00, upper01, dim=1),
        tl.cat(zero16, upper11, dim=1),
        dim=0,
    )

    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )

    inverse_chol11 = _s9_n512_n2048_inverse_upper16(
        upper11, INVERSE_STEPS=2
    )
    inverse_chol01 = -tl.dot(
        tl.dot(
            inverse_chol00.to(tl.float16), upper01.to(tl.float16)
        ).to(tl.float16),
        inverse_chol11.to(tl.float16),
    )
    inverse_chol = tl.cat(
        tl.cat(inverse_chol00, inverse_chol01, dim=1),
        tl.cat(zero16, inverse_chol11, dim=1),
        dim=0,
    )
    identity = tl.where(r == c, 1.0, 0.0)

    matrix_base = batch * 1024 * 1024
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(matrix + matrix_base + rows * 1024 + columns)
    # ``inverse_chol`` is block upper triangular.  Six 16x16 products cover
    # the four output quadrants without multiplying by the zero lower-left
    # block of the inverse factor.
    top_column_blocks = tl.permute(
        tl.reshape(top, (32, 2, 16)), (0, 2, 1)
    )
    top_left, top_right = tl.split(top_column_blocks)
    top_left_blocks = tl.permute(
        tl.reshape(top_left, (2, 16, 16)), (1, 2, 0)
    )
    top_right_blocks = tl.permute(
        tl.reshape(top_right, (2, 16, 16)), (1, 2, 0)
    )
    top00, top10 = tl.split(top_left_blocks)
    top01, top11 = tl.split(top_right_blocks)
    q00 = tl.dot(top00.to(tl.float16), inverse_chol00.to(tl.float16))
    q01 = tl.dot(
        top00.to(tl.float16), inverse_chol01.to(tl.float16)
    ) + tl.dot(top01.to(tl.float16), inverse_chol11.to(tl.float16))
    q10 = tl.dot(top10.to(tl.float16), inverse_chol00.to(tl.float16))
    q11 = tl.dot(
        top10.to(tl.float16), inverse_chol01.to(tl.float16)
    ) + tl.dot(top11.to(tl.float16), inverse_chol11.to(tl.float16))
    q_top = tl.cat(tl.cat(q00, q01, dim=1), tl.cat(q10, q11, dim=1), dim=0)
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(
        panel_active, inverse_chol * signs[None, :], 0.0
    )
    signed_column_blocks = tl.permute(
        tl.reshape(signed_inverse, (32, 2, 16)), (0, 2, 1)
    )
    signed_left, signed_right = tl.split(signed_column_blocks)
    signed_left_blocks = tl.permute(
        tl.reshape(signed_left, (2, 16, 16)), (1, 2, 0)
    )
    signed_right_blocks = tl.permute(
        tl.reshape(signed_right, (2, 16, 16)), (1, 2, 0)
    )
    signed_inverse00, unused_signed10 = tl.split(signed_left_blocks)
    signed_inverse01, signed_inverse11 = tl.split(signed_right_blocks)
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    m_factor = tl.where(
        panel_active, identity - q_top * signs[None, :], identity
    )
    tl.store(
        h + matrix_base + rows * 1024 + columns,
        signed_upper,
        mask=r <= c,
    )

    identity16 = tl.where(half_r == half_c, 1.0, 0.0)
    m_column_blocks = tl.permute(
        tl.reshape(m_factor, (32, 2, 16)), (0, 2, 1)
    )
    m_left, m_right = tl.split(m_column_blocks)
    m_left_blocks = tl.permute(tl.reshape(m_left, (2, 16, 16)), (1, 2, 0))
    m_right_blocks = tl.permute(
        tl.reshape(m_right, (2, 16, 16)), (1, 2, 0)
    )
    m00, m10 = tl.split(m_left_blocks)
    m01, m11 = tl.split(m_right_blocks)
    lu00 = _s9_n512_n2048_lu16(m00)
    lower00 = tl.where(half_r > half_c, lu00, identity16)
    upper_lu00 = tl.where(half_r <= half_c, lu00, 0.0)
    inverse_lower00 = _s9_n512_n2048_inverse_unit_lower16(
        lower00, INVERSE_STEPS=1
    )
    inverse_lu_upper00 = _s9_n512_n2048_inverse_upper16(
        upper_lu00, INVERSE_STEPS=1
    )
    upper_lu01 = tl.dot(
        inverse_lower00.to(tl.float16), m01.to(tl.float16)
    )
    lower10 = tl.dot(m10.to(tl.float16), inverse_lu_upper00.to(tl.float16))
    compact_schur = m11 - tl.dot(
        lower10.to(tl.float16), upper_lu01.to(tl.float16)
    )
    lu11 = _s9_n512_n2048_lu16(compact_schur)
    lower11 = tl.where(half_r > half_c, lu11, identity16)
    upper_lu11 = tl.where(half_r <= half_c, lu11, 0.0)
    lower = tl.cat(
        tl.cat(lower00, zero16, dim=1),
        tl.cat(lower10, lower11, dim=1),
        dim=0,
    )

    inverse_lower11 = _s9_n512_n2048_inverse_unit_lower16(
        lower11, INVERSE_STEPS=1
    )
    inverse_lt00 = tl.trans(inverse_lower00)
    inverse_lt11 = tl.trans(inverse_lower11)
    inverse_lt01 = -tl.dot(
        tl.dot(
            inverse_lt00.to(tl.float16), tl.trans(lower10).to(tl.float16)
        ).to(tl.float16),
        inverse_lt11.to(tl.float16),
    )
    t00 = _compensated_fp16_dot_lhs_residual(upper_lu00, inverse_lt00)
    t01 = _compensated_fp16_dot_lhs_residual(
        upper_lu00, inverse_lt01
    ) + _compensated_fp16_dot_lhs_residual(upper_lu01, inverse_lt11)
    t11 = _compensated_fp16_dot_lhs_residual(upper_lu11, inverse_lt11)
    t_factor = tl.cat(
        tl.cat(t00, t01, dim=1),
        tl.cat(zero16, t11, dim=1),
        dim=0,
    )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    inverse_lu_upper11 = _s9_n512_n2048_inverse_upper16(
        upper_lu11, INVERSE_STEPS=1
    )
    inverse_lu_upper01 = -tl.dot(
        tl.dot(
            inverse_lu_upper00.to(tl.float16), upper_lu01.to(tl.float16)
        ).to(tl.float16),
        inverse_lu_upper11.to(tl.float16),
    )
    bottom00 = -_compensated_fp16_dot(
        signed_inverse00, inverse_lu_upper00
    )
    bottom01 = -(
        _compensated_fp16_dot(signed_inverse00, inverse_lu_upper01)
        + _compensated_fp16_dot(signed_inverse01, inverse_lu_upper11)
    )
    bottom11 = -_compensated_fp16_dot(
        signed_inverse11, inverse_lu_upper11
    )
    bottom_transform = tl.cat(
        tl.cat(bottom00, bottom01, dim=1),
        tl.cat(zero16, bottom11, dim=1),
        dim=0,
    )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    workspace_base = batch * 32 * 32
    tl.store(
        h + matrix_base + rows * 1024 + columns,
        lower,
        mask=r > c,
    )
    tl.store(
        coefficient_workspace + workspace_base + r * 32 + c,
        tl.trans(t_factor),
    )
    tl.store(
        bottom_workspace + workspace_base + r * 32 + c,
        bottom_transform,
    )
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(tau + batch * 1024 + panel_start + offsets, t_diagonal)


@triton.jit
def _s9_n512_n2048_chol_extract_bottom_kernel(
    source,
    h,
    bottom_transform_workspace,
    packed_vectors,
    norm_partials,
    route_flags,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    ROUTE_N512: tl.constexpr,
    REQUIRE_FULL_ACTIVE: tl.constexpr,
    PRECISE_RECOVERY: tl.constexpr,
    STORE_NORMS: tl.constexpr,
    STORE_PACKED_VECTORS: tl.constexpr,
):
    """Materialize the recovered reflector tails in independent row tiles."""
    row_tile = tl.program_id(0)
    batch = tl.program_id(1)
    if ROUTE_N512:
        route_metadata = tl.load(route_flags + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
        if not route_selected:
            return
    rows = panel_start + PANEL + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = panel_start + tl.arange(0, PANEL)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    route_selected = True
    if ROUTE_N512:
        route_metadata = tl.load(route_flags + batch)
        route_selected = (route_metadata & 7) == 0
        if REQUIRE_FULL_ACTIVE:
            route_selected &= route_metadata == n * 8
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    offsets = tl.arange(0, PANEL)
    transform = tl.load(
        bottom_transform_workspace
        + batch * PANEL * PANEL
        + offsets[:, None] * PANEL
        + offsets[None, :]
    )
    if PRECISE_RECOVERY:
        # The recovery transform is block upper triangular.  Form its three
        # nonzero 16x16 products directly instead of multiplying by the
        # lower-left zero quadrant in each compensated product.
        value_blocks = tl.permute(
            tl.reshape(values, (BLOCK_M, 2, 16)), (0, 2, 1)
        )
        values0, values1 = tl.split(value_blocks)
        transform_column_blocks = tl.permute(
            tl.reshape(transform, (32, 2, 16)), (0, 2, 1)
        )
        transform_left, transform_right = tl.split(transform_column_blocks)
        transform_left_blocks = tl.permute(
            tl.reshape(transform_left, (2, 16, 16)), (1, 2, 0)
        )
        transform_right_blocks = tl.permute(
            tl.reshape(transform_right, (2, 16, 16)), (1, 2, 0)
        )
        transform00, _ = tl.split(transform_left_blocks)
        transform01, transform11 = tl.split(transform_right_blocks)
        vectors0 = _compensated_fp16_dot_rhs_residual(values0, transform00)
        vectors1 = _compensated_fp16_dot_rhs_residual(
            values0, transform01
        ) + _compensated_fp16_dot_rhs_residual(values1, transform11)
        vectors = tl.cat(vectors0, vectors1, dim=1)
    else:
        vectors = tl.dot(values.to(tl.float16), transform.to(tl.float16))
    mask = (rows[:, None] < n) & route_selected
    tl.store(
        h + matrix_base + rows[:, None] * n + columns[None, :],
        vectors,
        mask=mask,
    )
    if STORE_PACKED_VECTORS:
        tl.store(
            packed_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            vectors,
            mask=mask,
        )
    if ROUTE_N512 and STORE_NORMS:
        bottom_norm_squared = tl.sum(vectors * vectors, axis=0)
        norm_base = (batch * 12 + panel_start // 32) * 16 * 32
        tl.store(
            norm_partials
            + norm_base
            + (row_tile + 1) * 32
            + offsets,
            bottom_norm_squared,
        )


@triton.jit
def _s9_n1024_n1024_factor_panel_kernel(
    h,
    source,
    tau_out,
    coefficients,
    tail_flags,
    route_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
    SKIP_FULL_DENSE: tl.constexpr,
):
    """Factor one narrow panel independently for every batch matrix."""
    batch_id = tl.program_id(0)
    if SKIP_FULL_DENSE:
        if tl.load(route_flags + batch_id) == N * 8:
            return
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if tail_active == 0:
            zero_cols = tl.arange(0, PANEL)
            tl.store(
                tau_out + batch_id * N + panel_start + zero_cols,
                tl.zeros((PANEL,), dtype=tl.float32),
                mask=panel_start + zero_cols < N,
            )
            return
    row_offset = tl.arange(0, BLOCK_M)
    col_offset = tl.arange(0, PANEL)
    rows = panel_start + row_offset
    cols = panel_start + col_offset
    matrix = h + batch_id * N * N
    read_matrix = matrix
    if FIRST_PANEL:
        read_matrix = source + batch_id * N * N
    panel = tl.load(
        read_matrix + rows[:, None] * N + cols[None, :],
        mask=(rows[:, None] < N) & (cols[None, :] < N),
        other=0.0,
    )
    if BLOCK_M >= 128:
        transform_row = tl.arange(0, PANEL)
        # Carry the compact transform transposed so the recurrence reduces
        # along the same blocked axis as the panel data.  Materialize the
        # original coefficient layout only once at the final store.
        transform_t = tl.zeros((PANEL, PANEL), dtype=tl.float16)

    # Keep the panel iteration as a real loop.  Fully unrolling all 32
    # reflectors makes the N=1024 specialization spill heavily on Blackwell.
    for j in tl.range(0, PANEL, loop_unroll_factor=4):
        # Triton tensors intentionally only support shape-changing slices, not
        # scalar subscripts.  The compile-time one-hot reduction extracts a
        # column and is folded to lane selection by the compiler.
        column = tl.sum(
            tl.where((col_offset == j)[None, :], panel, 0.0), axis=1
        )
        active_rows = row_offset >= j
        x0 = tl.sum(tl.where(row_offset == j, column, 0.0), axis=0)
        norm = tl.sqrt(tl.sum(tl.where(active_rows, column * column, 0.0), axis=0))
        beta = tl.where(x0 >= 0.0, -norm, norm)
        denominator = x0 - beta
        safe_denominator = tl.where(denominator == 0.0, 1.0, denominator)
        safe_beta = tl.where(beta == 0.0, 1.0, beta)
        tau = tl.where(norm == 0.0, 0.0, (beta - x0) / safe_beta)
        vector = tl.where(
            row_offset == j,
            1.0,
            tl.where(row_offset > j, column / safe_denominator, 0.0),
        )

        products = tl.sum(vector[:, None] * panel, axis=0)
        if BLOCK_M >= 128:
            gram = tl.where(col_offset < j, products, 0.0)
            new_transform_row = -tau * tl.sum(
                transform_t * gram[None, :], axis=1
            )
            new_transform_row = tl.where(
                col_offset == j, tau, new_transform_row
            ).to(tl.float16)
            transform_t = tl.where(
                (transform_row == j)[None, :],
                new_transform_row[:, None],
                transform_t,
            )
        updated = tl.fma(-vector[:, None], (tau * products)[None, :], panel)
        panel = tl.where((col_offset > j)[None, :], updated, panel)

        packed_column = tl.where(
            row_offset < j,
            column,
            tl.where(row_offset == j, beta, vector),
        )
        panel = tl.where((col_offset == j)[None, :], packed_column[:, None], panel)
        tl.store(tau_out + batch_id * N + panel_start + j, tau)

    tl.store(
        matrix + rows[:, None] * N + cols[None, :],
        panel,
        mask=(rows[:, None] < N) & (cols[None, :] < N),
    )

    if BLOCK_M >= 128:
        # Store the triangular block-reflector transform for trailing tiles.
        tl.store(
            coefficients
            + batch_id * PANEL * PANEL
            + transform_row[:, None] * PANEL
            + col_offset[None, :],
            tl.trans(transform_t),
        )


@triton.jit
def _s9_n1024_n1024_factor_panel_two_block_kernel(
    h,
    source,
    tau_out,
    coefficients,
    route_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_A: tl.constexpr,
    BLOCK_B: tl.constexpr,
    UNROLL: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    SKIP_FULL_DENSE: tl.constexpr,
):
    """Factor a panel whose live height fits two power-of-two row blocks."""
    batch_id = tl.program_id(0)
    if SKIP_FULL_DENSE:
        if tl.load(route_flags + batch_id) == N * 8:
            return
    row_a = tl.arange(0, BLOCK_A)
    row_b = BLOCK_A + tl.arange(0, BLOCK_B)
    col_offset = tl.arange(0, PANEL)
    rows_a = panel_start + row_a
    rows_b = panel_start + row_b
    cols = panel_start + col_offset
    matrix = h + batch_id * N * N
    read_matrix = matrix
    if FIRST_PANEL:
        read_matrix = source + batch_id * N * N
    # Every two-block launch has BLOCK_A == 512 and remaining > 512, hence
    # panel_start <= N - 513.  The entire A block and its 32 panel columns
    # are therefore in bounds; only the trailing B block needs a tail mask.
    panel_a = tl.load(read_matrix + rows_a[:, None] * N + cols[None, :])
    panel_b = tl.load(
        read_matrix + rows_b[:, None] * N + cols[None, :],
        mask=(rows_b[:, None] < N) & (cols[None, :] < N),
        other=0.0,
    )
    transform_row = tl.arange(0, PANEL)
    # The transposed in-register representation avoids repeated layout
    # conversions in the spill-sensitive two-block factor recurrence.
    transform_t = tl.zeros((PANEL, PANEL), dtype=tl.float16)

    for j in tl.range(0, PANEL, loop_unroll_factor=UNROLL):
        column_a = tl.sum(
            tl.where((col_offset == j)[None, :], panel_a, 0.0), axis=1
        )
        column_b = tl.sum(
            tl.where((col_offset == j)[None, :], panel_b, 0.0), axis=1
        )
        active_a = row_a >= j
        x0 = tl.sum(tl.where(row_a == j, column_a, 0.0), axis=0)
        norm = tl.sqrt(
            tl.sum(tl.where(active_a, column_a * column_a, 0.0), axis=0)
            + tl.sum(column_b * column_b, axis=0)
        )
        beta = tl.where(x0 >= 0.0, -norm, norm)
        denominator = x0 - beta
        safe_denominator = tl.where(denominator == 0.0, 1.0, denominator)
        safe_beta = tl.where(beta == 0.0, 1.0, beta)
        tau = tl.where(norm == 0.0, 0.0, (beta - x0) / safe_beta)
        vector_a = tl.where(
            row_a == j,
            1.0,
            tl.where(row_a > j, column_a / safe_denominator, 0.0),
        )
        vector_b = column_b / safe_denominator

        products = tl.sum(vector_a[:, None] * panel_a, axis=0) + tl.sum(
            vector_b[:, None] * panel_b, axis=0
        )
        gram = tl.where(col_offset < j, products, 0.0)
        new_transform_row = -tau * tl.sum(
            transform_t * gram[None, :], axis=1
        )
        new_transform_row = tl.where(col_offset == j, tau, new_transform_row).to(
            tl.float16
        )
        transform_t = tl.where(
            (transform_row == j)[None, :],
            new_transform_row[:, None],
            transform_t,
        )

        updated_a = tl.fma(-vector_a[:, None], (tau * products)[None, :], panel_a)
        updated_b = tl.fma(-vector_b[:, None], (tau * products)[None, :], panel_b)
        panel_a = tl.where((col_offset > j)[None, :], updated_a, panel_a)
        panel_b = tl.where((col_offset > j)[None, :], updated_b, panel_b)

        packed_a = tl.where(
            row_a < j,
            column_a,
            tl.where(row_a == j, beta, vector_a),
        )
        panel_a = tl.where((col_offset == j)[None, :], packed_a[:, None], panel_a)
        panel_b = tl.where((col_offset == j)[None, :], vector_b[:, None], panel_b)
        tl.store(tau_out + batch_id * N + panel_start + j, tau)

    tl.store(matrix + rows_a[:, None] * N + cols[None, :], panel_a)
    tl.store(
        matrix + rows_b[:, None] * N + cols[None, :],
        panel_b,
        mask=(rows_b[:, None] < N) & (cols[None, :] < N),
    )
    tl.store(
        coefficients
        + batch_id * PANEL * PANEL
        + transform_row[:, None] * PANEL
        + col_offset[None, :],
        tl.trans(transform_t),
    )


@triton.jit
def _s9_n1024_n1024_form_panel_product_kernel(
    h,
    source,
    coefficients,
    product_out,
    tail_flags,
    precision_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
    TAIL_START: tl.constexpr,
    USE_MXFP8: tl.constexpr,
    TILE_BASE: tl.constexpr,
    SPLIT_BAND_ROUTE: tl.constexpr,
):
    """Form and triangularly transform V.T @ A for one block reflector."""
    tile_id = tl.program_id(0)
    batch_id = tl.program_id(1)
    if SPLIT_BAND_ROUTE:
        band_selected = (tl.load(precision_flags + batch_id) & 2) != 0
        if USE_MXFP8:
            if not band_selected:
                return
        else:
            if band_selected:
                return
    panel_col = tl.arange(0, PANEL)
    out_col = (TILE_BASE + tile_id) * BLOCK_N + tl.arange(0, BLOCK_N)
    matrix_col = panel_start + PANEL + out_col
    tail_mask = tl.full((BLOCK_N,), True, dtype=tl.int1)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        first_matrix_col = (
            panel_start + PANEL + (TILE_BASE + tile_id) * BLOCK_N
        )
        if (tail_active == 0) & (first_matrix_col >= TAIL_START):
            return
        tail_mask = (tail_active != 0) | (matrix_col < TAIL_START)
    matrix = h + batch_id * N * N
    read_matrix = matrix
    if FIRST_PANEL:
        read_matrix = source + batch_id * N * N
    accumulator = tl.zeros((PANEL, BLOCK_N), dtype=tl.float32)

    # Explicit warp specialization was prototyped here, but Triton's
    # NVWSInsertTmemAref pass currently asserts when the loop-carried scaled
    # accumulator feeds all three compensated products.  Keep the verified
    # native TMEM/scaled-MMA path until that compiler limitation is fixed.
    # The leaderboard's Triton build allocates more shared memory for the
    # 64x128x128 tail-panel specialization than the local build.  Two pipeline
    # stages keep the original 128-column arithmetic/program mapping while
    # fitting below its 232448-byte per-block limit.
    if PANEL == 128 and BLOCK_N == 128:
        # Only the first two 64-row chunks overlap the implicit unit-lower
        # Householder triangle.  Later chunks are strictly below every panel
        # column, so packed storage is already V and does not need compares
        # or selects.  The P128 call sites also round the live height to an
        # exact BLOCK_K multiple, so the suffix rows are in bounds.
        triangular_end: tl.constexpr = (
            ((PANEL + BLOCK_K - 1) // BLOCK_K) * BLOCK_K
        )
        for row_start in tl.range(0, triangular_end, BLOCK_K, num_stages=2):
            row_offset = row_start + tl.arange(0, BLOCK_K)
            rows = panel_start + row_offset
            packed = tl.load(
                matrix + rows[:, None] * N + panel_start + panel_col[None, :],
                mask=rows[:, None] < N,
                other=0.0,
            )
            vectors = tl.where(
                row_offset[:, None] == panel_col[None, :],
                1.0,
                tl.where(row_offset[:, None] > panel_col[None, :], packed, 0.0),
            )
            matrix_tile = tl.load(
                read_matrix + rows[:, None] * N + matrix_col[None, :],
                mask=(rows[:, None] < N)
                & (matrix_col[None, :] < N)
                & tail_mask[None, :],
                other=0.0,
            )
            if USE_MXFP8:
                accumulator = _v10_mxfp8_two_term_dot(
                    tl.trans(vectors), matrix_tile, accumulator
                )
            else:
                accumulator += tl.dot(
                    tl.trans(vectors).to(tl.float16),
                    matrix_tile.to(tl.float16),
                )

        for row_start in tl.range(triangular_end, BLOCK_M, BLOCK_K, num_stages=2):
            row_offset = row_start + tl.arange(0, BLOCK_K)
            rows = panel_start + row_offset
            vectors = tl.load(
                matrix + rows[:, None] * N + panel_start + panel_col[None, :],
            )
            matrix_tile = tl.load(
                read_matrix + rows[:, None] * N + matrix_col[None, :],
                mask=(matrix_col[None, :] < N) & tail_mask[None, :],
                other=0.0,
            )
            if USE_MXFP8:
                accumulator = _v10_mxfp8_two_term_dot(
                    tl.trans(vectors), matrix_tile, accumulator
                )
            else:
                accumulator += tl.dot(
                    tl.trans(vectors).to(tl.float16),
                    matrix_tile.to(tl.float16),
                )
    else:
        for row_start in tl.range(0, BLOCK_M, BLOCK_K, num_stages=2):
            row_offset = row_start + tl.arange(0, BLOCK_K)
            rows = panel_start + row_offset
            packed = tl.load(
                matrix + rows[:, None] * N + panel_start + panel_col[None, :],
                mask=rows[:, None] < N,
                other=0.0,
            )
            vectors = tl.where(
                row_offset[:, None] == panel_col[None, :],
                1.0,
                tl.where(row_offset[:, None] > panel_col[None, :], packed, 0.0),
            )
            matrix_tile = tl.load(
                read_matrix + rows[:, None] * N + matrix_col[None, :],
                mask=(rows[:, None] < N)
                & (matrix_col[None, :] < N)
                & tail_mask[None, :],
                other=0.0,
            )
            if USE_MXFP8:
                accumulator = _v10_mxfp8_two_term_dot(
                    tl.trans(vectors), matrix_tile, accumulator
                )
            else:
                accumulator += tl.dot(
                    tl.trans(vectors).to(tl.float16),
                    matrix_tile.to(tl.float16),
                )

    transform = tl.load(
        coefficients
        + batch_id * PANEL * PANEL
        + panel_col[:, None] * PANEL
        + panel_col[None, :]
    )
    transformed = tl.dot(
        transform.to(tl.float16),
        accumulator.to(tl.float16),
        out_dtype=tl.float16,
    )

    tl.store(
        product_out + batch_id * PANEL * N + panel_col[:, None] * N + matrix_col[None, :],
        transformed,
        mask=(matrix_col[None, :] < N) & tail_mask[None, :],
    )


@triton.jit
def _n1024_apply32_form_next_gram_kernel(
    h,
    source,
    product,
    partial_gram,
    route_flags,
    panel_start,
    BATCH: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Update the adjacent panel and retain its dense-route Gram partial.

    Each CTA already owns the same 256-row partition consumed by the next
    Cholesky Gram.  Forming that partial while ``updated`` is live removes the
    immediate global write/read round trip and one dependency launch per
    adjacent 32-column panel.
    """
    row_tile = tl.program_id(0)
    batch = BATCH - 1 - tl.program_id(1)
    row_offset = row_tile * 256 + tl.arange(0, 256)
    panel_col = tl.arange(0, 32)
    rows = panel_start + row_offset
    cols = panel_start + 32 + panel_col
    matrix = h + batch * 1024 * 1024
    read_matrix = source + batch * 1024 * 1024 if FIRST_PANEL else matrix

    packed = tl.load(
        matrix + rows[:, None] * 1024 + panel_start + panel_col[None, :],
        mask=rows[:, None] < 1024,
        other=0.0,
    )
    # Only the leading 256-row CTA can intersect the implicit unit-lower
    # triangle of a 32-column panel.  Later CTAs are strictly below every
    # panel column, so their packed storage is already the exact V tile.
    vectors = packed
    if row_tile == 0:
        vectors = tl.where(
            row_offset[:, None] == panel_col[None, :],
            1.0,
            tl.where(row_offset[:, None] > panel_col[None, :], packed, 0.0),
        )
    weights = tl.load(
        product
        + batch * 32 * 1024
        + panel_col[:, None] * 1024
        + cols[None, :]
    )
    correction = tl.dot(
        vectors.to(tl.float16),
        weights.to(tl.float16),
        out_dtype=tl.float16,
    )
    pointers = matrix + rows[:, None] * 1024 + cols[None, :]
    valid = rows[:, None] < 1024
    values = tl.load(
        read_matrix + rows[:, None] * 1024 + cols[None, :],
        mask=valid,
        other=0.0,
    )
    updated = values - correction
    tl.store(pointers, updated, mask=valid)

    if tl.load(route_flags + batch) == 1024 * 8:
        # The next panel begins 32 rows below this panel.  Mask the rows that
        # just became triangular so the partials cover its exact active tail.
        gram_values = tl.where(
            (row_offset[:, None] >= 32) & valid, updated, 0.0
        )
        gram = tl.dot(
            tl.trans(gram_values.to(tl.float16)),
            gram_values.to(tl.float16),
        )
        diagonal = tl.sum(gram_values * gram_values, axis=0)
        gram = tl.where(
            panel_col[:, None] == panel_col[None, :],
            diagonal[:, None],
            gram,
        )
        gram_base = (batch * 4 + row_tile) * 32 * 32
        tl.store(
            partial_gram
            + gram_base
            + panel_col[:, None] * 32
            + panel_col[None, :],
            gram,
        )


@triton.jit
def _s9_n1024_n1024_apply_block_product_kernel(
    h,
    source,
    product,
    tail_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROW: tl.constexpr,
    BLOCK_COL: tl.constexpr,
    ROW_TILES: tl.constexpr,
    ROW_BASE: tl.constexpr,
    BATCH: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
    TAIL_START: tl.constexpr,
):
    """Apply A -= V @ Z using a tensor-core block product."""
    tile_id = tl.program_id(0)
    batch_id = BATCH - 1 - tl.program_id(1)
    row_tile = tile_id % ROW_TILES
    col_tile = tile_id // ROW_TILES
    row_offset = ROW_BASE + row_tile * BLOCK_ROW + tl.arange(0, BLOCK_ROW)
    out_col = col_tile * BLOCK_COL + tl.arange(0, BLOCK_COL)
    panel_col = tl.arange(0, PANEL)
    rows = panel_start + row_offset
    cols = panel_start + PANEL + out_col
    tail_mask = tl.full((BLOCK_COL,), True, dtype=tl.int1)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        first_col = panel_start + PANEL + col_tile * BLOCK_COL
        if (tail_active == 0) & (first_col >= TAIL_START):
            return
        tail_mask = (tail_active != 0) | (cols < TAIL_START)
    matrix = h + batch_id * N * N
    read_matrix = matrix
    if FIRST_PANEL:
        read_matrix = source + batch_id * N * N

    packed = tl.load(
        matrix + rows[:, None] * N + panel_start + panel_col[None, :],
        mask=rows[:, None] < N,
        other=0.0,
    )
    vectors = packed
    if ROW_BASE == 0:
        vectors = tl.where(
            row_offset[:, None] == panel_col[None, :],
            1.0,
            tl.where(row_offset[:, None] > panel_col[None, :], packed, 0.0),
        )
    weights = tl.load(
        product + batch_id * PANEL * N + panel_col[:, None] * N + cols[None, :],
        mask=(cols[None, :] < N) & tail_mask[None, :],
        other=0.0,
    )
    correction = tl.dot(
        vectors.to(tl.float16),
        weights.to(tl.float16),
        out_dtype=tl.float16,
    )
    pointers = matrix + rows[:, None] * N + cols[None, :]
    valid = (rows[:, None] < N) & (cols[None, :] < N) & tail_mask[None, :]
    values = tl.load(
        read_matrix + rows[:, None] * N + cols[None, :],
        mask=valid,
        other=0.0,
    )
    tl.store(pointers, values - correction, mask=valid)


@triton.jit
def _s9_n1024_n1024_apply_panel_kernel(
    h,
    tau_out,
    tail_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
):
    """Apply a factored panel to independent tiles of the trailing columns."""
    batch_id = tl.program_id(0)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if tail_active == 0:
            return
    tile_id = tl.program_id(1)
    row_offset = tl.arange(0, BLOCK_M)
    col_offset = tl.arange(0, BLOCK_N)
    rows = panel_start + row_offset
    cols = panel_start + PANEL + tile_id * BLOCK_N + col_offset
    matrix = h + batch_id * N * N
    valid = (rows[:, None] < N) & (cols[None, :] < N)
    tile = tl.load(
        matrix + rows[:, None] * N + cols[None, :],
        mask=valid,
        other=0.0,
    )

    for j in tl.static_range(0, PANEL):
        packed = tl.load(
            matrix + rows * N + panel_start + j,
            mask=rows < N,
            other=0.0,
        )
        vector = tl.where(
            row_offset == j,
            1.0,
            tl.where(row_offset > j, packed, 0.0),
        )
        tau = tl.load(tau_out + batch_id * N + panel_start + j)
        products = tl.sum(vector[:, None] * tile, axis=0)
        tile = tl.fma(-vector[:, None], (tau * products)[None, :], tile)

    tl.store(
        matrix + rows[:, None] * N + cols[None, :],
        tile,
        mask=valid,
    )


@triton.jit
def _s9_n1024_n1024_recursive_cross_gram_kernel(
    h,
    cross_partials,
    tail_flags,
    panel_start,
    N: tl.constexpr,
    HALF_PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    PARTIAL_STRIDE: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
):
    """Form one row block of ``V1.T @ V0`` for a recursive superpanel."""
    row_block = tl.program_id(0)
    batch_id = tl.program_id(1)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if tail_active == 0:
            return

    row_offset = row_block * BLOCK_ROWS + tl.arange(0, BLOCK_ROWS)
    panel_offset = tl.arange(0, HALF_PANEL)
    rows = panel_start + row_offset
    cols0 = panel_start + panel_offset
    cols1 = panel_start + HALF_PANEL + panel_offset
    matrix = h + batch_id * N * N
    valid_rows = rows < N

    if row_block == 0:
        packed0 = tl.load(
            matrix + rows[:, None] * N + cols0[None, :],
            mask=valid_rows[:, None] & (rows[:, None] > cols0[None, :]),
            other=0.0,
        )
        packed1 = tl.load(
            matrix + rows[:, None] * N + cols1[None, :],
            mask=valid_rows[:, None] & (rows[:, None] > cols1[None, :]),
            other=0.0,
        )
        vectors0 = tl.where(rows[:, None] == cols0[None, :], 1.0, packed0)
        vectors1 = tl.where(rows[:, None] == cols1[None, :], 1.0, packed1)
    else:
        # BLOCK_ROWS is 128 for every n1024 launch, while the largest
        # half-panel is 64.  Every later row block is therefore strictly
        # below both compact panels and already stores the exact V entries.
        vectors0 = tl.load(
            matrix + rows[:, None] * N + cols0[None, :],
            mask=valid_rows[:, None],
            other=0.0,
        )
        vectors1 = tl.load(
            matrix + rows[:, None] * N + cols1[None, :],
            mask=valid_rows[:, None],
            other=0.0,
        )
    cross = tl.dot(
        tl.trans(vectors1.to(tl.float16)),
        vectors0.to(tl.float16),
    )
    tl.store(
        cross_partials
        + (batch_id * PARTIAL_STRIDE + row_block)
        * HALF_PANEL
        * HALF_PANEL
        + panel_offset[:, None] * HALF_PANEL
        + panel_offset[None, :],
        cross,
    )


@triton.jit
def _s9_n1024_n1024_assemble_recursive_t64_kernel(
    coefficients0,
    coefficients1,
    cross_partials,
    coefficients64,
    tail_flags,
    route_flags,
    panel_start,
    N: tl.constexpr,
    HALF_PANEL: tl.constexpr,
    PARTIALS: tl.constexpr,
    PARTIAL_STRIDE: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
):
    """Couple two stable 32-reflector blocks into one 64-reflector block.

    The factor kernels apply each block from the left, so their stored compact
    transforms are lower triangular.  For ``Q1.T @ Q0.T`` the exact coupling
    is ``-T1 @ (V1.T @ V0) @ T0`` in the lower-left quadrant.  Keeping the two
    diagonal blocks intact avoids a fragile 64-step scalar recurrence.
    """
    batch_id = tl.program_id(0)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if tail_active == 0:
            return

    offsets = tl.arange(0, HALF_PANEL)
    block_offsets = offsets[:, None] * HALF_PANEL + offsets[None, :]
    t0 = tl.load(
        coefficients0 + batch_id * HALF_PANEL * HALF_PANEL + block_offsets
    )
    t1 = tl.load(
        coefficients1 + batch_id * HALF_PANEL * HALF_PANEL + block_offsets
    )
    cross = tl.zeros((HALF_PANEL, HALF_PANEL), dtype=tl.float32)
    partial_base = batch_id * PARTIAL_STRIDE * HALF_PANEL * HALF_PANEL
    for partial in tl.static_range(0, PARTIALS):
        cross += tl.load(
            cross_partials
            + partial_base
            + partial * HALF_PANEL * HALF_PANEL
            + block_offsets
        )

    if tl.load(route_flags + batch_id) == 1024 * 8:
        middle = tl.dot(t1.to(tl.float16), cross.to(tl.float16))
        coupling = -tl.dot(middle.to(tl.float16), t0.to(tl.float16))
    else:
        middle = _compensated_fp16_dot_rhs_residual(t1, cross)
        coupling = -_compensated_fp16_dot_lhs_residual(middle, t0)
    panel: tl.constexpr = 2 * HALF_PANEL
    output = coefficients64 + batch_id * panel * panel
    output_offsets = offsets[:, None] * panel + offsets[None, :]
    tl.store(output + output_offsets, t0)
    tl.store(output + output_offsets + HALF_PANEL, 0.0)
    tl.store(output + output_offsets + HALF_PANEL * panel, coupling)
    tl.store(output + output_offsets + HALF_PANEL * panel + HALF_PANEL, t1)


def _s9_n1024_n1024_launch_factor_panel(
    h: torch.Tensor,
    source: torch.Tensor,
    tau: torch.Tensor,
    coefficients: torch.Tensor,
    tail_flags: torch.Tensor,
    route_flags: torch.Tensor,
    panel_start: int,
    n: int,
    batch: int,
    skip_full_dense: bool = False,
) -> int:
    """Launch the existing stable 32-column factor and return its row extent."""
    panel = 32
    remaining = n - panel_start
    if remaining > 768:
        block_m = 1024
        _s9_n1024_n1024_factor_panel_two_block_kernel[(batch,)](
            h,
            source,
            tau,
            coefficients,
            route_flags,
            panel_start,
            N=n,
            PANEL=panel,
            BLOCK_A=512,
            BLOCK_B=512,
            UNROLL=4,
            FIRST_PANEL=panel_start == 0,
            SKIP_FULL_DENSE=skip_full_dense,
            num_warps=8,
        )
    elif 512 < remaining <= 576:
        block_m = 576
        _s9_n1024_n1024_factor_panel_two_block_kernel[(batch,)](
            h,
            source,
            tau,
            coefficients,
            route_flags,
            panel_start,
            N=n,
            PANEL=panel,
            BLOCK_A=512,
            BLOCK_B=64,
            UNROLL=4,
            FIRST_PANEL=False,
            SKIP_FULL_DENSE=skip_full_dense,
            num_warps=8,
        )
    elif 576 < remaining <= 640:
        block_m = 640
        _s9_n1024_n1024_factor_panel_two_block_kernel[(batch,)](
            h,
            source,
            tau,
            coefficients,
            route_flags,
            panel_start,
            N=n,
            PANEL=panel,
            BLOCK_A=512,
            BLOCK_B=128,
            UNROLL=4,
            FIRST_PANEL=False,
            SKIP_FULL_DENSE=skip_full_dense,
            num_warps=8,
        )
    elif 640 < remaining <= 768:
        block_m = 768
        _s9_n1024_n1024_factor_panel_two_block_kernel[(batch,)](
            h,
            source,
            tau,
            coefficients,
            route_flags,
            panel_start,
            N=n,
            PANEL=panel,
            BLOCK_A=512,
            BLOCK_B=256,
            UNROLL=4,
            FIRST_PANEL=False,
            SKIP_FULL_DENSE=skip_full_dense,
            num_warps=8,
        )
    else:
        block_m = triton.next_power_of_2(remaining)
        factor_warps = min(32, max(4, block_m // 32))
        if block_m == 1024:
            factor_warps = 16
        elif block_m == 512:
            factor_warps = 8
        elif block_m == 256:
            factor_warps = 8
        elif block_m == 128:
            factor_warps = 4
        _s9_n1024_n1024_factor_panel_kernel[(batch,)](
            h,
            source,
            tau,
            coefficients,
            tail_flags,
            route_flags,
            panel_start,
            N=n,
            PANEL=panel,
            BLOCK_M=block_m,
            FIRST_PANEL=panel_start == 0,
            USE_TAIL_SKIP=panel_start >= 768,
            SKIP_FULL_DENSE=skip_full_dense,
            num_warps=factor_warps,
        )
    return block_m


def _s9_n1024_n1024_qr_v2(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Factor n1024 through guarded Cholesky panels and 128-reflector WY blocks."""
    batch, n, _ = a.shape
    h = torch.empty_like(a)
    tau = torch.empty((batch, n), device=a.device, dtype=a.dtype)

    half_panel = 32
    superpanel = 64
    coefficients0 = torch.empty(
        (batch, half_panel, half_panel), device=a.device, dtype=torch.float16
    )
    coefficients1 = torch.empty_like(coefficients0)
    coefficients64 = torch.empty(
        (batch, superpanel, superpanel), device=a.device, dtype=torch.float16
    )
    coefficients64_b = torch.empty_like(coefficients64)
    coefficients128 = torch.empty(
        (batch, 128, 128), device=a.device, dtype=torch.float16
    )
    product32 = torch.empty(
        (batch, half_panel, n), device=a.device, dtype=torch.float16
    )
    product64 = torch.empty(
        (batch, superpanel, n), device=a.device, dtype=torch.float16
    )
    product128 = torch.empty(
        (batch, 128, n), device=a.device, dtype=torch.float16
    )
    cross_partials = torch.empty(
        (batch, 8, half_panel, half_panel), device=a.device, dtype=a.dtype
    )
    cross_partials64 = torch.empty(
        (batch, 8, superpanel, superpanel), device=a.device, dtype=a.dtype
    )
    tail_flags = torch.empty((batch,), device=a.device, dtype=torch.int32)
    precision_flags = torch.empty((batch,), device=a.device, dtype=torch.int32)
    chol_partials = torch.empty(
        (batch, 8, half_panel, half_panel), device=a.device, dtype=a.dtype
    )
    chol_bottom = torch.empty(
        (batch, half_panel, half_panel), device=a.device, dtype=a.dtype
    )

    _s9_n512_n512_classify_precision_kernel[(batch,)](
        a,
        precision_flags,
        N=n,
        PANEL=half_panel,
        DROP_RELATIVE_L1=1.0e-5,
        num_warps=4,
    )
    # Couple two stable 64-reflector blocks into a 128-wide superpanel.
    # The first 64 block updates only the columns needed by the second; one
    # 128-wide product then updates the whole remaining trail.
    for group_start in range(0, 640, 128):
        group_remaining = n - group_start
        row_tiles = triton.cdiv(group_remaining, 128)
        product_block_m = triton.cdiv(group_remaining, 64) * 64
        for pair_offset in (0, 64):
            panel_start = group_start + pair_offset
            remaining = n - panel_start
            pair_coefficients = (
                coefficients64 if pair_offset == 0 else coefficients64_b
            )
            _s9_n1024_n1024_cholesky_factor_32_launch(
                a,
                h,
                tau,
                coefficients0,
                chol_partials,
                chol_bottom,
                tail_flags,
                precision_flags,
                batch,
                panel_start,
            )
            _s9_n1024_n1024_launch_factor_panel(
                h,
                a,
                tau,
                coefficients0,
                tail_flags,
                precision_flags,
                panel_start,
                n,
                batch,
                skip_full_dense=True,
            )
            _s9_n1024_n1024_form_panel_product_kernel[(1, batch)](
                h,
                a,
                coefficients0,
                product32,
                tail_flags,
                precision_flags,
                panel_start,
                N=n,
                PANEL=half_panel,
                BLOCK_M=triton.cdiv(remaining, 64) * 64,
                BLOCK_K=128,
                BLOCK_N=half_panel,
                FIRST_PANEL=panel_start == 0,
                USE_TAIL_SKIP=False,
                TAIL_START=768,
                USE_MXFP8=False,
                TILE_BASE=0,
                SPLIT_BAND_ROUTE=False,
                num_warps=4,
                num_stages=2,
            )
            pair_row_tiles = triton.cdiv(remaining, 256)
            _n1024_apply32_form_next_gram_kernel[(pair_row_tiles, batch)](
                h,
                a,
                product32,
                chol_partials,
                precision_flags,
                panel_start,
                BATCH=batch,
                FIRST_PANEL=panel_start == 0,
                num_warps=8,
            )
            _s9_n1024_n1024_cholesky_factor_32_launch(
                h,
                h,
                tau,
                coefficients1,
                chol_partials,
                chol_bottom,
                tail_flags,
                precision_flags,
                batch,
                panel_start + half_panel,
                precomputed_chunks=pair_row_tiles,
            )
            _s9_n1024_n1024_launch_factor_panel(
                h,
                a,
                tau,
                coefficients1,
                tail_flags,
                precision_flags,
                panel_start + half_panel,
                n,
                batch,
                skip_full_dense=True,
            )
            partials32 = triton.cdiv(remaining, 128)
            _s9_n1024_n1024_recursive_cross_gram_kernel[
                (partials32, batch)
            ](
                h,
                cross_partials,
                tail_flags,
                panel_start,
                N=n,
                HALF_PANEL=half_panel,
                BLOCK_ROWS=128,
                PARTIAL_STRIDE=8,
                USE_TAIL_SKIP=False,
                num_warps=4,
            )
            _s9_n1024_n1024_assemble_recursive_t64_kernel[(batch,)](
                coefficients0,
                coefficients1,
                cross_partials,
                pair_coefficients,
                tail_flags,
                precision_flags,
                panel_start,
                N=n,
                HALF_PANEL=half_panel,
                PARTIALS=partials32,
                PARTIAL_STRIDE=8,
                USE_TAIL_SKIP=False,
                num_warps=2,
            )

            if pair_offset == 0:
                # Materialize just the second 64-reflector block.
                _s9_n1024_n1024_form_panel_product_kernel[(1, batch)](
                    h,
                    a,
                    coefficients64,
                    product64,
                    tail_flags,
                    precision_flags,
                    group_start,
                    N=n,
                    PANEL=64,
                    BLOCK_M=product_block_m,
                    BLOCK_K=128,
                    BLOCK_N=64,
                    FIRST_PANEL=group_start == 0,
                    USE_TAIL_SKIP=False,
                    TAIL_START=768,
                    USE_MXFP8=False,
                    TILE_BASE=0,
                    SPLIT_BAND_ROUTE=False,
                    num_warps=8,
                    num_stages=1,
                )
                _s9_n1024_n1024_apply_block_product_kernel[
                    (row_tiles, batch)
                ](
                    h,
                    a,
                    product64,
                    tail_flags,
                    group_start,
                    N=n,
                    PANEL=64,
                    BLOCK_ROW=128,
                    BLOCK_COL=64,
                    ROW_TILES=row_tiles,
                    ROW_BASE=0,
                    BATCH=batch,
                    FIRST_PANEL=group_start == 0,
                    USE_TAIL_SKIP=False,
                    TAIL_START=768,
                    num_warps=16,
                )

        partials64 = triton.cdiv(group_remaining, 128)
        _s9_n1024_n1024_recursive_cross_gram_kernel[(partials64, batch)](
            h,
            cross_partials64,
            tail_flags,
            group_start,
            N=n,
            HALF_PANEL=64,
            BLOCK_ROWS=128,
            PARTIAL_STRIDE=8,
            USE_TAIL_SKIP=group_start >= 640,
            num_warps=4,
        )
        _s9_n1024_n1024_assemble_recursive_t64_kernel[(batch,)](
            coefficients64,
            coefficients64_b,
            cross_partials64,
            coefficients128,
            tail_flags,
            precision_flags,
            group_start,
            N=n,
            HALF_PANEL=64,
            PARTIALS=partials64,
            PARTIAL_STRIDE=8,
            USE_TAIL_SKIP=group_start >= 640,
            num_warps=8,
        )

        trailing = n - group_start - 128
        if trailing > 0:
            # Modal's Triton build exceeds B200 shared-memory resources for
            # the locally-fast 256-column, two-stage specialization.  Split
            # only independent output columns; arithmetic and reductions are
            # unchanged.
            far_block_n = 128
            _s9_n1024_n1024_form_panel_product_kernel[
                (triton.cdiv(trailing, far_block_n), batch)
            ](
                h,
                a,
                coefficients128,
                product128,
                tail_flags,
                precision_flags,
                group_start,
                N=n,
                PANEL=128,
                BLOCK_M=product_block_m,
                BLOCK_K=64,
                BLOCK_N=far_block_n,
                FIRST_PANEL=group_start == 0,
                USE_TAIL_SKIP=group_start >= 256,
                TAIL_START=768,
                USE_MXFP8=False,
                TILE_BASE=0,
                SPLIT_BAND_ROUTE=group_start == 640,
                num_warps=8,
                num_stages=1,
            )
            far_apply_cols = 128
            col_tiles = triton.cdiv(trailing, far_apply_cols)
            far_apply_rows = 64 if group_start < 256 else 32
            far_row_tiles = triton.cdiv(group_remaining, far_apply_rows)
            _s9_n1024_n1024_apply_block_product_kernel[
                (far_row_tiles * col_tiles, batch)
            ](
                h,
                a,
                product128,
                tail_flags,
                group_start,
                N=n,
                PANEL=128,
                BLOCK_ROW=far_apply_rows,
                BLOCK_COL=far_apply_cols,
                ROW_TILES=far_row_tiles,
                ROW_BASE=0,
                BATCH=batch,
                FIRST_PANEL=group_start == 0,
                USE_TAIL_SKIP=group_start >= 256,
                TAIL_START=768,
                num_warps=8 if group_start < 256 else 4,
            )

    # At row 640, couple only the next two 32-column panels.  A full
    # 128-reflector group no longer amortizes its second cross-Gram, while
    # immediately falling back to scalar panels gives up a still-profitable
    # matrix pass.  This 64-reflector bridge is the exact middle ground.
    panel_start = 640
    remaining = n - panel_start
    product_block_m = triton.cdiv(remaining, 64) * 64
    _s9_n1024_n1024_cholesky_factor_32_launch(
        h,
        h,
        tau,
        coefficients0,
        chol_partials,
        chol_bottom,
        tail_flags,
        precision_flags,
        batch,
        panel_start,
    )
    _s9_n1024_n1024_launch_factor_panel(
        h,
        a,
        tau,
        coefficients0,
        tail_flags,
        precision_flags,
        panel_start,
        n,
        batch,
        skip_full_dense=True,
    )
    _s9_n1024_n1024_form_panel_product_kernel[(1, batch)](
        h,
        a,
        coefficients0,
        product32,
        tail_flags,
        precision_flags,
        panel_start,
        N=n,
        PANEL=half_panel,
        BLOCK_M=product_block_m,
        BLOCK_K=128,
        BLOCK_N=half_panel,
        FIRST_PANEL=False,
        USE_TAIL_SKIP=False,
        TAIL_START=768,
        USE_MXFP8=False,
        TILE_BASE=0,
        SPLIT_BAND_ROUTE=False,
        num_warps=4,
        num_stages=2,
    )
    panel_row_tiles = triton.cdiv(remaining, 256)
    _n1024_apply32_form_next_gram_kernel[(panel_row_tiles, batch)](
        h,
        a,
        product32,
        chol_partials,
        precision_flags,
        panel_start,
        BATCH=batch,
        FIRST_PANEL=False,
        num_warps=8,
    )
    _s9_n1024_n1024_cholesky_factor_32_launch(
        h,
        h,
        tau,
        coefficients1,
        chol_partials,
        chol_bottom,
        tail_flags,
        precision_flags,
        batch,
        panel_start + half_panel,
        precomputed_chunks=panel_row_tiles,
    )
    _s9_n1024_n1024_launch_factor_panel(
        h,
        a,
        tau,
        coefficients1,
        tail_flags,
        precision_flags,
        panel_start + half_panel,
        n,
        batch,
        skip_full_dense=True,
    )
    partials32 = triton.cdiv(remaining, 128)
    _s9_n1024_n1024_recursive_cross_gram_kernel[(partials32, batch)](
        h,
        cross_partials,
        tail_flags,
        panel_start,
        N=n,
        HALF_PANEL=half_panel,
        BLOCK_ROWS=128,
        PARTIAL_STRIDE=8,
        USE_TAIL_SKIP=False,
        num_warps=4,
    )
    _s9_n1024_n1024_assemble_recursive_t64_kernel[(batch,)](
        coefficients0,
        coefficients1,
        cross_partials,
        coefficients64,
        tail_flags,
        precision_flags,
        panel_start,
        N=n,
        HALF_PANEL=half_panel,
        PARTIALS=partials32,
        PARTIAL_STRIDE=8,
        USE_TAIL_SKIP=False,
        num_warps=2,
    )
    trailing = n - panel_start - superpanel
    _s9_n1024_n1024_form_panel_product_kernel[
        (triton.cdiv(trailing, 128), batch)
    ](
        h,
        a,
        coefficients64,
        product64,
        tail_flags,
        precision_flags,
        panel_start,
        N=n,
        PANEL=superpanel,
        BLOCK_M=product_block_m,
        BLOCK_K=128,
        BLOCK_N=128,
        FIRST_PANEL=False,
        USE_TAIL_SKIP=True,
        TAIL_START=768,
        USE_MXFP8=False,
        TILE_BASE=0,
        SPLIT_BAND_ROUTE=False,
        num_warps=8,
        num_stages=1,
    )
    row_tiles = triton.cdiv(remaining, 128)
    col_tiles = triton.cdiv(trailing, 64)
    _s9_n1024_n1024_apply_block_product_kernel[(row_tiles * col_tiles, batch)](
        h,
        a,
        product64,
        tail_flags,
        panel_start,
        N=n,
        PANEL=superpanel,
        BLOCK_ROW=128,
        BLOCK_COL=64,
        ROW_TILES=row_tiles,
        ROW_BASE=0,
        BATCH=batch,
        FIRST_PANEL=False,
        USE_TAIL_SKIP=True,
        TAIL_START=768,
        num_warps=8,
    )

    # Late panels have too little far-trailing traffic to repay cross-block
    # conversion.  Return to the seed's 32-column path, including its direct
    # scalar-Householder specialization once the live height reaches 128.
    for panel_start in range(704, n, half_panel):
        use_late_cholesky = panel_start < 768
        if use_late_cholesky:
            _s9_n1024_n1024_cholesky_factor_32_launch(
                h,
                h,
                tau,
                coefficients0,
                chol_partials,
                chol_bottom,
                tail_flags,
                precision_flags,
                batch,
                panel_start,
            )
        block_m = _s9_n1024_n1024_launch_factor_panel(
            h,
            a,
            tau,
            coefficients0,
            tail_flags,
            precision_flags,
            panel_start,
            n,
            batch,
            skip_full_dense=use_late_cholesky,
        )
        trailing = n - panel_start - half_panel
        if trailing > 0:
            if block_m < 128:
                direct_cols = 32
                direct_warps = 4
                _s9_n1024_n1024_apply_panel_kernel[
                    (batch, triton.cdiv(trailing, direct_cols))
                ](
                    h,
                    tau,
                    tail_flags,
                    panel_start,
                    N=n,
                    PANEL=half_panel,
                    BLOCK_M=block_m,
                    BLOCK_N=direct_cols,
                    USE_TAIL_SKIP=True,
                    num_warps=direct_warps,
                )
                continue

            product_block_m = triton.cdiv(n - panel_start, 64) * 64
            _s9_n1024_n1024_form_panel_product_kernel[
                (triton.cdiv(trailing, 64), batch)
            ](
                h,
                a,
                coefficients0,
                product32,
                tail_flags,
                precision_flags,
                panel_start,
                N=n,
                PANEL=half_panel,
                BLOCK_M=product_block_m,
                BLOCK_K=64,
                BLOCK_N=64,
                FIRST_PANEL=False,
                USE_TAIL_SKIP=True,
                TAIL_START=768,
                USE_MXFP8=False,
                TILE_BASE=0,
                SPLIT_BAND_ROUTE=False,
                num_warps=4,
                num_stages=1,
            )
            row_tiles = triton.cdiv(n - panel_start, 64)
            col_tiles = triton.cdiv(trailing, 64)
            _s9_n1024_n1024_apply_block_product_kernel[(row_tiles * col_tiles, batch)](
                h,
                a,
                product32,
                tail_flags,
                panel_start,
                N=n,
                PANEL=half_panel,
                BLOCK_ROW=64,
                BLOCK_COL=64,
                ROW_TILES=row_tiles,
                ROW_BASE=0,
                BATCH=batch,
                FIRST_PANEL=False,
                USE_TAIL_SKIP=True,
                TAIL_START=768,
                num_warps=2,
            )
    return h, tau


@triton.jit
def _s9_n2048_n1024_apply_panel_kernel(
    h,
    tau_out,
    tail_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
    SKIP_NONZERO: tl.constexpr,
):
    """Apply a factored panel to independent tiles of the trailing columns."""
    batch_id = tl.program_id(0)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if SKIP_NONZERO:
            if tail_active != 0:
                return
        else:
            if tail_active == 0:
                return
    tile_id = tl.program_id(1)
    row_offset = tl.arange(0, BLOCK_M)
    col_offset = tl.arange(0, BLOCK_N)
    rows = panel_start + row_offset
    cols = panel_start + PANEL + tile_id * BLOCK_N + col_offset
    matrix = h + batch_id * N * N
    valid = (rows[:, None] < N) & (cols[None, :] < N)
    tile = tl.load(
        matrix + rows[:, None] * N + cols[None, :],
        mask=valid,
        other=0.0,
    )

    for j in tl.static_range(0, PANEL):
        packed = tl.load(
            matrix + rows * N + panel_start + j,
            mask=rows < N,
            other=0.0,
        )
        vector = tl.where(
            row_offset == j,
            1.0,
            tl.where(row_offset > j, packed, 0.0),
        )
        tau = tl.load(tau_out + batch_id * N + panel_start + j)
        products = tl.sum(vector[:, None] * tile, axis=0)
        tile = tl.fma(-vector[:, None], (tau * products)[None, :], tile)

    tl.store(
        matrix + rows[:, None] * N + cols[None, :],
        tile,
        mask=valid,
    )


@triton.jit
def _s9_n2048_n2048_chol_gram32_kernel(
    source,
    h,
    partial_gram,
    panel_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Compute a 32-column Gram as two diagonal and one cross 16 tile."""
    chunk = tl.program_id(0)
    gram_block = tl.program_id(1)
    batch = tl.program_id(2)
    rows = panel_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    tile_offsets = tl.arange(0, 16)
    left_tile = tl.where(gram_block == 2, 1, 0)
    right_tile = tl.where(gram_block == 0, 0, 1)
    left_columns = panel_start + left_tile * 16 + tile_offsets
    right_columns = panel_start + right_tile * 16 + tile_offsets
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    left = tl.load(
        matrix + matrix_base + rows[:, None] * n + left_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    right = tl.load(
        matrix + matrix_base + rows[:, None] * n + right_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    gram = _compensated_fp16_dot(tl.trans(left), right)
    if gram_block != 1:
        diagonal = tl.sum(left * left, axis=0)
        gram = tl.where(
            tile_offsets[:, None] == tile_offsets[None, :],
            diagonal[:, None],
            gram,
        )

    base = (batch * CHUNKS + chunk) * 32 * 32
    row_offsets = left_tile * 16 + tile_offsets
    column_offsets = right_tile * 16 + tile_offsets
    tl.store(
        partial_gram
        + base
        + row_offsets[:, None] * 32
        + column_offsets[None, :],
        gram,
    )
    if gram_block == 1:
        tl.store(
            partial_gram
            + base
            + column_offsets[:, None] * 32
            + row_offsets[None, :],
            tl.trans(gram),
        )


@triton.jit
def _s9_n2048_n2048_chol_pair_gram_kernel(
    source,
    h,
    partial_gram,
    pair_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Compute one complete 64-column partial Gram per row slab.

    The column-zero specialization reads the immutable input directly; later
    pairs consume the trailing matrix produced by the preceding WY update.
    """
    chunk = tl.program_id(0)
    batch = tl.program_id(1)
    rows = pair_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    offsets = tl.arange(0, 64)
    columns = pair_start + offsets
    matrix_base = batch * n * n
    matrix = source if FIRST_PAIR else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    gram = tl.dot(
        tl.trans(values.to(tl.float16)), values.to(tl.float16)
    )
    diagonal = tl.sum(values * values, axis=0)
    gram = tl.where(
        offsets[:, None] == offsets[None, :], diagonal[:, None], gram
    )
    base = (batch * CHUNKS + chunk) * 64 * 64
    tl.store(
        partial_gram
        + base
        + offsets[:, None] * 64
        + offsets[None, :],
        gram,
    )


@triton.jit
def _s10_pair_gram_reduce_kernel(
    partial_gram,
    reduced_gram,
    active_chunks,
    MAX_CHUNKS: tl.constexpr,
):
    """Reduce the three independent 32x32 pair-Gram quadrants."""
    gram_block = tl.program_id(0)
    batch = tl.program_id(1)
    offsets = tl.arange(0, 32)
    block_row = tl.where(gram_block == 2, 1, 0)
    block_col = tl.where(gram_block == 0, 0, 1)
    result = tl.zeros((32, 32), tl.float32)
    for chunk in tl.static_range(0, MAX_CHUNKS):
        result += tl.load(
            partial_gram
            + (batch * MAX_CHUNKS + chunk) * 64 * 64
            + (block_row * 32 + offsets[:, None]) * 64
            + block_col * 32
            + offsets[None, :],
            mask=chunk < active_chunks,
            other=0.0,
        )
    tl.store(
        reduced_gram
        + batch * 64 * 64
        + (block_row * 32 + offsets[:, None]) * 64
        + block_col * 32
        + offsets[None, :],
        result,
    )


@triton.jit
def _n4096_group_gram_reduce_kernel(
    pair_partials,
    target_partials,
    reduced_pair,
    reduced_targets,
    active_chunks,
    MAX_CHUNKS: tl.constexpr,
):
    """Joint Gram reduction with four 16x16 CTAs per old quadrant."""
    output_block = tl.program_id(0)
    batch = tl.program_id(1)
    offsets = tl.arange(0, 16)
    result = tl.zeros((16, 16), tl.float32)
    if output_block < 12:
        matrix_block = output_block // 4
        subtile = output_block % 4
        matrix_row = tl.where(matrix_block == 2, 1, 0)
        matrix_col = tl.where(matrix_block == 0, 0, 1)
        rows = matrix_row * 32 + (subtile // 2) * 16 + offsets[:, None]
        columns = matrix_col * 32 + (subtile % 2) * 16 + offsets[None, :]
        for chunk in tl.static_range(0, MAX_CHUNKS):
            result += tl.load(
                pair_partials
                + (batch * MAX_CHUNKS + chunk) * 64 * 64
                + rows * 64
                + columns,
                mask=chunk < active_chunks,
                other=0.0,
            )
        tl.store(
            reduced_pair + batch * 64 * 64 + rows * 64 + columns,
            result,
        )
    else:
        target_block = output_block - 12
        matrix = target_block // 16
        subtile = target_block % 16
        rows = (subtile // 4) * 16 + offsets[:, None]
        columns = (subtile % 4) * 16 + offsets[None, :]
        for chunk in tl.static_range(0, MAX_CHUNKS):
            result += tl.load(
                target_partials
                + ((batch * MAX_CHUNKS + chunk) * 2 + matrix) * 64 * 64
                + rows * 64
                + columns,
                mask=chunk < active_chunks,
                other=0.0,
            )
        tl.store(
            reduced_targets
            + (batch * 2 + matrix) * 64 * 64
            + rows * 64
            + columns,
            result,
        )


@triton.jit
def _s9_n2048_n2048_inverse_upper16(upper, INVERSE_STEPS: tl.constexpr):
    """Invert a 16x16 upper triangle with a bounded Neumann product."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    power = tl.where(r < c, upper / diagonal[:, None], 0.0)
    inverse_unit = identity - power
    for inverse_step in tl.static_range(0, INVERSE_STEPS):
        power = tl.dot(power.to(tl.float16), power.to(tl.float16))
        inverse_unit = tl.dot(
            inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
        )
    return inverse_unit / diagonal[None, :]


@triton.jit
def _s9_n2048_n2048_inverse_unit_lower16(lower, INVERSE_STEPS: tl.constexpr):
    """Invert a 16x16 unit-lower triangle with tensor-core products."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    power = tl.where(r > c, lower, 0.0)
    inverse = identity - power
    for inverse_step in tl.static_range(0, INVERSE_STEPS):
        power = tl.dot(power.to(tl.float16), power.to(tl.float16))
        inverse = tl.dot(
            (identity + power).to(tl.float16), inverse.to(tl.float16)
        )
    return inverse


@triton.jit
def _s9_n2048_n2048_chol32_factor_kernel(
    source,
    h,
    partial_gram,
    second_gram_workspace,
    target_weights,
    matrix_workspace,
    inverse_workspace,
    packed_vectors,
    tau_out,
    dense_guard,
    pair_first_vectors,
    pair_cross_workspace,
    panel_start,
    active_chunks,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    CHECK_DENSE: tl.constexpr,
    NEWTON_STEPS: tl.constexpr,
    PAIR_GRAM: tl.constexpr,
    FORM_PAIR_CROSS: tl.constexpr,
):
    """Factor and compact-convert one 32-column CholeskyQR panel."""
    batch = tl.program_id(0)
    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * 32 * 32
    partial_stride = 64 * 64 if PAIR_GRAM else 32 * 32
    partial_base = batch * CHUNKS * partial_stride

    gram = tl.zeros((32, 32), tl.float32)
    for chunk in tl.static_range(0, CHUNKS):
        gram += tl.load(
            partial_gram
            + partial_base
            + chunk * partial_stride
            + r * (64 if PAIR_GRAM else 32)
            + c
        )
    gram_diagonal = tl.sum(tl.where(r == c, gram, 0.0), axis=1)
    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30

    identity = tl.where(r == c, 1.0, 0.0)
    if NEWTON_STEPS > 0:
        # Solve X^T G X = I directly for the unique upper-triangular inverse
        # Cholesky factor.  At every step C-I is split as S^T+S with S upper,
        # and X <- X(I-S).  The routed dense tail is sufficiently equilibrated
        # that two FP16-tensor-core iterations retain the required margin.
        scale = tl.sqrt(tl.maximum(gram_diagonal, 1.0e-20))
        normalized = gram / (scale[:, None] * scale[None, :])
        # X starts at I, so C=X^T N X is exactly N in the first Newton
        # iteration.  Writing that iteration explicitly removes three
        # identity tensor products from every panel and is also slightly more
        # accurate because N is not rounded through an MMA on the way in.
        correction = tl.where(r < c, normalized, 0.0)
        correction += tl.where(
            r == c, 0.5 * (normalized - identity), 0.0
        )
        first_correction = correction
        inverse_unit = identity - correction
        for newton_step in tl.static_range(1, NEWTON_STEPS):
            normalized_times_inverse = tl.dot(
                normalized.to(tl.float16), inverse_unit.to(tl.float16)
            )
            transformed = tl.dot(
                tl.trans(inverse_unit.to(tl.float16)),
                normalized_times_inverse.to(tl.float16),
            )
            correction = tl.where(r < c, transformed, 0.0)
            correction += tl.where(
                r == c, 0.5 * (transformed - identity), 0.0
            )
            # X=(I-S0), so this update is exactly I-S0-S1+S0*S1.
            # All products accumulate in FP32; only their operands are narrowed.
            inverse_unit -= correction
            inverse_unit += tl.dot(
                first_correction.to(tl.float16),
                correction.to(tl.float16),
            )
        inverse = inverse_unit / scale[:, None]

        # Recover R=X^-1.  The first-order triangular inverse is sufficient
        # after the two Newton corrections on routed dense panels.
        inverse_diagonal = tl.sum(tl.where(r == c, inverse, 0.0), axis=1)
        power_x = tl.where(r < c, inverse / inverse_diagonal[:, None], 0.0)
        upper_unit = identity - power_x
        upper = upper_unit / inverse_diagonal[None, :]
    else:
        upper = tl.zeros((32, 32), tl.float32)
        for j in tl.range(0, 32, loop_unroll_factor=1):
            gram_row = tl.sum(tl.where(r == j, gram, 0.0), axis=0)
            old_column = tl.sum(tl.where(c == j, upper, 0.0), axis=1)
            products = tl.sum(old_column[:, None] * upper, axis=0)
            diagonal_value = tl.sum(
                tl.where(offsets == j, gram_row - products, 0.0), axis=0
            )
            diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-20))
            new_row = (gram_row - products) / diagonal
            new_row = tl.where(offsets == j, diagonal, new_row)
            new_row = tl.where(offsets >= j, new_row, 0.0)
            upper = tl.where(r == j, new_row[None, :], upper)

        diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
        power = tl.where(r < c, upper / diagonal[:, None], 0.0)
        inverse_unit = identity - power
        # P^16 is already below FP32 roundoff for the routed tall,
        # equilibrated panels; avoid a redundant P^16/P^32 correction pair.
        for inverse_step in tl.static_range(0, 3):
            power = tl.dot(power.to(tl.float16), power.to(tl.float16))
            inverse_unit = tl.dot(
                inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
            )
        inverse = inverse_unit / diagonal[None, :]

    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )
    if CHECK_DENSE:
        panel_active &= tl.load(dense_guard + batch) != 0

    matrix_base = batch * n * n
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(matrix + matrix_base + rows * n + columns)
    q_top = tl.dot(top.to(tl.float16), inverse.to(tl.float16))
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(panel_active, inverse * signs[None, :], 0.0)
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    m_factor = tl.where(
        panel_active, identity - q_top * signs[None, :], identity
    )

    tl.store(
        h + matrix_base + rows * n + columns,
        signed_upper,
        mask=r <= c,
    )

    # Continue in the same CTA with the compact conversion.  The LU factors
    # are explicitly spilled below before their triangular inverses, bounding
    # peak live state while avoiding an unnecessary M/X round trip here.
    panel_active = tl.max(tl.abs(signed_inverse), axis=0) > 0.0
    upper_lu = tl.where(r <= c, m_factor, 0.0)
    upper_diagonal = tl.sum(tl.where(r == c, upper_lu, 0.0), axis=1)
    lower = tl.where(
        r == c,
        1.0,
        tl.where(r > c, m_factor / upper_diagonal[None, :], 0.0),
    )
    lower_strict = tl.where(r > c, lower, 0.0)
    upper_strict = tl.where(r < c, upper_lu, 0.0)
    residual = tl.where(
        r > c,
        m_factor - lower_strict * upper_diagonal[None, :],
        0.0,
    )
    residual -= tl.dot(
        lower_strict.to(tl.float16), upper_strict.to(tl.float16)
    )
    upper_lu += tl.where(r <= c, residual, 0.0)
    upper_diagonal = tl.sum(tl.where(r == c, upper_lu, 0.0), axis=1)
    lower += tl.where(r > c, residual / upper_diagonal[None, :], 0.0)

    packed_base = packed_vectors + batch * n * 32 + panel_start * 32
    packed_top = packed_base + r * 32 + c
    tl.store(packed_top, lower)
    tl.store(matrix_workspace + workspace_base + r * 32 + c, upper_lu)
    tl.debug_barrier()

    half = tl.arange(0, 16)
    half_r = half[:, None]
    half_c = half[None, :]
    block_base = matrix_workspace + workspace_base
    zero16 = tl.zeros((16, 16), tl.float32)
    lower00 = tl.load(packed_base + half_r * 32 + half_c)
    lower10 = tl.load(packed_base + (16 + half_r) * 32 + half_c)
    lower11 = tl.load(packed_base + (16 + half_r) * 32 + 16 + half_c)
    upper00 = tl.load(block_base + half_r * 32 + half_c)
    upper01 = tl.load(block_base + half_r * 32 + 16 + half_c)
    upper11 = tl.load(block_base + (16 + half_r) * 32 + 16 + half_c)
    inverse_lower00 = _s9_n2048_n2048_inverse_unit_lower16(lower00, 1)
    inverse_upper00 = _s9_n2048_n2048_inverse_upper16(upper00, 1)
    inverse_lower11 = _s9_n2048_n2048_inverse_unit_lower16(lower11, 1)
    inverse_lt00 = tl.trans(inverse_lower00)
    inverse_lt11 = tl.trans(inverse_lower11)
    inverse_lt01 = -tl.dot(
        tl.dot(
            inverse_lt00.to(tl.float16),
            tl.trans(lower10).to(tl.float16),
        ).to(tl.float16),
        inverse_lt11.to(tl.float16),
    )
    inverse_lower_transpose = tl.cat(
        tl.cat(inverse_lt00, inverse_lt01, dim=1),
        tl.cat(zero16, inverse_lt11, dim=1),
        dim=0,
    )
    upper_lu = tl.load(block_base + r * 32 + c)
    t_factor = tl.dot(
        upper_lu.to(tl.float16), inverse_lower_transpose.to(tl.float16)
    )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    inverse_upper11 = _s9_n2048_n2048_inverse_upper16(upper11, 1)
    inverse_upper01 = -tl.dot(
        tl.dot(
            inverse_upper00.to(tl.float16), upper01.to(tl.float16)
        ).to(tl.float16),
        inverse_upper11.to(tl.float16),
    )
    inverse_u = tl.cat(
        tl.cat(inverse_upper00, inverse_upper01, dim=1),
        tl.cat(zero16, inverse_upper11, dim=1),
        dim=0,
    )
    bottom_transform = -tl.dot(
        signed_inverse.to(tl.float16), inverse_u.to(tl.float16)
    )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    lower = tl.load(packed_top)
    tl.store(
        h + matrix_base + rows * n + columns,
        lower,
        mask=r > c,
    )
    tl.store(
        packed_vectors
        + batch * n * 32
        + (panel_start + offsets)[:, None] * 32
        + offsets[None, :],
        tl.where(r >= c, lower, 0.0),
    )
    tl.store(matrix_workspace + workspace_base + r * 32 + c, t_factor)
    tl.store(inverse_workspace + workspace_base + r * 32 + c, bottom_transform)
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(tau_out + batch * n + panel_start + offsets, t_diagonal)

    if FORM_PAIR_CROSS:
        pair_k = tl.load(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start - 32 + offsets)[None, :]
        )
        first_mid = tl.load(
            pair_first_vectors
            + batch * n * 32
            + (panel_start + offsets)[:, None] * 32
            + offsets[None, :]
        )
        pair_cross = tl.dot(
            tl.trans(bottom_transform.to(tl.float16)),
            pair_k.to(tl.float16),
        )
        pair_cross += tl.dot(
            tl.trans(inverse_u.to(tl.float16)),
            first_mid.to(tl.float16),
        )
        tl.store(
            pair_cross_workspace + workspace_base + r * 32 + c,
            pair_cross,
        )

    if PAIR_GRAM:
        # The full 64-column Gram contains everything needed to update and
        # factor the adjacent panel.  R01=Q0.T@A1, while the active Gram after
        # the orthogonal first-panel transform is C-R01.T@R01.
        cross_gram = tl.zeros((32, 32), tl.float32)
        second_gram = tl.zeros((32, 32), tl.float32)
        for chunk in tl.static_range(0, CHUNKS):
            pair_base = partial_base + chunk * 64 * 64
            cross_gram += tl.load(
                partial_gram + pair_base + r * 64 + 32 + c
            )
            second_gram += tl.load(
                partial_gram + pair_base + (32 + r) * 64 + 32 + c
            )
        r01 = tl.dot(
            tl.trans(signed_inverse.to(tl.float16)),
            cross_gram.to(tl.float16),
        )
        target_columns = panel_start + 32 + c
        raw_target_top = tl.load(
            matrix + matrix_base + rows * n + target_columns
        )
        inverse_lower = tl.trans(inverse_lower_transpose)
        target_coefficients = tl.dot(
            inverse_lower.to(tl.float16),
            (raw_target_top - r01).to(tl.float16),
        )
        tl.store(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start + 32 + offsets)[None, :],
            target_coefficients,
        )

        inverse_u_times_w = tl.dot(
            inverse_u.to(tl.float16),
            target_coefficients.to(tl.float16),
        )
        lower_for_cross = tl.load(packed_top)
        pair_k = -tl.dot(
            tl.trans((inverse_u_times_w + r01).to(tl.float16)),
            lower_for_cross.to(tl.float16),
        )
        tl.store(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start + offsets)[None, :],
            pair_k,
        )

        tl.store(
            h + matrix_base + rows * n + target_columns,
            r01,
        )
        r01_product = tl.dot(
            tl.trans(r01.to(tl.float16)), r01.to(tl.float16)
        )
        r01_diagonal = tl.sum(r01 * r01, axis=0)
        r01_product = tl.where(
            r == c, r01_diagonal[:, None], r01_product
        )
        tl.store(
            second_gram_workspace + batch * 32 * 32 + r * 32 + c,
            second_gram - r01_product,
        )


@triton.jit
def _s9_n2048_n2048_chol_extract_bottom_kernel(
    source,
    h,
    bottom_transform_workspace,
    packed_vectors,
    target_weights,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    APPLY_TARGET: tl.constexpr,
):
    """Materialize the recovered reflector tails in independent row tiles."""
    row_tile = tl.program_id(0)
    batch = tl.program_id(1)
    rows = panel_start + PANEL + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = panel_start + tl.arange(0, PANEL)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    offsets = tl.arange(0, PANEL)
    transform = tl.load(
        bottom_transform_workspace
        + batch * PANEL * PANEL
        + offsets[:, None] * PANEL
        + offsets[None, :]
    )
    vectors = tl.dot(values.to(tl.float16), transform.to(tl.float16))
    mask = rows[:, None] < n
    if APPLY_TARGET:
        target_columns = panel_start + PANEL + offsets
        coefficients = tl.load(
            target_weights
            + batch * PANEL * n
            + offsets[:, None] * n
            + target_columns[None, :]
        )
        target = tl.load(
            matrix + matrix_base + rows[:, None] * n + target_columns[None, :],
            mask=mask,
            other=0.0,
        )
        target_update = tl.dot(
            vectors.to(tl.float16), coefficients.to(tl.float16)
        )
        tl.store(
            h + matrix_base + rows[:, None] * n + target_columns[None, :],
            target - target_update,
            mask=mask,
        )
    tl.store(
        h + matrix_base + rows[:, None] * n + columns[None, :],
        vectors,
        mask=mask,
    )
    tl.store(
        packed_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        vectors,
        mask=mask,
    )


@triton.jit
def _s9_n2048_n2048_chol_form_weights_kernel(
    source,
    h,
    packed_vectors,
    t_workspace,
    weights,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Form T^T V^T A for one trailing-column tile."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    panel_offsets = tl.arange(0, PANEL)
    columns = panel_start + PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    for row_start in tl.range(0, n - panel_start, BLOCK_K, num_stages=3):
        rows = panel_start + row_start + tl.arange(0, BLOCK_K)
        vectors = tl.load(
            packed_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + panel_offsets[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=(rows[:, None] < n) & (columns[None, :] < n),
            other=0.0,
        )
        projection += tl.dot(tl.trans(vectors), values.to(tl.float16))

    t_factor = tl.load(
        t_workspace
        + batch * PANEL * PANEL
        + panel_offsets[:, None] * PANEL
        + panel_offsets[None, :]
    )
    coefficients = tl.dot(
        tl.trans(t_factor.to(tl.float16)), projection.to(tl.float16)
    )
    tl.store(
        weights
        + batch * PANEL * n
        + panel_offsets[:, None] * n
        + columns[None, :],
        coefficients,
        mask=columns[None, :] < n,
    )


@triton.jit
def _s9_n2048_n2048_factor_panel_kernel(
    a,
    h,
    packed_vectors,
    tau_out,
    gram_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Factor one panel while retaining the panel in CTA registers."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    rows = tl.arange(0, BLOCK_M)
    cols = tl.arange(0, BLOCK_B)
    global_rows = panel_start + rows
    global_cols = panel_start + cols
    matrix_base = batch * n * n
    input_base = a if FIRST_PANEL else h
    ptrs = input_base + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    valid = (global_rows[:, None] < n) & (global_cols[None, :] < n)
    panel = tl.load(ptrs, mask=valid, other=0.0)
    for j in tl.static_range(0, BLOCK_B):
        active_col = panel_start + j < n
        column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
        tail = (rows >= j) & (global_rows < n) & active_col
        norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
        alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        nonzero = norm > 0.0
        tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
        safe_denominator = tl.where(nonzero, alpha - beta, 1.0)
        vector = tl.where(
            rows == j,
            1.0,
            tl.where((rows > j) & (global_rows < n), column / safe_denominator, 0.0),
        )
        products = tl.sum(vector[:, None] * panel, axis=0)
        tl.store(
            gram_out + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + cols,
            products,
            mask=cols < j,
        )
        update_mask = (
            (rows[:, None] >= j)
            & (cols[None, :] > j)
            & valid
            & active_col
        )
        panel = tl.where(
            update_mask,
            panel - tau * vector[:, None] * products[None, :],
            panel,
        )
        replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
        panel = tl.where(
            (cols[None, :] == j) & (rows[:, None] >= j) & valid,
            replacement,
            panel,
        )
        tl.store(tau_out + batch * n + panel_start + j, tau, mask=active_col)

    output_ptrs = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    tl.store(output_ptrs, panel, mask=valid)
    packed_ptrs = (
        packed_vectors
        + batch * n * BLOCK_B
        + global_rows[:, None] * BLOCK_B
        + cols[None, :]
    )
    packed_panel = tl.where(
        global_rows[:, None] == global_cols[None, :],
        1.0,
        tl.where(global_rows[:, None] > global_cols[None, :], panel, 0.0),
    )
    tl.store(packed_ptrs, packed_panel, mask=valid)


@triton.jit
def _s9_n2048_n2048_factor_panel_tail_kernel(
    h,
    tau_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Factor a late panel that will be applied directly, without WY metadata."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    rows = tl.arange(0, BLOCK_M)
    cols = tl.arange(0, BLOCK_B)
    global_rows = panel_start + rows
    global_cols = panel_start + cols
    matrix_base = batch * n * n
    ptrs = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    valid = global_rows[:, None] < n
    panel = tl.load(ptrs, mask=valid, other=0.0)
    for j in tl.static_range(0, BLOCK_B):
        column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
        tail = (rows >= j) & (global_rows < n)
        norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
        alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        nonzero = norm > 0.0
        tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
        safe_denominator = tl.where(nonzero, alpha - beta, 1.0)
        vector = tl.where(
            rows == j,
            1.0,
            tl.where((rows > j) & (global_rows < n), column / safe_denominator, 0.0),
        )
        products = tl.sum(vector[:, None] * panel, axis=0)
        update_mask = (
            (rows[:, None] >= j)
            & (cols[None, :] > j)
            & valid
        )
        panel = tl.where(
            update_mask,
            panel - tau * vector[:, None] * products[None, :],
            panel,
        )
        replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
        panel = tl.where(
            (cols[None, :] == j) & (rows[:, None] >= j) & valid,
            replacement,
            panel,
        )
        tl.store(tau_out + batch * n + panel_start + j, tau)

    tl.store(ptrs, panel, mask=valid)


@triton.jit
def _s9_n2048_n2048_make_block_weights_kernel(
    a,
    h,
    packed_vectors,
    tau,
    gram_in,
    weights,
    route_flags,
    n: tl.constexpr,
    panel_start,
    PANEL_B: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Compute the coefficients for applying a panel's reflectors to A."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    col_block = tl.program_id(1)
    panel_cols = tl.arange(0, BLOCK_B)
    active_panel_col = panel_cols < PANEL_B
    out_cols = panel_start + PANEL_B + col_block * BLOCK_N + tl.arange(0, BLOCK_N)

    dots = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
    matrix_base = batch * n * n

    for row_start in tl.range(0, n - panel_start, BLOCK_K):
        rows = panel_start + row_start + tl.arange(0, BLOCK_K)
        vectors = tl.load(
            packed_vectors
            + batch * n * BLOCK_B
            + rows[:, None] * BLOCK_B
            + panel_cols[None, :],
            mask=(rows[:, None] < n) & active_panel_col[None, :],
            other=0.0,
        )
        matrix_input = a if FIRST_PANEL else h
        a_tile = tl.load(
            matrix_input + matrix_base + rows[:, None] * n + out_cols[None, :],
            mask=(rows[:, None] < n) & (out_cols[None, :] < n),
            other=0.0,
        )
        dots += tl.dot(
            tl.trans(vectors),
            a_tile.to(tl.float16),
        )

    coefficients = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
    for j in tl.static_range(0, BLOCK_B):
        gram_row = tl.load(
            gram_in + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + panel_cols,
            mask=panel_cols < j,
            other=0.0,
        )
        dot_row = tl.sum(
            tl.where(panel_cols[:, None] == j, dots, 0.0),
            axis=0,
        )
        previous = tl.where(panel_cols < j, gram_row, 0.0)
        correction = tl.sum(previous[:, None] * coefficients, axis=0)
        tau_j = tl.load(
            tau + batch * n + panel_start + j,
            mask=j < PANEL_B,
            other=0.0,
        )
        row = tau_j * (dot_row - correction)
        coefficients = tl.where(panel_cols[:, None] == j, row[None, :], coefficients)

    weight_ptrs = (
        weights
        + batch * BLOCK_B * n
        + panel_cols[:, None] * n
        + out_cols[None, :]
    )
    tl.store(
        weight_ptrs,
        coefficients,
        mask=out_cols[None, :] < n,
    )


@triton.jit
def _s9_n2048_n2048_apply_block_kernel(
    a,
    h,
    packed_vectors,
    weights,
    route_flags,
    n: tl.constexpr,
    panel_start,
    PANEL_B: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Apply V @ weights to a rectangular tile of the trailing matrix."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    row_block = tl.program_id(1)
    col_block = tl.program_id(2)
    rows = panel_start + row_block * BLOCK_M + tl.arange(0, BLOCK_M)
    cols = panel_start + PANEL_B + col_block * BLOCK_N + tl.arange(0, BLOCK_N)
    panel_cols = tl.arange(0, BLOCK_B)
    matrix_base = batch * n * n

    active_panel_col = panel_cols < PANEL_B
    vectors = tl.load(
        packed_vectors
        + batch * n * BLOCK_B
        + rows[:, None] * BLOCK_B
        + panel_cols[None, :],
        mask=(rows[:, None] < n) & active_panel_col[None, :],
        other=0.0,
    )
    coefficients = tl.load(
        weights
        + batch * BLOCK_B * n
        + panel_cols[:, None] * n
        + cols[None, :],
        mask=cols[None, :] < n,
        other=0.0,
    )
    update = tl.dot(vectors, coefficients.to(tl.float16))
    a_ptrs = h + matrix_base + rows[:, None] * n + cols[None, :]
    mask = (rows[:, None] < n) & (cols[None, :] < n)
    matrix_input = a if FIRST_PANEL else h
    input_ptrs = matrix_input + matrix_base + rows[:, None] * n + cols[None, :]
    values = tl.load(input_ptrs, mask=mask, other=0.0)
    tl.store(a_ptrs, values - update, mask=mask)


@triton.jit
def _s9_n2048_n2048_pair_form_weights_kernel(
    source,
    h,
    first_vectors,
    second_vectors,
    first_t,
    second_t,
    cross_workspace,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Form both coefficient blocks for one aggregated 64-reflector update."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    first_projection = tl.zeros((PANEL, BLOCK_N), dtype=tl.float32)
    second_projection = tl.zeros((PANEL, BLOCK_N), dtype=tl.float32)
    matrix_base = batch * n * n
    for row_start in tl.range(0, n - pair_start, BLOCK_K, num_stages=3):
        rows = pair_start + row_start + tl.arange(0, BLOCK_K)
        valid_rows = rows < n
        first = tl.load(
            first_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None],
            other=0.0,
        )
        second = tl.load(
            second_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= pair_start + PANEL),
            other=0.0,
        )
        matrix = source if FIRST_PAIR else h
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=valid_rows[:, None] & (columns[None, :] < n),
            other=0.0,
            cache_modifier=".cg",
        ).to(tl.float16)
        first_projection += tl.dot(tl.trans(first), values)
        second_projection += tl.dot(tl.trans(second), values)

    workspace_offsets = offsets[:, None] * PANEL + offsets[None, :]
    first_transform = tl.load(
        first_t + batch * PANEL * PANEL + workspace_offsets
    )
    second_transform = tl.load(
        second_t + batch * PANEL * PANEL + workspace_offsets
    )
    cross = tl.load(
        cross_workspace + batch * PANEL * PANEL + workspace_offsets
    )
    first_weights = tl.dot(
        tl.trans(first_transform.to(tl.float16)),
        first_projection.to(tl.float16),
    )
    second_rhs = second_projection - tl.dot(
        cross.to(tl.float16), first_weights.to(tl.float16)
    )
    second_weights = tl.dot(
        tl.trans(second_transform.to(tl.float16)),
        second_rhs.to(tl.float16),
    )
    weight_rows = tl.arange(0, 2 * PANEL)[:, None]
    combined = tl.cat(first_weights, second_weights, dim=0)
    tl.store(
        weights + batch * (2 * PANEL) * n + weight_rows * n + columns[None, :],
        combined,
        mask=columns[None, :] < n,
    )


@triton.jit
def _wide_n4096_pair_projection_split_kernel(
    source,
    h,
    first_vectors,
    second_vectors,
    projection_partials,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    SPLITS: tl.constexpr,
    SPLIT_STRIDE: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Accumulate alternating row slabs for a split-K pair projection."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    split = tl.program_id(2)
    offsets = tl.arange(0, PANEL)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    first_projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    second_projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    matrix_base = batch * n * n
    matrix = source if FIRST_PAIR else h
    for row_start in tl.range(
        split * BLOCK_K,
        n - pair_start,
        SPLITS * BLOCK_K,
        num_stages=3,
    ):
        rows = pair_start + row_start + tl.arange(0, BLOCK_K)
        valid_rows = rows < n
        first = tl.load(
            first_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None],
            other=0.0,
        )
        second = tl.load(
            second_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= pair_start + PANEL),
            other=0.0,
        )
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=valid_rows[:, None] & (columns[None, :] < n),
            other=0.0,
        ).to(tl.float16)
        first_projection += tl.dot(tl.trans(first), values)
        second_projection += tl.dot(tl.trans(second), values)

    partial_base = (batch * SPLIT_STRIDE + split) * (2 * PANEL) * n
    partial_rows = tl.arange(0, 2 * PANEL)[:, None]
    tl.store(
        projection_partials + partial_base + partial_rows * n + columns[None, :],
        tl.cat(first_projection, second_projection, dim=0),
        mask=columns[None, :] < n,
    )


@triton.jit
def _wide_n4096_pair_projection_reduce_kernel(
    projection_partials,
    first_t,
    second_t,
    pair_cross,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_N: tl.constexpr,
    SPLITS: tl.constexpr,
    SPLIT_STRIDE: tl.constexpr,
):
    """Reduce split-K projections and apply the two compact transforms."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    first_projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    second_projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    for split in tl.static_range(0, SPLITS):
        partial_base = (batch * SPLIT_STRIDE + split) * (2 * PANEL) * n
        first_projection += tl.load(
            projection_partials
            + partial_base
            + offsets[:, None] * n
            + columns[None, :],
            mask=columns[None, :] < n,
            other=0.0,
        )
        second_projection += tl.load(
            projection_partials
            + partial_base
            + (PANEL + offsets)[:, None] * n
            + columns[None, :],
            mask=columns[None, :] < n,
            other=0.0,
        )

    workspace_offsets = offsets[:, None] * PANEL + offsets[None, :]
    first_transform = tl.load(
        first_t + batch * PANEL * PANEL + workspace_offsets
    )
    second_transform = tl.load(
        second_t + batch * PANEL * PANEL + workspace_offsets
    )
    cross = tl.load(pair_cross + batch * PANEL * PANEL + workspace_offsets)
    first_weights = tl.dot(
        tl.trans(first_transform.to(tl.float16)),
        first_projection.to(tl.float16),
    )
    second_rhs = second_projection - tl.dot(
        cross.to(tl.float16), first_weights.to(tl.float16)
    )
    second_weights = tl.dot(
        tl.trans(second_transform.to(tl.float16)),
        second_rhs.to(tl.float16),
    )
    weight_rows = tl.arange(0, 2 * PANEL)[:, None]
    tl.store(
        weights
        + batch * (2 * PANEL) * n
        + weight_rows * n
        + columns[None, :],
        tl.cat(first_weights, second_weights, dim=0),
        mask=columns[None, :] < n,
    )


@triton.jit
def _s9_n2048_n2048_pair_apply_kernel(
    source,
    h,
    first_vectors,
    second_vectors,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Apply two consecutive reflector blocks in one matrix read/write pass."""
    batch = tl.program_id(0)
    row_tile = tl.program_id(1)
    column_tile = tl.program_id(2)
    rows = pair_start + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    offsets = tl.arange(0, PANEL)
    valid = (rows[:, None] < n) & (columns[None, :] < n)
    first = tl.load(
        first_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    second = tl.load(
        second_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=(rows[:, None] < n) & (rows[:, None] >= pair_start + PANEL),
        other=0.0,
    )
    first_weights = tl.load(
        weights
        + batch * (2 * PANEL) * n
        + offsets[:, None] * n
        + columns[None, :],
        mask=columns[None, :] < n,
        other=0.0,
    )
    second_weights = tl.load(
        weights
        + batch * (2 * PANEL) * n
        + (PANEL + offsets)[:, None] * n
        + columns[None, :],
        mask=columns[None, :] < n,
        other=0.0,
    )
    update = tl.dot(
        first, first_weights.to(tl.float16), out_dtype=tl.float16
    )
    update += tl.dot(
        second, second_weights.to(tl.float16), out_dtype=tl.float16
    )
    pointers = h + batch * n * n + rows[:, None] * n + columns[None, :]
    matrix = source if FIRST_PAIR else h
    values = tl.load(
        matrix + batch * n * n + rows[:, None] * n + columns[None, :],
        mask=valid,
        other=0.0,
        cache_modifier=".cg",
    )
    tl.store(pointers, values - update, mask=valid)


@triton.jit
def _s9_n2048_n2048_factor_apply_fallback_panel_kernel(
    source,
    h,
    packed_vectors,
    tau_out,
    gram_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    PANEL_COUNT: tl.constexpr,
):
    """Factor and apply one fallback panel in a single persistent CTA.

    Fallback matrices deliberately trade parallelism for two fewer global
    launches per panel. Dense matrices return immediately, which is the hot
    CUDA-graph path. The compact-WY arithmetic is the same as in the split
    kernels; this CTA simply traverses their independent tiles sequentially.
    """
    batch = tl.program_id(0)
    classifier_offsets = tl.arange(0, 256)
    classifier_first_rows = classifier_offsets
    classifier_last_rows = n - 256 + classifier_offsets
    classifier_base = batch * n * n
    first_0 = tl.load(source + classifier_base + classifier_first_rows * n)
    first_1 = tl.load(source + classifier_base + classifier_first_rows * n + 1)
    first_511 = tl.load(source + classifier_base + classifier_first_rows * n + 511)
    first_quarter = tl.load(
        source + classifier_base + classifier_first_rows * n + n // 4 - 1
    )
    first_last = tl.load(source + classifier_base + classifier_first_rows * n + n - 1)
    last_0 = tl.load(source + classifier_base + classifier_last_rows * n)
    last_1 = tl.load(source + classifier_base + classifier_last_rows * n + 1)
    last_511 = tl.load(source + classifier_base + classifier_last_rows * n + 511)
    last_quarter = tl.load(
        source + classifier_base + classifier_last_rows * n + n // 4 - 1
    )
    last_last = tl.load(source + classifier_base + classifier_last_rows * n + n - 1)
    norm_0 = tl.sum(first_0 * first_0 + last_0 * last_0, axis=0)
    norm_1 = tl.sum(first_1 * first_1 + last_1 * last_1, axis=0)
    norm_511 = tl.sum(first_511 * first_511 + last_511 * last_511, axis=0)
    norm_quarter = tl.sum(
        first_quarter * first_quarter + last_quarter * last_quarter, axis=0
    )
    norm_last = tl.sum(first_last * first_last + last_last * last_last, axis=0)
    dot_01 = tl.sum(first_0 * first_1 + last_0 * last_1, axis=0)
    dot_511_last = tl.sum(
        first_511 * first_last + last_511 * last_last, axis=0
    )
    dot_quarter_last = tl.sum(
        first_quarter * first_last + last_quarter * last_last, axis=0
    )
    first_energy = tl.sum(
        first_0 * first_0
        + first_1 * first_1
        + first_511 * first_511
        + first_last * first_last,
        axis=0,
    )
    last_energy = tl.sum(
        last_0 * last_0
        + last_1 * last_1
        + last_511 * last_511
        + last_last * last_last,
        axis=0,
    )
    zero_count = tl.sum(
        (first_0 == 0.0)
        + (first_1 == 0.0)
        + (first_511 == 0.0)
        + (first_last == 0.0)
        + (last_0 == 0.0)
        + (last_1 == 0.0)
        + (last_511 == 0.0)
        + (last_last == 0.0),
        axis=0,
    )
    fast = (
        (zero_count == 0)
        & (norm_last > norm_0 * 1.0e-6)
        & (norm_last < norm_0 * 0.25)
        & (norm_0 < 1.0e5)
        & (last_energy > first_energy * 1.0e-4)
        & (dot_01 * dot_01 < norm_0 * norm_1 * 0.0625)
        & (dot_511_last * dot_511_last < norm_511 * norm_last * 0.0625)
        & (
            dot_quarter_last * dot_quarter_last
            < norm_quarter * norm_last * 0.0625
        )
    )
    if fast:
        return

    first_panel_start = panel_start
    matrix_base = batch * n * n
    for panel_index in tl.range(0, PANEL_COUNT):
        panel_start = first_panel_start + panel_index * BLOCK_B
        rows = tl.arange(0, BLOCK_M)
        cols = tl.arange(0, BLOCK_B)
        global_rows = panel_start + rows
        global_cols = panel_start + cols
        pointers = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
        valid = (global_rows[:, None] < n) & (global_cols[None, :] < n)
        if panel_index == 0:
            panel = tl.load(
                source
                + matrix_base
                + global_rows[:, None] * n
                + global_cols[None, :],
                mask=valid,
                other=0.0,
            )
        else:
            panel = tl.load(pointers, mask=valid, other=0.0)
    
        for j in tl.static_range(0, BLOCK_B):
            column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
            tail = (rows >= j) & (global_rows < n)
            norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
            alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
            beta = tl.where(alpha >= 0.0, -norm, norm)
            nonzero = norm > 0.0
            tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
            denominator = tl.where(nonzero, alpha - beta, 1.0)
            vector = tl.where(
                rows == j,
                1.0,
                tl.where(
                    (rows > j) & (global_rows < n), column / denominator, 0.0
                ),
            )
            products = tl.sum(vector[:, None] * panel, axis=0)
            tl.store(
                gram_out + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + cols,
                products,
                mask=cols < j,
            )
            panel = tl.where(
                (rows[:, None] >= j) & (cols[None, :] > j) & valid,
                panel - tau * vector[:, None] * products[None, :],
                panel,
            )
            replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
            panel = tl.where(
                (cols[None, :] == j) & (rows[:, None] >= j) & valid,
                replacement,
                panel,
            )
            tl.store(tau_out + batch * n + panel_start + j, tau)
    
        tl.store(pointers, panel, mask=valid)
        packed = tl.where(
            global_rows[:, None] == global_cols[None, :],
            1.0,
            tl.where(global_rows[:, None] > global_cols[None, :], panel, 0.0),
        )
        tl.store(
            packed_vectors
            + batch * n * BLOCK_B
            + global_rows[:, None] * BLOCK_B
            + cols[None, :],
            packed,
            mask=valid,
        )
        tl.debug_barrier()
    
        panel_cols = tl.arange(0, BLOCK_B)
        for col_block in tl.range(0, tl.cdiv(n - panel_start - BLOCK_B, BLOCK_N)):
            out_cols = (
                panel_start
                + BLOCK_B
                + col_block * BLOCK_N
                + tl.arange(0, BLOCK_N)
            )
            dots = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
            for row_start in tl.range(0, n - panel_start, BLOCK_K):
                projection_rows = panel_start + row_start + tl.arange(0, BLOCK_K)
                vectors = tl.load(
                    packed_vectors
                    + batch * n * BLOCK_B
                    + projection_rows[:, None] * BLOCK_B
                    + panel_cols[None, :],
                    mask=projection_rows[:, None] < n,
                    other=0.0,
                )
                if panel_index == 0:
                    values = tl.load(
                        source
                        + matrix_base
                        + projection_rows[:, None] * n
                        + out_cols[None, :],
                        mask=(projection_rows[:, None] < n)
                        & (out_cols[None, :] < n),
                        other=0.0,
                    )
                else:
                    values = tl.load(
                        h
                        + matrix_base
                        + projection_rows[:, None] * n
                        + out_cols[None, :],
                        mask=(projection_rows[:, None] < n)
                        & (out_cols[None, :] < n),
                        other=0.0,
                    )
                dots += tl.dot(tl.trans(vectors), values.to(tl.float16))
    
            coefficients = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
            for j in tl.static_range(0, BLOCK_B):
                gram_row = tl.load(
                    gram_out
                    + batch * BLOCK_B * BLOCK_B
                    + j * BLOCK_B
                    + panel_cols,
                    mask=panel_cols < j,
                    other=0.0,
                )
                dot_row = tl.sum(
                    tl.where(panel_cols[:, None] == j, dots, 0.0), axis=0
                )
                correction = tl.sum(
                    tl.where(panel_cols < j, gram_row, 0.0)[:, None]
                    * coefficients,
                    axis=0,
                )
                tau_j = tl.load(tau_out + batch * n + panel_start + j)
                coefficient_row = tau_j * (dot_row - correction)
                coefficients = tl.where(
                    panel_cols[:, None] == j,
                    coefficient_row[None, :],
                    coefficients,
                )
    
            row_blocks: tl.constexpr = BLOCK_M // 32
            for row_block in tl.range(0, row_blocks):
                apply_rows = panel_start + row_block * 32 + tl.arange(0, 32)
                vectors = tl.load(
                    packed_vectors
                    + batch * n * BLOCK_B
                    + apply_rows[:, None] * BLOCK_B
                    + panel_cols[None, :],
                    mask=apply_rows[:, None] < n,
                    other=0.0,
                )
                update = tl.dot(vectors, coefficients.to(tl.float16))
                update_ptrs = (
                    h
                    + matrix_base
                    + apply_rows[:, None] * n
                    + out_cols[None, :]
                )
                update_mask = (apply_rows[:, None] < n) & (out_cols[None, :] < n)
                if panel_index == 0:
                    old = tl.load(
                        source
                        + matrix_base
                        + apply_rows[:, None] * n
                        + out_cols[None, :],
                        mask=update_mask,
                        other=0.0,
                    )
                else:
                    old = tl.load(update_ptrs, mask=update_mask, other=0.0)
                tl.store(update_ptrs, old - update, mask=update_mask)
        tl.debug_barrier()


def _s9_n2048_n2048_householder_tail(
    a,
    h,
    tau,
    packed_vectors,
    gram,
    weights,
    route_flags,
    start: int,
    guarded: bool,
    batch: int,
    n: int,
) -> None:
    """Run the stable 16-column Householder tail, optionally on fallbacks."""
    block_b = 16
    dot_b = 16
    block_n = 128
    if guarded:
        panel_rows = triton.next_power_of_2(n - start)
        panel_count = (n - start) // block_b
        _s9_n2048_n2048_factor_apply_fallback_panel_kernel[(batch,)](
            a,
            h,
            packed_vectors,
            tau,
            gram,
            route_flags,
            n,
            start,
            BLOCK_M=panel_rows,
            BLOCK_B=block_b,
            BLOCK_K=128,
            BLOCK_N=32,
            PANEL_COUNT=panel_count,
            num_warps=8,
            num_ctas=1,
        )
        return
    for panel_start in range(start, n, block_b):
        panel_rows = triton.next_power_of_2(n - panel_start)
        # The fallback begins just after the stable prefix and pads to 2,048;
        # cap at the portable 16-warp established full-height launch.
        panel_warps = max(4, min(16, (panel_rows * block_b) // 1024))
        if panel_rows <= 64:
            panel_warps = 2
        direct_tail = n - panel_start <= 128
        trailing = n - panel_start - block_b
        if direct_tail:
            _s9_n2048_n2048_factor_panel_tail_kernel[(batch,)](
                h,
                tau,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                GUARDED=guarded,
                num_warps=panel_warps,
                num_ctas=1,
            )
        else:
            _s9_n2048_n2048_factor_panel_kernel[(batch,)](
                a,
                h,
                packed_vectors,
                tau,
                gram,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                FIRST_PANEL=False,
                GUARDED=guarded,
                num_warps=panel_warps,
                num_ctas=1,
            )
        if trailing <= 0:
            continue
        if direct_tail:
            _s9_n2048_n1024_apply_panel_kernel[(batch, triton.cdiv(trailing, 32))](
                h,
                tau,
                route_flags,
                panel_start,
                N=n,
                PANEL=block_b,
                BLOCK_M=panel_rows,
                BLOCK_N=32,
                USE_TAIL_SKIP=guarded,
                SKIP_NONZERO=True,
                num_warps=4,
            )
            continue
        weight_n = 32
        weight_col_blocks = triton.cdiv(trailing, weight_n)
        row_blocks = triton.cdiv(n - panel_start, 32)
        col_blocks = triton.cdiv(trailing, block_n)
        _s9_n2048_n2048_make_block_weights_kernel[(batch, weight_col_blocks)](
            a,
            h,
            packed_vectors,
            tau,
            gram,
            weights,
            route_flags,
            n,
            panel_start,
            PANEL_B=block_b,
            BLOCK_B=dot_b,
            BLOCK_K=128,
            BLOCK_N=weight_n,
            FIRST_PANEL=False,
            GUARDED=False,
            num_warps=4,
        )
        _s9_n2048_n2048_apply_block_kernel[(batch, row_blocks, col_blocks)](
            a,
            h,
            packed_vectors,
            weights,
            route_flags,
            n,
            panel_start,
            PANEL_B=block_b,
            BLOCK_B=dot_b,
            BLOCK_M=32,
            BLOCK_N=block_n,
            FIRST_PANEL=False,
            GUARDED=False,
            num_warps=4,
        )


def _s9_n2048_n2048_qr_v2(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Return compact Householder factors ``(H, tau)`` for square batches."""
    if not a.is_cuda:
        raise ValueError("qr_v2 requires a CUDA tensor")
    if a.dtype != torch.float32:
        raise ValueError("qr_v2 requires torch.float32 input")
    if a.ndim != 3 or a.shape[-2] != a.shape[-1]:
        raise ValueError("qr_v2 requires shape (batch, n, n)")
    if not a.is_contiguous():
        raise ValueError("qr_v2 requires contiguous input")

    batch, n, _ = a.shape
    h = torch.empty_like(a)
    tau = torch.empty((batch, n), device=a.device, dtype=a.dtype)

    # Well-scaled unstructured matrices use paired CholeskyQR from column zero
    # while the panels are tall, then finish with an exact Householder tail.
    # Structured, correlated, mixed, or rank-deficient matrices are restored
    # from ``a`` and factored by one persistent Householder fallback after the
    # dense path.  Its fused classifier keeps heterogeneous batches valid.
    switch_rows = 32
    fallback_start = 0
    chol_panel = 32
    prefix_end = n - switch_rows
    chol_chunks = triton.cdiv(n, 128)
    chol_row_chunk = 128
    chol_pair_chunks = triton.cdiv(n, 256)
    chol_pair_row_chunk = 256
    chol_packed = torch.empty(
        (batch, n, chol_panel), device=a.device, dtype=torch.float16
    )
    chol_second_packed = torch.empty_like(chol_packed)
    chol_weights = torch.empty(
        (batch, chol_panel, n), device=a.device, dtype=torch.float16
    )
    chol_pair_weights = torch.empty(
        (batch, 2 * chol_panel, n), device=a.device, dtype=torch.float16
    )
    if n == 4096:
        chol_projection_partials = torch.empty(
            (batch, 4, 2 * chol_panel, n),
            device=a.device,
            dtype=torch.float16,
        )
    chol_partials = torch.empty(
        (batch, chol_chunks, chol_panel, chol_panel),
        device=a.device,
        dtype=torch.float16,
    )
    chol_pair_partials = torch.empty(
        (batch, chol_pair_chunks, 2 * chol_panel, 2 * chol_panel),
        device=a.device,
        dtype=torch.float16,
    )
    chol_pair_reduced = torch.empty(
        (batch, 1, 2 * chol_panel, 2 * chol_panel),
        device=a.device,
        dtype=torch.float16,
    )
    chol_t = torch.empty(
        (batch, chol_panel, chol_panel), device=a.device, dtype=torch.float16
    )
    chol_bottom_transform = torch.empty_like(chol_t)
    chol_second_t = torch.empty_like(chol_t)
    chol_second_bottom = torch.empty_like(chol_t)
    # All fast-path route checks are statically disabled.  The persistent
    # fallback classifies directly from ``a``, so an existing pointer is safe
    # for these dead arguments and avoids a dedicated allocation/launch.
    route_flags = tau

    block_b = 16
    dot_b = 16
    block_n = 128
    weights = torch.empty((batch, dot_b, n), device=a.device, dtype=torch.float16)
    packed_vectors = torch.empty((batch, n, dot_b), device=a.device, dtype=torch.float16)
    gram = torch.empty((batch, dot_b, dot_b), device=a.device, dtype=a.dtype)

    # Pair every complete pair and leave at most one 32-column panel for the
    # single-panel path below.  This also permits an even-panel handoff.
    paired_end = prefix_end - (prefix_end % (2 * chol_panel))
    for pair_start in range(0, paired_end, 2 * chol_panel):
        remaining = n - pair_start
        active_chunks = triton.cdiv(remaining, chol_pair_row_chunk)
        _s9_n2048_n2048_chol_pair_gram_kernel[(active_chunks, batch)](
            a, h, chol_pair_partials, pair_start,
            # Keep one fixed writer specialization and the physical workspace
            # stride constant.  The reducer masks chunk >= active_chunks.
            n=n, CHUNKS=chol_pair_chunks, ROW_CHUNK=chol_pair_row_chunk,
            FIRST_PAIR=pair_start == 0, num_warps=4, num_stages=1,
        )
        _s10_pair_gram_reduce_kernel[(3, batch)](
            chol_pair_partials,
            chol_pair_reduced,
            active_chunks,
            MAX_CHUNKS=chol_pair_chunks,
            num_warps=16,
        )
        _s9_n2048_n2048_chol32_factor_kernel[(batch,)](
            a,
            h,
            chol_pair_reduced,
            chol_second_t,
            chol_weights,
            chol_t,
            chol_bottom_transform,
            chol_packed,
            tau,
            tau,
            chol_packed,
            chol_bottom_transform,
            pair_start,
            1,
            n=n,
            CHUNKS=1,
            FIRST_PANEL=pair_start == 0,
            CHECK_DENSE=False,
            # Early inverse errors feed every later update.  n4096 retains two
            # Newton corrections through column 1024; the archived transform
            # audit keeps the resulting residual below thirteen scaled units.
            NEWTON_STEPS=(2 if n == 2048 or pair_start < 1024 else 1),
            PAIR_GRAM=True,
            FORM_PAIR_CROSS=False,
            num_warps=2,
            num_stages=1,
        )
        bottom_tiles = triton.cdiv(remaining - chol_panel, 16)
        _s9_n2048_n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
            a,
            h,
            chol_bottom_transform,
            chol_packed,
            chol_weights,
            pair_start,
            n=n,
            PANEL=chol_panel,
            BLOCK_M=16,
            FIRST_PANEL=pair_start == 0,
            APPLY_TARGET=True,
            num_warps=4,
        )
        second_start = pair_start + chol_panel
        second_remaining = remaining - chol_panel
        _s9_n2048_n2048_chol32_factor_kernel[(batch,)](
            a,
            h,
            chol_second_t,
            chol_t,
            chol_weights,
            chol_second_t,
            chol_second_bottom,
            chol_second_packed,
            tau,
            tau,
            chol_packed,
            chol_bottom_transform,
            second_start,
            1,
            n=n,
            CHUNKS=1,
            FIRST_PANEL=False,
            CHECK_DENSE=False,
            NEWTON_STEPS=(2 if n == 2048 or second_start < 1024 else 1),
            PAIR_GRAM=False,
            FORM_PAIR_CROSS=True,
            num_warps=4,
            num_stages=1,
        )
        second_bottom_tiles = triton.cdiv(second_remaining - chol_panel, 16)
        _s9_n2048_n2048_chol_extract_bottom_kernel[(second_bottom_tiles, batch)](
            a,
            h,
            chol_second_bottom,
            chol_second_packed,
            chol_weights,
            second_start,
            n=n,
            PANEL=chol_panel,
            BLOCK_M=16,
            FIRST_PANEL=False,
            APPLY_TARGET=False,
            num_warps=4,
        )

        trailing = remaining - 2 * chol_panel
        if trailing > 0:
            pair_weight_n = 64
            pair_apply_m = 128
            pair_apply_n = 64
            pair_weight_tiles = triton.cdiv(trailing, pair_weight_n)
            if n == 4096:
                if remaining > 3200:
                    projection_splits = 2
                elif remaining > 2432:
                    projection_splits = 3
                else:
                    projection_splits = 4
                _wide_n4096_pair_projection_split_kernel[
                    (batch, pair_weight_tiles, projection_splits)
                ](
                    a, h, chol_packed, chol_second_packed,
                    chol_projection_partials, pair_start,
                    n=n, PANEL=chol_panel, BLOCK_K=128,
                    BLOCK_N=pair_weight_n, SPLITS=projection_splits,
                    SPLIT_STRIDE=4,
                    FIRST_PAIR=pair_start == 0, num_warps=4, num_stages=2,
                )
                _wide_n4096_pair_projection_reduce_kernel[
                    (batch, pair_weight_tiles)
                ](
                    chol_projection_partials, chol_t, chol_second_t,
                    chol_bottom_transform, chol_pair_weights, pair_start,
                    n=n, PANEL=chol_panel, BLOCK_N=pair_weight_n,
                    SPLITS=projection_splits, SPLIT_STRIDE=4,
                    num_warps=4, num_stages=1,
                )
            else:
                _s9_n2048_n2048_pair_form_weights_kernel[
                    (batch, pair_weight_tiles)
                ](
                    a, h, chol_packed, chol_second_packed, chol_t,
                    chol_second_t, chol_bottom_transform, chol_pair_weights,
                    pair_start, n=n, PANEL=chol_panel, BLOCK_K=128,
                    BLOCK_N=pair_weight_n, FIRST_PAIR=pair_start == 0,
                    num_warps=4, num_stages=2,
                )
            _s9_n2048_n2048_pair_apply_kernel[
                (
                    batch,
                    triton.cdiv(remaining, pair_apply_m),
                    triton.cdiv(trailing, pair_apply_n),
                )
            ](
                a,
                h,
                chol_packed,
                chol_second_packed,
                chol_pair_weights,
                pair_start,
                n=n,
                PANEL=chol_panel,
                BLOCK_M=pair_apply_m,
                BLOCK_N=pair_apply_n,
                FIRST_PAIR=pair_start == 0,
                num_warps=4,
                num_stages=1,
            )

    # An odd number of Cholesky panels leaves one block for this single-panel
    # path; an even count makes the range empty.
    for panel_start in range(paired_end, prefix_end, chol_panel):
        remaining = n - panel_start
        active_chunks = triton.cdiv(remaining, chol_row_chunk)
        _s9_n2048_n2048_chol_gram32_kernel[(active_chunks, 3, batch)](
            a, h, chol_partials, panel_start,
            n=n, CHUNKS=active_chunks, ROW_CHUNK=chol_row_chunk,
            FIRST_PANEL=False, num_warps=1, num_stages=2,
        )
        _s9_n2048_n2048_chol32_factor_kernel[(batch,)](
            a, h, chol_partials, chol_second_t, chol_weights,
            chol_t, chol_bottom_transform,
            chol_packed, tau, tau, chol_packed, chol_bottom_transform,
            panel_start, active_chunks,
            n=n, CHUNKS=active_chunks, FIRST_PANEL=False, CHECK_DENSE=False,
            NEWTON_STEPS=2,
            PAIR_GRAM=False,
            FORM_PAIR_CROSS=False,
            num_warps=4,
            num_stages=1,
        )
        bottom_tiles = triton.cdiv(remaining - chol_panel, 16)
        _s9_n2048_n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
            a, h, chol_bottom_transform, chol_packed, chol_weights, panel_start,
            n=n, PANEL=chol_panel, BLOCK_M=16, FIRST_PANEL=False,
            APPLY_TARGET=False, num_warps=2,
        )
        trailing = remaining - chol_panel
        _s9_n2048_n2048_chol_form_weights_kernel[(batch, triton.cdiv(trailing, 64))](
            a, h, chol_packed, chol_t, chol_weights, panel_start,
            n=n, PANEL=chol_panel, BLOCK_K=128, BLOCK_N=64,
            FIRST_PANEL=False, num_warps=4, num_stages=2,
        )
        _s9_n2048_n2048_apply_block_kernel[
            (batch, triton.cdiv(remaining, 32), triton.cdiv(trailing, 64))
        ](
            a, h, chol_packed, chol_weights, tau, n, panel_start,
            PANEL_B=chol_panel, BLOCK_B=chol_panel, BLOCK_M=32, BLOCK_N=64,
            FIRST_PANEL=False, GUARDED=False, num_warps=4,
        )

    _s9_n2048_n2048_householder_tail(
        a,
        h,
        tau,
        packed_vectors,
        gram,
        weights,
        route_flags,
        prefix_end,
        False,
        batch,
        n,
    )
    if prefix_end > 0:
        # Re-impose the exact scalar Householder normalization.  In exact
        # arithmetic this is the diagonal of T already; the reduction removes
        # the last Cholesky/triangular-solve rounding drift from the public tau.
        _s9_n2048_n4096_normalize_tau_tiled_kernel[
            (triton.cdiv(prefix_end, 32), batch)
        ](
            h,
            tau,
            0,
            n=n,
            row_block=1024,
            column_block=32,
            num_warps=8,
        )

    # Rejected matrices overwrite the speculative dense result directly from
    # the immutable source.  Dense matrices pay only this kernel's classifier.
    _s9_n2048_n2048_householder_tail(
        a,
        h,
        tau,
        packed_vectors,
        gram,
        weights,
        route_flags,
        fallback_start,
        True,
        batch,
        n,
    )

    return h, tau


@triton.jit
def _s9_n2048_n4096_normalize_tau_tiled_kernel(
    h,
    tau,
    column_offset: tl.constexpr,
    n: tl.constexpr,
    row_block: tl.constexpr,
    column_block: tl.constexpr,
):
    column_start = column_offset + tl.program_id(0) * column_block
    batch_id = tl.program_id(1)
    matrix_base = batch_id * n * n
    columns = column_start + tl.arange(0, column_block)[None, :]
    row_offsets = tl.arange(0, row_block)[:, None]
    norm_squared = tl.zeros((column_block,), tl.float32)
    for row_start in tl.static_range(0, n, row_block):
        # This kernel is launched only on full n4096 1024x32 tiles.  A slab
        # ending before the first column is wholly above the reflector
        # diagonal and contributes exact zeros, so avoid issuing its masked
        # matrix walk.
        if row_start + row_block > column_start:
            rows = row_start + row_offsets
            if row_start >= column_start + column_block:
                # Once the whole slab is below the right edge of this
                # 32-column tile, every reflector entry is active.
                values = tl.load(h + matrix_base + rows * n + columns)
            else:
                # Exactly one retained slab can intersect the diagonal.
                values = tl.load(
                    h + matrix_base + rows * n + columns,
                    mask=rows > columns,
                    other=0.0,
                )
            norm_squared += tl.sum(values * values, axis=0)
    vector_norm_squared = 1.0 + norm_squared
    tau_offsets = batch_id * n + column_start + tl.arange(0, column_block)
    old_tau = tl.load(tau + tau_offsets)
    tau_values = tl.where(old_tau != 0.0, 2.0 / vector_norm_squared, 0.0)
    tl.store(tau + tau_offsets, tau_values)


@triton.jit
def _compensated_fp16_dot_rhs_residual(lhs, rhs):
    """Two-product expansion retaining the residual of the right operand."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    rhs_residual = (rhs - rhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    return result + tl.dot(lhs_high, rhs_residual)


@triton.jit
def _compensated_fp16_dot_lhs_residual(lhs, rhs):
    """Two-product expansion retaining the residual of the left operand."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    lhs_residual = (lhs - lhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    return result + tl.dot(lhs_residual, rhs_high)


@triton.jit
def _v10_mxfp8_quantize_rows_two_term(values):
    """Return primary/residual E4M3 rows and their E8M0 block scales."""
    rows: tl.constexpr = values.shape[0]
    width: tl.constexpr = values.shape[1]
    groups: tl.constexpr = width // 32
    grouped = values.reshape(rows, groups, 32)
    maximum = tl.max(tl.abs(grouped), axis=2)
    maximum_bits = maximum.to(tl.uint32, bitcast=True)
    maximum_exponent = (maximum_bits >> 23) & 255
    maximum_mantissa = maximum_bits & 0x7FFFFF
    scale_exponent = (
        maximum_exponent.to(tl.int32)
        + (maximum_mantissa > 0x600000)
        - 8
    )
    scale_exponent = tl.maximum(1, tl.minimum(253, scale_exponent))
    scale_bits = (scale_exponent.to(tl.uint32) << 23)
    scale = scale_bits.to(tl.float32, bitcast=True)
    inverse_scale_bits = (254 - scale_exponent).to(tl.uint32) << 23
    inverse_scale = inverse_scale_bits.to(tl.float32, bitcast=True)
    primary = (grouped * inverse_scale[:, :, None]).to(tl.float8e4nv)
    residual = grouped - primary.to(tl.float32) * scale[:, :, None]

    residual_maximum = tl.max(tl.abs(residual), axis=2)
    residual_bits = residual_maximum.to(tl.uint32, bitcast=True)
    residual_exponent = (residual_bits >> 23) & 255
    residual_mantissa = residual_bits & 0x7FFFFF
    residual_scale_exponent = (
        residual_exponent.to(tl.int32)
        + (residual_mantissa > 0x600000)
        - 8
    )
    residual_scale_exponent = tl.maximum(
        1, tl.minimum(253, residual_scale_exponent)
    )
    residual_scale_bits = residual_scale_exponent.to(tl.uint32) << 23
    residual_scale = residual_scale_bits.to(tl.float32, bitcast=True)
    residual_inverse_bits = (
        (254 - residual_scale_exponent).to(tl.uint32) << 23
    )
    residual_inverse = residual_inverse_bits.to(tl.float32, bitcast=True)
    residual_value = (
        residual * residual_inverse[:, :, None]
    ).to(tl.float8e4nv)
    return (
        primary.reshape(rows, width),
        scale_exponent.to(tl.uint8),
        residual_value.reshape(rows, width),
        residual_scale_exponent.to(tl.uint8),
    )


@triton.jit
def _v10_mxfp8_two_term_dot(lhs, rhs, result):
    """MXFP8 primary product plus both first-order residual corrections."""
    lhs0, lhs_scale0, lhs1, lhs_scale1 = (
        _v10_mxfp8_quantize_rows_two_term(lhs)
    )
    rhs0_t, rhs_scale0, rhs1_t, rhs_scale1 = (
        _v10_mxfp8_quantize_rows_two_term(
            tl.trans(rhs)
        )
    )
    rhs0 = tl.trans(rhs0_t)
    rhs1 = tl.trans(rhs1_t)
    result = tl.dot_scaled(
        lhs0, lhs_scale0, "e4m3", rhs0, rhs_scale0, "e4m3", result
    )
    result = tl.dot_scaled(
        lhs0, lhs_scale0, "e4m3", rhs1, rhs_scale1, "e4m3", result
    )
    return tl.dot_scaled(
        lhs1, lhs_scale1, "e4m3", rhs0, rhs_scale0, "e4m3", result
    )


def _s9_n1024_n1024_cholesky_factor_32_launch(
    source: torch.Tensor,
    h: torch.Tensor,
    tau: torch.Tensor,
    coefficients: torch.Tensor,
    partials: torch.Tensor,
    bottom_transform: torch.Tensor,
    tail_flags: torch.Tensor,
    route_flags: torch.Tensor,
    batch: int,
    panel_start: int,
    precomputed_chunks: int = 0,
) -> None:
    """Parallel compact-Householder recovery for a guarded n1024 panel."""
    active_chunks = (
        precomputed_chunks
        if precomputed_chunks
        else triton.cdiv(1024 - panel_start, 256)
    )
    first_panel = panel_start == 0
    if not precomputed_chunks:
        _n1024_chol_gram32_kernel[(active_chunks, 1, batch)](
            source,
            h,
            route_flags,
            tail_flags,
            partials,
            panel_start,
            FIRST_PANEL=first_panel,
            DETECT_TAIL=panel_start == 256,
            UNIFORM_ROUTE=first_panel,
            BATCH=batch,
            num_warps=4,
            num_stages=1,
        )
    _n1024_chol32_factor_compact_kernel[(batch,)](
        source,
        h,
        partials,
        coefficients,
        bottom_transform,
        tau,
        route_flags,
        panel_start,
        active_chunks,
        FIRST_PANEL=first_panel,
        num_warps=1,
        num_stages=1,
    )
    bottom_tiles = triton.cdiv(1024 - panel_start - 32, 64)
    _s9_n512_n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
        source,
        h,
        bottom_transform,
        h,
        h,
        route_flags,
        panel_start,
        n=1024,
        PANEL=32,
        BLOCK_M=64,
        FIRST_PANEL=first_panel,
        ROUTE_N512=True,
        REQUIRE_FULL_ACTIVE=True,
        PRECISE_RECOVERY=True,
        STORE_NORMS=False,
        STORE_PACKED_VECTORS=False,
        num_warps=2,
    )


@triton.jit
def _safe512_round_to_tf32(value):
    # Plain TF32 dot inputs may be truncated.  Explicit round-to-nearest keeps
    # the one-MMA path unbiased without paying for the full 3xTF32 expansion.
    return tl.inline_asm_elementwise(
        "cvt.rna.tf32.f32 $0, $1;",
        "=r,r",
        [value],
        dtype=tl.float32,
        is_pure=True,
        pack=1,
    )


@triton.jit
def _safe512_dot_tf32x2_rhs(lhs, rhs, accumulator):
    lhs_big = _safe512_round_to_tf32(lhs)
    rhs_big = _safe512_round_to_tf32(rhs)
    rhs_small = rhs - rhs_big
    accumulator = tl.dot(
        lhs_big, rhs_big, accumulator, input_precision="tf32"
    )
    return tl.dot(
        lhs_big, rhs_small, accumulator, input_precision="tf32"
    )


@triton.jit
def _safe512_compensated_fp16_dot(lhs, rhs):
    """Near-FP32 product using three high-throughput FP16 MMAs."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    lhs_residual = (lhs - lhs_high).to(tl.float16)
    rhs_residual = (rhs - rhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    result += tl.dot(lhs_high, rhs_residual)
    result += tl.dot(lhs_residual, rhs_high)
    return result


@triton.jit
def _safe512_compensated_fp16_dot_rhs_residual(lhs, rhs):
    """Two-product expansion retaining the residual of the right operand."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    rhs_residual = (rhs - rhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    return result + tl.dot(lhs_high, rhs_residual)


@triton.jit
def _safe512_compensated_fp16_dot_lhs_residual(lhs, rhs):
    """Two-product expansion retaining the residual of the left operand."""
    lhs_high = lhs.to(tl.float16)
    rhs_high = rhs.to(tl.float16)
    lhs_residual = (lhs - lhs_high).to(tl.float16)
    result = tl.dot(lhs_high, rhs_high)
    return result + tl.dot(lhs_residual, rhs_high)


@triton.jit
def _safe512_v10_mxfp8_quantize_rows_two_term(values):
    """Return primary/residual E4M3 rows and their E8M0 block scales."""
    rows: tl.constexpr = values.shape[0]
    width: tl.constexpr = values.shape[1]
    groups: tl.constexpr = width // 32
    grouped = values.reshape(rows, groups, 32)
    maximum = tl.max(tl.abs(grouped), axis=2)
    maximum_bits = maximum.to(tl.uint32, bitcast=True)
    maximum_exponent = (maximum_bits >> 23) & 255
    maximum_mantissa = maximum_bits & 0x7FFFFF
    scale_exponent = (
        maximum_exponent.to(tl.int32)
        + (maximum_mantissa > 0x600000)
        - 8
    )
    scale_exponent = tl.maximum(1, tl.minimum(253, scale_exponent))
    scale_bits = (scale_exponent.to(tl.uint32) << 23)
    scale = scale_bits.to(tl.float32, bitcast=True)
    inverse_scale_bits = (254 - scale_exponent).to(tl.uint32) << 23
    inverse_scale = inverse_scale_bits.to(tl.float32, bitcast=True)
    primary = (grouped * inverse_scale[:, :, None]).to(tl.float8e4nv)
    residual = grouped - primary.to(tl.float32) * scale[:, :, None]

    residual_maximum = tl.max(tl.abs(residual), axis=2)
    residual_bits = residual_maximum.to(tl.uint32, bitcast=True)
    residual_exponent = (residual_bits >> 23) & 255
    residual_mantissa = residual_bits & 0x7FFFFF
    residual_scale_exponent = (
        residual_exponent.to(tl.int32)
        + (residual_mantissa > 0x600000)
        - 8
    )
    residual_scale_exponent = tl.maximum(
        1, tl.minimum(253, residual_scale_exponent)
    )
    residual_scale_bits = residual_scale_exponent.to(tl.uint32) << 23
    residual_scale = residual_scale_bits.to(tl.float32, bitcast=True)
    residual_inverse_bits = (
        (254 - residual_scale_exponent).to(tl.uint32) << 23
    )
    residual_inverse = residual_inverse_bits.to(tl.float32, bitcast=True)
    residual_value = (
        residual * residual_inverse[:, :, None]
    ).to(tl.float8e4nv)
    return (
        primary.reshape(rows, width),
        scale_exponent.to(tl.uint8),
        residual_value.reshape(rows, width),
        residual_scale_exponent.to(tl.uint8),
    )


@triton.jit
def _safe512_v10_mxfp8_two_term_dot(lhs, rhs, result):
    """MXFP8 primary product plus both first-order residual corrections."""
    lhs0, lhs_scale0, lhs1, lhs_scale1 = (
        _safe512_v10_mxfp8_quantize_rows_two_term(lhs)
    )
    rhs0_t, rhs_scale0, rhs1_t, rhs_scale1 = (
        _safe512_v10_mxfp8_quantize_rows_two_term(
            tl.trans(rhs)
        )
    )
    rhs0 = tl.trans(rhs0_t)
    rhs1 = tl.trans(rhs1_t)
    result = tl.dot_scaled(
        lhs0, lhs_scale0, "e4m3", rhs0, rhs_scale0, "e4m3", result
    )
    result = tl.dot_scaled(
        lhs0, lhs_scale0, "e4m3", rhs1, rhs_scale1, "e4m3", result
    )
    return tl.dot_scaled(
        lhs1, lhs_scale1, "e4m3", rhs0, rhs_scale0, "e4m3", result
    )


@triton.jit
def _safe512_apply_reflectors_transposed(
    block_t,
    h_ptr,
    tau_ptr,
    rows,
    row_offsets,
    row_mask,
    batch_id,
    base,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
):
    for i in tl.static_range(0, PANEL):
        reflector_col = panel_start + i
        v = tl.load(
            h_ptr + base + rows * N + reflector_col,
            mask=row_mask & (row_offsets > i),
            other=0.0,
        )
        v = tl.where(row_offsets == i, 1.0, v)
        v = tl.where(row_offsets < i, 0.0, v)
        tau = tl.load(tau_ptr + batch_id * N + reflector_col)
        dot = tl.sum(block_t * v[None, :], axis=1)
        block_t = block_t - (tau * dot)[:, None] * v[None, :]
    return block_t


@triton.jit
def _safe512_apply_panel_kernel(
    src_ptr,
    h_ptr,
    tau_ptr,
    residual_tiles_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    RESIDUAL_TILES: tl.constexpr,
    RECORD_COLUMN_ACTIVITY: tl.constexpr,
    SELECT_FLAG: tl.constexpr,
    ASSUME_ACTIVE: tl.constexpr,
):
    batch_id = tl.program_id(0)
    col_block = tl.program_id(1)
    base = batch_id * N * N
    panel_start = tl.multiple_of(panel_start, PANEL)

    row_offsets = tl.arange(0, BLOCK_M)
    col_offsets = tl.arange(0, BLOCK_N)
    rows = panel_start + row_offsets
    cols = panel_start + PANEL + col_block * BLOCK_N + col_offsets
    row_mask = rows < N
    col_mask = cols < N

    matrix_is_selected = True
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    precision_flag = (precision_metadata & 1) != 0
    band_flag = (precision_metadata & 2) != 0
    if SELECT_FLAG == 1:
        matrix_is_selected = (not precision_flag) and (not band_flag)
    elif SELECT_FLAG == 2:
        matrix_is_selected = precision_flag or band_flag
    elif SELECT_FLAG == 3:
        matrix_is_selected = not band_flag
    elif SELECT_FLAG == 4:
        matrix_is_selected = band_flag
    elif SELECT_FLAG == 5:
        matrix_is_selected = precision_flag and (not band_flag)
    elif SELECT_FLAG == 6:
        matrix_is_selected = (precision_metadata & 7) != 0
    if SELECT_FLAG == 2 or SELECT_FLAG == 4 or SELECT_FLAG == 6:
        row_mask = row_mask & ((not band_flag) | (row_offsets < 64))

    if matrix_is_selected:
        if RECORD_COLUMN_ACTIVITY:
            # Save a residual scale per tile, then place post-panel residual
            # norms in future tau entries.  Subsequent factor kernels overwrite
            # the metadata before tau is returned.
            tile_start = panel_start + PANEL + col_block * BLOCK_N
            active_column_end = precision_metadata >> 3
            tile_is_active = tile_start < active_column_end
            if band_flag:
                tile_is_active = tile_start < panel_start + 64
            if tile_is_active:
                block = tl.load(
                    src_ptr + base + rows[:, None] * N + cols[None, :],
                    mask=row_mask[:, None] & col_mask[None, :],
                    other=0.0,
                )
                block_t = tl.trans(block)
                block_t = _safe512_apply_reflectors_transposed(
                    block_t,
                    h_ptr,
                    tau_ptr,
                    rows,
                    row_offsets,
                    row_mask,
                    batch_id,
                    base,
                    panel_start,
                    N,
                    PANEL,
                )
                trailing_rows = row_offsets[None, :] >= PANEL
                residual_column_l1 = tl.sum(
                    tl.where(trailing_rows, tl.abs(block_t), 0.0),
                    axis=1,
                )
                tl.store(
                    tau_ptr + batch_id * N + cols,
                    residual_column_l1,
                    mask=col_mask,
                )
                tl.store(
                    h_ptr + base + rows[:, None] * N + cols[None, :],
                    tl.trans(block_t),
                    mask=row_mask[:, None] & col_mask[None, :],
                )
            else:
                metadata_value = tl.where(band_flag, 1.0, 0.0)
                metadata_repeats: tl.constexpr = BLOCK_N // 32
                for metadata_i in tl.static_range(0, metadata_repeats):
                    tl.store(
                        residual_tiles_ptr
                        + batch_id * RESIDUAL_TILES
                        + metadata_repeats * col_block
                        + metadata_i,
                        metadata_value,
                    )
                tl.store(
                    tau_ptr + batch_id * N + cols,
                    metadata_value,
                    mask=col_mask,
                )
                if not band_flag:
                    tl.store(
                        h_ptr + base + rows[:, None] * N + cols[None, :],
                        0.0,
                        mask=row_mask[:, None] & col_mask[None, :],
                    )
        else:
            tile_start = panel_start + PANEL + col_block * BLOCK_N
            if ASSUME_ACTIVE:
                tile_is_active = True
            else:
                active_column_end = tl.load(
                    tau_ptr + batch_id * N + panel_start + PANEL
                )
                tile_is_active = tile_start < active_column_end
                if band_flag:
                    tile_is_active = tile_start < panel_start + 64
            if tile_is_active:
                block = tl.load(
                    src_ptr + base + rows[:, None] * N + cols[None, :],
                    mask=row_mask[:, None] & col_mask[None, :],
                    other=0.0,
                )
                block_t = tl.trans(block)
                block_t = _safe512_apply_reflectors_transposed(
                    block_t,
                    h_ptr,
                    tau_ptr,
                    rows,
                    row_offsets,
                    row_mask,
                    batch_id,
                    base,
                    panel_start,
                    N,
                    PANEL,
                )
                tl.store(
                    h_ptr + base + rows[:, None] * N + cols[None, :],
                    tl.trans(block_t),
                    mask=row_mask[:, None] & col_mask[None, :],
                )


@triton.jit
def _safe512_apply_panel_wy_fused_kernel(
    src_ptr,
    h_ptr,
    tau_ptr,
    wy_ptr,
    v_half_ptr,
    precision_flags_ptr,
    residual_tiles_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    CHUNK_M: tl.constexpr,
    TF32_MODE: tl.constexpr,
    SELECT_FLAG: tl.constexpr,
    RESIDUAL_TILES: tl.constexpr,
    RECORD_COLUMN_ACTIVITY: tl.constexpr,
    HANDLE_BAND: tl.constexpr,
):
    """Apply ``I - W @ V.T`` while retaining one row chunk at a time."""
    col_block = tl.program_id(0)
    batch_id = tl.program_id(1)
    matrix_base = batch_id * N * N
    wy_base = batch_id * N * PANEL
    v_half_base = batch_id * N * PANEL
    panel_start = tl.multiple_of(panel_start, PANEL)

    chunk_offsets = tl.arange(0, CHUNK_M)
    col_offsets = tl.arange(0, BLOCK_N)
    reflector_offsets = tl.arange(0, PANEL)
    cols = panel_start + PANEL + col_block * BLOCK_N + col_offsets
    # Every caller in the private n512 schedule launches complete 32-column
    # tiles: the initial 32:64 tile, 64:512 precision tiles, one next-panel
    # tile per recursive group, and the terminal 480:512 tile.
    reflector_cols = panel_start + reflector_offsets

    tile_start = panel_start + PANEL + col_block * BLOCK_N
    matrix_is_selected = True
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    band_flag = (precision_metadata & 2) != 0
    precision_flag = (precision_metadata & 1) != 0
    if RECORD_COLUMN_ACTIVITY:
        active_column_end = precision_metadata >> 3
        tile_is_active = tile_start < active_column_end
        if band_flag:
            tile_is_active = tile_start < panel_start + 64
    else:
        active_column_end = tl.load(
            tau_ptr + batch_id * N + panel_start + PANEL
        )
        tile_is_active = tile_start < active_column_end
        if band_flag:
            tile_is_active = tile_start < panel_start + 64
    if SELECT_FLAG == 1:
        matrix_is_selected = (not precision_flag) and (not band_flag)
    elif SELECT_FLAG == 2:
        matrix_is_selected = precision_flag or band_flag
    elif SELECT_FLAG == 3:
        matrix_is_selected = not band_flag
    elif SELECT_FLAG == 4:
        matrix_is_selected = band_flag
    elif SELECT_FLAG == 5:
        matrix_is_selected = precision_flag and (not band_flag)

    if HANDLE_BAND and band_flag and tile_is_active:
        # This specialization is launched for exactly the next 32-column
        # tile.  Reuse it for the band route instead of paying for a second
        # one-CTA launch that reads the same panel and tile.
        rows = panel_start + chunk_offsets
        row_mask = rows < N
        block = tl.load(
            src_ptr + matrix_base + rows[:, None] * N + cols[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        block_t = _safe512_apply_reflectors_transposed(
            tl.trans(block),
            h_ptr,
            tau_ptr,
            rows,
            chunk_offsets,
            row_mask,
            batch_id,
            matrix_base,
            panel_start,
            N,
            PANEL,
        )
        tl.store(
            h_ptr + matrix_base + rows[:, None] * N + cols[None, :],
            tl.trans(block_t),
            mask=row_mask[:, None],
        )
    elif matrix_is_selected and tile_is_active:
        projection = tl.zeros((BLOCK_N, PANEL), dtype=tl.float32)
        for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
            rows = panel_start + row_block + chunk_offsets
            row_mask = rows < N
            block = tl.load(
                src_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                mask=row_mask[:, None],
                other=0.0,
            )
            block_t = tl.trans(block)
            wy = tl.load(
                wy_ptr
                + wy_base
                + rows[:, None] * PANEL
                + reflector_offsets[None, :],
                mask=row_mask[:, None],
                other=0.0,
            )
            if TF32_MODE == 0:
                projection = tl.dot(
                    block_t.to(tl.float16),
                    wy.to(tl.float16),
                    projection,
                )
            elif TF32_MODE == 3 or TF32_MODE == 5:
                projection = tl.dot(
                    block_t,
                    wy.to(tl.float32),
                    projection,
                    input_precision="tf32x3",
                )
            elif TF32_MODE == 2:
                projection = _safe512_dot_tf32x2_rhs(block_t, wy, projection)
            else:
                projection += tl.dot(
                    _safe512_round_to_tf32(block_t),
                    _safe512_round_to_tf32(wy),
                    input_precision="tf32",
                )

        residual_column_l1 = tl.zeros((BLOCK_N,), dtype=tl.float32)
        for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
            rows = panel_start + row_block + chunk_offsets
            row_mask = rows < N
            block = tl.load(
                src_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                mask=row_mask[:, None],
                other=0.0,
            )
            block_t = tl.trans(block)
            if TF32_MODE == 0:
                v_t = tl.load(
                    v_half_ptr
                    + v_half_base
                    + rows[None, :] * PANEL
                    + reflector_offsets[:, None],
                    mask=row_mask[None, :],
                    other=0.0,
                )
            else:
                v_t = tl.load(
                    h_ptr
                    + matrix_base
                    + rows[None, :] * N
                    + reflector_cols[:, None],
                    mask=(rows[None, :] > reflector_cols[:, None])
                    & row_mask[None, :],
                    other=0.0,
                )
                v_t = tl.where(
                    rows[None, :] == reflector_cols[:, None], 1.0, v_t
                )
            if TF32_MODE == 0:
                correction = tl.dot(
                    projection.to(tl.float16),
                    v_t.to(tl.float16),
                )
                block_t -= correction
            elif TF32_MODE == 3 or TF32_MODE == 5:
                correction = tl.dot(
                    projection, v_t.to(tl.float32), input_precision="tf32x3"
                )
                block_t -= correction
            elif TF32_MODE == 2:
                correction = _safe512_dot_tf32x2_rhs(
                    projection,
                    v_t,
                    tl.zeros((BLOCK_N, CHUNK_M), dtype=tl.float32),
                )
                block_t -= correction
            else:
                block_t -= tl.dot(
                    _safe512_round_to_tf32(projection),
                    _safe512_round_to_tf32(v_t),
                    input_precision="tf32",
                )

            if RECORD_COLUMN_ACTIVITY:
                residual_column_l1 += tl.sum(
                    tl.where(
                        rows[None, :] >= panel_start + PANEL,
                        tl.abs(block_t),
                        0.0,
                    ),
                    axis=1,
                )

            tl.store(
                h_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                tl.trans(block_t),
                mask=row_mask[:, None],
            )

        if RECORD_COLUMN_ACTIVITY:
            tl.store(
                tau_ptr + batch_id * N + cols,
                residual_column_l1,
            )
    elif RECORD_COLUMN_ACTIVITY and matrix_is_selected:
        metadata_repeats: tl.constexpr = BLOCK_N // 32
        metadata_value = tl.where(band_flag, 1.0, 0.0)
        for metadata_i in tl.static_range(0, metadata_repeats):
            tl.store(
                residual_tiles_ptr
                + batch_id * RESIDUAL_TILES
                + metadata_repeats * col_block
                + metadata_i,
                metadata_value,
            )
        tl.store(tau_ptr + batch_id * N + cols, metadata_value)
        if not band_flag:
            for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
                rows = panel_start + row_block + chunk_offsets
                row_mask = rows < N
                tl.store(
                    h_ptr + matrix_base + rows[None, :] * N + cols[:, None],
                    0.0,
                    mask=row_mask[None, :],
                )


@triton.jit
def _safe512_form_w_kernel(
    h_ptr,
    tau_ptr,
    t_ptr,
    wy_ptr,
    v_half_ptr,
    precision_w_ptr,
    precision_v_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    ASSUME_ACTIVE: tl.constexpr,
    TF32_MODE: tl.constexpr,
    SELECT_FLAG: tl.constexpr,
    SPLIT_PRECISION: tl.constexpr,
    PRECISE_FORM: tl.constexpr,
    ZERO_PREFIX: tl.constexpr,
):
    """Materialize ``W = V @ T`` once for reuse by all column tiles."""
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    matrix_base = batch_id * N * N
    t_base = batch_id * PANEL * PANEL
    wy_base = batch_id * N * PANEL
    panel_start = tl.multiple_of(panel_start, PANEL)

    row_offsets = tl.arange(0, BLOCK_ROWS)
    panel_offsets = tl.arange(0, PANEL)
    rows = panel_start + row_block * BLOCK_ROWS + row_offsets
    row_mask = rows < N
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    if ASSUME_ACTIVE:
        panel_is_active = True
    else:
        active_column_end = tl.load(
            tau_ptr + batch_id * N + panel_start + PANEL
        )
        panel_is_active = panel_start + PANEL < active_column_end
    if SELECT_FLAG == 1:
        panel_is_active &= (
            ((precision_metadata & 1) == 0)
            & ((precision_metadata & 2) == 0)
        )
    elif SELECT_FLAG == 2:
        panel_is_active &= (precision_metadata & 1) != 0
    elif SELECT_FLAG == 3:
        panel_is_active &= (precision_metadata & 2) == 0
    elif SELECT_FLAG == 6:
        panel_is_active &= (
            ((precision_metadata & 5) != 0)
            & ((precision_metadata & 2) == 0)
        )
    elif SELECT_FLAG == 7:
        panel_is_active &= (
            ((precision_metadata & 4) != 0)
            & ((precision_metadata & 2) == 0)
        )
    if panel_is_active:
        reflector_cols = panel_start + panel_offsets
        if (precision_metadata & 1) != 0:
            # Precision-routed matrices retain the original FP32 reflector
            # tails for their compensated W formation.
            v = tl.load(
                h_ptr
                + matrix_base
                + rows[:, None] * N
                + reflector_cols[None, :],
                mask=row_mask[:, None]
                & (rows[:, None] > reflector_cols[None, :]),
                other=0.0,
            )
            v = tl.where(
                rows[:, None] == reflector_cols[None, :], 1.0, v
            )
        else:
            # The Gram pass (or Cholesky recovery) has already materialized
            # the rounded reflector tile.  Reusing it avoids a second strided
            # read from H and preserves exactly the values used by the MMA.
            if SPLIT_PRECISION and (precision_metadata & 4) != 0:
                v = tl.load(
                    precision_v_ptr
                    + wy_base
                    + rows[:, None] * PANEL
                    + panel_offsets[None, :],
                    mask=row_mask[:, None],
                    other=0.0,
                ).to(tl.float32)
            else:
                v = tl.load(
                    v_half_ptr
                    + wy_base
                    + rows[:, None] * PANEL
                    + panel_offsets[None, :],
                    mask=row_mask[:, None],
                    other=0.0,
                ).to(tl.float32)
        t_factor = tl.load(
            t_ptr
            + t_base
            + panel_offsets[:, None] * PANEL
            + panel_offsets[None, :]
        )
        if TF32_MODE == 3:
            wy = tl.dot(v, t_factor, input_precision="tf32x3")
        elif TF32_MODE == 2:
            wy = _safe512_dot_tf32x2_rhs(
                v,
                t_factor,
                tl.zeros((BLOCK_ROWS, PANEL), dtype=tl.float32),
            )
        else:
            # W is stored as FP16 immediately below.  Classifier-approved
            # matrices can therefore form it directly at that precision;
            # rowscale-sensitive matrices retain rounded TF32 accumulation.
            if (precision_metadata & 1) != 0:
                if PRECISE_FORM:
                    wy = _safe512_dot_tf32x2_rhs(
                        v,
                        t_factor,
                        tl.zeros((BLOCK_ROWS, PANEL), dtype=tl.float32),
                    )
                else:
                    wy = tl.dot(
                        _safe512_round_to_tf32(v),
                        _safe512_round_to_tf32(t_factor),
                        input_precision="tf32",
                    )
            else:
                wy = tl.dot(v.to(tl.float16), t_factor.to(tl.float16))

        output_offsets = wy_base + rows[:, None] * PANEL + panel_offsets[None, :]
        if ZERO_PREFIX and row_block == 0:
            prefix_rows = panel_start - PANEL + panel_offsets
            prefix_offsets = (
                wy_base
                + prefix_rows[:, None] * PANEL
                + panel_offsets[None, :]
            )
            tl.store(wy_ptr + prefix_offsets, 0.0)
            tl.store(v_half_ptr + prefix_offsets, 0.0)
            if SPLIT_PRECISION:
                tl.store(precision_w_ptr + prefix_offsets, 0.0)
                tl.store(precision_v_ptr + prefix_offsets, 0.0)
        if SPLIT_PRECISION:
            precision_output = (precision_metadata & 1) != 0
            tl.store(
                wy_ptr + output_offsets,
                wy,
                mask=row_mask[:, None],
            )
            tl.store(
                v_half_ptr + output_offsets,
                v,
                mask=row_mask[:, None],
            )
            tl.store(
                precision_w_ptr + output_offsets,
                wy,
                mask=row_mask[:, None] & precision_output,
            )
            tl.store(
                precision_v_ptr + output_offsets,
                v,
                mask=row_mask[:, None] & precision_output,
            )
        else:
            tl.store(
                wy_ptr + output_offsets,
                wy,
                mask=row_mask[:, None],
            )


@triton.jit
def _safe512_cross_gram_partial_kernel(
    left_v_ptr,
    right_v_ptr,
    tau_ptr,
    partial_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    EXCLUDE_PRECISION: tl.constexpr,
):
    """Form row-block contributions to ``V0.T @ V1``.

    Two independently factored 32-reflector blocks are recursively coupled
    into one 64-reflector compact-WY transform.  Keeping the cross product in
    32x32 tiles avoids the register and layout hazards of a monolithic 64x64
    Gram tile.
    """
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    workspace_base = batch_id * N * PANEL
    partial_base = (batch_id * 4 + row_block) * PANEL * PANEL
    offsets = tl.arange(0, PANEL)
    row_offsets = tl.arange(0, BLOCK_ROWS)
    rows = panel_start + row_block * BLOCK_ROWS + row_offsets
    row_mask = rows < N

    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    active_column_end = tl.load(
        tau_ptr + batch_id * N + panel_start + 2 * PANEL
    )
    panel_is_active = (
        ((precision_metadata & 2) == 0)
        & ((not EXCLUDE_PRECISION) | ((precision_metadata & 1) == 0))
        & (panel_start + 2 * PANEL < active_column_end)
    )
    if panel_is_active:
        left = tl.load(
            left_v_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        right = tl.load(
            right_v_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        cross = tl.dot(tl.trans(left), right)
        tl.store(
            partial_ptr
            + partial_base
            + offsets[:, None] * PANEL
            + offsets[None, :],
            cross,
        )


@triton.jit
def _safe512_right_gram_and_cross_partial_kernel(
    h_ptr,
    left_v_ptr,
    right_v_ptr,
    tau_ptr,
    gram_partial_ptr,
    cross_partial_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    REUSE_RECOVERED_SAFE: tl.constexpr,
):
    """Share V1 loads between its Gram matrix and the V0/V1 cross Gram."""
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    matrix_base = batch_id * N * N
    partial_base = (batch_id * 4 + row_block) * PANEL * PANEL
    offsets = tl.arange(0, PANEL)
    rows = (
        panel_start
        + row_block * BLOCK_ROWS
        + tl.arange(0, BLOCK_ROWS)
    )
    right_start = panel_start + PANEL
    right_cols = right_start + offsets
    row_mask = rows < N
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    active_column_end = tl.load(
        tau_ptr + batch_id * N + panel_start + 2 * PANEL
    )
    panel_is_active = (
        ((precision_metadata & 2) == 0)
        & (panel_start + 2 * PANEL < active_column_end)
    )
    if panel_is_active:
        left = tl.load(
            left_v_ptr
            + batch_id * N * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        reuse_recovered = REUSE_RECOVERED_SAFE and (
            (precision_metadata & 7) == 0
        )
        right_v_offsets = (
            batch_id * N * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :]
        )
        if reuse_recovered:
            # Compact Cholesky recovery already packed V1.  Only its
            # mathematically zero prefix was not written.
            right_narrow = tl.load(
                right_v_ptr + right_v_offsets,
                mask=row_mask[:, None] & (rows[:, None] >= right_start),
                other=0.0,
            )
            tl.store(
                right_v_ptr + right_v_offsets,
                0.0,
                mask=row_mask[:, None] & (rows[:, None] < right_start),
            )
        else:
            right = tl.load(
                h_ptr
                + matrix_base
                + rows[:, None] * N
                + right_cols[None, :],
                mask=row_mask[:, None]
                & (rows[:, None] > right_cols[None, :]),
                other=0.0,
            )
            right = tl.where(
                rows[:, None] == right_cols[None, :], 1.0, right
            )
            right_narrow = right.to(tl.float16)
            tl.store(
                right_v_ptr + right_v_offsets,
                right_narrow,
                mask=row_mask[:, None],
            )
            right_gram = tl.dot(tl.trans(right_narrow), right_narrow)
            partial_offsets = offsets[:, None] * PANEL + offsets[None, :]
            tl.store(
                gram_partial_ptr + partial_base + partial_offsets,
                right_gram,
            )
        cross = tl.dot(tl.trans(left), right_narrow)
        partial_offsets = offsets[:, None] * PANEL + offsets[None, :]
        tl.store(
            cross_partial_ptr + partial_base + partial_offsets,
            cross,
        )


@triton.jit
def _safe512_aggregate_cross_t_kernel(
    partial_ptr,
    tau_ptr,
    right_t_ptr,
    coupling_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    PARTIALS: tl.constexpr,
    EXCLUDE_PRECISION: tl.constexpr,
):
    """Replace ``T1`` by the recursive coupling ``(V0.T @ V1) @ T1``."""
    batch_id = tl.program_id(0)
    offsets = tl.arange(0, PANEL)
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    active_column_end = tl.load(
        tau_ptr + batch_id * N + panel_start + 2 * PANEL
    )
    panel_is_active = (
        ((precision_metadata & 2) == 0)
        & ((not EXCLUDE_PRECISION) | ((precision_metadata & 1) == 0))
        & (panel_start + 2 * PANEL < active_column_end)
    )
    if panel_is_active:
        partial_base = batch_id * 4 * PANEL * PANEL
        pointers = (
            partial_ptr
            + partial_base
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        cross = tl.load(pointers)
        if PARTIALS >= 2:
            cross += tl.load(pointers + PANEL * PANEL)
        if PARTIALS >= 3:
            cross += tl.load(pointers + 2 * PANEL * PANEL)
        if PARTIALS >= 4:
            cross += tl.load(pointers + 3 * PANEL * PANEL)
        t_pointers = (
            right_t_ptr
            + batch_id * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        right_t = tl.load(t_pointers)
        if (precision_metadata & 1) != 0:
            coupling = _safe512_dot_tf32x2_rhs(
                cross,
                right_t,
                tl.zeros((PANEL, PANEL), dtype=tl.float32),
            )
        else:
            coupling = tl.dot(
                cross.to(tl.float16), right_t.to(tl.float16)
            )
        coupling_pointers = (
            coupling_ptr
            + batch_id * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        tl.store(coupling_pointers, coupling)


@triton.jit
def _safe512_couple_w_kernel(
    left_w_ptr,
    right_w_ptr,
    coupling_ptr,
    tau_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    EXCLUDE_PRECISION: tl.constexpr,
):
    """Materialize ``W1 - W0 @ (V0.T @ V1) @ T1`` in split form."""
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    row_offsets = tl.arange(0, BLOCK_ROWS)
    rows = panel_start + row_block * BLOCK_ROWS + row_offsets
    row_mask = rows < N
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    active_column_end = tl.load(
        tau_ptr + batch_id * N + panel_start + 2 * PANEL
    )
    panel_is_active = (
        ((precision_metadata & 2) == 0)
        & ((not EXCLUDE_PRECISION) | ((precision_metadata & 1) == 0))
        & (panel_start + 2 * PANEL < active_column_end)
    )
    if panel_is_active:
        workspace_base = batch_id * N * PANEL
        left_w = tl.load(
            left_w_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        # W1 starts 32 rows later.  Its missing top block is mathematically
        # zero before the recursive cross correction is applied.
        right_w = tl.load(
            right_w_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        coupling = tl.load(
            coupling_ptr
            + batch_id * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        correction = tl.dot(
            left_w.to(tl.float16),
            coupling.to(tl.float16),
        )
        tl.store(
            right_w_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            right_w - correction,
            mask=row_mask[:, None],
        )


@triton.jit
def _safe512_form_coupled_right_w_kernel(
    h_ptr,
    right_t_ptr,
    left_w_ptr,
    coupling_ptr,
    right_w_ptr,
    right_v_ptr,
    tau_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
):
    """Form the recursively coupled right W block without materializing W1."""
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    row_offsets = tl.arange(0, BLOCK_ROWS)
    rows = panel_start + row_block * BLOCK_ROWS + row_offsets
    right_start = panel_start + PANEL
    right_cols = right_start + offsets
    row_mask = rows < N
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    active_column_end = tl.load(
        tau_ptr + batch_id * N + panel_start + 2 * PANEL
    )
    panel_is_active = (
        ((precision_metadata & 2) == 0)
        & (panel_start + 2 * PANEL < active_column_end)
    )
    if panel_is_active:
        matrix_base = batch_id * N * N
        workspace_base = batch_id * N * PANEL
        if (precision_metadata & 1) != 0:
            v1 = tl.load(
                h_ptr
                + matrix_base
                + rows[:, None] * N
                + right_cols[None, :],
                mask=row_mask[:, None]
                & (rows[:, None] > right_cols[None, :]),
                other=0.0,
            )
            v1 = tl.where(rows[:, None] == right_cols[None, :], 1.0, v1)
        else:
            v1 = tl.load(
                right_v_ptr
                + workspace_base
                + rows[:, None] * PANEL
                + offsets[None, :],
                mask=row_mask[:, None],
                other=0.0,
            ).to(tl.float32)
        t1 = tl.load(
            right_t_ptr
            + batch_id * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        if (precision_metadata & 1) != 0:
            raw_w1 = _safe512_dot_tf32x2_rhs(
                v1,
                t1,
                tl.zeros((BLOCK_ROWS, PANEL), dtype=tl.float32),
            )
        else:
            raw_w1 = tl.dot(v1.to(tl.float16), t1.to(tl.float16))
        left_w = tl.load(
            left_w_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )
        coupling = tl.load(
            coupling_ptr
            + batch_id * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        correction = tl.dot(
            left_w.to(tl.float16), coupling.to(tl.float16)
        )
        tl.store(
            right_w_ptr
            + workspace_base
            + rows[:, None] * PANEL
            + offsets[None, :],
            raw_w1 - correction,
            mask=row_mask[:, None],
        )


@triton.jit
def _safe512_apply_recursive64_wy_kernel(
    src_ptr,
    h_ptr,
    tau_ptr,
    left_w_ptr,
    left_v_ptr,
    right_w_ptr,
    right_v_ptr,
    precision_flags_ptr,
    panel_start,
    STATIC_START: tl.constexpr,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    CHUNK_M: tl.constexpr,
    EXCLUDE_PRECISION: tl.constexpr,
    RECORD_COLUMN_ACTIVITY: tl.constexpr,
    ZERO_INACTIVE: tl.constexpr,
    BATCH: tl.constexpr,
):
    """Apply a recursively split 64-reflector compact-WY transform.

    Both 32-column projections share the same two reads of the matrix tile.
    This halves full trailing-matrix traffic relative to applying the two
    constituent compact-WY blocks in separate kernel launches.
    """
    col_block = tl.program_id(0)
    batch_id = BATCH - 1 - tl.program_id(1)
    matrix_base = batch_id * N * N
    workspace_base = batch_id * N * PANEL
    if STATIC_START >= 0:
        panel_start = STATIC_START
    else:
        panel_start = tl.multiple_of(panel_start, PANEL)
    chunk_offsets = tl.arange(0, CHUNK_M)
    col_offsets = tl.arange(0, BLOCK_N)
    reflector_offsets = tl.arange(0, PANEL)
    cols = panel_start + 2 * PANEL + col_block * BLOCK_N + col_offsets

    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    precision_flag = (precision_metadata & 1) != 0
    band_flag = (precision_metadata & 2) != 0
    if RECORD_COLUMN_ACTIVITY:
        active_column_end = precision_metadata >> 3
    else:
        active_column_end = tl.load(
            tau_ptr + batch_id * N + panel_start + 2 * PANEL
        )
    tile_is_active = (
        (not band_flag)
        & (panel_start + 2 * PANEL + col_block * BLOCK_N < active_column_end)
    )
    matrix_is_selected = not band_flag
    if EXCLUDE_PRECISION:
        matrix_is_selected &= not precision_flag
    if band_flag:
        # The first recursive column tile begins immediately after Q1, which
        # is also the sole live trailing tile for bandwidth-16 matrices.
        # Let that CTA perform Q1's exact update and eliminate a separate
        # band-only launch.  Other CTAs in the grid return immediately.
        band_tile_is_active = (
            (col_block == 0)
            & (panel_start + 2 * PANEL < active_column_end)
        )
        if band_tile_is_active:
            band_offsets = tl.arange(0, PANEL)
            band_rows = panel_start + PANEL + chunk_offsets
            band_cols = panel_start + 2 * PANEL + band_offsets
            band_row_mask = band_rows < N
            band_col_mask = band_cols < N
            band_block = tl.load(
                h_ptr
                + matrix_base
                + band_rows[:, None] * N
                + band_cols[None, :],
                mask=band_row_mask[:, None] & band_col_mask[None, :],
                other=0.0,
            )
            band_block_t = _safe512_apply_reflectors_transposed(
                tl.trans(band_block),
                h_ptr,
                tau_ptr,
                band_rows,
                chunk_offsets,
                band_row_mask,
                batch_id,
                matrix_base,
                panel_start + PANEL,
                N,
                PANEL,
            )
            tl.store(
                h_ptr
                + matrix_base
                + band_rows[:, None] * N
                + band_cols[None, :],
                tl.trans(band_block_t),
                mask=band_row_mask[:, None] & band_col_mask[None, :],
            )
    elif tile_is_active and matrix_is_selected:
        coupled_projection = tl.zeros((BLOCK_N, 2 * PANEL), dtype=tl.float32)
        residual_column_l1 = tl.zeros((BLOCK_N,), dtype=tl.float32)
        for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
            rows = panel_start + row_block + chunk_offsets
            block = tl.load(
                src_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                cache_modifier=".ca",
            )
            block_t = tl.trans(block)
            left_w = tl.load(
                left_w_ptr
                + workspace_base
                + rows[:, None] * PANEL
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            right_w = tl.load(
                right_w_ptr
                + workspace_base
                + rows[:, None] * PANEL
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            coupled_w = tl.cat(left_w, right_w, dim=1)
            coupled_projection = tl.dot(
                block_t.to(tl.float16),
                coupled_w.to(tl.float16),
                coupled_projection,
            )

        for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
            if BLOCK_M == 512 or BLOCK_M == 384:
                # Large tiles exceed useful L1 capacity. Revisit the newest
                # independent update chunk first without changing reduction
                # order in the projection pass above.
                rows = (
                    panel_start
                    + BLOCK_M
                    - CHUNK_M
                    - row_block
                    + chunk_offsets
                )
            else:
                rows = panel_start + row_block + chunk_offsets
            block = tl.load(
                src_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                cache_modifier=".ca",
            )
            block_t = tl.trans(block)
            left_v_t = tl.load(
                left_v_ptr
                + workspace_base
                + rows[None, :] * PANEL
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            right_v_t = tl.load(
                right_v_ptr
                + workspace_base
                + rows[None, :] * PANEL
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            coupled_v_t = tl.cat(left_v_t, right_v_t, dim=0)
            block_t -= tl.dot(
                coupled_projection.to(tl.float16),
                coupled_v_t.to(tl.float16),
            )
            if RECORD_COLUMN_ACTIVITY:
                residual_column_l1 += tl.sum(
                    tl.where(
                        rows[None, :] >= panel_start + 2 * PANEL,
                        tl.abs(block_t),
                        0.0,
                    ),
                    axis=1,
                )
            tl.store(
                h_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                tl.trans(block_t),
            )
        if RECORD_COLUMN_ACTIVITY:
            tl.store(
                tau_ptr + batch_id * N + cols,
                residual_column_l1,
            )
    elif (RECORD_COLUMN_ACTIVITY or ZERO_INACTIVE) and matrix_is_selected:
        if RECORD_COLUMN_ACTIVITY:
            tl.store(tau_ptr + batch_id * N + cols, 0.0)
        for row_block in tl.static_range(0, BLOCK_M, CHUNK_M):
            rows = panel_start + row_block + chunk_offsets
            tl.store(
                h_ptr + matrix_base + rows[:, None] * N + cols[None, :],
                0.0,
            )


@triton.jit
def _safe512_apply_selected_interleaved_wy_kernel(
    h_ptr,
    tau_ptr,
    left_w_ptr,
    left_v_ptr,
    right_w_ptr,
    right_v_ptr,
    precision_flags_ptr,
    GROUP_START: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BATCH: tl.constexpr,
):
    """Apply selected recursive WY groups with conflict-free operands.

    Keeping these statically shaped specializations separate prevents their
    interleaved shared layout from inflating compile time for the remaining
    recursive geometries.
    """
    col_block = tl.program_id(0)
    batch_id = BATCH - 1 - tl.program_id(1)
    matrix_base = batch_id * 512 * 512
    workspace_base = batch_id * 512 * 32
    chunk_offsets = tl.arange(0, 64)
    col_offsets = tl.arange(0, 64)
    reflector_offsets = tl.arange(0, 32)
    cols = GROUP_START + 64 + col_block * 64 + col_offsets

    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    band_flag = (precision_metadata & 2) != 0
    # Factor kernels seed every future tau slot with this exact endpoint and
    # overwrite only the two panels immediately preceding this update.
    # Reading the already-live classifier word avoids a redundant scalar
    # global dependency in every trailing-tile CTA.
    active_column_end = precision_metadata >> 3
    tile_is_active = (not band_flag) & (
        GROUP_START + 64 + col_block * 64 < active_column_end
    )

    if band_flag:
        # As in the generic kernel, tile zero performs Q1's sole live band
        # update; the remaining programs return without touching H.
        if (col_block == 0) & (GROUP_START + 64 < active_column_end):
            band_offsets = tl.arange(0, 32)
            band_rows = GROUP_START + 32 + chunk_offsets
            band_cols = GROUP_START + 64 + band_offsets
            band_block = tl.load(
                h_ptr
                + matrix_base
                + band_rows[:, None] * 512
                + band_cols[None, :]
            )
            band_block_t = _safe512_apply_reflectors_transposed(
                tl.trans(band_block),
                h_ptr,
                tau_ptr,
                band_rows,
                chunk_offsets,
                tl.full((64,), True, tl.int1),
                batch_id,
                matrix_base,
                GROUP_START + 32,
                512,
                32,
            )
            tl.store(
                h_ptr
                + matrix_base
                + band_rows[:, None] * 512
                + band_cols[None, :],
                tl.trans(band_block_t),
            )
    elif tile_is_active:
        coupled_projection = tl.zeros((64, 64), dtype=tl.float32)
        for row_block in tl.static_range(0, BLOCK_M, 64):
            rows = GROUP_START + row_block + chunk_offsets
            left_w = tl.load(
                left_w_ptr
                + workspace_base
                + rows[:, None] * 32
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            right_w = tl.load(
                right_w_ptr
                + workspace_base
                + rows[:, None] * 32
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            # Applying the same reflector permutation to W and V preserves
            # (C.T @ W) @ V.T while changing TCGen5's shared-store mapping.
            coupled_w = tl.interleave(left_w, right_w)
            block_t = tl.trans(tl.load(
                h_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                cache_modifier=".ca",
            ))
            coupled_projection = tl.dot(
                block_t.to(tl.float16),
                coupled_w.to(tl.float16),
                coupled_projection,
            )

        for row_block in tl.static_range(0, BLOCK_M, 64):
            rows = (
                GROUP_START + BLOCK_M - 64 - row_block + chunk_offsets
            )
            left_v_t = tl.load(
                left_v_ptr
                + workspace_base
                + rows[None, :] * 32
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            right_v_t = tl.load(
                right_v_ptr
                + workspace_base
                + rows[None, :] * 32
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            coupled_v_t = tl.trans(
                tl.interleave(tl.trans(left_v_t), tl.trans(right_v_t))
            )
            correction = tl.dot(
                coupled_projection.to(tl.float16),
                coupled_v_t.to(tl.float16),
            )
            block_t = tl.trans(tl.load(
                h_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                cache_modifier=".ca",
            ))
            block_t -= correction
            tl.store(
                h_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                tl.trans(block_t),
            )


@triton.jit
def _safe512_apply_initial_interleaved_wy_kernel(
    a_ptr,
    h_ptr,
    tau_ptr,
    left_w_ptr,
    left_v_ptr,
    right_w_ptr,
    right_v_ptr,
    precision_flags_ptr,
    BATCH: tl.constexpr,
):
    """Apply the initial paired WY block with its exact route semantics."""
    col_block = tl.program_id(0)
    batch_id = BATCH - 1 - tl.program_id(1)
    matrix_base = batch_id * 512 * 512
    workspace_base = batch_id * 512 * 32
    chunk_offsets = tl.arange(0, 64)
    col_offsets = tl.arange(0, 64)
    reflector_offsets = tl.arange(0, 32)
    cols = 64 + col_block * 64 + col_offsets

    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    precision_flag = (precision_metadata & 1) != 0
    band_flag = (precision_metadata & 2) != 0
    # tau[64] still contains the same seeded endpoint at this handoff.
    active_column_end = precision_metadata >> 3
    tile_is_active = (not band_flag) & (
        64 + col_block * 64 < active_column_end
    )
    matrix_is_selected = (not band_flag) & (not precision_flag)

    if band_flag:
        if (col_block == 0) & (64 < active_column_end):
            band_offsets = tl.arange(0, 32)
            band_rows = 32 + chunk_offsets
            band_cols = 64 + band_offsets
            band_block = tl.load(
                h_ptr
                + matrix_base
                + band_rows[:, None] * 512
                + band_cols[None, :]
            )
            band_block_t = _safe512_apply_reflectors_transposed(
                tl.trans(band_block),
                h_ptr,
                tau_ptr,
                band_rows,
                chunk_offsets,
                tl.full((64,), True, tl.int1),
                batch_id,
                matrix_base,
                32,
                512,
                32,
            )
            tl.store(
                h_ptr
                + matrix_base
                + band_rows[:, None] * 512
                + band_cols[None, :],
                tl.trans(band_block_t),
            )
    elif tile_is_active and matrix_is_selected:
        coupled_projection = tl.zeros((64, 64), dtype=tl.float32)
        for row_block in tl.static_range(0, 512, 64):
            rows = row_block + chunk_offsets
            left_w = tl.load(
                left_w_ptr
                + workspace_base
                + rows[:, None] * 32
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            right_w = tl.load(
                right_w_ptr
                + workspace_base
                + rows[:, None] * 32
                + reflector_offsets[None, :],
                cache_modifier=".cg",
            )
            coupled_w = tl.interleave(left_w, right_w)
            block_t = tl.trans(tl.load(
                a_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                cache_modifier=".ca",
            ))
            coupled_projection = tl.dot(
                block_t.to(tl.float16),
                coupled_w.to(tl.float16),
                coupled_projection,
            )

        for row_block in tl.static_range(0, 512, 64):
            rows = 512 - 64 - row_block + chunk_offsets
            left_v_t = tl.load(
                left_v_ptr
                + workspace_base
                + rows[None, :] * 32
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            right_v_t = tl.load(
                right_v_ptr
                + workspace_base
                + rows[None, :] * 32
                + reflector_offsets[:, None],
                cache_modifier=".cg",
            )
            coupled_v_t = tl.trans(
                tl.interleave(tl.trans(left_v_t), tl.trans(right_v_t))
            )
            correction = tl.dot(
                coupled_projection.to(tl.float16), coupled_v_t.to(tl.float16)
            )
            block_t = tl.trans(tl.load(
                a_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                cache_modifier=".ca",
            ))
            block_t -= correction
            tl.store(
                h_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                tl.trans(block_t),
            )
    elif matrix_is_selected:
        for row_block in tl.static_range(0, 512, 64):
            rows = row_block + chunk_offsets
            tl.store(
                h_ptr + matrix_base + rows[:, None] * 512 + cols[None, :],
                0.0,
            )


@triton.jit
def _safe512_gram_partial_kernel(
    h_ptr,
    packed_v_ptr,
    tau_ptr,
    partial_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_ROWS: tl.constexpr,
    ASSUME_ACTIVE: tl.constexpr,
    SKIP_FLAGS: tl.constexpr,
    REQUIRE_PRECISION: tl.constexpr,
):
    """Compute one row block of the panel Gram matrix with FP16 MMA.

    Householder tails are bounded, so FP16 needs no scale metadata here.  The
    sensitive reduction across row blocks and triangular recurrence stay FP32.
    """
    batch_id = tl.program_id(0)
    row_block = tl.program_id(1)
    base = batch_id * N * N
    partial_base = (batch_id * 4 + row_block) * PANEL * PANEL
    row_offsets = tl.arange(0, BLOCK_ROWS)
    panel_offsets = tl.arange(0, PANEL)
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    if (precision_metadata & SKIP_FLAGS) != 0:
        return
    if REQUIRE_PRECISION and (precision_metadata & 5) == 0:
        return
    rows = panel_start + row_block * BLOCK_ROWS + row_offsets
    cols = panel_start + panel_offsets
    row_mask = rows < N
    if ASSUME_ACTIVE:
        panel_is_active = True
    else:
        active_column_end = tl.load(
            tau_ptr + batch_id * N + panel_start + PANEL
        )
        panel_is_active = panel_start + PANEL < active_column_end

    if panel_is_active:
        v = tl.load(
            h_ptr + base + rows[:, None] * N + cols[None, :],
            mask=row_mask[:, None] & (rows[:, None] > cols[None, :]),
            other=0.0,
        )
        v = tl.where(rows[:, None] == cols[None, :], 1.0, v)
        v_narrow = v.to(tl.float16)
        tl.store(
            packed_v_ptr
            + batch_id * N * PANEL
            + rows[:, None] * PANEL
            + panel_offsets[None, :],
            v_narrow,
            mask=row_mask[:, None],
        )
        gram = tl.dot(tl.trans(v_narrow), v_narrow)
        tl.store(
            partial_ptr
            + partial_base
            + panel_offsets[:, None] * PANEL
            + panel_offsets[None, :],
            gram,
        )


@triton.jit
def _safe512_aggregate_t_kernel(
    partial_ptr,
    tau_ptr,
    t_ptr,
    cross_partial_ptr,
    coupling_ptr,
    precision_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    PARTIALS: tl.constexpr,
    ASSUME_ACTIVE: tl.constexpr,
    SKIP_FLAGS: tl.constexpr,
    REQUIRE_PRECISION: tl.constexpr,
    BLOCKED_T: tl.constexpr,
    BUILD_COUPLING: tl.constexpr,
    REUSE_RECOVERED_T: tl.constexpr,
):
    """Reduce partial Gram blocks and form the compact-WY triangular T."""
    batch_id = tl.program_id(0)
    partial_base = batch_id * 4 * PANEL * PANEL
    t_base = batch_id * PANEL * PANEL
    panel_offsets = tl.arange(0, PANEL)
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    if (precision_metadata & SKIP_FLAGS) != 0:
        return
    if (
        REQUIRE_PRECISION
        and (precision_metadata & 5) == 0
        and not (
            REUSE_RECOVERED_T and ((precision_metadata & 7) == 0)
        )
    ):
        return
    if ASSUME_ACTIVE:
        panel_is_active = True
    else:
        active_column_end = tl.load(
            tau_ptr + batch_id * N + panel_start + PANEL
        )
        panel_is_active = panel_start + PANEL < active_column_end

    if (
        REUSE_RECOVERED_T
        and panel_is_active
        and ((precision_metadata & 7) == 0)
    ):
        # Cholesky recovery already produced T. Only the recursive cross
        # coupling remains; avoid reducing an unused Gram and rebuilding T.
        if BUILD_COUPLING:
            cross_ptrs = (
                cross_partial_ptr
                + partial_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :]
            )
            cross = tl.load(cross_ptrs)
            if PARTIALS >= 2:
                cross += tl.load(cross_ptrs + PANEL * PANEL)
            if PARTIALS >= 3:
                cross += tl.load(cross_ptrs + 2 * PANEL * PANEL)
            if PARTIALS >= 4:
                cross += tl.load(cross_ptrs + 3 * PANEL * PANEL)
            recovered_t = tl.load(
                t_ptr
                + t_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :]
            )
            coupling = tl.dot(
                cross.to(tl.float16), recovered_t.to(tl.float16)
            )
            tl.store(
                coupling_ptr
                + t_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :],
                coupling,
            )
        return

    if panel_is_active:
        if BLOCKED_T:
            half = tl.arange(0, 16)
            half_r = half[:, None]
            half_c = half[None, :]
            gram00_ptrs = partial_ptr + partial_base + half_r * PANEL + half_c
            gram01_ptrs = gram00_ptrs + 16
            gram11_ptrs = (
                partial_ptr
                + partial_base
                + (16 + half_r) * PANEL
                + 16
                + half_c
            )
            gram00 = tl.load(gram00_ptrs)
            gram01 = tl.load(gram01_ptrs)
            gram11 = tl.load(gram11_ptrs)
            if PARTIALS >= 2:
                gram00 += tl.load(gram00_ptrs + PANEL * PANEL)
                gram01 += tl.load(gram01_ptrs + PANEL * PANEL)
                gram11 += tl.load(gram11_ptrs + PANEL * PANEL)
            if PARTIALS >= 3:
                gram00 += tl.load(gram00_ptrs + 2 * PANEL * PANEL)
                gram01 += tl.load(gram01_ptrs + 2 * PANEL * PANEL)
                gram11 += tl.load(gram11_ptrs + 2 * PANEL * PANEL)
            if PARTIALS >= 4:
                gram00 += tl.load(gram00_ptrs + 3 * PANEL * PANEL)
                gram01 += tl.load(gram01_ptrs + 3 * PANEL * PANEL)
                gram11 += tl.load(gram11_ptrs + 3 * PANEL * PANEL)

            t00 = tl.zeros((16, 16), dtype=tl.float32)
            t11 = tl.zeros((16, 16), dtype=tl.float32)
            for i in tl.static_range(0, 16):
                tau0 = tl.load(tau_ptr + batch_id * N + panel_start + i)
                gram_column0 = tl.sum(
                    tl.where(half[None, :] == i, gram00, 0.0), axis=1
                )
                t_column0 = tl.where(
                    half < i,
                    -tau0 * tl.sum(t00 * gram_column0[None, :], axis=1),
                    tl.where(half == i, tau0, 0.0),
                )
                t00 = tl.where(
                    half[None, :] == i, t_column0[:, None], t00
                )

                tau1 = tl.load(
                    tau_ptr + batch_id * N + panel_start + 16 + i
                )
                gram_column1 = tl.sum(
                    tl.where(half[None, :] == i, gram11, 0.0), axis=1
                )
                t_column1 = tl.where(
                    half < i,
                    -tau1 * tl.sum(t11 * gram_column1[None, :], axis=1),
                    tl.where(half == i, tau1, 0.0),
                )
                t11 = tl.where(
                    half[None, :] == i, t_column1[:, None], t11
                )

            off_diagonal = -tl.dot(
                tl.dot(t00, gram01, input_precision="tf32x3"),
                t11,
                input_precision="tf32x3",
            )
            zero16 = tl.zeros((16, 16), dtype=tl.float32)
            t_factor = tl.cat(
                tl.cat(t00, off_diagonal, dim=1),
                tl.cat(zero16, t11, dim=1),
                dim=0,
            )
        else:
            t_factor = tl.zeros((PANEL, PANEL), dtype=tl.float32)
            gram_ptrs = (
                partial_ptr
                + partial_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :]
            )
            gram = tl.load(gram_ptrs)
            if PARTIALS >= 2:
                gram += tl.load(gram_ptrs + PANEL * PANEL)
            if PARTIALS >= 3:
                gram += tl.load(gram_ptrs + 2 * PANEL * PANEL)
            if PARTIALS >= 4:
                gram += tl.load(gram_ptrs + 3 * PANEL * PANEL)
            for i in tl.static_range(0, PANEL):
                tau = tl.load(tau_ptr + batch_id * N + panel_start + i)
                gram_column = tl.sum(
                    tl.where(panel_offsets[None, :] == i, gram, 0.0), axis=1
                )
                t_column = tl.where(
                    panel_offsets < i,
                    -tau * tl.sum(t_factor * gram_column[None, :], axis=1),
                    tl.where(panel_offsets == i, tau, 0.0),
                )
                t_factor = tl.where(
                    panel_offsets[None, :] == i,
                    t_column[:, None],
                    t_factor,
                )
        if BUILD_COUPLING:
            cross_ptrs = (
                cross_partial_ptr
                + partial_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :]
            )
            cross = tl.load(cross_ptrs)
            if PARTIALS >= 2:
                cross += tl.load(cross_ptrs + PANEL * PANEL)
            if PARTIALS >= 3:
                cross += tl.load(cross_ptrs + 2 * PANEL * PANEL)
            if PARTIALS >= 4:
                cross += tl.load(cross_ptrs + 3 * PANEL * PANEL)
            if (precision_metadata & 1) != 0:
                coupling = _safe512_dot_tf32x2_rhs(
                    cross,
                    t_factor,
                    tl.zeros((PANEL, PANEL), dtype=tl.float32),
                )
            else:
                coupling = tl.dot(
                    cross.to(tl.float16), t_factor.to(tl.float16)
                )
            tl.store(
                coupling_ptr
                + t_base
                + panel_offsets[:, None] * PANEL
                + panel_offsets[None, :],
                coupling,
            )
        tl.store(
            t_ptr
            + t_base
            + panel_offsets[:, None] * PANEL
            + panel_offsets[None, :],
            t_factor,
        )


@triton.jit
def _safe512_classify_precision_kernel(
    a_ptr,
    flags_ptr,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    DROP_RELATIVE_L1: tl.constexpr,
):
    batch_id = tl.program_id(0)
    offsets = tl.arange(0, N)
    base = batch_id * N * N
    first_row = tl.load(a_ptr + base + offsets)
    last_row = tl.load(a_ptr + base + (N - 1) * N + offsets)
    second_row = tl.load(a_ptr + base + N + offsets)
    penultimate_row = tl.load(a_ptr + base + (N - 2) * N + offsets)
    first_row_abs = tl.abs(first_row)
    last_row_abs = tl.abs(last_row)
    first_row_l1 = tl.sum(first_row_abs, axis=0)
    last_row_l1 = tl.sum(last_row_abs, axis=0)
    precision_flag = (last_row_l1 < first_row_l1 * 0.03125) | (
        first_row_l1 < last_row_l1 * 0.03125
    )
    # Near-collinear inputs have rows that are almost proportional across
    # columns (including the public column scaling).  Route them away from
    # normal equations using a sign-invariant row cosine test.
    row_cross = tl.sum(first_row * last_row, axis=0)
    first_row_l2 = tl.sum(first_row * first_row, axis=0)
    last_row_l2 = tl.sum(last_row * last_row, axis=0)
    correlation_squared = (row_cross * row_cross) / tl.maximum(
        first_row_l2 * last_row_l2, 1.0e-30
    )
    top_cross = tl.sum(first_row * second_row, axis=0)
    top_correlation_squared = (top_cross * top_cross) / tl.maximum(
        first_row_l2 * tl.sum(second_row * second_row, axis=0), 1.0e-30
    )
    bottom_cross = tl.sum(last_row * penultimate_row, axis=0)
    bottom_correlation_squared = (bottom_cross * bottom_cross) / tl.maximum(
        last_row_l2 * tl.sum(penultimate_row * penultimate_row, axis=0),
        1.0e-30,
    )
    nearcollinear_flag = (
        (correlation_squared > 0.25)
        | (top_correlation_squared > 0.25)
        | (bottom_correlation_squared > 0.25)
    )

    # A near-rank profile can have a trailing column that is an almost scaled
    # copy of an early column while its rows still look unstructured.  Detect
    # that true mathematical dependency with a sign- and scale-invariant
    # column cosine and route it to the stable Householder implementation.
    rows = tl.arange(0, N)
    head_column = tl.load(a_ptr + base + rows * N)
    tail_column = tl.load(a_ptr + base + rows * N + 3 * (N // 4))
    column_cross = tl.sum(head_column * tail_column, axis=0)
    head_norm = tl.sum(head_column * head_column, axis=0)
    tail_norm = tl.sum(tail_column * tail_column, axis=0)
    column_correlation_squared = (column_cross * column_cross) / tl.maximum(
        head_norm * tail_norm, 1.0e-30
    )
    balanced_column_norms = tl.minimum(head_norm, tail_norm) > (
        tl.maximum(head_norm, tail_norm) * 0.015625
    )
    precision_flag |= (column_correlation_squared > 0.99) & balanced_column_norms

    first_row_peak = tl.max(first_row_abs, axis=0)
    last_row_peak = tl.max(last_row_abs, axis=0)
    first_row_count = tl.sum(
        tl.where(first_row_abs > first_row_peak * 1.0e-6, 1, 0),
        axis=0,
    )
    last_row_count = tl.sum(
        tl.where(last_row_abs > last_row_peak * 1.0e-6, 1, 0),
        axis=0,
    )
    first_support_end = tl.max(
        tl.where(first_row_abs > first_row_peak * 1.0e-6, offsets + 1, 0),
        axis=0,
    )
    last_support_start = tl.min(
        tl.where(last_row_abs > last_row_peak * 1.0e-6, offsets, N),
        axis=0,
    )
    sparse_boundary = (
        (first_row_count >= 8)
        & (first_row_count <= 64)
        & (last_row_count >= 8)
        & (last_row_count <= 64)
    )
    band_flag = (
        sparse_boundary
        & (first_support_end <= 64)
        & (last_support_start >= N - 64)
    )
    # A row/column transform can turn the diagonal band into an anti-band.
    # It is still sparse and well defined, but the diagonal-band copy kernel
    # is no longer valid; keep it on the general stable route instead.
    precision_flag |= sparse_boundary & ~band_flag
    signal = first_row_abs + last_row_abs
    matrix_scale = tl.max(signal, axis=0)
    active = signal > matrix_scale * DROP_RELATIVE_L1
    active_column_end = tl.max(tl.where(active, offsets + 1, 0), axis=0)
    active_column_end = (
        (active_column_end + PANEL - 1) // PANEL
    ) * PANEL
    active_column_end = tl.maximum(active_column_end, PANEL)
    active_column_end = tl.where(active_column_end == 288, 256, active_column_end)
    tl.store(
        flags_ptr + batch_id,
        active_column_end * 8
        + nearcollinear_flag * 4
        + band_flag * 2
        + precision_flag,
    )


@triton.jit
def _safe512_s9_n1024_n1024_require_uniform_dense_kernel(
    flags_ptr,
    N: tl.constexpr,
    BATCH: tl.constexpr,
    BLOCK: tl.constexpr,
):
    """Keep the fast route only when every matrix satisfies its guard."""
    offsets = tl.arange(0, BLOCK)
    valid = offsets < BATCH
    flags = tl.load(flags_ptr + offsets, mask=valid, other=N * 8)
    selected = (flags == N * 8) | ~valid
    all_selected = tl.min(selected.to(tl.int32), axis=0) != 0
    # Setting a low classifier bit makes both the Cholesky route and the
    # direct-route skip test choose the stable path for a heterogeneous batch.
    tl.store(
        flags_ptr + offsets,
        tl.where(all_selected, flags, flags | 1),
        mask=valid,
    )


@triton.jit
def _safe512_v10_n512_build_late_route_flags(
    source_flags,
    late_flags,
    BATCH: tl.constexpr,
    BLOCK: tl.constexpr,
):
    """Enable late normal equations only for a uniformly stable batch."""
    offsets = tl.arange(0, BLOCK)
    valid = offsets < BATCH
    flags = tl.load(source_flags + offsets, mask=valid, other=1)
    stable = ((flags & 7) == 0) & ((flags >> 3) > 384)
    all_stable = tl.min(tl.where(valid, stable, True).to(tl.int32), axis=0) != 0
    tl.store(
        late_flags + offsets,
        tl.where(all_stable, flags, flags | 1),
        mask=valid,
    )


@triton.jit
def _safe512_copy_band_initial_kernel(
    src_ptr,
    h_ptr,
    precision_flags_ptr,
    N: tl.constexpr,
    BLOCK_N: tl.constexpr,
    CHUNK_M: tl.constexpr,
):
    col_block = tl.program_id(0)
    batch_id = tl.program_id(1)
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    band_flag = (precision_metadata & 2) != 0
    if band_flag:
        base = batch_id * N * N
        col_offsets = tl.arange(0, BLOCK_N)
        row_offsets = tl.arange(0, CHUNK_M)
        cols = col_block * BLOCK_N + col_offsets
        col_mask = cols < N
        for row_block in tl.static_range(0, N, CHUNK_M):
            rows = row_block + row_offsets
            block = tl.load(
                src_ptr + base + rows[None, :] * N + cols[:, None],
                mask=col_mask[:, None]
                & (tl.abs(rows[None, :] - cols[:, None]) <= 16),
                other=0.0,
            )
            tl.store(
                h_ptr + base + rows[None, :] * N + cols[:, None],
                block,
                mask=col_mask[:, None],
            )


@triton.jit
def _safe512_factor_panel_kernel(
    src_ptr,
    h_ptr,
    tau_ptr,
    gram_ptr,
    precision_flags_ptr,
    normalization_flags_ptr,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    CHECK_COLUMN_ACTIVITY: tl.constexpr,
    BUILD_GRAM: tl.constexpr,
    SELECT_FLAG: tl.constexpr,
    SEED_ACTIVITY: tl.constexpr,
    NORMALIZE_CHOLESKY: tl.constexpr,
):
    batch_id = tl.program_id(0)
    base = batch_id * N * N
    panel_start = tl.multiple_of(panel_start, PANEL)
    row_offsets = tl.arange(0, BLOCK_M)
    col_offsets = tl.arange(0, PANEL)
    rows = panel_start + row_offsets
    cols = panel_start + col_offsets
    row_mask = rows < N
    precision_metadata = tl.load(precision_flags_ptr + batch_id)
    if SEED_ACTIVITY:
        # Seed future panel metadata through one ordinary rank-1 store.  Each
        # factor overwrites only its own 32 tau entries, leaving the next
        # panel's conservative endpoint intact.
        metadata_offsets = tl.arange(0, N)
        tl.store(
            tau_ptr + batch_id * N + metadata_offsets,
            precision_metadata >> 3,
            mask=metadata_offsets >= 32,
        )
    band_flag = (precision_metadata & 2) != 0
    matrix_is_selected = True
    if SELECT_FLAG == 2:
        matrix_is_selected = (precision_metadata & 7) != 0
    row_mask = row_mask & ((not band_flag) | (row_offsets < 64))
    gram_base = batch_id * PANEL * PANEL

    if CHECK_COLUMN_ACTIVITY:
        active_column_end = tl.load(tau_ptr + batch_id * N + panel_start)
        panel_is_active = (panel_start < active_column_end) & matrix_is_selected
    else:
        panel_is_active = matrix_is_selected

    if panel_is_active:
        panel = tl.load(
            src_ptr + base + rows[:, None] * N + cols[None, :],
            mask=row_mask[:, None],
            other=0.0,
        )

        for i in tl.static_range(0, PANEL):
            col_i = tl.sum(tl.where(col_offsets[None, :] == i, panel, 0.0), axis=1)
            active = row_offsets >= i
            alpha = tl.sum(tl.where(row_offsets == i, col_i, 0.0), axis=0)
            tail = tl.where(row_offsets > i, col_i, 0.0)
            tail_ss = tl.sum(tail * tail, axis=0)
            tail_norm = tl.sqrt(tail_ss)
            full_norm = tl.sqrt(alpha * alpha + tail_ss)

            beta_raw = tl.where(alpha >= 0.0, -full_norm, full_norm)
            beta = tl.where(tail_norm == 0.0, alpha, beta_raw)
            tau_raw = (beta - alpha) / beta
            tau = tl.where(tail_norm == 0.0, 0.0, tau_raw)
            denom = alpha - beta
            v_tail = tl.where(tail_norm == 0.0, 0.0, col_i / denom)
            stored = tl.where(
                row_offsets == i,
                beta,
                tl.where(tail_norm == 0.0, 0.0, v_tail),
            )
            new_col_i = tl.where(active, stored, col_i)
            panel = tl.where(col_offsets[None, :] == i, new_col_i[:, None], panel)

            v = tl.where(
                row_offsets == i,
                1.0,
                tl.where(row_offsets > i, v_tail, 0.0),
            )
            dot = tl.sum(panel * v[:, None], axis=0)
            if BUILD_GRAM:
                previous_dot = tl.where(col_offsets < i, dot, 0.0)
                tl.store(
                    gram_ptr
                    + gram_base
                    + col_offsets * PANEL
                    + i,
                    previous_dot,
                )
            updated = panel - (tau * dot)[None, :] * v[:, None]
            panel = tl.where(col_offsets[None, :] > i, updated, panel)
            tl.store(tau_ptr + batch_id * N + panel_start + i, tau)
        tl.store(
            h_ptr + base + rows[:, None] * N + cols[None, :],
            panel,
            mask=row_mask[:, None],
        )
    else:
        if matrix_is_selected:
            tl.store(tau_ptr + batch_id * N + cols, 0.0)
            if CHECK_COLUMN_ACTIVITY:
                if panel_start < active_column_end + 32:
                    tl.store(
                        h_ptr + base + rows[:, None] * N + cols[None, :],
                        0.0,
                        mask=row_mask[:, None],
                    )
    if NORMALIZE_CHOLESKY and (precision_metadata & 7) == 0:
        normalization_metadata = tl.load(
            normalization_flags_ptr + batch_id
        )
        partial_offsets = tl.arange(0, 16)[:, None]
        norm_columns = tl.arange(0, 32)[None, :]
        # This hybrid retains Cholesky recovery through column 319: eight
        # unconditional early panels plus two consensus-routed late panels.
        # The parent schedule continued through column 383 and therefore
        # normalized twelve slots; reading those final two uninitialized
        # norm slots produced NaNs after shortening the schedule.
        for panel_index in tl.static_range(0, 10):
            normalize_panel = (panel_index < 8) | (
                (normalization_metadata & 7) == 0
            )
            partial_count = (480 - panel_index * 32 + 63) // 64 + 1
            norm_values = tl.load(
                gram_ptr
                + (batch_id * 12 + panel_index) * 16 * 32
                + partial_offsets * 32
                + norm_columns,
                mask=(partial_offsets < partial_count) & normalize_panel,
                other=0.0,
            )
            norm_squared = tl.sum(norm_values, axis=0)
            tau_offsets = panel_index * 32 + tl.arange(0, 32)
            old_tau = tl.load(tau_ptr + batch_id * N + tau_offsets)
            normalized = tl.where(
                old_tau != 0.0, 2.0 / norm_squared, 0.0
            )
            tl.store(
                tau_ptr + batch_id * N + tau_offsets,
                normalized,
                mask=normalize_panel,
            )


@triton.jit
def _safe512_factor16_parent_order(
    panel,
    row_offsets,
    col_offsets,
    tau_ptr,
    tau_base,
    UNROLL: tl.constexpr,
):
    """Run the accepted 16-column Householder recurrence unchanged."""
    for i in tl.range(0, 16, loop_unroll_factor=UNROLL):
        col_i = tl.sum(
            tl.where(col_offsets[None, :] == i, panel, 0.0), axis=1
        )
        active = row_offsets >= i
        alpha = tl.sum(tl.where(row_offsets == i, col_i, 0.0), axis=0)
        tail = tl.where(row_offsets > i, col_i, 0.0)
        tail_ss = tl.sum(tail * tail, axis=0)
        tail_norm = tl.sqrt(tail_ss)
        full_norm = tl.sqrt(alpha * alpha + tail_ss)

        beta_raw = tl.where(alpha >= 0.0, -full_norm, full_norm)
        beta = tl.where(tail_norm == 0.0, alpha, beta_raw)
        tau_raw = (beta - alpha) / beta
        tau = tl.where(tail_norm == 0.0, 0.0, tau_raw)
        denom = alpha - beta
        v_tail = tl.where(tail_norm == 0.0, 0.0, col_i / denom)
        stored = tl.where(
            row_offsets == i,
            beta,
            tl.where(tail_norm == 0.0, 0.0, v_tail),
        )
        new_col_i = tl.where(active, stored, col_i)
        panel = tl.where(
            col_offsets[None, :] == i, new_col_i[:, None], panel
        )

        v = tl.where(
            row_offsets == i,
            1.0,
            tl.where(row_offsets > i, v_tail, 0.0),
        )
        dot = tl.sum(panel * v[:, None], axis=0)
        updated = panel - (tau * dot)[None, :] * v[:, None]
        panel = tl.where(col_offsets[None, :] > i, updated, panel)
        tl.store(tau_ptr + tau_base + i, tau)
    return panel


@triton.jit
def _safe512_apply16_parent_order_looped(
    block_t,
    h_ptr,
    tau_ptr,
    rows,
    row_offsets,
    row_mask,
    batch_id,
    base,
    panel_start,
    N: tl.constexpr,
):
    """Apply the accepted reflector recurrence without code-size unrolling."""
    for i in tl.range(0, 16, loop_unroll_factor=1):
        reflector_col = panel_start + i
        v = tl.load(
            h_ptr + base + rows * N + reflector_col,
            mask=row_mask & (row_offsets > i),
            other=0.0,
        )
        v = tl.where(row_offsets == i, 1.0, v)
        v = tl.where(row_offsets < i, 0.0, v)
        tau = tl.load(tau_ptr + batch_id * N + reflector_col)
        dot = tl.sum(block_t * v[None, :], axis=1)
        block_t = block_t - (tau * dot)[:, None] * v[None, :]
    return block_t


@triton.jit
def _safe512_factor_fallback32_phased_kernel(
    src_ptr,
    h_ptr,
    tau_ptr,
    precision_flags_ptr,
    gram_ptr,
    normalization_flags_ptr,
    panel_start,
    N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    SELECT_ALL: tl.constexpr,
    FACTOR_UNROLL: tl.constexpr,
    NORMALIZE_CHOLESKY: tl.constexpr,
):
    """Execute factor16 -> apply16 -> factor16 in one ordered CTA.

    Global stores and CTA barriers reproduce the accepted three-launch
    handoff while making each phase's registers dead before the next phase.
    """
    batch_id = tl.program_id(0)
    base = batch_id * N * N
    panel_start = tl.multiple_of(panel_start, 32)
    precision_metadata = tl.load(precision_flags_ptr + batch_id)

    if FIRST_PANEL:
        metadata_offsets = tl.arange(0, N)
        tl.store(
            tau_ptr + batch_id * N + metadata_offsets,
            precision_metadata >> 3,
            mask=metadata_offsets >= 32,
        )

    fallback_selected = SELECT_ALL or ((precision_metadata & 7) != 0)
    if not fallback_selected:
        return

    if FIRST_PANEL:
        panel_active = True
        active_column_end = N
    else:
        active_column_end = tl.load(tau_ptr + batch_id * N + panel_start)
        panel_active = panel_start < active_column_end

    row_offsets = tl.arange(0, BLOCK_M)
    col_offsets = tl.arange(0, 16)
    rows0 = panel_start + row_offsets
    cols0 = panel_start + col_offsets
    cols1 = panel_start + 16 + col_offsets
    band_flag = (precision_metadata & 2) != 0
    row_mask0 = (rows0 < N) & ((not band_flag) | (row_offsets < 64))
    tau_base = batch_id * N + panel_start

    if panel_active:
        first_panel = tl.load(
            src_ptr + base + rows0[:, None] * N + cols0[None, :],
            mask=row_mask0[:, None],
            other=0.0,
        )
        first_panel = _safe512_factor16_parent_order(
            first_panel,
            row_offsets,
            col_offsets,
            tau_ptr,
            tau_base,
            FACTOR_UNROLL,
        )
        tl.store(
            h_ptr + base + rows0[:, None] * N + cols0[None, :],
            first_panel,
            mask=row_mask0[:, None],
        )

        tl.debug_barrier()
        second_columns = tl.load(
            src_ptr + base + rows0[:, None] * N + cols1[None, :],
            mask=row_mask0[:, None],
            other=0.0,
        )
        second_columns_t = _safe512_apply16_parent_order_looped(
            tl.trans(second_columns),
            h_ptr,
            tau_ptr,
            rows0,
            row_offsets,
            row_mask0,
            batch_id,
            base,
            panel_start,
            N,
        )
        tl.store(
            h_ptr + base + rows0[:, None] * N + cols1[None, :],
            tl.trans(second_columns_t),
            mask=row_mask0[:, None],
        )

        tl.debug_barrier()
        rows1 = panel_start + 16 + row_offsets
        row_mask1 = (rows1 < N) & ((not band_flag) | (row_offsets < 64))
        second_panel = tl.load(
            h_ptr + base + rows1[:, None] * N + cols1[None, :],
            mask=row_mask1[:, None],
            other=0.0,
        )
        second_panel = _safe512_factor16_parent_order(
            second_panel,
            row_offsets,
            col_offsets,
            tau_ptr,
            tau_base + 16,
            FACTOR_UNROLL,
        )
        tl.store(
            h_ptr + base + rows1[:, None] * N + cols1[None, :],
            second_panel,
            mask=row_mask1[:, None],
        )
    else:
        tl.store(tau_ptr + tau_base + tl.arange(0, 32), 0.0)
        if panel_start < active_column_end + 32:
            zero_cols = panel_start + tl.arange(0, 32)
            tl.store(
                h_ptr + base + rows0[:, None] * N + zero_cols[None, :],
                0.0,
                mask=row_mask0[:, None],
            )

    if NORMALIZE_CHOLESKY and (precision_metadata & 7) == 0:
        tl.debug_barrier()
        normalization_metadata = tl.load(
            normalization_flags_ptr + batch_id
        )
        partial_offsets = tl.arange(0, 16)[:, None]
        norm_columns = tl.arange(0, 32)[None, :]
        for panel_index in tl.static_range(0, 10):
            normalize_panel = (panel_index < 8) | (
                (normalization_metadata & 7) == 0
            )
            partial_count = (480 - panel_index * 32 + 63) // 64 + 1
            norm_values = tl.load(
                gram_ptr
                + (batch_id * 12 + panel_index) * 16 * 32
                + partial_offsets * 32
                + norm_columns,
                mask=(partial_offsets < partial_count) & normalize_panel,
                other=0.0,
            )
            norm_squared = tl.sum(norm_values, axis=0)
            tau_offsets = panel_index * 32 + tl.arange(0, 32)
            old_tau = tl.load(tau_ptr + batch_id * N + tau_offsets)
            normalized = tl.where(
                old_tau != 0.0, 2.0 / norm_squared, 0.0
            )
            tl.store(
                tau_ptr + batch_id * N + tau_offsets,
                normalized,
                mask=normalize_panel,
            )


def _safe512_panel_configuration(panel_start: int) -> tuple[int, int]:
    """Return the live-row power of two and tuned factor warp count."""
    if panel_start < 256:
        return 512, 8
    if panel_start < 384:
        return 256, 2
    if panel_start < 448:
        return 128, 2
    if panel_start < 480:
        return 64, 1
    return 32, 1


def _safe512_factor_32_launch(
    source: torch.Tensor,
    h: torch.Tensor,
    tau: torch.Tensor,
    scratch: torch.Tensor,
    precision_flags: torch.Tensor,
    batch: int,
    panel_start: int,
    unsafe_only: bool = False,
    normalize_cholesky: bool = False,
    normalization_flags: torch.Tensor | None = None,
) -> None:
    """Factor a stable 32-column block as two 16-column subpanels."""
    if normalization_flags is None:
        normalization_flags = precision_flags
    block_m, panel_warps = _safe512_panel_configuration(panel_start)
    if panel_warps == 8:
        apply_warps = 4
    elif block_m <= 128:
        apply_warps = 1
    elif panel_warps == 4:
        apply_warps = 2
    else:
        apply_warps = panel_warps
    first_panel = panel_start == 0
    use_phased_factor = (
        (unsafe_only and panel_start < 320)
        or ((not unsafe_only) and panel_start < 512)
    )
    if use_phased_factor:
        factor_unroll = 4 if 256 <= panel_start < 384 else 8
        _safe512_factor_fallback32_phased_kernel[(batch,)](
            source,
            h,
            tau,
            precision_flags,
            scratch,
            normalization_flags,
            panel_start,
            N=512,
            BLOCK_M=block_m,
            FIRST_PANEL=first_panel,
            SELECT_ALL=not unsafe_only,
            FACTOR_UNROLL=factor_unroll,
            NORMALIZE_CHOLESKY=normalize_cholesky,
            num_warps=panel_warps,
            maxnreg=(168 if panel_warps == 2 else 128),
        )
        return
    _safe512_factor_panel_kernel[(batch,)](
        source,
        h,
        tau,
        scratch,
        precision_flags,
        normalization_flags,
        panel_start,
        N=512,
        PANEL=16,
        BLOCK_M=block_m,
        CHECK_COLUMN_ACTIVITY=not first_panel,
        BUILD_GRAM=False,
        SELECT_FLAG=(2 if unsafe_only else 0),
        SEED_ACTIVITY=first_panel,
        NORMALIZE_CHOLESKY=False,
        num_warps=panel_warps,
    )
    _safe512_apply_panel_kernel[(batch, 1)](
        source,
        h,
        tau,
        scratch,
        precision_flags,
        panel_start,
        N=512,
        PANEL=16,
        BLOCK_M=block_m,
        BLOCK_N=16,
        RESIDUAL_TILES=16,
        RECORD_COLUMN_ACTIVITY=False,
        SELECT_FLAG=(6 if unsafe_only else 0),
        ASSUME_ACTIVE=first_panel,
        num_warps=apply_warps,
    )
    _safe512_factor_panel_kernel[(batch,)](
        h,
        h,
        tau,
        scratch,
        precision_flags,
        normalization_flags,
        panel_start + 16,
        N=512,
        PANEL=16,
        BLOCK_M=block_m,
        CHECK_COLUMN_ACTIVITY=not first_panel,
        BUILD_GRAM=False,
        SELECT_FLAG=(2 if unsafe_only else 0),
        SEED_ACTIVITY=False,
        NORMALIZE_CHOLESKY=normalize_cholesky,
        num_warps=panel_warps,
    )


def _safe512_build_32_w_launch(
    h: torch.Tensor,
    tau: torch.Tensor,
    gram_partials: torch.Tensor,
    t_factor: torch.Tensor,
    w_half: torch.Tensor,
    v_half: torch.Tensor,
    precision_flags: torch.Tensor,
    batch: int,
    panel_start: int,
    precision_only: bool = False,
    form_w: bool = True,
) -> None:
    """Build one 32-reflector compact-WY block in split FP16 storage."""
    block_m, _ = _safe512_panel_configuration(panel_start)
    gram_block_rows = 64 if panel_start == 448 else 128
    setup_warps = 2 if panel_start >= 224 else 4
    partial_count = triton.cdiv(512 - panel_start, gram_block_rows)
    first_panel = panel_start == 0
    _safe512_gram_partial_kernel[(batch, partial_count)](
        h,
        v_half,
        tau,
        gram_partials,
        precision_flags,
        panel_start,
        N=512,
        PANEL=32,
        BLOCK_ROWS=gram_block_rows,
        ASSUME_ACTIVE=first_panel,
        SKIP_FLAGS=(3 if first_panel else 2),
        REQUIRE_PRECISION=precision_only,
        num_warps=setup_warps,
    )
    _safe512_aggregate_t_kernel[(batch,)](
        gram_partials,
        tau,
        t_factor,
        gram_partials,
        t_factor,
        precision_flags,
        panel_start,
        N=512,
        PANEL=32,
        PARTIALS=partial_count,
        ASSUME_ACTIVE=first_panel,
        SKIP_FLAGS=(3 if first_panel else 2),
        REQUIRE_PRECISION=precision_only,
        BLOCKED_T=True,
        BUILD_COUPLING=False,
        REUSE_RECOVERED_T=False,
        num_warps=1,
    )
    if form_w:
        live_rows = 512 - panel_start
        form_block_rows = (
            64
            if live_rows % 128 != 0 and live_rows % 128 <= 64
            else 128
        )
        form_warps = 2 if form_block_rows == 64 else setup_warps
        _safe512_form_w_kernel[
            (batch, triton.cdiv(live_rows, form_block_rows))
        ](
            h,
            tau,
            t_factor,
            w_half,
            v_half,
            w_half,
            v_half,
            precision_flags,
            panel_start,
            N=512,
            PANEL=32,
            BLOCK_ROWS=form_block_rows,
            ASSUME_ACTIVE=first_panel,
            TF32_MODE=1,
            SELECT_FLAG=(
                7
                if precision_only and first_panel
                else (6 if precision_only else (1 if first_panel else 3))
            ),
            SPLIT_PRECISION=False,
            PRECISE_FORM=False,
            ZERO_PREFIX=False,
            num_warps=form_warps,
        )


def _safe512_cholesky_factor_32_launch(
    source: torch.Tensor,
    h: torch.Tensor,
    tau: torch.Tensor,
    partials: torch.Tensor,
    t_factor: torch.Tensor,
    bottom_transform: torch.Tensor,
    packed_vectors: torch.Tensor,
    norm_partials: torch.Tensor,
    precision_flags: torch.Tensor,
    batch: int,
    panel_start: int,
) -> None:
    """Recover an ordinary compact-Householder block through CholeskyQR."""
    active_chunks = triton.cdiv(512 - panel_start, 128)
    first_panel = panel_start == 0
    _s9_n512_n2048_chol_gram32_kernel[(active_chunks, 1, batch)](
        source,
        h,
        precision_flags,
        partials,
        panel_start,
        n=512,
        CHUNKS=4,
        ROW_CHUNK=128,
        FIRST_PANEL=first_panel,
        ROUTE_N512=True,
        REQUIRE_FULL_ACTIVE=False,
        PRECISE_ROUTE=False,
        num_warps=4,
        num_stages=2,
    )
    _safe512_chol32_factor_compact_kernel[(batch,)](
        source,
        h,
        partials,
        t_factor,
        bottom_transform,
        packed_vectors,
        norm_partials,
        tau,
        precision_flags,
        panel_start,
        active_chunks,
        FIRST_PANEL=first_panel,
        num_warps=1,
        num_stages=1,
    )
    bottom_tiles = triton.cdiv(512 - panel_start - 32, 64)
    _s9_n512_n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
        source,
        h,
        bottom_transform,
        packed_vectors,
        norm_partials,
        precision_flags,
        panel_start,
        n=512,
        PANEL=32,
        BLOCK_M=64,
        FIRST_PANEL=first_panel,
        ROUTE_N512=True,
        REQUIRE_FULL_ACTIVE=False,
        PRECISE_RECOVERY=False,
        STORE_NORMS=True,
        STORE_PACKED_VECTORS=True,
        num_warps=1,
    )


def _safe512_form_known_32_w_launch(
    h: torch.Tensor,
    tau: torch.Tensor,
    t_factor: torch.Tensor,
    w_half: torch.Tensor,
    v_half: torch.Tensor,
    precision_flags: torch.Tensor,
    batch: int,
    panel_start: int,
    safe_only: bool = False,
    precision_w: torch.Tensor | None = None,
    precision_v: torch.Tensor | None = None,
    precise_form: bool = False,
) -> None:
    """Form W when compact-Householder recovery has already produced T."""
    block_m, _ = _safe512_panel_configuration(panel_start)
    live_rows = 512 - panel_start
    block_rows = (
        64 if live_rows % 128 != 0 and live_rows % 128 <= 64 else 128
    )
    setup_warps = 2 if panel_start >= 224 else 4
    form_warps = 2 if block_rows == 64 else setup_warps
    split_precision = precision_w is not None
    if precision_w is None:
        precision_w = w_half
        precision_v = v_half
    _safe512_form_w_kernel[
        (batch, triton.cdiv(512 - panel_start, block_rows))
    ](
        h,
        tau,
        t_factor,
        w_half,
        v_half,
        precision_w,
        precision_v,
        precision_flags,
        panel_start,
        N=512,
        PANEL=32,
        BLOCK_ROWS=block_rows,
        ASSUME_ACTIVE=panel_start == 0,
        TF32_MODE=1,
        SELECT_FLAG=(1 if safe_only else 3),
        SPLIT_PRECISION=split_precision,
        PRECISE_FORM=split_precision or precise_form,
        ZERO_PREFIX=panel_start == 32,
        num_warps=form_warps,
    )


def _safe512_qr_recursive64(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """QR using recursively coupled pairs of 32-reflector WY blocks.

    The first block retains conservative numerical routing and seeds the
    active endpoint selected by the input classifier.  Thereafter, each pair
    is coupled algebraically into a split 64-reflector compact-WY update,
    reducing full trailing-matrix round trips without constructing a
    compiler-hostile monolithic 64x64 register tile.
    """
    batch = a.shape[0]
    h = torch.empty_like(a)
    tau = torch.empty((batch, 512), device=a.device, dtype=a.dtype)
    residual_tiles = torch.empty((batch, 16), device=a.device, dtype=a.dtype)
    gram_partials = torch.empty((batch, 4, 32, 32), device=a.device, dtype=a.dtype)
    cross_gram_partials = torch.empty(
        (batch, 4, 32, 32), device=a.device, dtype=a.dtype
    )
    left_t = torch.empty((batch, 32, 32), device=a.device, dtype=a.dtype)
    right_t = torch.empty((batch, 32, 32), device=a.device, dtype=a.dtype)
    left_w = torch.empty((batch, 512, 32), device=a.device, dtype=torch.float16)
    left_v = torch.empty((batch, 512, 32), device=a.device, dtype=torch.float16)
    right_w = torch.empty((batch, 512, 32), device=a.device, dtype=torch.float16)
    right_v = torch.empty((batch, 512, 32), device=a.device, dtype=torch.float16)
    precision_w = torch.empty((batch, 512, 32), device=a.device, dtype=a.dtype)
    precision_v = torch.empty((batch, 512, 32), device=a.device, dtype=a.dtype)
    chol_bottom = torch.empty((batch, 32, 32), device=a.device, dtype=a.dtype)
    chol_norms = torch.empty((batch, 12, 16, 32), device=a.device, dtype=a.dtype)
    precision_flags = torch.empty((batch,), device=a.device, dtype=torch.int32)
    late_precision_flags = torch.empty_like(precision_flags)

    _safe512_classify_precision_kernel[(batch,)](
        a,
        precision_flags,
        N=512,
        PANEL=32,
        DROP_RELATIVE_L1=1.0e-5,
        num_warps=2,
    )
    _safe512_v10_n512_build_late_route_flags[(1,)](
        precision_flags,
        late_precision_flags,
        BATCH=batch,
        BLOCK=triton.next_power_of_2(batch),
        num_warps=1,
    )
    _safe512_copy_band_initial_kernel[(4, batch)](
        a,
        h,
        precision_flags,
        N=512,
        BLOCK_N=128,
        CHUNK_M=32,
        num_warps=2,
    )
    # Safe matrices recover their first compact-Householder block through
    # CholeskyQR; flagged matrices retain stable Householder factoring.
    _safe512_cholesky_factor_32_launch(
        a, h, tau, gram_partials, left_t, chol_bottom, left_v,
        chol_norms, precision_flags, batch, 0,
    )
    _safe512_factor_32_launch(
        a, h, tau, residual_tiles, precision_flags, batch, 0,
        unsafe_only=True,
    )
    _safe512_build_32_w_launch(
        h,
        tau,
        gram_partials,
        left_t,
        left_w,
        left_v,
        precision_flags,
        batch,
        0,
        precision_only=True,
    )
    _safe512_form_known_32_w_launch(
        h, tau, left_t, left_w, left_v, precision_flags, batch, 0,
        safe_only=True,
    )
    _safe512_apply_panel_wy_fused_kernel[(1, batch)](
        a,
        h,
        tau,
        left_w,
        left_v,
        precision_flags,
        residual_tiles,
        0,
        N=512,
        PANEL=32,
        BLOCK_M=512,
        BLOCK_N=32,
        CHUNK_M=64,
        TF32_MODE=0,
        SELECT_FLAG=1,
        RESIDUAL_TILES=16,
        RECORD_COLUMN_ACTIVITY=False,
        HANDLE_BAND=False,
        num_warps=2,
        num_stages=1,
        maxnreg=168,
    )
    _safe512_apply_panel_kernel[(batch, 8)](
        a,
        h,
        tau,
        residual_tiles,
        precision_flags,
        0,
        N=512,
        PANEL=32,
        BLOCK_M=512,
        BLOCK_N=64,
        RESIDUAL_TILES=16,
        RECORD_COLUMN_ACTIVITY=False,
        SELECT_FLAG=2,
        ASSUME_ACTIVE=False,
        num_warps=16,
    )
    _safe512_cholesky_factor_32_launch(
        h, h, tau, gram_partials, right_t, chol_bottom, right_v,
        chol_norms, precision_flags, batch, 32
    )
    _safe512_factor_32_launch(
        h, h, tau, residual_tiles, precision_flags, batch, 32,
        unsafe_only=True,
    )
    _safe512_build_32_w_launch(
        h, tau, gram_partials, right_t, precision_w, precision_v,
        precision_flags, batch, 32, precision_only=True, form_w=False,
    )
    _safe512_form_known_32_w_launch(
        h, tau, right_t, right_w, right_v,
        precision_flags, batch, 32,
        precision_w=precision_w, precision_v=precision_v,
    )
    first_cross_rows = 128
    first_cross_partials = 4
    _safe512_cross_gram_partial_kernel[(batch, first_cross_partials)](
        left_v,
        right_v,
        tau,
        cross_gram_partials,
        precision_flags,
        0,
        N=512,
        PANEL=32,
        BLOCK_ROWS=first_cross_rows,
        EXCLUDE_PRECISION=True,
        num_warps=4,
    )
    _safe512_aggregate_cross_t_kernel[(batch,)](
        cross_gram_partials,
        tau,
        right_t,
        left_t,
        precision_flags,
        0,
        N=512,
        PANEL=32,
        PARTIALS=first_cross_partials,
        EXCLUDE_PRECISION=True,
        num_warps=1,
    )
    _safe512_couple_w_kernel[(batch, 4)](
        left_w,
        right_w,
        left_t,
        tau,
        precision_flags,
        0,
        N=512,
        PANEL=32,
        BLOCK_ROWS=first_cross_rows,
        EXCLUDE_PRECISION=True,
        num_warps=4,
    )
    # The initial pair has fixed geometry.  The interleaved operand layout
    # preserves the coupled WY product while avoiding the conflicted shared
    # staging generated for the generic concatenation.
    _safe512_apply_initial_interleaved_wy_kernel[(7, batch)](
        a,
        h,
        tau,
        left_w,
        left_v,
        right_w,
        right_v,
        precision_flags,
        BATCH=batch,
        num_warps=4,
        num_stages=2,
        maxnreg=120,
    )
    _safe512_apply_panel_wy_fused_kernel[(14, batch)](
        h,
        h,
        tau,
        precision_w,
        precision_v,
        precision_flags,
        residual_tiles,
        32,
        N=512,
        PANEL=32,
        BLOCK_M=512,
        BLOCK_N=32,
        CHUNK_M=64,
        TF32_MODE=5,
        SELECT_FLAG=5,
        RESIDUAL_TILES=16,
        RECORD_COLUMN_ACTIVITY=False,
        HANDLE_BAND=False,
        num_warps=8,
        num_stages=2,
        maxnreg=128,
    )
    # Recursive 64-column groups.  Q0 is applied only to Q1's panel; their
    # coupled compact-WY representation updates all remaining columns.
    for group_start in range(64, 448, 64):
        factor_route_flags = (
            precision_flags if group_start < 256 else late_precision_flags
        )
        if group_start < 320:
            _safe512_cholesky_factor_32_launch(
                h,
                h,
                tau,
                gram_partials,
                left_t,
                chol_bottom,
                left_v,
                chol_norms,
                factor_route_flags,
                batch,
                group_start,
            )
            _safe512_factor_32_launch(
                h, h, tau, residual_tiles, factor_route_flags, batch,
                group_start, unsafe_only=True,
            )
        else:
            _safe512_factor_32_launch(
                h, h, tau, residual_tiles, precision_flags, batch, group_start,
            )
        # Build W0 before factoring Q1: the next panel's first tau slot still
        # carries the live-column metadata consumed by the WY setup kernels.
        left_block_m, _ = _safe512_panel_configuration(group_start)
        _safe512_build_32_w_launch(
            h, tau, gram_partials, left_t, left_w, left_v,
            factor_route_flags, batch, group_start,
            precision_only=group_start < 320,
            form_w=group_start >= 320,
        )
        if group_start < 320:
            if group_start == 64:
                _safe512_form_known_32_w_launch(
                    h, tau, left_t, left_w, left_v,
                    factor_route_flags, batch, group_start,
                    precise_form=True,
                )
            else:
                _safe512_form_known_32_w_launch(
                    h, tau, left_t, left_w, left_v,
                    factor_route_flags, batch, group_start,
                )
        _safe512_apply_panel_wy_fused_kernel[(1, batch)](
            h,
            h,
            tau,
            left_w,
            left_v,
            precision_flags,
            residual_tiles,
            group_start,
            N=512,
            PANEL=32,
            BLOCK_M=512 - group_start,
            BLOCK_N=32,
            CHUNK_M=64,
            TF32_MODE=0,
            SELECT_FLAG=3,
            RESIDUAL_TILES=16,
            RECORD_COLUMN_ACTIVITY=False,
            HANDLE_BAND=True,
            num_warps=2,
            num_stages=1,
            maxnreg=168,
        )
        right_start = group_start + 32
        if group_start < 320:
            _safe512_cholesky_factor_32_launch(
                h,
                h,
                tau,
                gram_partials,
                right_t,
                chol_bottom,
                right_v,
                chol_norms,
                factor_route_flags,
                batch,
                right_start,
            )
            _safe512_factor_32_launch(
                h, h, tau, residual_tiles, factor_route_flags, batch,
                right_start, unsafe_only=True,
            )
        else:
            _safe512_factor_32_launch(
                h, h, tau, residual_tiles, precision_flags, batch, right_start,
            )
        cross_rows = 128
        cross_partials = triton.cdiv(512 - group_start, cross_rows)
        cross_warps = 4 if group_start < 320 else 2
        aggregate_flags = (
            factor_route_flags if group_start == 256 else precision_flags
        )
        _safe512_right_gram_and_cross_partial_kernel[
            (batch, cross_partials)
        ](
            h,
            left_v,
            right_v,
            tau,
            gram_partials,
            cross_gram_partials,
            aggregate_flags,
            group_start,
            N=512,
            PANEL=32,
            BLOCK_ROWS=cross_rows,
            REUSE_RECOVERED_SAFE=group_start < 320,
            num_warps=cross_warps,
        )
        _safe512_aggregate_t_kernel[(batch,)](
            gram_partials,
            tau,
            right_t,
            cross_gram_partials,
            left_t,
            aggregate_flags,
            right_start,
            N=512,
            PANEL=32,
            PARTIALS=cross_partials,
            ASSUME_ACTIVE=False,
            SKIP_FLAGS=2,
            REQUIRE_PRECISION=group_start < 256,
            BLOCKED_T=True,
            BUILD_COUPLING=True,
            REUSE_RECOVERED_T=group_start < 320,
            num_warps=(2 if group_start < 256 else 1),
        )
        _safe512_form_coupled_right_w_kernel[
            (
                batch,
                triton.cdiv(
                    512 - group_start,
                    (
                        64
                        if (512 - group_start) % 128 != 0
                        and (512 - group_start) % 128 <= 64
                        else 128
                    ),
                ),
            )
        ](
            h,
            right_t,
            left_w,
            left_t,
            right_w,
            right_v,
            tau,
            precision_flags,
            group_start,
            N=512,
            PANEL=32,
            BLOCK_ROWS=(
                64
                if (512 - group_start) % 128 != 0
                and (512 - group_start) % 128 <= 64
                else 128
            ),
            num_warps=(
                2
                if (512 - group_start) % 128 != 0
                and (512 - group_start) % 128 <= 64
                else cross_warps
            ),
        )

        trailing = 512 - group_start - 64
        _safe512_apply_selected_interleaved_wy_kernel[
            (triton.cdiv(trailing, 64), batch)
        ](
            h,
            tau,
            left_w,
            left_v,
            right_w,
            right_v,
            precision_flags,
            GROUP_START=group_start,
            BLOCK_M=512 - group_start,
            BATCH=batch,
            num_warps=4,
            num_stages=2,
            maxnreg=120,
        )
    # The terminal panel retains compact-WY arithmetic for transformed
    # precision-sensitive inputs; direct application changes their rounding
    # enough to consume the residual margin.
    _safe512_factor_32_launch(
        h, h, tau, residual_tiles, precision_flags, batch, 448,
    )
    _safe512_build_32_w_launch(
        h,
        tau,
        gram_partials,
        left_t,
        left_w,
        left_v,
        precision_flags,
        batch,
        448,
    )
    _safe512_apply_panel_wy_fused_kernel[(1, batch)](
        h,
        h,
        tau,
        left_w,
        left_v,
        precision_flags,
        residual_tiles,
        448,
        N=512,
        PANEL=32,
        BLOCK_M=64,
        BLOCK_N=32,
        CHUNK_M=64,
        TF32_MODE=0,
        SELECT_FLAG=3,
        RESIDUAL_TILES=16,
        RECORD_COLUMN_ACTIVITY=False,
        HANDLE_BAND=True,
        num_warps=2,
        num_stages=1,
        maxnreg=168,
    )
    _safe512_factor_32_launch(
        h,
        h,
        tau,
        chol_norms,
        precision_flags,
        batch,
        480,
        normalize_cholesky=True,
        normalization_flags=late_precision_flags,
    )
    return h, tau


def _safe512_qr_v2(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Return compact Householder factors ``(H, tau)`` for 512x512 batches."""
    if a.shape[-1] == 512:
        return _safe512_qr_recursive64(a)
    return _safe512_qr_recursive64(a)


@triton.jit
def _n1024_apply_panel_kernel(
    h,
    tau_out,
    tail_flags,
    panel_start,
    N: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    USE_TAIL_SKIP: tl.constexpr,
    SKIP_NONZERO: tl.constexpr,
):
    """Apply a factored panel to independent tiles of the trailing columns."""
    batch_id = tl.program_id(0)
    if USE_TAIL_SKIP:
        tail_active = tl.load(tail_flags + batch_id)
        if SKIP_NONZERO:
            if tail_active != 0:
                return
        else:
            if tail_active == 0:
                return
    tile_id = tl.program_id(1)
    row_offset = tl.arange(0, BLOCK_M)
    col_offset = tl.arange(0, BLOCK_N)
    rows = panel_start + row_offset
    cols = panel_start + PANEL + tile_id * BLOCK_N + col_offset
    matrix = h + batch_id * N * N
    valid = (rows[:, None] < N) & (cols[None, :] < N)
    tile = tl.load(
        matrix + rows[:, None] * N + cols[None, :],
        mask=valid,
        other=0.0,
    )

    for j in tl.static_range(0, PANEL):
        packed = tl.load(
            matrix + rows * N + panel_start + j,
            mask=rows < N,
            other=0.0,
        )
        vector = tl.where(
            row_offset == j,
            1.0,
            tl.where(row_offset > j, packed, 0.0),
        )
        tau = tl.load(tau_out + batch_id * N + panel_start + j)
        products = tl.sum(vector[:, None] * tile, axis=0)
        tile = tl.fma(-vector[:, None], (tau * products)[None, :], tile)

    tl.store(
        matrix + rows[:, None] * N + cols[None, :],
        tile,
        mask=valid,
    )


@triton.jit
def _n2048_chol_gram_kernel(
    source,
    h,
    partial_gram,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Compute independent, accurate Gram contributions for a tall panel."""
    chunk = tl.program_id(0)
    batch = tl.program_id(1)
    rows = panel_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    columns = panel_start + tl.arange(0, PANEL)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )

    # tf32x3 is close enough to an IEEE FP32 product for the Cholesky
    # recovery, while retaining Blackwell tensor-core throughput.  Replacing
    # the diagonal with an ordinary FP32 reduction avoids downward norm error.
    gram = tl.dot(tl.trans(values), values, input_precision="tf32x3")
    diagonal = tl.sum(values * values, axis=0)
    offsets = tl.arange(0, PANEL)
    gram = tl.where(offsets[:, None] == offsets[None, :], diagonal[:, None], gram)
    base = (batch * CHUNKS + chunk) * PANEL * PANEL
    tl.store(
        partial_gram
        + base
        + offsets[:, None] * PANEL
        + offsets[None, :],
        gram,
    )


@triton.jit
def _n2048_chol_gram32_kernel(
    source,
    h,
    partial_gram,
    panel_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Compute a 32-column Gram as two diagonal and one cross 16 tile."""
    chunk = tl.program_id(0)
    gram_block = tl.program_id(1)
    batch = tl.program_id(2)
    rows = panel_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    tile_offsets = tl.arange(0, 16)
    left_tile = tl.where(gram_block == 2, 1, 0)
    right_tile = tl.where(gram_block == 0, 0, 1)
    left_columns = panel_start + left_tile * 16 + tile_offsets
    right_columns = panel_start + right_tile * 16 + tile_offsets
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    left = tl.load(
        matrix + matrix_base + rows[:, None] * n + left_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    right = tl.load(
        matrix + matrix_base + rows[:, None] * n + right_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    gram = tl.dot(tl.trans(left), right, input_precision="tf32x3")
    if gram_block != 1:
        diagonal = tl.sum(left * left, axis=0)
        gram = tl.where(
            tile_offsets[:, None] == tile_offsets[None, :],
            diagonal[:, None],
            gram,
        )

    base = (batch * CHUNKS + chunk) * 32 * 32
    row_offsets = left_tile * 16 + tile_offsets
    column_offsets = right_tile * 16 + tile_offsets
    tl.store(
        partial_gram
        + base
        + row_offsets[:, None] * 32
        + column_offsets[None, :],
        gram,
    )
    if gram_block == 1:
        tl.store(
            partial_gram
            + base
            + column_offsets[:, None] * 32
            + row_offsets[None, :],
            tl.trans(gram),
        )


@triton.jit
def _n2048_chol_pair_gram_kernel(
    source,
    h,
    partial_gram,
    pair_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Compute one complete 64-column partial Gram per row slab.

    The column-zero specialization reads the immutable input directly; later
    pairs consume the trailing matrix produced by the preceding WY update.
    """
    chunk = tl.program_id(0)
    batch = tl.program_id(1)
    rows = pair_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    offsets = tl.arange(0, 64)
    columns = pair_start + offsets
    matrix_base = batch * n * n
    matrix = source if FIRST_PAIR else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    gram = tl.dot(
        tl.trans(values.to(tl.float16)), values.to(tl.float16)
    )
    diagonal = tl.sum(values * values, axis=0)
    gram = tl.where(
        offsets[:, None] == offsets[None, :], diagonal[:, None], gram
    )
    base = (batch * CHUNKS + chunk) * 64 * 64
    tl.store(
        partial_gram
        + base
        + offsets[:, None] * 64
        + offsets[None, :],
        gram,
    )


@triton.jit
def _n2048_group_joint_gram_kernel(
    source,
    h,
    pair_partials,
    target_partials,
    pair_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
):
    """Form G00, G01, and G11 for two adjacent 64-column pairs."""
    chunk = tl.program_id(0)
    gram_block = tl.program_id(1)
    batch = tl.program_id(2)
    rows = pair_start + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    offsets = tl.arange(0, 64)
    left_pair = tl.where(gram_block == 2, 1, 0)
    right_pair = tl.where(gram_block == 0, 0, 1)
    left_columns = pair_start + left_pair * 64 + offsets
    right_columns = pair_start + right_pair * 64 + offsets
    matrix_base = batch * n * n
    matrix = source if FIRST_GROUP else h
    left = tl.load(
        matrix + matrix_base + rows[:, None] * n + left_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    right = tl.load(
        matrix + matrix_base + rows[:, None] * n + right_columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    gram = tl.dot(tl.trans(left.to(tl.float16)), right.to(tl.float16))
    if gram_block != 1:
        diagonal = tl.sum(left * left, axis=0)
        gram = tl.where(
            offsets[:, None] == offsets[None, :], diagonal[:, None], gram
        )
    matrix_offsets = offsets[:, None] * 64 + offsets[None, :]
    if gram_block == 0:
        base = (batch * CHUNKS + chunk) * 64 * 64
        tl.store(pair_partials + base + matrix_offsets, gram)
    else:
        base = ((batch * CHUNKS + chunk) * 2 + (gram_block - 1)) * 64 * 64
        tl.store(target_partials + base + matrix_offsets, gram)


@triton.jit
def _n2048_group_target_from_gram_kernel(
    source,
    h,
    target_partials,
    pair_partials,
    first_inverse,
    second_inverse,
    first_vectors,
    second_vectors,
    weights,
    pair_start,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
    QUADRATIC_LOWER: tl.constexpr,
):
    """Derive adjacent-pair WY weights and its active Schur Gram."""
    batch = tl.program_id(0)
    offsets = tl.arange(0, 32)
    target_offsets = tl.arange(0, 64)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * 32 * 32
    partial_base = batch * CHUNKS * 2 * 64 * 64
    cross0 = tl.zeros((32, 64), tl.float32)
    cross1 = tl.zeros((32, 64), tl.float32)
    target_gram = tl.zeros((64, 64), tl.float32)
    gram_rows = target_offsets[:, None]
    gram_cols = target_offsets[None, :]
    for chunk in tl.static_range(0, CHUNKS):
        chunk_base = partial_base + chunk * 2 * 64 * 64
        cross0 += tl.load(
            target_partials
            + chunk_base
            + offsets[:, None] * 64
            + target_offsets[None, :]
        )
        cross1 += tl.load(
            target_partials
            + chunk_base
            + (32 + offsets)[:, None] * 64
            + target_offsets[None, :]
        )
        target_gram += tl.load(
            target_partials
            + chunk_base
            + 64 * 64
            + gram_rows * 64
            + gram_cols
        )

    inverse0 = tl.load(first_inverse + workspace_base + r * 32 + c)
    inverse1 = tl.load(second_inverse + workspace_base + r * 32 + c)
    r0 = tl.dot(tl.trans(inverse0.to(tl.float16)), cross0.to(tl.float16))
    matrix_base = batch * n * n
    r01 = tl.load(
        h
        + matrix_base
        + (pair_start + offsets)[:, None] * n
        + (pair_start + 32 + offsets)[None, :]
    )
    adjusted_cross1 = cross1 - tl.dot(
        tl.trans(r01.to(tl.float16)), r0.to(tl.float16)
    )
    r1 = tl.dot(tl.trans(inverse1.to(tl.float16)), adjusted_cross1.to(tl.float16))

    identity = tl.where(r == c, 1.0, 0.0)
    lower0 = tl.load(
        first_vectors
        + batch * n * 32
        + (pair_start + offsets)[:, None] * 32
        + offsets[None, :]
    )
    lower1 = tl.load(
        second_vectors
        + batch * n * 32
        + (pair_start + 32 + offsets)[:, None] * 32
        + offsets[None, :]
    )
    lower0_strict = tl.where(r > c, lower0, 0.0)
    lower1_strict = tl.where(r > c, lower1, 0.0)
    inverse_lower0 = identity - lower0_strict
    inverse_lower1 = identity - lower1_strict
    if QUADRATIC_LOWER:
        inverse_lower0 += tl.dot(
            lower0_strict.to(tl.float16), lower0_strict.to(tl.float16)
        )
        inverse_lower1 += tl.dot(
            lower1_strict.to(tl.float16), lower1_strict.to(tl.float16)
        )

    matrix = source if FIRST_GROUP else h
    target_columns = pair_start + 64 + target_offsets
    top0 = tl.load(
        matrix
        + matrix_base
        + (pair_start + offsets)[:, None] * n
        + target_columns[None, :]
    )
    top1 = tl.load(
        matrix
        + matrix_base
        + (pair_start + 32 + offsets)[:, None] * n
        + target_columns[None, :]
    )
    w0 = tl.dot(inverse_lower0.to(tl.float16), (top0 - r0).to(tl.float16))
    first_middle = tl.load(
        first_vectors
        + batch * n * 32
        + (pair_start + 32 + offsets)[:, None] * 32
        + offsets[None, :]
    )
    rhs1 = top1 - r1 - tl.dot(
        first_middle.to(tl.float16), w0.to(tl.float16)
    )
    w1 = tl.dot(inverse_lower1.to(tl.float16), rhs1.to(tl.float16))
    combined_weights = tl.cat(w0, w1, dim=0)
    weight_rows = target_offsets[:, None]
    tl.store(
        weights
        + batch * 64 * n
        + weight_rows * n
        + target_columns[None, :],
        combined_weights,
    )

    r0_product = tl.dot(tl.trans(r0.to(tl.float16)), r0.to(tl.float16))
    r1_product = tl.dot(tl.trans(r1.to(tl.float16)), r1.to(tl.float16))
    r0_diagonal = tl.sum(r0 * r0, axis=0)
    r1_diagonal = tl.sum(r1 * r1, axis=0)
    r0_product = tl.where(gram_rows == gram_cols, r0_diagonal[:, None], r0_product)
    r1_product = tl.where(gram_rows == gram_cols, r1_diagonal[:, None], r1_product)
    active_gram = target_gram - r0_product - r1_product
    tl.store(
        pair_partials
        + batch * 64 * 64
        + gram_rows * 64
        + gram_cols,
        active_gram,
    )


@triton.jit
def _n2048_chol_recover_kernel(
    source,
    h,
    tau,
    partial_gram,
    t_workspace,
    bottom_transform_workspace,
    packed_vectors,
    panel_start,
    active_chunks,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    CHUNKS: tl.constexpr,
    INVERSE_STEPS: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Recover ordinary compact Householder data from a Cholesky panel.

    If ``Q = A R^-1 D`` and ``M = I - Q_top = L U``, then the compact-WY
    representation is ``V_top=L``, ``T=U L^-T`` and
    ``V_bottom=-Q_bottom U^-1``.  The older recovery path merely split the raw
    lower and upper triangles of M.  Performing the actual LU factorization is
    essential once more than one Cholesky panel is used.
    """
    batch = tl.program_id(0)
    offsets = tl.arange(0, PANEL)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * PANEL * PANEL

    gram = tl.zeros((PANEL, PANEL), tl.float32)
    partial_base = batch * CHUNKS * PANEL * PANEL
    for chunk in tl.static_range(0, CHUNKS):
        gram += tl.load(
            partial_gram
            + partial_base
            + chunk * PANEL * PANEL
            + r * PANEL
            + c,
            mask=chunk < active_chunks,
            other=0.0,
        )
    gram_diagonal = tl.sum(tl.where(r == c, gram, 0.0), axis=1)
    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30
    # Upper Cholesky factor, kept entirely in FP32 registers.
    upper = tl.zeros((PANEL, PANEL), tl.float32)
    for j in tl.range(0, PANEL, loop_unroll_factor=1):
        gram_row = tl.sum(tl.where(r == j, gram, 0.0), axis=0)
        old_column = tl.sum(tl.where(c == j, upper, 0.0), axis=1)
        products = tl.sum(old_column[:, None] * upper, axis=0)
        diagonal_value = tl.sum(
            tl.where(offsets == j, gram_row - products, 0.0), axis=0
        )
        diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-20))
        new_row = (gram_row - products) / diagonal
        new_row = tl.where(offsets == j, diagonal, new_row)
        new_row = tl.where(offsets >= j, new_row, 0.0)
        upper = tl.where(r == j, new_row[None, :], upper)

    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )

    # For a 16-wide triangular matrix the finite Neumann product below is an
    # exact inverse in real arithmetic: (I-N)(I+N^2)(I+N^4)(I+N^8).
    # tf32x3 preserves the accuracy of the serial FP32 substitution while
    # replacing sixteen dependent reductions with tensor-core products.
    identity = tl.where(r == c, 1.0, 0.0)
    upper_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    power_r = tl.where(r < c, upper / upper_diagonal[:, None], 0.0)
    inverse_unit_r = identity - power_r
    for r_step in tl.static_range(0, INVERSE_STEPS):
        power_r = tl.dot(power_r, power_r, input_precision="tf32x3")
        inverse_unit_r = tl.dot(
            inverse_unit_r, identity + power_r, input_precision="tf32x3"
        )
    inverse = inverse_unit_r / upper_diagonal[None, :]

    matrix_base = batch * n * n
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(
        matrix + matrix_base + rows * n + columns,
        mask=(rows < n) & (columns < n),
        other=0.0,
    )
    q_top = tl.dot(top, inverse, input_precision="tf32x3")
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(panel_active, inverse * signs[None, :], 0.0)
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    q_top *= signs[None, :]

    # Stable sign selection makes I-Q_top strongly diagonally biased.  A true
    # no-pivot LU is still required; raw triangle extraction is not an LU.
    lu = tl.where(panel_active, identity - q_top, identity)
    for j in tl.range(0, PANEL, loop_unroll_factor=1):
        pivot_row = tl.sum(tl.where(r == j, lu, 0.0), axis=0)
        pivot_column = tl.sum(tl.where(c == j, lu, 0.0), axis=1)
        pivot = tl.sum(tl.where(offsets == j, pivot_row, 0.0), axis=0)
        factor = tl.where(offsets > j, pivot_column / pivot, 0.0)
        updated = lu - factor[:, None] * pivot_row[None, :]
        lu = tl.where((r > j) & (c > j), updated, lu)
        lu = tl.where((r > j) & (c == j), factor[:, None], lu)

    lower = tl.where(r > c, lu, identity)
    upper_lu = tl.where(r <= c, lu, 0.0)

    lower_transpose_tail = tl.where(r < c, tl.trans(lower), 0.0)
    inverse_lower_transpose = identity - lower_transpose_tail
    for l_step in tl.static_range(0, INVERSE_STEPS):
        lower_transpose_tail = tl.dot(
            lower_transpose_tail,
            lower_transpose_tail,
            input_precision="tf32x3",
        )
        inverse_lower_transpose = tl.dot(
            inverse_lower_transpose,
            identity + lower_transpose_tail,
            input_precision="tf32x3",
        )
    t_factor = tl.dot(
        upper_lu, inverse_lower_transpose, input_precision="tf32x3"
    )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    # Form -R^-1 D U^-1 for the independently processed bottom rows.
    lu_diagonal = tl.sum(tl.where(r == c, upper_lu, 0.0), axis=1)
    power_u = tl.where(r < c, upper_lu / lu_diagonal[:, None], 0.0)
    inverse_unit_u = identity - power_u
    for u_step in tl.static_range(0, INVERSE_STEPS):
        power_u = tl.dot(power_u, power_u, input_precision="tf32x3")
        inverse_unit_u = tl.dot(
            inverse_unit_u, identity + power_u, input_precision="tf32x3"
        )
    inverse_u = inverse_unit_u / lu_diagonal[None, :]
    bottom_transform = -tl.dot(
        signed_inverse, inverse_u, input_precision="tf32x3"
    )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    compact_top = tl.where(panel_active, tl.where(r > c, lower, signed_upper), top)
    tl.store(h + matrix_base + rows * n + columns, compact_top)
    tl.store(
        packed_vectors
        + batch * n * PANEL
        + (panel_start + offsets)[:, None] * PANEL
        + offsets[None, :],
        tl.where(r >= c, lower, 0.0),
    )
    tl.store(t_workspace + workspace_base + r * PANEL + c, t_factor)
    tl.store(
        bottom_transform_workspace + workspace_base + r * PANEL + c,
        bottom_transform,
    )
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(tau + batch * n + panel_start + offsets, t_diagonal)


@triton.jit
def _n2048_inverse_upper16(upper, INVERSE_STEPS: tl.constexpr):
    """Invert a 16x16 upper triangle with a bounded Neumann product."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    power = tl.where(r < c, upper / diagonal[:, None], 0.0)
    inverse_unit = identity - power
    if INVERSE_STEPS == 1:
        # These compact-conversion factors are strongly diagonal.  Form the
        # qualified quadratic Neumann inverse directly; multiplying by
        # (I+P^2) spends another MMA only to append the tiny -P^3 term.
        inverse_unit += tl.dot(power.to(tl.float16), power.to(tl.float16))
    else:
        for inverse_step in tl.static_range(0, INVERSE_STEPS):
            power = tl.dot(power.to(tl.float16), power.to(tl.float16))
            inverse_unit = tl.dot(
                inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
            )
    return inverse_unit / diagonal[None, :]


@triton.jit
def _n2048_inverse_unit_lower16(lower, INVERSE_STEPS: tl.constexpr):
    """Invert a 16x16 unit-lower triangle with tensor-core products."""
    offsets = tl.arange(0, 16)
    r = offsets[:, None]
    c = offsets[None, :]
    identity = tl.where(r == c, 1.0, 0.0)
    power = tl.where(r > c, lower, 0.0)
    inverse = identity - power
    for inverse_step in tl.static_range(0, INVERSE_STEPS):
        power = tl.dot(power.to(tl.float16), power.to(tl.float16))
        inverse = tl.dot(
            (identity + power).to(tl.float16), inverse.to(tl.float16)
        )
    return inverse


@triton.jit
def _n2048_chol32_factor_kernel(
    source,
    h,
    partial_gram,
    second_gram_workspace,
    target_weights,
    matrix_workspace,
    inverse_workspace,
    packed_vectors,
    tau_out,
    dense_guard,
    pair_first_vectors,
    pair_cross_workspace,
    panel_start,
    active_chunks,
    n: tl.constexpr,
    CHUNKS: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    CHECK_DENSE: tl.constexpr,
    NEWTON_STEPS: tl.constexpr,
    PAIR_GRAM: tl.constexpr,
    FORM_PAIR_CROSS: tl.constexpr,
    STORE_INVERSE: tl.constexpr,
):
    """Factor and compact-convert one 32-column CholeskyQR panel."""
    batch = tl.program_id(0)
    offsets = tl.arange(0, 32)
    r = offsets[:, None]
    c = offsets[None, :]
    workspace_base = batch * 32 * 32
    partial_stride = 64 * 64 if PAIR_GRAM else 32 * 32
    partial_base = batch * CHUNKS * partial_stride

    gram = tl.zeros((32, 32), tl.float32)
    for chunk in tl.static_range(0, CHUNKS):
        gram += tl.load(
            partial_gram
            + partial_base
            + chunk * partial_stride
            + r * (64 if PAIR_GRAM else 32)
            + c
        )
    gram_diagonal = tl.sum(tl.where(r == c, gram, 0.0), axis=1)
    panel_active = tl.max(tl.abs(gram_diagonal), axis=0) > 1.0e-30

    identity = tl.where(r == c, 1.0, 0.0)
    if NEWTON_STEPS > 0:
        # Solve X^T G X = I directly for the unique upper-triangular inverse
        # Cholesky factor.  At every step C-I is split as S^T+S with S upper,
        # and X <- X(I-S).  The routed dense tail is sufficiently equilibrated
        # that two iterations reduce the error quadratically below FP32 noise.
        scale = tl.sqrt(tl.maximum(gram_diagonal, 1.0e-20))
        normalized = gram / (scale[:, None] * scale[None, :])
        # X starts at I, so C=X^T N X is exactly N in the first Newton
        # iteration.  Writing that iteration explicitly removes three
        # identity tensor products from every panel and is also slightly more
        # accurate because N is not rounded through an MMA on the way in.
        correction = tl.where(r < c, normalized, 0.0)
        correction += tl.where(
            r == c, 0.5 * (normalized - identity), 0.0
        )
        first_correction = correction
        inverse_unit = identity - correction
        for newton_step in tl.static_range(1, NEWTON_STEPS):
            normalized_times_inverse = tl.dot(
                normalized.to(tl.float16),
                inverse_unit.to(tl.float16),
            )
            transformed = tl.dot(
                tl.trans(inverse_unit.to(tl.float16)),
                normalized_times_inverse.to(tl.float16),
            )
            correction = tl.where(r < c, transformed, 0.0)
            correction += tl.where(
                r == c, 0.5 * (transformed - identity), 0.0
            )
            # X=(I-S0), so this update is exactly I-S0-S1+S0*S1.
            # The joint Gram already has FP16 off-diagonals.  Keeping this
            # normalized refinement on FP16 tensor cores avoids paying for
            # wider products that cannot restore information absent ups.tream.
            inverse_unit -= correction
            # Earlier panel errors propagate through every later update, so
            # retain their cubic composition.  In the short late tails the
            # product is fourth-order in the normalized Gram perturbation and
            # the broad transform sweep confirms it is below the error budget.
            if panel_start < 768:
                inverse_unit += tl.dot(
                    first_correction.to(tl.float16),
                    correction.to(tl.float16),
                )
        inverse = inverse_unit / scale[:, None]

        # Recover R=X^-1.  The first-order triangular inverse is sufficient
        # after the two Newton corrections on routed dense panels.
        inverse_diagonal = tl.sum(tl.where(r == c, inverse, 0.0), axis=1)
        power_x = tl.where(r < c, inverse / inverse_diagonal[:, None], 0.0)
        upper_unit = identity - power_x
        upper = upper_unit / inverse_diagonal[None, :]
    else:
        upper = tl.zeros((32, 32), tl.float32)
        for j in tl.range(0, 32, loop_unroll_factor=1):
            gram_row = tl.sum(tl.where(r == j, gram, 0.0), axis=0)
            old_column = tl.sum(tl.where(c == j, upper, 0.0), axis=1)
            products = tl.sum(old_column[:, None] * upper, axis=0)
            diagonal_value = tl.sum(
                tl.where(offsets == j, gram_row - products, 0.0), axis=0
            )
            diagonal = tl.sqrt(tl.maximum(diagonal_value, 1.0e-20))
            new_row = (gram_row - products) / diagonal
            new_row = tl.where(offsets == j, diagonal, new_row)
            new_row = tl.where(offsets >= j, new_row, 0.0)
            upper = tl.where(r == j, new_row[None, :], upper)

        diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
        power = tl.where(r < c, upper / diagonal[:, None], 0.0)
        inverse_unit = identity - power
        # P^16 is already below FP32 roundoff for the routed tall,
        # equilibrated panels; avoid a redundant P^16/P^32 correction pair.
        for inverse_step in tl.static_range(0, 3):
            power = tl.dot(power.to(tl.float16), power.to(tl.float16))
            inverse_unit = tl.dot(
                inverse_unit.to(tl.float16), (identity + power).to(tl.float16)
            )
        inverse = inverse_unit / diagonal[None, :]

    cholesky_diagonal = tl.sum(tl.where(r == c, upper, 0.0), axis=1)
    panel_scale = tl.sqrt(tl.max(tl.abs(gram_diagonal), axis=0))
    panel_active &= (
        tl.min(tl.abs(cholesky_diagonal), axis=0)
        > tl.maximum(panel_scale * 1.0e-4, 1.0e-20)
    )
    if CHECK_DENSE:
        panel_active &= tl.load(dense_guard + batch) != 0

    matrix_base = batch * n * n
    rows = panel_start + r
    columns = panel_start + c
    matrix = source if FIRST_PANEL else h
    top = tl.load(matrix + matrix_base + rows * n + columns)
    q_top = tl.dot(top, inverse, input_precision="tf32")
    q_diagonal = tl.sum(tl.where(r == c, q_top, 0.0), axis=1)
    signs = tl.where(q_diagonal >= 0.0, -1.0, 1.0)
    signed_inverse = tl.where(panel_active, inverse * signs[None, :], 0.0)
    if STORE_INVERSE:
        tl.store(
            dense_guard + workspace_base + r * 32 + c,
            signed_inverse,
        )
    signed_upper = tl.where(panel_active, signs[:, None] * upper, top)
    m_factor = tl.where(
        panel_active, identity - q_top * signs[None, :], identity
    )

    tl.store(
        h + matrix_base + rows * n + columns,
        signed_upper,
        mask=r <= c,
    )

    # Continue in the same CTA with the compact conversion.  The LU factors
    # are explicitly spilled below before their triangular inverses, bounding
    # peak live state while avoiding an unnecessary M/X round trip here.
    panel_active = tl.max(tl.abs(signed_inverse), axis=0) > 0.0
    upper_lu = tl.where(r <= c, m_factor, 0.0)
    upper_diagonal = tl.sum(tl.where(r == c, upper_lu, 0.0), axis=1)
    lower = tl.where(
        r == c,
        1.0,
        tl.where(r > c, m_factor / upper_diagonal[None, :], 0.0),
    )
    lower_strict = tl.where(r > c, lower, 0.0)
    upper_strict = tl.where(r < c, upper_lu, 0.0)
    residual = tl.where(
        r > c,
        m_factor - lower_strict * upper_diagonal[None, :],
        0.0,
    )
    residual -= tl.dot(
        lower_strict.to(tl.float16), upper_strict.to(tl.float16)
    )
    upper_lu += tl.where(r <= c, residual, 0.0)
    upper_diagonal = tl.sum(tl.where(r == c, upper_lu, 0.0), axis=1)
    lower += tl.where(r > c, residual / upper_diagonal[None, :], 0.0)

    packed_base = packed_vectors + batch * n * 32 + panel_start * 32
    packed_top = packed_base + r * 32 + c
    tl.store(packed_top, lower)
    tl.store(matrix_workspace + workspace_base + r * 32 + c, upper_lu)
    tl.debug_barrier()

    half = tl.arange(0, 16)
    half_r = half[:, None]
    half_c = half[None, :]
    block_base = matrix_workspace + workspace_base
    zero16 = tl.zeros((16, 16), tl.float32)
    lower00 = tl.load(packed_base + half_r * 32 + half_c)
    lower10 = tl.load(packed_base + (16 + half_r) * 32 + half_c)
    lower11 = tl.load(packed_base + (16 + half_r) * 32 + 16 + half_c)
    upper00 = tl.load(block_base + half_r * 32 + half_c)
    upper01 = tl.load(block_base + half_r * 32 + 16 + half_c)
    upper11 = tl.load(block_base + (16 + half_r) * 32 + 16 + half_c)
    # The first Neumann term is sufficient for the strongly diagonal lower
    # factors; upper inverses retain their quadratic correction.
    inverse_lower00 = _n2048_inverse_unit_lower16(lower00, 1)
    inverse_upper00 = _n2048_inverse_upper16(upper00, 1)
    inverse_lower11 = _n2048_inverse_unit_lower16(lower11, 1)
    inverse_lt00 = tl.trans(inverse_lower00)
    inverse_lt11 = tl.trans(inverse_lower11)
    inverse_lt01 = -tl.dot(
        tl.dot(
            inverse_lt00.to(tl.float16),
            tl.trans(lower10).to(tl.float16),
        ).to(tl.float16),
        inverse_lt11.to(tl.float16),
    )
    inverse_lower_transpose = tl.cat(
        tl.cat(inverse_lt00, inverse_lt01, dim=1),
        tl.cat(zero16, inverse_lt11, dim=1),
        dim=0,
    )
    upper_lu = tl.load(block_base + r * 32 + c)
    t_factor = tl.dot(
        upper_lu.to(tl.float16), inverse_lower_transpose.to(tl.float16)
    )
    t_factor = tl.where(panel_active, t_factor, 0.0)

    inverse_upper11 = _n2048_inverse_upper16(upper11, 1)
    inverse_upper01 = -tl.dot(
        tl.dot(
            inverse_upper00.to(tl.float16), upper01.to(tl.float16)
        ).to(tl.float16),
        inverse_upper11.to(tl.float16),
    )
    inverse_u = tl.cat(
        tl.cat(inverse_upper00, inverse_upper01, dim=1),
        tl.cat(zero16, inverse_upper11, dim=1),
        dim=0,
    )
    bottom_transform = -tl.dot(
        signed_inverse.to(tl.float16), inverse_u.to(tl.float16)
    )
    bottom_transform = tl.where(panel_active, bottom_transform, 0.0)

    lower = tl.load(packed_top)
    tl.store(
        h + matrix_base + rows * n + columns,
        lower,
        mask=r > c,
    )
    tl.store(
        packed_vectors
        + batch * n * 32
        + (panel_start + offsets)[:, None] * 32
        + offsets[None, :],
        tl.where(r >= c, lower, 0.0),
    )
    tl.store(matrix_workspace + workspace_base + r * 32 + c, t_factor)
    tl.store(inverse_workspace + workspace_base + r * 32 + c, bottom_transform)
    t_diagonal = tl.sum(tl.where(r == c, t_factor, 0.0), axis=1)
    tl.store(tau_out + batch * n + panel_start + offsets, t_diagonal)

    if FORM_PAIR_CROSS:
        pair_k = tl.load(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start - 32 + offsets)[None, :]
        )
        first_mid = tl.load(
            pair_first_vectors
            + batch * n * 32
            + (panel_start + offsets)[:, None] * 32
            + offsets[None, :]
        )
        pair_cross = tl.dot(
            tl.trans(bottom_transform.to(tl.float16)),
            pair_k.to(tl.float16),
        )
        pair_cross += tl.dot(
            tl.trans(inverse_u.to(tl.float16)),
            first_mid.to(tl.float16),
        )
        tl.store(
            pair_cross_workspace + workspace_base + r * 32 + c,
            pair_cross,
        )

    if PAIR_GRAM:
        # The full 64-column Gram contains everything needed to update and
        # factor the adjacent panel.  R01=Q0.T@A1, while the active Gram after
        # the orthogonal first-panel transform is C-R01.T@R01.
        cross_gram = tl.zeros((32, 32), tl.float32)
        second_gram = tl.zeros((32, 32), tl.float32)
        for chunk in tl.static_range(0, CHUNKS):
            pair_base = partial_base + chunk * 64 * 64
            cross_gram += tl.load(
                partial_gram + pair_base + r * 64 + 32 + c
            )
            second_gram += tl.load(
                partial_gram + pair_base + (32 + r) * 64 + 32 + c
            )
        r01 = tl.dot(
            tl.trans(signed_inverse.to(tl.float16)),
            cross_gram.to(tl.float16),
        )
        target_columns = panel_start + 32 + c
        raw_target_top = tl.load(
            matrix + matrix_base + rows * n + target_columns
        )
        inverse_lower = tl.trans(inverse_lower_transpose)
        target_coefficients = tl.dot(
            inverse_lower.to(tl.float16),
            (raw_target_top - r01).to(tl.float16),
        )
        tl.store(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start + 32 + offsets)[None, :],
            target_coefficients,
        )

        inverse_u_times_w = tl.dot(
            inverse_u.to(tl.float16),
            target_coefficients.to(tl.float16),
        )
        lower_for_cross = tl.load(packed_top)
        pair_k = -tl.dot(
            tl.trans((inverse_u_times_w + r01).to(tl.float16)),
            lower_for_cross.to(tl.float16),
        )
        tl.store(
            target_weights
            + batch * 32 * n
            + offsets[:, None] * n
            + (panel_start + offsets)[None, :],
            pair_k,
        )

        tl.store(
            h + matrix_base + rows * n + target_columns,
            r01,
        )
        r01_product = tl.dot(
            tl.trans(r01.to(tl.float16)), r01.to(tl.float16)
        )
        r01_diagonal = tl.sum(r01 * r01, axis=0)
        r01_product = tl.where(
            r == c, r01_diagonal[:, None], r01_product
        )
        tl.store(
            second_gram_workspace + batch * 32 * 32 + r * 32 + c,
            second_gram - r01_product,
        )

        # Only these 32 rows are a dependency of the adjacent factor.  Form
        # them here while the first-panel transform and target coefficients
        # are resident.  Far first-panel rows can then overlap the adjacent
        # factor; a follow-up launch recovers only the second reflector tail.
        middle_rows = panel_start + 32 + r
        raw_first_middle = tl.load(
            matrix + matrix_base + middle_rows * n + columns
        )
        first_middle = tl.dot(
            raw_first_middle.to(tl.float16),
            bottom_transform.to(tl.float16),
        )
        raw_target_middle = tl.load(
            matrix + matrix_base + middle_rows * n + target_columns
        )
        updated_target_middle = raw_target_middle - tl.dot(
            first_middle.to(tl.float16),
            target_coefficients.to(tl.float16),
        )
        tl.store(
            h + matrix_base + middle_rows * n + columns,
            first_middle,
        )
        tl.store(
            packed_vectors
            + batch * n * 32
            + middle_rows * 32
            + c,
            first_middle,
        )
        tl.store(
            h + matrix_base + middle_rows * n + target_columns,
            updated_target_middle,
        )


@triton.jit
def _n2048_second_factor_extract_first_kernel(
    source,
    h,
    second_gram,
    first_t,
    target_weights,
    second_t,
    first_bottom,
    second_bottom,
    first_vectors,
    second_vectors,
    tau,
    pair_cross,
    second_inverse,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
    STORE_INVERSE: tl.constexpr,
    NEWTON_STEPS: tl.constexpr,
):
    """Overlap the adjacent factor with independent first-tail recovery."""
    batch = tl.program_id(0)
    work_index = tl.program_id(1)
    if work_index == 0:
        _n2048_chol32_factor_kernel(
            source,
            h,
            second_gram,
            first_t,
            target_weights,
            second_t,
            second_bottom,
            second_vectors,
            tau,
            second_inverse,
            first_vectors,
            pair_cross,
            pair_start + PANEL,
            1,
            n=n,
            CHUNKS=1,
            FIRST_PANEL=False,
            CHECK_DENSE=False,
            NEWTON_STEPS=NEWTON_STEPS,
            PAIR_GRAM=False,
            FORM_PAIR_CROSS=True,
            STORE_INVERSE=STORE_INVERSE,
        )
    else:
        rows = (
            pair_start
            + 2 * PANEL
            + (work_index - 1) * BLOCK_M
            + tl.arange(0, BLOCK_M)
        )
        offsets = tl.arange(0, PANEL)
        first_columns = pair_start + offsets
        second_columns = pair_start + PANEL + offsets
        matrix_base = batch * n * n
        matrix = source if FIRST_PAIR else h
        if n == 4096:
            # The n4096 pair grid has no partial 16-row extraction tiles.
            raw_first = tl.load(
                matrix + matrix_base + rows[:, None] * n + first_columns[None, :]
            )
        else:
            mask = rows[:, None] < n
            raw_first = tl.load(
                matrix + matrix_base + rows[:, None] * n + first_columns[None, :],
                mask=mask,
                other=0.0,
            )
        first_transform = tl.load(
            first_bottom
            + batch * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        first = tl.dot(
            raw_first.to(tl.float16), first_transform.to(tl.float16)
        )
        coefficients = tl.load(
            target_weights
            + batch * PANEL * n
            + offsets[:, None] * n
            + second_columns[None, :]
        )
        if n == 4096:
            raw_second = tl.load(
                matrix + matrix_base + rows[:, None] * n + second_columns[None, :]
            )
        else:
            raw_second = tl.load(
                matrix + matrix_base + rows[:, None] * n + second_columns[None, :],
                mask=mask,
                other=0.0,
            )
        updated_second = raw_second - tl.dot(
            first.to(tl.float16), coefficients.to(tl.float16)
        )
        if n == 4096:
            tl.store(
                h + matrix_base + rows[:, None] * n + first_columns[None, :],
                first,
            )
            tl.store(
                h + matrix_base + rows[:, None] * n + second_columns[None, :],
                updated_second,
            )
            tl.store(
                first_vectors
                + batch * n * PANEL
                + rows[:, None] * PANEL
                + offsets[None, :],
                first,
            )
        else:
            tl.store(
                h + matrix_base + rows[:, None] * n + first_columns[None, :],
                first,
                mask=mask,
            )
            tl.store(
                h + matrix_base + rows[:, None] * n + second_columns[None, :],
                updated_second,
                mask=mask,
            )
            tl.store(
                first_vectors
                + batch * n * PANEL
                + rows[:, None] * PANEL
                + offsets[None, :],
                first,
                mask=mask,
            )


@triton.jit
def _n2048_chol_extract_bottom_kernel(
    source,
    h,
    bottom_transform_workspace,
    packed_vectors,
    target_weights,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    APPLY_TARGET: tl.constexpr,
):
    """Materialize the recovered reflector tails in independent row tiles."""
    row_tile = tl.program_id(0)
    batch = tl.program_id(1)
    rows = panel_start + PANEL + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = panel_start + tl.arange(0, PANEL)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    values = tl.load(
        matrix + matrix_base + rows[:, None] * n + columns[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    offsets = tl.arange(0, PANEL)
    transform = tl.load(
        bottom_transform_workspace
        + batch * PANEL * PANEL
        + offsets[:, None] * PANEL
        + offsets[None, :]
    )
    vectors = tl.dot(values.to(tl.float16), transform.to(tl.float16))
    mask = rows[:, None] < n
    if APPLY_TARGET:
        target_columns = panel_start + PANEL + offsets
        coefficients = tl.load(
            target_weights
            + batch * PANEL * n
            + offsets[:, None] * n
            + target_columns[None, :]
        )
        target = tl.load(
            matrix + matrix_base + rows[:, None] * n + target_columns[None, :],
            mask=mask,
            other=0.0,
        )
        target_update = tl.dot(
            vectors.to(tl.float16), coefficients.to(tl.float16)
        )
        tl.store(
            h + matrix_base + rows[:, None] * n + target_columns[None, :],
            target - target_update,
            mask=mask,
        )
    tl.store(
        h + matrix_base + rows[:, None] * n + columns[None, :],
        vectors,
        mask=mask,
    )
    tl.store(
        packed_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        vectors,
        mask=mask,
    )


@triton.jit
def _n2048_extract_second_group_target_kernel(
    source,
    h,
    target_partials,
    pair_partials,
    first_inverse,
    second_inverse,
    first_vectors,
    second_vectors,
    second_bottom,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    CHUNKS: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
    QUADRATIC_LOWER: tl.constexpr,
):
    """Overlap algebraic target conversion with second-tail extraction."""
    batch = tl.program_id(0)
    work_index = tl.program_id(1)
    if work_index == 0:
        _n2048_group_target_from_gram_kernel(
            source,
            h,
            target_partials,
            pair_partials,
            first_inverse,
            second_inverse,
            first_vectors,
            second_vectors,
            weights,
            pair_start,
            n=n,
            CHUNKS=CHUNKS,
            FIRST_GROUP=FIRST_GROUP,
            QUADRATIC_LOWER=QUADRATIC_LOWER,
        )
    else:
        panel_start = pair_start + PANEL
        rows = (
            panel_start
            + PANEL
            + (work_index - 1) * BLOCK_M
            + tl.arange(0, BLOCK_M)
        )
        offsets = tl.arange(0, PANEL)
        columns = panel_start + offsets
        matrix_base = batch * n * n
        values = tl.load(
            h + matrix_base + rows[:, None] * n + columns[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        transform = tl.load(
            second_bottom
            + batch * PANEL * PANEL
            + offsets[:, None] * PANEL
            + offsets[None, :]
        )
        vectors = tl.dot(values.to(tl.float16), transform.to(tl.float16))
        mask = rows[:, None] < n
        tl.store(
            h + matrix_base + rows[:, None] * n + columns[None, :],
            vectors,
            mask=mask,
        )
        tl.store(
            second_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            vectors,
            mask=mask,
        )


@triton.jit
def _n2048_chol_form_weights_kernel(
    source,
    h,
    packed_vectors,
    t_workspace,
    weights,
    panel_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
):
    """Form T^T V^T A for one trailing-column tile."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    panel_offsets = tl.arange(0, PANEL)
    columns = panel_start + PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    projection = tl.zeros((PANEL, BLOCK_N), tl.float32)
    matrix_base = batch * n * n
    matrix = source if FIRST_PANEL else h
    for row_start in tl.range(0, n - panel_start, BLOCK_K, num_stages=3):
        rows = panel_start + row_start + tl.arange(0, BLOCK_K)
        vectors = tl.load(
            packed_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + panel_offsets[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=(rows[:, None] < n) & (columns[None, :] < n),
            other=0.0,
        )
        projection += tl.dot(tl.trans(vectors), values.to(tl.float16))

    t_factor = tl.load(
        t_workspace
        + batch * PANEL * PANEL
        + panel_offsets[:, None] * PANEL
        + panel_offsets[None, :]
    )
    coefficients = tl.dot(
        tl.trans(t_factor.to(tl.float16)), projection.to(tl.float16)
    )
    tl.store(
        weights
        + batch * PANEL * n
        + panel_offsets[:, None] * n
        + columns[None, :],
        coefficients,
        mask=columns[None, :] < n,
    )


@triton.jit
def _n2048_factor_panel_kernel(
    a,
    h,
    packed_vectors,
    tau_out,
    gram_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Factor one panel while retaining the panel in CTA registers."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    rows = tl.arange(0, BLOCK_M)
    cols = tl.arange(0, BLOCK_B)
    global_rows = panel_start + rows
    global_cols = panel_start + cols
    matrix_base = batch * n * n
    input_base = a if FIRST_PANEL else h
    ptrs = input_base + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    valid = (global_rows[:, None] < n) & (global_cols[None, :] < n)
    panel = tl.load(ptrs, mask=valid, other=0.0)
    for j in tl.static_range(0, BLOCK_B):
        active_col = panel_start + j < n
        column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
        tail = (rows >= j) & (global_rows < n) & active_col
        norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
        alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        nonzero = norm > 0.0
        tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
        safe_denominator = tl.where(nonzero, alpha - beta, 1.0)
        vector = tl.where(
            rows == j,
            1.0,
            tl.where((rows > j) & (global_rows < n), column / safe_denominator, 0.0),
        )
        products = tl.sum(vector[:, None] * panel, axis=0)
        tl.store(
            gram_out + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + cols,
            products,
            mask=cols < j,
        )
        update_mask = (
            (rows[:, None] >= j)
            & (cols[None, :] > j)
            & valid
            & active_col
        )
        panel = tl.where(
            update_mask,
            panel - tau * vector[:, None] * products[None, :],
            panel,
        )
        replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
        panel = tl.where(
            (cols[None, :] == j) & (rows[:, None] >= j) & valid,
            replacement,
            panel,
        )
        tl.store(tau_out + batch * n + panel_start + j, tau, mask=active_col)

    output_ptrs = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    tl.store(output_ptrs, panel, mask=valid)
    packed_ptrs = (
        packed_vectors
        + batch * n * BLOCK_B
        + global_rows[:, None] * BLOCK_B
        + cols[None, :]
    )
    packed_panel = tl.where(
        global_rows[:, None] == global_cols[None, :],
        1.0,
        tl.where(global_rows[:, None] > global_cols[None, :], panel, 0.0),
    )
    tl.store(packed_ptrs, packed_panel, mask=valid)


@triton.jit
def _n2048_factor_panel_tail_kernel(
    h,
    tau_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Factor a late panel that will be applied directly, without WY metadata."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    rows = tl.arange(0, BLOCK_M)
    cols = tl.arange(0, BLOCK_B)
    global_rows = panel_start + rows
    global_cols = panel_start + cols
    matrix_base = batch * n * n
    ptrs = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
    valid = global_rows[:, None] < n
    panel = tl.load(ptrs, mask=valid, other=0.0)
    for j in tl.static_range(0, BLOCK_B):
        column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
        tail = (rows >= j) & (global_rows < n)
        norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
        alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
        beta = tl.where(alpha >= 0.0, -norm, norm)
        nonzero = norm > 0.0
        tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
        safe_denominator = tl.where(nonzero, alpha - beta, 1.0)
        vector = tl.where(
            rows == j,
            1.0,
            tl.where((rows > j) & (global_rows < n), column / safe_denominator, 0.0),
        )
        products = tl.sum(vector[:, None] * panel, axis=0)
        update_mask = (
            (rows[:, None] >= j)
            & (cols[None, :] > j)
            & valid
        )
        panel = tl.where(
            update_mask,
            panel - tau * vector[:, None] * products[None, :],
            panel,
        )
        replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
        panel = tl.where(
            (cols[None, :] == j) & (rows[:, None] >= j) & valid,
            replacement,
            panel,
        )
        tl.store(tau_out + batch * n + panel_start + j, tau)

    tl.store(ptrs, panel, mask=valid)


@triton.jit
def _n2048_factor_apply_fallback_tail_kernel(
    h,
    tau_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_N: tl.constexpr,
    PANEL_COUNT: tl.constexpr,
):
    """Fuse the exact scalar-reflector fallback at the short matrix tail."""
    batch = tl.program_id(0)
    if tl.load(route_flags + batch) != 0:
        return
    first_panel_start = panel_start
    matrix_base = batch * n * n
    for panel_index in tl.range(0, PANEL_COUNT):
        panel_start = first_panel_start + panel_index * BLOCK_B
        rows = tl.arange(0, BLOCK_M)
        cols = tl.arange(0, BLOCK_B)
        global_rows = panel_start + rows
        global_cols = panel_start + cols
        pointers = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
        valid = global_rows[:, None] < n
        panel = tl.load(pointers, mask=valid, other=0.0)
        tau_values = tl.zeros((BLOCK_B,), tl.float32)
    
        for j in tl.static_range(0, BLOCK_B):
            column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
            tail = (rows >= j) & (global_rows < n)
            norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
            alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
            beta = tl.where(alpha >= 0.0, -norm, norm)
            nonzero = norm > 0.0
            tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
            denominator = tl.where(nonzero, alpha - beta, 1.0)
            vector = tl.where(
                rows == j,
                1.0,
                tl.where(
                    (rows > j) & (global_rows < n), column / denominator, 0.0
                ),
            )
            products = tl.sum(vector[:, None] * panel, axis=0)
            panel = tl.where(
                (rows[:, None] >= j) & (cols[None, :] > j) & valid,
                panel - tau * vector[:, None] * products[None, :],
                panel,
            )
            replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
            panel = tl.where(
                (cols[None, :] == j) & (rows[:, None] >= j) & valid,
                replacement,
                panel,
            )
            tau_values = tl.where(cols == j, tau, tau_values)
    
        tl.store(pointers, panel, mask=valid)
        tl.store(tau_out + batch * n + panel_start + cols, tau_values)
        for col_block in tl.range(0, tl.cdiv(n - panel_start - BLOCK_B, BLOCK_N)):
            trailing_cols = (
                panel_start
                + BLOCK_B
                + col_block * BLOCK_N
                + tl.arange(0, BLOCK_N)
            )
            c = tl.load(
                h
                + matrix_base
                + global_rows[:, None] * n
                + trailing_cols[None, :],
                mask=(global_rows[:, None] < n) & (trailing_cols[None, :] < n),
                other=0.0,
            )
            for j in tl.static_range(0, BLOCK_B):
                packed_column = tl.sum(
                    tl.where(cols[None, :] == j, panel, 0.0), axis=1
                )
                reflector = tl.where(
                    rows == j,
                    1.0,
                    tl.where(rows > j, packed_column, 0.0),
                )
                tau = tl.sum(tl.where(cols == j, tau_values, 0.0), axis=0)
                product = tl.sum(reflector[:, None] * c, axis=0)
                c -= (tau * reflector)[:, None] * product[None, :]
            output_ptrs = (
                h
                + matrix_base
                + global_rows[:, None] * n
                + trailing_cols[None, :]
            )
            output_mask = (global_rows[:, None] < n) & (trailing_cols[None, :] < n)
            tl.store(output_ptrs, c, mask=output_mask)
        tl.debug_barrier()


@triton.jit
def _n2048_make_block_weights_kernel(
    a,
    h,
    packed_vectors,
    tau,
    gram_in,
    weights,
    route_flags,
    n: tl.constexpr,
    panel_start,
    PANEL_B: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Compute the coefficients for applying a panel's reflectors to A."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    col_block = tl.program_id(1)
    panel_cols = tl.arange(0, BLOCK_B)
    active_panel_col = panel_cols < PANEL_B
    out_cols = panel_start + PANEL_B + col_block * BLOCK_N + tl.arange(0, BLOCK_N)

    dots = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
    matrix_base = batch * n * n

    for row_start in tl.range(0, n - panel_start, BLOCK_K):
        rows = panel_start + row_start + tl.arange(0, BLOCK_K)
        vectors = tl.load(
            packed_vectors
            + batch * n * BLOCK_B
            + rows[:, None] * BLOCK_B
            + panel_cols[None, :],
            mask=(rows[:, None] < n) & active_panel_col[None, :],
            other=0.0,
        )
        matrix_input = a if FIRST_PANEL else h
        a_tile = tl.load(
            matrix_input + matrix_base + rows[:, None] * n + out_cols[None, :],
            mask=(rows[:, None] < n) & (out_cols[None, :] < n),
            other=0.0,
        )
        dots += tl.dot(
            tl.trans(vectors),
            a_tile.to(tl.float16),
        )

    coefficients = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
    for j in tl.static_range(0, BLOCK_B):
        gram_row = tl.load(
            gram_in + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + panel_cols,
            mask=panel_cols < j,
            other=0.0,
        )
        dot_row = tl.sum(
            tl.where(panel_cols[:, None] == j, dots, 0.0),
            axis=0,
        )
        previous = tl.where(panel_cols < j, gram_row, 0.0)
        correction = tl.sum(previous[:, None] * coefficients, axis=0)
        tau_j = tl.load(
            tau + batch * n + panel_start + j,
            mask=j < PANEL_B,
            other=0.0,
        )
        row = tau_j * (dot_row - correction)
        coefficients = tl.where(panel_cols[:, None] == j, row[None, :], coefficients)

    weight_ptrs = (
        weights
        + batch * BLOCK_B * n
        + panel_cols[:, None] * n
        + out_cols[None, :]
    )
    tl.store(
        weight_ptrs,
        coefficients,
        mask=out_cols[None, :] < n,
    )


@triton.jit
def _n2048_apply_block_kernel(
    a,
    h,
    packed_vectors,
    weights,
    route_flags,
    n: tl.constexpr,
    panel_start,
    PANEL_B: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PANEL: tl.constexpr,
    GUARDED: tl.constexpr,
):
    """Apply V @ weights to a rectangular tile of the trailing matrix."""
    batch = tl.program_id(0)
    if GUARDED:
        if tl.load(route_flags + batch) != 0:
            return
    row_block = tl.program_id(1)
    col_block = tl.program_id(2)
    rows = panel_start + row_block * BLOCK_M + tl.arange(0, BLOCK_M)
    cols = panel_start + PANEL_B + col_block * BLOCK_N + tl.arange(0, BLOCK_N)
    panel_cols = tl.arange(0, BLOCK_B)
    matrix_base = batch * n * n

    active_panel_col = panel_cols < PANEL_B
    vectors = tl.load(
        packed_vectors
        + batch * n * BLOCK_B
        + rows[:, None] * BLOCK_B
        + panel_cols[None, :],
        mask=(rows[:, None] < n) & active_panel_col[None, :],
        other=0.0,
    )
    coefficients = tl.load(
        weights
        + batch * BLOCK_B * n
        + panel_cols[:, None] * n
        + cols[None, :],
        mask=cols[None, :] < n,
        other=0.0,
    )
    update = tl.dot(vectors, coefficients.to(tl.float16))
    a_ptrs = h + matrix_base + rows[:, None] * n + cols[None, :]
    mask = (rows[:, None] < n) & (cols[None, :] < n)
    matrix_input = a if FIRST_PANEL else h
    input_ptrs = matrix_input + matrix_base + rows[:, None] * n + cols[None, :]
    values = tl.load(input_ptrs, mask=mask, other=0.0)
    tl.store(a_ptrs, values - update, mask=mask)


@triton.jit
def _n2048_pair_form_weights_kernel(
    source,
    h,
    first_vectors,
    second_vectors,
    first_t,
    second_t,
    cross_workspace,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Form both coefficient blocks for one aggregated 64-reflector update."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    projection_t = tl.zeros((BLOCK_N, 2 * PANEL), dtype=tl.float32)
    matrix_base = batch * n * n
    for row_start in tl.range(0, n - pair_start, BLOCK_K, num_stages=3):
        rows = pair_start + row_start + tl.arange(0, BLOCK_K)
        valid_rows = rows < n
        first = tl.load(
            first_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None],
            other=0.0,
        )
        second = tl.load(
            second_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= pair_start + PANEL),
            other=0.0,
        )
        matrix = source if FIRST_PAIR else h
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
            mask=valid_rows[:, None] & (columns[None, :] < n),
            other=0.0,
        ).to(tl.float16)
        # Project both panel-vector blocks with one logical Kx64 tensor-core
        # operand so the target tile is consumed by a single MMA sequence.
        combined_vectors = tl.cat(first, second, dim=1)
        projection_t += tl.dot(tl.trans(values), combined_vectors)

    projection_blocks_t = tl.permute(
        tl.reshape(projection_t, (BLOCK_N, 2, PANEL)), (0, 2, 1)
    )
    first_projection_t, second_projection_t = tl.split(projection_blocks_t)

    workspace_offsets = offsets[:, None] * PANEL + offsets[None, :]
    first_transform = tl.load(
        first_t + batch * PANEL * PANEL + workspace_offsets
    )
    second_transform = tl.load(
        second_t + batch * PANEL * PANEL + workspace_offsets
    )
    cross = tl.load(
        cross_workspace + batch * PANEL * PANEL + workspace_offsets
    )
    first_weights_t = tl.dot(
        first_projection_t.to(tl.float16),
        first_transform.to(tl.float16),
    )
    second_rhs_t = second_projection_t - tl.dot(
        first_weights_t.to(tl.float16), tl.trans(cross.to(tl.float16))
    )
    second_weights_t = tl.dot(
        second_rhs_t.to(tl.float16),
        second_transform.to(tl.float16),
    )
    weight_rows = tl.arange(0, 2 * PANEL)[None, :]
    combined_t = tl.cat(first_weights_t, second_weights_t, dim=1)
    tl.store(
        weights + batch * (2 * PANEL) * n + weight_rows * n + columns[:, None],
        combined_t,
        mask=columns[:, None] < n,
    )


@triton.jit
def _n2048_pair_apply_kernel(
    source,
    h,
    first_vectors,
    second_vectors,
    weights,
    pair_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_PAIR: tl.constexpr,
):
    """Apply two consecutive reflector blocks in one matrix read/write pass."""
    batch = tl.program_id(0)
    row_tile = tl.program_id(1)
    column_tile = tl.program_id(2)
    rows = pair_start + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = pair_start + 2 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    offsets = tl.arange(0, PANEL)
    valid = (rows[:, None] < n) & (columns[None, :] < n)
    first = tl.load(
        first_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    second = tl.load(
        second_vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=(rows[:, None] < n) & (rows[:, None] >= pair_start + PANEL),
        other=0.0,
    )
    first_weights = tl.load(
        weights
        + batch * (2 * PANEL) * n
        + offsets[:, None] * n
        + columns[None, :],
        mask=columns[None, :] < n,
        other=0.0,
    )
    second_weights = tl.load(
        weights
        + batch * (2 * PANEL) * n
        + (PANEL + offsets)[:, None] * n
        + columns[None, :],
        mask=columns[None, :] < n,
        other=0.0,
    )
    update = tl.dot(
        first, first_weights.to(tl.float16), out_dtype=tl.float16
    )
    update += tl.dot(
        second, second_weights.to(tl.float16), out_dtype=tl.float16
    )
    pointers = h + batch * n * n + rows[:, None] * n + columns[None, :]
    matrix = source if FIRST_PAIR else h
    values = tl.load(
        matrix + batch * n * n + rows[:, None] * n + columns[None, :],
        mask=valid,
        other=0.0,
    )
    tl.store(pointers, values - update, mask=valid)


@triton.jit
def _n2048_pair_group_cross_kernel(
    first0,
    first1,
    second0,
    second1,
    cross_workspace,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
):
    """Form independent partial cross blocks between two panel pairs."""
    chunk = tl.program_id(0)
    cross_block = tl.program_id(1)
    batch = tl.program_id(2)
    offsets = tl.arange(0, PANEL)
    rows = group_start + 2 * PANEL + chunk * ROW_CHUNK + tl.arange(0, ROW_CHUNK)
    valid = rows < n
    earlier_ptr = first0 if cross_block % 2 == 0 else first1
    later_ptr = second0 if cross_block < 2 else second1
    earlier = tl.load(
        earlier_ptr
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=valid[:, None], other=0.0,
    )
    later = tl.load(
        later_ptr
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=valid[:, None]
        & ((cross_block < 2) | (rows[:, None] >= group_start + 3 * PANEL)),
        other=0.0,
    )
    cross = tl.dot(tl.trans(later), earlier)

    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    base = (batch * CHUNKS + chunk) * 4 * PANEL * PANEL
    tl.store(
        cross_workspace + base + cross_block * PANEL * PANEL + matrix_offsets,
        cross,
    )


@triton.jit
def _n2048_pair_group_cross_reduce_kernel(
    partials,
    cross_workspace,
    active_chunks,
    PANEL: tl.constexpr,
    CHUNKS: tl.constexpr,
):
    cross_block = tl.program_id(0)
    batch = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    cross = tl.zeros((PANEL, PANEL), tl.float32)
    base = batch * CHUNKS * 4 * PANEL * PANEL
    for chunk in tl.static_range(0, CHUNKS):
        chunk_base = base + chunk * 4 * PANEL * PANEL
        mask = chunk < active_chunks
        cross += tl.load(
            partials
            + chunk_base
            + cross_block * PANEL * PANEL
            + matrix_offsets,
            mask=mask, other=0.0,
        )
    output_base = batch * 4 * PANEL * PANEL
    tl.store(
        cross_workspace
        + output_base
        + cross_block * PANEL * PANEL
        + matrix_offsets,
        cross,
    )


@triton.jit
def _n2048_precompose_group_cross_kernel(
    first0,
    first1,
    second0,
    second1,
    cross_workspace,
    transform0,
    transform1,
    transform2,
    transform3,
    within_cross0,
    within_cross1,
    pair_transforms,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    CHUNKS: tl.constexpr,
    ROW_CHUNK: tl.constexpr,
):
    """Form pair-to-pair crosses and precompose each adjacent pair.

    Cross partials occupy the first ``4 * CHUNKS`` tasks.  Two additional
    tasks build the block-lower transforms that map a pair's two projections
    directly to its two weight blocks.  Compensated FP16 products keep the
    composite off-diagonal accurate before its final FP16 storage.
    """
    task = tl.program_id(0)
    batch = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]

    if task >= 4 * CHUNKS:
        pair = task - 4 * CHUNKS
        if pair == 0:
            first_t = tl.load(
                transform0 + batch * PANEL * PANEL + matrix_offsets
            )
            second_t = tl.load(
                transform1 + batch * PANEL * PANEL + matrix_offsets
            )
            pair_cross = tl.load(
                within_cross0 + batch * PANEL * PANEL + matrix_offsets
            )
        else:
            first_t = tl.load(
                transform2 + batch * PANEL * PANEL + matrix_offsets
            )
            second_t = tl.load(
                transform3 + batch * PANEL * PANEL + matrix_offsets
            )
            pair_cross = tl.load(
                within_cross1 + batch * PANEL * PANEL + matrix_offsets
            )
        d0 = tl.trans(first_t)
        d1 = tl.trans(second_t)
        middle = tl.dot(pair_cross.to(tl.float16), d0.to(tl.float16))
        lower_left = -_compensated_fp16_dot_rhs_residual(d1, middle)
        pair_base = (batch * 2 + pair) * (2 * PANEL) * (2 * PANEL)
        leading = (
            pair_transforms
            + pair_base
            + offsets[:, None] * (2 * PANEL)
            + offsets[None, :]
        )
        tl.store(leading, d0)
        tl.store(leading + PANEL, 0.0)
        tl.store(leading + PANEL * (2 * PANEL), lower_left)
        tl.store(leading + PANEL * (2 * PANEL) + PANEL, d1)
    else:
        chunk = task // 4
        cross_block = task % 4
        rows = (
            group_start
            + 2 * PANEL
            + chunk * ROW_CHUNK
            + tl.arange(0, ROW_CHUNK)
        )
        valid = rows < n
        earlier_ptr = first0 if cross_block % 2 == 0 else first1
        later_ptr = second0 if cross_block < 2 else second1
        earlier = tl.load(
            earlier_ptr
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid[:, None],
            other=0.0,
        )
        later = tl.load(
            later_ptr
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=valid[:, None]
            & ((cross_block < 2) | (rows[:, None] >= group_start + 3 * PANEL)),
            other=0.0,
        )
        cross = tl.dot(tl.trans(later), earlier)
        base = (batch * CHUNKS + chunk) * 4 * PANEL * PANEL
        tl.store(
            cross_workspace
            + base
            + cross_block * PANEL * PANEL
            + matrix_offsets,
            cross,
        )


@triton.jit
def _n2048_precompose_group_form_weights_kernel(
    source,
    h,
    vectors0,
    vectors1,
    vectors2,
    vectors3,
    group_cross_partials,
    pair_transforms,
    weights,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    CHUNKS: tl.constexpr,
    PIPELINE_STAGES: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
):
    """Project once, then solve through two precomposed 64-wide pairs."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = (
        group_start
        + 4 * PANEL
        + column_tile * BLOCK_N
        + tl.arange(0, BLOCK_N)
    )
    # Keep the wide projection in column-major logical order.  The previous
    # row-major accumulator makes the tensor-core reduction legal, but forces
    # a costly layout transpose before the two 64-wide solves below.  Carrying
    # the transpose through those solves keeps the same products and
    # precision boundaries while shortening the accumulator live range.
    projection_t = tl.zeros((BLOCK_N, 4 * PANEL), tl.float32)
    matrix_base = batch * n * n
    matrix = source if FIRST_GROUP else h
    for row_start in tl.range(
        0, n - group_start, BLOCK_K, num_stages=PIPELINE_STAGES
    ):
        rows = group_start + row_start + tl.arange(0, BLOCK_K)
        v0 = tl.load(
            vectors0 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
        )
        v1 = tl.load(
            vectors1 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=rows[:, None] >= group_start + PANEL,
            other=0.0,
        )
        v2 = tl.load(
            vectors2 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=rows[:, None] >= group_start + 2 * PANEL,
            other=0.0,
        )
        v3 = tl.load(
            vectors3 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=rows[:, None] >= group_start + 3 * PANEL,
            other=0.0,
        )
        values = tl.load(
            matrix + matrix_base + rows[:, None] * n + columns[None, :],
        ).to(tl.float16)
        combined_vectors = tl.cat(
            tl.cat(v0, v1, dim=1), tl.cat(v2, v3, dim=1), dim=1
        )
        projection_t += tl.dot(tl.trans(values), combined_vectors)

    projection_blocks_t = tl.permute(
        tl.reshape(projection_t, (BLOCK_N, 4, PANEL)), (0, 2, 1)
    )
    projection_pairs_t = tl.reshape(
        projection_blocks_t, (BLOCK_N, PANEL, 2, 2)
    )
    even_blocks_t, odd_blocks_t = tl.split(projection_pairs_t)
    p0_t, p2_t = tl.split(even_blocks_t)
    p1_t, p3_t = tl.split(odd_blocks_t)
    first_projection_t = tl.cat(p0_t, p1_t, dim=1)
    second_projection_t = tl.cat(p2_t, p3_t, dim=1)

    wide_offsets = tl.arange(0, 2 * PANEL)
    wide_matrix_offsets = (
        wide_offsets[:, None] * (2 * PANEL) + wide_offsets[None, :]
    )
    first_transform = tl.load(
        pair_transforms
        + (batch * 2) * (2 * PANEL) * (2 * PANEL)
        + wide_matrix_offsets
    )
    second_transform = tl.load(
        pair_transforms
        + (batch * 2 + 1) * (2 * PANEL) * (2 * PANEL)
        + wide_matrix_offsets
    )
    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    cross_base = batch * CHUNKS * 4 * PANEL * PANEL
    c20 = tl.zeros((PANEL, PANEL), tl.float32)
    c21 = tl.zeros((PANEL, PANEL), tl.float32)
    c30 = tl.zeros((PANEL, PANEL), tl.float32)
    c31 = tl.zeros((PANEL, PANEL), tl.float32)
    for chunk in tl.static_range(0, CHUNKS):
        chunk_base = cross_base + chunk * 4 * PANEL * PANEL
        c20 += tl.load(group_cross_partials + chunk_base + matrix_offsets)
        c21 += tl.load(
            group_cross_partials + chunk_base + PANEL * PANEL + matrix_offsets
        )
        c30 += tl.load(
            group_cross_partials + chunk_base + 2 * PANEL * PANEL + matrix_offsets
        )
        c31 += tl.load(
            group_cross_partials + chunk_base + 3 * PANEL * PANEL + matrix_offsets
        )
    pair_cross = tl.cat(
        tl.cat(c20, c21, dim=1), tl.cat(c30, c31, dim=1), dim=0
    )

    first_weights_t = tl.dot(
        first_projection_t.to(tl.float16),
        tl.trans(first_transform.to(tl.float16)),
    )
    second_rhs_t = second_projection_t - tl.dot(
        first_weights_t.to(tl.float16),
        tl.trans(pair_cross.to(tl.float16)),
    )
    second_weights_t = tl.dot(
        second_rhs_t.to(tl.float16),
        tl.trans(second_transform.to(tl.float16)),
    )
    weight_rows = tl.arange(0, 4 * PANEL)[None, :]
    combined_t = tl.cat(first_weights_t, second_weights_t, dim=1)
    tl.store(
        weights
        + batch * (4 * PANEL) * n
        + weight_rows * n
        + columns[:, None],
        combined_t,
    )


@triton.jit
def _n2048_pair_group_form_weights_kernel(
    source,
    h,
    vectors0,
    vectors1,
    vectors2,
    vectors3,
    transform0,
    transform1,
    transform2,
    transform3,
    within_cross0,
    within_cross1,
    group_cross,
    weights,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    PIPELINE_STAGES: tl.constexpr,
    HALF_ACCUMULATION: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
):
    """Project a far tile through four sequential compact-WY blocks."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = group_start + 4 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    if HALF_ACCUMULATION:
        projection = tl.zeros((4 * PANEL, BLOCK_N), tl.float16)
    else:
        projection = tl.zeros((4 * PANEL, BLOCK_N), tl.float32)
    matrix_base = batch * n * n
    matrix = source if FIRST_GROUP else h
    for row_start in tl.range(
        0, n - group_start, BLOCK_K, num_stages=PIPELINE_STAGES
    ):
        rows = group_start + row_start + tl.arange(0, BLOCK_K)
        if n == 4096:
            # Every n4096 grouped launch starts on a 128-column boundary and
            # walks exact 64-row/64-column tiles.  Keep the lower-panel masks,
            # but do not predicate the already in-bounds matrix traffic.
            valid_rows = tl.full((BLOCK_K,), True, tl.int1)
        else:
            valid_rows = rows < n
        v0 = tl.load(
            vectors0 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid_rows[:, None], other=0.0,
        )
        v1 = tl.load(
            vectors1 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= group_start + PANEL),
            other=0.0,
        )
        v2 = tl.load(
            vectors2 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= group_start + 2 * PANEL),
            other=0.0,
        )
        v3 = tl.load(
            vectors3 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid_rows[:, None] & (rows[:, None] >= group_start + 3 * PANEL),
            other=0.0,
        )
        if n == 4096:
            values = tl.load(
                matrix + matrix_base + rows[:, None] * n + columns[None, :]
            ).to(tl.float16)
        else:
            values = tl.load(
                matrix + matrix_base + rows[:, None] * n + columns[None, :],
                mask=valid_rows[:, None] & (columns[None, :] < n),
                other=0.0,
            ).to(tl.float16)
        combined_vectors = tl.cat(
            tl.cat(v0, v1, dim=1), tl.cat(v2, v3, dim=1), dim=1
        )
        if HALF_ACCUMULATION:
            projection += tl.dot(
                tl.trans(combined_vectors), values, out_dtype=tl.float16
            )
        else:
            projection += tl.dot(tl.trans(combined_vectors), values)

    projection_blocks = tl.permute(
        tl.reshape(projection, (4, PANEL, BLOCK_N)), (1, 2, 0)
    )
    projection_pairs = tl.reshape(
        projection_blocks, (PANEL, BLOCK_N, 2, 2)
    )
    even_blocks, odd_blocks = tl.split(projection_pairs)
    p0, p2 = tl.split(even_blocks)
    p1, p3 = tl.split(odd_blocks)

    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    t0 = tl.load(transform0 + batch * PANEL * PANEL + matrix_offsets)
    t1 = tl.load(transform1 + batch * PANEL * PANEL + matrix_offsets)
    t2 = tl.load(transform2 + batch * PANEL * PANEL + matrix_offsets)
    t3 = tl.load(transform3 + batch * PANEL * PANEL + matrix_offsets)
    c10 = tl.load(within_cross0 + batch * PANEL * PANEL + matrix_offsets)
    c32 = tl.load(within_cross1 + batch * PANEL * PANEL + matrix_offsets)
    cross_base = batch * 4 * PANEL * PANEL
    c20 = tl.load(group_cross + cross_base + 0 * PANEL * PANEL + matrix_offsets)
    c21 = tl.load(group_cross + cross_base + 1 * PANEL * PANEL + matrix_offsets)
    c30 = tl.load(group_cross + cross_base + 2 * PANEL * PANEL + matrix_offsets)
    c31 = tl.load(group_cross + cross_base + 3 * PANEL * PANEL + matrix_offsets)

    w0 = tl.dot(tl.trans(t0.to(tl.float16)), p0.to(tl.float16))
    rhs1 = p1 - tl.dot(c10.to(tl.float16), w0.to(tl.float16))
    w1 = tl.dot(tl.trans(t1.to(tl.float16)), rhs1.to(tl.float16))
    rhs2 = p2
    rhs2 -= tl.dot(c20.to(tl.float16), w0.to(tl.float16))
    rhs2 -= tl.dot(c21.to(tl.float16), w1.to(tl.float16))
    w2 = tl.dot(tl.trans(t2.to(tl.float16)), rhs2.to(tl.float16))
    rhs3 = p3
    rhs3 -= tl.dot(c30.to(tl.float16), w0.to(tl.float16))
    rhs3 -= tl.dot(c31.to(tl.float16), w1.to(tl.float16))
    rhs3 -= tl.dot(c32.to(tl.float16), w2.to(tl.float16))
    w3 = tl.dot(tl.trans(t3.to(tl.float16)), rhs3.to(tl.float16))

    weight_rows = tl.arange(0, 4 * PANEL)[:, None]
    combined = tl.cat(tl.cat(w0, w1, dim=0), tl.cat(w2, w3, dim=0), dim=0)
    tl.store(
        weights + batch * (4 * PANEL) * n + weight_rows * n + columns[None, :],
        combined,
        mask=columns[None, :] < n,
    )


@triton.jit
def _n2048_pair_group_apply_kernel(
    source,
    h,
    vectors0,
    vectors1,
    vectors2,
    vectors3,
    weights,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    FIRST_GROUP: tl.constexpr,
):
    """Apply four compact-WY blocks with one far-matrix read/write pass."""
    batch = tl.program_id(0)
    row_tile = tl.program_id(1)
    column_tile = tl.program_id(2)
    rows = group_start + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = group_start + 4 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    offsets = tl.arange(0, PANEL)
    valid = (rows[:, None] < n) & (columns[None, :] < n)
    v0 = tl.load(
        vectors0 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
        mask=rows[:, None] < n, other=0.0,
    )
    w0 = tl.load(
        weights + batch * (4 * PANEL) * n + offsets[:, None] * n + columns[None, :],
        mask=columns[None, :] < n, other=0.0,
    )
    update = tl.dot(v0, w0.to(tl.float16), out_dtype=tl.float16)
    v1 = tl.load(
        vectors1 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
        mask=(rows[:, None] < n) & (rows[:, None] >= group_start + PANEL), other=0.0,
    )
    w1 = tl.load(
        weights + batch * (4 * PANEL) * n + (PANEL + offsets)[:, None] * n + columns[None, :],
        mask=columns[None, :] < n, other=0.0,
    )
    update += tl.dot(v1, w1.to(tl.float16), out_dtype=tl.float16)
    v2 = tl.load(
        vectors2 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
        mask=(rows[:, None] < n) & (rows[:, None] >= group_start + 2 * PANEL), other=0.0,
    )
    w2 = tl.load(
        weights + batch * (4 * PANEL) * n + (2 * PANEL + offsets)[:, None] * n + columns[None, :],
        mask=columns[None, :] < n, other=0.0,
    )
    update += tl.dot(v2, w2.to(tl.float16), out_dtype=tl.float16)
    v3 = tl.load(
        vectors3 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
        mask=(rows[:, None] < n) & (rows[:, None] >= group_start + 3 * PANEL), other=0.0,
    )
    w3 = tl.load(
        weights + batch * (4 * PANEL) * n + (3 * PANEL + offsets)[:, None] * n + columns[None, :],
        mask=columns[None, :] < n, other=0.0,
    )
    update += tl.dot(v3, w3.to(tl.float16), out_dtype=tl.float16)
    pointers = h + batch * n * n + rows[:, None] * n + columns[None, :]
    matrix = source if FIRST_GROUP else h
    values = tl.load(
        matrix + batch * n * n + rows[:, None] * n + columns[None, :],
        mask=valid, other=0.0,
    )
    tl.store(pointers, values - update, mask=valid)


@triton.jit
def _n2048_form_target_weights_kernel(
    h,
    vectors,
    transform,
    weights,
    panel_start,
    target_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """Form one panel's coefficients for an explicit adjacent target."""
    batch = tl.program_id(0)
    offsets = tl.arange(0, PANEL)
    columns = target_start + offsets
    projection = tl.zeros((PANEL, PANEL), tl.float32)
    matrix_base = batch * n * n
    for row_start in tl.range(0, n - panel_start, BLOCK_K):
        rows = panel_start + row_start + tl.arange(0, BLOCK_K)
        panel_vectors = tl.load(
            vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        values = tl.load(
            h + matrix_base + rows[:, None] * n + columns[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        projection += tl.dot(
            tl.trans(panel_vectors), values.to(tl.float16)
        )
    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    t_values = tl.load(transform + batch * PANEL * PANEL + matrix_offsets)
    coefficients = tl.dot(
        tl.trans(t_values), projection, input_precision="tf32x3"
    )
    tl.store(
        weights + batch * PANEL * n + offsets[:, None] * n + columns[None, :],
        coefficients,
    )


@triton.jit
def _n2048_apply_target_kernel(
    h,
    vectors,
    weights,
    panel_start,
    target_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
):
    """Apply one recovered panel to one explicit 16-column target."""
    batch = tl.program_id(0)
    row_tile = tl.program_id(1)
    rows = panel_start + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = target_start + tl.arange(0, PANEL)
    offsets = tl.arange(0, PANEL)
    panel_vectors = tl.load(
        vectors
        + batch * n * PANEL
        + rows[:, None] * PANEL
        + offsets[None, :],
        mask=rows[:, None] < n,
        other=0.0,
    )
    coefficients = tl.load(
        weights
        + batch * PANEL * n
        + offsets[:, None] * n
        + columns[None, :]
    )
    update = tl.dot(panel_vectors, coefficients.to(tl.float16))
    pointers = h + batch * n * n + rows[:, None] * n + columns[None, :]
    mask = rows[:, None] < n
    values = tl.load(pointers, mask=mask, other=0.0)
    tl.store(pointers, values - update, mask=mask)


@triton.jit
def _n2048_quad_cross_kernel(
    later_vectors,
    earlier_vectors,
    cross_workspace,
    group_start,
    later_offset,
    slot,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """Form one of the six lower block products Vj.T @ Vi."""
    batch = tl.program_id(0)
    offsets = tl.arange(0, PANEL)
    cross = tl.zeros((PANEL, PANEL), tl.float32)
    for row_start in tl.range(later_offset, n - group_start, BLOCK_K):
        rows = group_start + row_start + tl.arange(0, BLOCK_K)
        later = tl.load(
            later_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        earlier = tl.load(
            earlier_vectors
            + batch * n * PANEL
            + rows[:, None] * PANEL
            + offsets[None, :],
            mask=rows[:, None] < n,
            other=0.0,
        )
        cross += tl.dot(tl.trans(later), earlier)
    workspace_offsets = offsets[:, None] * PANEL + offsets[None, :]
    tl.store(
        cross_workspace
        + (batch * 6 + slot) * PANEL * PANEL
        + workspace_offsets,
        cross,
    )


@triton.jit
def _n2048_quad_form_weights_kernel(
    h,
    vectors0,
    vectors1,
    vectors2,
    vectors3,
    transform0,
    transform1,
    transform2,
    transform3,
    cross_workspace,
    weights,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    """Form four sequential coefficient blocks for a 64-reflector group."""
    batch = tl.program_id(0)
    column_tile = tl.program_id(1)
    offsets = tl.arange(0, PANEL)
    columns = group_start + 4 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    p0 = tl.zeros((PANEL, BLOCK_N), tl.float32)
    p1 = tl.zeros((PANEL, BLOCK_N), tl.float32)
    p2 = tl.zeros((PANEL, BLOCK_N), tl.float32)
    p3 = tl.zeros((PANEL, BLOCK_N), tl.float32)
    matrix_base = batch * n * n
    for row_start in tl.range(0, n - group_start, BLOCK_K):
        rows = group_start + row_start + tl.arange(0, BLOCK_K)
        valid = rows < n
        v0 = tl.load(
            vectors0 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid[:, None], other=0.0,
        )
        v1 = tl.load(
            vectors1 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid[:, None] & (rows[:, None] >= group_start + PANEL), other=0.0,
        )
        v2 = tl.load(
            vectors2 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid[:, None] & (rows[:, None] >= group_start + 2 * PANEL), other=0.0,
        )
        v3 = tl.load(
            vectors3 + batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :],
            mask=valid[:, None] & (rows[:, None] >= group_start + 3 * PANEL), other=0.0,
        )
        values = tl.load(
            h + matrix_base + rows[:, None] * n + columns[None, :],
            mask=valid[:, None] & (columns[None, :] < n), other=0.0,
        ).to(tl.float16)
        p0 += tl.dot(tl.trans(v0), values)
        p1 += tl.dot(tl.trans(v1), values)
        p2 += tl.dot(tl.trans(v2), values)
        p3 += tl.dot(tl.trans(v3), values)

    matrix_offsets = offsets[:, None] * PANEL + offsets[None, :]
    workspace_base = batch * PANEL * PANEL
    t0 = tl.load(transform0 + workspace_base + matrix_offsets)
    t1 = tl.load(transform1 + workspace_base + matrix_offsets)
    t2 = tl.load(transform2 + workspace_base + matrix_offsets)
    t3 = tl.load(transform3 + workspace_base + matrix_offsets)
    cross_base = batch * 6 * PANEL * PANEL
    c10 = tl.load(cross_workspace + cross_base + 0 * PANEL * PANEL + matrix_offsets)
    c20 = tl.load(cross_workspace + cross_base + 1 * PANEL * PANEL + matrix_offsets)
    c21 = tl.load(cross_workspace + cross_base + 2 * PANEL * PANEL + matrix_offsets)
    c30 = tl.load(cross_workspace + cross_base + 3 * PANEL * PANEL + matrix_offsets)
    c31 = tl.load(cross_workspace + cross_base + 4 * PANEL * PANEL + matrix_offsets)
    c32 = tl.load(cross_workspace + cross_base + 5 * PANEL * PANEL + matrix_offsets)
    w0 = tl.dot(tl.trans(t0), p0, input_precision="tf32x3")
    w1 = tl.dot(
        tl.trans(t1), p1 - tl.dot(c10, w0, input_precision="tf32x3"),
        input_precision="tf32x3",
    )
    rhs2 = p2 - tl.dot(c20, w0, input_precision="tf32x3")
    rhs2 -= tl.dot(c21, w1, input_precision="tf32x3")
    w2 = tl.dot(tl.trans(t2), rhs2, input_precision="tf32x3")
    rhs3 = p3 - tl.dot(c30, w0, input_precision="tf32x3")
    rhs3 -= tl.dot(c31, w1, input_precision="tf32x3")
    rhs3 -= tl.dot(c32, w2, input_precision="tf32x3")
    w3 = tl.dot(tl.trans(t3), rhs3, input_precision="tf32x3")
    combined01 = tl.cat(w0, w1, dim=0)
    combined23 = tl.cat(w2, w3, dim=0)
    combined = tl.cat(combined01, combined23, dim=0)
    weight_rows = tl.arange(0, 4 * PANEL)[:, None]
    tl.store(
        weights + batch * (4 * PANEL) * n + weight_rows * n + columns[None, :],
        combined,
        mask=columns[None, :] < n,
    )


@triton.jit
def _n2048_quad_apply_kernel(
    h,
    vectors0,
    vectors1,
    vectors2,
    vectors3,
    weights,
    group_start,
    n: tl.constexpr,
    PANEL: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    """Apply four 16-column panels with one global matrix round trip."""
    batch = tl.program_id(0)
    row_tile = tl.program_id(1)
    column_tile = tl.program_id(2)
    rows = group_start + row_tile * BLOCK_M + tl.arange(0, BLOCK_M)
    columns = group_start + 4 * PANEL + column_tile * BLOCK_N + tl.arange(0, BLOCK_N)
    offsets = tl.arange(0, PANEL)
    valid_rows = rows < n
    vector_offsets = batch * n * PANEL + rows[:, None] * PANEL + offsets[None, :]
    v0 = tl.load(
        vectors0 + vector_offsets,
        mask=valid_rows[:, None], other=0.0,
    )
    v1 = tl.load(
        vectors1 + vector_offsets,
        mask=valid_rows[:, None] & (rows[:, None] >= group_start + PANEL), other=0.0,
    )
    v2 = tl.load(
        vectors2 + vector_offsets,
        mask=valid_rows[:, None] & (rows[:, None] >= group_start + 2 * PANEL), other=0.0,
    )
    v3 = tl.load(
        vectors3 + vector_offsets,
        mask=valid_rows[:, None] & (rows[:, None] >= group_start + 3 * PANEL), other=0.0,
    )
    weight_base = batch * (4 * PANEL) * n
    w0 = tl.load(weights + weight_base + offsets[:, None] * n + columns[None, :], mask=columns[None, :] < n, other=0.0)
    w1 = tl.load(weights + weight_base + (PANEL + offsets)[:, None] * n + columns[None, :], mask=columns[None, :] < n, other=0.0)
    w2 = tl.load(weights + weight_base + (2 * PANEL + offsets)[:, None] * n + columns[None, :], mask=columns[None, :] < n, other=0.0)
    w3 = tl.load(weights + weight_base + (3 * PANEL + offsets)[:, None] * n + columns[None, :], mask=columns[None, :] < n, other=0.0)
    update = tl.dot(v0, w0.to(tl.float16))
    update += tl.dot(v1, w1.to(tl.float16))
    update += tl.dot(v2, w2.to(tl.float16))
    update += tl.dot(v3, w3.to(tl.float16))
    pointers = h + batch * n * n + rows[:, None] * n + columns[None, :]
    valid = valid_rows[:, None] & (columns[None, :] < n)
    values = tl.load(pointers, mask=valid, other=0.0)
    tl.store(pointers, values - update, mask=valid)


@triton.jit
def _n2048_is_dense(a, batch, n: tl.constexpr, BLOCK: tl.constexpr):
    """Classify one matrix by scale, sparsity, row energy, and correlation."""
    offsets = tl.arange(0, BLOCK)
    first_rows = offsets
    last_rows = n - BLOCK + offsets
    base = batch * n * n

    first_0 = tl.load(a + base + first_rows * n)
    first_1 = tl.load(a + base + first_rows * n + 1)
    first_511 = tl.load(a + base + first_rows * n + 511)
    first_last = tl.load(a + base + first_rows * n + n - 1)
    last_0 = tl.load(a + base + last_rows * n)
    last_1 = tl.load(a + base + last_rows * n + 1)
    last_511 = tl.load(a + base + last_rows * n + 511)
    last_last = tl.load(a + base + last_rows * n + n - 1)

    norm_0 = tl.sum(first_0 * first_0 + last_0 * last_0, axis=0)
    norm_1 = tl.sum(first_1 * first_1 + last_1 * last_1, axis=0)
    norm_511 = tl.sum(first_511 * first_511 + last_511 * last_511, axis=0)
    norm_last = tl.sum(first_last * first_last + last_last * last_last, axis=0)
    dot_01 = tl.sum(first_0 * first_1 + last_0 * last_1, axis=0)
    dot_511_last = tl.sum(
        first_511 * first_last + last_511 * last_last, axis=0
    )
    first_energy = tl.sum(
        first_0 * first_0
        + first_1 * first_1
        + first_511 * first_511
        + first_last * first_last,
        axis=0,
    )
    last_energy = tl.sum(
        last_0 * last_0
        + last_1 * last_1
        + last_511 * last_511
        + last_last * last_last,
        axis=0,
    )
    zero_count = tl.sum(
        (first_0 == 0.0)
        + (first_1 == 0.0)
        + (first_511 == 0.0)
        + (first_last == 0.0)
        + (last_0 == 0.0)
        + (last_1 == 0.0)
        + (last_511 == 0.0)
        + (last_last == 0.0),
        axis=0,
    )
    scale_ok = (norm_last > norm_0 * 1.0e-6) & (norm_last < norm_0 * 0.25)
    rows_ok = last_energy > first_energy * 1.0e-4
    corr_01_ok = dot_01 * dot_01 < norm_0 * norm_1 * 0.0625
    corr_tail_ok = dot_511_last * dot_511_last < norm_511 * norm_last * 0.0625
    return (zero_count == 0) & scale_ok & rows_ok & corr_01_ok & corr_tail_ok


@triton.jit
def _n2048_finalize_kernel(
    h,
    tau,
    fallback_h,
    fallback_tau,
    route_flags,
    n: tl.constexpr,
    NORMALIZE_COLUMNS: tl.constexpr,
    ROW_BLOCK: tl.constexpr,
    COLUMN_BLOCK: tl.constexpr,
    COPY_BLOCK: tl.constexpr,
    WORK_CTAS: tl.constexpr,
):
    """Normalize dense reflectors or restore a routed stable fallback."""
    batch = tl.program_id(0)
    work = tl.program_id(1)
    matrix_base = batch * n * n
    if tl.load(route_flags + batch) != 0:
        column_start = work * COLUMN_BLOCK
        if column_start >= NORMALIZE_COLUMNS:
            return
        # The active n2048 prefix ends on a whole COLUMN_BLOCK, so every
        # surviving work item owns a fully valid column tile.
        columns = column_start + tl.arange(0, COLUMN_BLOCK)[:, None]
        row_offsets = tl.arange(0, ROW_BLOCK)[None, :]
        norm_squared = tl.zeros((COLUMN_BLOCK,), tl.float32)
        # Rows in earlier full blocks are entirely on or above this column
        # tile's diagonal, so they contribute only masked zeros.  Guard them
        # out while retaining the same row order for every nonzero term.
        for row_start in tl.static_range(0, n, ROW_BLOCK):
            if row_start + ROW_BLOCK > column_start:
                rows = row_start + row_offsets
                values = tl.load(
                    h + matrix_base + rows * n + columns,
                    mask=rows > columns,
                    other=0.0,
                )
                norm_squared += tl.sum(values * values, axis=1)
        tau_offsets = batch * n + column_start + tl.arange(0, COLUMN_BLOCK)
        old_tau = tl.load(tau + tau_offsets)
        normalized = tl.where(
            old_tau != 0.0, 2.0 / (1.0 + norm_squared), 0.0
        )
        tl.store(tau + tau_offsets, normalized)
    else:
        elements = n * n
        tiles = tl.cdiv(elements, COPY_BLOCK)
        for tile in tl.range(work, tiles, WORK_CTAS):
            offsets = tile * COPY_BLOCK + tl.arange(0, COPY_BLOCK)
            mask = offsets < elements
            values = tl.load(fallback_h + matrix_base + offsets, mask=mask)
            tl.store(h + matrix_base + offsets, values, mask=mask)
        tau_offsets = work * COLUMN_BLOCK + tl.arange(0, COLUMN_BLOCK)
        tau_mask = tau_offsets < n
        tau_values = tl.load(
            fallback_tau + batch * n + tau_offsets,
            mask=tau_mask,
            other=0.0,
        )
        tl.store(tau + batch * n + tau_offsets, tau_values, mask=tau_mask)


@triton.jit
def _n2048_factor_apply_fallback_panel_kernel(
    source,
    h,
    packed_vectors,
    tau_out,
    gram_out,
    route_flags,
    n: tl.constexpr,
    panel_start,
    BLOCK_M: tl.constexpr,
    BLOCK_B: tl.constexpr,
    BLOCK_K: tl.constexpr,
    BLOCK_N: tl.constexpr,
    PANEL_COUNT: tl.constexpr,
    CLASSIFY_BLOCK: tl.constexpr,
):
    """Factor and apply one fallback panel in a single persistent CTA.

    Fallback matrices deliberately trade parallelism for two fewer global
    launches per panel. Dense matrices return immediately, which is the hot
    CUDA-graph path. The compact-WY arithmetic is the same as in the split
    kernels; this CTA simply traverses their independent tiles sequentially.
    """
    batch = tl.program_id(0)
    fast = _n2048_is_dense(source, batch, n, CLASSIFY_BLOCK)
    tl.store(route_flags + batch, fast.to(tl.int32))
    if fast:
        return

    first_panel_start = panel_start
    matrix_base = batch * n * n
    copy_offsets = tl.arange(0, CLASSIFY_BLOCK)
    for tile in tl.range(0, tl.cdiv(n * n, CLASSIFY_BLOCK)):
        linear = tile * CLASSIFY_BLOCK + copy_offsets
        values = tl.load(source + matrix_base + linear)
        tl.store(h + matrix_base + linear, values)
    tl.debug_barrier()
    for panel_index in tl.range(0, PANEL_COUNT):
        panel_start = first_panel_start + panel_index * BLOCK_B
        rows = tl.arange(0, BLOCK_M)
        cols = tl.arange(0, BLOCK_B)
        global_rows = panel_start + rows
        global_cols = panel_start + cols
        pointers = h + matrix_base + global_rows[:, None] * n + global_cols[None, :]
        valid = (global_rows[:, None] < n) & (global_cols[None, :] < n)
        panel = tl.load(pointers, mask=valid, other=0.0)
    
        for j in tl.static_range(0, BLOCK_B):
            column = tl.sum(tl.where(cols[None, :] == j, panel, 0.0), axis=1)
            tail = (rows >= j) & (global_rows < n)
            norm = tl.sqrt(tl.sum(tl.where(tail, column * column, 0.0), axis=0))
            alpha = tl.sum(tl.where(rows == j, column, 0.0), axis=0)
            beta = tl.where(alpha >= 0.0, -norm, norm)
            nonzero = norm > 0.0
            tau = tl.where(nonzero, (beta - alpha) / beta, 0.0)
            denominator = tl.where(nonzero, alpha - beta, 1.0)
            vector = tl.where(
                rows == j,
                1.0,
                tl.where(
                    (rows > j) & (global_rows < n), column / denominator, 0.0
                ),
            )
            products = tl.sum(vector[:, None] * panel, axis=0)
            tl.store(
                gram_out + batch * BLOCK_B * BLOCK_B + j * BLOCK_B + cols,
                products,
                mask=cols < j,
            )
            panel = tl.where(
                (rows[:, None] >= j) & (cols[None, :] > j) & valid,
                panel - tau * vector[:, None] * products[None, :],
                panel,
            )
            replacement = tl.where(rows[:, None] == j, beta, vector[:, None])
            panel = tl.where(
                (cols[None, :] == j) & (rows[:, None] >= j) & valid,
                replacement,
                panel,
            )
            tl.store(tau_out + batch * n + panel_start + j, tau)
    
        tl.store(pointers, panel, mask=valid)
        packed = tl.where(
            global_rows[:, None] == global_cols[None, :],
            1.0,
            tl.where(global_rows[:, None] > global_cols[None, :], panel, 0.0),
        )
        tl.store(
            packed_vectors
            + batch * n * BLOCK_B
            + global_rows[:, None] * BLOCK_B
            + cols[None, :],
            packed,
            mask=valid,
        )
        tl.debug_barrier()
    
        panel_cols = tl.arange(0, BLOCK_B)
        for col_block in tl.range(0, tl.cdiv(n - panel_start - BLOCK_B, BLOCK_N)):
            out_cols = (
                panel_start
                + BLOCK_B
                + col_block * BLOCK_N
                + tl.arange(0, BLOCK_N)
            )
            dots = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
            for row_start in tl.range(0, n - panel_start, BLOCK_K):
                projection_rows = panel_start + row_start + tl.arange(0, BLOCK_K)
                vectors = tl.load(
                    packed_vectors
                    + batch * n * BLOCK_B
                    + projection_rows[:, None] * BLOCK_B
                    + panel_cols[None, :],
                    mask=projection_rows[:, None] < n,
                    other=0.0,
                )
                values = tl.load(
                    h
                    + matrix_base
                    + projection_rows[:, None] * n
                    + out_cols[None, :],
                    mask=(projection_rows[:, None] < n)
                    & (out_cols[None, :] < n),
                    other=0.0,
                )
                dots += tl.dot(tl.trans(vectors), values.to(tl.float16))
    
            coefficients = tl.zeros((BLOCK_B, BLOCK_N), tl.float32)
            for j in tl.static_range(0, BLOCK_B):
                gram_row = tl.load(
                    gram_out
                    + batch * BLOCK_B * BLOCK_B
                    + j * BLOCK_B
                    + panel_cols,
                    mask=panel_cols < j,
                    other=0.0,
                )
                dot_row = tl.sum(
                    tl.where(panel_cols[:, None] == j, dots, 0.0), axis=0
                )
                correction = tl.sum(
                    tl.where(panel_cols < j, gram_row, 0.0)[:, None]
                    * coefficients,
                    axis=0,
                )
                tau_j = tl.load(tau_out + batch * n + panel_start + j)
                coefficient_row = tau_j * (dot_row - correction)
                coefficients = tl.where(
                    panel_cols[:, None] == j,
                    coefficient_row[None, :],
                    coefficients,
                )
    
            row_blocks: tl.constexpr = BLOCK_M // 32
            for row_block in tl.range(0, row_blocks):
                apply_rows = panel_start + row_block * 32 + tl.arange(0, 32)
                vectors = tl.load(
                    packed_vectors
                    + batch * n * BLOCK_B
                    + apply_rows[:, None] * BLOCK_B
                    + panel_cols[None, :],
                    mask=apply_rows[:, None] < n,
                    other=0.0,
                )
                update = tl.dot(vectors, coefficients.to(tl.float16))
                update_ptrs = (
                    h
                    + matrix_base
                    + apply_rows[:, None] * n
                    + out_cols[None, :]
                )
                update_mask = (apply_rows[:, None] < n) & (out_cols[None, :] < n)
                old = tl.load(update_ptrs, mask=update_mask, other=0.0)
                tl.store(update_ptrs, old - update, mask=update_mask)
        tl.debug_barrier()


def _n2048_householder_tail(
    a,
    h,
    tau,
    packed_vectors,
    gram,
    weights,
    route_flags,
    start: int,
    guarded: bool,
    batch: int,
    n: int,
) -> None:
    """Run the stable 16-column Householder tail, optionally on fallbacks."""
    block_b = 16
    dot_b = 16
    block_n = 128
    if guarded:
        if start == 0:
            # The dense route exits this kernel immediately.  Keeping the
            # complete correctness fallback in one persistent launch removes
            # seven empty hot-path launches while retaining the same scalar
            # Householder/compact-WY arithmetic for rejected matrices.
            _n2048_factor_apply_fallback_panel_kernel[(batch,)](
                a,
                h,
                packed_vectors,
                tau,
                gram,
                route_flags,
                n,
                0,
                BLOCK_M=2048,
                BLOCK_B=block_b,
                BLOCK_K=128,
                BLOCK_N=32,
                PANEL_COUNT=n // block_b,
                CLASSIFY_BLOCK=256,
                num_warps=8,
                num_ctas=1,
            )
            return
        panel_start = start
        direct_start = n - 128
        while panel_start < direct_start:
            panel_rows = triton.next_power_of_2(n - panel_start)
            group_end = min(direct_start, n - panel_rows // 2)
            panel_count = (group_end - panel_start) // block_b
            panel_warps = max(
                4, min(16, (panel_rows * block_b) // 1024)
            )
            _n2048_factor_apply_fallback_panel_kernel[(batch,)](
                a,
                h,
                packed_vectors,
                tau,
                gram,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                BLOCK_K=128,
                BLOCK_N=32,
                PANEL_COUNT=panel_count,
                CLASSIFY_BLOCK=256,
                num_warps=panel_warps,
                num_ctas=1,
            )
            panel_start = group_end
        panel_start = direct_start
        while panel_start < n:
            panel_rows = triton.next_power_of_2(n - panel_start)
            panel_warps = 4 if panel_rows > 64 else 2
            boundary = n - panel_rows // 2
            group_end = min(
                n,
                panel_start
                + triton.cdiv(boundary - panel_start, block_b) * block_b,
            )
            panel_count = (group_end - panel_start) // block_b
            _n2048_factor_apply_fallback_tail_kernel[(batch,)](
                h,
                tau,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                BLOCK_N=32,
                PANEL_COUNT=panel_count,
                num_warps=panel_warps,
                num_ctas=1,
            )
            panel_start = group_end
        return
    for panel_start in range(start, n, block_b):
        panel_rows = triton.next_power_of_2(n - panel_start)
        # The fallback begins just after the stable prefix and pads to 2,048;
        # cap at the portable 16-warp established full-height launch.
        panel_warps = max(4, min(16, (panel_rows * block_b) // 1024))
        if panel_rows <= 64:
            panel_warps = 2
        direct_tail = n - panel_start <= 128
        trailing = n - panel_start - block_b
        if direct_tail:
            _n2048_factor_panel_tail_kernel[(batch,)](
                h,
                tau,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                GUARDED=guarded,
                num_warps=panel_warps,
                num_ctas=1,
            )
        else:
            _n2048_factor_panel_kernel[(batch,)](
                a,
                h,
                packed_vectors,
                tau,
                gram,
                route_flags,
                n,
                panel_start,
                BLOCK_M=panel_rows,
                BLOCK_B=block_b,
                FIRST_PANEL=False,
                GUARDED=guarded,
                num_warps=panel_warps,
                num_ctas=1,
            )
        if trailing <= 0:
            continue
        if direct_tail:
            _n1024_apply_panel_kernel[(batch, triton.cdiv(trailing, 32))](
                h,
                tau,
                route_flags,
                panel_start,
                N=n,
                PANEL=block_b,
                BLOCK_M=panel_rows,
                BLOCK_N=32,
                USE_TAIL_SKIP=guarded,
                SKIP_NONZERO=True,
                num_warps=4,
            )
            continue
        weight_n = 32
        weight_col_blocks = triton.cdiv(trailing, weight_n)
        row_blocks = triton.cdiv(n - panel_start, 32)
        col_blocks = triton.cdiv(trailing, block_n)
        _n2048_make_block_weights_kernel[(batch, weight_col_blocks)](
            a,
            h,
            packed_vectors,
            tau,
            gram,
            weights,
            route_flags,
            n,
            panel_start,
            PANEL_B=block_b,
            BLOCK_B=dot_b,
            BLOCK_K=128,
            BLOCK_N=weight_n,
            FIRST_PANEL=False,
            GUARDED=False,
            num_warps=4,
        )
        _n2048_apply_block_kernel[(batch, row_blocks, col_blocks)](
            a,
            h,
            packed_vectors,
            weights,
            route_flags,
            n,
            panel_start,
            PANEL_B=block_b,
            BLOCK_B=dot_b,
            BLOCK_M=32,
            BLOCK_N=block_n,
            FIRST_PANEL=False,
            GUARDED=False,
            num_warps=4,
        )


def _n2048_qr_v2(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Return compact Householder factors ``(H, tau)`` for square batches."""
    if not a.is_cuda:
        raise ValueError("qr_v2 requires a CUDA tensor")
    if a.dtype != torch.float32:
        raise ValueError("qr_v2 requires torch.float32 input")
    if a.ndim != 3 or a.shape[-2] != a.shape[-1]:
        raise ValueError("qr_v2 requires shape (batch, n, n)")
    if not a.is_contiguous():
        raise ValueError("qr_v2 requires contiguous input")

    batch, n, _ = a.shape
    h = torch.empty_like(a)
    tau = torch.empty((batch, n), device=a.device, dtype=a.dtype)

    # Well-scaled unstructured matrices use paired CholeskyQR from column zero
    # while the panels are tall, then finish with an exact Householder tail.
    # Structured, correlated, mixed, or rank-deficient matrices retain a full
    # stable Householder factorization in ``fallback_h``.  The common structural
    # classifier selects per matrix, so heterogeneous batches remain valid.
    switch_rows = 64
    fallback_start = 0
    chol_panel = 32
    prefix_end = n - switch_rows
    chol_chunks = triton.cdiv(n, 128)
    chol_row_chunk = 128
    chol_pair_chunks = triton.cdiv(n, 256)
    chol_pair_row_chunk = 256
    chol_packed = torch.empty(
        (batch, n, chol_panel), device=a.device, dtype=torch.float16
    )
    chol_second_packed = torch.empty_like(chol_packed)
    chol_weights = torch.empty(
        (batch, chol_panel, n), device=a.device, dtype=torch.float16
    )
    chol_pair_weights = torch.empty(
        (batch, 2 * chol_panel, n), device=a.device, dtype=torch.float16
    )
    chol_group_weights = torch.empty(
        (batch, 4 * chol_panel, n), device=a.device, dtype=torch.float16
    )
    chol_partials = torch.empty(
        (batch, chol_chunks, chol_panel, chol_panel),
        device=a.device,
        dtype=a.dtype,
    )
    chol_pair_partials = torch.empty(
        (batch, chol_pair_chunks, 2 * chol_panel, 2 * chol_panel),
        device=a.device,
        dtype=a.dtype,
    )
    if n == 4096:
        chol_pair_reduced = torch.empty(
            (batch, 1, 2 * chol_panel, 2 * chol_panel),
            device=a.device,
            dtype=a.dtype,
        )
    chol_group_target_partials = torch.empty(
        (batch, chol_pair_chunks, 2, 2 * chol_panel, 2 * chol_panel),
        device=a.device,
        dtype=a.dtype,
    )
    if n == 4096:
        chol_group_target_reduced = torch.empty(
            (batch, 1, 2, 2 * chol_panel, 2 * chol_panel),
            device=a.device,
            dtype=a.dtype,
        )
    # First-panel T factors are only consumed as FP16 MMA operands on n4096.
    # Keep second-panel workspaces FP32 because they temporarily hold the
    # adjacent Schur complement before being overwritten by its T factor.
    chol_t = torch.empty(
        (batch, chol_panel, chol_panel),
        device=a.device,
        dtype=(torch.float16 if n == 4096 else a.dtype),
    )
    chol_bottom_transform = torch.empty(
        (batch, chol_panel, chol_panel), device=a.device, dtype=a.dtype
    )
    chol_second_t = torch.empty_like(chol_bottom_transform)
    chol_second_bottom = torch.empty_like(chol_bottom_transform)
    # Pair crosses have the same n4096 consumption boundary as the four-way
    # crosses below: every active consumer rounds them to FP16 for MMA.
    chol_pair_cross = torch.empty(
        (batch, chol_panel, chol_panel),
        device=a.device,
        dtype=(torch.float16 if n == 4096 else a.dtype),
    )
    chol_inverse = torch.empty_like(chol_bottom_transform)
    chol_second_inverse = torch.empty_like(chol_bottom_transform)
    chol_group_packed = torch.empty_like(chol_packed)
    chol_group_second_packed = torch.empty_like(chol_packed)
    chol_group_t = torch.empty_like(chol_t)
    chol_group_bottom = torch.empty_like(chol_bottom_transform)
    chol_group_second_t = torch.empty_like(chol_bottom_transform)
    chol_group_second_bottom = torch.empty_like(chol_bottom_transform)
    chol_group_pair_cross = torch.empty(
        (batch, chol_panel, chol_panel),
        device=a.device,
        dtype=(torch.float16 if n == 4096 else a.dtype),
    )
    chol_group_inverse = torch.empty_like(chol_bottom_transform)
    chol_group_second_inverse = torch.empty_like(chol_bottom_transform)
    # The n4096 grouped-form consumer immediately rounds these reduced
    # cross tiles to FP16 for every MMA.  Store that same rounded value once
    # at the one-consumer workspace boundary instead of reloading FP32 and
    # converting it in every far-column CTA.
    chol_group_cross = torch.empty(
        (batch, 4, chol_panel, chol_panel),
        device=a.device,
        dtype=(torch.float16 if n == 4096 else a.dtype),
    )
    chol_group_cross_partials = torch.empty(
        (batch, chol_pair_chunks, 4, chol_panel, chol_panel),
        device=a.device,
        dtype=a.dtype,
    )
    if n == 2048:
        chol_pair_transforms = torch.empty(
            (batch, 2, 2 * chol_panel, 2 * chol_panel),
            device=a.device,
            dtype=torch.float16,
        )
    if n == 4096:
        # The accepted n4096 persistent fallback classifies from the immutable
        # source and overwrites the speculative dense result in place. Reuse
        # tau for its dead fast-path flag arguments, as in the accepted route.
        route_flags = tau
    else:
        fallback_h = torch.empty_like(a)
        fallback_tau = torch.empty_like(tau)
        route_flags = torch.empty((batch,), device=a.device, dtype=torch.int32)

    block_b = 16
    dot_b = 16
    block_n = 128
    weights = torch.empty((batch, dot_b, n), device=a.device, dtype=torch.float16)
    packed_vectors = torch.empty((batch, n, dot_b), device=a.device, dtype=torch.float16)
    gram = torch.empty((batch, dot_b, dot_b), device=a.device, dtype=a.dtype)

    # Pair every complete pair and leave at most one 32-column panel for the
    # single-panel path below.  This also permits an even-panel handoff.
    paired_end = prefix_end - (prefix_end % (2 * chol_panel))
    quad_panel = 16
    # Disabled after profiling: the launch-rich four-panel prototype was
    # correct but slower than the two-panel path on GB200.
    for group_start in range(paired_end, paired_end, 4 * quad_panel):
        for subpanel in range(4):
            panel_start = group_start + subpanel * quad_panel
            remaining = n - panel_start
            active_chunks = triton.cdiv(remaining, chol_row_chunk)
            _n2048_chol_gram_kernel[(active_chunks, batch)](
                a,
                h,
                quad_partials,
                panel_start,
                n=n,
                PANEL=quad_panel,
                CHUNKS=chol_chunks,
                ROW_CHUNK=chol_row_chunk,
                FIRST_PANEL=False,
                num_warps=2,
                num_stages=2,
            )
            _n2048_chol_recover_kernel[(batch,)](
                a,
                h,
                tau,
                quad_partials,
                quad_t[subpanel],
                quad_bottom[subpanel],
                quad_packed[subpanel],
                panel_start,
                active_chunks,
                n=n,
                PANEL=quad_panel,
                CHUNKS=chol_chunks,
                INVERSE_STEPS=3,
                FIRST_PANEL=False,
                num_warps=2,
            )
            bottom_tiles = triton.cdiv(remaining - quad_panel, 16)
            _n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
                a,
                h,
                quad_bottom[subpanel],
                quad_packed[subpanel],
                weights,
                panel_start,
                n=n,
                PANEL=quad_panel,
                BLOCK_M=16,
                FIRST_PANEL=False,
                APPLY_TARGET=False,
                num_warps=2,
            )

            if subpanel < 3:
                target_start = panel_start + quad_panel
                for prior in range(subpanel + 1):
                    prior_start = group_start + prior * quad_panel
                    _n2048_form_target_weights_kernel[(batch,)](
                        h,
                        quad_packed[prior],
                        quad_t[prior],
                        weights,
                        prior_start,
                        target_start,
                        n=n,
                        PANEL=quad_panel,
                        BLOCK_K=128,
                        num_warps=2,
                        num_stages=2,
                    )
                    _n2048_apply_target_kernel[
                        (batch, triton.cdiv(n - prior_start, 32))
                    ](
                        h,
                        quad_packed[prior],
                        weights,
                        prior_start,
                        target_start,
                        n=n,
                        PANEL=quad_panel,
                        BLOCK_M=32,
                        num_warps=2,
                    )

        cross_pairs = (
            (1, 0, 0),
            (2, 0, 1),
            (2, 1, 2),
            (3, 0, 3),
            (3, 1, 4),
            (3, 2, 5),
        )
        for later, earlier, slot in cross_pairs:
            _n2048_quad_cross_kernel[(batch,)](
                quad_packed[later],
                quad_packed[earlier],
                quad_cross,
                group_start,
                later * quad_panel,
                slot,
                n=n,
                PANEL=quad_panel,
                BLOCK_K=256,
                num_warps=2,
                num_stages=2,
            )

        trailing = n - group_start - 4 * quad_panel
        quad_weight_tiles = triton.cdiv(trailing, 64)
        _n2048_quad_form_weights_kernel[(batch, quad_weight_tiles)](
            h,
            quad_packed[0],
            quad_packed[1],
            quad_packed[2],
            quad_packed[3],
            quad_t[0],
            quad_t[1],
            quad_t[2],
            quad_t[3],
            quad_cross,
            quad_weights,
            group_start,
            n=n,
            PANEL=quad_panel,
            BLOCK_K=128,
            BLOCK_N=64,
            num_warps=4,
            num_stages=2,
        )
        _n2048_quad_apply_kernel[
            (
                batch,
                triton.cdiv(n - group_start, 64),
                triton.cdiv(trailing, 64),
            )
        ](
            h,
            quad_packed[0],
            quad_packed[1],
            quad_packed[2],
            quad_packed[3],
            quad_weights,
            group_start,
            n=n,
            PANEL=quad_panel,
            BLOCK_M=64,
            BLOCK_N=64,
            num_warps=4,
        )

    grouped_end = (
        min(3072, paired_end - (paired_end % (4 * chol_panel)))
        if n == 4096
        else min(1280, paired_end - (paired_end % (4 * chol_panel)))
    )
    for pair_start in range(0, paired_end, 2 * chol_panel):
        grouped = pair_start < grouped_end
        second_pair = grouped and (pair_start % (4 * chol_panel) != 0)
        if second_pair:
            current_first_vectors = chol_group_packed
            current_second_vectors = chol_group_second_packed
            current_first_t = chol_group_t
            current_first_bottom = chol_group_bottom
            current_second_t = chol_group_second_t
            current_second_bottom = chol_group_second_bottom
            current_pair_cross = chol_group_pair_cross
            current_first_inverse = chol_group_inverse
            current_second_inverse = chol_group_second_inverse
        else:
            current_first_vectors = chol_packed
            current_second_vectors = chol_second_packed
            current_first_t = chol_t
            current_first_bottom = chol_bottom_transform
            current_second_t = chol_second_t
            current_second_bottom = chol_second_bottom
            current_pair_cross = chol_pair_cross
            current_first_inverse = chol_inverse
            current_second_inverse = chol_second_inverse
        remaining = n - pair_start
        active_chunks = triton.cdiv(remaining, chol_pair_row_chunk)
        if grouped and not second_pair:
            group_joint_row_chunk = 256 if n == 4096 else 512
            active_chunks = triton.cdiv(remaining, group_joint_row_chunk)
            _n2048_group_joint_gram_kernel[(active_chunks, 3, batch)](
                a,
                h,
                chol_pair_partials,
                chol_group_target_partials,
                pair_start,
                n=n,
                CHUNKS=active_chunks,
                ROW_CHUNK=group_joint_row_chunk,
                FIRST_GROUP=pair_start == 0,
                num_warps=8,
                num_stages=2,
            )
        elif second_pair:
            # The preceding pair's joint Gram kernel and target conversion
            # already produced this active 64x64 Schur complement.
            active_chunks = 1
        else:
            _n2048_chol_pair_gram_kernel[(active_chunks, batch)](
                a,
                h,
                chol_pair_partials,
                pair_start,
                n=n,
                CHUNKS=active_chunks,
                ROW_CHUNK=chol_pair_row_chunk,
                FIRST_PAIR=pair_start == 0,
                num_warps=4,
                num_stages=2,
            )
        factor_partials = chol_pair_partials
        factor_chunks = active_chunks
        group_target_partials = chol_group_target_partials
        group_target_chunks = active_chunks
        if n == 4096 and grouped and not second_pair:
            _n4096_group_gram_reduce_kernel[(44, batch)](
                chol_pair_partials,
                chol_group_target_partials,
                chol_pair_reduced,
                chol_group_target_reduced,
                active_chunks,
                MAX_CHUNKS=active_chunks,
                num_warps=8,
            )
            factor_partials = chol_pair_reduced
            factor_chunks = 1
            group_target_partials = chol_group_target_reduced
            group_target_chunks = 1
        _n2048_chol32_factor_kernel[(batch,)](
            a,
            h,
            factor_partials,
            current_second_t,
            chol_weights,
            current_first_t,
            current_first_bottom,
            current_first_vectors,
            tau,
            current_first_inverse,
            current_first_vectors,
            current_first_bottom,
            pair_start,
            factor_chunks,
            n=n,
            CHUNKS=factor_chunks,
            FIRST_PANEL=pair_start == 0,
            CHECK_DENSE=False,
            NEWTON_STEPS=(2 if n == 2048 or pair_start < 1024 else 1),
            PAIR_GRAM=True,
            FORM_PAIR_CROSS=False,
            STORE_INVERSE=grouped and not second_pair,
            num_warps=4,
            num_stages=1,
        )
        second_start = pair_start + chol_panel
        second_remaining = remaining - chol_panel
        pair_bottom_tiles = triton.cdiv(second_remaining - chol_panel, 16)
        _n2048_second_factor_extract_first_kernel[
            (batch, pair_bottom_tiles + 1)
        ](
            a,
            h,
            current_second_t,
            current_first_t,
            chol_weights,
            current_second_t,
            current_first_bottom,
            current_second_bottom,
            current_first_vectors,
            current_second_vectors,
            tau,
            current_pair_cross,
            current_second_inverse,
            pair_start,
            n=n,
            PANEL=chol_panel,
            BLOCK_M=16,
            FIRST_PAIR=pair_start == 0,
            STORE_INVERSE=grouped and not second_pair,
            NEWTON_STEPS=(2 if n == 2048 or second_start < 1024 else 1),
            num_warps=4,
            num_stages=1,
        )
        if grouped and not second_pair:
            _n2048_extract_second_group_target_kernel[
                (batch, pair_bottom_tiles + 1)
            ](
                a,
                h,
                group_target_partials,
                chol_pair_partials,
                current_first_inverse,
                current_second_inverse,
                current_first_vectors,
                current_second_vectors,
                current_second_bottom,
                chol_pair_weights,
                pair_start,
                n=n,
                PANEL=chol_panel,
                BLOCK_M=16,
                CHUNKS=group_target_chunks,
                FIRST_GROUP=pair_start == 0,
                QUADRATIC_LOWER=pair_start >= 1152,
                num_warps=8,
                num_stages=1,
            )
        else:
            _n2048_chol_extract_bottom_kernel[(pair_bottom_tiles, batch)](
                a,
                h,
                current_second_bottom,
                current_second_vectors,
                chol_weights,
                second_start,
                n=n,
                PANEL=chol_panel,
                BLOCK_M=16,
                FIRST_PANEL=False,
                APPLY_TARGET=False,
                num_warps=2,
            )

        trailing = remaining - 2 * chol_panel
        if trailing > 0:
            pair_weight_n = 64
            if grouped and not second_pair:
                # Only the adjacent pair is a dependency of its factor. The
                # far matrix remains untouched until all four panels exist.
                _n2048_pair_apply_kernel[
                    (batch, triton.cdiv(remaining, 128), 1)
                ](
                    a, h, current_first_vectors, current_second_vectors,
                    chol_pair_weights, pair_start,
                    n=n, PANEL=chol_panel, BLOCK_M=128,
                    BLOCK_N=pair_weight_n, FIRST_PAIR=pair_start == 0,
                    num_warps=4,
                )
            elif second_pair:
                group_start = pair_start - 2 * chol_panel
                group_cross_max_chunks = triton.cdiv(n, 256) if n == 4096 else 4
                group_cross_row_chunk = 256 if n == 4096 else 512
                group_cross_chunks = triton.cdiv(
                    n - group_start - 2 * chol_panel, group_cross_row_chunk
                )
                if n == 2048:
                    _n2048_precompose_group_cross_kernel[
                        (4 * group_cross_chunks + 2, batch)
                    ](
                        chol_packed, chol_second_packed,
                        chol_group_packed, chol_group_second_packed,
                        chol_group_cross_partials,
                        chol_t, chol_second_t,
                        chol_group_t, chol_group_second_t,
                        chol_pair_cross, chol_group_pair_cross,
                        chol_pair_transforms, group_start,
                        n=n, PANEL=chol_panel,
                        CHUNKS=group_cross_chunks,
                        ROW_CHUNK=group_cross_row_chunk,
                        num_warps=4, num_stages=2,
                    )
                else:
                    _n2048_pair_group_cross_kernel[
                        (group_cross_chunks, 4, batch)
                    ](
                        chol_packed, chol_second_packed,
                        chol_group_packed, chol_group_second_packed,
                        chol_group_cross_partials, group_start,
                        n=n, PANEL=chol_panel, CHUNKS=group_cross_max_chunks,
                        ROW_CHUNK=group_cross_row_chunk,
                        num_warps=8, num_stages=2,
                    )
                    _n2048_pair_group_cross_reduce_kernel[(4, batch)](
                        chol_group_cross_partials, chol_group_cross,
                        group_cross_chunks,
                        PANEL=chol_panel, CHUNKS=group_cross_max_chunks,
                        num_warps=16,
                    )
                # N=64 keeps the 128-row projection below the register cliff
                # and provides enough CTAs for batch two. On n4096 a six-stage
                # K pipeline hides global latency while remaining below the
                # 232448-byte portable shared-memory ceiling. Early panels
                # retain FP32 accumulation for error propagation margin; late
                # projections can accumulate in FP16 because their weights are
                # immediately quantized and have few remaining update steps.
                # On n2048 the group at 768 still has enough trailing work for
                # the 128-wide tile; later groups switch to the 64-wide tile.
                group_weight_n = 64 if n == 4096 else (
                    128 if group_start < 896 else 64
                )
                group_weight_tiles = triton.cdiv(trailing, group_weight_n)
                if n == 2048:
                    _n2048_precompose_group_form_weights_kernel[
                        (batch, group_weight_tiles)
                    ](
                        a, h,
                        chol_packed, chol_second_packed,
                        chol_group_packed, chol_group_second_packed,
                        chol_group_cross_partials,
                        chol_pair_transforms,
                        chol_group_weights, group_start,
                        n=n, PANEL=chol_panel, BLOCK_K=64,
                        BLOCK_N=group_weight_n,
                        CHUNKS=group_cross_chunks,
                        PIPELINE_STAGES=4,
                        FIRST_GROUP=group_start == 0,
                        num_warps=4, num_stages=2,
                    )
                else:
                    _n2048_pair_group_form_weights_kernel[
                        (batch, group_weight_tiles)
                    ](
                        a, h,
                        chol_packed, chol_second_packed,
                        chol_group_packed, chol_group_second_packed,
                        chol_t, chol_second_t,
                        chol_group_t, chol_group_second_t,
                        chol_pair_cross, chol_group_pair_cross,
                        chol_group_cross, chol_group_weights, group_start,
                        n=n, PANEL=chol_panel, BLOCK_K=64,
                        BLOCK_N=group_weight_n,
                        PIPELINE_STAGES=6,
                        HALF_ACCUMULATION=group_start >= 1024,
                        FIRST_GROUP=group_start == 0,
                        num_warps=4, num_stages=2,
                    )
                _n2048_pair_group_apply_kernel[
                    (
                        batch,
                        triton.cdiv(n - group_start, 128),
                        triton.cdiv(trailing, pair_weight_n),
                    )
                ](
                    a, h,
                    chol_packed, chol_second_packed,
                    chol_group_packed, chol_group_second_packed,
                    chol_group_weights, group_start,
                    n=n, PANEL=chol_panel, BLOCK_M=128,
                    BLOCK_N=pair_weight_n, FIRST_GROUP=group_start == 0,
                    num_warps=4,
                )
            else:
                pair_weight_tiles = triton.cdiv(trailing, pair_weight_n)
                _n2048_pair_form_weights_kernel[(batch, pair_weight_tiles)](
                    a, h,
                    current_first_vectors, current_second_vectors,
                    current_first_t, current_second_t, current_pair_cross,
                    chol_pair_weights, pair_start,
                    n=n, PANEL=chol_panel, BLOCK_K=128,
                    BLOCK_N=pair_weight_n, FIRST_PAIR=pair_start == 0,
                    num_warps=4, num_stages=2,
                )
                _n2048_pair_apply_kernel[
                    (
                        batch, triton.cdiv(remaining, 128),
                        pair_weight_tiles,
                    )
                ](
                    a, h, current_first_vectors, current_second_vectors,
                    chol_pair_weights, pair_start,
                    n=n, PANEL=chol_panel, BLOCK_M=128,
                    BLOCK_N=pair_weight_n, FIRST_PAIR=pair_start == 0,
                    num_warps=4,
                )

    # An odd number of Cholesky panels leaves one block for this single-panel
    # path; an even count makes the range empty.
    for panel_start in range(paired_end, prefix_end, chol_panel):
        remaining = n - panel_start
        active_chunks = triton.cdiv(remaining, chol_row_chunk)
        _n2048_chol_gram32_kernel[(active_chunks, 3, batch)](
            a, h, chol_partials, panel_start,
            n=n, CHUNKS=active_chunks, ROW_CHUNK=chol_row_chunk,
            FIRST_PANEL=False, num_warps=1, num_stages=2,
        )
        _n2048_chol32_factor_kernel[(batch,)](
            a, h, chol_partials, chol_second_t, chol_weights,
            chol_t, chol_bottom_transform,
            chol_packed, tau, tau, chol_packed, chol_bottom_transform,
            panel_start, active_chunks,
            n=n, CHUNKS=active_chunks, FIRST_PANEL=False, CHECK_DENSE=False,
            NEWTON_STEPS=2,
            PAIR_GRAM=False,
            FORM_PAIR_CROSS=False,
            STORE_INVERSE=False,
            num_warps=4,
            num_stages=1,
        )
        bottom_tiles = triton.cdiv(remaining - chol_panel, 16)
        _n2048_chol_extract_bottom_kernel[(bottom_tiles, batch)](
            a, h, chol_bottom_transform, chol_packed, chol_weights, panel_start,
            n=n, PANEL=chol_panel, BLOCK_M=16, FIRST_PANEL=False,
            APPLY_TARGET=False, num_warps=2,
        )
        trailing = remaining - chol_panel
        _n2048_chol_form_weights_kernel[(batch, triton.cdiv(trailing, 64))](
            a, h, chol_packed, chol_t, chol_weights, panel_start,
            n=n, PANEL=chol_panel, BLOCK_K=128, BLOCK_N=64,
            FIRST_PANEL=False, num_warps=4, num_stages=2,
        )
        _n2048_apply_block_kernel[
            (batch, triton.cdiv(remaining, 32), triton.cdiv(trailing, 64))
        ](
            a, h, chol_packed, chol_weights, tau, n, panel_start,
            PANEL_B=chol_panel, BLOCK_B=chol_panel, BLOCK_M=32, BLOCK_N=64,
            FIRST_PANEL=False, GUARDED=False, num_warps=4,
        )

    _n2048_householder_tail(
        a,
        h,
        tau,
        packed_vectors,
        gram,
        weights,
        route_flags,
        prefix_end,
        False,
        batch,
        n,
    )
    if n == 4096:
        _s9_n2048_n4096_normalize_tau_tiled_kernel[
            (triton.cdiv(prefix_end, 32), batch)
        ](
            h,
            tau,
            0,
            n=n,
            row_block=1024,
            column_block=32,
            num_warps=8,
        )
        _s9_n2048_n2048_householder_tail(
            a,
            h,
            tau,
            packed_vectors,
            gram,
            weights,
            tau,
            fallback_start,
            True,
            batch,
            n,
        )
    else:
        _n2048_householder_tail(
            a,
            fallback_h,
            fallback_tau,
            packed_vectors,
            gram,
            weights,
            route_flags,
            fallback_start,
            True,
            batch,
            n,
        )

        # Re-impose exact Householder normalization on the dense route, or
        # copy the stable fallback result. The branches are mutually exclusive
        # per matrix, so one routed launch replaces three empty/active launches.
        _n2048_finalize_kernel[(batch, 64)](
            h,
            tau,
            fallback_h,
            fallback_tau,
            route_flags,
            n=n,
            NORMALIZE_COLUMNS=prefix_end,
            ROW_BLOCK=512,
            COLUMN_BLOCK=32,
            COPY_BLOCK=256,
            WORK_CTAS=64,
            num_warps=4,
        )

    return h, tau


def custom_kernel(a):
    n = a.shape[-1]
    if n in (32, 176, 352):
        return _small_qr_v2(a)
    if n == 512:
        return _safe512_qr_v2(a)
    if n == 1024:
        return _s9_n1024_n1024_qr_v2(a)
    if n == 2048:
        return _n2048_qr_v2(a)
    if n == 4096:
        return _n2048_qr_v2(a)
    raise ValueError(f"unsupported QR size {n}")


qr_v2 = custom_kernel
scrolls · 12663 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