Skip to content
KernelIndex
Search⌘K

submission 418500

shiyegao · 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.

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA A100
18.6ms
#46 of 69
2026-02-01

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

shared-memorysmem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)

Kernel source

submission.py1277 lines
from __future__ import annotations

from typing import Any, Dict, Tuple

import torch

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, range_constexpr








class _LayerNormLastDimF32ToF16:
    def __init__(self, threads: int = 256) -> None:
        self.threads = int(threads)
        self.warps = self.threads // 32

    @cute.jit
    def __call__(
        self,
        x_ptr: "cute.Pointer",
        w_ptr: "cute.Pointer",
        b_ptr: "cute.Pointer",
        y_ptr: "cute.Pointer",
        problem: tuple,
    ):
        bs, n, d = problem

        stride_bs = n * n * d
        stride_i = n * d
        stride_j = d

        x = cute.make_tensor(
            x_ptr,
            cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
        )
        w = cute.make_tensor(w_ptr, cute.make_layout((d,), stride=(1,)))
        b = cute.make_tensor(b_ptr, cute.make_layout((d,), stride=(1,)))
        y = cute.make_tensor(
            y_ptr,
            cute.make_layout((bs, n, n, d), stride=(stride_bs, stride_i, stride_j, 1)),
        )

        total = bs * n * n
        grid = (total + self.warps - 1) // self.warps

        self.kernel(x, w, b, y, bs, n, d).launch(
            grid=[grid, 1, 1],
            block=[self.threads, 1, 1],
        )
        return

    @cute.kernel
    def kernel(
        self,
        x: "cute.Tensor",
        w: "cute.Tensor",
        b: "cute.Tensor",
        y: "cute.Tensor",
        bs: int,
        n: int,
        d: int,
    ):
        tx, _, _ = cute.arch.thread_idx()
        bx, _, _ = cute.arch.block_idx()

        warp_id = tx >> 5
        lane = tx & 31

        idx = bx * self.warps + warp_id

        def _do_one():
            j = idx % n
            t0 = idx // n
            i = t0 % n
            bb = t0 // n

            sum0_init = cutlass.Float32(0.0)
            sumsq0_init = cutlass.Float32(0.0)
            for dd, acc, acc_out in for_generate(lane, d, 32, iter_args=[sum0_init, sumsq0_init]):
                sum0_it = acc[0]
                sumsq0_it = acc[1]
                v = x[bb, i, j, dd]
                sum0_it = sum0_it + v
                sumsq0_it = sumsq0_it + v * v
                yield_out([sum0_it, sumsq0_it])

            sum0 = cute.arch.warp_reduction_sum(acc_out[0])
            sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])

            inv_d = cutlass.Float32(1.0) / cutlass.Float32(d)
            mean = sum0 * inv_d
            var = sumsq0 * inv_d - mean * mean
            inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)

            for dd in for_generate(lane, d, 32):
                v = x[bb, i, j, dd]
                nrm = (v - mean) * inv_std
                out = nrm * w[dd] + b[dd]
                y[bb, i, j, dd] = out.to(cutlass.Float16)
                yield_out()

        if_generate(idx < bs * n * n, _do_one)









