Skip to content
KernelIndex
Search⌘K

submission 417990

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417990?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
20.5ms
#57 of 69
2026-02-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:514f4e001f4516e1e3d3d713667b197b3f57f92f8a5271c81329b02a835107fb
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_stride = _TK + 1

Kernel source

submission.py345 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


_TMN = 32
_TK = 32
_THREADS = 256


class _BatchedGemmF32_32x32x32:
    def __init__(self) -> None:
        self.threads = _THREADS

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

        stride_batch = n * n
        stride_i = n

        a = cute.make_tensor(
            a_ptr,
            cute.make_layout(
                (batch, n, n),
                stride=(stride_batch, stride_i, 1),
            ),
        )
        b = cute.make_tensor(
            b_ptr,
            cute.make_layout(
                (batch, n, n),
                stride=(stride_batch, stride_i, 1),
            ),
        )
        c = cute.make_tensor(
            c_ptr,
            cute.make_layout(
                (batch, n, n),
                stride=(stride_batch, stride_i, 1),
            ),
        )

        if (n & (_TMN - 1)) == 0:
            grid_x = n // _TMN
            grid_y = n // _TMN
            self.kernel_fast(a, b, c, n).launch(
                grid=[grid_x, grid_y, batch],
                block=[self.threads, 1, 1],
            )
        else:
            grid_x = (n + _TMN - 1) // _TMN
            grid_y = (n + _TMN - 1) // _TMN
            self.kernel_bnd(a, b, c, n).launch(
                grid=[grid_x, grid_y, batch],
                block=[self.threads, 1, 1],
            )
        return

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

        tid = tx
        lane_m = tid >> 4
        lane_n = tid & 15

        base_m = by * _TMN
        base_n = bx * _TMN

        i0 = base_m + lane_m
        j0 = base_n + lane_n
        i1 = i0 + 16
        j1 = j0 + 16

        acc00 = cutlass.Float32(0.0)
        acc01 = cutlass.Float32(0.0)
        acc10 = cutlass.Float32(0.0)
        acc11 = cutlass.Float32(0.0)

        smem_stride = _TK + 1
        smem_a_elems = _TMN * smem_stride
        smem_b_elems = _TMN * smem_stride
        smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)

        sA = cute.make_tensor(
            smem_ptr,
            cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
        )
        sB = cute.make_tensor(
            smem_ptr + smem_a_elems,
            cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
        )

        for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
            0,
            n,
            _TK,
            iter_args=[acc00, acc01, acc10, acc11],
        ):
            base = tid << 2
            for t in range(4):
                idx = base + t
                row = idx >> 5
                col = idx & 31
                sA[row, col] = a[bz, base_m + row, k0 + col]
                sB[row, col] = b[bz, base_n + row, k0 + col]

            cute.arch.sync_threads()

            for kk in range(_TK):
                a0 = sA[lane_m, kk]
                a1 = sA[lane_m + 16, kk]
                b0 = sB[lane_n, kk]
                b1 = sB[lane_n + 16, kk]

                acc00 = acc00 + a0 * b0
                acc01 = acc01 + a0 * b1
                acc10 = acc10 + a1 * b0
                acc11 = acc11 + a1 * b1

            cute.arch.sync_threads()
            yield_out([acc00, acc01, acc10, acc11])

        c[bz, i0, j0] = acc00_out
        c[bz, i0, j1] = acc01_out
        c[bz, i1, j0] = acc10_out
        c[bz, i1, j1] = acc11_out

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

        tid = tx
        lane_m = tid >> 4
        lane_n = tid & 15

        base_m = by * _TMN
        base_n = bx * _TMN

        i0 = base_m + lane_m
        j0 = base_n + lane_n
        i1 = i0 + 16
        j1 = j0 + 16

        acc00 = cutlass.Float32(0.0)
        acc01 = cutlass.Float32(0.0)
        acc10 = cutlass.Float32(0.0)
        acc11 = cutlass.Float32(0.0)

        smem_stride = _TK + 1
        smem_a_elems = _TMN * smem_stride
        smem_b_elems = _TMN * smem_stride
        smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)

        sA = cute.make_tensor(
            smem_ptr,
            cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
        )
        sB = cute.make_tensor(
            smem_ptr + smem_a_elems,
            cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
        )

        for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
            0,
            n,
            _TK,
            iter_args=[acc00, acc01, acc10, acc11],
        ):
            base = tid << 2
            for t in range(4):
                idx = base + t
                row = idx >> 5
                col = idx & 31

                gi_a = base_m + row
                gk_a = k0 + col
                gj_b = base_n + row
                gk_b = k0 + col

                def _ld_a():
                    sA[row, col] = a[bz, gi_a, gk_a]

                def _ze_a():
                    sA[row, col] = cutlass.Float32(0.0)

                def _ld_b():
                    sB[row, col] = b[bz, gj_b, gk_b]

                def _ze_b():
                    sB[row, col] = cutlass.Float32(0.0)

                if_generate((gi_a < n) & (gk_a < n), _ld_a, _ze_a)
                if_generate((gj_b < n) & (gk_b < n), _ld_b, _ze_b)

            cute.arch.sync_threads()

            for kk in range(_TK):
                a0 = sA[lane_m, kk]
                a1 = sA[lane_m + 16, kk]
                b0 = sB[lane_n, kk]
                b1 = sB[lane_n + 16, kk]

                acc00 = acc00 + a0 * b0
                acc01 = acc01 + a0 * b1
                acc10 = acc10 + a1 * b0
                acc11 = acc11 + a1 * b1

            cute.arch.sync_threads()
            yield_out([acc00, acc01, acc10, acc11])

        def _st00():
            c[bz, i0, j0] = acc00_out

        def _st01():
            c[bz, i0, j1] = acc01_out

        def _st10():
            c[bz, i1, j0] = acc10_out

        def _st11():
            c[bz, i1, j1] = acc11_out

        if_generate((i0 < n) & (j0 < n), _st00)
        if_generate((i0 < n) & (j1 < n), _st01)
        if_generate((i1 < n) & (j0 < n), _st10)
        if_generate((i1 < n) & (j1 < n), _st11)


