Skip to content
KernelIndex
Search⌘K

submission 835902

UjasShah · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_panel_bucket_all.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-835902?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
8.42ms
#257 of 515
2026-06-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4e0ec8a6d9909518baad7be54ff9d4dca6dd61352066499f1209fa9a126cb1c8
license declaredunknown
license concludedunknown
authorsUjasShah
imported2026-08-26

Techniques

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

num-warps = 1_qr32_kernel[(data.shape[0],)](data, h, tau, num_warps=1)

Kernel source

submission_panel_bucket_all.py662 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl
from task import input_t, output_t


@triton.jit
def _qr32_kernel(a_ptr, h_ptr, tau_ptr):
    b = tl.program_id(0)
    offs = tl.arange(0, 32)
    rows = offs[:, None]
    cols = offs[None, :]
    base = b * 1024

    mat = tl.load(a_ptr + base + rows * 32 + cols)

    for k in tl.static_range(0, 32):
        col_k = tl.sum(tl.where(cols == k, mat, 0.0), axis=1)
        alpha = tl.sum(tl.where(offs == k, col_k, 0.0), axis=0)

        below = offs > k
        sigma = tl.sum(tl.where(below, col_k * col_k, 0.0), axis=0)
        active = sigma > 0.0

        norm = tl.sqrt(alpha * alpha + sigma)
        beta_raw = tl.where(alpha <= 0.0, norm, -norm)
        beta = tl.where(active, beta_raw, alpha)
        safe_beta = tl.where(active, beta_raw, 1.0)
        tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)

        denom = tl.where(active, alpha - beta_raw, 1.0)
        v_tail = tl.where(below, col_k / denom, 0.0)
        v = tl.where(offs == k, 1.0, v_tail)

        dots = tl.sum(v[:, None] * mat, axis=0)
        updated = mat - tau_k * v[:, None] * dots[None, :]
        mat = tl.where(cols > k, updated, mat)

        diag_mask = (rows == k) & (cols == k)
        tail_mask = (rows > k) & (cols == k)
        mat = tl.where(diag_mask, beta, mat)
        mat = tl.where(tail_mask, v_tail[:, None], mat)

        tl.store(tau_ptr + b * 32 + k, tau_k)

    tl.store(h_ptr + base + rows * 32 + cols, mat)


def _qr32(data: torch.Tensor) -> output_t:
    h = torch.empty_like(data)
    tau = torch.empty((data.shape[0], 32), device=data.device, dtype=data.dtype)
    _qr32_kernel[(data.shape[0],)](data, h, tau, num_warps=1)
    return h, tau


@triton.jit
def _qr_factor_kernel(h_ptr, tau_ptr, k, n: tl.constexpr, block_m: tl.constexpr):
    b = tl.program_id(0)
    rows = tl.arange(0, block_m)
    base = b * n * n

    col_k = tl.load(h_ptr + base + rows * n + k, mask=rows < n, other=0.0)
    alpha = tl.load(h_ptr + base + k * n + k)

    below = rows > k
    sigma = tl.sum(tl.where(below, col_k * col_k, 0.0), axis=0)
    active = sigma > 0.0

    norm = tl.sqrt(alpha * alpha + sigma)
    beta_raw = tl.where(alpha <= 0.0, norm, -norm)
    beta = tl.where(active, beta_raw, alpha)
    safe_beta = tl.where(active, beta_raw, 1.0)
    tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)

    denom = tl.where(active, alpha - beta_raw, 1.0)
    v_tail = tl.where(below, col_k / denom, 0.0)

    tl.store(tau_ptr + b * n + k, tau_k)
    tl.store(h_ptr + base + k * n + k, beta)
    tl.store(
        h_ptr + base + rows * n + k,
        v_tail,
        mask=(rows > k) & (rows < n),
    )


@triton.jit
def _qr_apply_kernel(h_ptr, tau_ptr, k, n: tl.constexpr, block_m: tl.constexpr, block_n: tl.constexpr):
    b = tl.program_id(0)
    col_block = tl.program_id(1)

    rows = tl.arange(0, block_m)
    cols = col_block * block_n + tl.arange(0, block_n)
    base = b * n * n

    stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
    v = tl.where(rows == k, 1.0, 0.0)
    v = tl.where(rows > k, stored_tail, v)
    tau_k = tl.load(tau_ptr + b * n + k)

    tile = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
        other=0.0,
    )
    dots = tl.sum(v[:, None] * tile, axis=0)
    updated = tile - tau_k * v[:, None] * dots[None, :]
    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        updated,
        mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
    )


