Skip to content
KernelIndex
Search⌘K

submission 839047

benfattori · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-839047?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
3.62ms
#115 of 515
2026-06-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:289c318462d714efcd7afc2139ff155eb57422913b30c48994ed4f33aaec8f24
license declaredunknown
license concludedunknown
authorsbenfattori
imported2026-08-26

Techniques

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

autotunetriton.Config(
fused-epilogueDO_EPILOGUE: tl.constexpr,
mmaacc = tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)
num-warps = 4num_warps=4,
split-kdef _prune_bmm_t_split_k_configs(configs, named_args, **_kwargs):
stages = 3num_stages=3,

Kernel source

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


_BMM_ACC_AUTOTUNE_CONFIGS = [
    triton.Config(
        {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 256},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 128},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 256},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 32},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 64},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 128},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256},
        num_warps=8,
        num_stages=3,
    ),
]


_BMM_T_AUTOTUNE_CONFIGS = [
    triton.Config(
        {"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 128},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 16, "BLOCK_SIZE_N": 256},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 128},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 32, "BLOCK_SIZE_N": 256},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 32},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 64},
        num_warps=4,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 128},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 64, "BLOCK_SIZE_N": 256},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 32},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 64},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 128},
        num_warps=8,
        num_stages=3,
    ),
    triton.Config(
        {"BLOCK_SIZE_K": 128, "BLOCK_SIZE_N": 256},
        num_warps=8,
        num_stages=3,
    ),
]


def _prune_bmm_acc_configs(configs, named_args, **_kwargs):
    dim_m = named_args["dim_m"]
    dim_n = named_args["dim_n"]
    pruned = []
    for config in configs:
        block_m = config.kwargs["BLOCK_SIZE_M"]
        block_n = config.kwargs["BLOCK_SIZE_N"]

        if dim_m <= 16 and block_m != 16:
            continue
        if dim_m > 16 and block_m < 32:
            continue
        if dim_n <= 16 and block_n > 16:
            continue
        if dim_n <= 32 and block_n > 32:
            continue
        if dim_n <= 64 and block_n > 64:
            continue
        pruned.append(config)
    return pruned or configs[:1]


def _prune_bmm_t_configs(configs, named_args, **_kwargs):
    dim_k = named_args["dim_k"]
    dim_n = named_args["dim_n"]
    pruned = []
    for config in configs:
        block_k = config.kwargs["BLOCK_SIZE_K"]
        block_n = config.kwargs["BLOCK_SIZE_N"]

        if dim_k <= 16 and block_k != 16:
            continue
        if dim_k <= 32 and block_k > 32:
            continue
        if dim_k <= 64 and block_k > 64:
            continue
        if dim_k >= 128 and block_k < 32:
            continue
        if dim_n <= 16 and block_n > 16:
            continue
        if dim_n <= 32 and block_n > 32:
            continue
        if dim_n <= 64 and block_n > 64:
            continue
        pruned.append(config)
    return pruned or configs[:1]


def _prune_bmm_t_split_k_configs(configs, named_args, **_kwargs):
    dim_k = named_args["tune_split_k"]
    dim_n = named_args["dim_n"]
    pruned = []
    for config in configs:
        block_k = config.kwargs["BLOCK_SIZE_K"]
        block_n = config.kwargs["BLOCK_SIZE_N"]

        if dim_k <= 16 and block_k != 16:
            continue
        if dim_k <= 32 and block_k > 32:
            continue
        if dim_k <= 64 and block_k > 64:
            continue
        if dim_k >= 128 and block_k < 32:
            continue
        if dim_n <= 16 and block_n > 16:
            continue
        if dim_n <= 32 and block_n > 32:
            continue
        if dim_n <= 64 and block_n > 64:
            continue
        pruned.append(config)
    return pruned or configs[:1]


def _shape_bucket(value: int) -> int:
    return triton.next_power_of_2(value)