class _GemmF16F16ToF16:
    def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
        self.threads = int(threads)
        self.tile_m = int(tile_m)
        self.tile_n = int(tile_n)
        self.tile_k = int(tile_k)
        self.warps = self.threads // 32
        self.rows_per_warp = self.tile_m // self.warps

    @cute.jit
    def __call__(
        self,
        a_ptr: "cute.Pointer",
        b_ptr: "cute.Pointer",
        c_ptr: "cute.Pointer",
        problem: tuple,
    ):
        m, n, k = problem

        a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
        b = cute.make_tensor(b_ptr, cute.make_layout((n, k), stride=(k, 1)))
        c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))

        grid_n = (n + self.tile_n - 1) // self.tile_n
        grid_m = (m + self.tile_m - 1) // self.tile_m

        self.kernel(a, b, c, m, n, k).launch(
            grid=[grid_n, grid_m, 1],
            block=[self.threads, 1, 1],
        )
        return

    @cute.kernel
    def kernel(
        self,
        a: "cute.Tensor",
        b: "cute.Tensor",
        c: "cute.Tensor",
        m: int,
        n: int,
        k: int,
    ):
        tx, _, _ = cute.arch.thread_idx()
        bx, by, _ = cute.arch.block_idx()

        warp_id = tx >> 5
        lane = tx & 31

        base_m = by * self.tile_m
        base_n = bx * self.tile_n

        row0 = base_m + warp_id * self.rows_per_warp
        col0 = base_n + lane
        col1 = col0 + 32

        acc00 = cutlass.Float32(0.0)
        acc01 = cutlass.Float32(0.0)
        acc02 = cutlass.Float32(0.0)
        acc03 = cutlass.Float32(0.0)
        acc04 = cutlass.Float32(0.0)
        acc05 = cutlass.Float32(0.0)
        acc06 = cutlass.Float32(0.0)
        acc07 = cutlass.Float32(0.0)
        acc10 = cutlass.Float32(0.0)
        acc11 = cutlass.Float32(0.0)
        acc12 = cutlass.Float32(0.0)
        acc13 = cutlass.Float32(0.0)
        acc14 = cutlass.Float32(0.0)
        acc15 = cutlass.Float32(0.0)
        acc16 = cutlass.Float32(0.0)
        acc17 = cutlass.Float32(0.0)

        smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
        smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_n * (self.tile_k + 2), alignment=16)
        smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
        smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_n, self.tile_k + 2), stride=(self.tile_k + 2, 1)))

        for k0, acc, acc_out in for_generate(
            0,
            k,
            self.tile_k,
            iter_args=[
                acc00,
                acc01,
                acc02,
                acc03,
                acc04,
                acc05,
                acc06,
                acc07,
                acc10,
                acc11,
                acc12,
                acc13,
                acc14,
                acc15,
                acc16,
                acc17,
            ],
        ):
            acc00 = acc[0]
            acc01 = acc[1]
            acc02 = acc[2]
            acc03 = acc[3]
            acc04 = acc[4]
            acc05 = acc[5]
            acc06 = acc[6]
            acc07 = acc[7]
            acc10 = acc[8]
            acc11 = acc[9]
            acc12 = acc[10]
            acc13 = acc[11]
            acc14 = acc[12]
            acc15 = acc[13]
            acc16 = acc[14]
            acc17 = acc[15]

            for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
                mm = idx // self.tile_k
                kk = idx - mm * self.tile_k
                gm = base_m + mm
                gk = k0 + kk

                def _ld():
                    smem_a[mm, kk] = a[gm, gk]

                def _stz():
                    smem_a[mm, kk] = cutlass.Float16(0.0)

                if_generate((gm < m) & (gk < k), _ld, _stz)
                yield_out()

            for idx in for_generate(tx, self.tile_n * self.tile_k, self.threads):
                nn = idx // self.tile_k
                kk = idx - nn * self.tile_k
                gn = base_n + nn
                gk = k0 + kk

                def _ld():
                    smem_b[nn, kk] = b[gn, gk]

                def _stz():
                    smem_b[nn, kk] = cutlass.Float16(0.0)

                if_generate((gn < n) & (gk < k), _ld, _stz)
                yield_out()

            cute.arch.sync_threads()

            for kk in range_constexpr(128):
                bv0 = smem_b[lane, kk].to(cutlass.Float32)
                bv1 = smem_b[lane + 32, kk].to(cutlass.Float32)

                a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
                a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
                a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
                a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
                a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
                a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
                a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
                a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)

                acc00 = acc00 + a0 * bv0
                acc01 = acc01 + a1 * bv0
                acc02 = acc02 + a2 * bv0
                acc03 = acc03 + a3 * bv0
                acc04 = acc04 + a4 * bv0
                acc05 = acc05 + a5 * bv0
                acc06 = acc06 + a6 * bv0
                acc07 = acc07 + a7 * bv0

                acc10 = acc10 + a0 * bv1
                acc11 = acc11 + a1 * bv1
                acc12 = acc12 + a2 * bv1
                acc13 = acc13 + a3 * bv1
                acc14 = acc14 + a4 * bv1
                acc15 = acc15 + a5 * bv1
                acc16 = acc16 + a6 * bv1
                acc17 = acc17 + a7 * bv1

            cute.arch.sync_threads()
            yield_out(
                [
                    acc00,
                    acc01,
                    acc02,
                    acc03,
                    acc04,
                    acc05,
                    acc06,
                    acc07,
                    acc10,
                    acc11,
                    acc12,
                    acc13,
                    acc14,
                    acc15,
                    acc16,
                    acc17,
                ]
            )

        def _st_row(r: int, col: int, val: cutlass.Float32):
            gm = row0 + r
            if_generate((gm < m) & (col < n), lambda: c.__setitem__((gm, col), val.to(cutlass.Float16)))

        _st_row(0, col0, acc_out[0])
        _st_row(1, col0, acc_out[1])
        _st_row(2, col0, acc_out[2])
        _st_row(3, col0, acc_out[3])
        _st_row(4, col0, acc_out[4])
        _st_row(5, col0, acc_out[5])
        _st_row(6, col0, acc_out[6])
        _st_row(7, col0, acc_out[7])

        _st_row(0, col1, acc_out[8])
        _st_row(1, col1, acc_out[9])
        _st_row(2, col1, acc_out[10])
        _st_row(3, col1, acc_out[11])
        _st_row(4, col1, acc_out[12])
        _st_row(5, col1, acc_out[13])
        _st_row(6, col1, acc_out[14])
        _st_row(7, col1, acc_out[15])


