Skip to content
KernelIndex
Search⌘K

submission 844791

arseni_ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

blackwell_qr.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-844791?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
4.08ms
#138 of 515
2026-06-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c07ee29592c7fe93db8b81af770d3780fcbab7dca3c6ae79167535fe34ed715e
license declaredunknown
license concludedunknown
authorsarseni_ivanov
imported2026-08-26

Techniques

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

shared-memory_coop_smem_cache = {}

Kernel source

blackwell_qr.py1101 lines
#!POPCORN leaderboard qr_v2
import cutlass
import cutlass.cute as cute
import torch
import cutlass.torch as cutlass_torch
from cutlass.cute.runtime import from_dlpack
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True

from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm, vector
from cutlass.cute.typing import (
    Float32,
)

_NB = 64                      # panel / block width
NMAX = 4096

_ASM_FAST = r"""
{
    .reg .f32  %asq, %full, %nfull, %beta, %amb, %sc, %rb, %tau, %zero;
    .reg .pred %neg, %active;
    mov.f32         %zero, 0f00000000;
    fma.rn.f32      %asq,  $3, $3, $4;            // alpha^2 + sumsq
    sqrt.approx.f32 %full, %asq;                  // ||x||                    (SFU)
    neg.f32         %nfull, %full;
    setp.lt.f32     %neg,  $3, %zero;             // alpha < 0 ?
    selp.f32        %beta, %full, %nfull, %neg;   // beta = (alpha<0)? +||x|| : -||x||
    sub.f32         %amb,  $3, %beta;             // alpha - beta
    rcp.approx.f32  %sc,   %amb;                  // scale = 1/(alpha-beta)    (SFU)
    rcp.approx.f32  %rb,   %beta;                 // 1/beta                    (SFU)
    mul.f32         %tau,  %amb, %rb;
    neg.f32         %tau,  %tau;                  // tau = (beta-alpha)/beta = -(amb/beta)
    setp.gt.f32     %active, $4, %zero;           // sumsq > 0 ?
    selp.f32        $0, %beta, $3,    %active;    // new_diag = active? beta : alpha
    selp.f32        $1, %sc,   %zero, %active;    // scale    = active? sc   : 0
    selp.f32        $2, %tau,  %zero, %active;    // tau      = active? tau  : 0
}
"""

@dsl_user_op
def householder_solve(alpha, sumsq, *, loc=None, ip=None):
    f32 = Float32.mlir_type
    res = llvm.inline_asm(
        llvm.StructType.get_literal([f32, f32, f32]),
        [alpha.ir_value(), sumsq.ir_value()],
        _ASM_FAST,
        "=f,=f,=f,f,f",
        has_side_effects=False,
        is_align_stack=False,
        asm_dialect=llvm.AsmDialect.AD_ATT,
        loc=loc, ip=ip,
    )
    return tuple(cutlass.Float32(llvm.extractvalue(f32, res, [i], loc=loc, ip=ip)) for i in range(3))

@dsl_user_op
def load4(ptr, loc=None, ip=None):
    """128-bit vectorized load: ptr -> (f0,f1,f2,f3). ptr must be 16B-aligned. The loaded
    vector indexes directly, so no extractelement/wrapping is needed."""
    v4 = ir.VectorType.get([4], Float32.mlir_type, loc=loc)
    vv = cute.arch.load(ptr, v4, loc=loc, ip=ip)
    return vv[0], vv[1], vv[2], vv[3]


@dsl_user_op
def store4(ptr, a, b, c, d, loc=None, ip=None):
    """128-bit vectorized store of 4 Float32 to ptr (must be 16B-aligned)."""
    v4 = ir.VectorType.get([4], Float32.mlir_type, loc=loc)
    vec = vector.from_elements(v4, [a.ir_value(), b.ir_value(), c.ir_value(), d.ir_value()], loc=loc, ip=ip)
    cute.arch.store(ptr, vec, loc=loc, ip=ip)