@triton.autotune(
    configs=_BMM_ACC_AUTOTUNE_CONFIGS,
    key=["op_kind", "tune_m", "tune_k", "tune_n"],
    prune_configs_by={"early_config_prune": _prune_bmm_acc_configs},
    restore_value=["c_out_ptr"],
    warmup=5,
    rep=15,
)
@triton.jit
def _bmm_tf32x3_kernel(
    a_ptr,
    b_ptr,
    c_in_ptr,
    c_out_ptr,
    bounds_ptr,
    panel_end,
    a_b_stride,
    a_m_stride,
    a_k_stride,
    b_b_stride,
    b_k_stride,
    b_n_stride,
    c_in_b_stride,
    c_in_m_stride,
    c_in_n_stride,
    c_out_b_stride,
    c_out_m_stride,
    c_out_n_stride,
    dim_m,
    dim_k,
    dim_n,
    op_kind,
    tune_m,
    tune_k,
    tune_n,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    specialize_constexpr: tl.constexpr,
):
    pid = tl.program_id(0)
    b_pid = tl.program_id(1)

    active_n = dim_n
    if specialize_constexpr:
        bounds_ptr += b_pid
        bound = tl.load(bounds_ptr)
        active_n = bound - panel_end
        if active_n <= 0:
            return

    num_pid_n = tl.cdiv(dim_n, BLOCK_SIZE_N)

    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n
    n_tile_start = pid_n * BLOCK_SIZE_N
    if specialize_constexpr:
        if n_tile_start >= active_n:
            return

    a_base = a_ptr + b_pid * a_b_stride
    b_base = b_ptr + b_pid * b_b_stride
    c_in_base = c_in_ptr + b_pid * c_in_b_stride
    c_out_base = c_out_ptr + b_pid * c_out_b_stride

    a_block_ptr = tl.make_block_ptr(
        base=a_base,
        shape=(dim_m, dim_k),
        strides=(a_m_stride, a_k_stride),
        offsets=(pid_m * BLOCK_SIZE_M, 0),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_K),
        order=(1, 0),
    )
    b_block_ptr = tl.make_block_ptr(
        base=b_base,
        shape=(dim_k, dim_n),
        strides=(b_k_stride, b_n_stride),
        offsets=(0, n_tile_start),
        block_shape=(BLOCK_SIZE_K, BLOCK_SIZE_N),
        order=(1, 0),
    )

    a = tl.load(a_block_ptr, boundary_check=(0, 1))
    b = tl.load(b_block_ptr, boundary_check=(0, 1))

    acc = tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)

    c_in_block_ptr = tl.make_block_ptr(
        base=c_in_base,
        shape=(dim_m, dim_n),
        strides=(c_in_m_stride, c_in_n_stride),
        offsets=(pid_m * BLOCK_SIZE_M, n_tile_start),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
        order=(1, 0),
    )
    c = tl.load(c_in_block_ptr, boundary_check=(0, 1))
    acc = c - acc

    c_out_block_ptr = tl.make_block_ptr(
        base=c_out_base,
        shape=(dim_m, dim_n),
        strides=(c_out_m_stride, c_out_n_stride),
        offsets=(pid_m * BLOCK_SIZE_M, n_tile_start),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
        order=(1, 0),
    )
    tl.store(c_out_block_ptr, acc, boundary_check=(0, 1))


@triton.autotune(
    configs=_BMM_T_AUTOTUNE_CONFIGS,
    key=["tune_m", "tune_k", "tune_n"],
    prune_configs_by={"early_config_prune": _prune_bmm_t_configs},
    warmup=5,
    rep=15,
)
@triton.jit
def bmm_then_t_tf32x3_kernel(
    a_ptr,
    b_ptr,
    t_ptr,
    c_ptr,
    bounds_ptr,
    panel_end,
    a_b_stride,
    a_m_stride,
    a_k_stride,
    b_b_stride,
    b_k_stride,
    b_n_stride,
    t_b_stride,
    t_m_stride,
    t_k_stride,
    c_b_stride,
    c_m_stride,
    c_n_stride,
    dim_m,
    dim_k,
    dim_n,
    tune_m,
    tune_k,
    tune_n,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    specialize_constexpr: tl.constexpr,
):
    pid = tl.program_id(0)
    b_pid = tl.program_id(1)

    active_n = dim_n
    if specialize_constexpr:
        bounds_ptr += b_pid
        bound = tl.load(bounds_ptr)
        active_n = bound - panel_end
        if active_n <= 0:
            return

    # dim_m is the panel width.  One program handles all panel rows for one
    # output-column tile so the following T multiply has the full small block.
    pid_m = 0  # noqa
    pid_n = pid
    n_tile_start = pid_n * BLOCK_SIZE_N
    if specialize_constexpr:
        if n_tile_start >= active_n:
            return

    a_base = a_ptr + b_pid * a_b_stride
    b_base = b_ptr + b_pid * b_b_stride
    t_base = t_ptr + b_pid * t_b_stride
    c_base = c_ptr + b_pid * c_b_stride

    a_block_ptr = tl.make_block_ptr(
        base=a_base,
        shape=(dim_m, dim_k),
        strides=(a_m_stride, a_k_stride),
        offsets=(0, 0),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_K),
        order=(1, 0),
    )
    b_block_ptr = tl.make_block_ptr(
        base=b_base,
        shape=(dim_k, dim_n),
        strides=(b_k_stride, b_n_stride),
        offsets=(0, n_tile_start),
        block_shape=(BLOCK_SIZE_K, BLOCK_SIZE_N),
        order=(1, 0),
    )

    acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    for _ in range(0, tl.cdiv(dim_k, BLOCK_SIZE_K)):
        a = tl.load(a_block_ptr, boundary_check=(0, 1))
        b = tl.load(b_block_ptr, boundary_check=(0, 1))

        acc += tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)

        a_block_ptr = tl.advance(a_block_ptr, (0, BLOCK_SIZE_K))
        b_block_ptr = tl.advance(b_block_ptr, (BLOCK_SIZE_K, 0))

    t_block_ptr = tl.make_block_ptr(
        base=t_base,
        shape=(dim_m, dim_m),
        strides=(t_m_stride, t_k_stride),
        offsets=(0, 0),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_M),
        order=(1, 0),
    )
    t = tl.load(t_block_ptr, boundary_check=(0, 1))

    acc = tl.dot(t, acc, input_precision="tf32x3", out_dtype=tl.float32)

    c_block_ptr = tl.make_block_ptr(
        base=c_base,
        shape=(dim_m, dim_n),
        strides=(c_m_stride, c_n_stride),
        offsets=(0, n_tile_start),
        block_shape=(BLOCK_SIZE_M, BLOCK_SIZE_N),
        order=(1, 0),
    )
    tl.store(c_block_ptr, acc, boundary_check=(0, 1))