@triton.jit
def _qr_apply_range_kernel(
    h_ptr,
    tau_ptr,
    k,
    col_start: tl.constexpr,
    n: tl.constexpr,
    block_m: tl.constexpr,
    block_n: tl.constexpr,
):
    b = tl.program_id(0)
    rows = tl.arange(0, block_m)
    cols = col_start + tl.arange(0, block_n)
    base = b * n * n

    stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
    v = tl.where(rows == k, 1.0, 0.0)
    v = tl.where(rows > k, stored_tail, v)
    tau_k = tl.load(tau_ptr + b * n + k)

    tile = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
        other=0.0,
    )
    dots = tl.sum(v[:, None] * tile, axis=0)
    updated = tile - tau_k * v[:, None] * dots[None, :]
    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        updated,
        mask=(rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] > k),
    )


@triton.jit
def _qr_panel4_apply_kernel(
    h_ptr,
    tau_ptr,
    panel_start,
    n: tl.constexpr,
    block_m: tl.constexpr,
    block_n: tl.constexpr,
):
    b = tl.program_id(0)
    col_block = tl.program_id(1)
    rows = tl.arange(0, block_m)
    cols = panel_start + 4 + col_block * block_n + tl.arange(0, block_n)
    base = b * n * n

    tile = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n),
        other=0.0,
    )

    for j in tl.static_range(0, 4):
        k = panel_start + j
        stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
        v = tl.where(rows == k, 1.0, 0.0)
        v = tl.where(rows > k, stored_tail, v)
        tau_k = tl.load(tau_ptr + b * n + k)
        dots = tl.sum(v[:, None] * tile, axis=0)
        tile = tile - tau_k * v[:, None] * dots[None, :]

    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        tile,
        mask=(rows[:, None] < n) & (cols[None, :] < n),
    )


@triton.jit
def _qr_panel8_apply_kernel(
    h_ptr,
    tau_ptr,
    panel_start,
    n: tl.constexpr,
    block_m: tl.constexpr,
    block_n: tl.constexpr,
):
    b = tl.program_id(0)
    col_block = tl.program_id(1)
    rows = tl.arange(0, block_m)
    cols = panel_start + 8 + col_block * block_n + tl.arange(0, block_n)
    base = b * n * n

    tile = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n),
        other=0.0,
    )

    for j in tl.static_range(0, 8):
        k = panel_start + j
        stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
        v = tl.where(rows == k, 1.0, 0.0)
        v = tl.where(rows > k, stored_tail, v)
        tau_k = tl.load(tau_ptr + b * n + k)
        dots = tl.sum(v[:, None] * tile, axis=0)
        tile = tile - tau_k * v[:, None] * dots[None, :]

    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        tile,
        mask=(rows[:, None] < n) & (cols[None, :] < n),
    )


@triton.jit
def _qr_panel8_apply_tail_kernel(
    h_ptr,
    tau_ptr,
    panel_start,
    n: tl.constexpr,
    block_m: tl.constexpr,
    block_n: tl.constexpr,
):
    b = tl.program_id(0)
    col_block = tl.program_id(1)
    rows = panel_start + tl.arange(0, block_m)
    cols = panel_start + 8 + col_block * block_n + tl.arange(0, block_n)
    base = b * n * n

    tile = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n),
        other=0.0,
    )

    for j in tl.static_range(0, 8):
        k = panel_start + j
        stored_tail = tl.load(h_ptr + base + rows * n + k, mask=(rows > k) & (rows < n), other=0.0)
        v = tl.where(rows == k, 1.0, 0.0)
        v = tl.where(rows > k, stored_tail, v)
        tau_k = tl.load(tau_ptr + b * n + k)
        dots = tl.sum(v[:, None] * tile, axis=0)
        tile = tile - tau_k * v[:, None] * dots[None, :]

    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        tile,
        mask=(rows[:, None] < n) & (cols[None, :] < n),
    )


