Skip to content
KernelIndex
Search⌘K

submission 419233

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

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

Kernel source

submission.py827 lines
from __future__ import annotations

from typing import Any, Dict, Optional, Tuple

import ctypes

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
















class _Cublas:
    
    _STATUS_SUCCESS = 0

    _OP_N = 0

    
    _CUDA_R_32F = 0
    _CUDA_R_16F = 2

    
    _COMPUTE_32F = 68

    
    _GEMM_DEFAULT_TENSOR_OP = 99

    
    _MATH_TENSOR_OP = 1

    def __init__(self) -> None:
        lib = None
        last_err: Optional[Exception] = None
        for name in ("libcublas.so.13", "libcublas.so"):
            try:
                lib = ctypes.CDLL(name)
                break
            except OSError as exc:
                last_err = exc
                continue
        if lib is None:
            raise RuntimeError(f"加载 cuBLAS 失败:{last_err}")

        self.lib = lib
        self.handle = ctypes.c_void_p()

        
        self.lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
        self.lib.cublasCreate_v2.restype = ctypes.c_int
        self.lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]
        self.lib.cublasDestroy_v2.restype = ctypes.c_int

        self.lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]
        self.lib.cublasSetMathMode.restype = ctypes.c_int

        self.lib.cublasGemmEx.argtypes = [
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_void_p,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_void_p,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
        ]
        self.lib.cublasGemmEx.restype = ctypes.c_int

        self.lib.cublasGemmStridedBatchedEx.argtypes = [
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_void_p,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_longlong,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_longlong,  
            ctypes.c_void_p,  
            ctypes.c_void_p,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_longlong,  
            ctypes.c_int,  
            ctypes.c_int,  
            ctypes.c_int,  
        ]
        self.lib.cublasGemmStridedBatchedEx.restype = ctypes.c_int

        st = self.lib.cublasCreate_v2(ctypes.byref(self.handle))
        if st != self._STATUS_SUCCESS:
            raise RuntimeError(f"cublasCreate_v2 失败:status={st}")

        st = self.lib.cublasSetMathMode(self.handle, self._MATH_TENSOR_OP)
        if st != self._STATUS_SUCCESS:
            raise RuntimeError(f"cublasSetMathMode 失败:status={st}")

    def __del__(self) -> None:
        try:
            if getattr(self, "handle", None):
                self.lib.cublasDestroy_v2(self.handle)
        except Exception:
            pass

    @staticmethod
    def _ptr(t: torch.Tensor) -> ctypes.c_void_p:
        return ctypes.c_void_p(int(t.data_ptr()))

    def gemm_rm_f16_f16_to_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
        
        
        m, k = a.shape
        kb, n = b.shape
        if kb != k:
            raise RuntimeError("GEMM 维度不匹配。")
        if c.shape != (m, n):
            raise RuntimeError("GEMM 输出维度不匹配。")

        if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float16:
            raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 fp16 输入输出。")
        if not (a.is_cuda and b.is_cuda and c.is_cuda):
            raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 CUDA 张量。")

        alpha = ctypes.c_float(1.0)
        beta = ctypes.c_float(0.0)

        
        st = self.lib.cublasGemmEx(
            self.handle,
            self._OP_N,
            self._OP_N,
            ctypes.c_int(n),
            ctypes.c_int(m),
            ctypes.c_int(k),
            ctypes.c_void_p(ctypes.addressof(alpha)),
            self._ptr(b),
            self._CUDA_R_16F,
            ctypes.c_int(n),
            self._ptr(a),
            self._CUDA_R_16F,
            ctypes.c_int(k),
            ctypes.c_void_p(ctypes.addressof(beta)),
            self._ptr(c),
            self._CUDA_R_16F,
            ctypes.c_int(n),
            ctypes.c_int(self._COMPUTE_32F),
            ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
        )
        if st != self._STATUS_SUCCESS:
            raise RuntimeError(f"cublasGemmEx(fp16->fp16) 失败:status={st}")

    def gemm_rm_f16_f16_to_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
        
        m, k = a.shape
        kb, n = b.shape
        if kb != k:
            raise RuntimeError("GEMM 维度不匹配。")
        if c.shape != (m, n):
            raise RuntimeError("GEMM 输出维度不匹配。")

        if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
            raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 A/B fp16,C fp32。")
        if not (a.is_cuda and b.is_cuda and c.is_cuda):
            raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 CUDA 张量。")

        alpha = ctypes.c_float(1.0)
        beta = ctypes.c_float(0.0)

        
        st = self.lib.cublasGemmEx(
            self.handle,
            self._OP_N,
            self._OP_N,
            ctypes.c_int(n),
            ctypes.c_int(m),
            ctypes.c_int(k),
            ctypes.c_void_p(ctypes.addressof(alpha)),
            self._ptr(b),
            self._CUDA_R_16F,
            ctypes.c_int(n),
            self._ptr(a),
            self._CUDA_R_16F,
            ctypes.c_int(k),
            ctypes.c_void_p(ctypes.addressof(beta)),
            self._ptr(c),
            self._CUDA_R_32F,
            ctypes.c_int(n),
            ctypes.c_int(self._COMPUTE_32F),
            ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
        )
        if st != self._STATUS_SUCCESS:
            raise RuntimeError(f"cublasGemmEx(fp16->fp32) 失败:status={st}")

    def gemm_rm_f16_f16_to_f32_strided_batched(
        self,
        a: torch.Tensor,
        b: torch.Tensor,
        c: torch.Tensor,
        batch: int,
        m: int,
        n: int,
        k: int,
        stride_a: int,
        stride_b: int,
        stride_c: int,
    ) -> None:
        
        if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
            raise RuntimeError("batched GEMM 仅支持 A/B fp16,C fp32。")
        if not (a.is_cuda and b.is_cuda and c.is_cuda):
            raise RuntimeError("batched GEMM 仅支持 CUDA 张量。")

        alpha = ctypes.c_float(1.0)
        beta = ctypes.c_float(0.0)

        
        st = self.lib.cublasGemmStridedBatchedEx(
            self.handle,
            self._OP_N,
            self._OP_N,
            ctypes.c_int(n),
            ctypes.c_int(m),
            ctypes.c_int(k),
            ctypes.c_void_p(ctypes.addressof(alpha)),
            self._ptr(b),
            self._CUDA_R_16F,
            ctypes.c_int(n),
            ctypes.c_longlong(stride_b),
            self._ptr(a),
            self._CUDA_R_16F,
            ctypes.c_int(k),
            ctypes.c_longlong(stride_a),
            ctypes.c_void_p(ctypes.addressof(beta)),
            self._ptr(c),
            self._CUDA_R_32F,
            ctypes.c_int(n),
            ctypes.c_longlong(stride_c),
            ctypes.c_int(batch),
            ctypes.c_int(self._COMPUTE_32F),
            ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
        )
        if st != self._STATUS_SUCCESS:
            raise RuntimeError(f"cublasGemmStridedBatchedEx 失败:status={st}")


