Skip to content
KernelIndex
Search⌘K

submission 417619

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

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

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

shared-memory__shared__ float shm_sum[256];
vector-width = float4const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);

Kernel source

submission.py695 lines
from __future__ import annotations

from typing import Any, Dict, Tuple

import os

import torch





_EXT = None
_EXT_LOCK = None


def _lazy_import_extension_utils():
    
    from torch.utils.cpp_extension import load_inline  

    return load_inline


def _get_ext():
    global _EXT, _EXT_LOCK
    if _EXT is not None:
        return _EXT
    if _EXT_LOCK is None:
        import threading  

        _EXT_LOCK = threading.Lock()
    with _EXT_LOCK:
        if _EXT is not None:
            return _EXT

        load_inline = _lazy_import_extension_utils()

        
        if "TORCH_CUDA_ARCH_LIST" not in os.environ:
            os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"

        
        

        cpp_src = r"""
#include <torch/extension.h>

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);

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("fwd", &trimul_fwd, "trimul forward (cuda)");
}
"""

        cuda_src = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>

#include <mutex>

namespace {

static inline void checkCuda(cudaError_t e, const char* msg) {
  if (e != cudaSuccess) {
    throw std::runtime_error(std::string(msg) + ": " + cudaGetErrorString(e));
  }
}

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));
  }
}

struct CublasHandleHolder {
  cublasHandle_t handle = nullptr;
  CublasHandleHolder() {
    checkCublas(cublasCreate(&handle), "cublasCreate");
    // 默认数学模式即可;这里不做额外设置,减少环境依赖
  }
  ~CublasHandleHolder() {
    if (handle) {
      cublasDestroy(handle);
      handle = nullptr;
    }
  }
};

static CublasHandleHolder* get_cublas() {
  static std::once_flag once;
  static CublasHandleHolder* holder = nullptr;
  std::call_once(once, []() { holder = new CublasHandleHolder(); });
  return holder;
}

__device__ __forceinline__ float warp_sum(float v) {
  for (int d = 16; d > 0; d >>= 1) {
    v += __shfl_down_sync(0xffffffff, v, d);
  }
  return v;
}

__device__ __forceinline__ float fast_sigmoid(float x) {
  // 使用 __expf,精度在题面容忍范围内通常足够
  float z = __expf(-x);
  return 1.0f / (1.0f + z);
}

__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

  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);
  // warp_sum 仅保证 lane0 得到全和,需要广播到全 warp
  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);

  const float4* w4 = reinterpret_cast<const float4*>(w);
  const float4* b4 = reinterpret_cast<const float4*>(b);
  float4 gw = w4[lane];
  float4 gb = b4[lane];

  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;

  half2 h0 = __floats2half2_rn(y0, y1);
  half2 h1 = __floats2half2_rn(y2, y3);

  half2* y2p = reinterpret_cast<half2*>(y + row * 128 + lane * 4);
  y2p[0] = h0;
  y2p[1] = h1;
}

__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;

  // 计算均值与二阶矩
  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;
  }

  __shared__ float shm_sum[256];
  __shared__ float shm_sq[256];
  int t = (int)threadIdx.x;
  shm_sum[t] = sum;
  shm_sq[t] = sq;
  __syncthreads();

  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();
  }

  float mean = shm_sum[0] / (float)dim;
  float var = shm_sq[0] / (float)dim - mean * mean;
  float inv = rsqrtf(var + 1e-5f);

  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);
  }
}