class _GemmF16F16ToF32:
    def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
        self.threads = int(threads)
        self.tile_m = int(tile_m)
        self.tile_n = int(tile_n)
        self.tile_k = int(tile_k)
        self.warps = self.threads // 32
        self.rows_per_warp = self.tile_m // self.warps

    @cute.jit
    def __call__(
        self,
        a_ptr: "cute.Pointer",
        b_ptr: "cute.Pointer",
        c_ptr: "cute.Pointer",
        problem: tuple,
    ):
        m, n, k = problem

        a = cute.make_tensor(a_ptr, cute.make_layout((m, k), stride=(k, 1)))
        b = cute.make_tensor(b_ptr, cute.make_layout((k, n), stride=(n, 1)))
        c = cute.make_tensor(c_ptr, cute.make_layout((m, n), stride=(n, 1)))

        grid_n = (n + self.tile_n - 1) // self.tile_n
        grid_m = (m + self.tile_m - 1) // self.tile_m

        self.kernel(a, b, c, m, n, k).launch(
            grid=[grid_n, grid_m, 1],
            block=[self.threads, 1, 1],
        )
        return

    @cute.kernel
    def kernel(
        self,
        a: "cute.Tensor",
        b: "cute.Tensor",
        c: "cute.Tensor",
        m: int,
        n: int,
        k: int,
    ):
        tx, _, _ = cute.arch.thread_idx()
        bx, by, _ = cute.arch.block_idx()

        warp_id = tx >> 5
        lane = tx & 31

        base_m = by * self.tile_m
        base_n = bx * self.tile_n

        row0 = base_m + warp_id * self.rows_per_warp
        col0 = base_n + lane
        col1 = col0 + 32

        acc00 = cutlass.Float32(0.0)
        acc01 = cutlass.Float32(0.0)
        acc02 = cutlass.Float32(0.0)
        acc03 = cutlass.Float32(0.0)
        acc04 = cutlass.Float32(0.0)
        acc05 = cutlass.Float32(0.0)
        acc06 = cutlass.Float32(0.0)
        acc07 = cutlass.Float32(0.0)
        acc10 = cutlass.Float32(0.0)
        acc11 = cutlass.Float32(0.0)
        acc12 = cutlass.Float32(0.0)
        acc13 = cutlass.Float32(0.0)
        acc14 = cutlass.Float32(0.0)
        acc15 = cutlass.Float32(0.0)
        acc16 = cutlass.Float32(0.0)
        acc17 = cutlass.Float32(0.0)

        smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
        smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
        smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
        smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))

        for k0, acc, acc_out in for_generate(
            0,
            k,
            self.tile_k,
            iter_args=[
                acc00,
                acc01,
                acc02,
                acc03,
                acc04,
                acc05,
                acc06,
                acc07,
                acc10,
                acc11,
                acc12,
                acc13,
                acc14,
                acc15,
                acc16,
                acc17,
            ],
        ):
            acc00 = acc[0]
            acc01 = acc[1]
            acc02 = acc[2]
            acc03 = acc[3]
            acc04 = acc[4]
            acc05 = acc[5]
            acc06 = acc[6]
            acc07 = acc[7]
            acc10 = acc[8]
            acc11 = acc[9]
            acc12 = acc[10]
            acc13 = acc[11]
            acc14 = acc[12]
            acc15 = acc[13]
            acc16 = acc[14]
            acc17 = acc[15]

            for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
                mm = idx // self.tile_k
                kk = idx - mm * self.tile_k
                gm = base_m + mm
                gk = k0 + kk

                def _ld():
                    smem_a[mm, kk] = a[gm, gk]

                def _stz():
                    smem_a[mm, kk] = cutlass.Float16(0.0)

                if_generate((gm < m) & (gk < k), _ld, _stz)
                yield_out()

            for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
                kk = idx // self.tile_n
                nn = idx - kk * self.tile_n
                gk = k0 + kk
                gn = base_n + nn

                def _ld():
                    smem_b[kk, nn] = b[gk, gn]

                def _stz():
                    smem_b[kk, nn] = cutlass.Float16(0.0)

                if_generate((gk < k) & (gn < n), _ld, _stz)
                yield_out()

            cute.arch.sync_threads()

            for kk in range_constexpr(128):
                bv0 = smem_b[kk, lane].to(cutlass.Float32)
                bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)

                a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
                a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
                a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
                a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
                a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
                a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
                a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
                a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)

                acc00 = acc00 + a0 * bv0
                acc01 = acc01 + a1 * bv0
                acc02 = acc02 + a2 * bv0
                acc03 = acc03 + a3 * bv0
                acc04 = acc04 + a4 * bv0
                acc05 = acc05 + a5 * bv0
                acc06 = acc06 + a6 * bv0
                acc07 = acc07 + a7 * bv0

                acc10 = acc10 + a0 * bv1
                acc11 = acc11 + a1 * bv1
                acc12 = acc12 + a2 * bv1
                acc13 = acc13 + a3 * bv1
                acc14 = acc14 + a4 * bv1
                acc15 = acc15 + a5 * bv1
                acc16 = acc16 + a6 * bv1
                acc17 = acc17 + a7 * bv1

            cute.arch.sync_threads()
            yield_out(
                [
                    acc00,
                    acc01,
                    acc02,
                    acc03,
                    acc04,
                    acc05,
                    acc06,
                    acc07,
                    acc10,
                    acc11,
                    acc12,
                    acc13,
                    acc14,
                    acc15,
                    acc16,
                    acc17,
                ]
            )

        def _st_row(r: int, col: int, val: cutlass.Float32):
            gm = row0 + r
            if_generate((gm < m) & (col < n), lambda: c.__setitem__((gm, col), val))

        _st_row(0, col0, acc_out[0])
        _st_row(1, col0, acc_out[1])
        _st_row(2, col0, acc_out[2])
        _st_row(3, col0, acc_out[3])
        _st_row(4, col0, acc_out[4])
        _st_row(5, col0, acc_out[5])
        _st_row(6, col0, acc_out[6])
        _st_row(7, col0, acc_out[7])

        _st_row(0, col1, acc_out[8])
        _st_row(1, col1, acc_out[9])
        _st_row(2, col1, acc_out[10])
        _st_row(3, col1, acc_out[11])
        _st_row(4, col1, acc_out[12])
        _st_row(5, col1, acc_out[13])
        _st_row(6, col1, acc_out[14])
        _st_row(7, col1, acc_out[15])