@triton.autotune(
    configs=_BMM_T_AUTOTUNE_CONFIGS,
    key=["tune_m", "tune_split_k", "tune_n", "split_k_slices"],
    prune_configs_by={"early_config_prune": _prune_bmm_t_split_k_configs},
    restore_value=["c_ptr"],
    warmup=5,
    rep=15,
)
@triton.jit
def bmm_then_t_split_k_tf32x3_kernel(
    a_ptr,
    b_ptr,
    t_ptr,
    c_ptr,
    bounds_ptr,
    panel_end,
    a_b_stride,
    a_m_stride,
    a_k_stride,
    b_b_stride,
    b_k_stride,
    b_n_stride,
    t_b_stride,
    t_m_stride,
    t_k_stride,
    c_b_stride,
    c_m_stride,
    c_n_stride,
    dim_m,
    dim_k,
    dim_n,
    split_k_slices,
    tune_m,
    tune_split_k,
    tune_n,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    specialize_constexpr: tl.constexpr,
):
    pid_n = tl.program_id(0)
    b_pid = tl.program_id(1)
    split_pid = tl.program_id(2)
    n_tile_start = pid_n * BLOCK_SIZE_N

    active_n = dim_n
    if specialize_constexpr:
        bounds_ptr += b_pid
        bound = tl.load(bounds_ptr)
        active_n = bound - panel_end
        if active_n <= 0:
            return
        if n_tile_start >= active_n:
            return

    split_size = tl.cdiv(dim_k, split_k_slices)
    split_start = split_pid * split_size
    split_end = min(split_start + split_size, dim_k)

    a_base = a_ptr + b_pid * a_b_stride
    b_base = b_ptr + b_pid * b_b_stride
    t_base = t_ptr + b_pid * t_b_stride
    c_base = c_ptr + b_pid * c_b_stride

    offs_m = tl.arange(0, BLOCK_SIZE_M)
    offs_k = tl.arange(0, BLOCK_SIZE_K)
    offs_n = n_tile_start + tl.arange(0, BLOCK_SIZE_N)

    acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    for k_offset in range(0, tl.cdiv(split_size, BLOCK_SIZE_K)):
        k_idxs = split_start + k_offset * BLOCK_SIZE_K + offs_k

        a = tl.load(
            a_base + offs_m[:, None] * a_m_stride + k_idxs[None, :] * a_k_stride,
            mask=(offs_m[:, None] < dim_m) & (k_idxs[None, :] < split_end),
            other=0.0,
        )
        b = tl.load(
            b_base + k_idxs[:, None] * b_k_stride + offs_n[None, :] * b_n_stride,
            mask=(k_idxs[:, None] < split_end) & (offs_n[None, :] < active_n),
            other=0.0,
        )

        acc += tl.dot(a, b, input_precision="tf32x3", out_dtype=tl.float32)

    offs_t = tl.arange(0, BLOCK_SIZE_M)
    t = tl.load(
        t_base + offs_m[:, None] * t_m_stride + offs_t[None, :] * t_k_stride,
        mask=(offs_m[:, None] < dim_m) & (offs_t[None, :] < dim_m),
        other=0.0,
    )
    acc = tl.dot(t, acc, input_precision="tf32x3", out_dtype=tl.float32)

    tl.atomic_add(
        c_base + offs_m[:, None] * c_m_stride + offs_n[None, :] * c_n_stride,
        acc,
        sem="relaxed",
        mask=(offs_m[:, None] < dim_m) & (offs_n[None, :] < active_n),
    )