__global__ void pack_proj_f16(
    const half* __restrict__ proj,
    const float* __restrict__ mask,
    half* __restrict__ left,
    half* __restrict__ right,
    half* __restrict__ ogate,
    int bs,
    int n,
    int hidden) {
  int64_t row = (int64_t)blockIdx.x;
  int d = (int)threadIdx.x;
  if (d >= hidden) return;

  int64_t nn = (int64_t)n * (int64_t)n;
  int b0 = (int)(row / nn);
  int64_t rem = row - (int64_t)b0 * nn;
  int i = (int)(rem / n);
  int j = (int)(rem - (int64_t)i * n);

  float m = mask[row];

  int out = hidden * 5;
  const half* p = proj + row * out;

  float l = __half2float(p[d]);
  float r = __half2float(p[hidden + d]);
  float gl = fast_sigmoid(__half2float(p[2 * hidden + d]));
  float gr = fast_sigmoid(__half2float(p[3 * hidden + d]));
  float go = fast_sigmoid(__half2float(p[4 * hidden + d]));

  float l2 = l * gl * m;
  float r2 = r * gr * m;

  int64_t base = (((int64_t)b0 * hidden + d) * n + i) * n + j;
  left[base] = __float2half_rn(l2);
  right[base] = __float2half_rn(r2);
  ogate[base] = __float2half_rn(go);
}

