Skip to content
KernelIndex
Search⌘K

submission 842146

yeehaw2567 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_b200_constk_n512_tg_currentstack.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-842146?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
2.66ms
#58 of 515
2026-06-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8f8ab33d710f3a3d010c754edbe93de5360e8b1e4535a70c39725e6ee8275e14
license declaredunknown
license concludedunknown
authorsyeehaw2567
imported2026-08-26

Techniques

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

mmagmat += tl.dot(tl.trans(v), v, input_precision="tf32x3")
num-warps = 4_qr32[(batch,)](data, h, tau, stride_ab, stride_am, stride_an, num_warps=4)
tile-k = 128BLOCK_K=128,
tile-m = 32_copy_input[copy_grid](data, h, n, stride_ab, stride_am, stride_an, h_sm, h_sn, BLOCK_M=32, BLOCK_N=32, num_warps=4)
tile-n = 32_copy_input[copy_grid](data, h, n, stride_ab, stride_am, stride_an, h_sm, h_sn, BLOCK_M=32, BLOCK_N=32, num_warps=4)

Kernel source

submission_b200_constk_n512_tg_currentstack.py2011 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl

from task import input_t, output_t


@triton.jit
def _copy_input(
    A,
    H,
    n: tl.constexpr,
    stride_ab: tl.constexpr,
    stride_am: tl.constexpr,
    stride_an: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    offs_m = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    rows = tl.multiple_of(pid_m * BLOCK_M, BLOCK_M) + offs_m
    cols = tl.multiple_of(pid_n * BLOCK_N, BLOCK_N) + offs_n
    mask = (rows[:, None] < n) & (cols[None, :] < n)
    vals = tl.load(
        A + batch * stride_ab + rows[:, None] * stride_am + cols[None, :] * stride_an,
        mask=mask,
        other=0.0,
    ).to(tl.float32)
    tl.store(
        H + batch * n * n + rows[:, None] * h_sm + cols[None, :] * h_sn,
        vals.to(tl.float16) if (n == 2048 or n == 4096) else vals,
        mask=mask,
    )


@triton.jit
def _copy_h_to_float(
    H,
    O,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    o_sm: tl.constexpr,
    o_sn: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    rows = pid_m * BLOCK_M + tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = pid_n * BLOCK_N + tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    mask = (rows[:, None] < n) & (cols[None, :] < n)
    vals = tl.load(
        H + batch * n * n + rows[:, None] * h_sm + cols[None, :] * h_sn,
        mask=mask,
        other=0.0,
    ).to(tl.float32)
    tl.store(O + batch * n * n + rows[:, None] * o_sm + cols[None, :] * o_sn, vals, mask=mask)


@triton.jit
def _qr32(
    A,
    H,
    tau,
    stride_ab: tl.constexpr,
    stride_am: tl.constexpr,
    stride_an: tl.constexpr,
):
    batch = tl.program_id(0)
    offs = tl.arange(0, 32)
    rows = offs[:, None]
    cols = offs[None, :]
    vals = tl.load(A + batch * stride_ab + rows * stride_am + cols * stride_an).to(tl.float32)
    taus = tl.zeros((32,), dtype=tl.float32)

    for j in tl.static_range(0, 32):
        col = tl.sum(tl.where(cols == j, vals, 0.0), axis=1)
        alpha = tl.sum(tl.where(offs == j, col, 0.0), axis=0)
        tail_abs = tl.max(tl.where(offs > j, tl.abs(col), 0.0), axis=0)
        scale = tl.maximum(tl.abs(alpha), tail_abs)
        safe_scale = tl.where(scale > 0.0, scale, 1.0)
        scaled = col / safe_scale
        sumsq = tl.sum(tl.where(offs >= j, scaled * scaled, 0.0), axis=0)
        norm = scale * tl.sqrt(sumsq)
        beta = tl.where(alpha < 0.0, norm, -norm)
        live = tail_abs != 0.0
        tau_j = tl.where(live, (beta - alpha) / tl.where(live, beta, 1.0), 0.0)
        inv = tl.where(live, 1.0 / tl.where(live, alpha - beta, 1.0), 0.0)
        v = tl.where(offs == j, 1.0, tl.where(offs > j, col * inv, 0.0))
        dots = tl.sum(v[:, None] * vals, axis=0)
        updated = vals - tau_j * v[:, None] * dots[None, :]
        vals = tl.where((cols > j) & (rows >= j) & live, updated, vals)
        vals = tl.where((cols == j) & (rows == j) & live, beta, vals)
        vals = tl.where((cols == j) & (rows > j) & live, col[:, None] * inv, vals)
        taus = tl.where(offs == j, tau_j, taus)

    tl.store(H + batch * 32 * 32 + rows * 32 + cols, vals)
    tl.store(tau + batch * 32 + offs, taus)


@triton.jit
def _factor_panel(
    H,
    tau,
    V,
    k,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    PANEL_WIDTH: tl.constexpr,
    PANEL_N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    V_WIDTH: tl.constexpr,
    V_OFFSET: tl.constexpr,
    v_sr: tl.constexpr,
    v_sc: tl.constexpr,
    STORE_V: tl.constexpr,
):
    batch = tl.program_id(0)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = tl.max_contiguous(tl.arange(0, PANEL_WIDTH), PANEL_WIDTH)
    m = n - k
    vals = tl.load(
        H + batch * n * n + (k + rows[:, None]) * h_sm + (k + cols[None, :]) * h_sn,
        mask=(rows[:, None] < m) & (cols[None, :] < PANEL_N),
        other=0.0,
    ).to(tl.float32)
    taus = tl.zeros((PANEL_WIDTH,), dtype=tl.float32)

    for j in tl.static_range(0, PANEL_WIDTH):
        active = j < PANEL_N
        col = tl.sum(tl.where(cols[None, :] == j, vals, 0.0), axis=1)
        alpha = tl.sum(tl.where(rows == j, col, 0.0), axis=0)
        tail_sumsq = tl.sum(tl.where((rows > j) & (rows < m), col * col, 0.0), axis=0)
        sumsq = alpha * alpha + tail_sumsq
        norm = tl.sqrt(sumsq)
        beta = tl.where(alpha < 0.0, norm, -norm)
        live = active & (tail_sumsq != 0.0)
        tau_j = tl.where(live, (beta - alpha) / tl.where(live, beta, 1.0), 0.0)
        inv = tl.where(live, 1.0 / tl.where(live, alpha - beta, 1.0), 0.0)
        v = tl.where(rows == j, 1.0, tl.where((rows > j) & (rows < m), col * inv, 0.0))
        dots = tl.sum(v[:, None] * vals, axis=0)
        updated = vals - tau_j * v[:, None] * dots[None, :]
        vals = tl.where(
            (cols[None, :] > j) & (rows[:, None] >= j) & (rows[:, None] < m) & live,
            updated,
            vals,
        )
        vals = tl.where((cols[None, :] == j) & (rows[:, None] == j) & live, beta, vals)
        vals = tl.where(
            (cols[None, :] == j) & (rows[:, None] > j) & (rows[:, None] < m) & live,
            col[:, None] * inv,
            vals,
        )
        taus = tl.where(cols == j, tau_j, taus)

    packed = tl.where(
        rows[:, None] == cols[None, :],
        1.0,
        tl.where((rows[:, None] > cols[None, :]) & (rows[:, None] < m), vals, 0.0),
    )
    tl.store(
        H + batch * n * n + (k + rows[:, None]) * h_sm + (k + cols[None, :]) * h_sn,
        vals.to(tl.float16) if (n == 2048 or n == 4096) else vals,
        mask=(rows[:, None] < m) & (cols[None, :] < PANEL_N),
    )
    if STORE_V:
        tl.store(
            V + batch * n * V_WIDTH + rows[:, None] * v_sr + (V_OFFSET + cols[None, :]) * v_sc,
            packed,
            mask=(rows[:, None] < m) & (cols[None, :] < PANEL_N),
        )
    tl.store(tau + batch * n + k + cols, taus, mask=cols < PANEL_N)


@triton.jit
def _apply_panel(
    H,
    tau,
    V,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    PANEL_WIDTH: tl.constexpr,
    PANEL_N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    V_WIDTH: tl.constexpr,
    V_OFFSET: tl.constexpr,
    v_sr: tl.constexpr,
    v_sc: tl.constexpr,
    V_FROM_H: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    cols = tl.multiple_of(pid_n * BLOCK_N, BLOCK_N) + offs_n
    global_cols = k + PANEL_N + cols
    m = n - k
    vals = tl.load(
        H + batch * n * n + (k + rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        mask=(rows[:, None] < m) & (cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)

    for j in tl.static_range(0, PANEL_WIDTH):
        tau_j = tl.load(tau + batch * n + k + j, mask=j < PANEL_N, other=0.0).to(tl.float32)
        if V_FROM_H:
            raw = tl.load(
                H + batch * n * n + (k + rows) * h_sm + (k + V_OFFSET + j) * h_sn,
                mask=(rows < m) & (j < PANEL_N),
                other=0.0,
            ).to(tl.float32)
            v = tl.where(rows == V_OFFSET + j, 1.0, tl.where(rows > V_OFFSET + j, raw, 0.0))
        else:
            v = tl.load(
                V + batch * n * V_WIDTH + rows * v_sr + (V_OFFSET + j) * v_sc,
                mask=(rows < m) & (j < PANEL_N),
                other=0.0,
            ).to(tl.float32)
        dots = tl.sum(v[:, None] * vals, axis=0)
        vals = vals - tau_j * v[:, None] * dots[None, :]

    tl.store(
        H + batch * n * n + (k + rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        vals.to(tl.float16) if (n == 2048 or n == 4096) else vals,
        mask=(rows[:, None] < m) & (cols[None, :] < ntrail),
    )


@triton.jit
def _apply_pair(
    H,
    tau,
    V,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    PANEL_WIDTH: tl.constexpr,
    PANEL0_N: tl.constexpr,
    PANEL1_N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    V_WIDTH: tl.constexpr,
    v_sr: tl.constexpr,
    v_sc: tl.constexpr,
    V_FROM_H: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    cols = tl.multiple_of(pid_n * BLOCK_N, BLOCK_N) + offs_n
    far_cols = k + PANEL0_N + PANEL1_N + cols
    m0 = n - k
    m1 = m0 - PANEL0_N
    vals = tl.load(
        H + batch * n * n + (k + rows[:, None]) * h_sm + far_cols[None, :] * h_sn,
        mask=(rows[:, None] < m0) & (cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)

    for j in tl.static_range(0, PANEL_WIDTH):
        tau_j = tl.load(tau + batch * n + k + j, mask=j < PANEL0_N, other=0.0).to(tl.float32)
        if V_FROM_H:
            raw = tl.load(
                H + batch * n * n + (k + rows) * h_sm + (k + j) * h_sn,
                mask=(rows < m0) & (j < PANEL0_N),
                other=0.0,
            ).to(tl.float32)
            v = tl.where(rows == j, 1.0, tl.where(rows > j, raw, 0.0))
        else:
            v = tl.load(
                V + batch * n * V_WIDTH + rows * v_sr + j * v_sc,
                mask=(rows < m0) & (j < PANEL0_N),
                other=0.0,
            ).to(tl.float32)
        dots = tl.sum(v[:, None] * vals, axis=0)
        vals = vals - tau_j * v[:, None] * dots[None, :]

    for j in tl.static_range(0, PANEL_WIDTH):
        row1 = rows - PANEL0_N
        row1_safe = tl.where(rows >= PANEL0_N, row1, 0)
        tau_j = tl.load(tau + batch * n + k + PANEL0_N + j, mask=j < PANEL1_N, other=0.0).to(tl.float32)
        if V_FROM_H:
            raw = tl.load(
                H + batch * n * n + (k + rows) * h_sm + (k + PANEL0_N + j) * h_sn,
                mask=(rows >= PANEL0_N) & (row1 < m1) & (j < PANEL1_N),
                other=0.0,
            ).to(tl.float32)
            v = tl.where(rows == PANEL0_N + j, 1.0, tl.where(rows > PANEL0_N + j, raw, 0.0))
        else:
            stored = tl.load(
                V + batch * n * V_WIDTH + row1_safe * v_sr + (PANEL_WIDTH + j) * v_sc,
                mask=(rows >= PANEL0_N) & (row1 < m1) & (j < PANEL1_N),
                other=0.0,
            ).to(tl.float32)
            v = tl.where(rows >= PANEL0_N, stored, 0.0)
        dots = tl.sum(v[:, None] * vals, axis=0)
        vals = vals - tau_j * v[:, None] * dots[None, :]

    tl.store(
        H + batch * n * n + (k + rows[:, None]) * h_sm + far_cols[None, :] * h_sn,
        vals.to(tl.float16) if (n == 2048 or n == 4096) else vals,
        mask=(rows[:, None] < m0) & (cols[None, :] < ntrail),
    )


@triton.jit
def _build_t32_from_h(
    H,
    tau,
    T,
    k,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.arange(0, BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    tmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for i in tl.static_range(0, 32):
        tau_i = tl.load(tau + batch * n + k + i).to(tl.float32)
        g = tl.zeros((32,), dtype=tl.float32)
        for off in range(0, m, BLOCK_K):
            rel = off + kidx
            raw_i = tl.load(
                H + batch * n * n + (k + rel) * h_sm + (k + i) * h_sn,
                mask=rel < m,
                other=0.0,
            ).to(tl.float32)
            vi = tl.where(rel == i, 1.0, tl.where(rel > i, raw_i, 0.0))
            raw_j = tl.load(
                H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
                mask=rel[:, None] < m,
                other=0.0,
            ).to(tl.float32)
            vj = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw_j, 0.0))
            g += tl.sum(vi[:, None] * vj, axis=0)

        prod = tl.sum(tl.where(tr < i, g[:, None] * tmat, 0.0), axis=0)
        row = -tau_i * prod
        tmat = tl.where((tr == i) & (tc < i), row[None, :], tmat)
        tmat = tl.where((tr == i) & (tc == i), tau_i, tmat)

    tl.store(T + batch * 32 * 32 + tr * 32 + tc, tmat)


@triton.jit
def _build_t32_dot_from_h(
    H,
    tau,
    T,
    k: tl.constexpr,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    gmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for off in range(0, m, BLOCK_K):
        rel = off + kidx
        raw = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw, 0.0))
        gmat += tl.dot(tl.trans(v), v, input_precision="tf32x3")

    tmat = tl.zeros((32, 32), dtype=tl.float32)
    for i in tl.static_range(0, 32):
        tau_i = tl.load(tau + batch * n + k + i).to(tl.float32)
        g = tl.sum(tl.where(tr == i, gmat, 0.0), axis=0)
        prod = tl.sum(tl.where(tr < i, g[:, None] * tmat, 0.0), axis=0)
        row = -tau_i * prod
        tmat = tl.where((tr == i) & (tc < i), row[None, :], tmat)
        tmat = tl.where((tr == i) & (tc == i), tau_i, tmat)

    tl.store(T + batch * 32 * 32 + tr * 32 + tc, tmat)


@triton.jit
def _build_t32_dot_tf32_from_h(
    H,
    tau,
    T,
    k: tl.constexpr,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    gmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for off in range(0, m, BLOCK_K):
        rel = off + kidx
        raw = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw, 0.0))
        gmat += tl.dot(tl.trans(v), v, input_precision="tf32")

    tmat = tl.zeros((32, 32), dtype=tl.float32)
    for i in tl.static_range(0, 32):
        tau_i = tl.load(tau + batch * n + k + i).to(tl.float32)
        g = tl.sum(tl.where(tr == i, gmat, 0.0), axis=0)
        prod = tl.sum(tl.where(tr < i, g[:, None] * tmat, 0.0), axis=0)
        row = -tau_i * prod
        tmat = tl.where((tr == i) & (tc < i), row[None, :], tmat)
        tmat = tl.where((tr == i) & (tc == i), tau_i, tmat)

    tl.store(T + batch * 32 * 32 + tr * 32 + tc, tmat)


@triton.jit
def _build_t32_dot_f16_from_h(
    H,
    tau,
    T,
    k,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    gmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for off in range(0, m, BLOCK_K):
        rel = off + kidx
        raw = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw, 0.0))
        gmat += tl.dot(tl.trans(v.to(tl.float16)), v.to(tl.float16))

    tmat = tl.zeros((32, 32), dtype=tl.float32)
    for i in tl.static_range(0, 32):
        tau_i = tl.load(tau + batch * n + k + i).to(tl.float32)
        g = tl.sum(tl.where(tr == i, gmat, 0.0), axis=0)
        prod = tl.sum(tl.where(tr < i, g[:, None] * tmat, 0.0), axis=0)
        row = -tau_i * prod
        tmat = tl.where((tr == i) & (tc < i), row[None, :], tmat)
        tmat = tl.where((tr == i) & (tc == i), tau_i, tmat)

    tl.store(T + batch * 32 * 32 + tr * 32 + tc, tmat)


@triton.jit
def _wy32_make_y(
    H,
    T,
    Y,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    offs_k = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    offs_r = tl.arange(0, 32)
    cols = pid_n * BLOCK_N + offs_n
    global_cols = k + 32 + cols
    m = n - k
    wt = tl.zeros((BLOCK_N, 32), dtype=tl.float32)

    for off in range(0, m, BLOCK_K):
        rel = off + offs_k
        c_t = tl.load(
            H + batch * n * n + (k + rel[None, :]) * h_sm + global_cols[:, None] * h_sn,
            mask=(rel[None, :] < m) & (cols[:, None] < ntrail),
            other=0.0,
        ).to(tl.float32)
        raw_v = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v = tl.where(rel[:, None] == offs_r[None, :], 1.0, tl.where(rel[:, None] > offs_r[None, :], raw_v, 0.0))
        wt += tl.dot(c_t.to(tl.float16), v.to(tl.float16))

    tt = tl.load(
        T + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None],
    ).to(tl.float32)
    yt = tl.dot(wt, tt, input_precision="tf32")
    tl.store(
        Y + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + offs_n[:, None] * 32 + offs_r[None, :],
        yt.to(tl.float16),
        mask=cols[:, None] < ntrail,
    )


@triton.jit
def _wy32_update(
    H,
    Y,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    rel_rows = pid_m * BLOCK_M + rows
    rel_cols = pid_n * BLOCK_N + cols
    global_cols = k + 32 + rel_cols
    offs_r = tl.arange(0, 32)
    m = n - k

    raw_v = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw_v, 0.0))
    y = tl.load(
        Y + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    delta = tl.dot(v.to(tl.float16), y.to(tl.float16))
    vals = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)
    tl.store(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        (vals - delta).to(tl.float16) if (n == 2048 or n == 4096) else vals - delta,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
    )


@triton.jit
def _build_g32_from_h(
    H,
    G,
    k: tl.constexpr,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.arange(0, BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    gmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for off in range(0, m, BLOCK_K):
        rel = off + kidx
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + ridx[None, :], 1.0, tl.where(rel[:, None] > 32 + ridx[None, :], raw1, 0.0))
        gmat += tl.dot(tl.trans(v0), v1, input_precision="tf32x3")

    tl.store(G + batch * 32 * 32 + tr * 32 + tc, gmat)


@triton.jit
def _build_g32_tf32_from_h(
    H,
    G,
    k: tl.constexpr,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    ridx = tl.arange(0, 32)
    kidx = tl.arange(0, BLOCK_K)
    tr = ridx[:, None]
    tc = ridx[None, :]
    gmat = tl.zeros((32, 32), dtype=tl.float32)
    m = n - k

    for off in range(0, m, BLOCK_K):
        rel = off + kidx
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == ridx[None, :], 1.0, tl.where(rel[:, None] > ridx[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + ridx[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + ridx[None, :], 1.0, tl.where(rel[:, None] > 32 + ridx[None, :], raw1, 0.0))
        gmat += tl.dot(tl.trans(v0), v1, input_precision="tf32")

    tl.store(G + batch * 32 * 32 + tr * 32 + tc, gmat)


@triton.jit
def _wy64_make_y(
    H,
    T0,
    T1,
    G,
    Y0,
    Y1,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    offs_k = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    offs_r = tl.arange(0, 32)
    cols = pid_n * BLOCK_N + offs_n
    global_cols = k + 64 + cols
    m = n - k
    wt0 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)
    wt1 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)

    for off in range(0, m, BLOCK_K):
        rel = off + offs_k
        c_t = tl.load(
            H + batch * n * n + (k + rel[None, :]) * h_sm + global_cols[:, None] * h_sn,
            mask=(rel[None, :] < m) & (cols[:, None] < ntrail),
            other=0.0,
        ).to(tl.float32)
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == offs_r[None, :], 1.0, tl.where(rel[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        wt0 += tl.dot(c_t, v0, input_precision="tf32x3")
        wt1 += tl.dot(c_t, v1, input_precision="tf32x3")

    tt0 = tl.load(T0 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    tt1 = tl.load(T1 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    g = tl.load(G + batch * 32 * 32 + offs_r[:, None] * 32 + offs_r[None, :]).to(tl.float32)
    y0 = tl.dot(wt0, tt0, input_precision="tf32x3")
    wt1_corr = wt1 - tl.dot(y0, g, input_precision="tf32x3")
    y1 = tl.dot(wt1_corr, tt1, input_precision="tf32x3")
    tl.store(
        Y0 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + offs_n[:, None] * 32 + offs_r[None, :],
        y0.to(tl.float16),
        mask=cols[:, None] < ntrail,
    )
    tl.store(
        Y1 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + offs_n[:, None] * 32 + offs_r[None, :],
        y1.to(tl.float16),
        mask=cols[:, None] < ntrail,
    )


@triton.jit
def _wy64_make_y_tf32(
    H,
    T0,
    T1,
    G,
    Y0,
    Y1,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    offs_k = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    offs_r = tl.arange(0, 32)
    cols = pid_n * BLOCK_N + offs_n
    global_cols = k + 64 + cols
    m = n - k
    wt0 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)
    wt1 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)

    for off in range(0, m, BLOCK_K):
        rel = off + offs_k
        c_t = tl.load(
            H + batch * n * n + (k + rel[None, :]) * h_sm + global_cols[:, None] * h_sn,
            mask=(rel[None, :] < m) & (cols[:, None] < ntrail),
            other=0.0,
        ).to(tl.float32)
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == offs_r[None, :], 1.0, tl.where(rel[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        wt0 += tl.dot(c_t, v0, input_precision="tf32")
        wt1 += tl.dot(c_t, v1, input_precision="tf32")

    tt0 = tl.load(T0 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    tt1 = tl.load(T1 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    g = tl.load(G + batch * 32 * 32 + offs_r[:, None] * 32 + offs_r[None, :]).to(tl.float32)
    y0 = tl.dot(wt0, tt0, input_precision="tf32")
    wt1_corr = wt1 - tl.dot(y0, g, input_precision="tf32")
    y1 = tl.dot(wt1_corr, tt1, input_precision="tf32")
    tl.store(
        Y0 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + offs_n[:, None] * 32 + offs_r[None, :],
        y0.to(tl.float16),
        mask=cols[:, None] < ntrail,
    )
    tl.store(
        Y1 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + offs_n[:, None] * 32 + offs_r[None, :],
        y1.to(tl.float16),
        mask=cols[:, None] < ntrail,
    )


@triton.jit
def _wy64_update(
    H,
    Y0,
    Y1,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    rel_rows = pid_m * BLOCK_M + rows
    rel_cols = pid_n * BLOCK_N + cols
    global_cols = k + 64 + rel_cols
    offs_r = tl.arange(0, 32)
    m = n - k

    raw0 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v0 = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw0, 0.0))
    raw1 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v1 = tl.where(rel_rows[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > 32 + offs_r[None, :], raw1, 0.0))
    y0 = tl.load(
        Y0 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    y1 = tl.load(
        Y1 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    delta = tl.dot(v0, y0, input_precision="tf32x3") + tl.dot(v1, y1, input_precision="tf32x3")
    vals = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)
    tl.store(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        vals - delta,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
    )


@triton.jit
def _wy64_update_tf32(
    H,
    Y0,
    Y1,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    rel_rows = pid_m * BLOCK_M + rows
    rel_cols = pid_n * BLOCK_N + cols
    global_cols = k + 64 + rel_cols
    offs_r = tl.arange(0, 32)
    m = n - k

    raw0 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v0 = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw0, 0.0))
    raw1 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v1 = tl.where(rel_rows[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > 32 + offs_r[None, :], raw1, 0.0))
    y0 = tl.load(
        Y0 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    y1 = tl.load(
        Y1 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    delta = tl.dot(v0, y0, input_precision="tf32") + tl.dot(v1, y1, input_precision="tf32")
    vals = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)
    tl.store(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        vals - delta,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
    )


@triton.jit
def _wy64_update_f16(
    H,
    Y0,
    Y1,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    Y_BLOCKS: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    batch = tl.program_id(2)
    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    cols = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    rel_rows = pid_m * BLOCK_M + rows
    rel_cols = pid_n * BLOCK_N + cols
    global_cols = k + 64 + rel_cols
    offs_r = tl.arange(0, 32)
    m = n - k

    raw0 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v0 = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw0, 0.0))
    raw1 = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
        mask=rel_rows[:, None] < m,
        other=0.0,
    ).to(tl.float32)
    v1 = tl.where(rel_rows[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > 32 + offs_r[None, :], raw1, 0.0))
    y0 = tl.load(
        Y0 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    y1 = tl.load(
        Y1 + batch * Y_BLOCKS * BLOCK_N * 32 + pid_n * BLOCK_N * 32 + cols[None, :] * 32 + offs_r[:, None],
        mask=rel_cols[None, :] < ntrail,
        other=0.0,
    ).to(tl.float32)
    delta = tl.dot(v0.to(tl.float16), y0.to(tl.float16)) + tl.dot(v1.to(tl.float16), y1.to(tl.float16))
    vals = tl.load(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
        other=0.0,
    ).to(tl.float32)
    tl.store(
        H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
        vals - delta,
        mask=(rel_rows[:, None] < m) & (rel_cols[None, :] < ntrail),
    )


@triton.jit
def _wy64_fused_apply_f16(
    H,
    T0,
    T1,
    G,
    k,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    offs_k = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    offs_r = tl.arange(0, 32)
    cols = pid_n * BLOCK_N + offs_n
    global_cols = k + 64 + cols
    m = n - k
    wt0 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)
    wt1 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)

    for off in range(0, m, BLOCK_K):
        rel = off + offs_k
        c_t = tl.load(
            H + batch * n * n + (k + rel[None, :]) * h_sm + global_cols[:, None] * h_sn,
            mask=(rel[None, :] < m) & (cols[:, None] < ntrail),
            other=0.0,
        ).to(tl.float32)
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == offs_r[None, :], 1.0, tl.where(rel[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        wt0 += tl.dot(c_t, v0, input_precision="tf32x3")
        wt1 += tl.dot(c_t, v1, input_precision="tf32x3")

    tt0 = tl.load(T0 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    tt1 = tl.load(T1 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    g = tl.load(G + batch * 32 * 32 + offs_r[:, None] * 32 + offs_r[None, :]).to(tl.float32)
    y0 = tl.dot(wt0, tt0, input_precision="tf32x3")
    wt1_corr = wt1 - tl.dot(y0, g, input_precision="tf32x3")
    y1 = tl.dot(wt1_corr, tt1, input_precision="tf32x3")
    y0_t = tl.trans(y0.to(tl.float16))
    y1_t = tl.trans(y1.to(tl.float16))

    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    for row_base in range(0, m, BLOCK_M):
        rel_rows = row_base + rows
        raw0 = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel_rows[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel_rows[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel_rows[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        delta = tl.dot(v0.to(tl.float16), y0_t) + tl.dot(v1.to(tl.float16), y1_t)
        vals = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
            mask=(rel_rows[:, None] < m) & (cols[None, :] < ntrail),
            other=0.0,
        ).to(tl.float32)
        tl.store(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
            vals - delta,
            mask=(rel_rows[:, None] < m) & (cols[None, :] < ntrail),
        )


@triton.jit
def _wy64_fused_apply_f16_constk(
    H,
    T0,
    T1,
    G,
    k: tl.constexpr,
    ntrail,
    n: tl.constexpr,
    h_sm: tl.constexpr,
    h_sn: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    batch = tl.program_id(1)
    offs_n = tl.max_contiguous(tl.arange(0, BLOCK_N), BLOCK_N)
    offs_k = tl.max_contiguous(tl.arange(0, BLOCK_K), BLOCK_K)
    offs_r = tl.arange(0, 32)
    cols = pid_n * BLOCK_N + offs_n
    global_cols = k + 64 + cols
    m = n - k
    wt0 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)
    wt1 = tl.zeros((BLOCK_N, 32), dtype=tl.float32)

    for off in range(0, m, BLOCK_K):
        rel = off + offs_k
        c_t = tl.load(
            H + batch * n * n + (k + rel[None, :]) * h_sm + global_cols[:, None] * h_sn,
            mask=(rel[None, :] < m) & (cols[:, None] < ntrail),
            other=0.0,
        ).to(tl.float32)
        raw0 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel[:, None] == offs_r[None, :], 1.0, tl.where(rel[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        wt0 += tl.dot(c_t, v0, input_precision="tf32x3")
        wt1 += tl.dot(c_t, v1, input_precision="tf32x3")

    tt0 = tl.load(T0 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    tt1 = tl.load(T1 + batch * 32 * 32 + offs_r[None, :] * 32 + offs_r[:, None]).to(tl.float32)
    g = tl.load(G + batch * 32 * 32 + offs_r[:, None] * 32 + offs_r[None, :]).to(tl.float32)
    y0 = tl.dot(wt0, tt0, input_precision="tf32x3")
    wt1_corr = wt1 - tl.dot(y0, g, input_precision="tf32x3")
    y1 = tl.dot(wt1_corr, tt1, input_precision="tf32x3")
    y0_t = tl.trans(y0.to(tl.float16))
    y1_t = tl.trans(y1.to(tl.float16))

    rows = tl.max_contiguous(tl.arange(0, BLOCK_M), BLOCK_M)
    for row_base in range(0, m, BLOCK_M):
        rel_rows = row_base + rows
        raw0 = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + offs_r[None, :]) * h_sn,
            mask=rel_rows[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v0 = tl.where(rel_rows[:, None] == offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > offs_r[None, :], raw0, 0.0))
        raw1 = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + (k + 32 + offs_r[None, :]) * h_sn,
            mask=rel_rows[:, None] < m,
            other=0.0,
        ).to(tl.float32)
        v1 = tl.where(rel_rows[:, None] == 32 + offs_r[None, :], 1.0, tl.where(rel_rows[:, None] > 32 + offs_r[None, :], raw1, 0.0))
        delta = tl.dot(v0.to(tl.float16), y0_t) + tl.dot(v1.to(tl.float16), y1_t)
        vals = tl.load(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
            mask=(rel_rows[:, None] < m) & (cols[None, :] < ntrail),
            other=0.0,
        ).to(tl.float32)
        tl.store(
            H + batch * n * n + (k + rel_rows[:, None]) * h_sm + global_cols[None, :] * h_sn,
            vals - delta,
            mask=(rel_rows[:, None] < m) & (cols[None, :] < ntrail),
        )


def _block_m(m: int) -> int:
    if m <= 64:
        return 64
    if m <= 128:
        return 128
    if m <= 256:
        return 256
    if m <= 512:
        return 512
    if m <= 1024:
        return 1024
    if m <= 2048:
        return 2048
    return 4096


def _panel_width(n: int) -> int:
    if n == 2048:
        return 8
    return 16 if n <= 2048 else 8


def _block_n(n: int) -> int:
    if n <= 512:
        return 16
    return 8


def _factor_warps(n: int) -> int:
    if n >= 4096:
        return 8
    return 4 if n <= 1024 else 8


def _apply_warps(n: int) -> int:
    if n >= 4096:
        return 8
    return 4 if n <= 512 else 8


def _custom_kernel_impl(data: input_t, h=None, input_ready: bool = False, out_h=None) -> output_t:
    if data.ndim != 3 or data.shape[-1] != data.shape[-2]:
        raise RuntimeError("expected a batch of square matrices")

    batch = data.shape[0]
    n = data.shape[-1]
    h_dtype = torch.float16 if n == 2048 or n == 4096 else torch.float32
    if h is None:
        if n >= 352:
            h = torch.empty_strided((batch, n, n), (n * n, 1, n), device=data.device, dtype=h_dtype)
        else:
            h = torch.empty((batch, n, n), device=data.device, dtype=h_dtype)
    else:
        h_dtype = h.dtype
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    if batch == 0:
        return h.to(torch.float32), tau

    stride_ab, stride_am, stride_an = data.stride()
    if n == 32:
        _qr32[(batch,)](data, h, tau, stride_ab, stride_am, stride_an, num_warps=4)
        return h, tau

    panel_width = _panel_width(n)
    block_n = _block_n(n)
    h_sm, h_sn = h.stride()[1:]
    v_width = panel_width * 2
    use_wy64 = n == 512 or n == 1024
    use_wy32x4 = n == 2048 or n == 4096
    use_h_reflectors = n == 352 or n == 512 or use_wy64 or n >= 4096
    use_wy32 = (n == 512 or n == 1024 or n == 2048) and not use_wy64
    if use_h_reflectors:
        v_panel = h
        v_sr = 1
        v_sc = 1
    elif n >= 512:
        v_panel = torch.empty((batch, v_width, n), device=data.device, dtype=torch.float32)
        v_sr = 1
        v_sc = n
    else:
        v_panel = torch.empty((batch, n, v_width), device=data.device, dtype=torch.float32)
        v_sr = v_width
        v_sc = 1
    if use_wy64:
        t_panel = torch.empty((batch, 32, 32), device=data.device, dtype=torch.float32)
        t_panel1 = torch.empty((batch, 32, 32), device=data.device, dtype=torch.float32)
        g_panel = torch.empty((batch, 32, 32), device=data.device, dtype=torch.float32)
        y_blocks = (n + 63) // 64
        y_panel = torch.empty((batch, y_blocks, 64, 32), device=data.device, dtype=torch.float16)
        y_panel1 = torch.empty((batch, y_blocks, 64, 32), device=data.device, dtype=torch.float16)
    elif use_wy32 or use_wy32x4:
        t_panel = torch.empty((batch, 32, 32), device=data.device, dtype=torch.float32)
        t_panel1 = h
        g_panel = h
        y_blocks = (n + 63) // 64
        y_panel = torch.empty((batch, y_blocks, 64, 32), device=data.device, dtype=torch.float16)
        y_panel1 = h
    else:
        t_panel = h
        t_panel1 = h
        g_panel = h
        y_blocks = 1
        y_panel = h
        y_panel1 = h
    if not input_ready:
        copy_grid = (triton.cdiv(n, 32), triton.cdiv(n, 32), batch)
        _copy_input[copy_grid](data, h, n, stride_ab, stride_am, stride_an, h_sm, h_sn, BLOCK_M=32, BLOCK_N=32, num_warps=4)

    step = panel_width * (4 if use_wy64 or use_wy32x4 else 2)
    for k in range(0, n, step):
        if use_wy32x4 and panel_width == 8 and n - k >= 32:
            block_m0 = _block_m(n - k)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m0,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(24, block_n), batch)](
                h,
                tau,
                v_panel,
                k,
                24,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m0,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k1 = k + 8
            block_m1 = _block_m(n - k1)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k1,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m1,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(16, block_n), batch)](
                h,
                tau,
                v_panel,
                k1,
                16,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m1,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k2 = k + 16
            block_m2 = _block_m(n - k2)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k2,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m2,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(8, block_n), batch)](
                h,
                tau,
                v_panel,
                k2,
                8,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m2,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k3 = k + 24
            block_m3 = _block_m(n - k3)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k3,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=8,
                BLOCK_M=block_m3,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )

            ntrail32 = n - k - 32
            if ntrail32:
                _build_t32_dot_f16_from_h[(batch,)](
                    h,
                    tau,
                    t_panel,
                    k,
                    n,
                    h_sm,
                    h_sn,
                    BLOCK_K=128,
                    num_warps=8,
                )
                _wy32_make_y[(triton.cdiv(ntrail32, 64), batch)](
                    h,
                    t_panel,
                    y_panel,
                    k,
                    ntrail32,
                    n,
                    h_sm,
                    h_sn,
                    Y_BLOCKS=y_blocks,
                    BLOCK_N=64,
                    BLOCK_K=128,
                    num_warps=8,
                )
                _wy32_update[(triton.cdiv(n - k, 128), triton.cdiv(ntrail32, 64), batch)](
                    h,
                    y_panel,
                    k,
                    ntrail32,
                    n,
                    h_sm,
                    h_sn,
                    Y_BLOCKS=y_blocks,
                    BLOCK_M=128,
                    BLOCK_N=64,
                    num_warps=4,
                )
            continue

        if use_wy64 and panel_width == 16 and n - k >= 64:
            block_m0 = _block_m(n - k)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m0,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(48, block_n), batch)](
                h,
                tau,
                v_panel,
                k,
                48,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m0,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k1 = k + 16
            block_m1 = _block_m(n - k1)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k1,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m1,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(32, block_n), batch)](
                h,
                tau,
                v_panel,
                k1,
                32,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m1,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k2 = k + 32
            block_m2 = _block_m(n - k2)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k2,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m2,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )
            _apply_panel[(triton.cdiv(16, block_n), batch)](
                h,
                tau,
                v_panel,
                k2,
                16,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m2,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=True,
                num_warps=_apply_warps(n),
            )
            k3 = k + 48
            block_m3 = _block_m(n - k3)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k3,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=16,
                BLOCK_M=block_m3,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=False,
                num_warps=_factor_warps(n),
            )

            ntrail64 = n - k - 64
            if ntrail64:
                if n == 1024 or n == 2048 or (n == 512 and k >= 256):
                    _build_t32_dot_tf32_from_h[(batch,)](
                        h,
                        tau,
                        t_panel,
                        k,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _build_t32_dot_tf32_from_h[(batch,)](
                        h,
                        tau,
                        t_panel1,
                        k + 32,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _build_g32_tf32_from_h[(batch,)](
                        h,
                        g_panel,
                        k,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                else:
                    _build_t32_dot_from_h[(batch,)](
                        h,
                        tau,
                        t_panel,
                        k,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _build_t32_dot_from_h[(batch,)](
                        h,
                        tau,
                        t_panel1,
                        k + 32,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _build_g32_from_h[(batch,)](
                        h,
                        g_panel,
                        k,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                if n == 512:
                    _wy64_fused_apply_f16_constk[(triton.cdiv(ntrail64, 64), batch)](
                        h,
                        t_panel,
                        t_panel1,
                        g_panel,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_M=128,
                        BLOCK_N=64,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                elif n == 1024:
                    _wy64_fused_apply_f16[(triton.cdiv(ntrail64, 64), batch)](
                        h,
                        t_panel,
                        t_panel1,
                        g_panel,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_M=128,
                        BLOCK_N=64,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                elif n == 2048 or (n == 1024 and k >= 640):
                    _wy64_make_y_tf32[(triton.cdiv(ntrail64, 64), batch)](
                        h,
                        t_panel,
                        t_panel1,
                        g_panel,
                        y_panel,
                        y_panel1,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_N=64,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                else:
                    _wy64_make_y[(triton.cdiv(ntrail64, 64), batch)](
                        h,
                        t_panel,
                        t_panel1,
                        g_panel,
                        y_panel,
                        y_panel1,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_N=64,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                if n == 512 or n == 1024:
                    pass
                elif n == 2048:
                    _wy64_update_f16[(triton.cdiv(n - k, 128), triton.cdiv(ntrail64, 64), batch)](
                        h,
                        y_panel,
                        y_panel1,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_M=128,
                        BLOCK_N=64,
                        num_warps=4,
                    )
                else:
                    _wy64_update[(triton.cdiv(n - k, 128), triton.cdiv(ntrail64, 64), batch)](
                        h,
                        y_panel,
                        y_panel1,
                        k,
                        ntrail64,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_M=128,
                        BLOCK_N=64,
                        num_warps=4,
                    )
            continue

        panel_n = min(panel_width, n - k)
        block_m = _block_m(n - k)
        _factor_panel[(batch,)](
            h,
            tau,
            v_panel,
            k,
            n,
            h_sm,
            h_sn,
            PANEL_WIDTH=panel_width,
            PANEL_N=panel_n,
            BLOCK_M=block_m,
            V_WIDTH=v_width,
            V_OFFSET=0,
            v_sr=v_sr,
            v_sc=v_sc,
            STORE_V=not use_h_reflectors,
            num_warps=_factor_warps(n),
        )
        panel1_n = min(panel_width, n - k - panel_n)
        if panel1_n:
            _apply_panel[(triton.cdiv(panel1_n, block_n), batch)](
                h,
                tau,
                v_panel,
                k,
                panel1_n,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=panel_n,
                BLOCK_M=block_m,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=use_h_reflectors,
                num_warps=_apply_warps(n),
            )
            k1 = k + panel_n
            block_m1 = _block_m(n - k1)
            _factor_panel[(batch,)](
                h,
                tau,
                v_panel,
                k1,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=panel1_n,
                BLOCK_M=block_m1,
                V_WIDTH=v_width,
                V_OFFSET=panel_width,
                v_sr=v_sr,
                v_sc=v_sc,
                STORE_V=not use_h_reflectors,
                num_warps=_factor_warps(n),
            )
            ntrail = n - k - panel_n - panel1_n
            if ntrail:
                if use_wy32 and panel_n == 16 and panel1_n == 16:
                    _build_t32_dot_from_h[(batch,)](
                        h,
                        tau,
                        t_panel,
                        k,
                        n,
                        h_sm,
                        h_sn,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _wy32_make_y[(triton.cdiv(ntrail, 64), batch)](
                        h,
                        t_panel,
                        y_panel,
                        k,
                        ntrail,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_N=64,
                        BLOCK_K=64,
                        num_warps=8,
                    )
                    _wy32_update[(triton.cdiv(n - k, 128), triton.cdiv(ntrail, 64), batch)](
                        h,
                        y_panel,
                        k,
                        ntrail,
                        n,
                        h_sm,
                        h_sn,
                        Y_BLOCKS=y_blocks,
                        BLOCK_M=128,
                        BLOCK_N=64,
                        num_warps=4,
                    )
                else:
                    _apply_pair[(triton.cdiv(ntrail, block_n), batch)](
                        h,
                        tau,
                        v_panel,
                        k,
                        ntrail,
                        n,
                        h_sm,
                        h_sn,
                        PANEL_WIDTH=panel_width,
                        PANEL0_N=panel_n,
                        PANEL1_N=panel1_n,
                        BLOCK_M=block_m,
                        BLOCK_N=block_n,
                        V_WIDTH=v_width,
                        v_sr=v_sr,
                        v_sc=v_sc,
                        V_FROM_H=use_h_reflectors,
                        num_warps=_apply_warps(n),
                    )
            continue

        ntrail = n - k - panel_n
        if ntrail:
            _apply_panel[(triton.cdiv(ntrail, block_n), batch)](
                h,
                tau,
                v_panel,
                k,
                ntrail,
                n,
                h_sm,
                h_sn,
                PANEL_WIDTH=panel_width,
                PANEL_N=panel_n,
                BLOCK_M=block_m,
                BLOCK_N=block_n,
                V_WIDTH=v_width,
                V_OFFSET=0,
                v_sr=v_sr,
                v_sc=v_sc,
                V_FROM_H=use_h_reflectors,
                num_warps=_apply_warps(n),
            )

    if h_dtype is torch.float16:
        if out_h is None:
            out_h = torch.empty_strided((batch, n, n), (n * n, 1, n), device=data.device, dtype=torch.float32)
        _copy_h_to_float[(triton.cdiv(n, 64), triton.cdiv(n, 32), batch)](
            h,
            out_h,
            n,
            h_sm,
            h_sn,
            out_h.stride(1),
            out_h.stride(2),
            BLOCK_M=64,
            BLOCK_N=32,
            num_warps=4,
        )
        return out_h, tau
    if out_h is not None:
        _copy_h_to_float[(triton.cdiv(n, 64), triton.cdiv(n, 32), batch)](
            h,
            out_h,
            n,
            h_sm,
            h_sn,
            out_h.stride(1),
            out_h.stride(2),
            BLOCK_M=64,
            BLOCK_N=32,
            num_warps=4,
        )
        return out_h, tau
    return h, tau


_GRAPH_CACHE = {}
_STATIC_GRAPH_CACHE = {}
_GRAPH_DISABLED = False
# ponytail: two output buffers avoid immediate aliasing; add more only if benchmark proves it.
_STATIC_GRAPH_RING = 2


def custom_kernel(data: input_t) -> output_t:
    global _GRAPH_DISABLED
    if (
        _GRAPH_DISABLED
        or not data.is_cuda
        or data.ndim != 3
        or data.shape[-1] != data.shape[-2]
        or data.shape[-1] not in (176, 352, 512, 1024, 2048, 4096)
    ):
        return _custom_kernel_impl(data)

    if data.shape[-1] in (176, 352, 2048, 4096):
        key = (data.device.index, data.dtype, tuple(data.shape), tuple(data.stride()))
        cached = _STATIC_GRAPH_CACHE.get(key)
        if cached is not None:
            idx, entries = cached
            if len(entries) >= _STATIC_GRAPH_RING:
                graph, static_h, h, tau = entries[idx]
                _STATIC_GRAPH_CACHE[key] = ((idx + 1) % len(entries), entries)
                static_h.copy_(data)
                graph.replay()
                if data.shape[-1] == 176 or data.shape[-1] == 352:
                    return h.clone(), tau.clone()
                return h, tau
        else:
            entries = []

        try:
            n = data.shape[-1]
            h_dtype = torch.float16 if n == 2048 or n == 4096 else torch.float32
            h_stride = (n * n, 1, n) if n >= 352 else (n * n, n, 1)
            static_h = torch.empty_strided(tuple(data.shape), h_stride, device=data.device, dtype=h_dtype)
            out_h = None
            static_h.copy_(data)
            graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(graph):
                h, tau = _custom_kernel_impl(static_h, static_h, True, out_h)
            static_h.copy_(data)
            entries.append((graph, static_h, h, tau))
            _STATIC_GRAPH_CACHE[key] = (len(entries) % _STATIC_GRAPH_RING, entries)
            graph.replay()
            if n == 176 or n == 352:
                return h.clone(), tau.clone()
            return h, tau
        except Exception:
            return _custom_kernel_impl(data)

    key = (data.device.index, data.dtype, tuple(data.shape), tuple(data.stride()), data.data_ptr())
    cached = _GRAPH_CACHE.get(key)
    if cached is not None:
        graph, h, tau = cached
        graph.replay()
        return h, tau

    try:
        _custom_kernel_impl(data)
        torch.cuda.synchronize(data.device)
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            h, tau = _custom_kernel_impl(data)
        _GRAPH_CACHE[key] = (graph, h, tau)
        graph.replay()
        torch.cuda.synchronize(data.device)
        return h, tau
    except Exception:
        _GRAPH_DISABLED = True
        return _custom_kernel_impl(data)
scrolls · 2011 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