Skip to content
KernelIndex
Search⌘K

submission 417907

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-417907?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
1.91ms
#25 of 71
2026-01-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b2048c34eb32d34b892732101fc87e8c9b666b9bdb64636f59cd5be390bd547e
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.py806 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 <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) {
  // --use_fast_math 下的 __expf 通常足够快;精度仍能满足题面容忍度
  float z = __expf(-x);
  return 1.0f / (1.0f + z);
}

// ---------------- LN1 ----------------

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

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

// --------------- 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 数量。

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

  int b = (int)blockIdx.y;
  int64_t p_base = (int64_t)blockIdx.x * (int64_t)P_TILE;

  int tx = (int)threadIdx.x; // 0..31
  int ty = (int)threadIdx.y; // 0..BLOCK_ROWS-1

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

  // 共享内存第二维 +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];

  __syncthreads();

  int out_ch = hidden * 5;

  // 逐块处理 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);

      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;

        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;

        hl = __float2half_rn(l2);
        hr = __float2half_rn(r2);
        hg = __float2half_rn(go);
      }

      sh_l[p_l][tx] = hl;
      sh_r[p_l][tx] = hr;
      sh_g[p_l][tx] = hg;
    }

    __syncthreads();

    // 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];
      }
    }

    __syncthreads();
  }
}

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 维)。

