Skip to content
KernelIndex
Search⌘K

submission 843441

Subho Ghosh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-843441?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.56ms
#52 of 515
2026-06-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:03592dacd34bf7c56a34dbcc32d074636444a25d4fac2110147a97101d6e2c3f
license declaredunknown
license concludedunknown
authorsSubho Ghosh
imported2026-08-26

Techniques

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

clusterdef _panel_cluster_launch(
mbarrier"st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"
shared-memorydef _set_block_rank(smem_ptr, peer, *, loc=None, ip=None):

Kernel source

submission.py1084 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import operator

import torch

import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32
from cutlass._mlir.dialects import llvm
from cutlass.cute.runtime import from_dlpack
from cutlass.cutlass_dsl import T, dsl_user_op

from task import input_t, output_t

_NB = 32
_CN = 64
_TPB = 512
_PANEL_NW = _TPB // 32

_GEQR2_TPB = 256
_GEQR2_TDIM = 16

_PC_TPB = 512
_PC_NW = _PC_TPB // 32
_PC_C = 8


def _t2c(t, align=32):
    return from_dlpack(t, assumed_align=align)




@cute.jit
def _block_sum(val, red, warp, lane):
    NW = cutlass.const_expr(cute.size(red))
    v = cute.arch.warp_reduction(val, operator.add)
    if lane == 0:
        red[warp] = v
    cute.arch.barrier()
    total = cutlass.Float32(0.0)
    for w in cutlass.range_constexpr(NW):
        total = total + red[w]
    cute.arch.barrier()
    return total


@dsl_user_op
def _set_block_rank(smem_ptr, peer, *, loc=None, ip=None):
    p = smem_ptr.toint(loc=loc, ip=ip).ir_value()
    return Int32(
        llvm.inline_asm(
            T.i32(), [p, peer.ir_value()],
            "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r",
            has_side_effects=False, is_align_stack=False,
            asm_dialect=llvm.AsmDialect.AD_ATT,
        )
    )


@dsl_user_op
def _load_remote(smem_ptr, peer, *, loc=None, ip=None):
    addr = _set_block_rank(smem_ptr, peer, loc=loc, ip=ip).ir_value()
    return Float32(
        llvm.inline_asm(
            T.f32(), [addr],
            "ld.shared::cluster.f32 $0, [$1];", "=f,r",
            has_side_effects=True, is_align_stack=False,
            asm_dialect=llvm.AsmDialect.AD_ATT,
        )
    )


@dsl_user_op
def _store_remote_v4(
    v0: Float32,
    v1: Float32,
    v2: Float32,
    v3: Float32,
    smem_ptr,
    mbar_ptr,
    peer,
    *,
    loc=None,
    ip=None,
):
    dst = _set_block_rank(smem_ptr, peer, loc=loc, ip=ip).ir_value()
    bar = _set_block_rank(mbar_ptr, peer, loc=loc, ip=ip).ir_value()
    llvm.inline_asm(
        None,
        [
            dst,
            bar,
            v0.ir_value(loc=loc, ip=ip),
            v1.ir_value(loc=loc, ip=ip),
            v2.ir_value(loc=loc, ip=ip),
            v3.ir_value(loc=loc, ip=ip),
        ],
        "{\n\t"
        ".reg .v4 .f32 abcd;\n\t"
        "mov.f32 abcd.x, $2;\n\t"
        "mov.f32 abcd.y, $3;\n\t"
        "mov.f32 abcd.z, $4;\n\t"
        "mov.f32 abcd.w, $5;\n\t"
        "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"
        "}\n",
        "r,r,f,f,f,f",
        has_side_effects=True,
        is_align_stack=False,
        asm_dialect=llvm.AsmDialect.AD_ATT,
    )


@cute.jit
def _cluster_vsum(redbuf, total, cnt, rank, tidx, C: cutlass.Constexpr):
    cute.arch.barrier()
    cute.arch.cluster_arrive_relaxed()
    cute.arch.cluster_wait()
    for v in cutlass.range(tidx, cnt, _PC_TPB):
        s = redbuf[v]
        for p in cutlass.range_constexpr(C):
            if p != rank:
                s = s + _load_remote(redbuf.iterator + v, Int32(p))
        total[v] = s
    cute.arch.barrier()
    cute.arch.cluster_arrive_relaxed()
    cute.arch.cluster_wait()


@cute.jit
def _cluster_vsum_db(redbuf, total, buf, first, rank, tidx, C: cutlass.Constexpr):
    cute.arch.barrier()
    cute.arch.cluster_arrive_relaxed()
    cute.arch.cluster_wait()
    for q in cutlass.range(tidx, 2 * _NB, _PC_TPB):
        v = buf + q
        s = redbuf[v]
        for p in cutlass.range_constexpr(C):
            if p != rank:
                s = s + _load_remote(redbuf.iterator + v, Int32(p))
        total[v] = s
    cute.arch.barrier()


@cute.jit
def _cluster_vsum_async(
    redbuf, total, recv, mbars, buf, jj, rank, tidx,
    C: cutlass.Constexpr, TPB: cutlass.Constexpr,
):
    cute.arch.barrier()
    slot = jj & 1
    if tidx < 16:
        q = 4 * tidx
        dst = recv.iterator + (slot * C + rank) * (2 * _NB) + q
        for peer in cutlass.range_constexpr(C):
            _store_remote_v4(
                redbuf[buf + q],
                redbuf[buf + q + 1],
                redbuf[buf + q + 2],
                redbuf[buf + q + 3],
                dst,
                mbars + slot,
                Int32(peer),
            )
    cute.arch.mbarrier_wait(mbars + slot, (jj // 2) & 1)
    for q in cutlass.range(tidx, 2 * _NB, TPB):
        s = cutlass.Float32(0.0)
        for peer in cutlass.range_constexpr(C):
            s = s + recv[slot, peer, q]
        total[buf + q] = s
    cute.arch.barrier()
    if tidx == 0 and jj + 2 < _NB:
        cute.arch.mbarrier_arrive_and_expect_tx(
            mbars + slot, C * (2 * _NB) * 4
        )


@cute.kernel
def _geqr2_small_kernel(
    mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor,
    tiled_copy: cute.TiledCopy, rows: cutlass.Constexpr,
    cols: cutlass.Constexpr, TPB: cutlass.Constexpr,
    NW: cutlass.Constexpr,
):
    tidx, _, _ = cute.arch.thread_idx()
    b, _, _ = cute.arch.block_idx()
    warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
    lane = cute.arch.lane_idx()
    smem = cutlass.utils.SmemAllocator()
    sH = smem.allocate_tensor(
        cutlass.Float32,
        cute.make_layout((rows, cols), stride=(cols + 1, 1)),
        byte_alignment=16,
    )
    red = smem.allocate_tensor(cutlass.Float32, cute.make_layout(NW), byte_alignment=16)

    gA = mA[b, None, None]
    gH = mH[b, None, None]
    tc = tiled_copy.get_slice(tidx)

    if cutlass.const_expr(rows == cols):
        g_src = tc.partition_S(gA)
        frag = cute.make_rmem_tensor_like(g_src)
        cute.copy(tiled_copy, g_src, frag)
        cute.copy(tiled_copy, frag, tc.partition_D(sH))
    else:
        for idx in cutlass.range(tidx, rows * cols, TPB):
            sH[idx // cols, idx % cols] = gA[idx // cols, idx % cols]
    cute.arch.barrier()

    for j in cutlass.range(0, cols):
        m = rows - j
        p = cols - j
        sT = cute.domain_offset((j, j), sH)
        local = cutlass.Float32(0.0)
        for i in cutlass.range(1 + tidx, m, TPB):
            x = sT[i, 0]
            local = local + x * x
        xnorm2 = _block_sum(local, red, warp, lane)
        alpha = sT[0, 0]
        nrm = cute.math.sqrt(alpha * alpha + xnorm2)
        beta = nrm
        if alpha >= 0.0:
            beta = -nrm
        tau_j = cutlass.Float32(0.0)
        scale = cutlass.Float32(0.0)
        beta_diag = alpha
        if xnorm2 > 0.0:
            tau_j = (beta - alpha) / beta
            scale = 1.0 / (alpha - beta)
            beta_diag = beta
        if tidx == 0:
            sT[0, 0] = beta_diag
            mTau[b, j] = tau_j
        for i in cutlass.range(1 + tidx, m, TPB):
            sT[i, 0] = sT[i, 0] * scale
        cute.arch.barrier()
        for c in cutlass.range(1 + warp, p, NW):
            dot = cutlass.Float32(0.0)
            for i in cutlass.range(1 + lane, m, 32):
                dot = dot + sT[i, 0] * sT[i, c]
            dot = cute.arch.warp_reduction(dot, operator.add)
            tw = tau_j * (dot + sT[0, c])
            if lane == 0:
                sT[0, c] = sT[0, c] - tw
            for i in cutlass.range(1 + lane, m, 32):
                sT[i, c] = sT[i, c] - sT[i, 0] * tw
        cute.arch.barrier()

    if cutlass.const_expr(rows == cols):
        s_src = tc.partition_S(sH)
        frag2 = cute.make_rmem_tensor_like(s_src)
        cute.copy(tiled_copy, s_src, frag2)
        cute.copy(tiled_copy, frag2, tc.partition_D(gH))
    else:
        for idx in cutlass.range(tidx, rows * cols, TPB):
            gH[idx // cols, idx % cols] = sH[idx // cols, idx % cols]


@cute.jit
def _geqr2_launch(mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor):
    rows = mA.shape[1]
    cols = mA.shape[2]
    tpb = 1024 if rows == 128 and cols == 128 else (512 if rows >= 64 else 256)
    TR = 32 if tpb == 1024 else _GEQR2_TDIM
    TC = 32 if tpb >= 512 else _GEQR2_TDIM
    VR = cutlass.const_expr(rows // TR)
    VC = cutlass.const_expr(cols // TC)
    thr = cute.make_ordered_layout((TR, TC), order=(1, 0))
    val = cute.make_ordered_layout((VR, VC), order=(1, 0))
    atom = cute.make_copy_atom(
        cute.nvgpu.CopyUniversalOp(), cutlass.Float32, num_bits_per_copy=32
    )
    tiled_copy = cute.make_tiled_copy_tv(atom, thr, val)
    _geqr2_small_kernel(mA, mH, mTau, tiled_copy, rows, cols, tpb, tpb // 32).launch(
        grid=[mA.shape[0], 1, 1], block=[tpb, 1, 1]
    )


_geqr2_cache: dict = {}
_tail_cache: dict = {}


def geqr2_small(A: torch.Tensor):
    batch, n, _ = A.shape
    H = torch.empty_like(A)
    tau = torch.empty((batch, n), device=A.device, dtype=torch.float32)
    key = (batch, n)
    if key not in _geqr2_cache:
        _geqr2_cache[key] = cute.compile(
            _geqr2_launch, _t2c(A, 16), _t2c(H, 16), _t2c(tau, 16)
        )
    _geqr2_cache[key](_t2c(A, 16), _t2c(H, 16), _t2c(tau, 16))
    return H, tau


def _geqr2_tail(H: torch.Tensor, tau: torch.Tensor, start: int, stop: int):
    tail = H[:, start:, start:stop]
    tail_tau = tau[:, start:stop]
    key = (H.shape[0], H.shape[1], start, stop)
    mTail = _t2c(tail, 16)
    mTau = _t2c(tail_tau, 16)
    if key not in _tail_cache:
        _tail_cache[key] = cute.compile(_geqr2_launch, mTail, mTail, mTau)
    _tail_cache[key](mTail, mTail, mTau)


@cute.kernel
def _panel_resident_kernel(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mV: cute.Tensor,
    k: cutlass.Int32, n: cutlass.Constexpr,
    form_t: cutlass.Constexpr, TPB: cutlass.Constexpr,
    NW: cutlass.Constexpr,
):
    tidx, _, _ = cute.arch.thread_idx()
    b, _, _ = cute.arch.block_idx()
    k = cute.assume(k, divby=_NB)
    warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
    lane = cute.arch.lane_idx()

    m = n - k
    gP = cute.domain_offset((k, k), mH[b, None, None])
    gT = mT[b, k // _NB, None, None]

    smem = cutlass.utils.SmemAllocator()
    sP = smem.allocate_tensor(cutlass.Float32, cute.make_layout((n, _NB), stride=(_NB + 1, 1)), byte_alignment=16)
    sT = smem.allocate_tensor(cutlass.Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
    sS = smem.allocate_tensor(cutlass.Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
    red = smem.allocate_tensor(cutlass.Float32, cute.make_layout(NW), byte_alignment=16)
    s_tau = smem.allocate_tensor(cutlass.Float32, cute.make_layout(_NB), byte_alignment=16)
    s_sc = smem.allocate_tensor(cutlass.Float32, cute.make_layout(2), byte_alignment=16)

    for idx in cutlass.range(tidx, m * _NB, TPB):
        sP[idx // _NB, idx % _NB] = gP[idx // _NB, idx % _NB]
    for idx in cutlass.range(tidx, _NB * _NB, TPB):
        sT[idx // _NB, idx % _NB] = 0.0
    cute.arch.barrier()

    for jj in cutlass.range(0, _NB, 1, unroll=1):
        local = cutlass.Float32(0.0)
        for i in cutlass.range(jj + 1 + tidx, m, TPB):
            x = sP[i, jj]
            local = local + x * x
        partial = cute.arch.warp_reduction(local, operator.add)
        if lane == 0:
            red[warp] = partial
        cute.arch.barrier()
        if tidx == 0:
            xnorm2 = cutlass.Float32(0.0)
            for w in cutlass.range_constexpr(NW):
                xnorm2 = xnorm2 + red[w]
            alpha = sP[jj, jj]
            nrm = cute.math.sqrt(alpha * alpha + xnorm2)
            beta = nrm
            if alpha >= 0.0:
                beta = -nrm
            tau_j = cutlass.Float32(0.0)
            scale = cutlass.Float32(0.0)
            beta_diag = alpha
            if xnorm2 > 0.0:
                tau_j = (beta - alpha) / beta
                scale = 1.0 / (alpha - beta)
                beta_diag = beta
            sP[jj, jj] = beta_diag
            mTau[b, k + jj] = tau_j
            s_tau[jj] = tau_j
            s_sc[0] = tau_j
            s_sc[1] = scale
        cute.arch.barrier()
        tau_j = s_sc[0]
        scale = s_sc[1]
        for i in cutlass.range(jj + 1 + tidx, m, TPB):
            sP[i, jj] = sP[i, jj] * scale
        cute.arch.barrier()
        for c in cutlass.range(jj + 1 + warp, _NB, NW):
            dot = cutlass.Float32(0.0)
            for i in cutlass.range(jj + 1 + lane, m, 32):
                dot = dot + sP[i, jj] * sP[i, c]
            dot = cute.arch.warp_reduction(dot, operator.add)
            tw = tau_j * (dot + sP[jj, c])
            if lane == 0:
                sP[jj, c] = sP[jj, c] - tw
            for i in cutlass.range(jj + 1 + lane, m, 32):
                sP[i, c] = sP[i, c] - sP[i, jj] * tw
        cute.arch.barrier()

    if cutlass.const_expr(form_t):
        for idx in cutlass.range(tidx, _NB * _NB, TPB):
            l = idx // _NB
            jc = idx % _NB
            if l < jc:
                s = sP[jc, l]
                for r in cutlass.range(jc + 1, m, 1):
                    s = s + sP[r, l] * sP[r, jc]
                sS[l, jc] = s
        cute.arch.barrier()
        for i in cutlass.range_constexpr(_NB):
            if tidx < _NB:
                l = tidx
                if l == i:
                    sT[i, i] = s_tau[i]
                elif l < i:
                    acc = cutlass.Float32(0.0)
                    for p in cutlass.range(l, i, 1):
                        acc = acc + sT[l, p] * sS[p, i]
                    sT[l, i] = -s_tau[i] * acc
            cute.arch.barrier()

    for idx in cutlass.range(tidx, m * _NB, TPB):
        r = idx // _NB
        c = idx % _NB
        value = sP[r, c]
        logical = value
        if r < _NB:
            if r < c:
                logical = cutlass.Float32(0.0)
            elif r == c:
                logical = cutlass.Float32(1.0)
        gP[r, c] = value
        mV[b, r, c] = logical
    if cutlass.const_expr(form_t):
        for idx in cutlass.range(tidx, _NB * _NB, TPB):
            gT[idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]


@cute.jit
def _panel_resident_launch(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mV: cute.Tensor,
    k: cutlass.Int32,
):
    tpb = 1024 if mH.shape[1] == 352 else 512
    _panel_resident_kernel(mH, mTau, mT, mV, k, mH.shape[1], True, tpb, tpb // 32).launch(
        grid=[mH.shape[0], 1, 1], block=[tpb, 1, 1]
    )


@cute.jit
def _panel_resident_no_t_launch(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mV: cute.Tensor,
    k: cutlass.Int32,
):
    tpb = 1024 if mH.shape[1] == 352 else 512
    _panel_resident_kernel(mH, mTau, mT, mV, k, mH.shape[1], False, tpb, tpb // 32).launch(
        grid=[mH.shape[0], 1, 1], block=[tpb, 1, 1]
    )


@cute.kernel
def _larft_from_gram_kernel(
    mGram: cute.Tensor,
    mTau: cute.Tensor,
    mT: cute.Tensor,
    k: cutlass.Int32,
):
    tidx, _, _ = cute.arch.thread_idx()
    b, _, _ = cute.arch.block_idx()
    smem = cutlass.utils.SmemAllocator()
    sT = smem.allocate_tensor(
        cutlass.Float32,
        cute.make_layout((_NB, _NB), stride=(_NB, 1)),
        byte_alignment=128,
    )
    for idx in cutlass.range(tidx, _NB * _NB, 32):
        sT[idx // _NB, idx % _NB] = 0.0
    cute.arch.barrier()
    for col in cutlass.range_constexpr(_NB):
        row = tidx
        if row == col:
            sT[row, col] = mTau[b, k + col]
        elif row < col:
            value = cutlass.Float32(0.0)
            for p in cutlass.range(row, col, 1):
                value = value + sT[row, p] * mGram[b, p, col]
            sT[row, col] = -mTau[b, k + col] * value
        cute.arch.sync_warp()
    for idx in cutlass.range(tidx, _NB * _NB, 32):
        mT[b, k // _NB, idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]


@cute.jit
def _larft_from_gram_launch(
    mGram: cute.Tensor,
    mTau: cute.Tensor,
    mT: cute.Tensor,
    k: cutlass.Int32,
):
    _larft_from_gram_kernel(mGram, mTau, mT, k).launch(
        grid=[mGram.shape[0], 1, 1], block=[32, 1, 1]
    )


@cute.kernel
def _diag_prepare_kernel(
    mH: cute.Tensor,
    mR: cute.Tensor,
    k: cutlass.Int32,
):
    tidx, _, _ = cute.arch.thread_idx()
    b, _, _ = cute.arch.block_idx()
    p = k // _NB
    for idx in cutlass.range(tidx, _NB * _NB, _TPB):
        i = idx // _NB
        j = idx % _NB
        value = mH[b, k + i, k + j]
        mR[b, p, i, j] = value
        if i < j:
            value = cutlass.Float32(0.0)
        elif i == j:
            value = cutlass.Float32(1.0)
        mH[b, k + i, k + j] = value


@cute.jit
def _diag_prepare_launch(
    mH: cute.Tensor,
    mR: cute.Tensor,
    k: cutlass.Int32,
):
    _diag_prepare_kernel(mH, mR, k).launch(
        grid=[mH.shape[0], 1, 1], block=[_TPB, 1, 1]
    )


@cute.kernel
def _diag_restore_kernel(
    mH: cute.Tensor,
    mR: cute.Tensor,
):
    tidx, _, _ = cute.arch.thread_idx()
    p, b, _ = cute.arch.block_idx()
    k = p * _NB
    for idx in cutlass.range(tidx, _NB * _NB, _TPB):
        i = idx // _NB
        j = idx % _NB
        mH[b, k + i, k + j] = mR[b, p, i, j]


@cute.jit
def _diag_restore_launch(
    mH: cute.Tensor,
    mR: cute.Tensor,
):
    _diag_restore_kernel(mH, mR).launch(
        grid=[mR.shape[1], mH.shape[0], 1], block=[_TPB, 1, 1]
    )


@cute.kernel
def _panel_cluster_kernel(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mR: cute.Tensor,
    k: cutlass.Int32, n: cutlass.Constexpr, C: cutlass.Constexpr,
    rows_cap: cutlass.Constexpr,
    form_t: cutlass.Constexpr, TPB: cutlass.Constexpr,
    NW: cutlass.Constexpr,
):
    tidx, _, _ = cute.arch.thread_idx()
    rank = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster())
    _, b, _ = cute.arch.block_idx()
    k = cute.assume(k, divby=_NB)
    warp = cute.arch.make_warp_uniform(cute.arch.warp_idx())
    lane = cute.arch.lane_idx()

    m = n - k
    S = cute.ceil_div(m, C)
    row0 = rank * S
    avail = m - row0
    rows = avail if avail < S else S
    rows = rows if rows > 0 else 0

    gP = cute.domain_offset((k, k), mH[b, None, None])
    gT = mT[b, k // _NB, None, None]

    smem = cutlass.utils.SmemAllocator()
    sP = smem.allocate_tensor(Float32, cute.make_layout((rows_cap, _NB), stride=(_NB + 1, 1)), byte_alignment=16)
    sT = smem.allocate_tensor(Float32, cute.make_layout((_NB, _NB), stride=(_NB, 1)), byte_alignment=16)
    redbuf = smem.allocate_tensor(Float32, cute.make_layout(2 * _NB * _NB), byte_alignment=16)
    total = smem.allocate_tensor(Float32, cute.make_layout(2 * _NB * _NB), byte_alignment=16)
    recv = smem.allocate_tensor(
        Float32,
        cute.make_layout(
            (2, C, 2 * _NB),
            stride=(C * 2 * _NB, 2 * _NB, 1),
        ),
        byte_alignment=16,
    )
    mbars = smem.allocate_array(cutlass.Int64, num_elems=2)
    red = smem.allocate_tensor(Float32, cute.make_layout(NW), byte_alignment=16)
    s_tau = smem.allocate_tensor(Float32, cute.make_layout(_NB), byte_alignment=16)
    s_sc = smem.allocate_tensor(Float32, cute.make_layout(2), byte_alignment=16)

    for idx in cutlass.range(tidx, rows * _NB, TPB):
        sP[idx // _NB, idx % _NB] = gP[row0 + idx // _NB, idx % _NB]
    for idx in cutlass.range(tidx, _NB * _NB, TPB):
        sT[idx // _NB, idx % _NB] = 0.0
    if tidx < 2:
        cute.arch.mbarrier_init(mbars + tidx, 1)
        cute.arch.mbarrier_arrive_and_expect_tx(
            mbars + tidx, C * (2 * _NB) * 4
        )
    cute.arch.mbarrier_init_fence()
    cute.arch.cluster_arrive_relaxed()
    cute.arch.cluster_wait()
    cute.arch.barrier()

    for jj in cutlass.range(0, _NB, 1, unroll=1):
        rbuf = (jj & 1) * _NB * _NB
        jloc = jj - row0
        owns = (jloc >= 0) and (jloc < rows)
        lo = jj + 1 - row0
        lo = lo if lo > 0 else 0

        for c in cutlass.range(jj + warp, _NB, NW):
            dot = cutlass.Float32(0.0)
            for i in cutlass.range(lo + lane, rows, 32):
                dot = dot + sP[i, jj] * sP[i, c]
            dot = cute.arch.warp_reduction(dot, operator.add)
            if lane == 0:
                redbuf[rbuf + c] = dot
                redbuf[rbuf + _NB + c] = sP[jloc, c] if owns else cutlass.Float32(0.0)
        _cluster_vsum_async(redbuf, total, recv, mbars, rbuf, jj, rank, tidx, C, TPB)
        xnorm2 = total[rbuf + jj]
        alpha = total[rbuf + _NB + jj]

        nrm = cute.math.sqrt(alpha * alpha + xnorm2)
        beta = nrm
        if alpha >= 0.0:
            beta = -nrm
        tau_j = cutlass.Float32(0.0)
        scale = cutlass.Float32(0.0)
        if xnorm2 > 0.0:
            tau_j = (beta - alpha) / beta
            scale = 1.0 / (alpha - beta)
        if tidx == 0:
            s_tau[jj] = tau_j
            s_sc[0] = tau_j
            s_sc[1] = scale
            if owns:
                sP[jloc, jj] = beta if xnorm2 > 0.0 else alpha
                mTau[b, k + jj] = tau_j
        cute.arch.barrier()
        tau_j = s_sc[0]
        scale = s_sc[1]

        for i in cutlass.range(lo + tidx, rows, TPB):
            sP[i, jj] = sP[i, jj] * scale
        cute.arch.barrier()

        for c in cutlass.range(jj + 1 + warp, _NB, NW):
            tw = tau_j * (total[rbuf + _NB + c] + scale * total[rbuf + c])
            if owns and lane == 0:
                sP[jloc, c] = sP[jloc, c] - tw
            for i in cutlass.range(lo + lane, rows, 32):
                sP[i, c] = sP[i, c] - sP[i, jj] * tw
        cute.arch.barrier()

    if cutlass.const_expr(form_t):
        cute.arch.cluster_arrive_relaxed()
        cute.arch.cluster_wait()
        for idx in cutlass.range(tidx, _NB * _NB, TPB):
            redbuf[idx] = 0.0
        cute.arch.barrier()
        for idx in cutlass.range(tidx, _NB * _NB, TPB):
            l = idx // _NB
            jc = idx % _NB
            if l < jc:
                s = cutlass.Float32(0.0)
                jc_loc = jc - row0
                if (jc_loc >= 0) and (jc_loc < rows):
                    s = s + sP[jc_loc, l]
                lo2 = jc + 1 - row0
                lo2 = lo2 if lo2 > 0 else 0
                for r in cutlass.range(lo2, rows, 1):
                    s = s + sP[r, l] * sP[r, jc]
                redbuf[idx] = s
        cute.arch.barrier()
        cute.arch.cluster_arrive_relaxed()
        cute.arch.cluster_wait()
        if rank == 0:
            for v in cutlass.range(tidx, _NB * _NB, TPB):
                s = redbuf[v]
                for p in cutlass.range_constexpr(C):
                    if p != 0:
                        s = s + _load_remote(redbuf.iterator + v, Int32(p))
                total[v] = s
            cute.arch.barrier()
            for i in cutlass.range_constexpr(_NB):
                if tidx < _NB:
                    l = tidx
                    if l == i:
                        sT[i, i] = s_tau[i]
                    elif l < i:
                        acc = cutlass.Float32(0.0)
                        for p in cutlass.range(l, i, 1):
                            acc = acc + sT[l, p] * total[p * _NB + i]
                        sT[l, i] = -s_tau[i] * acc
                cute.arch.sync_warp()
        cute.arch.cluster_arrive_relaxed()
        cute.arch.cluster_wait()

    for idx in cutlass.range(tidx, rows * _NB, TPB):
        r = idx // _NB
        c = idx % _NB
        gr = row0 + r
        value = sP[r, c]
        if gr < _NB:
            logical = value
            if gr < c:
                logical = cutlass.Float32(0.0)
            elif gr == c:
                logical = cutlass.Float32(1.0)
            mR[b, gr, c] = logical
        else:
            mR[b, gr, c] = value
        gP[gr, c] = value
    if rank == 0 and cutlass.const_expr(form_t):
        for idx in cutlass.range(tidx, _NB * _NB, TPB):
            gT[idx // _NB, idx % _NB] = sT[idx // _NB, idx % _NB]


@cute.jit
def _panel_cluster_launch(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mR: cute.Tensor,
    k: cutlass.Int32, C: cutlass.Constexpr,
    rows_cap: cutlass.Constexpr,
):
    n = mH.shape[1]
    batch = mH.shape[0]
    tpb = 1024 if n >= 2048 else 512
    _panel_cluster_kernel(mH, mTau, mT, mR, k, n, C, rows_cap, True, tpb, tpb // 32).launch(
        grid=[C, batch, 1], block=[tpb, 1, 1], cluster=(C, 1, 1)
    )


@cute.jit
def _panel_cluster_no_t_launch(
    mH: cute.Tensor, mTau: cute.Tensor, mT: cute.Tensor,
    mR: cute.Tensor,
    k: cutlass.Int32, C: cutlass.Constexpr,
    rows_cap: cutlass.Constexpr,
):
    n = mH.shape[1]
    batch = mH.shape[0]
    tpb = 1024 if n >= 2048 else 512
    _panel_cluster_kernel(mH, mTau, mT, mR, k, n, C, rows_cap, False, tpb, tpb // 32).launch(
        grid=[C, batch, 1], block=[tpb, 1, 1], cluster=(C, 1, 1)
    )


def _alloc_vg(batch, n):
    return torch.empty((batch, n, _NB), device="cuda", dtype=torch.float32)


def _materialize_vg(H, Vg, k, strict_lower=None, eye=None):
    n = H.shape[1]
    m = n - k
    if strict_lower is None:
        ii = torch.arange(_NB, device=H.device)
        strict_lower = (ii[:, None] > ii[None, :]).to(torch.float32)
    if eye is None:
        eye = torch.eye(_NB, device=H.device, dtype=torch.float32)
    Vg[:, :_NB, :] = H[:, k : k + _NB, k : k + _NB] * strict_lower + eye
    if m > _NB:
        Vg[:, _NB:m, :] = H[:, k + _NB : n, k : k + _NB]


_panel_cache: dict = {}
_larft_cache: dict = {}
_diag_cache: dict = {}
_buf_cache: dict = {}


def _effective_factor_size(data: torch.Tensor):
    batch, n, _ = data.shape
    if n == 512 and batch >= 16:
        edge_by_matrix = torch.maximum(
            data[:, :8, -1].abs().amax(dim=1),
            data[:, -8:, -1].abs().amax(dim=1),
        )
        edge_min_t, edge_max_t = torch.aminmax(edge_by_matrix)
        edge_min = edge_min_t.item()
        edge_max = edge_max_t.item()
        if edge_max == 0.0:
            return 384, "rankdef"
        if edge_max < 1.0e-4:
            first_col_offdiag = data[:, 1:8, 0].abs().amax().item()
            if first_col_offdiag > 1.0e-4:
                return 256, "clustered"
        if batch >= 640:
            if edge_min > 1.0e-4:
                return n, "dense512"
            if edge_min == 0.0:
                return n, "mixed512"
    if n == 1024 and batch >= 4:
        tail_error = (data[:, 255, -1] - data[:, 255, 255]).abs().amax().item()
        if tail_error < 2.0e-4:
            return 768, "nearrank"
    return n, None


def _nbo(n):
    if n <= 352:
        return _NB
    if n == 512:
        return _NB
    return _NB


def _cluster_schedule(n):
    if n == 1024:
        return ((0, 4, 256),)
    if n == 2048:
        return ((0, 8, 256),)
    if n == 4096:
        return ((0, 16, 256),)
    return ((0, 8, (n + 7) // 8),)


def _get_buf(batch, n, NBO, device):
    key = (batch, n, NBO)
    if key not in _buf_cache:
        slL = torch.tril(torch.ones(NBO, NBO, device=device), -1)
        eyeL = torch.eye(NBO, device=device, dtype=torch.float32)
        _buf_cache[key] = (
            torch.empty(batch, n, NBO, device=device, dtype=torch.float32),
            torch.empty(batch, NBO, NBO, device=device, dtype=torch.float32),
            slL, eyeL,
        )
    return _buf_cache[key]


def _blocked_qr(data: torch.Tensor):
    batch, n, _ = data.shape
    factor_size, structure = _effective_factor_size(data)
    update_size = factor_size if structure is not None else n
    tail_start = (
        64 if n == 192 and factor_size == n
        # The batched 64x64 tail kernel is nondeterministic at the n=512,
        # batch=640 scored shapes.  Finish those with the resident panel path.
        else n - 64 if n >= 352 and n != 512 and factor_size == n
        else n
    )
    if n >= 1024:
        use_tf32 = True
    elif n == 512:
        use_tf32 = False
    elif n >= 512:
        use_tf32 = True
    else:
        use_tf32 = False
    torch.backends.cuda.matmul.allow_tf32 = use_tf32
    tf32_state = use_tf32
    H = data.clone()
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    T = torch.empty((batch, (n + _NB - 1) // _NB, _NB, _NB), device=data.device, dtype=torch.float32)
    Vpanel = _alloc_vg(batch, n)
    mH, mTau, mT = _t2c(H), _t2c(tau), _t2c(T)
    mR = _t2c(Vpanel)
    panel_schedule = None
    if n >= 1024:
        panel_schedule = []
        for start, C, rows_cap in _cluster_schedule(n):
            full_key = (batch, n, "cluster", C, rows_cap, True)
            no_t_key = (batch, n, "cluster", C, rows_cap, False)
            if full_key not in _panel_cache:
                _panel_cache[full_key] = cute.compile(
                    _panel_cluster_launch, mH, mTau, mT, mR,
                    cutlass.Int32(0), C, rows_cap,
                )
            if no_t_key not in _panel_cache:
                _panel_cache[no_t_key] = cute.compile(
                    _panel_cluster_no_t_launch, mH, mTau, mT, mR,
                    cutlass.Int32(0), C, rows_cap,
                )
            panel_schedule.append(
                (start, _panel_cache[full_key], _panel_cache[no_t_key])
            )
        panel = panel_schedule[0][1]
        last_panel = panel_schedule[0][2]
    else:
        key = (batch, n, "no_t" if n == 512 else "full")
        if key not in _panel_cache:
            launch = _panel_resident_no_t_launch if n == 512 else _panel_resident_launch
            _panel_cache[key] = cute.compile(
                launch, mH, mTau, mT, mR, cutlass.Int32(0)
            )
        panel = _panel_cache[key]
        if n == 512:
            last_panel = panel
        else:
            last_key = (batch, n, "last_no_t")
            if last_key not in _panel_cache:
                _panel_cache[last_key] = cute.compile(
                    _panel_resident_no_t_launch,
                    mH, mTau, mT, mR, cutlass.Int32(0),
                )
            last_panel = _panel_cache[last_key]

    NBO = _nbo(n)
    if NBO == _NB:
        Vg = Vpanel
        Gram = (
            torch.empty((batch, _NB, _NB), device=data.device, dtype=torch.float32)
            if n == 512
            else None
        )
        if n == 512:
            mGram = _t2c(Gram)
            lkey = (batch, n)
            if lkey not in _larft_cache:
                _larft_cache[lkey] = cute.compile(
                    _larft_from_gram_launch,
                    mGram,
                    mTau,
                    mT,
                    cutlass.Int32(0),
                )
            larft = _larft_cache[lkey]
        Wbuf = torch.empty(
            (batch, _NB, n), device=data.device, dtype=torch.float32
        )
        Auxbuf = torch.empty(
            (batch, _NB, n) if n == 512 else (batch, n, _NB),
            device=data.device,
            dtype=torch.float32,
        )
        for k in range(0, n, _NB):
            if k >= tail_start:
                break
            if k >= factor_size:
                break
            if panel_schedule is not None:
                for start, full_panel, no_t_panel in panel_schedule:
                    if k >= start:
                        panel = full_panel
                        last_panel = no_t_panel
            active_panel = last_panel if k + _NB >= factor_size else panel
            active_panel(mH, mTau, mT, mR, cutlass.Int32(k))
            if k + _NB >= factor_size:
                break
            mm = n - k
            V = Vpanel[:, :mm, :]
            A22 = H[:, k:n, k + _NB : update_size]
            if n == 512:
                full_tf32 = batch >= 640 and (
                    k >= 32
                    or structure in ("dense512", "rankdef", "clustered")
                )
                projection_only_tf32 = structure == "mixed512" and k == 0
                if full_tf32 or projection_only_tf32:
                    torch.backends.cuda.matmul.allow_tf32 = True
                    tf32_state = full_tf32
                cols = update_size - k - _NB
                W = Wbuf[:, :, :cols]
                gram_tf32 = structure in (
                    "dense512", "rankdef", "clustered"
                )
                if gram_tf32 and not tf32_state:
                    torch.backends.cuda.matmul.allow_tf32 = True
                    tf32_state = True
                elif not gram_tf32 and tf32_state:
                    torch.backends.cuda.matmul.allow_tf32 = False
                    tf32_state = False
                torch.bmm(V.transpose(1, 2), V, out=Gram)
                larft(mGram, mTau, mT, cutlass.Int32(k))
                if full_tf32 or projection_only_tf32:
                    torch.backends.cuda.matmul.allow_tf32 = True
                    tf32_state = full_tf32
                torch.bmm(V.transpose(1, 2), A22, out=W)
                transform_tf32 = full_tf32 and structure in (
                    "dense512", "rankdef", "clustered"
                )
                if (tf32_state or projection_only_tf32) and not transform_tf32:
                    torch.backends.cuda.matmul.allow_tf32 = False
                W2 = Auxbuf[:, :, :cols]
                torch.bmm(T[:, k // _NB].transpose(1, 2), W, out=W2)
                if tf32_state and not transform_tf32:
                    torch.backends.cuda.matmul.allow_tf32 = True
                torch.baddbmm(A22, V, W2, beta=1.0, alpha=-1.0, out=A22)
            elif n == 352:
                W = torch.bmm(V.transpose(1, 2), A22)
                W2 = torch.bmm(T[:, k // _NB].transpose(1, 2), W)
                torch.baddbmm(A22, V, W2, beta=1.0, alpha=-1.0, out=A22)
            else:
                cols = update_size - k - _NB
                W = Wbuf[:, :, :cols]
                YT = Auxbuf[:, :mm, :]
                torch.bmm(V, T[:, k // _NB].transpose(1, 2), out=YT)
                torch.bmm(V.transpose(1, 2), A22, out=W)
                torch.baddbmm(A22, YT, W, beta=1.0, alpha=-1.0, out=A22)
        if tail_start < factor_size:
            _geqr2_tail(H, tau, tail_start, factor_size)
        if factor_size < n:
            tau[:, factor_size:].zero_()
            if structure == "clustered":
                H[:, :, factor_size:].zero_()
            elif structure == "nearrank":
                H[:, :, factor_size:].zero_()
                H[:, :256, factor_size:].copy_(torch.triu(H[:, :256, :256]))
        return H, tau

    ii = torch.arange(_NB, device=data.device)
    slS = (ii[:, None] > ii[None, :]).to(torch.float32)
    eyeS = torch.eye(_NB, device=data.device, dtype=torch.float32)
    Vbuf, Tbuf, slL, eyeL = _get_buf(batch, n, NBO, data.device)
    for kb in range(0, n, NBO):
        if kb >= factor_size:
            break
        bw = min(NBO, n - kb)
        nsub = (bw + _NB - 1) // _NB
        for s in range(nsub):
            ki = kb + s * _NB
            panel(mH, mTau, mT, mR, cutlass.Int32(ki))
            ie = kb + bw
            if ki + _NB < ie:
                if n == 512 and batch >= 640:
                    desired_tf32 = kb >= 2 * NBO
                    if desired_tf32 != tf32_state:
                        torch.backends.cuda.matmul.allow_tf32 = desired_tf32
                        tf32_state = desired_tf32
                mm = n - ki
                V = Vbuf[:, :mm, :_NB]
                V[:, :_NB, :].copy_(H[:, ki:ki + _NB, ki:ki + _NB])
                V[:, :_NB, :].mul_(slS).add_(eyeS)
                if mm > _NB:
                    V[:, _NB:mm, :].copy_(H[:, ki + _NB:n, ki:ki + _NB])
                Ain = H[:, ki:n, ki + _NB : ie]
                YT = torch.bmm(V, T[:, ki // _NB].transpose(1, 2))
                W = torch.bmm(V.transpose(1, 2), Ain)
                torch.baddbmm(Ain, YT, W, beta=1.0, alpha=-1.0, out=Ain)
        if kb + bw >= n:
            break
        if n == 512 and batch >= 640:
            desired_tf32 = kb >= 2 * NBO
            if desired_tf32 != tf32_state:
                torch.backends.cuda.matmul.allow_tf32 = desired_tf32
                tf32_state = desired_tf32
        mm = n - kb
        Vblk = Vbuf[:, :mm, :bw]
        Vblk[:, :bw, :].copy_(H[:, kb:kb + bw, kb:kb + bw])
        Vblk[:, :bw, :].mul_(slL[:bw, :bw]).add_(eyeL[:bw, :bw])
        if mm > bw:
            Vblk[:, bw:mm, :].copy_(H[:, kb + bw : n, kb:kb + bw])
        Sb = torch.bmm(Vblk.transpose(1, 2), Vblk)
        Tb = Tbuf[:, :bw, :bw]
        W = min(_NB, bw)
        Tb[:, :W, :W] = T[:, kb // _NB][:, :W, :W]
        while W < bw:
            w = min(_NB, bw - W)
            o = W
            Tb[:, o:o + w, :o].zero_()
            Tb[:, o:o + w, o:o + w] = T[:, (kb + o) // _NB][:, :w, :w]
            Cross = Sb[:, :W, o:o + w]
            Tb[:, :W, o:o + w] = -torch.bmm(torch.bmm(Tb[:, :W, :W], Cross), Tb[:, o:o + w, o:o + w])
            W += w
        Af = H[:, kb:n, kb + bw : n]
        YTb = torch.bmm(Vblk, Tb.transpose(1, 2))
        W = torch.bmm(Vblk.transpose(1, 2), Af)
        torch.baddbmm(Af, YTb, W, beta=1.0, alpha=-1.0, out=Af)
    if factor_size < n:
        tau[:, factor_size:].zero_()
    return H, tau


def _geqr176_padded(data: torch.Tensor):
    padded = torch.nn.functional.pad(data, (0, 16, 0, 16))
    H, tau = _blocked_qr(padded)
    return H[:, :176, :176], tau[:, :176]


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]
    if n == 176:
        return _geqr176_padded(data)
    if n <= 176 and n % _GEQR2_TDIM == 0:
        return geqr2_small(data)
    return _blocked_qr(data)
scrolls · 1084 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