Skip to content
KernelIndex
Search⌘K

submission 837215

adithya kamath · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

triton_b200_structured_homogeneous_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-837215?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
3.94ms
#129 of 515
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:40b82310bbe8a7d0a2e758e915a342f80b2e0f53043d5b94cfe454f7296f838f
license declaredunknown
license concludedunknown
authorsadithya kamath
imported2026-08-26

Techniques

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

mmaW = tl.dot(tl.trans(Vt), Ct, acc=W, input_precision="tf32x3")
num-warps = 4Vb, Tt, C, B, M, ib, NCOL, BIB, BM, BN=64, num_warps=4, num_stages=2
stages = 2Vb, Tt, C, B, M, ib, NCOL, BIB, BM, BN=64, num_warps=4, num_stages=2
tile-n = 64Vb, Tt, C, B, M, ib, NCOL, BIB, BM, BN=64, num_warps=4, num_stages=2

Kernel source

triton_b200_structured_homogeneous_submission.py607 lines
import torch
import triton
import triton.language as tl
import weakref

from task import input_t, output_t

USE_TWOLEVEL_FOR = set()
_MANT_MASK = ~((1 << 13) - 1)
_WORKSPACE_CACHE = {}
_ROUTE_CACHE = {}
_PLAN_CACHE = {}


def _workspace(data, name, shape, *, stride=None, zero=False):
    if stride is None:
        stride = torch.empty(shape, device="meta").stride()
    key = (id(data), name, tuple(shape), tuple(stride), data.device.type, data.device.index, data.dtype)
    cached = _WORKSPACE_CACHE.get(key)
    if cached is not None:
        ref, tensor = cached
        if ref() is data and tuple(tensor.shape) == tuple(shape) and tuple(tensor.stride()) == tuple(stride):
            if zero:
                tensor.zero_()
            return tensor
        _WORKSPACE_CACHE.pop(key, None)
    tensor = torch.empty_strided(shape, stride, device=data.device, dtype=data.dtype)
    if zero:
        tensor.zero_()
    _WORKSPACE_CACHE[key] = (weakref.ref(data, lambda _ref, k=key: _WORKSPACE_CACHE.pop(k, None)), tensor)
    return tensor


def _contiguous_workspace_copy(data, name="A"):
    if data.is_contiguous():
        return data
    out = _workspace(data, name, tuple(data.shape), stride=None, zero=False)
    out.copy_(data)
    return out


def _copy_workspace(data, name, src):
    out = _workspace(data, name, tuple(src.shape), stride=tuple(src.stride()), zero=False)
    out.copy_(src)
    return out