template<int MAX_H>
__global__ void ln2_gate_store_p32_f16(
    const float* __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 = 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 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(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_32");
  } else if (hidden <= 64) {
    ln2_gate_store_p32_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(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_64");
  } else {
    ln2_gate_store_p32_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(),
        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_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]
    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;

  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;
  int64_t nn = (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;

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

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

  // LN2 + gate:输出为 [M, hidden] half,适配后续 GEMM2
  auto out_norm = torch::empty({M, hidden}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  // og view: [bs*hidden, n, n] 等价的连续指针
  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_mod3"

        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 · 806 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 417882.

⋯ 6 unchanged lines
import torch
+
+
+
_EXT = None
_EXT_LOCK = None
⋯ 18 unchanged lines
load_inline = _lazy_import_extension_utils()
+
if "TORCH_CUDA_ARCH_LIST" not in os.environ:
os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"
⋯ 43 unchanged lines
cublasHandle_t handle = nullptr;
CublasHandleHolder() {
checkCublas(cublasCreate(&handle), "cublasCreate");
- checkCublas(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode");
}
~CublasHandleHolder() {
if (handle) {
⋯ 11 unchanged lines
}
__device__ __forceinline__ float warp_sum(float v) {
- v += __shfl_down_sync(0xffffffff, v, 16);
- v += __shfl_down_sync(0xffffffff, v, 8);
- v += __shfl_down_sync(0xffffffff, v, 4);
- v += __shfl_down_sync(0xffffffff, v, 2);
- v += __shfl_down_sync(0xffffffff, v, 1);
+ for (int d = 16; d > 0; d >>= 1) {
+ v += __shfl_down_sync(0xffffffff, v, d);
+ }
return v;
}
__device__ __forceinline__ float fast_sigmoid(float x) {
+ // --use_fast_math 下的 __expf 通常足够快;精度仍能满足题面容忍度
float z = __expf(-x);
return 1.0f / (1.0f + z);
}
⋯ 8 unchanged lines
int64_t rows) {
int64_t row = (int64_t)blockIdx.x;
if (row >= rows) return;
- int lane = (int)threadIdx.x;
+ int lane = (int)threadIdx.x; // 0..31
const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);
float4 v = x4[lane];
⋯ 19 unchanged lines
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_384_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 tid = (int)threadIdx.x; // 0..95
- int lane = tid & 31;
- int warp = tid >> 5; // 0..2
-
- const float4* x4 = reinterpret_cast<const float4*>(x + row * 384);
- float4 v = x4[tid];
- 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);
-
- __shared__ float sh_sum[3];
- __shared__ float sh_sq[3];
- if (lane == 0) {
- sh_sum[warp] = s;
- sh_sq[warp] = ss;
- }
- __syncthreads();
-
- __shared__ float mean_sh;
- __shared__ float inv_sh;
- if (tid == 0) {
- float sum = sh_sum[0] + sh_sum[1] + sh_sum[2];
- float sq = sh_sq[0] + sh_sq[1] + sh_sq[2];
- float mean = sum * (1.0f / 384.0f);
- float var = sq * (1.0f / 384.0f) - mean * mean;
- mean_sh = mean;
- inv_sh = rsqrtf(var + 1e-5f);
- }
- __syncthreads();
-
- float mean = mean_sh;
- float inv = inv_sh;
-
- const float4* w4 = reinterpret_cast<const float4*>(w);
- const float4* b4 = reinterpret_cast<const float4*>(b);
- float4 gw = w4[tid];
- float4 gb = b4[tid];
-
- 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;
-
- int64_t off = row * 384 + (int64_t)tid * 4;
- half2 h0 = __floats2half2_rn(y0, y1);
- half2 h1 = __floats2half2_rn(y2, y3);
- half2* y2p = reinterpret_cast<half2*>(y + off);
- y2p[0] = h0;
- y2p[1] = h1;
- }
-
__global__ void ln1_generic_f16(
const float* __restrict__ x,
const float* __restrict__ w,
⋯ 52 unchanged lines
(half*)y.data_ptr(),
(int64_t)rows);
checkCuda(cudaGetLastError(), "ln1_128_f16");
- } else if (dim == 384) {
- dim3 block(96, 1, 1);
- dim3 grid((unsigned)rows, 1, 1);
- ln1_384_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_384_f16");
} else {
dim3 block(256, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
⋯ 9 unchanged lines
}
// --------------- 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 数量。
- template<int D_TILE, int TILE_P>
- __global__ void pack_proj_tiled2_f16(
- const half* __restrict__ proj,
- const float* __restrict__ mask,
- half* __restrict__ left,
- half* __restrict__ right,
- half* __restrict__ ogate,
+ __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;
+
int b = (int)blockIdx.y;
- int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;
+ int64_t p_base = (int64_t)blockIdx.x * (int64_t)P_TILE;
- int tid = (int)threadIdx.x;
- int p_l = tid / D_TILE;
- int d0 = tid - p_l * D_TILE;
- int64_t p = p0 + (int64_t)p_l;
+ int tx = (int)threadIdx.x; // 0..31
+ int ty = (int)threadIdx.y; // 0..BLOCK_ROWS-1
- __shared__ float m_sh[TILE_P];
- if (d0 == 0) {
+ __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[p_l] = mv;
+ m_sh[tx] = mv;
}
- __syncthreads();
- __shared__ half sh_l[128 * TILE_P];
- __shared__ half sh_r[128 * TILE_P];
- __shared__ half sh_g[128 * TILE_P];
+ // 共享内存第二维 +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];
- half out_l0 = __float2half_rn(0.0f);
- half out_r0 = __float2half_rn(0.0f);
- half out_g0 = __float2half_rn(0.0f);
- half out_l1 = __float2half_rn(0.0f);
- half out_r1 = __float2half_rn(0.0f);
- half out_g1 = __float2half_rn(0.0f);
+ __syncthreads();
- if (p < nn) {
- int64_t row = (int64_t)b * nn + p;
- int out_ch = hidden * 5;
- const half* base = proj + row * (int64_t)out_ch;
- float m = m_sh[p_l];
+ int out_ch = hidden * 5;
- int d = d0;
- if (d < hidden) {
- 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]));
- out_l0 = __float2half_rn(l * gl * m);
- out_r0 = __float2half_rn(r * gr * m);
- out_g0 = __float2half_rn(go);
- }
+ // 逐块处理 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;
- int d1 = d0 + 64;
- if (d1 < hidden) {
- float l = __half2float(base[d1]);
- float r = __half2float(base[hidden + d1]);
- float gl = fast_sigmoid(__half2float(base[2 * hidden + d1]));
- float gr = fast_sigmoid(__half2float(base[3 * hidden + d1]));
- float go = fast_sigmoid(__half2float(base[4 * hidden + d1]));
- out_l1 = __float2half_rn(l * gl * m);
- out_r1 = __float2half_rn(r * gr * m);
- out_g1 = __float2half_rn(go);
- }
- }
+ half hl = __float2half_rn(0.0f);
+ half hr = __float2half_rn(0.0f);
+ half hg = __float2half_rn(0.0f);
- sh_l[d0 * TILE_P + p_l] = out_l0;
- sh_r[d0 * TILE_P + p_l] = out_r0;
- sh_g[d0 * TILE_P + p_l] = out_g0;
- sh_l[(d0 + 64) * TILE_P + p_l] = out_l1;
- sh_r[(d0 + 64) * TILE_P + p_l] = out_r1;
- sh_g[(d0 + 64) * TILE_P + p_l] = out_g1;
- __syncthreads();
+ 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;
- int d2 = tid / TILE_P; // 0..63
- int p2 = tid - d2 * TILE_P;
- int64_t p_out = p0 + (int64_t)p2;
- if (p_out < nn) {
- int d = d2;
- if (d < hidden) {
- int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d) * nn + p_out;
- left[out_idx] = sh_l[d * TILE_P + p2];
- right[out_idx] = sh_r[d * TILE_P + p2];
- ogate[out_idx] = sh_g[d * TILE_P + p2];
- }
- int d3 = d2 + 64;
- if (d3 < hidden) {
- int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d3) * nn + p_out;
- left[out_idx] = sh_l[d3 * TILE_P + p2];
- right[out_idx] = sh_r[d3 * TILE_P + p2];
- ogate[out_idx] = sh_g[d3 * TILE_P + p2];
- }
- }
- }
+ 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]));
- template<int MAX_H, int TILE_P>
- __global__ void pack_proj_tiled_f16(
- const half* __restrict__ proj,
- const float* __restrict__ mask,
- half* __restrict__ left,
- half* __restrict__ right,
- half* __restrict__ ogate,
- int64_t nn,
- int hidden) {
- int b = (int)blockIdx.y;
- int64_t p0 = (int64_t)blockIdx.x * (int64_t)TILE_P;
+ float m = m_sh[p_l];
+ float l2 = l * gl * m;
+ float r2 = r * gr * m;
- int tid = (int)threadIdx.x;
- int p_l = tid / MAX_H;
- int d = tid - p_l * MAX_H;
- int64_t p = p0 + (int64_t)p_l;
+ hl = __float2half_rn(l2);
+ hr = __float2half_rn(r2);
+ hg = __float2half_rn(go);
+ }
- __shared__ float m_sh[TILE_P];
- if (d == 0) {
- float mv = 0.0f;
- if (p < nn) {
- mv = mask[(int64_t)b * nn + p];
+ sh_l[p_l][tx] = hl;
+ sh_r[p_l][tx] = hr;
+ sh_g[p_l][tx] = hg;
}
- m_sh[p_l] = mv;
- }
- __syncthreads();
- __shared__ half sh_l[MAX_H * TILE_P];
- __shared__ half sh_r[MAX_H * TILE_P];
- __shared__ half sh_g[MAX_H * TILE_P];
+ __syncthreads();
- half out_l = __float2half_rn(0.0f);
- half out_r = __float2half_rn(0.0f);
- half out_g = __float2half_rn(0.0f);
+ // 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];
+ }
+ }
- if (d < hidden && p < nn) {
- int64_t row = (int64_t)b * nn + p;
- int out_ch = hidden * 5;
- const half* base = proj + row * (int64_t)out_ch;
-
- 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];
- out_l = __float2half_rn(l * gl * m);
- out_r = __float2half_rn(r * gr * m);
- out_g = __float2half_rn(go);
+ __syncthreads();
}
-
- sh_l[d * TILE_P + p_l] = out_l;
- sh_r[d * TILE_P + p_l] = out_r;
- sh_g[d * TILE_P + p_l] = out_g;
- __syncthreads();
-
- int d2 = tid / TILE_P;
- int p2 = tid - d2 * TILE_P;
- int64_t p_out = p0 + (int64_t)p2;
- if (d2 < hidden && p_out < nn) {
- int64_t out_idx = ((int64_t)b * (int64_t)hidden + (int64_t)d2) * nn + p_out;
- left[out_idx] = sh_l[d2 * TILE_P + p2];
- right[out_idx] = sh_r[d2 * TILE_P + p2];
- ogate[out_idx] = sh_g[d2 * TILE_P + p2];
- }
}
static void launch_pack(
⋯ 6 unchanged lines
int n,
int hidden) {
int64_t nn = (int64_t)n * (int64_t)n;
- if (hidden <= 32) {
- constexpr int MAX_H = 32;
- constexpr int TILE_P = 16;
- dim3 block(MAX_H * TILE_P, 1, 1);
- dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
- pack_proj_tiled_f16<MAX_H, TILE_P><<<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_tiled_f16_32");
- } else if (hidden <= 64) {
- constexpr int MAX_H = 64;
- constexpr int TILE_P = 8;
- dim3 block(MAX_H * TILE_P, 1, 1);
- dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
- pack_proj_tiled_f16<MAX_H, TILE_P><<<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_tiled_f16_64");
- } else if (hidden <= 128) {
- constexpr int D_TILE = 64;
- constexpr int TILE_P = 8;
- dim3 block(D_TILE * TILE_P, 1, 1);
- dim3 grid((unsigned)((nn + TILE_P - 1) / TILE_P), (unsigned)bs, 1);
- pack_proj_tiled2_f16<D_TILE, TILE_P><<<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_tiled2_f16_128");
- } else {
+ 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 维)。
- template<int HMAX, int PTILE>
- __global__ void ln2_gate_store_warp_f16(
- const half* __restrict__ out_acc,
- const half* __restrict__ ogate,
- const float* __restrict__ w,
- const float* __restrict__ b,
- half* __restrict__ out_norm,
+ template<int MAX_H>
+ __global__ void ln2_gate_store_p32_f16(
+ const float* __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 p0 = (int64_t)blockIdx.x * (int64_t)PTILE;
+ int64_t p_base = (int64_t)blockIdx.x * (int64_t)P_TILE;
- int tid = (int)threadIdx.x;
- int warp = tid >> 5;
+ int tid = (int)threadIdx.x; // 0..255
int lane = tid & 31;
+ int warp = tid >> 5;
+ int num_warp = (int)(blockDim.x >> 5);
- constexpr int WP = HMAX / 32;
- int p_l = warp / WP;
- int w_in = warp - p_l * WP;
- int64_t p = p0 + (int64_t)p_l;
- int d = w_in * 32 + lane;
+ // 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];
- float x = 0.0f;
- float g = 0.0f;
- if (p < nn && d < hidden) {
- int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;
- x = __half2float(out_acc[idx]);
- g = __half2float(ogate[idx]);
+ 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;
+ }
}
- float s = (d < hidden && p < nn) ? x : 0.0f;
- float ss = (d < hidden && p < nn) ? x * x : 0.0f;
- s = warp_sum(s);
- ss = warp_sum(ss);
-
- __shared__ float sh_sum[PTILE * WP];
- __shared__ float sh_sq[PTILE * WP];
- if (lane == 0) {
- sh_sum[p_l * WP + w_in] = s;
- sh_sq[p_l * WP + w_in] = ss;
+ // 每个 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 = out_acc[idx];
+ g = ogate[idx];
+ }
+ sh_x[d][lane] = x;
+ sh_g[d][lane] = g;
}
+
__syncthreads();
- __shared__ float mean_sh[PTILE];
- __shared__ float inv_sh[PTILE];
- if (lane == 0 && w_in == 0) {
+ // 计算每个 p 的 mean/var(沿 hidden 归约)。只用一个 warp 处理 32 个 p。
+ if (warp == 0) {
+ int p_l = lane; // 0..31
float sum = 0.0f;
float sq = 0.0f;
- #pragma unroll
- for (int t = 0; t < WP; ++t) {
- sum += sh_sum[p_l * WP + t];
- sq += sh_sq[p_l * WP + t];
+ 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;
⋯ 1 unchanged lines
mean_sh[p_l] = mean;
inv_sh[p_l] = rsqrtf(var + 1e-5f);
}
+
__syncthreads();
- if (p < nn && d < hidden) {
- float mean = mean_sh[p_l];
- float inv = inv_sh[p_l];
- float y = (x - mean) * inv * w[d] + b[d];
- float yg = y * g;
- int64_t row = (int64_t)bb * nn + p;
- out_norm[row * (int64_t)hidden + (int64_t)d] = __float2half_rn(yg);
+ // 写回 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;
+ }
}
}
⋯ 7 unchanged lines
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) {
- constexpr int HMAX = 32;
- constexpr int PTILE = 16;
- dim3 block(HMAX * PTILE, 1, 1);
- dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
- ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
+ ln2_gate_store_p32_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(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_32");
+ checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_32");
} else if (hidden <= 64) {
- constexpr int HMAX = 64;
- constexpr int PTILE = 8;
- dim3 block(HMAX * PTILE, 1, 1);
- dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
- ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
+ ln2_gate_store_p32_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(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_64");
- } else if (hidden <= 128) {
- constexpr int HMAX = 128;
- constexpr int PTILE = 4;
- dim3 block(HMAX * PTILE, 1, 1);
- dim3 grid((unsigned)((nn + PTILE - 1) / PTILE), (unsigned)bs, 1);
- ln2_gate_store_warp_f16<HMAX, PTILE><<<grid, block>>>(
- (const half*)out_acc.data_ptr(),
+ checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_64");
+ } else {
+ ln2_gate_store_p32_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(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_store_warp_f16_128");
- } else {
- throw std::runtime_error("hidden_dim too large");
+ checkCuda(cudaGetLastError(), "ln2_gate_store_p32_f16_128");
}
}
⋯ 1 unchanged lines
static void gemm_x_wt_f16_f16(
cublasHandle_t h,
- const half* x_row,
- const half* w_row,
- half* y_row,
+ 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) {
⋯ 15 unchanged lines
static void gemm_x_wt_f16_f32(
cublasHandle_t h,
- const half* x_row,
- const half* w_row,
- float* y_row,
+ 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) {
⋯ 13 unchanged lines
"cublasGemmEx");
}
- static void gemm_contract_batched_f16_f16(
+ static void gemm_contract_batched_f16_f32(
cublasHandle_t h,
- const half* left_row,
- const half* right_row,
- half* out_row,
+ const half* left_row, // row-major [B,M,K]
+ const half* right_row, // row-major [B,N,K]
+ float* out_row, // row-major [B,M,N]
int batch,
int n) {
float alpha = 1.0f;
⋯ 11 unchanged lines
right_row, CUDA_R_16F, n, strideB,
left_row, CUDA_R_16F, n, strideA,
&beta,
- out_row, CUDA_R_16F, n, strideC,
+ out_row, CUDA_R_32F, n, strideC,
batch,
CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
"cublasGemmStridedBatchedEx");
⋯ 45 unchanged lines
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);
+ // 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});
- auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
+ auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat32));
- gemm_contract_batched_f16_f16(
+ gemm_contract_batched_f16_f32(
h,
(const half*)left3.data_ptr(),
(const half*)right3.data_ptr(),
- (half*)out_acc.data_ptr(),
+ (float*)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));
+ // og view: [bs*hidden, n, n] 等价的连续指针
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);
⋯ 145 unchanged lines
__all__ = ["custom_kernel"]
-
scrolls · 757 diff lines total

Best evidence level for this revision: reported

JSON