template<int MAX_H>
__global__ void ln2_gate_store_f16(
    const float* __restrict__ out_acc,
    const half* __restrict__ ogate,
    const float* __restrict__ w,
    const float* __restrict__ b,
    half* __restrict__ out_norm,
    int bs,
    int n,
    int hidden) {
  // blockIdx.x 对应 (b,i,j)
  int64_t row = (int64_t)blockIdx.x;
  int tid = (int)threadIdx.x;
  if (tid >= MAX_H) return;

  int64_t nn = (int64_t)n * (int64_t)n;
  int b0 = (int)(row / nn);
  int64_t rem = row - (int64_t)b0 * nn;
  int i = (int)(rem / n);
  int j = (int)(rem - (int64_t)i * n);

  float v = 0.0f;
  float vv = 0.0f;
  if (tid < hidden) {
    int64_t idx = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
    float x = out_acc[idx];
    v = x;
    vv = x * x;
  }

  // 归约:MAX_H 固定,使用共享内存
  __shared__ float shm_sum[MAX_H];
  __shared__ float shm_sq[MAX_H];
  shm_sum[tid] = v;
  shm_sq[tid] = vv;
  __syncthreads();

  for (int stride = MAX_H / 2; stride > 0; stride >>= 1) {
    if (tid < stride) {
      shm_sum[tid] += shm_sum[tid + stride];
      shm_sq[tid] += shm_sq[tid + stride];
    }
    __syncthreads();
  }

  float mean = shm_sum[0] / (float)hidden;
  float var = shm_sq[0] / (float)hidden - mean * mean;
  float inv = rsqrtf(var + 1e-5f);

  if (tid < hidden) {
    int64_t idx_in = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
    float x = out_acc[idx_in];
    float y = (x - mean) * inv * w[tid] + b[tid];
    float g = __half2float(ogate[idx_in]);
    float z = y * g;

    int64_t idx_out = ((row * hidden) + tid);
    out_norm[idx_out] = __float2half_rn(z);
  }
}

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");
  }
}

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 rows = (int64_t)bs * (int64_t)n * (int64_t)n;
  dim3 block((unsigned)hidden, 1, 1);
  dim3 grid((unsigned)rows, 1, 1);
  pack_proj_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(),
      bs, n, hidden);
  checkCuda(cudaGetLastError(), "pack_proj_f16");
}

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 rows = (int64_t)bs * (int64_t)n * (int64_t)n;
  if (hidden <= 32) {
    dim3 block(32, 1, 1);
    dim3 grid((unsigned)rows, 1, 1);
    ln2_gate_store_f16<32><<<grid, block>>>(
        (const float*)out_acc.data_ptr(),
        (const half*)og.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm.data_ptr(),
        bs, n, hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_store_f16_32");
  } else if (hidden <= 64) {
    dim3 block(64, 1, 1);
    dim3 grid((unsigned)rows, 1, 1);
    ln2_gate_store_f16<64><<<grid, block>>>(
        (const float*)out_acc.data_ptr(),
        (const half*)og.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm.data_ptr(),
        bs, n, hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_store_f16_64");
  } else if (hidden <= 128) {
    dim3 block(128, 1, 1);
    dim3 grid((unsigned)rows, 1, 1);
    ln2_gate_store_f16<128><<<grid, block>>>(
        (const float*)out_acc.data_ptr(),
        (const half*)og.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm.data_ptr(),
        bs, n, hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_store_f16_128");
  } else {
    throw std::runtime_error("hidden_dim too large");
  }
}

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) {
  // 使用列主序 trick:输出按列主序 (N x M) 写入,即等价于 row-major (M x N)
  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_x_wt_f16_f32(
    cublasHandle_t h,
    const half* x_row, // row-major [M,K]
    const half* w_row, // row-major [N,K]
    float* 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_32F, (int)N,
          CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
      "cublasGemmEx");
}

static void gemm_contract_batched_f16_f32(
    cublasHandle_t h,
    const half* left_row,   // row-major [B,M,K]
    const half* right_row,  // row-major [B,N,K]  (这里 N=M=K=n)
    float* 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;

  // 计算 C = L * R^T
  // 采用列主序 trick:C^T = R * L^T
  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_32F, 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 (w_cat.scalar_type() != torch::kFloat16 || w_out.scalar_type() != torch::kFloat16) {
    throw std::runtime_error("weights must be float16");
  }

  int bs = (int)x.size(0);
  int n = (int)x.size(1);
  int64_t M = (int64_t)bs * (int64_t)n * (int64_t)n;

  auto x2d = x.view({M, dim});
  auto xhat = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  launch_ln1(x2d, ln1_w, ln1_b, xhat);

  int64_t out_ch = hidden * 5;
  auto proj = torch::empty({M, out_ch}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));

  auto* holder = get_cublas();
  cublasHandle_t h = holder->handle;

  // gemm1: proj = xhat @ w_cat^T
  gemm_x_wt_f16_f16(h, (const half*)xhat.data_ptr(), (const half*)w_cat.data_ptr(), (half*)proj.data_ptr(), M, out_ch, dim);

  auto left = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  auto right = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  auto og = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  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});
  auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));

  gemm_contract_batched_f16_f32(
      h,
      (const half*)left3.data_ptr(),
      (const half*)right3.data_ptr(),
      (float*)out_acc.data_ptr(),
      bs * (int)hidden, n);

  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
  auto y = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
  gemm_x_wt_f16_f32(h, (const half*)out_norm.data_ptr(), (const half*)w_out.data_ptr(), (float*)y.data_ptr(), M, dim, hidden);

  return y.view({bs, n, n, dim});
}
"""

        name = "trimul_ext_mod"

        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"])

    if not x.is_cuda:
        raise RuntimeError("x must be CUDA tensor")
    if not mask.is_cuda:
        raise RuntimeError("mask must be CUDA tensor")

    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()

    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")

    
    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")

    
    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")

    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")

    
    w_cat, w_out = _prepare_weights(weights, dim, hidden)

    
    ext = _get_ext()
    return ext.fwd(x, mask, ln1_w, ln1_b, w_cat, ln2_w, ln2_b, w_out, dim, hidden)


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

⋯ 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, yield_out
+ _EXT = None
+ _EXT_LOCK = None
+ def _lazy_import_extension_utils():
+
+ from torch.utils.cpp_extension import load_inline
+ return load_inline
- _TILE_MN = 32
- _TILE_K = 32
- _THREADS = 256
+ def _get_ext():
+ global _EXT, _EXT_LOCK
+ if _EXT is not None:
+ return _EXT
+ if _EXT_LOCK is None:
+ import threading
- class _OutgoingContractBatchedKernel:
- def __init__(self) -> None:
- self.threads = _THREADS
+ _EXT_LOCK = threading.Lock()
+ with _EXT_LOCK:
+ if _EXT is not None:
+ return _EXT
- @cute.jit
- def __call__(
- self,
- a_ptr: "cute.Pointer",
- b_ptr: "cute.Pointer",
- c_ptr: "cute.Pointer",
- problem: tuple,
- ):
- bh, n = problem
+ load_inline = _lazy_import_extension_utils()
- stride_bh = n * n
- stride_m = n
- stride_n = 1
+
+ if "TORCH_CUDA_ARCH_LIST" not in os.environ:
+ os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"
- a = cute.make_tensor(
- a_ptr,
- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
- )
- b = cute.make_tensor(
- b_ptr,
- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
- )
- c = cute.make_tensor(
- c_ptr,
- cute.make_layout((bh, n, n), stride=(stride_bh, stride_m, stride_n)),
- )
+
+
- grid_x = n // _TILE_MN
- grid_y = n // _TILE_MN
- grid_z = bh
- self.kernel(a, b, c, n).launch(
- grid=[grid_x, grid_y, grid_z],
- block=[self.threads, 1, 1],
- )
- return
+ cpp_src = r"""
+ #include <torch/extension.h>
- @cute.kernel
- def kernel(
- self,
- a: "cute.Tensor",
- b: "cute.Tensor",
- c: "cute.Tensor",
- n: int,
- ):
- tx, _, _ = cute.arch.thread_idx()
- bx, by, bz = cute.arch.block_idx()
+ 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);
- tid = tx
- lane_m = tid >> 4
- lane_n = tid & 15
+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
+ m.def("fwd", &trimul_fwd, "trimul forward (cuda)");
+ }
+ """
- base_m = by * _TILE_MN
- base_n = bx * _TILE_MN
+ cuda_src = r"""
+ #include <torch/extension.h>
+ #include <ATen/cuda/CUDAContext.h>
+ #include <cuda.h>
+ #include <cuda_fp16.h>
+ #include <cublas_v2.h>
- i0 = base_m + lane_m
- j0 = base_n + lane_n
- i1 = i0 + 16
- j1 = j0 + 16
+ #include <mutex>
- acc00 = cutlass.Float32(0.0)
- acc01 = cutlass.Float32(0.0)
- acc10 = cutlass.Float32(0.0)
- acc11 = cutlass.Float32(0.0)
+ namespace {
-
- smem_stride = _TILE_K + 1
- smem_a_elems = _TILE_MN * smem_stride
- smem_b_elems = _TILE_MN * smem_stride
- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems)
+ static inline void checkCuda(cudaError_t e, const char* msg) {
+ if (e != cudaSuccess) {
+ throw std::runtime_error(std::string(msg) + ": " + cudaGetErrorString(e));
+ }
+ }
- sA = cute.make_tensor(
- smem_ptr,
- cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),
- )
- sB = cute.make_tensor(
- smem_ptr + smem_a_elems,
- cute.make_layout((_TILE_MN, _TILE_K), stride=(smem_stride, 1)),
- )
+ 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));
+ }
+ }
- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
- 0,
- n,
- _TILE_K,
- 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]
+ struct CublasHandleHolder {
+ cublasHandle_t handle = nullptr;
+ CublasHandleHolder() {
+ checkCublas(cublasCreate(&handle), "cublasCreate");
+ // 默认数学模式即可;这里不做额外设置,减少环境依赖
+ }
+ ~CublasHandleHolder() {
+ if (handle) {
+ cublasDestroy(handle);
+ handle = nullptr;
+ }
+ }
+ };
- cute.arch.sync_threads()
+ static CublasHandleHolder* get_cublas() {
+ static std::once_flag once;
+ static CublasHandleHolder* holder = nullptr;
+ std::call_once(once, []() { holder = new CublasHandleHolder(); });
+ return holder;
+ }
-
- for kk in range(_TILE_K):
- a0 = sA[lane_m, kk]
- a1 = sA[lane_m + 16, kk]
- b0 = sB[lane_n, kk]
- b1 = sB[lane_n + 16, kk]
+ __device__ __forceinline__ float warp_sum(float v) {
+ for (int d = 16; d > 0; d >>= 1) {
+ v += __shfl_down_sync(0xffffffff, v, d);
+ }
+ return v;
+ }
- acc00 = acc00 + a0 * b0
- acc01 = acc01 + a0 * b1
- acc10 = acc10 + a1 * b0
- acc11 = acc11 + a1 * b1
+ __device__ __forceinline__ float fast_sigmoid(float x) {
+ // 使用 __expf,精度在题面容忍范围内通常足够
+ float z = __expf(-x);
+ return 1.0f / (1.0f + z);
+ }
- cute.arch.sync_threads()
- yield_out([acc00, acc01, acc10, acc11])
+ __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
- c[bz, i0, j0] = acc00_out
- c[bz, i0, j1] = acc01_out
- c[bz, i1, j0] = acc10_out
- c[bz, i1, j1] = acc11_out
+ 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);
+ // warp_sum 仅保证 lane0 得到全和,需要广播到全 warp
+ 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);
+ const float4* w4 = reinterpret_cast<const float4*>(w);
+ const float4* b4 = reinterpret_cast<const float4*>(b);
+ float4 gw = w4[lane];
+ float4 gb = b4[lane];
- _CONTRACT = _OutgoingContractBatchedKernel()
- _CONTRACT_COMPILED = None
+ 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;
+ half2 h0 = __floats2half2_rn(y0, y1);
+ half2 h1 = __floats2half2_rn(y2, y3);
- def _get_contract_compiled():
- global _CONTRACT_COMPILED
- if _CONTRACT_COMPILED is not None:
- return _CONTRACT_COMPILED
+ half2* y2p = reinterpret_cast<half2*>(y + row * 128 + lane * 4);
+ y2p[0] = h0;
+ y2p[1] = h1;
+ }
- 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_COMPILED = cute.compile(
- _CONTRACT,
- a_ptr,
- b_ptr,
- c_ptr,
- (0, 0),
- options="--opt-level 3",
- )
- return _CONTRACT_COMPILED
+ __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;
+ // 计算均值与二阶矩
+ 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 _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
- bs, n, _, hidden = left.shape
+ __shared__ float shm_sum[256];
+ __shared__ float shm_sq[256];
+ int t = (int)threadIdx.x;
+ shm_sum[t] = sum;
+ shm_sq[t] = sq;
+ __syncthreads();
- 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。")
+ 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 hidden != 128:
- raise RuntimeError("仅支持 hidden_dim=128 的特化路径。")
+ float mean = shm_sum[0] / (float)dim;
+ float var = shm_sq[0] / (float)dim - mean * mean;
+ float inv = rsqrtf(var + 1e-5f);
- if (n & (_TILE_MN - 1)) != 0:
- raise RuntimeError("仅支持 N 为 32 的倍数的特化路径。")
+ 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);
+ }
+ }
+ __global__ void pack_proj_f16(
+ const half* __restrict__ proj,
+ const float* __restrict__ mask,
+ half* __restrict__ left,
+ half* __restrict__ right,
+ half* __restrict__ ogate,
+ int bs,
+ int n,
+ int hidden) {
+ int64_t row = (int64_t)blockIdx.x;
+ int d = (int)threadIdx.x;
+ if (d >= hidden) return;
+
+ int64_t nn = (int64_t)n * (int64_t)n;
+ int b0 = (int)(row / nn);
+ int64_t rem = row - (int64_t)b0 * nn;
+ int i = (int)(rem / n);
+ int j = (int)(rem - (int64_t)i * n);
+
+ float m = mask[row];
+
+ int out = hidden * 5;
+ const half* p = proj + row * out;
+
+ float l = __half2float(p[d]);
+ float r = __half2float(p[hidden + d]);
+ float gl = fast_sigmoid(__half2float(p[2 * hidden + d]));
+ float gr = fast_sigmoid(__half2float(p[3 * hidden + d]));
+ float go = fast_sigmoid(__half2float(p[4 * hidden + d]));
+
+ float l2 = l * gl * m;
+ float r2 = r * gr * m;
+
+ int64_t base = (((int64_t)b0 * hidden + d) * n + i) * n + j;
+ left[base] = __float2half_rn(l2);
+ right[base] = __float2half_rn(r2);
+ ogate[base] = __float2half_rn(go);
+ }
+
+ template<int MAX_H>
+ __global__ void ln2_gate_store_f16(
+ const float* __restrict__ out_acc,
+ const half* __restrict__ ogate,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ half* __restrict__ out_norm,
+ int bs,
+ int n,
+ int hidden) {
+ // blockIdx.x 对应 (b,i,j)
+ int64_t row = (int64_t)blockIdx.x;
+ int tid = (int)threadIdx.x;
+ if (tid >= MAX_H) return;
+
+ int64_t nn = (int64_t)n * (int64_t)n;
+ int b0 = (int)(row / nn);
+ int64_t rem = row - (int64_t)b0 * nn;
+ int i = (int)(rem / n);
+ int j = (int)(rem - (int64_t)i * n);
+
+ float v = 0.0f;
+ float vv = 0.0f;
+ if (tid < hidden) {
+ int64_t idx = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
+ float x = out_acc[idx];
+ v = x;
+ vv = x * x;
+ }
+
+ // 归约:MAX_H 固定,使用共享内存
+ __shared__ float shm_sum[MAX_H];
+ __shared__ float shm_sq[MAX_H];
+ shm_sum[tid] = v;
+ shm_sq[tid] = vv;
+ __syncthreads();
+
+ for (int stride = MAX_H / 2; stride > 0; stride >>= 1) {
+ if (tid < stride) {
+ shm_sum[tid] += shm_sum[tid + stride];
+ shm_sq[tid] += shm_sq[tid + stride];
+ }
+ __syncthreads();
+ }
+
+ float mean = shm_sum[0] / (float)hidden;
+ float var = shm_sq[0] / (float)hidden - mean * mean;
+ float inv = rsqrtf(var + 1e-5f);
+
+ if (tid < hidden) {
+ int64_t idx_in = (((int64_t)b0 * hidden + tid) * n + i) * n + j;
+ float x = out_acc[idx_in];
+ float y = (x - mean) * inv * w[tid] + b[tid];
+ float g = __half2float(ogate[idx_in]);
+ float z = y * g;
+
+ int64_t idx_out = ((row * hidden) + tid);
+ out_norm[idx_out] = __float2half_rn(z);
+ }
+ }
+
+ 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");
+ }
+ }
+
+ 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 rows = (int64_t)bs * (int64_t)n * (int64_t)n;
+ dim3 block((unsigned)hidden, 1, 1);
+ dim3 grid((unsigned)rows, 1, 1);
+ pack_proj_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(),
+ bs, n, hidden);
+ checkCuda(cudaGetLastError(), "pack_proj_f16");
+ }
+
+ 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 rows = (int64_t)bs * (int64_t)n * (int64_t)n;
+ if (hidden <= 32) {
+ dim3 block(32, 1, 1);
+ dim3 grid((unsigned)rows, 1, 1);
+ ln2_gate_store_f16<32><<<grid, block>>>(
+ (const float*)out_acc.data_ptr(),
+ (const half*)og.data_ptr(),
+ (const float*)w.data_ptr(),
+ (const float*)b.data_ptr(),
+ (half*)out_norm.data_ptr(),
+ bs, n, hidden);
+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_32");
+ } else if (hidden <= 64) {
+ dim3 block(64, 1, 1);
+ dim3 grid((unsigned)rows, 1, 1);
+ ln2_gate_store_f16<64><<<grid, block>>>(
+ (const float*)out_acc.data_ptr(),
+ (const half*)og.data_ptr(),
+ (const float*)w.data_ptr(),
+ (const float*)b.data_ptr(),
+ (half*)out_norm.data_ptr(),
+ bs, n, hidden);
+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_64");
+ } else if (hidden <= 128) {
+ dim3 block(128, 1, 1);
+ dim3 grid((unsigned)rows, 1, 1);
+ ln2_gate_store_f16<128><<<grid, block>>>(
+ (const float*)out_acc.data_ptr(),
+ (const half*)og.data_ptr(),
+ (const float*)w.data_ptr(),
+ (const float*)b.data_ptr(),
+ (half*)out_norm.data_ptr(),
+ bs, n, hidden);
+ checkCuda(cudaGetLastError(), "ln2_gate_store_f16_128");
+ } else {
+ throw std::runtime_error("hidden_dim too large");
+ }
+ }
+
+ 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) {
+ // 使用列主序 trick:输出按列主序 (N x M) 写入,即等价于 row-major (M x N)
+ 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_x_wt_f16_f32(
+ cublasHandle_t h,
+ const half* x_row, // row-major [M,K]
+ const half* w_row, // row-major [N,K]
+ float* 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_32F, (int)N,
+ CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
+ "cublasGemmEx");
+ }
+
+ static void gemm_contract_batched_f16_f32(
+ cublasHandle_t h,
+ const half* left_row, // row-major [B,M,K]
+ const half* right_row, // row-major [B,N,K] (这里 N=M=K=n)
+ float* 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;
+
+ // 计算 C = L * R^T
+ // 采用列主序 trick:C^T = R * L^T
+ 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_32F, 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 (w_cat.scalar_type() != torch::kFloat16 || w_out.scalar_type() != torch::kFloat16) {
+ throw std::runtime_error("weights must be float16");
+ }
+
+ int bs = (int)x.size(0);
+ int n = (int)x.size(1);
+ int64_t M = (int64_t)bs * (int64_t)n * (int64_t)n;
+
+ auto x2d = x.view({M, dim});
+ auto xhat = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+ launch_ln1(x2d, ln1_w, ln1_b, xhat);
+
+ int64_t out_ch = hidden * 5;
+ auto proj = torch::empty({M, out_ch}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+
+ auto* holder = get_cublas();
+ cublasHandle_t h = holder->handle;
+
+ // gemm1: proj = xhat @ w_cat^T
+ gemm_x_wt_f16_f16(h, (const half*)xhat.data_ptr(), (const half*)w_cat.data_ptr(), (half*)proj.data_ptr(), M, out_ch, dim);
+
+ auto left = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+ auto right = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+ auto og = torch::empty({bs, hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+ 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});
+ auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
+
+ gemm_contract_batched_f16_f32(
+ h,
+ (const half*)left3.data_ptr(),
+ (const half*)right3.data_ptr(),
+ (float*)out_acc.data_ptr(),
+ bs * (int)hidden, n);
+
+ 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
+ auto y = torch::empty({M, dim}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
+ gemm_x_wt_f16_f32(h, (const half*)out_norm.data_ptr(), (const half*)w_out.data_ptr(), (float*)y.data_ptr(), M, dim, hidden);
+
+ return y.view({bs, n, n, dim});
+ }
+ """
+
+ name = "trimul_ext_mod"
+
+ 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):
- 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)
+ 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
- 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))
+ 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"]
- out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
- return out
+ 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_dim = int(config["hidden_dim"])
+ hidden = 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")
+
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
- torch.backends.cuda.matmul.allow_tf32 = True
- torch.backends.cudnn.allow_tf32 = True
+ if mask.dtype != torch.float32:
+ mask = mask.to(dtype=torch.float32)
- x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
+ x = x.contiguous()
+ mask = mask.contiguous()
- left = F.linear(x, weights["left_proj.weight"], None)
- right = F.linear(x, weights["right_proj.weight"], None)
+ 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")
- mask_f = mask.unsqueeze(-1)
- if mask_f.dtype != left.dtype:
- mask_f = mask_f.to(dtype=left.dtype)
- left = left * mask_f
- right = right * mask_f
+
+ 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_gate = torch.sigmoid(F.linear(x, weights["left_gate.weight"], None))
- right_gate = torch.sigmoid(F.linear(x, weights["right_gate.weight"], None))
- out_gate = torch.sigmoid(F.linear(x, weights["out_gate.weight"], None))
+
+ 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")
- left = left * left_gate
- right = right * right_gate
+ 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")
- out = _contract_outgoing(left, right)
- out = F.layer_norm(
- out,
- (hidden_dim,),
- weights["to_out_norm.weight"],
- weights["to_out_norm.bias"],
- 1e-5,
- )
- out = out * out_gate
- out = F.linear(out, weights["to_out.weight"], None)
- return out
+
+ w_cat, w_out = _prepare_weights(weights, dim, hidden)
+
+ ext = _get_ext()
+ return ext.fwd(x, mask, ln1_w, ln1_b, w_cat, ln2_w, ln2_b, w_out, dim, hidden)
+
__all__ = ["custom_kernel"]
scrolls · 874 diff lines total

Best evidence level for this revision: reported

JSON