_CONTRACT = _BatchedGemmF32_32x32x32()
_CONTRACT_C = None


def _get_contract_compiled():
    global _CONTRACT_C
    if _CONTRACT_C is not None:
        return _CONTRACT_C

    a_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)
    c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
    _CONTRACT_C = cute.compile(
        _CONTRACT,
        a_ptr,
        b_ptr,
        c_ptr,
        (0, 0),
        options="--opt-level 3",
    )
    return _CONTRACT_C


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("该实现仅支持 CUDA 张量。")
    if left.dtype != torch.float32 or right.dtype != torch.float32:
        raise RuntimeError("该实现期望 left/right 为 float32。")

    
    left_t = left.permute(0, 3, 1, 2).contiguous()
    right_t = right.permute(0, 3, 1, 2).contiguous()

    bh = bs * hidden
    a = left_t.view(bh, n, n)
    b = right_t.view(bh, n, n)
    c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)

    compiled = _get_contract_compiled()
    a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    compiled(a_ptr, b_ptr, c_ptr, (bh, n))

    out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
    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 not x.is_cuda:
        raise RuntimeError("该实现要求 x 在 CUDA 上。")

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

    try:
        torch.backends.cuda.matmul.allow_tf32 = True
        torch.backends.cudnn.allow_tf32 = True
    except Exception:
        pass

    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.mul_(mask_f)
    right.mul_(mask_f)

    left_gate = F.linear(x, weights["left_gate.weight"], None)
    right_gate = F.linear(x, weights["right_gate.weight"], None)
    out_gate = F.linear(x, weights["out_gate.weight"], None)
    left_gate.sigmoid_()
    right_gate.sigmoid_()
    out_gate.sigmoid_()
    left.mul_(left_gate)
    right.mul_(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.mul_(out_gate)
    out = F.linear(out, weights["to_out.weight"], None)
    return out


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

⋯ 1 unchanged lines
from typing import Any, Dict, Tuple
- import os
-
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
+ _TMN = 32
+ _TK = 32
+ _THREADS = 256
- _EXT = None
- _EXT_LOCK = None
+ class _BatchedGemmF32_32x32x32:
+ def __init__(self) -> None:
+ self.threads = _THREADS
+ @cute.jit
+ def __call__(
+ self,
+ a_ptr: "cute.Pointer",
+ b_ptr: "cute.Pointer",
+ c_ptr: "cute.Pointer",
+ problem: tuple,
+ ):
+ batch, n = problem
- def _lazy_import_extension_utils():
- from torch.utils.cpp_extension import load_inline
+ stride_batch = n * n
+ stride_i = n
- return load_inline
+ a = cute.make_tensor(
+ a_ptr,
+ cute.make_layout(
+ (batch, n, n),
+ stride=(stride_batch, stride_i, 1),
+ ),
+ )
+ b = cute.make_tensor(
+ b_ptr,
+ cute.make_layout(
+ (batch, n, n),
+ stride=(stride_batch, stride_i, 1),
+ ),
+ )
+ c = cute.make_tensor(
+ c_ptr,
+ cute.make_layout(
+ (batch, n, n),
+ stride=(stride_batch, stride_i, 1),
+ ),
+ )
+ if (n & (_TMN - 1)) == 0:
+ grid_x = n // _TMN
+ grid_y = n // _TMN
+ self.kernel_fast(a, b, c, n).launch(
+ grid=[grid_x, grid_y, batch],
+ block=[self.threads, 1, 1],
+ )
+ else:
+ grid_x = (n + _TMN - 1) // _TMN
+ grid_y = (n + _TMN - 1) // _TMN
+ self.kernel_bnd(a, b, c, n).launch(
+ grid=[grid_x, grid_y, batch],
+ block=[self.threads, 1, 1],
+ )
+ return
- def _get_ext():
- global _EXT, _EXT_LOCK
- if _EXT is not None:
- return _EXT
- if _EXT_LOCK is None:
- import threading
+ @cute.kernel
+ def kernel_fast(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):
+ tx, _, _ = cute.arch.thread_idx()
+ bx, by, bz = cute.arch.block_idx()
- _EXT_LOCK = threading.Lock()
- with _EXT_LOCK:
- if _EXT is not None:
- return _EXT
+ tid = tx
+ lane_m = tid >> 4
+ lane_n = tid & 15
- load_inline = _lazy_import_extension_utils()
+ base_m = by * _TMN
+ base_n = bx * _TMN
-
- if "TORCH_CUDA_ARCH_LIST" not in os.environ:
- os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"
+ i0 = base_m + lane_m
+ j0 = base_n + lane_n
+ i1 = i0 + 16
+ j1 = j0 + 16
- cpp_src = r"""
- #include <torch/extension.h>
+ acc00 = cutlass.Float32(0.0)
+ acc01 = cutlass.Float32(0.0)
+ acc10 = cutlass.Float32(0.0)
+ acc11 = cutlass.Float32(0.0)
- torch::Tensor trimul_fwd(
- torch::Tensor x,
- torch::Tensor mask,
- torch::Tensor ln1_w,
- torch::Tensor ln1_b,
- torch::Tensor w_cat,
- torch::Tensor ln2_w,
- torch::Tensor ln2_b,
- torch::Tensor w_out,
- int64_t dim,
- int64_t hidden);
+ smem_stride = _TK + 1
+ smem_a_elems = _TMN * smem_stride
+ smem_b_elems = _TMN * smem_stride
+ smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)
- PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
- m.def("fwd", &trimul_fwd, "trimul forward (cuda)");
- }
- """
+ sA = cute.make_tensor(
+ smem_ptr,
+ cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
+ )
+ sB = cute.make_tensor(
+ smem_ptr + smem_a_elems,
+ cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
+ )
- cuda_src = r"""
- #include <torch/extension.h>
- #include <cuda.h>
- #include <cuda_fp16.h>
- #include <cublas_v2.h>
+ for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
+ 0,
+ n,
+ _TK,
+ iter_args=[acc00, acc01, acc10, acc11],
+ ):
+ base = tid << 2
+ for t in range(4):
+ idx = base + t
+ row = idx >> 5
+ col = idx & 31
+ sA[row, col] = a[bz, base_m + row, k0 + col]
+ sB[row, col] = b[bz, base_n + row, k0 + col]
- #include <mutex>
+ cute.arch.sync_threads()
- namespace {
+ for kk in range(_TK):
+ a0 = sA[lane_m, kk]
+ a1 = sA[lane_m + 16, kk]
+ b0 = sB[lane_n, kk]
+ b1 = sB[lane_n + 16, kk]
- static inline void checkCuda(cudaError_t e, const char* msg) {
- if (e != cudaSuccess) {
- throw std::runtime_error(std::string(msg) + ": " + cudaGetErrorString(e));
- }
- }
+ acc00 = acc00 + a0 * b0
+ acc01 = acc01 + a0 * b1
+ acc10 = acc10 + a1 * b0
+ acc11 = acc11 + a1 * b1
- static inline void checkCublas(cublasStatus_t s, const char* msg) {
- if (s != CUBLAS_STATUS_SUCCESS) {
- throw std::runtime_error(std::string(msg) + ": cublas status=" + std::to_string((int)s));
- }
- }
+ cute.arch.sync_threads()
+ yield_out([acc00, acc01, acc10, acc11])
- struct CublasHandleHolder {
- cublasHandle_t handle = nullptr;
- CublasHandleHolder() {
- checkCublas(cublasCreate(&handle), "cublasCreate");
- // 强制开启 Tensor Core 路线(不依赖环境默认值)
- cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH);
- }
- ~CublasHandleHolder() {
- if (handle) {
- cublasDestroy(handle);
- handle = nullptr;
- }
- }
- };
+ c[bz, i0, j0] = acc00_out
+ c[bz, i0, j1] = acc01_out
+ c[bz, i1, j0] = acc10_out
+ c[bz, i1, j1] = acc11_out
- static CublasHandleHolder* get_cublas() {
- static std::once_flag once;
- static CublasHandleHolder* holder = nullptr;
- std::call_once(once, []() { holder = new CublasHandleHolder(); });
- return holder;
- }
+ @cute.kernel
+ def kernel_bnd(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):
+ tx, _, _ = cute.arch.thread_idx()
+ bx, by, bz = cute.arch.block_idx()
- __device__ __forceinline__ float warp_sum(float v) {
- for (int d = 16; d > 0; d >>= 1) {
- v += __shfl_down_sync(0xffffffff, v, d);
- }
- return v;
- }
+ tid = tx
+ lane_m = tid >> 4
+ lane_n = tid & 15
- __device__ __forceinline__ float fast_sigmoid(float x) {
- // --use_fast_math 下的 __expf 通常足够快;精度仍能满足题面容忍度
- float z = __expf(-x);
- return 1.0f / (1.0f + z);
- }
+ base_m = by * _TMN
+ base_n = bx * _TMN
- // ---------------- LN1 ----------------
+ i0 = base_m + lane_m
+ j0 = base_n + lane_n
+ i1 = i0 + 16
+ j1 = j0 + 16
- __global__ void ln1_128_f16(
- const float* __restrict__ x,
- const float* __restrict__ w,
- const float* __restrict__ b,
- half* __restrict__ y,
- int64_t rows) {
- int64_t row = (int64_t)blockIdx.x;
- if (row >= rows) return;
- int lane = (int)threadIdx.x; // 0..31
+ acc00 = cutlass.Float32(0.0)
+ acc01 = cutlass.Float32(0.0)
+ acc10 = cutlass.Float32(0.0)
+ acc11 = cutlass.Float32(0.0)
- const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);
- float4 v = x4[lane];
- float s = v.x + v.y + v.z + v.w;
- float ss = v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;
- s = warp_sum(s);
- ss = warp_sum(ss);
- s = __shfl_sync(0xffffffff, s, 0);
- ss = __shfl_sync(0xffffffff, ss, 0);
- float mean = s * (1.0f / 128.0f);
- float var = ss * (1.0f / 128.0f) - mean * mean;
- float inv = rsqrtf(var + 1e-5f);
+ smem_stride = _TK + 1
+ smem_a_elems = _TMN * smem_stride
+ smem_b_elems = _TMN * smem_stride
+ smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)
- const float4* w4 = reinterpret_cast<const float4*>(w);
- const float4* b4 = reinterpret_cast<const float4*>(b);
- float4 gw = w4[lane];
- float4 gb = b4[lane];
+ sA = cute.make_tensor(
+ smem_ptr,
+ cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
+ )
+ sB = cute.make_tensor(
+ smem_ptr + smem_a_elems,
+ cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
+ )
- float y0 = (v.x - mean) * inv * gw.x + gb.x;
- float y1 = (v.y - mean) * inv * gw.y + gb.y;
- float y2 = (v.z - mean) * inv * gw.z + gb.z;
- float y3 = (v.w - mean) * inv * gw.w + gb.w;
+ for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
+ 0,
+ n,
+ _TK,
+ iter_args=[acc00, acc01, acc10, acc11],
+ ):
+ base = tid << 2
+ for t in range(4):
+ idx = base + t
+ row = idx >> 5
+ col = idx & 31
- half2 h0 = __floats2half2_rn(y0, y1);
- half2 h1 = __floats2half2_rn(y2, y3);
+ gi_a = base_m + row
+ gk_a = k0 + col
+ gj_b = base_n + row
+ gk_b = k0 + col
- half2* y2p = reinterpret_cast<half2*>(y + row * 128 + lane * 4);
- y2p[0] = h0;
- y2p[1] = h1;
- }
+ def _ld_a():
+ sA[row, col] = a[bz, gi_a, gk_a]
- __global__ void ln1_generic_f16(
- const float* __restrict__ x,
- const float* __restrict__ w,
- const float* __restrict__ b,
- half* __restrict__ y,
- int dim,
- int64_t rows) {
- int64_t row = (int64_t)blockIdx.x;
- if (row >= rows) return;
+ def _ze_a():
+ sA[row, col] = cutlass.Float32(0.0)
- float sum = 0.0f;
- float sq = 0.0f;
- int64_t base = row * (int64_t)dim;
- for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {
- float v = x[base + c];
- sum += v;
- sq += v * v;
- }
+ def _ld_b():
+ sB[row, col] = b[bz, gj_b, gk_b]
- __shared__ float shm_sum[256];
- __shared__ float shm_sq[256];
- int t = (int)threadIdx.x;
- shm_sum[t] = sum;
- shm_sq[t] = sq;
- __syncthreads();
+ def _ze_b():
+ sB[row, col] = cutlass.Float32(0.0)
- for (int stride = ((int)blockDim.x) / 2; stride > 0; stride >>= 1) {
- if (t < stride) {
- shm_sum[t] += shm_sum[t + stride];
- shm_sq[t] += shm_sq[t + stride];
- }
- __syncthreads();
- }
+ if_generate((gi_a < n) & (gk_a < n), _ld_a, _ze_a)
+ if_generate((gj_b < n) & (gk_b < n), _ld_b, _ze_b)
- float mean = shm_sum[0] / (float)dim;
- float var = shm_sq[0] / (float)dim - mean * mean;
- float inv = rsqrtf(var + 1e-5f);
+ cute.arch.sync_threads()
- for (int c = (int)threadIdx.x; c < dim; c += (int)blockDim.x) {
- float v = x[base + c];
- float yv = (v - mean) * inv * w[c] + b[c];
- y[base + c] = __float2half_rn(yv);
- }
- }
+ for kk in range(_TK):
+ a0 = sA[lane_m, kk]
+ a1 = sA[lane_m + 16, kk]
+ b0 = sB[lane_n, kk]
+ b1 = sB[lane_n + 16, kk]
- static void launch_ln1(torch::Tensor x, torch::Tensor w, torch::Tensor b, torch::Tensor y) {
- int dim = (int)x.size(1);
- auto rows = x.size(0);
- if (dim == 128) {
- dim3 block(32, 1, 1);
- dim3 grid((unsigned)rows, 1, 1);
- ln1_128_f16<<<grid, block>>>(
- (const float*)x.data_ptr(),
- (const float*)w.data_ptr(),
- (const float*)b.data_ptr(),
- (half*)y.data_ptr(),
- (int64_t)rows);
- checkCuda(cudaGetLastError(), "ln1_128_f16");
- } else {
- dim3 block(256, 1, 1);
- dim3 grid((unsigned)rows, 1, 1);
- ln1_generic_f16<<<grid, block>>>(
- (const float*)x.data_ptr(),
- (const float*)w.data_ptr(),
- (const float*)b.data_ptr(),
- (half*)y.data_ptr(),
- dim,
- (int64_t)rows);
- checkCuda(cudaGetLastError(), "ln1_generic_f16");
- }
- }
+ acc00 = acc00 + a0 * b0
+ acc01 = acc01 + a0 * b1
+ acc10 = acc10 + a1 * b0
+ acc11 = acc11 + a1 * b1
- // --------------- pack (proj -> left/right/ogate) ---------------
- // 目标:同时满足
- // - 读:proj 为 [p, d] row-major(d 连续),用 32×32 tile 合并读取
- // - 写:left/right/og 为 [d, p](p 连续),用 shared 转置后合并写回
- // 额外融合:mask + sigmoid + gate,减少全局访存与 kernel 数量。
+ cute.arch.sync_threads()
+ yield_out([acc00, acc01, acc10, acc11])
- __global__ void pack_proj_p32_f16(
- const half* __restrict__ proj, // [M, 5H]
- const float* __restrict__ mask, // [M]
- half* __restrict__ left, // [B*H, nn]
- half* __restrict__ right, // [B*H, nn]
- half* __restrict__ ogate, // [B*H, nn]
- int64_t nn,
- int hidden) {
- constexpr int P_TILE = 32;
- constexpr int D_TILE = 32;
- constexpr int BLOCK_ROWS = 8;
+ def _st00():
+ c[bz, i0, j0] = acc00_out
- int b = (int)blockIdx.y;
- int64_t p_base = (int64_t)blockIdx.x * (int64_t)P_TILE;
+ def _st01():
+ c[bz, i0, j1] = acc01_out
- int tx = (int)threadIdx.x; // 0..31
- int ty = (int)threadIdx.y; // 0..BLOCK_ROWS-1
+ def _st10():
+ c[bz, i1, j0] = acc10_out
- __shared__ float m_sh[P_TILE];
- if (ty == 0) {
- int64_t p = p_base + (int64_t)tx;
- float mv = 0.0f;
- if (p < nn) {
- mv = mask[(int64_t)b * nn + p];
- }
- m_sh[tx] = mv;
- }
+ def _st11():
+ c[bz, i1, j1] = acc11_out
- // 共享内存第二维 +1 padding,避免 bank conflict
- __shared__ half sh_l[P_TILE][D_TILE + 1];
- __shared__ half sh_r[P_TILE][D_TILE + 1];
- __shared__ half sh_g[P_TILE][D_TILE + 1];
+ if_generate((i0 < n) & (j0 < n), _st00)
+ if_generate((i0 < n) & (j1 < n), _st01)
+ if_generate((i1 < n) & (j0 < n), _st10)
+ if_generate((i1 < n) & (j1 < n), _st11)
- __syncthreads();
- int out_ch = hidden * 5;
+ _CONTRACT = _BatchedGemmF32_32x32x32()
+ _CONTRACT_C = None
- // 逐块处理 d 维(每次 32 个通道),每块内做一次转置写回
- for (int d0 = 0; d0 < hidden; d0 += D_TILE) {
- // load+compute:写入 shared[p_local][d_local]
- #pragma unroll
- for (int pj = 0; pj < P_TILE; pj += BLOCK_ROWS) {
- int p_l = ty + pj; // 0..31
- int d = d0 + tx; // 真实通道
- int64_t p = p_base + (int64_t)p_l;
- half hl = __float2half_rn(0.0f);
- half hr = __float2half_rn(0.0f);
- half hg = __float2half_rn(0.0f);
+ def _get_contract_compiled():
+ global _CONTRACT_C
+ if _CONTRACT_C is not None:
+ return _CONTRACT_C
- if (p < nn && d < hidden) {
- int64_t row = (int64_t)b * nn + p; // 0..M-1
- const half* base = proj + row * (int64_t)out_ch;
+ a_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)
+ c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
+ _CONTRACT_C = cute.compile(
+ _CONTRACT,
+ a_ptr,
+ b_ptr,
+ c_ptr,
+ (0, 0),
+ options="--opt-level 3",
+ )
+ return _CONTRACT_C
- float l = __half2float(base[d]);
- float r = __half2float(base[hidden + d]);
- float gl = fast_sigmoid(__half2float(base[2 * hidden + d]));
- float gr = fast_sigmoid(__half2float(base[3 * hidden + d]));
- float go = fast_sigmoid(__half2float(base[4 * hidden + d]));
- float m = m_sh[p_l];
- float l2 = l * gl * m;
- float r2 = r * gr * m;
+ def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
+ bs, n, _, hidden = left.shape
- hl = __float2half_rn(l2);
- hr = __float2half_rn(r2);
- hg = __float2half_rn(go);
- }
+ if not (left.is_cuda and right.is_cuda):
+ raise RuntimeError("该实现仅支持 CUDA 张量。")
+ if left.dtype != torch.float32 or right.dtype != torch.float32:
+ raise RuntimeError("该实现期望 left/right 为 float32。")
- sh_l[p_l][tx] = hl;
- sh_r[p_l][tx] = hr;
- sh_g[p_l][tx] = hg;
- }
+
+ left_t = left.permute(0, 3, 1, 2).contiguous()
+ right_t = right.permute(0, 3, 1, 2).contiguous()
- __syncthreads();
+ bh = bs * hidden
+ a = left_t.view(bh, n, n)
+ b = right_t.view(bh, n, n)
+ c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)
- // store:读 shared 转置,写回到 [d, p]
- #pragma unroll
- for (int dj = 0; dj < D_TILE; dj += BLOCK_ROWS) {
- int d_l = ty + dj; // 0..31(tile 内通道)
- int d = d0 + d_l; // 真实通道
- int64_t p = p_base + (int64_t)tx; // tile 内位置
- if (p < nn && d < hidden) {
- int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d) * nn + p;
- left[out_idx] = sh_l[tx][d_l];
- right[out_idx] = sh_r[tx][d_l];
- ogate[out_idx] = sh_g[tx][d_l];
- }
- }
+ compiled = _get_contract_compiled()
+ a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
+ b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
+ c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
+ compiled(a_ptr, b_ptr, c_ptr, (bh, n))
- __syncthreads();
- }
- }
+ out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
+ return out
- static void launch_pack(
- torch::Tensor proj,
- torch::Tensor mask,
- torch::Tensor left,
- torch::Tensor right,
- torch::Tensor og,
- int bs,
- int n,
- int hidden) {
- int64_t nn = (int64_t)n * (int64_t)n;
- if (hidden > 128) {
- throw std::runtime_error("hidden_dim too large");
- }
- constexpr int P_TILE = 32;
- constexpr int BLOCK_ROWS = 8;
- dim3 block(32, BLOCK_ROWS, 1); // 256 threads
- dim3 grid((unsigned)((nn + P_TILE - 1) / P_TILE), (unsigned)bs, 1);
- pack_proj_p32_f16<<<grid, block>>>(
- (const half*)proj.data_ptr(),
- (const float*)mask.data_ptr(),
- (half*)left.data_ptr(),
- (half*)right.data_ptr(),
- (half*)og.data_ptr(),
- nn,
- hidden);
- checkCuda(cudaGetLastError(), "pack_proj_p32_f16");
- }
-
- // --------------- LN2 + gate + store ---------------
- // 输入 out_acc / ogate 为 [B*H, nn];输出 out_norm 为 [B*nn, H](连续 H 维)。
- //
- // 与 node28 的差异:
- // - out_acc 改为 half 存储(仍由 GEMM 以 fp32 累加生成),降低超大中间张量的 HBM 带宽压力。
-
- template<int MAX_H>
- __global__ void ln2_gate_store_p32_f16(
- const half* __restrict__ out_acc, // [B*H, nn]
- const half* __restrict__ ogate, // [B*H, nn]
- const float* __restrict__ w, // [H]
- const float* __restrict__ b, // [H]
- half* __restrict__ out_norm, // [B*nn, H]
- int64_t nn,
- int hidden) {
- constexpr int P_TILE = 32;
- constexpr int PAD = 1;
-
- int bb = (int)blockIdx.y;
- int64_t p_base = (int64_t)blockIdx.x * (int64_t)P_TILE;
-
- int tid = (int)threadIdx.x; // 0..255
- int lane = tid & 31;
- int warp = tid >> 5;
- int num_warp = (int)(blockDim.x >> 5);
-
- // shared:按 [d][p] 存(p 连续),从全局读取时形成 128B 合并事务
- __shared__ float sh_x[MAX_H][P_TILE + PAD];
- __shared__ half sh_g[MAX_H][P_TILE + PAD];
- __shared__ float mean_sh[P_TILE];
- __shared__ float inv_sh[P_TILE];
- __shared__ float w_sh[MAX_H];
- __shared__ float b_sh[MAX_H];
-
- if (tid < MAX_H) {
- if (tid < hidden) {
- w_sh[tid] = w[tid];
- b_sh[tid] = b[tid];
- } else {
- w_sh[tid] = 0.0f;
- b_sh[tid] = 0.0f;
- }
- }
-
- // 每个 warp 负责多个 d(步长=warp 数),lane 对应 p_local(0..31)
- for (int d = warp; d < hidden; d += num_warp) {
- int64_t p = p_base + (int64_t)lane;
- float x = 0.0f;
- half g = __float2half_rn(0.0f);
- if (p < nn) {
- int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;
- x = __half2float(out_acc[idx]);
- g = ogate[idx];
- }
- sh_x[d][lane] = x;
- sh_g[d][lane] = g;
- }
-
- __syncthreads();
-
- // 计算每个 p 的 mean/var(沿 hidden 归约)。只用一个 warp 处理 32 个 p。
- if (warp == 0) {
- int p_l = lane; // 0..31
- float sum = 0.0f;
- float sq = 0.0f;
- if (p_base + (int64_t)p_l < nn) {
- for (int d = 0; d < hidden; ++d) {
- float v = sh_x[d][p_l];
- sum += v;
- sq += v * v;
- }
- }
- float inv_n = 1.0f / (float)hidden;
- float mean = sum * inv_n;
- float var = sq * inv_n - mean * mean;
- mean_sh[p_l] = mean;
- inv_sh[p_l] = rsqrtf(var + 1e-5f);
- }
-
- __syncthreads();
-
- // 写回 out_norm:按 (p,d) 线性遍历,保证 d 连续写
- for (int e = tid; e < MAX_H * P_TILE; e += (int)blockDim.x) {
- int p_l = e / MAX_H; // 0..31
- int d = e - p_l * MAX_H; // 0..MAX_H-1
- int64_t p = p_base + (int64_t)p_l;
- if (d < hidden && p < nn) {
- float x = sh_x[d][p_l];
- float mean = mean_sh[p_l];
- float inv = inv_sh[p_l];
- float y = (x - mean) * inv * w_sh[d] + b_sh[d];
- float g = __half2float(sh_g[d][p_l]);
- half out_h = __float2half_rn(y * g);
- int64_t row = (int64_t)bb * nn + p;
- out_norm[row * (int64_t)hidden + (int64_t)d] = out_h;
- }
- }
- }
-
- static void launch_ln2(
- torch::Tensor out_acc,
- torch::Tensor og,
- torch::Tensor w,
- torch::Tensor b,
- torch::Tensor out_norm,
- int bs,
- int n,
- int hidden) {
- int64_t nn = (int64_t)n * (int64_t)n;
- if (hidden > 128) {
- throw std::runtime_error("hidden_dim too large");
- }
-
- constexpr int P_TILE = 32;
- dim3 block(256, 1, 1);
- dim3 grid((unsigned)((nn + P_TILE - 1) / P_TILE), (unsigned)bs, 1);
-
- if (hidden <= 32) {
- ln2_gate_store_p32_f16<32><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
- (const half*)og.data_ptr(),
- (const float*)w.data_ptr(),
- (const float*)b.data_ptr(),
- (half*)out_norm.data_ptr(),
- nn,
- hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_32");
- } else if (hidden <= 64) {
- ln2_gate_store_p32_f16<64><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
- (const half*)og.data_ptr(),
- (const float*)w.data_ptr(),
- (const float*)b.data_ptr(),
- (half*)out_norm.data_ptr(),
- nn,
- hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_64");
- } else {
- ln2_gate_store_p32_f16<128><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
- (const half*)og.data_ptr(),
- (const float*)w.data_ptr(),
- (const float*)b.data_ptr(),
- (half*)out_norm.data_ptr(),
- nn,
- hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_128");
- }
- }
-
- // ---------------- GEMM helpers ----------------
-
- static void gemm_x_wt_f16_f16(
- cublasHandle_t h,
- const half* x_row, // row-major [M,K]
- const half* w_row, // row-major [N,K]
- half* y_row, // row-major [M,N]
- int64_t M,
- int64_t N,
- int64_t K) {
- float alpha = 1.0f;
- float beta = 0.0f;
- checkCublas(
- cublasGemmEx(
- h,
- CUBLAS_OP_T, CUBLAS_OP_N,
- (int)N, (int)M, (int)K,
- &alpha,
- w_row, CUDA_R_16F, (int)K,
- x_row, CUDA_R_16F, (int)K,
- &beta,
- y_row, CUDA_R_16F, (int)N,
- CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
- "cublasGemmEx");
- }
-
- static void gemm_contract_batched_f16_f16(
- cublasHandle_t h,
- const half* left_row, // row-major [B,M,K]
- const half* right_row, // row-major [B,N,K]
- half* out_row, // row-major [B,M,N]
- int batch,
- int n) {
- float alpha = 1.0f;
- float beta = 0.0f;
- long long strideA = (long long)n * (long long)n;
- long long strideB = (long long)n * (long long)n;
- long long strideC = (long long)n * (long long)n;
-
- checkCublas(
- cublasGemmStridedBatchedEx(
- h,
- CUBLAS_OP_T, CUBLAS_OP_N,
- n, n, n,
- &alpha,
- right_row, CUDA_R_16F, n, strideB,
- left_row, CUDA_R_16F, n, strideA,
- &beta,
- out_row, CUDA_R_16F, n, strideC,
- batch,
- CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
- "cublasGemmStridedBatchedEx");
- }
-
- } // namespace
-
- torch::Tensor trimul_fwd(
- torch::Tensor x,
- torch::Tensor mask,
- torch::Tensor ln1_w,
- torch::Tensor ln1_b,
- torch::Tensor w_cat,
- torch::Tensor ln2_w,
- torch::Tensor ln2_b,
- torch::Tensor w_out,
- int64_t dim,
- int64_t hidden) {
- if (!x.is_cuda() || !mask.is_cuda()) {
- throw std::runtime_error("cuda only");
- }
- if (x.scalar_type() != torch::kFloat32) {
- throw std::runtime_error("x must be float32");
- }
- if (mask.scalar_type() != torch::kFloat32) {
- throw std::runtime_error("mask must be float32");
- }
- if (dim != x.size(3)) {
- throw std::runtime_error("dim mismatch");
- }
- if (hidden <= 0 || hidden > 128) {
- throw std::runtime_error("hidden_dim invalid");
- }
-
- int bs = (int)x.size(0);
- int n = (int)x.size(1);
- int64_t nn = (int64_t)n * (int64_t)n;
- int64_t M = (int64_t)bs * nn;
-
- auto h = get_cublas()->handle;
-
- // LN1: x[M,dim] -> x_norm[M,dim] half
- auto x2 = x.view({M, dim});
- auto x_norm = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- launch_ln1(x2, ln1_w, ln1_b, x_norm);
-
- // gemm1: proj = x_norm @ w_cat^T, 输出 half [M,5H]
- int out_ch = (int)hidden * 5;
- auto proj = torch::empty({M, out_ch}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- gemm_x_wt_f16_f16(h, (const half*)x_norm.data_ptr(), (const half*)w_cat.data_ptr(), (half*)proj.data_ptr(), M, out_ch, dim);
-
- // left/right/ogate: [bs, hidden, n, n] -> 实际内存等价于 [bs*hidden, nn]
- auto left = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- auto right = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- auto og = torch::empty({bs, (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
-
- // pack:proj[M,5H] + mask[M] -> left/right/og[bs*H,nn]
- launch_pack(proj, mask.view({M}), left, right, og, bs, n, (int)hidden);
-
- auto left3 = left.view({bs * (int)hidden, n, n});
- auto right3 = right.view({bs * (int)hidden, n, n});
-
- // contraction:输出 half(仍 fp32 累加),显著降低后续 LN2 的带宽压力
- auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- gemm_contract_batched_f16_f16(
- h,
- (const half*)left3.data_ptr(),
- (const half*)right3.data_ptr(),
- (half*)out_acc.data_ptr(),
- bs * (int)hidden, n);
-
- // LN2 + gate:输出为 [M, hidden] half,适配后续 GEMM2
- auto out_norm = torch::empty({M, hidden}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- launch_ln2(out_acc, og.view({bs * (int)hidden, n, n}), ln2_w, ln2_b, out_norm, bs, n, (int)hidden);
-
- // gemm2: y = out_norm @ w_out^T,输出也改为 half(精度在题面阈值内显式折衷)
- auto y = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- gemm_x_wt_f16_f16(h, (const half*)out_norm.data_ptr(), (const half*)w_out.data_ptr(), (half*)y.data_ptr(), M, dim, hidden);
-
- return y.view({bs, n, n, dim});
- }
- """
-
- name = "trimul_ext_mod4"
-
- extra_cuda_cflags = [
- "-O3",
- "--use_fast_math",
- ]
- extra_cflags = [
- "-O3",
- ]
-
- extra_ldflags = [
- "-lcublas",
- ]
-
- _EXT = load_inline(
- name=name,
- cpp_sources=cpp_src,
- cuda_sources=cuda_src,
- functions=None,
- extra_cflags=extra_cflags,
- extra_cuda_cflags=extra_cuda_cflags,
- extra_ldflags=extra_ldflags,
- with_cuda=True,
- verbose=False,
- )
- return _EXT
-
-
- class _WeightCache:
- __slots__ = ("key", "w_cat", "w_out")
-
- def __init__(self) -> None:
- self.key = None
- self.w_cat = None
- self.w_out = None
-
-
- _W_CACHE = _WeightCache()
-
-
- def _prepare_weights(weights: Dict[str, torch.Tensor], dim: int, hidden: int):
- k = (
- int(weights["left_proj.weight"].data_ptr()),
- int(weights["right_proj.weight"].data_ptr()),
- int(weights["left_gate.weight"].data_ptr()),
- int(weights["right_gate.weight"].data_ptr()),
- int(weights["out_gate.weight"].data_ptr()),
- int(weights["to_out.weight"].data_ptr()),
- )
- if _W_CACHE.key == k and _W_CACHE.w_cat is not None and _W_CACHE.w_out is not None:
- return _W_CACHE.w_cat, _W_CACHE.w_out
-
- w_left = weights["left_proj.weight"]
- w_right = weights["right_proj.weight"]
- w_lg = weights["left_gate.weight"]
- w_rg = weights["right_gate.weight"]
- w_og = weights["out_gate.weight"]
- w_out = weights["to_out.weight"]
-
- if w_left.shape != (hidden, dim):
- raise RuntimeError("left_proj.weight shape mismatch")
- if w_right.shape != (hidden, dim):
- raise RuntimeError("right_proj.weight shape mismatch")
- if w_lg.shape != (hidden, dim):
- raise RuntimeError("left_gate.weight shape mismatch")
- if w_rg.shape != (hidden, dim):
- raise RuntimeError("right_gate.weight shape mismatch")
- if w_og.shape != (hidden, dim):
- raise RuntimeError("out_gate.weight shape mismatch")
- if w_out.shape != (dim, hidden):
- raise RuntimeError("to_out.weight shape mismatch")
-
- w_cat = torch.cat([w_left, w_right, w_lg, w_rg, w_og], dim=0).contiguous().to(dtype=torch.float16)
- w_out_h = w_out.contiguous().to(dtype=torch.float16)
-
- _W_CACHE.key = k
- _W_CACHE.w_cat = w_cat
- _W_CACHE.w_out = w_out_h
- return w_cat, w_out_h
-
-
@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 = int(config["hidden_dim"])
+ hidden_dim = int(config["hidden_dim"])
if not x.is_cuda:
- raise RuntimeError("x must be CUDA tensor")
- if not mask.is_cuda:
- raise RuntimeError("mask must be CUDA tensor")
+ raise RuntimeError("该实现要求 x 在 CUDA 上。")
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
- if mask.dtype != torch.float32:
- mask = mask.to(dtype=torch.float32)
- x = x.contiguous()
- mask = mask.contiguous()
+ try:
+ torch.backends.cuda.matmul.allow_tf32 = True
+ torch.backends.cudnn.allow_tf32 = True
+ except Exception:
+ pass
- if x.ndim != 4:
- raise RuntimeError("x must be 4D")
- if mask.ndim != 3:
- raise RuntimeError("mask must be 3D")
- if x.shape[:3] != mask.shape:
- raise RuntimeError("x/mask shape mismatch")
- if x.shape[3] != dim:
- raise RuntimeError("dim mismatch")
+ x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
- for k in (
- "norm.weight",
- "norm.bias",
- "left_proj.weight",
- "right_proj.weight",
- "left_gate.weight",
- "right_gate.weight",
- "out_gate.weight",
- "to_out_norm.weight",
- "to_out_norm.bias",
- "to_out.weight",
- ):
- if not weights[k].is_cuda:
- raise RuntimeError(f"weight {k} must be CUDA tensor")
- if weights[k].dtype != torch.float32:
- raise RuntimeError(f"weight {k} must be float32")
+ left = F.linear(x, weights["left_proj.weight"], None)
+ right = F.linear(x, weights["right_proj.weight"], None)
- ln1_w = weights["norm.weight"].contiguous()
- ln1_b = weights["norm.bias"].contiguous()
- if ln1_w.shape != (dim,) or ln1_b.shape != (dim,):
- raise RuntimeError("norm params shape mismatch")
+ mask_f = mask.unsqueeze(-1)
+ if mask_f.dtype != left.dtype:
+ mask_f = mask_f.to(dtype=left.dtype)
+ left.mul_(mask_f)
+ right.mul_(mask_f)
- ln2_w = weights["to_out_norm.weight"].contiguous()
- ln2_b = weights["to_out_norm.bias"].contiguous()
- if ln2_w.shape != (hidden,) or ln2_b.shape != (hidden,):
- raise RuntimeError("to_out_norm params shape mismatch")
+ left_gate = F.linear(x, weights["left_gate.weight"], None)
+ right_gate = F.linear(x, weights["right_gate.weight"], None)
+ out_gate = F.linear(x, weights["out_gate.weight"], None)
+ left_gate.sigmoid_()
+ right_gate.sigmoid_()
+ out_gate.sigmoid_()
+ left.mul_(left_gate)
+ right.mul_(right_gate)
- w_cat, w_out = _prepare_weights(weights, dim, hidden)
+ out = _contract_outgoing(left, right)
- ext = _get_ext()
- return ext.fwd(x, mask, ln1_w, ln1_b, w_cat, ln2_w, ln2_b, w_out, dim, hidden)
+ out = F.layer_norm(out, (hidden_dim,), weights["to_out_norm.weight"], weights["to_out_norm.bias"], 1e-5)
+ out.mul_(out_gate)
+ out = F.linear(out, weights["to_out.weight"], None)
+ return out
__all__ = ["custom_kernel"]
-
scrolls · 1047 diff lines total

Best evidence level for this revision: reported

JSON