def bmm_then_t_tf32x3(
    v: torch.Tensor,
    c: torch.Tensor,
    t: torch.Tensor,
    local_n_bounds: torch.Tensor,
    panel_end: int,
    specialize_constexpr: bool,
    split_k: bool = False,
    split_k_slices: int = 4,
) -> torch.Tensor:
    B, K, M = v.shape
    N = c.shape[2]
    if split_k and split_k_slices > 1:
        actual_split_k = min(split_k_slices, K)
        split_k_size = triton.cdiv(K, actual_split_k)
        out = torch.empty((B, M, N), device=v.device, dtype=v.dtype)
        out.zero_()
        grid = lambda META: (  # noqa
            triton.cdiv(N, META["BLOCK_SIZE_N"]),
            B,
            actual_split_k,
        )
        bmm_then_t_split_k_tf32x3_kernel[grid](
            v,
            c,
            t,
            out,
            local_n_bounds,
            panel_end,
            v.stride(0),
            v.stride(2),
            v.stride(1),
            c.stride(0),
            c.stride(1),
            c.stride(2),
            t.stride(0),
            t.stride(2),
            t.stride(1),
            out.stride(0),
            out.stride(1),
            out.stride(2),
            M,
            K,
            N,
            actual_split_k,
            _shape_bucket(M),
            _shape_bucket(split_k_size),
            _shape_bucket(N),
            BLOCK_SIZE_M=_shape_bucket(M),
            specialize_constexpr=specialize_constexpr,
        )
        return out

    out = torch.empty((B, M, N), device=v.device, dtype=v.dtype)
    grid = lambda META: (  # noqa
        triton.cdiv(N, META["BLOCK_SIZE_N"]),
        B,
    )
    bmm_then_t_tf32x3_kernel[grid](
        v,
        c,
        t,
        out,
        local_n_bounds,
        panel_end,
        v.stride(0),
        v.stride(2),
        v.stride(1),
        c.stride(0),
        c.stride(1),
        c.stride(2),
        t.stride(0),
        t.stride(2),
        t.stride(1),
        out.stride(0),
        out.stride(1),
        out.stride(2),
        M,
        K,
        N,
        _shape_bucket(M),
        _shape_bucket(K),
        _shape_bucket(N),
        BLOCK_SIZE_M=_shape_bucket(M),
        specialize_constexpr=specialize_constexpr,
    )
    return out


def baddbmm_sub_tf32x3(
    c: torch.Tensor,
    a: torch.Tensor,
    b: torch.Tensor,
    local_bounds: torch.Tensor,
    panel_end: int,
    specialize_constexpr: bool,
    c_source: torch.Tensor | None = None,
) -> None:
    if c_source is None:
        c_source = c
    B, M, K = a.shape
    N = b.shape[2]
    grid = lambda META: (  # noqa
        triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),
        B,
    )
    _bmm_tf32x3_kernel[grid](
        a,
        b,
        c_source,
        c,
        local_bounds,
        panel_end,
        a.stride(0),
        a.stride(1),
        a.stride(2),
        b.stride(0),
        b.stride(1),
        b.stride(2),
        c_source.stride(0),
        c_source.stride(1),
        c_source.stride(2),
        c.stride(0),
        c.stride(1),
        c.stride(2),
        M,
        K,
        N,
        1,
        _shape_bucket(M),
        _shape_bucket(K),
        _shape_bucket(N),
        BLOCK_SIZE_K=_shape_bucket(K),
        specialize_constexpr=specialize_constexpr,
    )