def _sigmoid_f16(x: cutlass.Float16) -> cutlass.Float16:
    xx = x.to(cutlass.Float32)
    ee = cute.exp(cutlass.Float32(0.0) - xx, fastmath=True)
    yy = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + ee)
    return yy.to(cutlass.Float16)


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

    @cute.jit
    def __call__(
        self,
        proj_ptr: "cute.Pointer",
        mask_ptr: "cute.Pointer",
        left_ptr: "cute.Pointer",
        right_ptr: "cute.Pointer",
        gate_ptr: "cute.Pointer",
        problem: tuple,
    ):
        bs, n, h = problem
        stride_proj_bs = n * n * (5 * h)
        stride_proj_i = n * (5 * h)
        stride_proj_j = 5 * h

        proj = cute.make_tensor(
            proj_ptr,
            cute.make_layout(
                (bs, n, n, 5 * h),
                stride=(stride_proj_bs, stride_proj_i, stride_proj_j, 1),
            ),
        )
        mask = cute.make_tensor(mask_ptr, cute.make_layout((bs, n, n), stride=(n * n, n, 1)))
        left = cute.make_tensor(left_ptr, cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)))
        right = cute.make_tensor(right_ptr, cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)))
        gate = cute.make_tensor(gate_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))

        total = bs * n * n * h
        grid = (total + self.threads - 1) // self.threads

        self.kernel(proj, mask, left, right, gate, bs, n, h).launch(
            grid=[grid, 1, 1],
            block=[self.threads, 1, 1],
        )
        return

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

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

        def _do_one():
            hh = idx % h
            t0 = idx // h
            j = t0 % n
            t1 = t0 // n
            i = t1 % n
            bb = t1 // n

            m = mask[bb, i, j].to(cutlass.Float16)

            lp = proj[bb, i, j, hh]
            rp = proj[bb, i, j, hh + h]
            lg = proj[bb, i, j, hh + 2 * h]
            rg = proj[bb, i, j, hh + 3 * h]
            og = proj[bb, i, j, hh + 4 * h]

            gl = _sigmoid_f16(lg)
            gr = _sigmoid_f16(rg)
            go = _sigmoid_f16(og)

            left[bb, i, hh, j] = (lp * gl * m).to(cutlass.Float16)
            right[bb, hh, j, i] = (rp * gr * m).to(cutlass.Float16)
            gate[bb, i, j, hh] = go

        if_generate(idx < total, _do_one)