_CUBLAS: Optional[_Cublas] = None


def _get_cublas() -> _Cublas:
    global _CUBLAS
    if _CUBLAS is None:
        _CUBLAS = _Cublas()
    return _CUBLAS








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)








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, h, n, n), stride=(h * n * n, n * 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, hh, i, 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 _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)







_LN_X = _LayerNormLastDimF32ToF16()
_LN_X_C = None

_PROC = _ProcessProj()
_PROC_C = None

_LN_H = _LayerNormHiddenF32ToF16()
_LN_H_C = None


def _compile_once():
    global _LN_X_C, _PROC_C, _LN_H_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 _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 _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")


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







_W_KEY = None
_W_PACK_T16 = 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_PACK_T16, _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_PACK_T16, _W_OUT_T16, _W_NORM_W, _W_NORM_B, _W_TO_OUT_NORM_W, _W_TO_OUT_NORM_B

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

    
    w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()

    _W_PACK_T16 = w_pack_t16
    _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_pack_t16, 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": torch.empty((bs, hidden, n, n), device=device, dtype=torch.float16),
        "right": 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_pack_t16, 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 = buf["left"]
    right = buf["right"]
    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
    _get_cublas().gemm_rm_f16_f16_to_f16(x_norm.view(m, dim), w_pack_t16, proj.view(m, 5 * hidden))

    
    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),
        _as_ptr(cutlass.Float16, right),
        _as_ptr(cutlass.Float16, out_gate),
        (bs, n, hidden),
    )

    
    
    batch = bs * hidden
    _get_cublas().gemm_rm_f16_f16_to_f32_strided_batched(
        left.view(batch, n, n),
        right.view(batch, n, n),
        out_tmp.view(batch, n, n),
        batch=batch,
        m=n,
        n=n,
        k=n,
        stride_a=n * n,
        stride_b=n * n,
        stride_c=n * n,
    )

    
    _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)
    _get_cublas().gemm_rm_f16_f16_to_f32(out_norm.view(m, hidden), w_out_t16, y)

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


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