@triton.jit
def tiled_qr_naive_kernel(
    panel_src_ptr,
    H_ptr,
    tau_ptr,
    V_ptr,
    T_ptr,
    bounds_ptr,
    panel_src_b_stride,
    panel_src_r_stride,
    panel_src_c_stride,
    H_b_stride,
    H_r_stride,
    H_c_stride,
    V_b_stride,
    V_r_stride,
    V_c_stride,
    T_b_stride,
    T_r_stride,
    T_c_stride,
    tau_b_stride,
    tau_r_stride,
    col_start,
    row_start,
    n,
    valid_cols,
    TILE_SIZE: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    DO_EPILOGUE: tl.constexpr,
    specialize_constexpr: tl.constexpr,
):
    b_pid = tl.program_id(0)

    if specialize_constexpr:
        bounds_ptr += b_pid
        bound = tl.load(bounds_ptr)
        if col_start >= bound:
            return

    # offset to the start of the loop
    panel_src_ptr += (
        b_pid * panel_src_b_stride
        + row_start * panel_src_r_stride
        + col_start * panel_src_c_stride
    )
    H_ptr += b_pid * H_b_stride + row_start * H_r_stride + col_start * H_c_stride
    tau_ptr += b_pid * tau_b_stride + col_start * tau_r_stride
    V_ptr += b_pid * V_b_stride
    T_ptr += b_pid * T_b_stride

    row_offsets = tl.arange(0, BLOCK_SIZE)
    col_offsets = tl.arange(0, TILE_SIZE)
    rows = row_offsets[:, None]
    cols = col_offsets[None, :]

    m = n - col_start

    tau_smem = tl.zeros((TILE_SIZE,), dtype=tl.float32)

    mask = (rows < m) & (cols < valid_cols)

    H_panel = tl.load(
        panel_src_ptr + rows * panel_src_r_stride + cols * panel_src_c_stride,
        mask=mask,
        other=0.0,
    )  # [BLOCK_SIZE, TILE_SIZE]

    for j in range(TILE_SIZE):
        H_col = tl.sum(tl.where(cols == j, H_panel, 0.0), axis=1)

        alpha = tl.sum(tl.where(row_offsets == j, H_col, 0.0))

        tail = tl.where(row_offsets > j, H_col, 0.0)

        tail_norm = tl.sum(tail * tail)

        active = tail_norm != 0

        x_norm = tl.sqrt(alpha * alpha + tail_norm)

        beta_new = tl.where(alpha >= 0, -x_norm, x_norm)
        beta = tl.where(active, beta_new, alpha)

        denom = alpha - beta_new

        denom_safe = tl.where(active, denom, 1.0)

        tau_col = tl.where(active, (beta_new - alpha) / beta_new, 0.0)

        scale = tl.where(active, 1.0 / denom_safe, 0.0)

        col_value_tail = tail * scale

        col_value = tl.where(row_offsets > j, col_value_tail, H_col)
        col_value = tl.where(row_offsets == j, beta, col_value)

        H_panel = tl.where(cols == j, col_value[:, None], H_panel)

        tau_smem = tl.where(col_offsets == j, tau_col, tau_smem)

        v = tl.where(row_offsets == j, 1.0, col_value_tail)
        vtC = tl.sum(tl.where(cols > j, v[:, None] * H_panel, 0.0), axis=0)
        H_applied = H_panel - tau_col * v[:, None] * vtC[None, :]
        H_panel = tl.where(cols > j, H_applied, H_panel)

    tl.store(H_ptr + rows * H_r_stride + cols * H_c_stride, mask=mask, value=H_panel)

    store_offs = tl.arange(0, TILE_SIZE)
    tl.store(
        tau_ptr + store_offs * tau_r_stride, tau_smem, mask=store_offs < valid_cols
    )

    V_panel = tl.where(
        rows == cols,
        1.0,
        tl.where(rows > cols, H_panel, 0.0),
    )

    tl.store(
        V_ptr + rows * V_r_stride + cols * V_c_stride,
        mask=mask,
        value=V_panel,
    )

    if DO_EPILOGUE:
        T_panel = tl.zeros([TILE_SIZE, TILE_SIZE], dtype=tl.float32)

        T_rows = col_offsets[:, None]
        T_cols = col_offsets[None, :]

        T_panel = tl.where(T_rows == T_cols, tau_smem, T_panel)

        for j in range(1, TILE_SIZE):
            Vi = tl.sum(tl.where(cols == j, V_panel, 0.0), axis=1)
            tau_col = tl.sum(tl.where(col_offsets == j, tau_smem, 0.0))
            dots = tl.sum(V_panel * Vi[:, None], axis=0)
            z = tl.where(col_offsets < j, -tau_col * dots, 0.0)
            t = tl.sum(T_panel * z[None, :], axis=1)

            T_panel = tl.where(
                (T_rows < j) & (T_cols == j),
                t[:, None],
                T_panel,
            )

        tl.store(
            T_ptr + T_rows * T_r_stride + T_cols * T_c_stride,
            mask=(T_rows < valid_cols) & (T_cols < valid_cols),
            value=T_panel,
        )