class PanelQRGmem:
    @cute.jit
    def __call__(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
        self.kernel(mH, mtau, mV, k, b, n).launch(
            grid=[cute.size(mH, mode=[0]), 1, 1], block=[THREADS, 1, 1]
        )

    @cute.jit
    def warp_reduce(self, val):
        for i in range(5):
            val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
        return val

    @cute.kernel
    def kernel(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
        tidx, _, _ = cute.arch.thread_idx()
        bi, _, _ = cute.arch.block_idx()
        warp_id = tidx >> 5
        lane_id = tidx & 31
        m = n - k

        smem = cutlass.utils.SmemAllocator()
        vsm = smem.allocate_tensor(cutlass.Float32, NMAX)
        red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
        wsm = smem.allocate_tensor(cutlass.Float32, _NB)
        scal = smem.allocate_tensor(cutlass.Float32, 4)

        for j in cutlass.range(0, b, 1, unroll=1):
            col_len = m - j
            base = k + j
            ssq = cutlass.Float32(0.0)
            r = tidx
            if r == 0:
                vsm[0] = mH[bi, base, base]
                r += THREADS
            while r < col_len:
                x = mH[bi, base + r, base]
                vsm[r] = x
                ssq += x * x
                r += THREADS
            cute.arch.sync_threads()
            ssq = self.warp_reduce(ssq)
            if lane_id == 0:
                red[warp_id] = ssq
            cute.arch.sync_threads()
            if warp_id == 0:
                v0 = red[lane_id] if lane_id < WARPS else cutlass.Float32(0.0)
                v0 = self.warp_reduce(v0)
                if lane_id == 0:
                    red[WARPS] = v0
            cute.arch.sync_threads()
            sumsq = red[WARPS]

            if tidx == 0:
                alpha = vsm[0]
                new_diag, scale, tau_j = householder_solve(alpha, sumsq)
                scal[0] = new_diag
                scal[1] = tau_j
                scal[2] = scale
                mtau[bi, base] = tau_j
            cute.arch.sync_threads()
            new_diag = scal[0]
            tau_j = scal[1]
            scale = scal[2]

            r = tidx
            if r == 0:
                vsm[0] = cutlass.Float32(1.0)
                mH[bi, base, base] = new_diag
                mV[bi, j, j] = cutlass.Float32(1.0)
                r += THREADS
            while r < col_len:
                vv = vsm[r] * scale
                vsm[r] = vv
                mH[bi, base + r, base] = vv
                mV[bi, j + r, j] = vv
                r += THREADS
            r = tidx
            while r < j:
                mV[bi, r, j] = cutlass.Float32(0.0)
                r += THREADS
            cute.arch.sync_threads()

            ncols = b - 1 - j
            cc = tidx
            while cc < ncols:
                acc = cutlass.Float32(0.0)
                r = 0
                while r < col_len:
                    acc += vsm[r] * mH[bi, base + r, base + 1 + cc]
                    r += 1
                wsm[cc] = acc
                cc += THREADS
            cute.arch.sync_threads()
            r = tidx
            while r < col_len:
                vr = vsm[r]
                row = base + r
                cc = 0
                while cc < ncols:
                    mH[bi, row, base + 1 + cc] -= tau_j * wsm[cc] * vr
                    cc += 1
                r += THREADS
            cute.arch.sync_threads()

THREADS = 256
WARPS = THREADS // 32
WPCTA = 8                     # warps per CTA for the tiny T-from-G kernel
_PF_MAX = 57000               # max panel-cache floats (B200 ~232KB smem)

class PanelQRv2:
    def __init__(self, panel_floats, nb, use_tpc=True, wide_red=False, vec=False, fast_update=False, fuse_update=False, use_regblock=False):
        self.PF = panel_floats
        self.NB = nb            # actual panel width -> right-size wacc / wred / wsm / scales
        self.use_tpc = use_tpc
        self.wide_red = wide_red   # TPC reduce: True -> 1 warp/col (conflict-free, n<=512); False -> 8 thr/col
        self.vec = vec             # 128-bit vectorized gmem staging/writeback (high-occupancy shapes
                                    # only; at low batch the latency-bound panel sees vector-op overhead)
        self.fast_update = fast_update  # no-B3 per-warp column trailing update for small n where it wins
        self.fuse_update = fuse_update  # fused reduction+update: no wsm round-trip, no B3, wval in register
        self.use_regblock = use_regblock  # register-block 4 columns in fused tpc32 (1 row-pass vs 4)
    @cute.jit
    def __call__(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
        self.kernel(mH, mtau, mV, k, b, n).launch(
            grid=[cute.size(mH, mode=[0]), 1, 1], block=[THREADS, 1, 1]
        )

    @cute.jit
    def warp_reduce(self, val):
        for i in range(5):
            val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
        return val

    @cute.jit
    def _reduce_tpc8(self, sp, wsm, j, m, ncols, sbp, scale, tidx):
        g = tidx // 8
        gl = tidx % 8
        cc = g
        acc = cutlass.Float32(0.0)
        if cc < ncols:
            r = j + 1 + gl
            while r < m:
                acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
                r += 8
        for o in (1, 2, 4):
            acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
        if cc < ncols:
            if gl == 0:
                wsm[cc] = sp[j * sbp + j + 1 + cc] + scale * acc

    @cute.jit
    def _reduce_tpc32(self, sp, wsm, j, m, ncols, sbp, scale, tidx):
        g = tidx // 32
        gl = tidx % 32
        cc = g
        while cc < ncols:
            acc = cutlass.Float32(0.0)
            r = j + 1 + gl
            while r < m:
                acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
                r += 32
            for i in range(5):             # butterfly offsets 1,2,4,8,16
                acc += cute.arch.shuffle_sync_bfly(acc, offset=1 << i)
            if gl == 0:
                wsm[cc] = sp[j * sbp + j + 1 + cc] + scale * acc
            cc += WARPS

    @cute.jit
    def _reduce_update_tpc32(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
        g = tidx // 32
        gl = tidx % 32
        cc = g
        tsc = tau_j * scale
        while cc < ncols:
            acc = cutlass.Float32(0.0)
            r = j + 1 + gl
            while r < m:
                acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
                r += 32
            for i in range(5):
                acc += cute.arch.shuffle_sync_bfly(acc, offset=1 << i)
            wval = sp[j * sbp + j + 1 + cc] + scale * acc
            if gl == 0:
                sp[j * sbp + j + 1 + cc] -= tau_j * wval
            r = j + 1 + gl
            while r < m:
                tvr = tsc * sp[r * sbp + j]
                sp[r * sbp + j + 1 + cc] -= tvr * wval
                r += 32
            cc += WARPS

    @cute.jit
    def _reduce_update_tpc32_regblock(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
        g = tidx // 32
        gl = tidx % 32
        tsc = tau_j * scale
        cc = g
        while cc + 24 < ncols:
            cc0 = cc
            cc1 = cc + 8
            cc2 = cc + 16
            cc3 = cc + 24
            # --- accumulate 4 columns in one row pass ---
            acc0 = cutlass.Float32(0.0)
            acc1 = cutlass.Float32(0.0)
            acc2 = cutlass.Float32(0.0)
            acc3 = cutlass.Float32(0.0)
            r = j + 1 + gl
            while r < m:
                vv = sp[r * sbp + j]
                acc0 += vv * sp[r * sbp + j + 1 + cc0]
                acc1 += vv * sp[r * sbp + j + 1 + cc1]
                acc2 += vv * sp[r * sbp + j + 1 + cc2]
                acc3 += vv * sp[r * sbp + j + 1 + cc3]
                r += 32
            for i in range(5):
                acc0 += cute.arch.shuffle_sync_bfly(acc0, offset=1 << i)
                acc1 += cute.arch.shuffle_sync_bfly(acc1, offset=1 << i)
                acc2 += cute.arch.shuffle_sync_bfly(acc2, offset=1 << i)
                acc3 += cute.arch.shuffle_sync_bfly(acc3, offset=1 << i)
            wval0 = sp[j * sbp + j + 1 + cc0] + scale * acc0
            wval1 = sp[j * sbp + j + 1 + cc1] + scale * acc1
            wval2 = sp[j * sbp + j + 1 + cc2] + scale * acc2
            wval3 = sp[j * sbp + j + 1 + cc3] + scale * acc3
            if gl == 0:
                sp[j * sbp + j + 1 + cc0] -= tau_j * wval0
                sp[j * sbp + j + 1 + cc1] -= tau_j * wval1
                sp[j * sbp + j + 1 + cc2] -= tau_j * wval2
                sp[j * sbp + j + 1 + cc3] -= tau_j * wval3
            r = j + 1 + gl
            while r < m:
                tvr = tsc * sp[r * sbp + j]
                sp[r * sbp + j + 1 + cc0] -= tvr * wval0
                sp[r * sbp + j + 1 + cc1] -= tvr * wval1
                sp[r * sbp + j + 1 + cc2] -= tvr * wval2
                sp[r * sbp + j + 1 + cc3] -= tvr * wval3
                r += 32
            cc += 32
        # Tail: remaining 0-3 columns, single-column fallback
        while cc < ncols:
            acc = cutlass.Float32(0.0)
            r = j + 1 + gl
            while r < m:
                acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
                r += 32
            for o in (1, 2, 4, 8, 16):
                acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
            wval = sp[j * sbp + j + 1 + cc] + scale * acc
            if gl == 0:
                sp[j * sbp + j + 1 + cc] -= tau_j * wval
            r = j + 1 + gl
            while r < m:
                tvr = tsc * sp[r * sbp + j]
                sp[r * sbp + j + 1 + cc] -= tvr * wval
                r += 32
            cc += 8

    @cute.jit
    def _reduce_update_tpc8(self, sp, j, m, ncols, sbp, scale, tau_j, tidx):
        g = tidx // 8
        gl = tidx % 8
        cc = g
        tsc = tau_j * scale
        acc = cutlass.Float32(0.0)
        if cc < ncols:
            r = j + 1 + gl
            while r < m:
                acc += sp[r * sbp + j] * sp[r * sbp + j + 1 + cc]
                r += 8
        for o in (1, 2, 4):
            acc += cute.arch.shuffle_sync_bfly(acc, offset=o)
        if cc < ncols:
            wval = sp[j * sbp + j + 1 + cc] + scale * acc
            if gl == 0:
                sp[j * sbp + j + 1 + cc] -= tau_j * wval
            r = j + 1 + gl
            while r < m:
                tvr = tsc * sp[r * sbp + j]
                sp[r * sbp + j + 1 + cc] -= tvr * wval
                r += 8

    @cute.jit
    def _reduce_wr(self, sp, wsm, wred, j, m, ncols, sbp, scale, tidx, warp_id, lane_id):
        wacc = cute.make_rmem_tensor(self.NB, cutlass.Float32)
        r = tidx
        while r < m:
            if r >= j:
                vr = cutlass.Float32(1.0) if r == j else scale * sp[r * sbp + j]   # scale inline
                base = r * sbp + j + 1
                cc = 0
                while cc < ncols:
                    wacc[cc] += vr * sp[base + cc]
                    cc += 1
            r += THREADS
        cc = 0
        while cc < ncols:
            wv = self.warp_reduce(wacc[cc])
            if lane_id == 0:
                wred[warp_id * self.NB + cc] = wv
            cc += 1
        if warp_id == 0:
            cc = lane_id
            while cc < ncols:
                acc = cutlass.Float32(0.0)
                w = 0
                while w < WARPS:
                    acc += wred[w * self.NB + cc]
                    w += 1
                wsm[cc] = acc
                cc += 32

    @cute.kernel
    def kernel(self, mH, mtau, mV, k: cutlass.Int32, b: cutlass.Int32, n: cutlass.Int32):
        tidx, _, _ = cute.arch.thread_idx()
        bi, _, _ = cute.arch.block_idx()
        warp_id = tidx >> 5
        lane_id = tidx & 31
        m = n - k
        sbp = b + 1

        smem = cutlass.utils.SmemAllocator()
        sp = smem.allocate_tensor(cutlass.Float32, self.PF)
        red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
        wsm = smem.allocate_tensor(cutlass.Float32, self.NB)
        scales = smem.allocate_tensor(cutlass.Float32, self.NB)   # per-column Householder scale
        if cutlass.const_expr(not self.use_tpc):
            wred = smem.allocate_tensor(cutlass.Float32, WARPS * self.NB)

        tot = m * b                        # used by the writeback loop at the end
        b4 = b >> 2
        rem = b - (b4 << 2)                 # leftover cols (0..3) when b not a multiple of 4
        # ---- stage the m x b panel into smem ----
        if cutlass.const_expr(self.vec):
            idx = tidx
            tot4 = m * b4
            while idx < tot4:
                r = idx // b4
                c = (idx - r * b4) << 2
                f0, f1, f2, f3 = load4(cute.domain_offset((bi, k + r, k + c), mH).iterator)
                base = r * sbp + c
                sp[base] = f0
                sp[base + 1] = f1
                sp[base + 2] = f2
                sp[base + 3] = f3
                idx += THREADS
            if rem > 0:
                c0 = b4 << 2
                idx = tidx
                totr = m * rem
                while idx < totr:
                    r = idx // rem
                    c = c0 + (idx - r * rem)
                    sp[r * sbp + c] = mH[bi, k + r, k + c]
                    idx += THREADS
        else:
            idx = tidx
            while idx < tot:
                r = idx // b
                c = idx - r * b
                sp[r * sbp + c] = mH[bi, k + r, k + c]
                idx += THREADS
        cute.arch.sync_threads()

        for j in cutlass.range(0, b, 1, unroll=1):
            # ---- column norm (single barrier) -------------------------
            alpha = sp[j * sbp + j]
            ssq = cutlass.Float32(0.0)
            r = tidx
            while r < m:
                if r > j:
                    x = sp[r * sbp + j]
                    ssq += x * x
                r += THREADS
            ssq = self.warp_reduce(ssq)
            if lane_id == 0:
                red[warp_id] = ssq
            cute.arch.sync_threads()                          # B1
            sumsq = cutlass.Float32(0.0)
            w = 0
            while w < WARPS:
                sumsq += red[w]
                w += 1
            new_diag, scale, tau_j = householder_solve(alpha, sumsq)
            if tidx == 0:
                scales[j] = scale
                sp[j * sbp + j] = new_diag
                mtau[bi, k + j] = tau_j

            # ---- w = v^T C (reduction) and trailing update (fused) -----------------
            ncols = b - 1 - j
            if cutlass.const_expr(self.fuse_update):
                if cutlass.const_expr(self.use_regblock):
                    self._reduce_update_tpc32_regblock(sp, j, m, ncols, sbp, scale, tau_j, tidx)
                elif cutlass.const_expr(self.wide_red):
                    self._reduce_update_tpc32(sp, j, m, ncols, sbp, scale, tau_j, tidx)
                else:
                    self._reduce_update_tpc8(sp, j, m, ncols, sbp, scale, tau_j, tidx)
            elif cutlass.const_expr(self.use_tpc):
                if cutlass.const_expr(self.wide_red):
                    self._reduce_tpc32(sp, wsm, j, m, ncols, sbp, scale, tidx)
                else:
                    self._reduce_tpc8(sp, wsm, j, m, ncols, sbp, scale, tidx)
                if cutlass.const_expr(self.fast_update):
                    g = warp_id
                    gl = lane_id
                    tsc = tau_j * scale
                    col = g
                    while col < ncols:
                        wval = wsm[col]
                        if gl == 0:
                            sp[j * sbp + j + 1 + col] -= tau_j * wval
                        r = j + 1 + gl
                        while r < m:
                            tvr = tsc * sp[r * sbp + j]
                            sp[r * sbp + j + 1 + col] -= tvr * wval
                            r += 32
                        col += WARPS
                else:
                    cute.arch.sync_threads()                      # B3
                    tsc = tau_j * scale
                    r = tidx
                    while r < m:
                        if r >= j:
                            tvr = tau_j if r == j else tsc * sp[r * sbp + j]
                            base = r * sbp + j + 1
                            cc = 0
                            while cc + 3 < ncols:
                                w0 = wsm[cc]; w1 = wsm[cc+1]; w2 = wsm[cc+2]; w3 = wsm[cc+3]
                                sp[base + cc]     -= tvr * w0
                                sp[base + cc + 1] -= tvr * w1
                                sp[base + cc + 2] -= tvr * w2
                                sp[base + cc + 3] -= tvr * w3
                                cc += 4
                            while cc < ncols:
                                sp[base + cc] -= tvr * wsm[cc]
                                cc += 1
                        r += THREADS
            else:
                self._reduce_wr(sp, wsm, wred, j, m, ncols, sbp, scale, tidx, warp_id, lane_id)
                cute.arch.sync_threads()                          # B3
                tsc = tau_j * scale
                r = tidx
                while r < m:
                    if r >= j:
                        tvr = tau_j if r == j else tsc * sp[r * sbp + j]
                        base = r * sbp + j + 1
                        cc = 0
                        while cc + 3 < ncols:
                            w0 = wsm[cc]; w1 = wsm[cc+1]; w2 = wsm[cc+2]; w3 = wsm[cc+3]
                            sp[base + cc]     -= tvr * w0
                            sp[base + cc + 1] -= tvr * w1
                            sp[base + cc + 2] -= tvr * w2
                            sp[base + cc + 3] -= tvr * w3
                            cc += 4
                        while cc < ncols:
                            sp[base + cc] -= tvr * wsm[cc]
                            cc += 1
                    r += THREADS
            cute.arch.sync_threads()                              # B4

        # ---- writeback: apply the deferred Householder scale to the stored v's ----
        if cutlass.const_expr(self.vec):
            nbr = m - b
            if nbr > 0:
                tot4 = nbr * b4
                idx = tidx
                while idx < tot4:
                    rr = idx // b4
                    r = b + rr
                    c = (idx - rr * b4) << 2
                    base = r * sbp + c
                    v0 = sp[base] * scales[c]
                    v1 = sp[base + 1] * scales[c + 1]
                    v2 = sp[base + 2] * scales[c + 2]
                    v3 = sp[base + 3] * scales[c + 3]
                    store4(cute.domain_offset((bi, k + r, k + c), mH).iterator, v0, v1, v2, v3)
                    store4(cute.domain_offset((bi, r, c), mV).iterator, v0, v1, v2, v3)
                    idx += THREADS
                if rem > 0:
                    c0 = b4 << 2
                    idx = tidx
                    totr = nbr * rem
                    while idx < totr:
                        rr = idx // rem
                        r = b + rr
                        c = c0 + (idx - rr * rem)
                        val = sp[r * sbp + c] * scales[c]
                        mH[bi, k + r, k + c] = val
                        mV[bi, r, c] = val
                        idx += THREADS
            # top b x b block (r < b): diagonal / above-diagonal -> scalar branchy path
            idx = tidx
            while idx < b * b:
                r = idx // b
                c = idx - r * b
                val = sp[r * sbp + c]
                if r > c:
                    mH[bi, k + r, k + c] = val * scales[c]
                    mV[bi, r, c] = val * scales[c]
                elif r == c:
                    mH[bi, k + r, k + c] = val                # beta (R diagonal)
                    mV[bi, r, c] = cutlass.Float32(1.0)
                else:
                    mH[bi, k + r, k + c] = val                # R (above diagonal)
                    mV[bi, r, c] = cutlass.Float32(0.0)
                idx += THREADS
        else:
            idx = tidx
            while idx < tot:
                r = idx // b
                c = idx - r * b
                val = sp[r * sbp + c]
                if r > c:
                    val = val * scales[c]
                    mH[bi, k + r, k + c] = val
                    mV[bi, r, c] = val
                elif r == c:
                    mH[bi, k + r, k + c] = val                # beta (R diagonal)
                    mV[bi, r, c] = cutlass.Float32(1.0)
                else:
                    mH[bi, k + r, k + c] = val                # R (above diagonal)
                    mV[bi, r, c] = cutlass.Float32(0.0)
                idx += THREADS

class PanelQRCoopSmem:
    def __init__(self, cpm, pf, nb):
        self.CPM = cpm
        self.PF = pf
        self.NB = nb

    @cute.jit
    def __call__(self, mH, mtau, mV, mpart, mbar, k: cutlass.Int32, b: cutlass.Int32,
                 n: cutlass.Int32, base: cutlass.Int32):
        batch = cute.size(mH, mode=[0])
        grid = batch * self.CPM
        self.kernel(mH, mtau, mV, mpart, mbar, k, b, n, base, batch).launch(
            grid=[grid, 1, 1], block=[THREADS, 1, 1], cooperative=True
        )

    @cute.jit
    def warp_reduce(self, val):
        for i in range(5):
            val = val + cute.arch.shuffle_sync_bfly(val, offset=1 << i)
        return val

    @cute.jit
    def cta_reduce(self, val, red, warp_id, lane_id):
        val = self.warp_reduce(val)
        if lane_id == 0:
            red[warp_id] = val
        cute.arch.sync_threads()
        if warp_id == 0:
            v0 = red[lane_id] if lane_id < WARPS else cutlass.Float32(0.0)
            v0 = self.warp_reduce(v0)
            if lane_id == 0:
                red[WARPS] = v0
        cute.arch.sync_threads()
        return red[WARPS]

    @cute.jit
    def gbar(self, mbar, slot, total, tidx):
        cute.arch.sync_threads()
        cute.arch.fence_acq_rel_gpu()
        if tidx == 0:
            cute.arch.atomic_add(mbar.iterator + slot, cutlass.Int32(1), sem="release", scope="gpu")
            done = cutlass.Int32(0)
            while done == 0:
                # acquire LOAD (not atom.add 0): a read doesn't serialize on the L2 atomic unit.
                cur = cute.arch.load(mbar.iterator + slot, cutlass.Int32, sem="acquire", scope="gpu")
                if cur >= total:
                    done = cutlass.Int32(1)
        cute.arch.sync_threads()
        cute.arch.fence_acq_rel_gpu()

    @cute.kernel
    def kernel(self, mH, mtau, mV, mpart, mbar, k: cutlass.Int32, b: cutlass.Int32,
               n: cutlass.Int32, base: cutlass.Int32, batch: cutlass.Int32):
        tidx, _, _ = cute.arch.thread_idx()
        g, _, _ = cute.arch.block_idx()
        bi = g // self.CPM
        sub = g - bi * self.CPM
        warp_id = tidx >> 5
        lane_id = tidx & 31
        m = n - k
        sbp = b + 1
        total = batch * self.CPM
        rpc = (m + self.CPM - 1) // self.CPM
        r0 = sub * rpc
        rend = r0 + rpc
        if rend > m:
            rend = m
        nloc = rend - r0
        if nloc < 0:
            nloc = cutlass.Int32(0)

        smem = cutlass.utils.SmemAllocator()
        sp = smem.allocate_tensor(cutlass.Float32, self.PF)
        red = smem.allocate_tensor(cutlass.Float32, WARPS + 1)
        wred = smem.allocate_tensor(cutlass.Float32, WARPS * self.NB)
        wsm = smem.allocate_tensor(cutlass.Float32, self.NB)

        tot = nloc * b
        idx = tidx
        while idx < tot:
            rr = idx // b
            c = idx - rr * b
            sp[rr * sbp + c] = mH[bi, k + r0 + rr, k + c]
            idx += THREADS
        cute.arch.sync_threads()

        for j in cutlass.range(0, b, 1, unroll=1):
            owner = j // rpc
            ncols = b - 1 - j

            # ---- local norm AND w_raw (off-diagonal, unscaled v_raw) ----
            ssq = cutlass.Float32(0.0)
            wacc = cute.make_fragment(self.NB, cutlass.Float32)
            for ci in range(self.NB):
                wacc[ci] = cutlass.Float32(0.0)

            rr = tidx
            while rr < nloc:
                gr = r0 + rr
                if gr > j:
                    v_raw = sp[rr * sbp + j]
                    ssq += v_raw * v_raw
                    bcol = rr * sbp + j + 1
                    cc = 0
                    while cc < ncols:
                        wacc[cc] += v_raw * sp[bcol + cc]
                        cc += 1
                rr += THREADS
            ssq = self.cta_reduce(ssq, red, warp_id, lane_id)
            # w_raw warp-reduce + cross-warp combine
            for ci in range(self.NB):
                wv = self.warp_reduce(wacc[ci])
                if lane_id == 0:
                    wred[warp_id * self.NB + ci] = wv
            cute.arch.sync_threads()
            if warp_id == 0:
                ci = lane_id
                while ci < self.NB:
                    acc = cutlass.Float32(0.0)
                    w = 0
                    while w < WARPS:
                        acc += wred[w * self.NB + ci]
                        w += 1
                    if ci < ncols:
                        mpart[bi, sub, 1 + ci] = acc
                    ci += 32
            cute.arch.sync_threads()

            # ---- write to mpart and single gbar ----
            if tidx == 0:
                mpart[bi, sub, 0] = ssq
                if sub == owner:
                    mpart[bi, sub, self.NB + 1] = sp[(j - r0) * sbp + j]
                    # store C[j, j+1+cc] for later w combination
                    ci = 0
                    while ci < ncols:
                        mpart[bi, sub, self.NB + 2 + ci] = sp[(j - r0) * sbp + j + 1 + ci]
                        ci += 1
            self.gbar(mbar, base + j, total, tidx)      # single barrier (was 2)

            # ---- after barrier: combine ssq + w_raw, compute tau/scale/w, update ----
            sumsq = cutlass.Float32(0.0)
            s = 0
            while s < self.CPM:
                sumsq += mpart[bi, s, 0]
                s += 1
            alpha = mpart[bi, owner, self.NB + 1]
            new_diag, scale, tau_j = householder_solve(alpha, sumsq)

            cc = tidx
            while cc < ncols:
                w_raw_global = cutlass.Float32(0.0)
                s = 0
                while s < self.CPM:
                    w_raw_global += mpart[bi, s, 1 + cc]
                    s += 1
                c_diag = mpart[bi, owner, self.NB + 2 + cc]
                wsm[cc] = c_diag + scale * w_raw_global
                cc += THREADS
            cute.arch.sync_threads()

            # ---- scale v in-place ----
            rr = tidx
            while rr < nloc:
                gr = r0 + rr
                if gr > j:
                    sp[rr * sbp + j] = sp[rr * sbp + j] * scale
                elif gr == j:
                    sp[rr * sbp + j] = new_diag
                rr += THREADS
            if sub == 0 and tidx == 0:
                mtau[bi, k + j] = tau_j
            cute.arch.sync_threads()

            # ---- trailing update ----
            rr = tidx
            while rr < nloc:
                gr = r0 + rr
                if gr >= j:
                    vr = cutlass.Float32(1.0) if gr == j else sp[rr * sbp + j]
                    tvr = tau_j * vr
                    bcol = rr * sbp + j + 1
                    cc = 0
                    while cc < ncols:
                        sp[bcol + cc] -= tvr * wsm[cc]
                        cc += 1
                rr += THREADS
            cute.arch.sync_threads()

        idx = tidx
        while idx < tot:
            rr = idx // b
            c = idx - rr * b
            gr = r0 + rr
            val = sp[rr * sbp + c]
            if gr > c:
                mH[bi, k + gr, k + c] = val
                mV[bi, gr, c] = val
            elif gr == c:
                mH[bi, k + gr, k + c] = val
                mV[bi, gr, c] = cutlass.Float32(1.0)
            else:
                mH[bi, k + gr, k + c] = val
                mV[bi, gr, c] = cutlass.Float32(0.0)
            idx += THREADS

class TFromG:
    def __init__(self, b):
        self.B = b

    @cute.jit
    def __call__(self, mG, mtau, mT, k: cutlass.Int32, b: cutlass.Int32):
        batch = cute.size(mG, mode=[0])
        grid = (batch + WPCTA - 1) // WPCTA
        self.kernel(mG, mtau, mT, k, b, batch).launch(grid=[grid, 1, 1], block=[WPCTA * 32, 1, 1])

    @cute.kernel
    def kernel(self, mG, mtau, mT,
            k: cutlass.Int32,
            b: cutlass.Int32,
            batch: cutlass.Int32):

        B = self.B
        sbp = B + 1

        tidx, _, _ = cute.arch.thread_idx()
        bidx, _, _ = cute.arch.block_idx()

        warp = tidx >> 5
        lane = tidx & 31
        bi = bidx * WPCTA + warp

        smem = cutlass.utils.SmemAllocator()

        gsm = smem.allocate_tensor(cutlass.Float32, WPCTA * B * sbp)
        tsm = smem.allocate_tensor(cutlass.Float32, WPCTA * B * sbp)
        taus = smem.allocate_tensor(cutlass.Float32, WPCTA * B)

        if bi < batch:
            base = warp * B * sbp
            tau_base = warp * B
            #
            # Stage tau once
            #
            if lane < b:
                taus[tau_base + lane] = mtau[bi, k + lane]
            #
            # Vectorized stage of G
            #
            b4 = b >> 2
            rem = b - (b4 << 2)
            idx = lane
            tot4 = b * b4

            while idx < tot4:
                r = idx // b4
                c = (idx - r * b4) << 2
                f0, f1, f2, f3 = load4(
                    cute.domain_offset((bi, r, c), mG).iterator
                )
                row = base + r * sbp + c

                gsm[row]     = f0
                gsm[row + 1] = f1
                gsm[row + 2] = f2
                gsm[row + 3] = f3

                tsm[row]     = cutlass.Float32(0.0)
                tsm[row + 1] = cutlass.Float32(0.0)
                tsm[row + 2] = cutlass.Float32(0.0)
                tsm[row + 3] = cutlass.Float32(0.0)

                idx += 32

            if rem > 0:
                c0 = b4 << 2
                idx = lane
                totr = b * rem
                while idx < totr:
                    r = idx // rem
                    c = c0 + (idx - r * rem)
                    row = base + r * sbp + c
                    gsm[row] = mG[bi, r, c]
                    tsm[row] = cutlass.Float32(0.0)
                    idx += 32

            if lane == 0:
                tsm[base] = taus[tau_base]

            cute.arch.sync_warp()

            #
            # WY recurrence
            #
            i = 1

            while i < b:
                tau_i = taus[tau_base + i]
                r = lane

                while r < i:
                    grow = base + r * sbp
                    trow = grow
                    tau_r = taus[tau_base + r]
                    acc = tau_r * gsm[grow + i]
                    c = r + 1

                    while c + 3 < i:
                        acc += tsm[trow + c]     * gsm[base + (c    ) * sbp + i]
                        acc += tsm[trow + c + 1] * gsm[base + (c + 1) * sbp + i]
                        acc += tsm[trow + c + 2] * gsm[base + (c + 2) * sbp + i]
                        acc += tsm[trow + c + 3] * gsm[base + (c + 3) * sbp + i]
                        c += 4

                    while c < i:
                        acc += tsm[trow + c] * gsm[base + c * sbp + i]
                        c += 1

                    tsm[trow + i] = -tau_i * acc
                    r += 32

                if lane == 0:
                    tsm[base + i * sbp + i] = tau_i
                i += 1

            cute.arch.sync_warp()

            #
            # Vectorized writeback
            #
            idx = lane

            while idx < tot4:
                r = idx // b4
                c = (idx - r * b4) << 2
                row = base + r * sbp + c

                store4(
                    cute.domain_offset((bi, r, c), mT).iterator,
                    tsm[row],
                    tsm[row + 1],
                    tsm[row + 2],
                    tsm[row + 3],
                )
                idx += 32

            if rem > 0:
                c0 = b4 << 2
                idx = lane
                totr = b * rem
                while idx < totr:
                    r = idx // rem
                    c = c0 + (idx - r * rem)
                    mT[bi, r, c] = tsm[base + r * sbp + c]
                    idx += 32

def _make_cute(t):
    """Wrap a (M, K, L) / (M, N, L) batch-last torch tensor as an fp32 cute tensor."""
    ct = from_dlpack(t, assumed_align=16)
    ct.element_type = cutlass.Float32
    ct = ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
    return ct

def _make_cute_i(t):
    ct = from_dlpack(t, assumed_align=16)
    ct.element_type = cutlass.Int32
    ct = ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
    return ct

_panel_cache = {}

def _panel(cH, ctau, cV, k, b, n, pf, nb, use_tpc=True, wide_red=False, vec=False, fast_update=False, fuse_update=False, use_regblock=False):
    ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n)
    key = (pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock)
    fn = _panel_cache.get(key)
    if fn is None:
        fn = cute.compile(PanelQRv2(pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock), cH, ctau, cV, ck, cb, cn)
        _panel_cache[key] = fn
    fn(cH, ctau, cV, ck, cb, cn)


_coop_smem_cache = {}

def _panel_coop_smem(H, tau, V, part, bar, cpm, pf, nb, k, b, n, base):
    ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n); cbase = cutlass.Int32(base)
    key = (cpm, pf, nb)
    fn = _coop_smem_cache.get(key)
    if fn is None:
        fn = cute.compile(PanelQRCoopSmem(cpm, pf, nb), _make_cute(H), _make_cute(tau), _make_cute(V),
                          _make_cute(part), _make_cute_i(bar), ck, cb, cn, cbase)
        _coop_smem_cache[key] = fn
    fn(_make_cute(H), _make_cute(tau), _make_cute(V), _make_cute(part), _make_cute_i(bar),
       ck, cb, cn, cbase)

_COOP_NB = 32

def _coop_config(n):
    if n not in (2048, 4096):
        return None
    rpc = 128 if n == 2048 else 256
    nb = _COOP_NB
    cpm = max(1, (n + rpc - 1) // rpc)
    if rpc * (nb + 1) > _PF_MAX or rpc < nb:
        return None
    return cpm, nb

_tfromg_cache = {}

def _tfromg(cG, ctau, cT, k, b):
    ck = cutlass.Int32(k); cb = cutlass.Int32(b)
    fn = _tfromg_cache.get(b)          # G is now b×b (varying layout) -> key per b; b -> smem size
    if fn is None:
        fn = cute.compile(TFromG(b), cG, ctau, cT, ck, cb)
        _tfromg_cache[b] = fn
    fn(cG, ctau, cT, ck, cb)

_panel_gmem_cache = {}

def _panel_gmem(mH, mtau, mV, k, b, n):
    ck = cutlass.Int32(k); cb = cutlass.Int32(b); cn = cutlass.Int32(n)
    key = (mH.shape[0], n)
    fn = _panel_gmem_cache.get(key)
    if fn is None:
        fn = cute.compile(
            PanelQRGmem(), _make_cute(mH), _make_cute(mtau), _make_cute(mV), ck, cb, cn
        )
        _panel_gmem_cache[key] = fn
    fn(_make_cute(mH), _make_cute(mtau), _make_cute(mV), ck, cb, cn)

def _choose_nb(n):
    for nb in (32, 16, 8):
        if nb <= _NB and n * (nb + 1) <= _PF_MAX:
            return nb
    return 4

def custom_kernel(data: input_t) -> output_t:
    A = data
    batch, n, _ = A.shape

    #---Set up branches and flags for various shapes---

    # Large low-batch shapes need different CTA split
    coop_cfg = _coop_config(n)
    # Always use SMEM when size allows
    cached = (n * (16 + 1) <= _PF_MAX)
    # Small panels do differently depending on amount of threads participating in reduction
    use_tpc = n <= 2048
    wide_red = (n <= 512) or (n == 2048)
    # Large batches benefit from vectorized loads, small don't somehow
    vec = batch >= 48
    # Use SMEM to do local rank-1 updates without sync with other warps, works for small shapes
    fast_update = n <= 352
    # Fused reduction+update: wval in RMEM
    fuse_update = wide_red
    use_regblock = fuse_update  # all fused-tpc32 shapes use register blocking

    if coop_cfg is not None:
        coop_cpm, nb = coop_cfg
        coop_pf = ((n + coop_cpm - 1) // coop_cpm) * (nb + 1)
    elif cached:
        nb = _choose_nb(n)
        pf = n * (nb + 1)
    else:
        nb = 32
        pf = 0

    #---Prep tensors---

    H = A.clone()
    tau = torch.empty(batch, n, device=A.device, dtype=torch.float32)  # panel writes all of tau
    Vbuf = torch.zeros(batch, n, _NB, device=A.device, dtype=torch.float32)
    Tbuf = torch.zeros(batch, _NB, _NB, device=A.device, dtype=torch.float32)
    if coop_cfg is not None:
        coop_part = torch.zeros(batch, coop_cpm, 2 * _NB + 2, device=A.device, dtype=torch.float32)
        coop_bar = torch.zeros(n + 1, device=A.device, dtype=torch.int32)
    torch.backends.cuda.matmul.allow_tf32 = False

    #---Compute loop---
    k = 0
    while k < n:
        b = min(nb, n - k)
        m = n - k
        if coop_cfg is not None:
            _panel_coop_smem(H, tau, Vbuf, coop_part, coop_bar, coop_cpm, coop_pf, nb,
                             k, b, n, k)
        elif cached:
            _panel(H, tau, Vbuf, k, b, n, pf, nb, use_tpc, wide_red, vec, fast_update, fuse_update, use_regblock)
        else:
            _panel_gmem(H, tau, Vbuf, k, b, n)
        if k + b < n:
            V = Vbuf[:, :m, :b]
            G = torch.matmul(V.transpose(-1, -2), V)         # (l,b,b) FP32
            _tfromg(G, tau, Tbuf[:, :b, :b], k, b)
            C0 = H[:, k:, k + b:]
            T = Tbuf[:, :b, :b]
            torch.backends.cuda.matmul.allow_tf32 = True
            W = torch.matmul(V.transpose(-1, -2), C0)        # (l,b,rest) — W only needs TF32x1
            torch.backends.cuda.matmul.allow_tf32 = False
            Z = torch.matmul(T.transpose(-1, -2), W)         # (l,b,rest) FP32
            C0.baddbmm_(V, Z, beta=1, alpha=-1)

        k += b #Move to next panel
    torch.backends.cuda.matmul.allow_tf32 = True
    return H, tau
scrolls · 1101 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