Skip to content
KernelIndex
Search⌘K

submission 418692

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-418692?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.24ms
#8 of 71
2026-02-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f888329ba9ffc1639966bc9d7a47d0253845f93e2b08aa7de6e9d28e56207217
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 warp_s[2][3];
vector-width = float4const float4* x4 = reinterpret_cast<const float4*>(x + row * 128);

Kernel source

submission.py936 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_h,
    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");
    // 强制启用 tensor op math(half 输入的 GEMM 明确走张量核)
    checkCublas(cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH), "cublasSetMathMode");
  }
  ~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) {
  float z = __expf(-x);
  return 1.0f / (1.0f + z);
}

// ---------------- LN1 ----------------
// 关键优化:减少 block 数量(每个 block 处理多个 row),显著降低大 M 场景下的调度/启动开销。

__global__ void ln1_128_f16_warp8(
    const float* __restrict__ x,
    const float* __restrict__ w,
    const float* __restrict__ b,
    half* __restrict__ y,
    int64_t rows) {
  int tid = (int)threadIdx.x;     // 0..255
  int lane = tid & 31;            // 0..31
  int warp_id = tid >> 5;         // 0..7
  int64_t row = (int64_t)blockIdx.x * 8 + (int64_t)warp_id;
  if (row >= rows) return;

  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_384_f16_rows2(
    const float* __restrict__ x,
    const float* __restrict__ w,
    const float* __restrict__ b,
    half* __restrict__ y,
    int64_t rows) {
  int tid = (int)threadIdx.x;     // 0..191
  int lane = tid & 31;            // 0..31
  int warp_id = tid >> 5;         // 0..5

  int row_in_block = warp_id / 3;       // 0..1
  int warp_in_row = warp_id - row_in_block * 3; // 0..2
  int tid_row = warp_in_row * 32 + lane; // 0..95

  int64_t row = (int64_t)blockIdx.x * 2 + (int64_t)row_in_block;
  bool row_ok = row < rows;

  float4 v;
  if (row_ok) {
    const float4* x4 = reinterpret_cast<const float4*>(x + row * 384);
    v = x4[tid_row];
  } else {
    v.x = 0.0f;
    v.y = 0.0f;
    v.z = 0.0f;
    v.w = 0.0f;
  }

  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 warp_s[2][3];
  __shared__ float warp_ss[2][3];
  __shared__ float tot_s[2];
  __shared__ float tot_ss[2];

  if (lane == 0) {
    warp_s[row_in_block][warp_in_row] = s;
    warp_ss[row_in_block][warp_in_row] = ss;
  }
  __syncthreads();

  if (warp_in_row == 0) {
    float sum = 0.0f;
    float sq = 0.0f;
    if (lane < 3) {
      sum = warp_s[row_in_block][lane];
      sq = warp_ss[row_in_block][lane];
    }
    sum = warp_sum(sum);
    sq = warp_sum(sq);
    if (lane == 0) {
      tot_s[row_in_block] = sum;
      tot_ss[row_in_block] = sq;
    }
  }
  __syncthreads();

  float sum = tot_s[row_in_block];
  float sq = tot_ss[row_in_block];

  float mean = sum * (1.0f / 384.0f);
  float var = sq * (1.0f / 384.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[tid_row];
  float4 gb = b4[tid_row];

  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;

  if (row_ok) {
    half2 h0 = __floats2half2_rn(y0, y1);
    half2 h1 = __floats2half2_rn(y2, y3);
    half2* y2p = reinterpret_cast<half2*>(y + row * 384 + tid_row * 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) {
    // 8 warps/CTA;每个 warp 处理 1 个 row
    dim3 block(256, 1, 1);
    dim3 grid((unsigned)((rows + 7) / 8), 1, 1);
    ln1_128_f16_warp8<<<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_warp8");
  } else if (dim == 384) {
    // 6 warps/CTA;每 3 个 warp 处理 1 个 row(单 CTA 处理 2 个 row)
    dim3 block(192, 1, 1);
    dim3 grid((unsigned)((rows + 1) / 2), 1, 1);
    ln1_384_f16_rows2<<<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_rows2");
  } 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 (projT -> left/right) ---------------
// projT 形状为 [5H, M](row-major,最后一维 M 连续),避免额外转置。
// mask_h 为 half,节省读带宽(0/1 掩码不会引入额外误差)。

__global__ void pack4_lr_dmaj_f16(
    const half* __restrict__ projT,  // [5H, M]
    const half* __restrict__ mask_h, // [bs, nn]
    half* __restrict__ left,         // [bs*H, nn]
    half* __restrict__ right,        // [bs*H, nn]
    int64_t nn,
    int64_t M,
    int hidden) {
  int bd = (int)blockIdx.y; // 0..bs*hidden-1
  int d = bd - (bd / hidden) * hidden;
  int b = bd / hidden;

  int64_t p0 = (int64_t)blockIdx.x * (int64_t)blockDim.x * 2 + (int64_t)threadIdx.x * 2;
  if (p0 >= nn) return;

  const half* mask_row = mask_h + (int64_t)b * nn;
  int64_t idx0 = (int64_t)b * nn + p0;

  const half* row_l = projT + (int64_t)d * M;
  const half* row_r = projT + (int64_t)(hidden + d) * M;
  const half* row_gl = projT + (int64_t)(2 * hidden + d) * M;
  const half* row_gr = projT + (int64_t)(3 * hidden + d) * M;

  half* out_l = left + (int64_t)bd * nn;
  half* out_r = right + (int64_t)bd * nn;

  // 优先走 half2 向量化路径(p0 偶数对齐;末尾不足 2 个元素时走标量路径)
  if (p0 + 1 < nn) {
    half2 m2 = *reinterpret_cast<const half2*>(mask_row + p0);
    half2 l2 = *reinterpret_cast<const half2*>(row_l + idx0);
    half2 r2 = *reinterpret_cast<const half2*>(row_r + idx0);
    half2 gl2 = *reinterpret_cast<const half2*>(row_gl + idx0);
    half2 gr2 = *reinterpret_cast<const half2*>(row_gr + idx0);

    float2 mf = __half22float2(m2);
    float2 lf = __half22float2(l2);
    float2 rf = __half22float2(r2);
    float2 glf = __half22float2(gl2);
    float2 grf = __half22float2(gr2);

    glf.x = fast_sigmoid(glf.x);
    glf.y = fast_sigmoid(glf.y);
    grf.x = fast_sigmoid(grf.x);
    grf.y = fast_sigmoid(grf.y);

    float lo0 = lf.x * glf.x * mf.x;
    float lo1 = lf.y * glf.y * mf.y;
    float ro0 = rf.x * grf.x * mf.x;
    float ro1 = rf.y * grf.y * mf.y;

    *reinterpret_cast<half2*>(out_l + p0) = __floats2half2_rn(lo0, lo1);
    *reinterpret_cast<half2*>(out_r + p0) = __floats2half2_rn(ro0, ro1);
    return;
  }

  // 标量收尾
  float m0 = __half2float(mask_row[p0]);
  float l0 = __half2float(row_l[idx0]);
  float r0 = __half2float(row_r[idx0]);
  float gl0 = fast_sigmoid(__half2float(row_gl[idx0]));
  float gr0 = fast_sigmoid(__half2float(row_gr[idx0]));
  out_l[p0] = __float2half_rn(l0 * gl0 * m0);
  out_r[p0] = __float2half_rn(r0 * gr0 * m0);
}

static void launch_pack4(
    torch::Tensor projT,
    torch::Tensor mask_h,
    torch::Tensor left,
    torch::Tensor right,
    int bs,
    int n,
    int hidden) {
  int64_t nn = (int64_t)n * (int64_t)n;
  int64_t M = (int64_t)bs * nn;
  dim3 block(256, 1, 1);
  dim3 grid((unsigned)((nn + (int64_t)block.x * 2 - 1) / ((int64_t)block.x * 2)), (unsigned)(bs * hidden), 1);
  pack4_lr_dmaj_f16<<<grid, block>>>(
      (const half*)projT.data_ptr(),
      (const half*)mask_h.data_ptr(),
      (half*)left.data_ptr(),
      (half*)right.data_ptr(),
      nn,
      M,
      hidden);
  checkCuda(cudaGetLastError(), "pack4_lr_dmaj_f16");
}

// --------------- LN2 + gate ---------------
// out_acc: [bs*H, nn];gate 原始值来自 projT 的第 5 段(out_gate),避免 ogate 中间张量写回/读回。
// 输出 out_norm_T: [H, M](row-major,最后一维 M 连续)。
// 关键点:按 p 连续加载;用少量同步做跨 warp 规约。
// 本版本将 out_acc 以 half 存储,减轻 LN2 的带宽压力与 shared footprint。

template<int WARPS>
__global__ void ln2_gate_tile64_f16_h2(
    const half* __restrict__ out_acc, // [bs*H, nn] (half)
    const half* __restrict__ projT,   // [5H, M] (half)
    const float* __restrict__ w,      // [H]
    const float* __restrict__ b,      // [H]
    half* __restrict__ out_norm_T,    // [H, M]
    int64_t nn,
    int hidden) {
  int bb = (int)blockIdx.y;
  int64_t p0 = (int64_t)blockIdx.x * 64;
  int tid = (int)threadIdx.x;
  int lane = tid & 31;
  int wid = tid >> 5;
  int64_t p = p0 + (int64_t)lane * 2;
  if (p >= nn) return;
  bool pair_ok = (p + 1 < nn);

  int64_t M = nn * (int64_t)gridDim.y;
  int64_t out_p = (int64_t)bb * nn + p;

  extern __shared__ unsigned char smem_u8[];
  half2* sh_x2 = (half2*)smem_u8; // hidden*32 (half2)
  float2* sh_sum2 = (float2*)(smem_u8 + (size_t)((int64_t)hidden * 32) * sizeof(half2));
  float2* sh_sq2 = sh_sum2 + WARPS * 32;
  float2* sh_mean2 = sh_sq2 + WARPS * 32;
  float2* sh_inv2 = sh_mean2 + 32;

  float2 sum2;
  sum2.x = 0.0f;
  sum2.y = 0.0f;
  float2 sq2;
  sq2.x = 0.0f;
  sq2.y = 0.0f;

  for (int d = wid; d < hidden; d += WARPS) {
    int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;
    half2 hx2;
    if (pair_ok) {
      hx2 = *reinterpret_cast<const half2*>(out_acc + idx);
    } else {
      hx2 = __halves2half2(out_acc[idx], __float2half_rn(0.0f));
    }
    sh_x2[(int64_t)d * 32 + lane] = hx2;
    float2 xf = __half22float2(hx2);
    sum2.x += xf.x;
    sum2.y += xf.y;
    sq2.x += xf.x * xf.x;
    sq2.y += xf.y * xf.y;
  }

  sh_sum2[wid * 32 + lane] = sum2;
  sh_sq2[wid * 32 + lane] = sq2;
  __syncthreads();

  if (wid == 0) {
    float2 tot;
    tot.x = 0.0f;
    tot.y = 0.0f;
    float2 tot_sq;
    tot_sq.x = 0.0f;
    tot_sq.y = 0.0f;
#pragma unroll
    for (int w_id = 0; w_id < WARPS; ++w_id) {
      float2 s = sh_sum2[w_id * 32 + lane];
      float2 q = sh_sq2[w_id * 32 + lane];
      tot.x += s.x;
      tot.y += s.y;
      tot_sq.x += q.x;
      tot_sq.y += q.y;
    }
    float inv_n = 1.0f / (float)hidden;
    float2 mean;
    mean.x = tot.x * inv_n;
    mean.y = tot.y * inv_n;
    float2 var;
    var.x = tot_sq.x * inv_n - mean.x * mean.x;
    var.y = tot_sq.y * inv_n - mean.y * mean.y;
    sh_mean2[lane] = mean;
    float2 inv;
    inv.x = rsqrtf(var.x + 1e-5f);
    inv.y = rsqrtf(var.y + 1e-5f);
    sh_inv2[lane] = inv;
  }
  __syncthreads();

  float2 mean = sh_mean2[lane];
  float2 inv = sh_inv2[lane];

  for (int d = wid; d < hidden; d += WARPS) {
    float wv = w[d];
    float bv = b[d];
    float2 xf = __half22float2(sh_x2[(int64_t)d * 32 + lane]);
    float2 y;
    y.x = (xf.x - mean.x) * inv.x * wv + bv;
    y.y = (xf.y - mean.y) * inv.y * wv + bv;

    int64_t gate_idx = ((int64_t)(4 * hidden + d) * M) + out_p;
    half2 g2;
    if (pair_ok) {
      g2 = *reinterpret_cast<const half2*>(projT + gate_idx);
    } else {
      g2 = __halves2half2(projT[gate_idx], __float2half_rn(0.0f));
    }
    float2 gf = __half22float2(g2);
    gf.x = fast_sigmoid(gf.x);
    gf.y = fast_sigmoid(gf.y);
    y.x *= gf.x;
    y.y *= gf.y;

    int64_t out_idx = (int64_t)d * M + out_p;
    if (pair_ok) {
      *reinterpret_cast<half2*>(out_norm_T + out_idx) = __floats2half2_rn(y.x, y.y);
    } else {
      out_norm_T[out_idx] = __float2half_rn(y.x);
    }
  }
}

static void launch_ln2(
    torch::Tensor out_acc,
    torch::Tensor projT,
    torch::Tensor w,
    torch::Tensor b,
    torch::Tensor out_norm_T,
    int bs,
    int n,
    int hidden) {
  int64_t nn = (int64_t)n * (int64_t)n;
  dim3 grid((unsigned)((nn + 63) / 64), (unsigned)bs, 1);

  if (hidden <= 32) {
    constexpr int WARPS = 1;
    dim3 block(WARPS * 32, 1, 1);
    size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
                 + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
    ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
        (const half*)out_acc.data_ptr(),
        (const half*)projT.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm_T.data_ptr(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_1w");
  } else if (hidden <= 64) {
    constexpr int WARPS = 2;
    dim3 block(WARPS * 32, 1, 1);
    size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
                 + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
    ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
        (const half*)out_acc.data_ptr(),
        (const half*)projT.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm_T.data_ptr(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_2w");
  } else if (hidden <= 128) {
    constexpr int WARPS = 4;
    dim3 block(WARPS * 32, 1, 1);
    size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
                 + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
    ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
        (const half*)out_acc.data_ptr(),
        (const half*)projT.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm_T.data_ptr(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_4w");
  } else if (hidden <= 256) {
    constexpr int WARPS = 8;
    dim3 block(WARPS * 32, 1, 1);
    size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
                 + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
    ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
        (const half*)out_acc.data_ptr(),
        (const half*)projT.data_ptr(),
        (const float*)w.data_ptr(),
        (const float*)b.data_ptr(),
        (half*)out_norm_T.data_ptr(),
        nn,
        hidden);
    checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_8w");
  } else {
    throw std::runtime_error("hidden_dim too large");
  }
}

// ---------------- GEMM helpers ----------------

// 约定:所有矩阵都来自 row-major Tensor,但用 cuBLAS 的 column-major 语义解释,
// 通过精确设置 m/n/k 与 lda/ldb/ldc 得到想要的布局,避免额外转置核。

static void gemm1_x_wt_to_dmaj_f16(
    cublasHandle_t h,
    const half* x_rm,  // [M, K] row-major
    const half* w_rm,  // [N, K] row-major
    half* c_dmaj_rm,   // [N, M] row-major (等价于 column-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)M, (int)N, (int)K,
          &alpha,
          x_rm, CUDA_R_16F, (int)K,
          w_rm, CUDA_R_16F, (int)K,
          &beta,
          c_dmaj_rm, CUDA_R_16F, (int)M,
          CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
      "cublasGemmEx_gemm1");
}

static void gemm2_dmaj_to_y_t_f16_f16(
    cublasHandle_t h,
    const half* a_dmaj_rm, // [K, M] row-major (等价于 column-major [M, K])
    const half* w_rm,      // [N, K] row-major (等价于 column-major [K, N])
    half* y_t_rm,          // [N, M] row-major (等价于 column-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_N, CUBLAS_OP_N,
          (int)M, (int)N, (int)K,
          &alpha,
          a_dmaj_rm, CUDA_R_16F, (int)M,
          w_rm, CUDA_R_16F, (int)K,
          &beta,
          y_t_rm, CUDA_R_16F, (int)M,
          CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
      "cublasGemmEx_gemm2");
}

static void gemm_contract_batched_f16_f16(
    cublasHandle_t h,
    const half* left_row,   // [B,M,K] row-major [batch,n,n]
    const half* right_row,  // [B,N,K] row-major [batch,n,n]
    half* out_row,          // [B,M,N] row-major [batch,n,n] (half)
    int batch,
    int n) {
  float alpha = 1.0f;
  float beta = 0.0f;
  long long strideA = (long long)n * (long long)n;
  long long strideB = (long long)n * (long long)n;
  long long strideC = (long long)n * (long long)n;

  checkCublas(
      cublasGemmStridedBatchedEx(
          h,
          CUBLAS_OP_T, CUBLAS_OP_N,
          n, n, n,
          &alpha,
          right_row, CUDA_R_16F, n, strideB,
          left_row, CUDA_R_16F, n, strideA,
          &beta,
          out_row, CUDA_R_16F, n, strideC,
          batch,
          CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
      "cublasGemmStridedBatchedEx_contract");
}

} // namespace

torch::Tensor trimul_fwd(
    torch::Tensor x,
    torch::Tensor mask_h,
    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_h.is_cuda()) {
    throw std::runtime_error("cuda only");
  }
  if (x.scalar_type() != torch::kFloat32) {
    throw std::runtime_error("x must be float32");
  }
  if (mask_h.scalar_type() != torch::kFloat16) {
    throw std::runtime_error("mask must be float16");
  }
  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 nn = (int64_t)n * (int64_t)n;
  int64_t M = (int64_t)bs * nn;

  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;
  // projT: [5H, M](d-major,最后一维连续)
  auto projT = torch::empty({out_ch, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));

  auto* holder = get_cublas();
  cublasHandle_t h = holder->handle;
  gemm1_x_wt_to_dmaj_f16(
      h,
      (const half*)xhat.data_ptr(),
      (const half*)w_cat.data_ptr(),
      (half*)projT.data_ptr(),
      M,
      out_ch,
      dim);

  auto left = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  auto right = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));

  // pack:projT[前 4 段, M] + mask[bs,nn] -> left/right[bs*H,nn]
  // out_gate(第 5 段)留在 projT 中,后续由 LN2 kernel 直接读取并 sigmoid。
  launch_pack4(projT, mask_h.view({bs, nn}), left, right, bs, n, (int)hidden);

  auto left3 = left.view({bs * (int)hidden, n, n});
  auto right3 = right.view({bs * (int)hidden, n, n});
  // out_acc: half(仍由 GEMM 做 FP32 累加,只降低写回精度)
  auto out_acc = torch::empty({bs * (int)hidden, n, n}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));

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

  // LN2 + gate:输出 out_norm_T[H,M] half
  auto out_norm_T = torch::empty({hidden, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  launch_ln2(out_acc, projT, ln2_w, ln2_b, out_norm_T, bs, n, (int)hidden);

  // gemm2:y_T[dim,M] half(降低最终写回带宽;仍 FP32 累加)
  auto y_T = torch::empty({dim, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
  gemm2_dmaj_to_y_t_f16_f16(
      h,
      (const half*)out_norm_T.data_ptr(),
      (const half*)w_out.data_ptr(),
      (half*)y_T.data_ptr(),
      M,
      dim,
      hidden);

  return y_T.view({dim, bs, n, n}).permute({1, 2, 3, 0});
}
"""

        name = "trimul_ext_mod8"

        
        
        extra_cuda_cflags = [
            "-O2",
            "--use_fast_math",
        ]
        extra_cflags = [
            "-O2",
        ]

        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.float16:
        mask = mask.to(dtype=torch.float16)

    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 · 936 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 418614.

⋯ 110 unchanged lines
}
// ---------------- LN1 ----------------
+ // 关键优化:减少 block 数量(每个 block 处理多个 row),显著降低大 M 场景下的调度/启动开销。
- __global__ void ln1_128_f16(
+ __global__ void ln1_128_f16_warp8(
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;
+ int tid = (int)threadIdx.x; // 0..255
+ int lane = tid & 31; // 0..31
+ int warp_id = tid >> 5; // 0..7
+ int64_t row = (int64_t)blockIdx.x * 8 + (int64_t)warp_id;
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];
⋯ 25 unchanged lines
y2p[1] = h1;
}
- __global__ void ln1_384_f16(
+ __global__ void ln1_384_f16_rows2(
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 tid = (int)threadIdx.x; // 0..191
int lane = tid & 31; // 0..31
- int warp_id = tid >> 5; // 0..2
+ int warp_id = tid >> 5; // 0..5
- const float4* x4 = reinterpret_cast<const float4*>(x + row * 384);
- float4 v = x4[tid];
+ int row_in_block = warp_id / 3; // 0..1
+ int warp_in_row = warp_id - row_in_block * 3; // 0..2
+ int tid_row = warp_in_row * 32 + lane; // 0..95
+
+ int64_t row = (int64_t)blockIdx.x * 2 + (int64_t)row_in_block;
+ bool row_ok = row < rows;
+
+ float4 v;
+ if (row_ok) {
+ const float4* x4 = reinterpret_cast<const float4*>(x + row * 384);
+ v = x4[tid_row];
+ } else {
+ v.x = 0.0f;
+ v.y = 0.0f;
+ v.z = 0.0f;
+ v.w = 0.0f;
+ }
+
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 warp_s[3];
- __shared__ float warp_ss[3];
- __shared__ float tot_s;
- __shared__ float tot_ss;
+ __shared__ float warp_s[2][3];
+ __shared__ float warp_ss[2][3];
+ __shared__ float tot_s[2];
+ __shared__ float tot_ss[2];
+
if (lane == 0) {
- warp_s[warp_id] = s;
- warp_ss[warp_id] = ss;
+ warp_s[row_in_block][warp_in_row] = s;
+ warp_ss[row_in_block][warp_in_row] = ss;
}
__syncthreads();
- float sum = 0.0f;
- float sq = 0.0f;
- if (warp_id == 0) {
+ if (warp_in_row == 0) {
+ float sum = 0.0f;
+ float sq = 0.0f;
if (lane < 3) {
- sum = warp_s[lane];
- sq = warp_ss[lane];
+ sum = warp_s[row_in_block][lane];
+ sq = warp_ss[row_in_block][lane];
}
sum = warp_sum(sum);
sq = warp_sum(sq);
+ if (lane == 0) {
+ tot_s[row_in_block] = sum;
+ tot_ss[row_in_block] = sq;
+ }
}
- if (warp_id == 0 && lane == 0) {
- tot_s = sum;
- tot_ss = sq;
- }
__syncthreads();
- sum = tot_s;
- sq = tot_ss;
+ float sum = tot_s[row_in_block];
+ float sq = tot_ss[row_in_block];
+
float mean = sum * (1.0f / 384.0f);
float var = sq * (1.0f / 384.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[tid];
- float4 gb = b4[tid];
+ float4 gw = w4[tid_row];
+ float4 gb = b4[tid_row];
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 * 384 + tid * 4);
- y2p[0] = h0;
- y2p[1] = h1;
+ if (row_ok) {
+ half2 h0 = __floats2half2_rn(y0, y1);
+ half2 h1 = __floats2half2_rn(y2, y3);
+ half2* y2p = reinterpret_cast<half2*>(y + row * 384 + tid_row * 4);
+ y2p[0] = h0;
+ y2p[1] = h1;
+ }
}
__global__ void ln1_generic_f16(
⋯ 45 unchanged lines
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>>>(
+ // 8 warps/CTA;每个 warp 处理 1 个 row
+ dim3 block(256, 1, 1);
+ dim3 grid((unsigned)((rows + 7) / 8), 1, 1);
+ ln1_128_f16_warp8<<<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");
+ checkCuda(cudaGetLastError(), "ln1_128_f16_warp8");
} else if (dim == 384) {
- dim3 block(96, 1, 1);
- dim3 grid((unsigned)rows, 1, 1);
- ln1_384_f16<<<grid, block>>>(
+ // 6 warps/CTA;每 3 个 warp 处理 1 个 row(单 CTA 处理 2 个 row)
+ dim3 block(192, 1, 1);
+ dim3 grid((unsigned)((rows + 1) / 2), 1, 1);
+ ln1_384_f16_rows2<<<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");
+ checkCuda(cudaGetLastError(), "ln1_384_f16_rows2");
} else {
dim3 block(256, 1, 1);
dim3 grid((unsigned)rows, 1, 1);
⋯ 10 unchanged lines
// --------------- pack (projT -> left/right) ---------------
// projT 形状为 [5H, M](row-major,最后一维 M 连续),避免额外转置。
- // 关键点:不再物化 out_gate(ogate)到单独张量,减少一次全量读写;
- // LN2 阶段直接从 projT 的第 5 组读取并做 sigmoid。
+ // mask_h 为 half,节省读带宽(0/1 掩码不会引入额外误差)。
- __global__ void pack4_dmaj_f16(
+ __global__ void pack4_lr_dmaj_f16(
const half* __restrict__ projT, // [5H, M]
const half* __restrict__ mask_h, // [bs, nn]
half* __restrict__ left, // [bs*H, nn]
half* __restrict__ right, // [bs*H, nn]
- int nn,
- int M,
+ int64_t nn,
+ int64_t M,
int hidden) {
int bd = (int)blockIdx.y; // 0..bs*hidden-1
int d = bd - (bd / hidden) * hidden;
int b = bd / hidden;
- int p0 = (int)blockIdx.x * ((int)blockDim.x * 2) + (int)threadIdx.x * 2;
+ int64_t p0 = (int64_t)blockIdx.x * (int64_t)blockDim.x * 2 + (int64_t)threadIdx.x * 2;
if (p0 >= nn) return;
- const half* mask_row = mask_h + b * nn;
- int idx0 = b * nn + p0;
+ const half* mask_row = mask_h + (int64_t)b * nn;
+ int64_t idx0 = (int64_t)b * nn + p0;
- const half* row_l = projT + (d * M);
- const half* row_r = projT + ((hidden + d) * M);
- const half* row_gl = projT + ((2 * hidden + d) * M);
- const half* row_gr = projT + ((3 * hidden + d) * M);
+ const half* row_l = projT + (int64_t)d * M;
+ const half* row_r = projT + (int64_t)(hidden + d) * M;
+ const half* row_gl = projT + (int64_t)(2 * hidden + d) * M;
+ const half* row_gr = projT + (int64_t)(3 * hidden + d) * M;
- half* out_l = left + bd * nn;
- half* out_r = right + bd * nn;
+ half* out_l = left + (int64_t)bd * nn;
+ half* out_r = right + (int64_t)bd * nn;
+ // 优先走 half2 向量化路径(p0 偶数对齐;末尾不足 2 个元素时走标量路径)
if (p0 + 1 < nn) {
half2 m2 = *reinterpret_cast<const half2*>(mask_row + p0);
half2 l2 = *reinterpret_cast<const half2*>(row_l + idx0);
⋯ 22 unchanged lines
return;
}
+ // 标量收尾
float m0 = __half2float(mask_row[p0]);
float l0 = __half2float(row_l[idx0]);
float r0 = __half2float(row_r[idx0]);
⋯ 11 unchanged lines
int bs,
int n,
int hidden) {
- int nn = n * n;
- int M = bs * nn;
+ int64_t nn = (int64_t)n * (int64_t)n;
+ int64_t M = (int64_t)bs * nn;
dim3 block(256, 1, 1);
- dim3 grid((unsigned)((nn + (int)block.x * 2 - 1) / ((int)block.x * 2)), (unsigned)(bs * hidden), 1);
- pack4_dmaj_f16<<<grid, block>>>(
+ dim3 grid((unsigned)((nn + (int64_t)block.x * 2 - 1) / ((int64_t)block.x * 2)), (unsigned)(bs * hidden), 1);
+ pack4_lr_dmaj_f16<<<grid, block>>>(
(const half*)projT.data_ptr(),
(const half*)mask_h.data_ptr(),
(half*)left.data_ptr(),
⋯ 1 unchanged lines
nn,
M,
hidden);
- checkCuda(cudaGetLastError(), "pack4_dmaj_f16");
+ checkCuda(cudaGetLastError(), "pack4_lr_dmaj_f16");
}
// --------------- LN2 + gate ---------------
- // out_acc: [bs*H, nn];go_T: [H, M](d-major,最后一维 M 连续,存的是 pre-sigmoid)
+ // out_acc: [bs*H, nn];gate 原始值来自 projT 的第 5 段(out_gate),避免 ogate 中间张量写回/读回。
// 输出 out_norm_T: [H, M](row-major,最后一维 M 连续)。
- // 本版本不读取/存储 ogate 中间张量,直接从 go_T 做 sigmoid。
+ // 关键点:按 p 连续加载;用少量同步做跨 warp 规约。
+ // 本版本将 out_acc 以 half 存储,减轻 LN2 的带宽压力与 shared footprint。
+
template<int WARPS>
- __global__ void ln2_gate_tile32_f16(
+ __global__ void ln2_gate_tile64_f16_h2(
const half* __restrict__ out_acc, // [bs*H, nn] (half)
- const half* __restrict__ go_T, // [H, M] (half, pre-sigmoid)
+ const half* __restrict__ projT, // [5H, M] (half)
const float* __restrict__ w, // [H]
const float* __restrict__ b, // [H]
half* __restrict__ out_norm_T, // [H, M]
- int nn,
+ int64_t nn,
int hidden) {
int bb = (int)blockIdx.y;
- int p0 = (int)blockIdx.x * 32;
+ int64_t p0 = (int64_t)blockIdx.x * 64;
int tid = (int)threadIdx.x;
int lane = tid & 31;
int wid = tid >> 5;
- int p = p0 + lane;
+ int64_t p = p0 + (int64_t)lane * 2;
+ if (p >= nn) return;
+ bool pair_ok = (p + 1 < nn);
- int M = nn * (int)gridDim.y;
- int out_p = bb * nn + p;
+ int64_t M = nn * (int64_t)gridDim.y;
+ int64_t out_p = (int64_t)bb * nn + p;
extern __shared__ unsigned char smem_u8[];
- half* sh_x = (half*)smem_u8; // hidden*32
- float* sh_sum = (float*)(smem_u8 + (size_t)((int64_t)hidden * 32) * sizeof(half));
- float* sh_sq = sh_sum + WARPS * 32;
- float* sh_mean = sh_sq + WARPS * 32;
- float* sh_inv = sh_mean + 32;
+ half2* sh_x2 = (half2*)smem_u8; // hidden*32 (half2)
+ float2* sh_sum2 = (float2*)(smem_u8 + (size_t)((int64_t)hidden * 32) * sizeof(half2));
+ float2* sh_sq2 = sh_sum2 + WARPS * 32;
+ float2* sh_mean2 = sh_sq2 + WARPS * 32;
+ float2* sh_inv2 = sh_mean2 + 32;
- float sum = 0.0f;
- float sq = 0.0f;
+ float2 sum2;
+ sum2.x = 0.0f;
+ sum2.y = 0.0f;
+ float2 sq2;
+ sq2.x = 0.0f;
+ sq2.y = 0.0f;
for (int d = wid; d < hidden; d += WARPS) {
- half hx = __float2half_rn(0.0f);
- if (p < nn) {
- int idx = (bb * hidden + d) * nn + p;
- hx = out_acc[idx];
+ int64_t idx = ((int64_t)bb * (int64_t)hidden + (int64_t)d) * nn + p;
+ half2 hx2;
+ if (pair_ok) {
+ hx2 = *reinterpret_cast<const half2*>(out_acc + idx);
+ } else {
+ hx2 = __halves2half2(out_acc[idx], __float2half_rn(0.0f));
}
- sh_x[(int64_t)d * 32 + lane] = hx;
- float x = __half2float(hx);
- sum += x;
- sq += x * x;
+ sh_x2[(int64_t)d * 32 + lane] = hx2;
+ float2 xf = __half22float2(hx2);
+ sum2.x += xf.x;
+ sum2.y += xf.y;
+ sq2.x += xf.x * xf.x;
+ sq2.y += xf.y * xf.y;
}
- sh_sum[wid * 32 + lane] = sum;
- sh_sq[wid * 32 + lane] = sq;
+ sh_sum2[wid * 32 + lane] = sum2;
+ sh_sq2[wid * 32 + lane] = sq2;
__syncthreads();
if (wid == 0) {
- float tot = 0.0f;
- float tot_sq = 0.0f;
+ float2 tot;
+ tot.x = 0.0f;
+ tot.y = 0.0f;
+ float2 tot_sq;
+ tot_sq.x = 0.0f;
+ tot_sq.y = 0.0f;
#pragma unroll
for (int w_id = 0; w_id < WARPS; ++w_id) {
- tot += sh_sum[w_id * 32 + lane];
- tot_sq += sh_sq[w_id * 32 + lane];
+ float2 s = sh_sum2[w_id * 32 + lane];
+ float2 q = sh_sq2[w_id * 32 + lane];
+ tot.x += s.x;
+ tot.y += s.y;
+ tot_sq.x += q.x;
+ tot_sq.y += q.y;
}
float inv_n = 1.0f / (float)hidden;
- float mean = tot * inv_n;
- float var = tot_sq * inv_n - mean * mean;
- sh_mean[lane] = mean;
- sh_inv[lane] = rsqrtf(var + 1e-5f);
+ float2 mean;
+ mean.x = tot.x * inv_n;
+ mean.y = tot.y * inv_n;
+ float2 var;
+ var.x = tot_sq.x * inv_n - mean.x * mean.x;
+ var.y = tot_sq.y * inv_n - mean.y * mean.y;
+ sh_mean2[lane] = mean;
+ float2 inv;
+ inv.x = rsqrtf(var.x + 1e-5f);
+ inv.y = rsqrtf(var.y + 1e-5f);
+ sh_inv2[lane] = inv;
}
__syncthreads();
- if (p >= nn) return;
+ float2 mean = sh_mean2[lane];
+ float2 inv = sh_inv2[lane];
- float mean = sh_mean[lane];
- float inv = sh_inv[lane];
-
for (int d = wid; d < hidden; d += WARPS) {
- float x = __half2float(sh_x[(int64_t)d * 32 + lane]);
- float y = (x - mean) * inv * w[d] + b[d];
- float g = fast_sigmoid(__half2float(go_T[d * M + out_p]));
- out_norm_T[d * M + out_p] = __float2half_rn(y * g);
+ float wv = w[d];
+ float bv = b[d];
+ float2 xf = __half22float2(sh_x2[(int64_t)d * 32 + lane]);
+ float2 y;
+ y.x = (xf.x - mean.x) * inv.x * wv + bv;
+ y.y = (xf.y - mean.y) * inv.y * wv + bv;
+
+ int64_t gate_idx = ((int64_t)(4 * hidden + d) * M) + out_p;
+ half2 g2;
+ if (pair_ok) {
+ g2 = *reinterpret_cast<const half2*>(projT + gate_idx);
+ } else {
+ g2 = __halves2half2(projT[gate_idx], __float2half_rn(0.0f));
+ }
+ float2 gf = __half22float2(g2);
+ gf.x = fast_sigmoid(gf.x);
+ gf.y = fast_sigmoid(gf.y);
+ y.x *= gf.x;
+ y.y *= gf.y;
+
+ int64_t out_idx = (int64_t)d * M + out_p;
+ if (pair_ok) {
+ *reinterpret_cast<half2*>(out_norm_T + out_idx) = __floats2half2_rn(y.x, y.y);
+ } else {
+ out_norm_T[out_idx] = __float2half_rn(y.x);
+ }
}
}
static void launch_ln2(
torch::Tensor out_acc,
- torch::Tensor go_T,
+ torch::Tensor projT,
torch::Tensor w,
torch::Tensor b,
torch::Tensor out_norm_T,
int bs,
int n,
int hidden) {
- int nn = n * n;
- dim3 grid((unsigned)((nn + 31) / 32), (unsigned)bs, 1);
+ int64_t nn = (int64_t)n * (int64_t)n;
+ dim3 grid((unsigned)((nn + 63) / 64), (unsigned)bs, 1);
if (hidden <= 32) {
constexpr int WARPS = 1;
dim3 block(WARPS * 32, 1, 1);
- size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
- + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
- ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
+ size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
+ + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
+ ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
- (const half*)go_T.data_ptr(),
+ (const half*)projT.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_1w");
+ checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_1w");
} else if (hidden <= 64) {
constexpr int WARPS = 2;
dim3 block(WARPS * 32, 1, 1);
- size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
- + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
- ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
+ size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
+ + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
+ ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
- (const half*)go_T.data_ptr(),
+ (const half*)projT.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_2w");
+ checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_2w");
} else if (hidden <= 128) {
constexpr int WARPS = 4;
dim3 block(WARPS * 32, 1, 1);
- size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
- + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
- ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
+ size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
+ + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
+ ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
- (const half*)go_T.data_ptr(),
+ (const half*)projT.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_4w");
+ checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_4w");
} else if (hidden <= 256) {
constexpr int WARPS = 8;
dim3 block(WARPS * 32, 1, 1);
- size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half)
- + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float);
- ln2_gate_tile32_f16<WARPS><<<grid, block, shmem>>>(
+ size_t shmem = (size_t)((int64_t)hidden * 32) * sizeof(half2)
+ + (size_t)((int64_t)WARPS * 32 * 2 + 64) * sizeof(float2);
+ ln2_gate_tile64_f16_h2<WARPS><<<grid, block, shmem>>>(
(const half*)out_acc.data_ptr(),
- (const half*)go_T.data_ptr(),
+ (const half*)projT.data_ptr(),
(const float*)w.data_ptr(),
(const float*)b.data_ptr(),
(half*)out_norm_T.data_ptr(),
nn,
hidden);
- checkCuda(cudaGetLastError(), "ln2_gate_tile32_f16_8w");
+ checkCuda(cudaGetLastError(), "ln2_gate_tile64_f16_h2_8w");
} else {
throw std::runtime_error("hidden_dim too large");
}
⋯ 136 unchanged lines
auto left = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
auto right = torch::empty({bs * (int)hidden, nn}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- // pack:projT[5H,M] + mask[bs,nn] -> left/right[bs*H,nn]
+ // pack:projT[前 4 段, M] + mask[bs,nn] -> left/right[bs*H,nn]
+ // out_gate(第 5 段)留在 projT 中,后续由 LN2 kernel 直接读取并 sigmoid。
launch_pack4(projT, mask_h.view({bs, nn}), left, right, bs, n, (int)hidden);
auto left3 = left.view({bs * (int)hidden, n, n});
⋯ 8 unchanged lines
(half*)out_acc.data_ptr(),
bs * (int)hidden, n);
- // go_T: [H, M],从 projT 的第 5 组切片得到(pre-sigmoid)
- auto go_T = projT.narrow(0, (int64_t)4 * hidden, hidden);
-
// LN2 + gate:输出 out_norm_T[H,M] half
auto out_norm_T = torch::empty({hidden, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
- launch_ln2(out_acc, go_T, ln2_w, ln2_b, out_norm_T, bs, n, (int)hidden);
+ launch_ln2(out_acc, projT, ln2_w, ln2_b, out_norm_T, bs, n, (int)hidden);
// gemm2:y_T[dim,M] half(降低最终写回带宽;仍 FP32 累加)
auto y_T = torch::empty({dim, M}, torch::TensorOptions().device(x.device()).dtype(torch::kFloat16));
⋯ 10 unchanged lines
}
"""
- name = "trimul_ext_mod9"
+ name = "trimul_ext_mod8"
+
+
extra_cuda_cflags = [
- "-O3",
+ "-O2",
"--use_fast_math",
]
extra_cflags = [
- "-O3",
+ "-O2",
]
extra_ldflags = [
⋯ 130 unchanged lines
__all__ = ["custom_kernel"]
-
scrolls · 576 diff lines total

Best evidence level for this revision: reported

JSON