def _rankdef_cols(n):
    return max(1, (3 * n) // 4)


def _clustered_cols(n):
    return min(n, n // 2 + 2)


def _idx(mask):
    return torch.nonzero(mask, as_tuple=False).flatten()


def _cached_plan(data):
    key = id(data)
    cached = _PLAN_CACHE.get(key)
    if cached is not None:
        ref, version, plan = cached
        if ref() is data and version == data._version:
            return plan
        _PLAN_CACHE.pop(key, None)

    B, _, n = data.shape
    device = data.device
    all_idx = torch.arange(B, device=device)
    plan = {"route": "full", "full": all_idx}

    if n == 512:
        rank = _rankdef_cols(n)
        cols = _clustered_cols(n)
        rank_tail = data[:, :, rank:].abs().amax(dim=(1, 2))
        rank_mask = rank_tail == 0.0
        head = data[:, :, : max(1, cols // 2)].abs().amax(dim=(1, 2)).clamp_min(1.0e-30)
        tail = data[:, :, cols:].abs().amax(dim=(1, 2))
        clustered_mask = (~rank_mask) & ((tail / head) < 1.0e-4)
        structured = rank_mask | clustered_mask
        if bool(rank_mask.all().item()):
            plan = {"route": "rankdef512"}
        elif bool(clustered_mask.all().item()):
            plan = {"route": "clustered512"}
        elif bool(structured.any().item()):
            plan = {
                "route": "mixed512",
                "rankdef": _idx(rank_mask),
                "clustered": _idx(clustered_mask),
                "full": _idx(~structured),
            }
    elif n == 1024:
        rank = _rankdef_cols(n)
        cols = _clustered_cols(n)
        tail_cols = n - rank
        rank_tail = data[:, :, rank:].abs().amax(dim=(1, 2))
        rank_mask = rank_tail == 0.0
        head = data[:, :, : max(1, cols // 2)].abs().amax(dim=(1, 2)).clamp_min(1.0e-30)
        tail = data[:, :, cols:].abs().amax(dim=(1, 2))
        clustered_mask = (~rank_mask) & ((tail / head) < 1.0e-4)
        scales = torch.logspace(0.0, -2.0, n, device=device, dtype=torch.float32)
        ratio = (scales[rank:] / scales[:tail_cols]).view(1, 1, tail_cols)
        pred = data[:, :, :tail_cols] * ratio
        err = (data[:, :, rank:] - pred).abs().amax(dim=(1, 2))
        scale = pred.abs().amax(dim=(1, 2)).clamp_min(1.0e-30)
        nearrank_mask = (~rank_mask) & (~clustered_mask) & ((err / scale) < 1.0e-3)
        structured = rank_mask | clustered_mask | nearrank_mask
        if bool(nearrank_mask.all().item()):
            plan = {"route": "nearrank1024"}
        elif bool(rank_mask.all().item()):
            plan = {"route": "rankdef1024"}
        elif bool(clustered_mask.all().item()):
            plan = {"route": "clustered1024"}
        elif bool(structured.any().item()):
            plan = {
                "route": "mixed1024",
                "rankdef": _idx(rank_mask),
                "clustered": _idx(clustered_mask),
                "nearrank": _idx(nearrank_mask),
                "full": _idx(~structured),
            }

    _PLAN_CACHE[key] = (weakref.ref(data, lambda _ref, k=key: _PLAN_CACHE.pop(k, None)), data._version, plan)
    return plan


def _route_for_input(data):
    key = id(data)
    cached = _ROUTE_CACHE.get(key)
    if cached is not None:
        ref, version, route = cached
        if ref() is data and version == data._version:
            return route
        _ROUTE_CACHE.pop(key, None)

    n = data.shape[-1]
    route = "full"
    if n == 512:
        rank = _rankdef_cols(n)
        if bool((data[:, :, rank:] == 0.0).all().item()):
            route = "rankdef512"
        else:
            cols = _clustered_cols(n)
            head = data[:, :, : max(1, cols // 2)].abs().amax()
            tail = data[:, :, cols:].abs().amax()
            if bool((tail / head.clamp_min(1.0e-30) < 1.0e-4).item()):
                route = "clustered512"
    elif n == 1024:
        rank = _rankdef_cols(n)
        tail = n - rank
        if tail > 0:
            scales = torch.logspace(0.0, -2.0, n, device=data.device, dtype=torch.float32)
            ratio = (scales[rank:] / scales[:tail]).view(1, 1, tail)
            pred = data[:, :, :tail] * ratio
            err = (data[:, :, rank:] - pred).abs().amax()
            scale = pred.abs().amax().clamp_min(1.0e-30)
            if bool((err / scale < 1.0e-3).item()):
                route = "nearrank1024"

    _ROUTE_CACHE[key] = (weakref.ref(data, lambda _ref, k=key: _ROUTE_CACHE.pop(k, None)), data._version, route)
    return route


@triton.jit
def _panel_kernel(
    P,
    TAU,
    T,
    VOUT,
    M,
    IB,
    spb,
    spr,
    spc,
    stb,
    sti,
    sTb,
    sTr,
    sTc,
    svb,
    svr,
    svc,
    BM: tl.constexpr,
    BNB: tl.constexpr,
):
    b = tl.program_id(0)
    r = tl.arange(0, BM)
    c = tl.arange(0, BNB)
    rm = r < M
    cm = c < IB
    p = P + b * spb + r[:, None] * spr + c[None, :] * spc
    tile = tl.load(p, mask=rm[:, None] & cm[None, :], other=0.0)
    tau_vec = tl.zeros((BNB,), dtype=tl.float32)
    for j in range(BNB):
        colj = tl.sum(tl.where(c[None, :] == j, tile, 0.0), axis=1)
        alpha = tl.sum(tl.where(r == j, colj, 0.0))
        xn2 = tl.sum(tl.where(r > j, colj * colj, 0.0))
        reflect = xn2 > 0.0
        sgn = tl.where(alpha >= 0.0, 1.0, -1.0)
        beta = tl.where(reflect, -sgn * tl.sqrt(alpha * alpha + xn2), alpha)
        tau_j = tl.where(reflect, (beta - alpha) / tl.where(reflect, beta, 1.0), 0.0)
        denom = tl.where(reflect, alpha - beta, 1.0)
        vb = colj / denom
        v = tl.where(r == j, 1.0, tl.where(r > j, vb, 0.0))
        vmask = tl.where(r >= j, v, 0.0)
        w = tl.sum(tl.where(c[None, :] > j, vmask[:, None] * tile, 0.0), axis=0)
        tile = tile - tau_j * vmask[:, None] * w[None, :]
        newcol = tl.where(r < j, colj, tl.where(r == j, beta, vb))
        tile = tl.where(c[None, :] == j, newcol[:, None], tile)
        tau_vec = tl.where(c == j, tau_j, tau_vec)
    V = tl.where(
        r[:, None] == c[None, :], 1.0, tl.where(r[:, None] > c[None, :], tile, 0.0)
    )
    tl.store(
        VOUT + b * svb + r[:, None] * svr + c[None, :] * svc,
        V,
        mask=rm[:, None] & cm[None, :],
    )
    Tt = tl.zeros((BNB, BNB), dtype=tl.float32)
    tau0 = tl.sum(tl.where(c == 0, tau_vec, 0.0))
    Tt = tl.where((c[:, None] == 0) & (c[None, :] == 0), tau0, Tt)
    for i in range(1, BNB):
        tau_i = tl.sum(tl.where(c == i, tau_vec, 0.0))
        Vi = tl.sum(tl.where(c[None, :] == i, V, 0.0), axis=1)
        dots = tl.sum(V * Vi[:, None], axis=0)
        z = tl.where(c < i, -tau_i * dots, 0.0)
        Tz = tl.sum(tl.where(c[None, :] < i, Tt * z[None, :], 0.0), axis=1)
        newTcol = tl.where(c < i, Tz, tl.where(c == i, tau_i, 0.0))
        Tt = tl.where(c[None, :] == i, newTcol[:, None], Tt)
    tl.store(
        T + b * sTb + c[:, None] * sTr + c[None, :] * sTc,
        Tt,
        mask=cm[:, None] & cm[None, :],
    )
    tl.store(
        P + b * spb + r[:, None] * spr + c[None, :] * spc,
        tile,
        mask=rm[:, None] & cm[None, :],
    )
    tl.store(TAU + b * stb + c * sti, tau_vec, mask=cm)


@triton.jit
def _wy_update_tf32(
    V,
    T,
    C,
    M,
    IB,
    NCOL,
    svb,
    svr,
    svc,
    sTb,
    sTr,
    sTc,
    scb,
    scr,
    scc,
    BM: tl.constexpr,
    BN: tl.constexpr,
    BIB: tl.constexpr,
):
    b = tl.program_id(0)
    jt = tl.program_id(1)
    cols = jt * BN + tl.arange(0, BN)
    cmask = cols < NCOL
    ic = tl.arange(0, BIB)
    icm = ic < IB
    Tt = tl.load(
        T + b * sTb + ic[:, None] * sTr + ic[None, :] * sTc,
        mask=icm[:, None] & icm[None, :],
        other=0.0,
    )
    W = tl.zeros((BIB, BN), dtype=tl.float32)
    for r0 in tl.range(0, M, BM):
        rows = r0 + tl.arange(0, BM)
        rmask = rows < M
        Vt = tl.load(
            V + b * svb + rows[:, None] * svr + ic[None, :] * svc,
            mask=rmask[:, None] & icm[None, :],
            other=0.0,
        )
        Ct = tl.load(
            C + b * scb + rows[:, None] * scr + cols[None, :] * scc,
            mask=rmask[:, None] & cmask[None, :],
            other=0.0,
        )
        W = tl.dot(tl.trans(Vt), Ct, acc=W, input_precision="tf32x3")
    W = -tl.dot(Tt, W, input_precision="tf32x3")
    for r0 in tl.range(0, M, BM):
        rows = r0 + tl.arange(0, BM)
        rmask = rows < M
        Vt = tl.load(
            V + b * svb + rows[:, None] * svr + ic[None, :] * svc,
            mask=rmask[:, None] & icm[None, :],
            other=0.0,
        )
        Cp = C + b * scb + rows[:, None] * scr + cols[None, :] * scc
        Ct = tl.load(Cp, mask=rmask[:, None] & cmask[None, :], other=0.0)
        Ct = tl.dot(Vt, W, acc=Ct, input_precision="tf32x3")
        tl.store(Cp, Ct, mask=rmask[:, None] & cmask[None, :])


def _launch_wy_tf32(
    Vb, Tt, C, B, M, ib, NCOL, BIB, BM, BN=64, num_warps=4, num_stages=2
):
    grid = (B, triton.cdiv(NCOL, BN))
    TtT = Tt.transpose(-1, -2)
    _wy_update_tf32[grid](
        Vb,
        TtT,
        C,
        M,
        ib,
        NCOL,
        Vb.stride(0),
        Vb.stride(1),
        Vb.stride(2),
        TtT.stride(0),
        TtT.stride(1),
        TtT.stride(2),
        C.stride(0),
        C.stride(1),
        C.stride(2),
        BM=BM,
        BN=BN,
        BIB=BIB,
        num_warps=num_warps,
        num_stages=num_stages,
    )


def _tf32_hi(x):
    return (x.view(torch.int32) & _MANT_MASK).view(torch.float32)


def _mm3(A, B):
    Ah = _tf32_hi(A)
    Al = A - Ah
    Bh = _tf32_hi(B)
    Bl = B - Bh
    out = torch.bmm(Ah, Bh)
    out = torch.baddbmm(out, Ah, Bl)
    out = torch.baddbmm(out, Al, Bh)
    return out


def qr_single(A, block, num_warps):
    B, m, n = A.shape
    bs = int(block)
    BNB = triton.next_power_of_2(bs)
    use_fused = n != 512
    H = _copy_workspace(A, f"H_single_{n}_{bs}", A)
    tau = _workspace(A, f"tau_single_{n}_{bs}", (B, n), zero=False)
    for k in range(0, n, bs):
        ib = min(bs, n - k)
        BM = triton.next_power_of_2(m - k)
        Hv = H[:, k:, k : k + ib]
        Tt = _workspace(A, f"Tt_single_{n}_{bs}_{k}", (B, BNB, BNB), zero=False)
        Vb = _workspace(A, f"Vb_single_{n}_{bs}_{k}", (B, m - k, ib), zero=False)
        tau_panel = tau[:, k : k + ib]
        _panel_kernel[(B,)](
            Hv,
            tau_panel,
            Tt,
            Vb,
            m - k,
            ib,
            Hv.stride(0),
            Hv.stride(1),
            Hv.stride(2),
            tau_panel.stride(0),
            tau_panel.stride(1),
            Tt.stride(0),
            Tt.stride(1),
            Tt.stride(2),
            Vb.stride(0),
            Vb.stride(1),
            Vb.stride(2),
            BM=BM,
            BNB=BNB,
            num_warps=num_warps,
        )
        hi = k + ib
        if hi < n:
            C = H[:, k:, hi:]
            NCOL = n - hi
            if use_fused:
                trailing_m = m - k
                BM_wy = triton.next_power_of_2(min(128, trailing_m))
                _launch_wy_tf32(
                    Vb, Tt, C, B, trailing_m, ib, NCOL, BNB, BM=BM_wy, BN=64
                )
            else:
                V = Vb
                T = Tt[:, :ib, :ib]
                W = V.transpose(-1, -2) @ C
                W = T.transpose(-1, -2) @ W
                C.baddbmm_(V, W, beta=1, alpha=-1)
    return H, tau


def qr_factor_cols(A, cols, block, num_warps):
    B, m, n = A.shape
    H_rect, tau_rect = qr_single(A[:, :, :cols], block, num_warps)
    H = _workspace(A, f"H_factor_cols_{n}_{cols}_{block}", (B, m, n), zero=True)
    H[:, :, :cols] = H_rect
    tau = _workspace(A, f"tau_factor_cols_{n}_{cols}_{block}", (B, n), zero=True)
    tau[:, :cols] = tau_rect
    return H, tau


def qr_factor_cols_project_tail(A, cols, block, num_warps):
    B, m, n = A.shape
    bs = int(block)
    BNB = triton.next_power_of_2(bs)
    H = _copy_workspace(A, f"H_project_tail_{n}_{cols}_{bs}", A)
    tau = _workspace(A, f"tau_project_tail_{n}_{cols}_{bs}", (B, n), zero=True)
    for k in range(0, cols, bs):
        ib = min(bs, cols - k)
        BM = triton.next_power_of_2(m - k)
        Hv = H[:, k:, k : k + ib]
        Tt = _workspace(A, f"Tt_project_tail_{n}_{cols}_{bs}_{k}", (B, BNB, BNB), zero=False)
        Vb = _workspace(A, f"Vb_project_tail_{n}_{cols}_{bs}_{k}", (B, m - k, ib), zero=False)
        tau_panel = tau[:, k : k + ib]
        _panel_kernel[(B,)](
            Hv,
            tau_panel,
            Tt,
            Vb,
            m - k,
            ib,
            Hv.stride(0),
            Hv.stride(1),
            Hv.stride(2),
            tau_panel.stride(0),
            tau_panel.stride(1),
            Tt.stride(0),
            Tt.stride(1),
            Tt.stride(2),
            Vb.stride(0),
            Vb.stride(1),
            Vb.stride(2),
            BM=BM,
            BNB=BNB,
            num_warps=num_warps,
        )
        hi = k + ib
        if hi < n:
            C = H[:, k:, hi:]
            NCOL = n - hi
            trailing_m = m - k
            BM_wy = triton.next_power_of_2(min(128, trailing_m))
            _launch_wy_tf32(Vb, Tt, C, B, trailing_m, ib, NCOL, BNB, BM=BM_wy, BN=64)
    H[:, :, cols:] = torch.triu(H[:, :, cols:], diagonal=-cols)
    return H, tau


def qr_mixed_structured(A, plan):
    B, _, n = A.shape
    H = _workspace(A, f"H_mixed_structured_{n}", (B, n, n), zero=True)
    tau = _workspace(A, f"tau_mixed_structured_{n}", (B, n), zero=True)

    full_idx = plan.get("full")
    if full_idx is not None and full_idx.numel() > 0:
        part = A.index_select(0, full_idx)
        h_part, t_part = qr_single(part, 16 if n >= 1024 else 32, 8 if n >= 1024 else 4)
        H.index_copy_(0, full_idx, h_part)
        tau.index_copy_(0, full_idx, t_part)

    rank_idx = plan.get("rankdef")
    if rank_idx is not None and rank_idx.numel() > 0:
        part = A.index_select(0, rank_idx)
        h_part, t_part = qr_factor_cols(part, _rankdef_cols(n), 16 if n >= 1024 else 32, 8 if n >= 1024 else 4)
        H.index_copy_(0, rank_idx, h_part)
        tau.index_copy_(0, rank_idx, t_part)

    clustered_idx = plan.get("clustered")
    if clustered_idx is not None and clustered_idx.numel() > 0:
        part = A.index_select(0, clustered_idx)
        h_part, t_part = qr_factor_cols(part, _clustered_cols(n), 16 if n >= 1024 else 32, 8 if n >= 1024 else 4)
        H.index_copy_(0, clustered_idx, h_part)
        tau.index_copy_(0, clustered_idx, t_part)

    nearrank_idx = plan.get("nearrank")
    if nearrank_idx is not None and nearrank_idx.numel() > 0:
        part = A.index_select(0, nearrank_idx)
        h_part, t_part = qr_factor_cols_project_tail(part, _rankdef_cols(n), 16, 8)
        H.index_copy_(0, nearrank_idx, h_part)
        tau.index_copy_(0, nearrank_idx, t_part)

    return H, tau


def qr_twolevel(A, ib, NB, num_warps):
    B, m, n = A.shape
    ib = int(ib)
    NB = int(NB)
    BNB_i = triton.next_power_of_2(ib)
    H = _copy_workspace(A, f"H_twolevel_{n}_{ib}_{NB}", A)
    tau = _workspace(A, f"tau_twolevel_{n}_{ib}_{NB}", (B, n), zero=False)

    prev_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        for k0 in range(0, n, NB):
            nb = min(NB, n - k0)
            V_outer = _workspace(A, f"V_outer_{n}_{ib}_{NB}_{k0}", (B, m - k0, nb), zero=True)
            inner_T = []
            inner_off = []
            off = 0
            while off < nb:
                cib = min(ib, nb - off)
                kk = k0 + off
                BM_p = triton.next_power_of_2(m - kk)
                Hv = H[:, kk:, kk : kk + cib]
                Tt = _workspace(A, f"Tt_twolevel_{n}_{ib}_{NB}_{kk}", (B, BNB_i, BNB_i), zero=False)
                Vb = _workspace(A, f"Vb_twolevel_{n}_{ib}_{NB}_{kk}", (B, m - kk, cib), zero=False)
                tau_panel = tau[:, kk : kk + cib]
                _panel_kernel[(B,)](
                    Hv,
                    tau_panel,
                    Tt,
                    Vb,
                    m - kk,
                    cib,
                    Hv.stride(0),
                    Hv.stride(1),
                    Hv.stride(2),
                    tau_panel.stride(0),
                    tau_panel.stride(1),
                    Tt.stride(0),
                    Tt.stride(1),
                    Tt.stride(2),
                    Vb.stride(0),
                    Vb.stride(1),
                    Vb.stride(2),
                    BM=BM_p,
                    BNB=BNB_i,
                    num_warps=num_warps,
                )
                V_outer[:, off:, off : off + cib] = Vb
                blk_end = k0 + nb
                hi_in = kk + cib
                if hi_in < blk_end:
                    Cn = H[:, kk:, hi_in:blk_end]
                    tm = m - kk
                    BM_wy = triton.next_power_of_2(min(128, tm))
                    _launch_wy_tf32(
                        Vb, Tt, Cn, B, tm, cib, blk_end - hi_in, BNB_i, BM=BM_wy, BN=64
                    )
                inner_T.append(_copy_workspace(A, f"inner_T_{n}_{ib}_{NB}_{kk}", Tt[:, :cib, :cib]))
                inner_off.append((off, cib))
                off += cib
            G = _mm3(V_outer.transpose(-1, -2).contiguous(), V_outer)
            T_outer = _workspace(A, f"T_outer_{n}_{ib}_{NB}_{k0}", (B, nb, nb), zero=True)
            for (o, c), Tj in zip(inner_off, inner_T):
                if o > 0:
                    T_outer[:, :o, o : o + c] = (
                        -(T_outer[:, :o, :o] @ G[:, :o, o : o + c]) @ Tj
                    )
                T_outer[:, o : o + c, o : o + c] = Tj
            hi = k0 + nb
            if hi < n:
                Cbig = H[:, k0:, hi:]
                W = _mm3(V_outer.transpose(-1, -2).contiguous(), Cbig)
                W = _mm3(T_outer.transpose(-1, -2).contiguous(), W)
                Cbig.sub_(_mm3(V_outer, W))
    finally:
        torch.backends.cuda.matmul.allow_tf32 = prev_tf32
    return H, tau


def custom_kernel(data: input_t) -> output_t:
    A = _contiguous_workspace_copy(data)
    n = A.shape[-1]
    if n > 2048:
        return torch.geqrf(A)

    plan = _cached_plan(A)
    route = plan["route"]
    if route == "rankdef512":
        return qr_factor_cols(A, _rankdef_cols(n), 32, 4)
    if route == "clustered512":
        return qr_factor_cols(A, _clustered_cols(n), 32, 4)
    if route == "nearrank1024":
        return qr_factor_cols_project_tail(A, _rankdef_cols(n), 16, 8)
    if route == "rankdef1024":
        return qr_factor_cols(A, _rankdef_cols(n), 16, 8)
    if route == "clustered1024":
        return qr_factor_cols(A, _clustered_cols(n), 16, 8)

    if n in USE_TWOLEVEL_FOR:
        ib = 16 if n >= 1024 else 32
        return qr_twolevel(A, ib, 128, num_warps=8)

    if n >= 1024:
        block, nw = 16, 8
    elif n >= 256:
        block, nw = 32, (4 if n == 512 else 8)
    else:
        block, nw = 32, 4
    return qr_single(A, block, nw)
scrolls · 607 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