@triton.jit
def _qr_panel8_factor_kernel(h_ptr, tau_ptr, panel_start, n: tl.constexpr, block_m: tl.constexpr):
    b = tl.program_id(0)
    rows = tl.arange(0, block_m)
    pcols = tl.arange(0, 8)
    cols = panel_start + pcols
    base = b * n * n

    panel = tl.load(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        mask=(rows[:, None] < n) & (cols[None, :] < n),
        other=0.0,
    )

    for j in tl.static_range(0, 8):
        k = panel_start + j
        col_j = tl.sum(tl.where(pcols[None, :] == j, panel, 0.0), axis=1)
        alpha = tl.sum(tl.where(rows == k, col_j, 0.0), axis=0)

        below = rows > k
        sigma = tl.sum(tl.where(below, col_j * col_j, 0.0), axis=0)
        active = sigma > 0.0

        norm = tl.sqrt(alpha * alpha + sigma)
        beta_raw = tl.where(alpha <= 0.0, norm, -norm)
        beta = tl.where(active, beta_raw, alpha)
        safe_beta = tl.where(active, beta_raw, 1.0)
        tau_k = tl.where(active, (beta_raw - alpha) / safe_beta, 0.0)

        denom = tl.where(active, alpha - beta_raw, 1.0)
        v_tail = tl.where(below, col_j / denom, 0.0)
        v = tl.where(rows == k, 1.0, 0.0)
        v = tl.where(rows > k, v_tail, v)

        dots = tl.sum(v[:, None] * panel, axis=0)
        updated = panel - tau_k * v[:, None] * dots[None, :]
        panel = tl.where(pcols[None, :] > j, updated, panel)

        diag_mask = (rows[:, None] == k) & (pcols[None, :] == j)
        tail_mask = (rows[:, None] > k) & (pcols[None, :] == j)
        panel = tl.where(diag_mask, beta, panel)
        panel = tl.where(tail_mask, v_tail[:, None], panel)

        tl.store(tau_ptr + b * n + k, tau_k)

    tl.store(
        h_ptr + base + rows[:, None] * n + cols[None, :],
        panel,
        mask=(rows[:, None] < n) & (cols[None, :] < n),
    )


def _qr176(data: torch.Tensor) -> output_t:
    n = 176
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    apply_grid = (data.shape[0], triton.cdiv(n, 16))
    for k in range(n):
        _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=256, num_warps=8)
        _qr_apply_kernel[apply_grid](h, tau, k, n, block_m=256, block_n=16, num_warps=8)
    return h, tau


def _qr176_panel8(data: torch.Tensor) -> output_t:
    n = 176
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        _qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=256, num_warps=4)
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            active_rows = n - panel_start
            if active_rows > 128:
                _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
            elif active_rows > 64:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
            elif active_rows > 32:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
            else:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
    return h, tau


def _qr352(data: torch.Tensor) -> output_t:
    n = 352
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    apply_grid = (data.shape[0], triton.cdiv(n, 16))
    for k in range(n):
        _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=512, num_warps=8)
        _qr_apply_kernel[apply_grid](h, tau, k, n, block_m=512, block_n=16, num_warps=8)
    return h, tau


def _qr352_panel8(data: torch.Tensor) -> output_t:
    n = 352
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        _qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=512, num_warps=4)
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            active_rows = n - panel_start
            if active_rows > 256:
                _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=4)
            elif active_rows > 128:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
            elif active_rows > 64:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
            elif active_rows > 32:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
            else:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
    return h, tau


def _qr512(data: torch.Tensor) -> output_t:
    n = 512
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    apply_grid = (data.shape[0], triton.cdiv(n, 32))
    for k in range(n):
        _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
        _qr_apply_kernel[apply_grid](h, tau, k, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr512_blocked4(data: torch.Tensor) -> output_t:
    n = 512
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 4):
        for j in range(4):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
            if j < 3:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=1024,
                    block_n=4,
                    num_warps=8,
                )
        trailing = n - panel_start - 4
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel4_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr512_blocked8(data: torch.Tensor) -> output_t:
    n = 512
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        for j in range(8):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
            if j < 7:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=1024,
                    block_n=8,
                    num_warps=8,
                )
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr512_panel8(data: torch.Tensor) -> output_t:
    n = 512
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        _qr_panel8_factor_kernel[(data.shape[0],)](h, tau, panel_start, n, block_m=512, num_warps=4)
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            active_rows = n - panel_start
            if active_rows > 256:
                _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=4)
            elif active_rows > 128:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
            elif active_rows > 64:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
            elif active_rows > 32:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
            else:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
    return h, tau