class _ContractHiddenGemm:
    def __init__(self, threads: int = 256, tile_m: int = 64, tile_n: int = 64, tile_k: int = 128) -> None:
        self.threads = int(threads)
        self.tile_m = int(tile_m)
        self.tile_n = int(tile_n)
        self.tile_k = int(tile_k)
        self.warps = self.threads // 32
        self.rows_per_warp = self.tile_m // self.warps

    @cute.jit
    def __call__(
        self,
        left_ptr: "cute.Pointer",
        right_ptr: "cute.Pointer",
        out_ptr: "cute.Pointer",
        problem: tuple,
    ):
        bs, n, h = problem

        left = cute.make_tensor(
            left_ptr,
            cute.make_layout((bs, n, h, n), stride=(n * h * n, h * n, n, 1)),
        )
        right = cute.make_tensor(
            right_ptr,
            cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
        )
        out = cute.make_tensor(
            out_ptr,
            cute.make_layout((bs, h, n, n), stride=(h * n * n, n * n, n, 1)),
        )

        grid_n = (n + self.tile_n - 1) // self.tile_n
        grid_m = (n + self.tile_m - 1) // self.tile_m
        grid_z = bs * h

        self.kernel(left, right, out, bs, n, h).launch(
            grid=[grid_n, grid_m, grid_z],
            block=[self.threads, 1, 1],
        )
        return

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

        warp_id = tx >> 5
        lane = tx & 31

        bb = bz // h
        hh = bz - bb * h

        base_m = by * self.tile_m
        base_n = bx * self.tile_n

        row0 = base_m + warp_id * self.rows_per_warp
        col0 = base_n + lane
        col1 = col0 + 32

        acc00 = cutlass.Float32(0.0)
        acc01 = cutlass.Float32(0.0)
        acc02 = cutlass.Float32(0.0)
        acc03 = cutlass.Float32(0.0)
        acc04 = cutlass.Float32(0.0)
        acc05 = cutlass.Float32(0.0)
        acc06 = cutlass.Float32(0.0)
        acc07 = cutlass.Float32(0.0)
        acc10 = cutlass.Float32(0.0)
        acc11 = cutlass.Float32(0.0)
        acc12 = cutlass.Float32(0.0)
        acc13 = cutlass.Float32(0.0)
        acc14 = cutlass.Float32(0.0)
        acc15 = cutlass.Float32(0.0)
        acc16 = cutlass.Float32(0.0)
        acc17 = cutlass.Float32(0.0)

        smem_a_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_m * self.tile_k, alignment=16)
        smem_b_ptr = cute.arch.alloc_smem(cutlass.Float16, self.tile_k * self.tile_n, alignment=16)
        smem_a = cute.make_tensor(smem_a_ptr, cute.make_layout((self.tile_m, self.tile_k), stride=(self.tile_k, 1)))
        smem_b = cute.make_tensor(smem_b_ptr, cute.make_layout((self.tile_k, self.tile_n), stride=(self.tile_n, 1)))

        for k0, acc, acc_out in for_generate(
            0,
            n,
            self.tile_k,
            iter_args=[
                acc00,
                acc01,
                acc02,
                acc03,
                acc04,
                acc05,
                acc06,
                acc07,
                acc10,
                acc11,
                acc12,
                acc13,
                acc14,
                acc15,
                acc16,
                acc17,
            ],
        ):
            acc00 = acc[0]
            acc01 = acc[1]
            acc02 = acc[2]
            acc03 = acc[3]
            acc04 = acc[4]
            acc05 = acc[5]
            acc06 = acc[6]
            acc07 = acc[7]
            acc10 = acc[8]
            acc11 = acc[9]
            acc12 = acc[10]
            acc13 = acc[11]
            acc14 = acc[12]
            acc15 = acc[13]
            acc16 = acc[14]
            acc17 = acc[15]

            for idx in for_generate(tx, self.tile_m * self.tile_k, self.threads):
                mm = idx // self.tile_k
                kk = idx - mm * self.tile_k
                gi = base_m + mm
                gk = k0 + kk

                def _ld():
                    smem_a[mm, kk] = left[bb, gi, hh, gk]

                def _stz():
                    smem_a[mm, kk] = cutlass.Float16(0.0)

                if_generate((gi < n) & (gk < n), _ld, _stz)
                yield_out()

            for idx in for_generate(tx, self.tile_k * self.tile_n, self.threads):
                kk = idx // self.tile_n
                nn = idx - kk * self.tile_n
                gk = k0 + kk
                gj = base_n + nn

                def _ld():
                    smem_b[kk, nn] = right[bb, hh, gk, gj]

                def _stz():
                    smem_b[kk, nn] = cutlass.Float16(0.0)

                if_generate((gk < n) & (gj < n), _ld, _stz)
                yield_out()

            cute.arch.sync_threads()

            for kk in range_constexpr(128):
                bv0 = smem_b[kk, lane].to(cutlass.Float32)
                bv1 = smem_b[kk, lane + 32].to(cutlass.Float32)

                a0 = smem_a[warp_id * self.rows_per_warp + 0, kk].to(cutlass.Float32)
                a1 = smem_a[warp_id * self.rows_per_warp + 1, kk].to(cutlass.Float32)
                a2 = smem_a[warp_id * self.rows_per_warp + 2, kk].to(cutlass.Float32)
                a3 = smem_a[warp_id * self.rows_per_warp + 3, kk].to(cutlass.Float32)
                a4 = smem_a[warp_id * self.rows_per_warp + 4, kk].to(cutlass.Float32)
                a5 = smem_a[warp_id * self.rows_per_warp + 5, kk].to(cutlass.Float32)
                a6 = smem_a[warp_id * self.rows_per_warp + 6, kk].to(cutlass.Float32)
                a7 = smem_a[warp_id * self.rows_per_warp + 7, kk].to(cutlass.Float32)

                acc00 = acc00 + a0 * bv0
                acc01 = acc01 + a1 * bv0
                acc02 = acc02 + a2 * bv0
                acc03 = acc03 + a3 * bv0
                acc04 = acc04 + a4 * bv0
                acc05 = acc05 + a5 * bv0
                acc06 = acc06 + a6 * bv0
                acc07 = acc07 + a7 * bv0

                acc10 = acc10 + a0 * bv1
                acc11 = acc11 + a1 * bv1
                acc12 = acc12 + a2 * bv1
                acc13 = acc13 + a3 * bv1
                acc14 = acc14 + a4 * bv1
                acc15 = acc15 + a5 * bv1
                acc16 = acc16 + a6 * bv1
                acc17 = acc17 + a7 * bv1

            cute.arch.sync_threads()
            yield_out(
                [
                    acc00,
                    acc01,
                    acc02,
                    acc03,
                    acc04,
                    acc05,
                    acc06,
                    acc07,
                    acc10,
                    acc11,
                    acc12,
                    acc13,
                    acc14,
                    acc15,
                    acc16,
                    acc17,
                ]
            )

        def _st_row(r: int, col: int, val: cutlass.Float32):
            gi = row0 + r
            if_generate((gi < n) & (col < n), lambda: out.__setitem__((bb, hh, gi, col), val))

        _st_row(0, col0, acc_out[0])
        _st_row(1, col0, acc_out[1])
        _st_row(2, col0, acc_out[2])
        _st_row(3, col0, acc_out[3])
        _st_row(4, col0, acc_out[4])
        _st_row(5, col0, acc_out[5])
        _st_row(6, col0, acc_out[6])
        _st_row(7, col0, acc_out[7])

        _st_row(0, col1, acc_out[8])
        _st_row(1, col1, acc_out[9])
        _st_row(2, col1, acc_out[10])
        _st_row(3, col1, acc_out[11])
        _st_row(4, col1, acc_out[12])
        _st_row(5, col1, acc_out[13])
        _st_row(6, col1, acc_out[14])
        _st_row(7, col1, acc_out[15])