# this is only for the n = 32 case
@triton.jit
def tiled_qr_naive_kernel_smol_h(
    panel_src_ptr,
    H_ptr,
    tau_ptr,
    panel_src_b_stride,
    panel_src_r_stride,
    panel_src_c_stride,
    H_b_stride,
    H_r_stride,
    H_c_stride,
    tau_b_stride,
    tau_r_stride,
    col_start,
    row_start,
    n,
    valid_cols,
    TILE_SIZE: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    b_pid = tl.program_id(0)
    n = n.to(tl.int32)
    col_start = col_start.to(tl.int32)
    valid_cols = valid_cols.to(tl.int32)

    # offset to the start of the loop
    panel_src_ptr += (
        b_pid * panel_src_b_stride
        + row_start * panel_src_r_stride
        + col_start * panel_src_c_stride
    )
    H_ptr += b_pid * H_b_stride + row_start * H_r_stride + col_start * H_c_stride
    tau_ptr += b_pid * tau_b_stride + col_start * tau_r_stride

    row_offsets = tl.arange(0, BLOCK_SIZE).to(tl.int32)
    col_offsets = tl.arange(0, TILE_SIZE).to(tl.int32)
    rows = row_offsets[:, None]
    cols = col_offsets[None, :]

    m = n - col_start

    tau_smem = tl.zeros((TILE_SIZE,), dtype=tl.float32)

    mask = (rows < m) & (cols < valid_cols)

    H_panel = tl.load(
        panel_src_ptr + rows * panel_src_r_stride + cols * panel_src_c_stride,
        mask=mask,
        other=0.0,
    )  # [BLOCK_SIZE, TILE_SIZE]

    for j in range(TILE_SIZE):
        H_col = tl.sum(tl.where(cols == j, H_panel, 0.0), axis=1)

        alpha = tl.sum(tl.where(row_offsets == j, H_col, 0.0))

        tail = tl.where(row_offsets > j, H_col, 0.0)

        tail_norm = tl.sum(tail * tail)

        active = tail_norm != 0

        x_norm = tl.sqrt(alpha * alpha + tail_norm)

        beta_new = tl.where(alpha >= 0, -x_norm, x_norm)
        beta = tl.where(active, beta_new, alpha)

        denom = alpha - beta_new

        denom_safe = tl.where(active, denom, 1.0)

        tau_col = tl.where(active, (beta_new - alpha) / beta_new, 0.0)

        scale = tl.where(active, 1.0 / denom_safe, 0.0)

        col_value_tail = tail * scale

        col_value = tl.where(row_offsets > j, col_value_tail, H_col)
        col_value = tl.where(row_offsets == j, beta, col_value)

        H_panel = tl.where(cols == j, col_value[:, None], H_panel)

        tau_smem = tl.where(col_offsets == j, tau_col, tau_smem)

        v = tl.where(row_offsets == j, 1.0, col_value_tail)
        vtC = tl.sum(tl.where(cols > j, v[:, None] * H_panel, 0.0), axis=0)
        H_applied = H_panel - tau_col * v[:, None] * vtC[None, :]
        H_panel = tl.where(cols > j, H_applied, H_panel)

    tl.store(H_ptr + rows * H_r_stride + cols * H_c_stride, mask=mask, value=H_panel)

    store_offs = tl.arange(0, TILE_SIZE).to(tl.int32)
    tl.store(
        tau_ptr + store_offs * tau_r_stride, tau_smem, mask=store_offs < valid_cols
    )


def launch_tiled_qr_panel(
    H: torch.Tensor,  # [b,n,n]
    tau: torch.Tensor,  # [b,n]
    V: torch.Tensor,  # [b, n - col_start, tile_size]
    T: torch.Tensor,  # [b,tile_size,tile_size]
    local_bounds: torch.Tensor,  # [b]
    col_start: int,
    specialize_constexpr: bool,
    tile_size: int = 16,
    panel_source: torch.Tensor | None = None,
):
    B, n, _ = H.shape
    grid = (B,)
    if panel_source is None:
        panel_source = H

    row_start = col_start
    actual_block_size = n - col_start
    block_size = triton.next_power_of_2(actual_block_size)
    valid_cols = min(tile_size, n - col_start)

    if n == 32:
        tiled_qr_naive_kernel_smol_h[grid](
            panel_source,
            H,
            tau,
            panel_source.stride(0),
            panel_source.stride(1),
            panel_source.stride(2),
            H.stride(0),
            H.stride(1),
            H.stride(2),
            tau.stride(0),
            tau.stride(1),
            col_start,
            row_start,
            n,
            valid_cols,
            TILE_SIZE=tile_size,  # type: ignore
            BLOCK_SIZE=block_size,  # type: ignore
            num_warps=1,
        )
        return

    # best config a quick scan
    if block_size <= 128:
        num_warps = 1
    elif block_size <= 256:
        num_warps = 2
    elif block_size < 1024:
        num_warps = 4
    else:
        num_warps = 8

    # fmt: off
    tiled_qr_naive_kernel[grid](
        panel_source, H, tau, V, T, local_bounds,
        panel_source.stride(0), panel_source.stride(1), panel_source.stride(2),
        H.stride(0), H.stride(1), H.stride(2),
        V.stride(0), V.stride(1), V.stride(2),
        T.stride(0), T.stride(1), T.stride(2),
        tau.stride(0), tau.stride(1),
        col_start, row_start, n, valid_cols,
        TILE_SIZE=tile_size,  # type: ignore
        BLOCK_SIZE=block_size,  # type: ignore
        DO_EPILOGUE=((col_start + valid_cols) < n),  # type: ignore
        specialize_constexpr = specialize_constexpr, # type: ignore
        num_warps = num_warps, # type: ignore
    )
    # fmt: on


@triton.jit
def check_structure_kernel(
    A_ptr,
    local_bounds_ptr,
    structure_meta_ptr,
    A_stride_b,
    A_stride_r,
    A_stride_c,
    local_stride_b,
    BLOCK_SIZE_ROWS: tl.constexpr,
    BLOCK_SIZE_COLS: tl.constexpr,
    ROWS: tl.constexpr,
    RANK_OFFSET: tl.constexpr,
    HEAD_BOUND: tl.constexpr,
    CLUSTER_BOUND: tl.constexpr,
    CLUSTERED_BOUND: tl.constexpr,
):
    pid = tl.program_id(0)
    A_ptr += A_stride_b * pid
    local_bounds_ptr += local_stride_b * pid

    offs_rows = tl.arange(0, BLOCK_SIZE_ROWS)[:, None]
    offs_cols = tl.arange(0, BLOCK_SIZE_COLS)[None, :]

    head_max = tl.full((), 0.0, tl.float32)
    cluster_tail_max = tl.full((), 0.0, tl.float32)
    rank_tail_max = tl.full((), 0.0, tl.float32)

    for row_start in range(0, ROWS, BLOCK_SIZE_ROWS):
        rows = row_start + offs_rows
        row_mask = rows < ROWS

        head_cols = offs_cols
        head_mask = row_mask & (head_cols < HEAD_BOUND)
        head = tl.load(
            A_ptr + rows * A_stride_r + head_cols * A_stride_c,
            mask=head_mask,
            other=0.0,
        )
        head_abs = tl.where(head_mask, tl.abs(head), 0.0)
        head_max = tl.maximum(head_max, tl.max(tl.max(head_abs, axis=0), axis=0))

        tail_cols = CLUSTER_BOUND + offs_cols
        tail_mask = row_mask & (tail_cols < ROWS)
        tail = tl.load(
            A_ptr + rows * A_stride_r + tail_cols * A_stride_c,
            mask=tail_mask,
            other=0.0,
        )
        tail_abs = tl.where(tail_mask, tl.abs(tail), 0.0)
        cluster_tail_max = tl.maximum(
            cluster_tail_max, tl.max(tl.max(tail_abs, axis=0), axis=0)
        )

        rank_tail_abs = tl.where(tail_cols >= RANK_OFFSET, tail_abs, 0.0)
        rank_tail_max = tl.maximum(
            rank_tail_max, tl.max(tl.max(rank_tail_abs, axis=0), axis=0)
        )

    bound = tl.where(rank_tail_max == 0.0, RANK_OFFSET, ROWS)
    clustered = cluster_tail_max <= tl.maximum(head_max, 1.0e-30) * 1.0e-5
    bound = tl.where(clustered, tl.minimum(bound, CLUSTERED_BOUND), bound)
    tl.store(local_bounds_ptr, bound)
    tl.atomic_max(structure_meta_ptr, bound, sem="relaxed")
    tl.atomic_max(structure_meta_ptr + 1, tl.where(bound < ROWS, 1, 0), sem="relaxed")


@triton.jit
def cleanup_structure_tail_kernel(
    H_ptr,
    tau_ptr,
    local_bounds_ptr,
    H_b_stride,
    H_r_stride,
    H_c_stride,
    tau_b_stride,
    tau_r_stride,
    local_stride_b,
    n,
    BLOCK_SIZE_ROWS: tl.constexpr,
    BLOCK_SIZE_COLS: tl.constexpr,
):
    b_pid = tl.program_id(0)
    row_pid = tl.program_id(1)
    col_pid = tl.program_id(2)

    bound = tl.load(local_bounds_ptr + b_pid * local_stride_b)
    col_start = col_pid * BLOCK_SIZE_COLS
    if col_start + BLOCK_SIZE_COLS <= bound:
        return

    rows = row_pid * BLOCK_SIZE_ROWS + tl.arange(0, BLOCK_SIZE_ROWS)
    cols = col_start + tl.arange(0, BLOCK_SIZE_COLS)

    mask = (rows[:, None] < n) & (cols[None, :] < n) & (cols[None, :] >= bound)
    tl.store(
        H_ptr
        + b_pid * H_b_stride
        + rows[:, None] * H_r_stride
        + cols[None, :] * H_c_stride,
        0.0,
        mask=mask,
    )

    if row_pid == 0:
        tau_mask = (cols < n) & (cols >= bound)
        tl.store(
            tau_ptr + b_pid * tau_b_stride + cols * tau_r_stride,
            0.0,
            mask=tau_mask,
        )


def geqrf_blocked_batched(
    A: torch.Tensor,
    tile_size: int = 16,
    split_k: bool = False,
    split_k_slices: int = 4,
    exploit_structure: bool = True,
):
    if A.ndim != 3 or A.shape[-1] != A.shape[-2]:
        raise ValueError("Expected A with shape (batch, n, n)")
    if tile_size <= 0:
        raise ValueError("tile_size must be positive")

    H = A.new_empty(A.shape)
    B, n, _ = H.shape
    tau = H.new_empty(B, n)

    # super early exit on n = 32 case, we can save some allocations and writes (V & T)
    if n == 32:
        launch_tiled_qr_panel(
            H,
            tau,
            None,  # type: ignore
            None,  # type: ignore
            None,  # type: ignore
            0,
            False,
            tile_size,
            panel_source=A,
        )

        return H, tau

    if n >= 2048:
        tile_size = 16

    if n > 2048:
        split_k_slices = 8

    T = H.new_empty(B, tile_size, tile_size)
    if n in {512}:
        split_k = False

    # these are the cases we scan for rankdef, clustered, or mixed
    cases_to_specialize = {512}

    if exploit_structure and n in cases_to_specialize:
        local_n_bounds = A.new_empty((B,), dtype=torch.long)
    else:
        local_n_bounds = A.new_empty((1,), dtype=torch.long)

    maybe_local_bounds = False
    structure_meta = None

    if exploit_structure and n in cases_to_specialize:
        rank = (3 * n) // 4
        cluster_bound = n // 2 + 2
        head_bound = n // 2 - 2
        clustered_bound = min(n, triton.cdiv(cluster_bound, tile_size) * tile_size)
        block_size_cols = triton.next_power_of_2(n - cluster_bound)
        structure_meta = A.new_zeros((2,), dtype=torch.long)

        grid = (B,)
        check_structure_kernel[grid](
            A,
            local_n_bounds,
            structure_meta,
            A.stride(0),
            A.stride(1),
            A.stride(2),
            local_n_bounds.stride(0),
            BLOCK_SIZE_ROWS=64,  # type: ignore
            BLOCK_SIZE_COLS=block_size_cols,  # type: ignore
            ROWS=n,  # type: ignore
            RANK_OFFSET=rank,  # type: ignore
            HEAD_BOUND=head_bound,  # type: ignore
            CLUSTER_BOUND=cluster_bound,  # type: ignore
            CLUSTERED_BOUND=clustered_bound,  # type: ignore
            num_warps=8,  # type: ignore
        )
        maybe_local_bounds = True

    max_n_bound = int(structure_meta[0].item()) if maybe_local_bounds else n

    specialize_constexpr = maybe_local_bounds

    for k in range(0, max_n_bound, tile_size):
        ib = min(tile_size, n - k)
        panel_end = k + ib
        is_first_panel = k == 0

        V = H.new_empty(B, n - k, ib)

        launch_tiled_qr_panel(
            H,
            tau,
            V,
            T,
            local_n_bounds,
            k,
            specialize_constexpr,
            tile_size,
            panel_source=A if is_first_panel else H,
        )

        if panel_end < n:
            C = H[:, k:, panel_end:n]
            C_source = A[:, k:, panel_end:n] if is_first_panel else C

            split_k_panel = False
            if split_k:
                trailing_n = max_n_bound - panel_end
                block_n_est = 128
                num_w_ctas = B * triton.cdiv(trailing_n, block_n_est)
                split_k_panel = split_k and (num_w_ctas < 300) and (trailing_n >= 512)

            W = bmm_then_t_tf32x3(
                V,
                C_source,
                T,
                local_n_bounds,
                panel_end,
                specialize_constexpr,
                split_k=split_k_panel,
                split_k_slices=split_k_slices,
            )

            baddbmm_sub_tf32x3(
                C,
                V,
                W,
                local_n_bounds,
                panel_end,
                specialize_constexpr,
                c_source=C_source,
            )

    if maybe_local_bounds and int(structure_meta[1].item()) != 0:
        cleanup_grid = (
            B,
            triton.cdiv(n, 16),
            triton.cdiv(n, 32),
        )
        cleanup_structure_tail_kernel[cleanup_grid](
            H,
            tau,
            local_n_bounds,
            H.stride(0),
            H.stride(1),
            H.stride(2),
            tau.stride(0),
            tau.stride(1),
            local_n_bounds.stride(0),
            n,
            BLOCK_SIZE_ROWS=16,  # type: ignore
            BLOCK_SIZE_COLS=32,  # type: ignore
            num_warps=4,  # type: ignore
        )

    return H, tau


def custom_kernel(data: input_t) -> output_t:
    return geqrf_blocked_batched(
        data,
        tile_size=32,
        split_k=True,
        split_k_slices=4,
        exploit_structure=True,
    )
scrolls · 1317 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