def _qr1024(data: torch.Tensor) -> output_t:
    n = 1024
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    apply_grid = (data.shape[0], triton.cdiv(n, 32))
    for k in range(n):
        _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
        _qr_apply_kernel[apply_grid](h, tau, k, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr1024_blocked4(data: torch.Tensor) -> output_t:
    n = 1024
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 4):
        for j in range(4):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
            if j < 3:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=1024,
                    block_n=4,
                    num_warps=8,
                )
        trailing = n - panel_start - 4
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel4_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr1024_blocked8(data: torch.Tensor) -> output_t:
    n = 1024
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        for j in range(8):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
            if j < 7:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=1024,
                    block_n=8,
                    num_warps=8,
                )
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
    return h, tau


def _qr1024_blocked8_bucket(data: torch.Tensor) -> output_t:
    n = 1024
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        for j in range(8):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=1024, num_warps=8)
            if j < 7:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=1024,
                    block_n=8,
                    num_warps=8,
                )
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            active_rows = n - panel_start
            if active_rows > 512:
                _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=1024, block_n=32, num_warps=8)
            elif active_rows > 256:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=512, block_n=32, num_warps=8)
            elif active_rows > 128:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=256, block_n=32, num_warps=4)
            elif active_rows > 64:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=128, block_n=32, num_warps=4)
            elif active_rows > 32:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=64, block_n=32, num_warps=2)
            else:
                _qr_panel8_apply_tail_kernel[apply_grid](h, tau, panel_start, n, block_m=32, block_n=32, num_warps=1)
    return h, tau


def _qr2048_blocked8(data: torch.Tensor) -> output_t:
    n = 2048
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        for j in range(8):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=2048, num_warps=8)
            if j < 7:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=2048,
                    block_n=8,
                    num_warps=8,
                )
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=2048, block_n=32, num_warps=8)
    return h, tau


def _qr4096_blocked8(data: torch.Tensor) -> output_t:
    n = 4096
    h = data.clone()
    tau = torch.empty((data.shape[0], n), device=data.device, dtype=data.dtype)
    for panel_start in range(0, n, 8):
        for j in range(8):
            k = panel_start + j
            _qr_factor_kernel[(data.shape[0],)](h, tau, k, n, block_m=4096, num_warps=8)
            if j < 7:
                _qr_apply_range_kernel[(data.shape[0],)](
                    h,
                    tau,
                    k,
                    panel_start,
                    n,
                    block_m=4096,
                    block_n=8,
                    num_warps=8,
                )
        trailing = n - panel_start - 8
        if trailing:
            apply_grid = (data.shape[0], triton.cdiv(trailing, 32))
            _qr_panel8_apply_kernel[apply_grid](h, tau, panel_start, n, block_m=4096, block_n=32, num_warps=8)
    return h, tau


def custom_kernel(data: input_t) -> output_t:
    if (
        data.shape[-1] == 32
        and data.shape[-2] == 32
        and data.dtype == torch.float32
        and data.is_contiguous()
    ):
        return _qr32(data)

    if (
        data.shape[-1] == 176
        and data.shape[-2] == 176
        and data.dtype == torch.float32
        and data.is_contiguous()
    ):
        return _qr176_panel8(data)

    if (
        data.shape[-1] == 352
        and data.shape[-2] == 352
        and data.dtype == torch.float32
        and data.is_contiguous()
    ):
        return _qr352_panel8(data)

    if (
        data.shape[-1] == 512
        and data.shape[-2] == 512
        and data.dtype == torch.float32
        and data.is_contiguous()
    ):
        return _qr512_panel8(data)

    if (
        data.shape[-1] == 1024
        and data.shape[-2] == 1024
        and data.dtype == torch.float32
        and data.is_contiguous()
    ):
        return _qr1024_blocked8_bucket(data)

    return torch.geqrf(data)
scrolls · 662 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