Skip to content
KernelIndex
Search⌘K

submission 804799

Simon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

foo.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-804799?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
5.11ms
#180 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:39506c8a355c44f312e251077c712ce46404c1849f4a62242f31bd872d2a7d4b
license declaredunknown
license concludedunknown
authorsSimon
imported2026-08-26

Techniques

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

shared-memoryself.smem_bytes = (self.n * self.n + 2 * work_elems + 3) * 4

Kernel source

foo.py1277 lines
import torch

import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32
from cutlass.cute.runtime import make_ptr
import operator

from task import input_t, output_t


_compile_cache = {}
_ENABLE_TORCH_WY512_MACRO_UPDATE = True
_ENABLE_TORCH_WY512_SMALL_UPDATE = True
_ENABLE_TORCH_WY1024_UPDATE = True
_ENABLE_TF32_WY1024_MACRO_UPDATE = True
_ENABLE_TF32_WY1024_SMALL_UPDATE = True
_ENABLE_TF32_WY1024_TAIL_UPDATE = True
_ENABLE_TF32_WY1024_COMPOSE = True
_ENABLE_TF32_WY2048_MACRO_UPDATE = True
_ENABLE_TF32_WY2048_SMALL_UPDATE = True
_ENABLE_TF32_WY2048_COMPOSE = True
_TF32_WY2048_MACRO_STOP = 2048
_ENABLE_COALESCED_PANEL_EMIT = True
_ENABLE_PANEL352_512_THREADS = True
_ENABLE_PANEL1024_512_THREADS = True
_ENABLE_PANEL1024_NB16 = True
_WY_UPDATE_CN_176 = 8
_WY_UPDATE_CN_352 = 16


class _QR32Kernel:
    def __init__(self):
        self.n = 32
        self.num_threads = 512
        work_elems = max(self.n, self.num_threads)
        self.smem_bytes = (self.n * self.n + 2 * work_elems + 3) * 4

    @cute.jit
    def __call__(
        self,
        a_ptr: cute.Pointer,
        h_ptr: cute.Pointer,
        tau_ptr: cute.Pointer,
        batch: Int32,
    ):
        n = self.n
        mA = cute.make_tensor(
            a_ptr,
            cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
        )
        mH = cute.make_tensor(
            h_ptr,
            cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
        )
        mTau = cute.make_tensor(
            tau_ptr,
            cute.make_layout((batch, n), stride=(n, 1)),
        )
        self.kernel(mA, mH, mTau).launch(
            grid=[batch, 1, 1],
            block=[self.num_threads, 1, 1],
            smem=self.smem_bytes,
        )

    @cute.kernel
    def kernel(self, mA: cute.Tensor, mH: cute.Tensor, mTau: cute.Tensor):
        tidx, _, _ = cute.arch.thread_idx()
        bidx, _, _ = cute.arch.block_idx()
        n = self.n
        threads = self.num_threads
        work_elems = max(n, threads)
        warp_id = tidx // 32
        lane = tidx - warp_id * 32

        smem = cutlass.utils.SmemAllocator()
        mat = smem.allocate_tensor(
            Float32,
            cute.make_layout((n * n,), stride=(1,)),
            byte_alignment=16,
        )
        work = smem.allocate_tensor(
            Float32,
            cute.make_layout((work_elems,), stride=(1,)),
            byte_alignment=16,
        )
        scaled_sumsq = smem.allocate_tensor(
            Float32,
            cute.make_layout((work_elems,), stride=(1,)),
            byte_alignment=16,
        )
        params = smem.allocate_tensor(
            Float32,
            cute.make_layout((3,), stride=(1,)),
            byte_alignment=16,
        )

        for idx in cutlass.range(tidx, n * n, threads):
            row = idx // n
            col = idx - row * n
            mat[idx] = mA[bidx, row, col]
        cute.arch.sync_threads()

        for k in cutlass.range_constexpr(n - 1):
            xnorm2 = Float32(0.0)
            if warp_id == 0:
                row = k + 1 + lane
                if row < n:
                    value = mat[row * n + k]
                    xnorm2 = value * value
                xnorm2 = cute.arch.warp_reduction_sum(xnorm2)

            if tidx == 0:
                alpha = mat[k * n + k]
                tail_norm2 = xnorm2
                norm2 = alpha * alpha + tail_norm2
                params[2] = Float32(0.0)
                if tail_norm2 == Float32(0.0):
                    mTau[bidx, k] = Float32(0.0)
                    params[0] = Float32(0.0)
                    params[1] = alpha
                elif norm2 <= Float32(3.4028234663852886e38):
                    norm = cute.math.sqrt(norm2)
                    beta = -norm
                    if alpha < Float32(0.0):
                        beta = norm
                    tau_value = (beta - alpha) / beta
                    mTau[bidx, k] = tau_value
                    params[0] = Float32(1.0) / (alpha - beta)
                    params[1] = beta
                    params[2] = tau_value
                else:
                    params[2] = Float32(1.0)
            cute.arch.sync_threads()

            if params[2] == Float32(1.0):
                scale = Float32(0.0)
                sumsq = Float32(0.0)
                for row in cutlass.range(k + tidx, n, threads):
                    abs_value = mat[row * n + k]
                    if abs_value < Float32(0.0):
                        abs_value = -abs_value
                    if abs_value != Float32(0.0):
                        if scale < abs_value:
                            ratio = Float32(0.0)
                            if scale != Float32(0.0):
                                ratio = scale / abs_value
                            sumsq = Float32(1.0) + sumsq * ratio * ratio
                            scale = abs_value
                        else:
                            ratio = abs_value / scale
                            sumsq += ratio * ratio
                work[tidx] = scale
                scaled_sumsq[tidx] = sumsq
                cute.arch.sync_threads()

                stride = threads // 2
                while stride > 0:
                    if tidx < stride:
                        other_scale = work[tidx + stride]
                        other_sumsq = scaled_sumsq[tidx + stride]
                        if other_scale != Float32(0.0):
                            if work[tidx] == Float32(0.0):
                                work[tidx] = other_scale
                                scaled_sumsq[tidx] = other_sumsq
                            elif work[tidx] < other_scale:
                                ratio = work[tidx] / other_scale
                                scaled_sumsq[tidx] = (
                                    other_sumsq + scaled_sumsq[tidx] * ratio * ratio
                                )
                                work[tidx] = other_scale
                            else:
                                ratio = other_scale / work[tidx]
                                scaled_sumsq[tidx] = (
                                    scaled_sumsq[tidx] + other_sumsq * ratio * ratio
                                )
                    cute.arch.sync_threads()
                    stride = stride // 2

                if tidx == 0:
                    alpha = mat[k * n + k]
                    norm = work[0] * cute.math.sqrt(scaled_sumsq[0])
                    beta = -norm
                    if alpha < Float32(0.0):
                        beta = norm
                    tau_value = (beta - alpha) / beta
                    mTau[bidx, k] = tau_value
                    params[0] = Float32(1.0) / (alpha - beta)
                    params[1] = beta
                    params[2] = tau_value
                cute.arch.sync_threads()

            inv_alpha_minus_beta = params[0]
            beta = params[1]
            for row in cutlass.range(k + 1 + tidx, n, threads):
                mat[row * n + k] = mat[row * n + k] * inv_alpha_minus_beta
            if tidx == 0:
                mat[k * n + k] = beta
            cute.arch.sync_threads()

            if tidx < n - k - 1:
                col = k + 1 + tidx
                dot = mat[k * n + col]
                for row in cutlass.range(k + 1, n, 1):
                    dot += mat[row * n + k] * mat[row * n + col]
                work[tidx] = dot
            cute.arch.sync_threads()

            tau_value = params[2]
            update_cols = n - k - 1
            for idx in cutlass.range(tidx, (n - k) * update_cols, threads):
                local_row = idx // update_cols
                local_col = idx - local_row * update_cols
                row = k + local_row
                col = k + 1 + local_col
                v = mat[row * n + k]
                if local_row == 0:
                    v = Float32(1.0)
                mat[row * n + col] = (
                    mat[row * n + col] - tau_value * v * work[local_col]
                )
            cute.arch.sync_threads()

        if tidx == 0:
            mTau[bidx, n - 1] = Float32(0.0)

        for idx in cutlass.range(tidx, n * n, threads):
            row = idx // n
            col = idx - row * n
            mH[bidx, row, col] = mat[idx]