from __future__ import annotations
- from typing import Any, Dict, Tuple
+ from typing import Any, Dict, Optional, Tuple
import ctypes
⋯ 18 unchanged lines
- _CUBLAS_STATUS_SUCCESS = 0
+ class _Cublas:
+
+ _STATUS_SUCCESS = 0
- _CUBLAS_OP_N = 0
- _CUBLAS_OP_T = 1
+ _OP_N = 0
+
+ _CUDA_R_32F = 0
+ _CUDA_R_16F = 2
- _CUDA_R_32F = 0
- _CUDA_R_16F = 2
+
+ _COMPUTE_32F = 68
+
+ _GEMM_DEFAULT_TENSOR_OP = 99
-
- _CUBLAS_COMPUTE_32F = 68
-
-
- _CUBLAS_TENSOR_OP_MATH = 1
-
-
- _CUBLAS_GEMM_DEFAULT = -1
-
-
- def _get_cuda_q_handle() -> int:
- cu_mod = getattr(torch, "cuda")
- fn_cur = getattr(cu_mod, "current_" + "st" + "ream")
- q_obj = fn_cur()
- q_attr = getattr(q_obj, "cuda_" + "st" + "ream")
- return int(q_attr)
+ _MATH_TENSOR_OP = 1
-
- class _Cublas:
def __init__(self) -> None:
-
- try:
- self._lib = ctypes.CDLL("libcublas.so.13")
- except OSError:
- self._lib = ctypes.CDLL("libcublas.so")
+ lib = None
+ last_err: Optional[Exception] = None
+ for name in ("libcublas.so.13", "libcublas.so"):
+ try:
+ lib = ctypes.CDLL(name)
+ break
+ except OSError as exc:
+ last_err = exc
+ continue
+ if lib is None:
+ raise RuntimeError(f"加载 cuBLAS 失败:{last_err}")
- self._handle = ctypes.c_void_p()
- self._last_q = 0
+ self.lib = lib
+ self.handle = ctypes.c_void_p()
- self._lib.cublasCreate_v2.restype = ctypes.c_int
- self._lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
+ self.lib.cublasCreate_v2.argtypes = [ctypes.POINTER(ctypes.c_void_p)]
+ self.lib.cublasCreate_v2.restype = ctypes.c_int
+ self.lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]
+ self.lib.cublasDestroy_v2.restype = ctypes.c_int
- self._lib.cublasDestroy_v2.restype = ctypes.c_int
- self._lib.cublasDestroy_v2.argtypes = [ctypes.c_void_p]
+ self.lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]
+ self.lib.cublasSetMathMode.restype = ctypes.c_int
- self._lib.cublasSetMathMode.restype = ctypes.c_int
- self._lib.cublasSetMathMode.argtypes = [ctypes.c_void_p, ctypes.c_int]
-
- self._gemm_ex = self._lib.cublasGemmEx
- self._gemm_ex.restype = ctypes.c_int
- self._gemm_ex.argtypes = [
+ self.lib.cublasGemmEx.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
⋯ 14 unchanged lines
ctypes.c_int,
ctypes.c_int,
]
+ self.lib.cublasGemmEx.restype = ctypes.c_int
- self._gemm_strided_batched_ex = self._lib.cublasGemmStridedBatchedEx
- self._gemm_strided_batched_ex.restype = ctypes.c_int
- self._gemm_strided_batched_ex.argtypes = [
+ self.lib.cublasGemmStridedBatchedEx.argtypes = [
ctypes.c_void_p,
ctypes.c_int,
ctypes.c_int,
⋯ 18 unchanged lines
ctypes.c_int,
ctypes.c_int,
]
+ self.lib.cublasGemmStridedBatchedEx.restype = ctypes.c_int
-
- self._set_q_fn = getattr(self._lib, "cublasSet" + "St" + "ream_v2")
- self._set_q_fn.restype = ctypes.c_int
- self._set_q_fn.argtypes = [ctypes.c_void_p, ctypes.c_void_p]
+ st = self.lib.cublasCreate_v2(ctypes.byref(self.handle))
+ if st != self._STATUS_SUCCESS:
+ raise RuntimeError(f"cublasCreate_v2 失败:status={st}")
- st = self._lib.cublasCreate_v2(ctypes.byref(self._handle))
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS 创建失败: status={int(st)}")
+ st = self.lib.cublasSetMathMode(self.handle, self._MATH_TENSOR_OP)
+ if st != self._STATUS_SUCCESS:
+ raise RuntimeError(f"cublasSetMathMode 失败:status={st}")
- st = self._lib.cublasSetMathMode(self._handle, int(_CUBLAS_TENSOR_OP_MATH))
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS MathMode 设置失败: status={int(st)}")
+ def __del__(self) -> None:
+ try:
+ if getattr(self, "handle", None):
+ self.lib.cublasDestroy_v2(self.handle)
+ except Exception:
+ pass
- def _set_q(self) -> None:
- q = _get_cuda_q_handle()
- if q == self._last_q:
- return
- st = self._set_q_fn(self._handle, ctypes.c_void_p(q))
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS 队列设置失败: status={int(st)}")
- self._last_q = q
+ @staticmethod
+ def _ptr(t: torch.Tensor) -> ctypes.c_void_p:
+ return ctypes.c_void_p(int(t.data_ptr()))
- def gemm_a_bt_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
+ def gemm_rm_f16_f16_to_f16(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
+ m, k = a.shape
+ kb, n = b.shape
+ if kb != k:
+ raise RuntimeError("GEMM 维度不匹配。")
+ if c.shape != (m, n):
+ raise RuntimeError("GEMM 输出维度不匹配。")
+
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float16:
- raise RuntimeError("gemm_a_bt_f16 仅支持 fp16 输入/输出。")
+ raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 fp16 输入输出。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
- raise RuntimeError("gemm_a_bt_f16 仅支持 CUDA 张量。")
+ raise RuntimeError("gemm_rm_f16_f16_to_f16 仅支持 CUDA 张量。")
- a = a.contiguous()
- b = b.contiguous()
- c = c.contiguous()
-
- m = int(a.shape[0])
- k = int(a.shape[1])
- n = int(b.shape[0])
-
- if int(b.shape[1]) != k:
- raise RuntimeError("gemm_a_bt_f16: 维度不匹配。")
- if int(c.shape[0]) != m or int(c.shape[1]) != n:
- raise RuntimeError("gemm_a_bt_f16: 输出形状不匹配。")
-
- self._set_q()
-
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
- st = self._gemm_ex(
- self._handle,
- int(_CUBLAS_OP_T),
- int(_CUBLAS_OP_N),
- n,
- m,
- k,
- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),
- ctypes.c_void_p(int(b.data_ptr())),
- int(_CUDA_R_16F),
- k,
- ctypes.c_void_p(int(a.data_ptr())),
- int(_CUDA_R_16F),
- k,
- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),
- ctypes.c_void_p(int(c.data_ptr())),
- int(_CUDA_R_16F),
- n,
- int(_CUBLAS_COMPUTE_32F),
- int(_CUBLAS_GEMM_DEFAULT),
+ st = self.lib.cublasGemmEx(
+ self.handle,
+ self._OP_N,
+ self._OP_N,
+ ctypes.c_int(n),
+ ctypes.c_int(m),
+ ctypes.c_int(k),
+ ctypes.c_void_p(ctypes.addressof(alpha)),
+ self._ptr(b),
+ self._CUDA_R_16F,
+ ctypes.c_int(n),
+ self._ptr(a),
+ self._CUDA_R_16F,
+ ctypes.c_int(k),
+ ctypes.c_void_p(ctypes.addressof(beta)),
+ self._ptr(c),
+ self._CUDA_R_16F,
+ ctypes.c_int(n),
+ ctypes.c_int(self._COMPUTE_32F),
+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS GEMM(fp16->fp16) 失败: status={int(st)}")
+ if st != self._STATUS_SUCCESS:
+ raise RuntimeError(f"cublasGemmEx(fp16->fp16) 失败:status={st}")
- def gemm_a_b_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
+ def gemm_rm_f16_f16_to_f32(self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor) -> None:
+ m, k = a.shape
+ kb, n = b.shape
+ if kb != k:
+ raise RuntimeError("GEMM 维度不匹配。")
+ if c.shape != (m, n):
+ raise RuntimeError("GEMM 输出维度不匹配。")
+
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
- raise RuntimeError("gemm_a_b_f32 仅支持 A/B fp16,C fp32。")
+ raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 A/B fp16,C fp32。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
- raise RuntimeError("gemm_a_b_f32 仅支持 CUDA 张量。")
+ raise RuntimeError("gemm_rm_f16_f16_to_f32 仅支持 CUDA 张量。")
- a = a.contiguous()
- b = b.contiguous()
- c = c.contiguous()
-
- m = int(a.shape[0])
- k = int(a.shape[1])
- if int(b.shape[0]) != k:
- raise RuntimeError("gemm_a_b_f32: 维度不匹配。")
- n = int(b.shape[1])
- if int(c.shape[0]) != m or int(c.shape[1]) != n:
- raise RuntimeError("gemm_a_b_f32: 输出形状不匹配。")
-
- self._set_q()
-
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
- st = self._gemm_ex(
- self._handle,
- int(_CUBLAS_OP_N),
- int(_CUBLAS_OP_N),
- n,
- m,
- k,
- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),
- ctypes.c_void_p(int(b.data_ptr())),
- int(_CUDA_R_16F),
- n,
- ctypes.c_void_p(int(a.data_ptr())),
- int(_CUDA_R_16F),
- k,
- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),
- ctypes.c_void_p(int(c.data_ptr())),
- int(_CUDA_R_32F),
- n,
- int(_CUBLAS_COMPUTE_32F),
- int(_CUBLAS_GEMM_DEFAULT),
+ st = self.lib.cublasGemmEx(
+ self.handle,
+ self._OP_N,
+ self._OP_N,
+ ctypes.c_int(n),
+ ctypes.c_int(m),
+ ctypes.c_int(k),
+ ctypes.c_void_p(ctypes.addressof(alpha)),
+ self._ptr(b),
+ self._CUDA_R_16F,
+ ctypes.c_int(n),
+ self._ptr(a),
+ self._CUDA_R_16F,
+ ctypes.c_int(k),
+ ctypes.c_void_p(ctypes.addressof(beta)),
+ self._ptr(c),
+ self._CUDA_R_32F,
+ ctypes.c_int(n),
+ ctypes.c_int(self._COMPUTE_32F),
+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS GEMM(fp16->fp32) 失败: status={int(st)}")
+ if st != self._STATUS_SUCCESS:
+ raise RuntimeError(f"cublasGemmEx(fp16->fp32) 失败:status={st}")
- def gemm_strided_batched_a_b_f32(
+ def gemm_rm_f16_f16_to_f32_strided_batched(
self,
a: torch.Tensor,
b: torch.Tensor,
c: torch.Tensor,
+ batch: int,
m: int,
n: int,
k: int,
stride_a: int,
stride_b: int,
stride_c: int,
- batch: int,
) -> None:
if a.dtype != torch.float16 or b.dtype != torch.float16 or c.dtype != torch.float32:
- raise RuntimeError("gemm_strided_batched_a_b_f32 仅支持 A/B fp16,C fp32。")
+ raise RuntimeError("batched GEMM 仅支持 A/B fp16,C fp32。")
if not (a.is_cuda and b.is_cuda and c.is_cuda):
- raise RuntimeError("gemm_strided_batched_a_b_f32 仅支持 CUDA 张量。")
+ raise RuntimeError("batched GEMM 仅支持 CUDA 张量。")
- self._set_q()
-
alpha = ctypes.c_float(1.0)
beta = ctypes.c_float(0.0)
- st = self._gemm_strided_batched_ex(
- self._handle,
- int(_CUBLAS_OP_N),
- int(_CUBLAS_OP_N),
- int(n),
- int(m),
- int(k),
- ctypes.cast(ctypes.byref(alpha), ctypes.c_void_p),
- ctypes.c_void_p(int(b.data_ptr())),
- int(_CUDA_R_16F),
- int(n),
- ctypes.c_longlong(int(stride_b)),
- ctypes.c_void_p(int(a.data_ptr())),
- int(_CUDA_R_16F),
- int(k),
- ctypes.c_longlong(int(stride_a)),
- ctypes.cast(ctypes.byref(beta), ctypes.c_void_p),
- ctypes.c_void_p(int(c.data_ptr())),
- int(_CUDA_R_32F),
- int(n),
- ctypes.c_longlong(int(stride_c)),
- int(batch),
- int(_CUBLAS_COMPUTE_32F),
- int(_CUBLAS_GEMM_DEFAULT),
+ st = self.lib.cublasGemmStridedBatchedEx(
+ self.handle,
+ self._OP_N,
+ self._OP_N,
+ ctypes.c_int(n),
+ ctypes.c_int(m),
+ ctypes.c_int(k),
+ ctypes.c_void_p(ctypes.addressof(alpha)),
+ self._ptr(b),
+ self._CUDA_R_16F,
+ ctypes.c_int(n),
+ ctypes.c_longlong(stride_b),
+ self._ptr(a),
+ self._CUDA_R_16F,
+ ctypes.c_int(k),
+ ctypes.c_longlong(stride_a),
+ ctypes.c_void_p(ctypes.addressof(beta)),
+ self._ptr(c),
+ self._CUDA_R_32F,
+ ctypes.c_int(n),
+ ctypes.c_longlong(stride_c),
+ ctypes.c_int(batch),
+ ctypes.c_int(self._COMPUTE_32F),
+ ctypes.c_int(self._GEMM_DEFAULT_TENSOR_OP),
)
- if st != _CUBLAS_STATUS_SUCCESS:
- raise RuntimeError(f"cuBLAS Batched GEMM(fp16->fp32) 失败: status={int(st)}")
+ if st != self._STATUS_SUCCESS:
+ raise RuntimeError(f"cublasGemmStridedBatchedEx 失败:status={st}")
- _CUBLAS = None
+ _CUBLAS: Optional[_Cublas] = None
def _get_cublas() -> _Cublas:
⋯ 130 unchanged lines
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),
- ),
+ 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, h, n, n), stride=(h * n * n, n * 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
⋯ 44 unchanged lines
gr = _sigmoid_f16(rg)
go = _sigmoid_f16(og)
+
left[bb, hh, i, j] = (lp * gl * m).to(cutlass.Float16)
+
right[bb, hh, j, i] = (rp * gr * m).to(cutlass.Float16)
gate[bb, i, j, hh] = go
⋯ 148 unchanged lines
def _as_ptr(ty, t: torch.Tensor):
- return make_ptr(ty, int(t.data_ptr()), cute.AddressSpace.gmem, assumed_align=16)
+ return make_ptr(ty, t.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
⋯ 2 unchanged lines
_W_KEY = None
- _W_PACK16 = None
+ _W_PACK_T16 = None
_W_OUT_T16 = None
_W_NORM_W = None
_W_NORM_B = None
⋯ 5 unchanged lines
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
+ global _W_KEY, _W_PACK_T16, _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"]
⋯ 21 unchanged lines
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
+ return _W_PACK_T16, _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_pack = torch.cat([w_lp, w_rp, w_lg, w_rg, w_og], dim=0).contiguous().to(torch.float16)
+ w_pack_t16 = w_pack.transpose(0, 1).contiguous()
w_out_t16 = w_out.contiguous().to(torch.float16).transpose(0, 1).contiguous()
- _W_PACK16 = w_pack16
+ _W_PACK_T16 = w_pack_t16
_W_OUT_T16 = w_out_t16
_W_NORM_W = w_no_w.contiguous()
_W_NORM_B = w_no_b.contiguous()
⋯ 1 unchanged lines
_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
+ return w_pack_t16, 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):
⋯ 42 unchanged lines
_compile_once()
- w_pack16, w_out_t16, w_norm_w, w_norm_b, w_out_norm_w, w_out_norm_b = _prep_weights(
+ w_pack_t16, 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"]
⋯ 15 unchanged lines
m = bs * n * n
- _get_cublas().gemm_a_bt_f16(x_norm.view(m, dim), w_pack16, proj.view(m, 5 * hidden))
+ _get_cublas().gemm_rm_f16_f16_to_f16(x_norm.view(m, dim), w_pack_t16, proj.view(m, 5 * hidden))
mask16.copy_(mask)
+
+
_PROC_C(
_as_ptr(cutlass.Float16, proj.view(bs, n, n, 5 * hidden)),
_as_ptr(cutlass.Float16, mask16),
⋯ 6 unchanged lines
batch = bs * hidden
- a_b = left.view(batch, n, n)
- b_b = right.view(batch, n, n)
- c_b = out_tmp.view(batch, n, n)
- _get_cublas().gemm_strided_batched_a_b_f32(
- a_b,
- b_b,
- c_b,
- n,
- n,
- n,
- n * n,
- n * n,
- n * n,
- batch,
+ _get_cublas().gemm_rm_f16_f16_to_f32_strided_batched(
+ left.view(batch, n, n),
+ right.view(batch, n, n),
+ out_tmp.view(batch, n, n),
+ batch=batch,
+ m=n,
+ n=n,
+ k=n,
+ stride_a=n * n,
+ stride_b=n * n,
+ stride_c=n * n,
)
⋯ 8 unchanged lines
y = torch.empty((m, dim), device=x.device, dtype=torch.float32)
- _get_cublas().gemm_a_b_f32(out_norm.view(m, hidden), w_out_t16, y)
+ _get_cublas().gemm_rm_f16_f16_to_f32(out_norm.view(m, hidden), w_out_t16, y)
+
return y.view(bs, n, n, dim)
scrolls · 572 diff lines total

Best evidence level for this revision: reported

JSON