class _LayerNormHiddenF32ToF16:
    def __init__(self, threads: int = 256) -> None:
        self.threads = int(threads)
        self.warps = self.threads // 32

    @cute.jit
    def __call__(
        self,
        x_ptr: "cute.Pointer",
        w_ptr: "cute.Pointer",
        b_ptr: "cute.Pointer",
        g_ptr: "cute.Pointer",
        y_ptr: "cute.Pointer",
        problem: tuple,
    ):
        bs, n, h = problem

        stride_bs = h * n * n
        stride_h = n * n
        stride_i = n

        x = cute.make_tensor(x_ptr, cute.make_layout((bs, h, n, n), stride=(stride_bs, stride_h, stride_i, 1)))
        w = cute.make_tensor(w_ptr, cute.make_layout((h,), stride=(1,)))
        b = cute.make_tensor(b_ptr, cute.make_layout((h,), stride=(1,)))
        g = cute.make_tensor(g_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))
        y = cute.make_tensor(y_ptr, cute.make_layout((bs, n, n, h), stride=(n * n * h, n * h, h, 1)))

        total = bs * n * n
        grid = (total + self.warps - 1) // self.warps

        self.kernel(x, w, b, g, y, bs, n, h).launch(
            grid=[grid, 1, 1],
            block=[self.threads, 1, 1],
        )
        return

    @cute.kernel
    def kernel(
        self,
        x: "cute.Tensor",
        w: "cute.Tensor",
        b: "cute.Tensor",
        g: "cute.Tensor",
        y: "cute.Tensor",
        bs: int,
        n: int,
        h: int,
    ):
        tx, _, _ = cute.arch.thread_idx()
        bx, _, _ = cute.arch.block_idx()

        warp_id = tx >> 5
        lane = tx & 31

        idx = bx * self.warps + warp_id

        def _do_one():
            j = idx % n
            t0 = idx // n
            i = t0 % n
            bb = t0 // n

            sum0_init = cutlass.Float32(0.0)
            sumsq0_init = cutlass.Float32(0.0)
            for hh, acc, acc_out in for_generate(lane, h, 32, iter_args=[sum0_init, sumsq0_init]):
                sum0_it = acc[0]
                sumsq0_it = acc[1]
                v = x[bb, hh, i, j]
                sum0_it = sum0_it + v
                sumsq0_it = sumsq0_it + v * v
                yield_out([sum0_it, sumsq0_it])

            sum0 = cute.arch.warp_reduction_sum(acc_out[0])
            sumsq0 = cute.arch.warp_reduction_sum(acc_out[1])

            inv_h = cutlass.Float32(1.0) / cutlass.Float32(h)
            mean = sum0 * inv_h
            var = sumsq0 * inv_h - mean * mean
            inv_std = cute.rsqrt(var + cutlass.Float32(1e-5), fastmath=True)

            for hh in for_generate(lane, h, 32):
                v = x[bb, hh, i, j]
                nrm = (v - mean) * inv_std
                out = (nrm * w[hh] + b[hh]).to(cutlass.Float16)
                y[bb, i, j, hh] = (out * g[bb, i, j, hh]).to(cutlass.Float16)
                yield_out()

        if_generate(idx < bs * n * n, _do_one)