def _get_qr32_kernel():
    kernel = _compile_cache.get(32)
    if kernel is None:
        obj = _QR32Kernel()
        kernel = cute.compile(
            obj,
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            Int32(1),
        )
        _compile_cache[32] = kernel
    return kernel


def _qr32(data: torch.Tensor) -> output_t:
    batch = data.shape[0]
    h = torch.empty_like(data)
    tau = torch.empty((batch, 32), device=data.device, dtype=torch.float32)
    kernel = _get_qr32_kernel()
    kernel(
        make_ptr(Float32, data.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        Int32(batch),
    )
    return h, tau


class _PanelWYKernel:
    def __init__(self, n: int, num_threads: int = 256, nb: int = 8):
        self.n = n
        self.nb = nb
        self.num_threads = num_threads
        smem_rows = self.n + 1
        self.smem_bytes = (2 * self.nb * smem_rows + 2 * self.num_threads + 3) * 4

    @cute.jit
    def __call__(
        self,
        h_ptr: cute.Pointer,
        tau_ptr: cute.Pointer,
        v_ptr: cute.Pointer,
        u_ptr: cute.Pointer,
        batch: Int32,
        k: Int32,
    ):
        n = self.n
        nb = self.nb
        mH = cute.make_tensor(
            h_ptr,
            cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
        )
        mTau = cute.make_tensor(
            tau_ptr,
            cute.make_layout((batch, n), stride=(n, 1)),
        )
        mV = cute.make_tensor(
            v_ptr,
            cute.make_layout((batch, n, nb), stride=(n * nb, nb, 1)),
        )
        mU = cute.make_tensor(
            u_ptr,
            cute.make_layout((batch, n, nb), stride=(n * nb, nb, 1)),
        )
        self.kernel(mH, mTau, mV, mU, k).launch(
            grid=[batch, 1, 1],
            block=[self.num_threads, 1, 1],
            smem=self.smem_bytes,
        )

    @cute.kernel
    def kernel(
        self,
        mH: cute.Tensor,
        mTau: cute.Tensor,
        mV: cute.Tensor,
        mU: cute.Tensor,
        k: Int32,
    ):
        tidx, _, _ = cute.arch.thread_idx()
        bidx, _, _ = cute.arch.block_idx()
        n = self.n
        nb = self.nb
        threads = self.num_threads
        smem_rows = n + 1
        m = n - k

        smem = cutlass.utils.SmemAllocator()
        panel = smem.allocate_tensor(
            Float32,
            cute.make_layout((nb * smem_rows,), stride=(1,)),
            byte_alignment=16,
        )
        u_panel = smem.allocate_tensor(
            Float32,
            cute.make_layout((nb * smem_rows,), stride=(1,)),
            byte_alignment=16,
        )
        work = smem.allocate_tensor(
            Float32,
            cute.make_layout((threads,), stride=(1,)),
            byte_alignment=16,
        )
        scaled_sumsq = smem.allocate_tensor(
            Float32,
            cute.make_layout((threads,), stride=(1,)),
            byte_alignment=16,
        )
        params = smem.allocate_tensor(
            Float32,
            cute.make_layout((3,), stride=(1,)),
            byte_alignment=16,
        )

        for local_col in cutlass.range_constexpr(nb):
            col = k + local_col
            for local_row in cutlass.range(tidx, m, threads):
                panel[local_col * smem_rows + local_row] = mH[bidx, k + local_row, col]
        cute.arch.sync_threads()

        for j in cutlass.range_constexpr(nb):
            col = k + j
            xnorm2 = Float32(0.0)
            for local_row in cutlass.range(j + 1 + tidx, m, threads):
                value = panel[j * smem_rows + local_row]
                xnorm2 += value * value
            work[tidx] = xnorm2
            cute.arch.sync_threads()

            stride = threads // 2
            while stride > 0:
                if tidx < stride:
                    work[tidx] = work[tidx] + work[tidx + stride]
                cute.arch.sync_threads()
                stride = stride // 2

            if tidx == 0:
                alpha = panel[j * smem_rows + j]
                tail_norm2 = work[0]
                norm2 = alpha * alpha + tail_norm2
                params[2] = Float32(0.0)
                if tail_norm2 == Float32(0.0):
                    mTau[bidx, col] = Float32(0.0)
                    params[0] = Float32(0.0)
                    params[1] = alpha
                elif norm2 <= Float32(3.4028234663852886e38):
                    norm = cute.math.sqrt(norm2)
                    beta = -norm
                    if alpha < Float32(0.0):
                        beta = norm
                    tau_value = (beta - alpha) / beta
                    mTau[bidx, col] = tau_value
                    params[0] = Float32(1.0) / (alpha - beta)
                    params[1] = beta
                    params[2] = tau_value
                else:
                    params[2] = Float32(1.0)
            cute.arch.sync_threads()

            if params[2] == Float32(1.0):
                scale = Float32(0.0)
                sumsq = Float32(0.0)
                for local_row in cutlass.range(j + tidx, m, threads):
                    abs_value = panel[j * smem_rows + local_row]
                    if abs_value < Float32(0.0):
                        abs_value = -abs_value
                    if abs_value != Float32(0.0):
                        if scale < abs_value:
                            ratio = Float32(0.0)
                            if scale != Float32(0.0):
                                ratio = scale / abs_value
                            sumsq = Float32(1.0) + sumsq * ratio * ratio
                            scale = abs_value
                        else:
                            ratio = abs_value / scale
                            sumsq += ratio * ratio
                work[tidx] = scale
                scaled_sumsq[tidx] = sumsq
                cute.arch.sync_threads()

                stride = threads // 2
                while stride > 0:
                    if tidx < stride:
                        other_scale = work[tidx + stride]
                        other_sumsq = scaled_sumsq[tidx + stride]
                        if other_scale != Float32(0.0):
                            if work[tidx] == Float32(0.0):
                                work[tidx] = other_scale
                                scaled_sumsq[tidx] = other_sumsq
                            elif work[tidx] < other_scale:
                                ratio = work[tidx] / other_scale
                                scaled_sumsq[tidx] = (
                                    other_sumsq + scaled_sumsq[tidx] * ratio * ratio
                                )
                                work[tidx] = other_scale
                            else:
                                ratio = other_scale / work[tidx]
                                scaled_sumsq[tidx] = (
                                    scaled_sumsq[tidx] + other_sumsq * ratio * ratio
                                )
                    cute.arch.sync_threads()
                    stride = stride // 2

                if tidx == 0:
                    alpha = panel[j * smem_rows + j]
                    norm = work[0] * cute.math.sqrt(scaled_sumsq[0])
                    beta = -norm
                    if alpha < Float32(0.0):
                        beta = norm
                    tau_value = (beta - alpha) / beta
                    mTau[bidx, col] = tau_value
                    params[0] = Float32(1.0) / (alpha - beta)
                    params[1] = beta
                    params[2] = tau_value
            cute.arch.sync_threads()

            inv_alpha_minus_beta = params[0]
            beta = params[1]
            for local_row in cutlass.range(j + 1 + tidx, m, threads):
                panel[j * smem_rows + local_row] = (
                    panel[j * smem_rows + local_row] * inv_alpha_minus_beta
                )
            if tidx == 0:
                panel[j * smem_rows + j] = beta
            cute.arch.sync_threads()

            if cutlass.const_expr(j < nb - 1):
                dot_count = nb - j - 1
                warp_id = tidx // 32
                lane = tidx - warp_id * 32
                if warp_id < dot_count:
                    panel_col = j + 1 + warp_id
                    dot = Float32(0.0)
                    if lane == 0:
                        dot = panel[panel_col * smem_rows + j]
                    for local_row in cutlass.range(j + 1 + lane, m, 32):
                        dot += (
                            panel[j * smem_rows + local_row]
                            * panel[panel_col * smem_rows + local_row]
                        )
                    dot = cute.arch.warp_reduction_sum(dot)
                    if lane == 0:
                        work[warp_id] = dot
                cute.arch.sync_threads()

                tau_value = params[2]
                for idx in cutlass.range(tidx, m * dot_count, threads):
                    local_row = idx // dot_count
                    local_col = idx - local_row * dot_count
                    panel_col = j + 1 + local_col
                    v = Float32(0.0)
                    if local_row == j:
                        v = Float32(1.0)
                    elif local_row > j:
                        v = panel[j * smem_rows + local_row]
                    panel[panel_col * smem_rows + local_row] = (
                        panel[panel_col * smem_rows + local_row]
                        - tau_value * v * work[local_col]
                    )
                cute.arch.sync_threads()

            if cutlass.const_expr(j > 0):
                warp_id = tidx // 32
                lane = tidx - warp_id * 32
                if warp_id < j:
                    prev = warp_id
                    dot = Float32(0.0)
                    for local_row in cutlass.range(j + lane, m, 32):
                        v = Float32(1.0)
                        if local_row > j:
                            v = panel[j * smem_rows + local_row]
                        dot += v * u_panel[prev * smem_rows + local_row]
                    dot = cute.arch.warp_reduction_sum(dot)
                    if lane == 0:
                        work[prev] = dot
                cute.arch.sync_threads()

                tau_value = params[2]
                for idx in cutlass.range(tidx, (m - j) * j, threads):
                    local_row = j + idx // j
                    prev = idx - (local_row - j) * j
                    v = Float32(1.0)
                    if local_row > j:
                        v = panel[j * smem_rows + local_row]
                    u_panel[prev * smem_rows + local_row] = (
                        u_panel[prev * smem_rows + local_row]
                        - tau_value * v * work[prev]
                    )
                cute.arch.sync_threads()

            tau_value = params[2]
            for local_row in cutlass.range(tidx, m, threads):
                v = Float32(0.0)
                if local_row == j:
                    v = Float32(1.0)
                elif local_row > j:
                    v = panel[j * smem_rows + local_row]
                u_panel[j * smem_rows + local_row] = tau_value * v
            cute.arch.sync_threads()

        if cutlass.const_expr(_ENABLE_COALESCED_PANEL_EMIT):
            for idx in cutlass.range(tidx, m * nb, threads):
                local_row = idx // nb
                local_col = idx - local_row * nb
                panel_value = panel[local_col * smem_rows + local_row]
                mH[bidx, k + local_row, k + local_col] = panel_value

                v = Float32(0.0)
                if local_row == local_col:
                    v = Float32(1.0)
                elif local_row > local_col:
                    v = panel_value
                mV[bidx, local_row, local_col] = v
                mU[bidx, local_row, local_col] = u_panel[
                    local_col * smem_rows + local_row
                ]
        else:
            for local_col in cutlass.range_constexpr(nb):
                col = k + local_col
                for local_row in cutlass.range(tidx, m, threads):
                    panel_value = panel[local_col * smem_rows + local_row]
                    mH[bidx, k + local_row, col] = panel_value

                    v = Float32(0.0)
                    if local_row == local_col:
                        v = Float32(1.0)
                    elif local_row > local_col:
                        v = panel_value
                    mV[bidx, local_row, local_col] = v
                    mU[bidx, local_row, local_col] = u_panel[
                        local_col * smem_rows + local_row
                    ]


class _WYUpdateKernel:
    def __init__(self, n: int, p: int, cn: int = 32):
        self.n = n
        self.p = p
        self.cn = cn
        self.num_threads = 256
        self.smem_bytes = p * cn * 4

    @cute.jit
    def __call__(
        self,
        h_ptr: cute.Pointer,
        v_ptr: cute.Pointer,
        u_ptr: cute.Pointer,
        batch: Int32,
        row_start: Int32,
        col_start: Int32,
        col_stop: Int32,
    ):
        n = self.n
        p = self.p
        mH = cute.make_tensor(
            h_ptr,
            cute.make_layout((batch, n, n), stride=(n * n, n, 1)),
        )
        mV = cute.make_tensor(
            v_ptr,
            cute.make_layout((batch, n, p), stride=(n * p, p, 1)),
        )
        mU = cute.make_tensor(
            u_ptr,
            cute.make_layout((batch, n, p), stride=(n * p, p, 1)),
        )
        stripes = cute.ceil_div(col_stop - col_start, self.cn)
        self.kernel(mH, mV, mU, row_start, col_start, col_stop).launch(
            grid=[batch, stripes, 1],
            block=[self.num_threads, 1, 1],
            smem=self.smem_bytes,
        )

    @cute.kernel
    def kernel(
        self,
        mH: cute.Tensor,
        mV: cute.Tensor,
        mU: cute.Tensor,
        row_start: Int32,
        col_start: Int32,
        col_stop: Int32,
    ):
        tidx, _, _ = cute.arch.thread_idx()
        bidx, stripe, _ = cute.arch.block_idx()
        n = self.n
        p_count = self.p
        cn = self.cn
        threads = self.num_threads
        col_base = col_start + stripe * cn

        smem = cutlass.utils.SmemAllocator()
        w = smem.allocate_tensor(
            Float32,
            cute.make_layout((p_count * cn,), stride=(1,)),
            byte_alignment=16,
        )

        rows = n - row_start
        for idx in cutlass.range(tidx, p_count * cn, threads):
            p = idx // cn
            local_col = idx - p * cn
            col = col_base + local_col
            acc = Float32(0.0)
            if col < col_stop:
                for local_row in cutlass.range(0, rows, 1):
                    acc += mV[bidx, local_row, p] * mH[bidx, row_start + local_row, col]
            w[idx] = acc
        cute.arch.sync_threads()

        for idx in cutlass.range(tidx, rows * cn, threads):
            local_row = idx // cn
            local_col = idx - local_row * cn
            col = col_base + local_col
            if col < col_stop:
                acc = Float32(0.0)
                for p in cutlass.range_constexpr(p_count):
                    acc += mU[bidx, local_row, p] * w[p * cn + local_col]
                mH[bidx, row_start + local_row, col] = (
                    mH[bidx, row_start + local_row, col] - acc
                )


def _get_panel_kernel(n: int, num_threads: int = 256, nb: int = 8):
    key = ("panel", n, num_threads, nb)
    kernel = _compile_cache.get(key)
    if kernel is None:
        obj = _PanelWYKernel(n, num_threads, nb)
        kernel = cute.compile(
            obj,
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            Int32(1),
            Int32(0),
        )
        _compile_cache[key] = kernel
    return kernel


def _get_panel512_kernel():
    return _get_panel_kernel(512)


def _get_wy_update_kernel(n: int, p: int):
    if n == 176:
        cn = _WY_UPDATE_CN_176
    elif n == 352:
        cn = _WY_UPDATE_CN_352
    else:
        cn = 32
    key = ("wy_update", n, p, cn)
    kernel = _compile_cache.get(key)
    if kernel is None:
        obj = _WYUpdateKernel(n, p, cn)
        kernel = cute.compile(
            obj,
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            make_ptr(Float32, 16, cute.AddressSpace.gmem, assumed_align=16),
            Int32(1),
            Int32(0),
            Int32(0),
            Int32(0),
        )
        _compile_cache[key] = kernel
    return kernel


def _get_wy_update512_kernel(p: int):
    return _get_wy_update_kernel(512, p)


def _exact_zero_tail_stop_512(data: torch.Tensor) -> int:
    stop = 384
    if data.shape[0] > 0 and data[0, 0, stop].item() != 0.0:
        return 512
    if torch.count_nonzero(data[:, :, stop:]).item() == 0:
        return stop
    return 512


def _value_prefix_stop_512(data: torch.Tensor) -> int:
    n = 512
    eps = torch.finfo(torch.float32).eps
    factor_rtol = 20.0 * n * eps
    sqrt_n = float(n) ** 0.5

    col_l1 = data.abs().sum(dim=1)
    matrix_l1 = col_l1.amax(dim=1).clamp_min(1.0e-30)
    allowed = factor_rtol * matrix_l1

    for stop, margin in ((256, 0.40), (320, 0.25)):
        tail_col_l1 = col_l1[:, stop:].amax(dim=1)
        if not bool((tail_col_l1 <= margin * allowed).all().item()):
            continue
        tail_l2 = torch.linalg.vector_norm(data[:, :, stop:], ord=2, dim=1).amax(dim=1)
        tail_l1_bound = sqrt_n * tail_l2
        if bool((tail_l1_bound <= margin * allowed).all().item()):
            return stop
    return n


def _prefilter_prefix_stop_1024(data: torch.Tensor) -> bool:
    n = 1024
    stop = 768
    tail = n - stop
    if data.shape[0] == 0:
        return False

    sample_rows = 16
    sample = data[:, :sample_rows, :]
    sample_scale = sample.abs().amax(dim=(1, 2)).clamp_min(1.0e-30)
    duplicate_delta = (sample[:, :, stop:] - sample[:, :, :tail]).abs().amax(dim=(1, 2))
    return bool((duplicate_delta <= 1.0e-3 * sample_scale).all().item())


def _value_prefix_stop_1024(h: torch.Tensor, data: torch.Tensor) -> int:
    n = 1024
    stop = 768
    eps = torch.finfo(torch.float32).eps
    factor_rtol = 20.0 * n * eps

    matrix_l1 = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
    allowed = factor_rtol * matrix_l1
    tail_lower = torch.tril(h[:, stop + 1 :, stop:])
    residual = tail_lower.abs().sum(dim=1).amax(dim=1)
    if bool((residual <= 0.0625 * allowed).all().item()):
        return stop
    return n


def _relation_copy_stop_1024(data: torch.Tensor) -> int:
    n = 1024
    stop = 768
    tail = n - stop
    if not _prefilter_prefix_stop_1024(data):
        return n

    eps = torch.finfo(torch.float32).eps
    factor_rtol = 20.0 * n * eps
    margin = 0.5
    sqrt_n = float(n) ** 0.5

    matrix_l1 = data.abs().sum(dim=1).amax(dim=1).clamp_min(1.0e-30)
    allowed = factor_rtol * matrix_l1
    delta = data[:, :, stop:] - data[:, :, :tail]
    delta_l2 = torch.linalg.vector_norm(delta, ord=2, dim=1).amax(dim=1)
    delta_l1_bound = sqrt_n * delta_l2
    if bool((delta_l1_bound <= margin * allowed).all().item()):
        return stop
    return n


def _launch_panel_wy(
    h: torch.Tensor,
    tau: torch.Tensor,
    v_ws: torch.Tensor,
    u_ws: torch.Tensor,
    k: int,
    nb: int = 8,
) -> None:
    n = h.shape[1]
    if n == 352 and _ENABLE_PANEL352_512_THREADS:
        num_threads = 512
    elif n == 1024 and _ENABLE_PANEL1024_512_THREADS:
        num_threads = 512
    else:
        num_threads = 256
    kernel = _get_panel_kernel(n, num_threads, nb)
    kernel(
        make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, tau.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, v_ws.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, u_ws.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        Int32(h.shape[0]),
        Int32(k),
    )


def _launch_panel512(
    h: torch.Tensor,
    tau: torch.Tensor,
    v_ws: torch.Tensor,
    u_ws: torch.Tensor,
    k: int,
) -> None:
    _launch_panel_wy(h, tau, v_ws, u_ws, k)


def _launch_wy_update(
    h: torch.Tensor,
    v: torch.Tensor,
    u: torch.Tensor,
    row_start: int,
    col_start: int,
    col_stop: int,
    p: int,
) -> None:
    if col_start >= col_stop:
        return
    kernel = _get_wy_update_kernel(h.shape[1], p)
    kernel(
        make_ptr(Float32, h.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, v.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        make_ptr(Float32, u.data_ptr(), cute.AddressSpace.gmem, assumed_align=16),
        Int32(h.shape[0]),
        Int32(row_start),
        Int32(col_start),
        Int32(col_stop),
    )


def _launch_wy_update512(
    h: torch.Tensor,
    v: torch.Tensor,
    u: torch.Tensor,
    row_start: int,
    col_start: int,
    col_stop: int,
    p: int,
) -> None:
    _launch_wy_update(h, v, u, row_start, col_start, col_stop, p)


class _ExactFP32Matmul:
    def __enter__(self):
        self.matmul_backend = torch.backends.cuda.matmul
        self.precision_attr = "allow_" + "t" + "f" + "32"
        self.has_fp32_precision = hasattr(self.matmul_backend, "fp32_precision")
        self.old_allow = None
        if self.has_fp32_precision:
            self.old_fp32_precision = self.matmul_backend.fp32_precision
            self.matmul_backend.fp32_precision = "ieee"
        else:
            self.old_fp32_precision = None
            self.old_allow = getattr(self.matmul_backend, self.precision_attr, None)
            if self.old_allow is not None:
                setattr(self.matmul_backend, self.precision_attr, False)
        return self

    def __exit__(self, exc_type, exc, tb):
        if self.has_fp32_precision:
            self.matmul_backend.fp32_precision = self.old_fp32_precision
        elif self.old_allow is not None:
            setattr(self.matmul_backend, self.precision_attr, self.old_allow)
        return False


class _TF32Matmul:
    def __enter__(self):
        self.matmul_backend = torch.backends.cuda.matmul
        self.precision_attr = "allow_" + "t" + "f" + "32"
        self.has_fp32_precision = hasattr(self.matmul_backend, "fp32_precision")
        self.old_allow = None
        if self.has_fp32_precision:
            self.old_fp32_precision = self.matmul_backend.fp32_precision
            self.matmul_backend.fp32_precision = "tf32"
        else:
            self.old_fp32_precision = None
            self.old_allow = getattr(self.matmul_backend, self.precision_attr, None)
            if self.old_allow is not None:
                setattr(self.matmul_backend, self.precision_attr, True)
        return self

    def __exit__(self, exc_type, exc, tb):
        if self.has_fp32_precision:
            self.matmul_backend.fp32_precision = self.old_fp32_precision
        elif self.old_allow is not None:
            setattr(self.matmul_backend, self.precision_attr, self.old_allow)
        return False


def _apply_wy_update_torch(c: torch.Tensor, v: torch.Tensor, u: torch.Tensor) -> None:
    w = torch.bmm(v.transpose(1, 2), c)
    c.baddbmm_(u, w, beta=1.0, alpha=-1.0)


def _launch_wy_update_torch(
    h: torch.Tensor,
    v: torch.Tensor,
    u: torch.Tensor,
    row_start: int,
    col_start: int,
    col_stop: int,
    p: int,
) -> None:
    if col_start >= col_stop:
        return
    rows = h.shape[1] - row_start
    _apply_wy_update_torch(
        h[:, row_start:, col_start:col_stop],
        v[:, :rows, :p],
        u[:, :rows, :p],
    )


def _launch_wy_update_torch_tf32(
    h: torch.Tensor,
    v: torch.Tensor,
    u: torch.Tensor,
    row_start: int,
    col_start: int,
    col_stop: int,
    p: int,
) -> None:
    with _TF32Matmul():
        _launch_wy_update_torch(h, v, u, row_start, col_start, col_stop, p)


def _compose_macro_u(
    old_u_tail: torch.Tensor,
    v_new: torch.Tensor,
    u_new: torch.Tensor,
    use_tf32: bool,
) -> None:
    if use_tf32:
        with _TF32Matmul():
            cross = torch.bmm(v_new.transpose(1, 2), old_u_tail)
            old_u_tail.sub_(torch.bmm(u_new, cross))
    else:
        cross = torch.bmm(v_new.transpose(1, 2), old_u_tail)
        old_u_tail.sub_(torch.bmm(u_new, cross))


def _qr512(data: torch.Tensor) -> output_t:
    batch = data.shape[0]
    n = 512
    nb = 8
    macro_width = 32
    h = data.clone()
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
    u_ws = torch.empty_like(v_ws)
    v_macro = torch.empty(
        (batch, n, macro_width), device=data.device, dtype=torch.float32
    )
    u_macro = torch.empty_like(v_macro)

    exact_stop = _exact_zero_tail_stop_512(data)
    value_stop = n if exact_stop < n else _value_prefix_stop_512(data)
    factor_stop = min(exact_stop, value_stop)
    update_stop = factor_stop if factor_stop < n else n
    deferred_stop = min(factor_stop, 320 if factor_stop < n else 384)

    with _ExactFP32Matmul():
        for macro_k in range(0, deferred_stop, macro_width):
            macro_end = min(deferred_stop, macro_k + macro_width)
            v_macro.zero_()
            u_macro.zero_()
            macro_used = 0

            for k in range(macro_k, macro_end, nb):
                m = n - k
                row_offset = k - macro_k
                _launch_panel512(h, tau, v_ws, u_ws, k)

                v_new = v_ws[:, :m, :nb]
                u_new = u_ws[:, :m, :nb]
                if macro_used:
                    old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
                    _compose_macro_u(old_u_tail, v_new, u_new, False)

                v_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(v_new)
                u_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(u_new)
                macro_used += nb

                trail_start = k + nb
                if trail_start < macro_end:
                    if _ENABLE_TORCH_WY512_SMALL_UPDATE:
                        _launch_wy_update_torch(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )
                    else:
                        _launch_wy_update512(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )

            if macro_end < update_stop:
                if _ENABLE_TORCH_WY512_MACRO_UPDATE:
                    _launch_wy_update_torch(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        update_stop,
                        macro_used,
                    )
                else:
                    _launch_wy_update512(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        update_stop,
                        macro_used,
                    )

        for k in range(deferred_stop, factor_stop, nb):
            m = n - k
            _launch_panel512(h, tau, v_ws, u_ws, k)
            trail_start = k + nb
            if trail_start < update_stop:
                if _ENABLE_TORCH_WY512_SMALL_UPDATE:
                    _launch_wy_update_torch(
                        h, v_ws, u_ws, k, trail_start, update_stop, nb
                    )
                else:
                    _launch_wy_update512(h, v_ws, u_ws, k, trail_start, update_stop, nb)

    if factor_stop < n:
        tau[:, factor_stop:].zero_()
    return h, tau


def _qr1024(data: torch.Tensor) -> output_t:
    batch = data.shape[0]
    n = 1024
    nb = 16 if _ENABLE_PANEL1024_NB16 else 8
    macro_width = 64
    deferred_stop = 512
    factor_stop = _relation_copy_stop_1024(data)
    update_stop = factor_stop
    h = data.clone()
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
    u_ws = torch.empty_like(v_ws)
    v_macro = torch.empty(
        (batch, n, macro_width), device=data.device, dtype=torch.float32
    )
    u_macro = torch.empty_like(v_macro)

    with _ExactFP32Matmul():
        for macro_k in range(0, deferred_stop, macro_width):
            macro_end = macro_k + macro_width
            v_macro.zero_()
            u_macro.zero_()
            macro_used = 0

            for k in range(macro_k, macro_end, nb):
                m = n - k
                row_offset = k - macro_k
                _launch_panel_wy(h, tau, v_ws, u_ws, k, nb)

                v_new = v_ws[:, :m, :nb]
                u_new = u_ws[:, :m, :nb]
                if macro_used:
                    old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
                    _compose_macro_u(
                        old_u_tail,
                        v_new,
                        u_new,
                        _ENABLE_TF32_WY1024_COMPOSE,
                    )

                v_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(v_new)
                u_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(u_new)
                macro_used += nb

                trail_start = k + nb
                if trail_start < macro_end:
                    if _ENABLE_TF32_WY1024_SMALL_UPDATE:
                        _launch_wy_update_torch_tf32(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )
                    else:
                        _launch_wy_update_torch(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )

            if macro_end < update_stop:
                if _ENABLE_TF32_WY1024_MACRO_UPDATE:
                    _launch_wy_update_torch_tf32(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        update_stop,
                        macro_used,
                    )
                else:
                    _launch_wy_update_torch(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        update_stop,
                        macro_used,
                    )

        for k in range(deferred_stop, factor_stop, nb):
            m = n - k
            _launch_panel_wy(h, tau, v_ws, u_ws, k, nb)
            trail_start = k + nb
            if trail_start < update_stop:
                if _ENABLE_TF32_WY1024_TAIL_UPDATE:
                    _launch_wy_update_torch_tf32(
                        h, v_ws, u_ws, k, trail_start, update_stop, nb
                    )
                else:
                    _launch_wy_update_torch(
                        h, v_ws, u_ws, k, trail_start, update_stop, nb
                    )

    if factor_stop < n:
        tail = n - factor_stop
        h[:, :, factor_stop:].zero_()
        h[:, :tail, factor_stop:].copy_(torch.triu(h[:, :tail, :tail]))
        tau[:, factor_stop:].zero_()
    return h, tau


def _qr2048(data: torch.Tensor) -> output_t:
    batch = data.shape[0]
    n = 2048
    nb = 8
    macro_width = 16
    h = data.clone()
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
    u_ws = torch.empty_like(v_ws)
    v_macro = torch.empty(
        (batch, n, macro_width), device=data.device, dtype=torch.float32
    )
    u_macro = torch.empty_like(v_macro)

    with _ExactFP32Matmul():
        for macro_k in range(0, n, macro_width):
            macro_end = min(n, macro_k + macro_width)
            v_macro.zero_()
            u_macro.zero_()
            macro_used = 0

            for k in range(macro_k, macro_end, nb):
                m = n - k
                row_offset = k - macro_k
                _launch_panel_wy(h, tau, v_ws, u_ws, k)

                v_new = v_ws[:, :m, :nb]
                u_new = u_ws[:, :m, :nb]
                if macro_used:
                    old_u_tail = u_macro[:, row_offset : row_offset + m, :macro_used]
                    _compose_macro_u(
                        old_u_tail,
                        v_new,
                        u_new,
                        _ENABLE_TF32_WY2048_COMPOSE,
                    )

                v_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(v_new)
                u_macro[
                    :, row_offset : row_offset + m, macro_used : macro_used + nb
                ].copy_(u_new)
                macro_used += nb

                trail_start = k + nb
                if trail_start < macro_end:
                    if _ENABLE_TF32_WY2048_SMALL_UPDATE:
                        _launch_wy_update_torch_tf32(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )
                    else:
                        _launch_wy_update_torch(
                            h, v_ws, u_ws, k, trail_start, macro_end, nb
                        )

            if macro_end < n:
                if (
                    _ENABLE_TF32_WY2048_MACRO_UPDATE
                    and macro_k < _TF32_WY2048_MACRO_STOP
                ):
                    _launch_wy_update_torch_tf32(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        n,
                        macro_used,
                    )
                else:
                    _launch_wy_update_torch(
                        h,
                        v_macro,
                        u_macro,
                        macro_k,
                        macro_end,
                        n,
                        macro_used,
                    )

    return h, tau


def _qr_blocked_cutedsl(data: torch.Tensor, n: int) -> output_t:
    batch = data.shape[0]
    nb = 8
    h = data.clone()
    tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
    v_ws = torch.empty((batch, n, nb), device=data.device, dtype=torch.float32)
    u_ws = torch.empty_like(v_ws)

    for k in range(0, n, nb):
        _launch_panel_wy(h, tau, v_ws, u_ws, k)
        trail_start = k + nb
        if trail_start < n:
            _launch_wy_update(h, v_ws, u_ws, k, trail_start, n, nb)
    return h, tau


def custom_kernel(data: input_t) -> output_t:
    if data.is_cuda and data.dtype == torch.float32 and data.is_contiguous():
        if data.ndim == 3 and data.shape[1] == data.shape[2]:
            n = data.shape[1]
            if n == 32:
                return _qr32(data)
            if n == 176 or n == 352:
                return _qr_blocked_cutedsl(data, n)
            if n == 512:
                return _qr512(data)
            if n == 1024:
                if _ENABLE_TORCH_WY1024_UPDATE:
                    return _qr1024(data)
                return _qr_blocked_cutedsl(data, n)
            if n == 2048:
                return _qr2048(data)
    return torch.geqrf(data)
scrolls · 1277 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