Skip to content
KernelIndex
Search⌘K

submission 417249

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417249?include=source"
interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA H100
40.7ms
#66 of 71
2026-01-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d56fc10559eaa965babd403f4ea41b816979c9542a8ec739662a1ce08ec904c6
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15

Kernel source

submission.py179 lines
from __future__ import annotations

from typing import Any, Dict, Tuple

import torch
import torch.nn.functional as F

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
from cutlass.cutlass_dsl import for_generate, if_generate, yield_out


class _OutgoingContractKernel:
    def __init__(self, threads: int = 256) -> None:
        self.threads = int(threads)

    @cute.jit
    def __call__(
        self,
        left_ptr: "cute.Pointer",
        right_ptr: "cute.Pointer",
        out_ptr: "cute.Pointer",
        problem: tuple,
    ):
        bs, n, hidden = problem
        stride_bs = n * n * hidden
        stride_i = n * hidden
        stride_k = hidden

        left = cute.make_tensor(
            left_ptr,
            cute.make_layout(
                (bs, n, n, hidden),
                stride=(stride_bs, stride_i, stride_k, 1),
            ),
        )
        right = cute.make_tensor(
            right_ptr,
            cute.make_layout(
                (bs, n, n, hidden),
                stride=(stride_bs, stride_i, stride_k, 1),
            ),
        )
        out = cute.make_tensor(
            out_ptr,
            cute.make_layout(
                (bs, n, n, hidden),
                stride=(stride_bs, stride_i, stride_k, 1),
            ),
        )

        total = bs * n * n * hidden
        grid = (total + self.threads - 1) // self.threads
        self.kernel(left, right, out, bs, n, hidden).launch(
            grid=[grid, 1, 1],
            block=[self.threads, 1, 1],
        )
        return

    @cute.kernel
    def kernel(
        self,
        left: "cute.Tensor",
        right: "cute.Tensor",
        out: "cute.Tensor",
        bs: int,
        n: int,
        hidden: int,
    ):
        tx, _, _ = cute.arch.thread_idx()
        bx, _, _ = cute.arch.block_idx()
        bdx, _, _ = cute.arch.block_dim()

        idx = bx * bdx + tx
        total = bs * n * n * hidden

        def _do_one():
            h = idx % hidden
            t0 = idx // hidden
            j = t0 % n
            t1 = t0 // n
            i = t1 % n
            b = t1 // n

            acc0 = cutlass.Float32(0.0)
            for k, acc, acc_out in for_generate(0, n, iter_args=[acc0]):
                acc = acc + left[b, i, k, h] * right[b, j, k, h]
                yield_out([acc])
            out[b, i, j, h] = acc_out

        if_generate(idx < total, _do_one)


_CONTRACT = _OutgoingContractKernel()
_CONTRACT_COMPILED = None


def _get_contract_compiled():
    global _CONTRACT_COMPILED
    if _CONTRACT_COMPILED is not None:
        return _CONTRACT_COMPILED
    left_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    right_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    _CONTRACT_COMPILED = cute.compile(
        _CONTRACT,
        left_ptr,
        right_ptr,
        out_ptr,
        (0, 0, 0),
        options="--opt-level 3",
    )
    return _CONTRACT_COMPILED


def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
    bs, n, _, hidden = left.shape

    if not (left.is_cuda and right.is_cuda):
        raise RuntimeError("This kernel requires CUDA tensors.")
    if left.dtype != torch.float32 or right.dtype != torch.float32:
        raise RuntimeError("This kernel expects float32 tensors.")

    left = left.contiguous()
    right = right.contiguous()

    out = torch.empty((bs, n, n, hidden), device=left.device, dtype=torch.float32)
    compiled = _get_contract_compiled()
    left_ptr = make_ptr(cutlass.Float32, left.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    right_ptr = make_ptr(cutlass.Float32, right.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    out_ptr = make_ptr(cutlass.Float32, out.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    compiled(left_ptr, right_ptr, out_ptr, (bs, n, hidden))
    return out


@torch.inference_mode()
def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:
    x, mask, weights, config = data

    dim = int(config["dim"])
    hidden_dim = int(config["hidden_dim"])

    if x.dtype != torch.float32:
        x = x.to(dtype=torch.float32)

    x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)

    left = F.linear(x, weights["left_proj.weight"], None)
    right = F.linear(x, weights["right_proj.weight"], None)

    mask_f = mask.unsqueeze(-1)
    if mask_f.dtype != left.dtype:
        mask_f = mask_f.to(dtype=left.dtype)
    left = left * mask_f
    right = right * mask_f

    left_gate = torch.sigmoid(F.linear(x, weights["left_gate.weight"], None))
    right_gate = torch.sigmoid(F.linear(x, weights["right_gate.weight"], None))
    out_gate = torch.sigmoid(F.linear(x, weights["out_gate.weight"], None))

    left = left * left_gate
    right = right * right_gate

    out = _contract_outgoing(left, right)
    out = F.layer_norm(
        out,
        (hidden_dim,),
        weights["to_out_norm.weight"],
        weights["to_out_norm.bias"],
        1e-5,
    )
    out = out * out_gate
    out = F.linear(out, weights["to_out.weight"], None)
    return out


__all__ = ["custom_kernel"]
scrolls · 179 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 417247.

Best evidence level for this revision: reported

JSON