class _FinalLinear:
    def __init__(self) -> None:
        self.gemm = _GemmF16F16ToF32()

    def __call__(self, a: torch.Tensor, w: torch.Tensor, y: torch.Tensor) -> None:
        m, k = a.shape
        n = y.shape[1]
        self.gemm(a, w, y, (m, n, k))







_LN_X = _LayerNormLastDimF32ToF16()
_LN_X_C = None

_GEMM_PROJ = _GemmF16F16ToF16()
_GEMM_PROJ_C = None

_PROC = _ProcessProj()
_PROC_C = None

_CONTRACT = _ContractHiddenGemm()
_CONTRACT_C = None

_LN_H = _LayerNormHiddenF32ToF16()
_LN_H_C = None

_GEMM_OUT = _GemmF16F16ToF32()
_GEMM_OUT_C = None


def _compile_once():
    global _LN_X_C, _GEMM_PROJ_C, _PROC_C, _CONTRACT_C, _LN_H_C, _GEMM_OUT_C

    if _LN_X_C is None:
        x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        w_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        _LN_X_C = cute.compile(_LN_X, x_ptr, w_ptr, b_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")

    if _GEMM_PROJ_C is None:
        a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        c_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        _GEMM_PROJ_C = cute.compile(_GEMM_PROJ, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")

    if _PROC_C is None:
        proj_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        mask_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        gate_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        _PROC_C = cute.compile(
            _PROC,
            proj_ptr,
            mask_ptr,
            left_ptr,
            right_ptr,
            gate_ptr,
            (0, 0, 0),
            options="--opt-level 3",
        )

    if _CONTRACT_C is None:
        left_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        right_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        out_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        _CONTRACT_C = cute.compile(_CONTRACT, left_ptr, right_ptr, out_ptr, (0, 0, 0), options="--opt-level 3")

    if _LN_H_C is None:
        x_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        w_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        g_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        y_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        _LN_H_C = cute.compile(_LN_H, x_ptr, w_ptr, b_ptr, g_ptr, y_ptr, (0, 0, 0), options="--opt-level 3")

    if _GEMM_OUT_C is None:
        a_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        b_ptr = make_ptr(cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=16)
        c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
        _GEMM_OUT_C = cute.compile(_GEMM_OUT, a_ptr, b_ptr, c_ptr, (0, 0, 0), options="--opt-level 3")


def _as_ptr(ty, t: torch.Tensor):
    return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)







_W_KEY = None
_W_PACK16 = None
_W_OUT_T16 = None
_W_NORM_W = None
_W_NORM_B = None
_W_TO_OUT_NORM_W = None
_W_TO_OUT_NORM_B = None

_BUF_KEY = None
_BUF = None


def _prep_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int, device: torch.device):
    global _W_KEY, _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B

    w_lp = weights["left_proj.weight"]
    w_rp = weights["right_proj.weight"]
    w_lg = weights["left_gate.weight"]
    w_rg = weights["right_gate.weight"]
    w_og = weights["out_gate.weight"]
    w_no_w = weights["norm.weight"]
    w_no_b = weights["norm.bias"]
    w_o_nw = weights["to_out_norm.weight"]
    w_o_nb = weights["to_out_norm.bias"]
    w_out = weights["to_out.weight"]

    key = (
        device,
        dim,
        hidden,
        int(w_lp.data_ptr()),
        int(w_rp.data_ptr()),
        int(w_lg.data_ptr()),
        int(w_rg.data_ptr()),
        int(w_og.data_ptr()),
        int(w_no_w.data_ptr()),
        int(w_no_b.data_ptr()),
        int(w_o_nw.data_ptr()),
        int(w_o_nb.data_ptr()),
        int(w_out.data_ptr()),
    )

    if key == _W_KEY:
        return _W_PACK16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B

    w_pack16 = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)
    w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()

    _W_PACK16 = w_pack16
    _W_OUT_T16 = w_out_t16
    _W_NORM_W = w_no_w.contiguous()
    _W_NORM_B = w_no_b.contiguous()
    _W_TO_OUT_NORM_W = w_o_nw.contiguous()
    _W_TO_OUT_NORM_B = w_o_nb.contiguous()
    _W_KEY = key

    return w_pack16, w_out_t16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B


def _get_buf(device: torch.device, bs: int, n: int, dim: int, hidden: int):
    global _BUF_KEY, _BUF

    key = (device, bs, n, dim, hidden)
    if key == _BUF_KEY and _BUF is not None:
        return _BUF

    m = bs * n * n
    buf = {
        "x_norm": torch.empty((bs, n, n, dim), device=device, dtype=torch.float16),
        "proj": torch.empty((m, 5 * hidden), device=device, dtype=torch.float16),
        "mask16": torch.empty((bs, n, n), device=device, dtype=torch.float16),
        "left_t": torch.empty((bs, n, hidden, n), device=device, dtype=torch.float16),
        "right_t": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float16),
        "out_gate": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
        "out_tmp": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float32),
        "out_norm": torch.empty((bs, n, n, hidden), device=device, dtype=torch.float16),
    }

    _BUF_KEY = key
    _BUF = buf
    return buf


@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

    if not x.is_cuda:
        raise RuntimeError("仅支持 CUDA 张量。")
    if x.dtype != torch.float32:
        x = x.to(torch.float32)
    x = x.contiguous()

    bs, n, n2, dim = x.shape
    if n != n2:
        raise RuntimeError("输入必须是 [bs, N, N, dim] 的方阵。")

    dim_cfg = int(config["dim"])
    hidden = int(config["hidden_dim"])
    if dim_cfg != dim:
        raise RuntimeError("config['dim'] 与 x.shape[-1] 不一致。")

    _compile_once()

    w_pack16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(
        weights, dim, hidden, x.device
    )
    buf = _get_buf(x.device, bs, n, dim, hidden)
    x_norm = buf["x_norm"]
    proj = buf["proj"]
    mask16 = buf["mask16"]
    left_t = buf["left_t"]
    right_t = buf["right_t"]
    out_gate = buf["out_gate"]
    out_tmp = buf["out_tmp"]
    out_norm = buf["out_norm"]

    _LN_X_C(
        _as_ptr(cutlass.Float32, x),
        _as_ptr(cutlass.Float32, w_norm_w),
        _as_ptr(cutlass.Float32, w_norm_b),
        _as_ptr(cutlass.Float16, x_norm),
        (bs, n, dim),
    )

    m = bs * n * n
    _GEMM_PROJ_C(
        _as_ptr(cutlass.Float16, x_norm.view(m, dim)),
        _as_ptr(cutlass.Float16, w_pack16),
        _as_ptr(cutlass.Float16, proj),
        (m, 5 * hidden, dim),
    )

    mask16.copy_(mask)
    _PROC_C(
        _as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),
        _as_ptr(cutlass.Float16, mask16),
        _as_ptr(cutlass.Float16, left_t),
        _as_ptr(cutlass.Float16, right_t),
        _as_ptr(cutlass.Float16, out_gate),
        (bs, n, hidden),
    )

    _CONTRACT_C(
        _as_ptr(cutlass.Float16, left_t),
        _as_ptr(cutlass.Float16, right_t),
        _as_ptr(cutlass.Float32, out_tmp),
        (bs, n, hidden),
    )

    _LN_H_C(
        _as_ptr(cutlass.Float32, out_tmp),
        _as_ptr(cutlass.Float32, w_out_norm_w),
        _as_ptr(cutlass.Float32, w_out_norm_b),
        _as_ptr(cutlass.Float16, out_gate),
        _as_ptr(cutlass.Float16, out_norm),
        (bs, n, hidden),
    )

    y = torch.empty((m, dim), device=x.device, dtype=torch.float32)
    _GEMM_OUT_C(
        _as_ptr(cutlass.Float16, out_norm.view(m, hidden)),
        _as_ptr(cutlass.Float16, w_out_t16),
        _as_ptr(cutlass.Float32, y),
        (m, dim, hidden),
    )

    return y.view(bs, n, n, dim)


__all__ = ["custom_kernel"]
scrolls · 1277 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 418358.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON