Skip to content
KernelIndex
Search⌘K

submission 415223

novo_force · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
589.4µs
#6 of 43
2026-01-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:56e0db772e58325ba800a07d699f1a6c99ff56b7359b59966be83c21d5286abf
license declaredunknown
license concludedunknown
authorsnovo_force
imported2026-08-15

Techniques

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

shared-memory__shared__ float warp_sum[4];
vector-width = float2__device__ __forceinline__ float2 _sigmoid_f2(float2 v) {

Kernel source

submission.py1131 lines
from __future__ import annotations

from typing import Any, Dict, Tuple

import torch

_EXT = None


def _get_ext():
    global _EXT
    if _EXT is not None:
        return _EXT

    from torch.utils.cpp_extension import load_inline

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

#include <type_traits>

static inline void _ck(bool ok, const char* msg) {
  if (!ok) { throw std::runtime_error(msg); }
}

static inline void _ck_tensor_cuda_contig(const torch::Tensor& t) {
  _ck(t.is_cuda(), "tensor must be CUDA");
  _ck(t.is_contiguous(), "tensor must be contiguous");
}

static inline void _ck_cublas(cublasStatus_t st) {
  if (st != CUBLAS_STATUS_SUCCESS) {
    throw std::runtime_error("cublas call failed");
  }
}

static inline cublasHandle_t _get_handle_tc() {
  cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
  static thread_local cublasHandle_t last = nullptr;
  if (h != last) {
    _ck_cublas(cublasSetMathMode(h, CUBLAS_TENSOR_OP_MATH));
    last = h;
  }
  return h;
}

static inline cublasComputeType_t _get_ct_fast() {
#if defined(CUBLAS_COMPUTE_32F_FAST_16F)
  return CUBLAS_COMPUTE_32F_FAST_16F;
#else
  return CUBLAS_COMPUTE_32F;
#endif
}

// Sigmoid:保持与参考实现一致的 fast-math 路径
__device__ __forceinline__ float _sigmoid_f(float x) {
  return __fdividef(1.0f, 1.0f + __expf(-x));
}

__device__ __forceinline__ float2 _sigmoid_f2(float2 v) {
  v.x = _sigmoid_f(v.x);
  v.y = _sigmoid_f(v.y);
  return v;
}

template <typename MaskT>
__device__ __forceinline__ float _mask_to_f32(MaskT v) {
  return static_cast<float>(v);
}

template <>
__device__ __forceinline__ float _mask_to_f32<__half>(__half v) {
  return __half2float(v);
}

template <>
__device__ __forceinline__ float _mask_to_f32<bool>(bool v) {
  return v ? 1.0f : 0.0f;
}

template <typename MaskT>
__global__ void _mask_gate_lr_fuse_f16_vec4(
    __half* __restrict__ left,
    __half* __restrict__ right,
    const __half* __restrict__ left_gate,
    const __half* __restrict__ right_gate,
    const MaskT* __restrict__ mask,
    int inner) {
  const int d = (int)blockIdx.y;
  const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
  const int col = t << 2;
  if (col >= inner) return;
  const int idx = d * inner + col;

  if (col + 3 < inner) {
    float m0, m1, m2, m3;
    if constexpr (std::is_same<MaskT, float>::value) {
      const float4 mv = *(const float4*)(mask + col);
      m0 = mv.x; m1 = mv.y; m2 = mv.z; m3 = mv.w;
    } else {
      m0 = _mask_to_f32<MaskT>(mask[col]);
      m1 = _mask_to_f32<MaskT>(mask[col + 1]);
      m2 = _mask_to_f32<MaskT>(mask[col + 2]);
      m3 = _mask_to_f32<MaskT>(mask[col + 3]);
    }

    const __half2 l2_0 = *(const __half2*)(left + idx);
    const __half2 l2_1 = *(const __half2*)(left + idx + 2);
    const __half2 r2_0 = *(const __half2*)(right + idx);
    const __half2 r2_1 = *(const __half2*)(right + idx + 2);
    const __half2 lg2_0 = *(const __half2*)(left_gate + idx);
    const __half2 lg2_1 = *(const __half2*)(left_gate + idx + 2);
    const __half2 rg2_0 = *(const __half2*)(right_gate + idx);
    const __half2 rg2_1 = *(const __half2*)(right_gate + idx + 2);

    const float2 gl0 = _sigmoid_f2(__half22float2(lg2_0));
    const float2 gl1 = _sigmoid_f2(__half22float2(lg2_1));
    const float2 gr0 = _sigmoid_f2(__half22float2(rg2_0));
    const float2 gr1 = _sigmoid_f2(__half22float2(rg2_1));

    float2 lv0 = __half22float2(l2_0);
    float2 lv1 = __half22float2(l2_1);
    float2 rv0 = __half22float2(r2_0);
    float2 rv1 = __half22float2(r2_1);

    lv0.x = lv0.x * m0 * gl0.x;
    lv0.y = lv0.y * m1 * gl0.y;
    lv1.x = lv1.x * m2 * gl1.x;
    lv1.y = lv1.y * m3 * gl1.y;

    rv0.x = rv0.x * m0 * gr0.x;
    rv0.y = rv0.y * m1 * gr0.y;
    rv1.x = rv1.x * m2 * gr1.x;
    rv1.y = rv1.y * m3 * gr1.y;

    *(__half2*)(left + idx) = __floats2half2_rn(lv0.x, lv0.y);
    *(__half2*)(left + idx + 2) = __floats2half2_rn(lv1.x, lv1.y);
    *(__half2*)(right + idx) = __floats2half2_rn(rv0.x, rv0.y);
    *(__half2*)(right + idx + 2) = __floats2half2_rn(rv1.x, rv1.y);
  } else {
    #pragma unroll
    for (int off = 0; off < 4; ++off) {
      const int c = col + off;
      if (c < inner) {
        const float m = _mask_to_f32<MaskT>(mask[c]);
        const int id = idx + off;
        float l = __half2float(left[id]) * m;
        float r = __half2float(right[id]) * m;
        const float gl = _sigmoid_f(__half2float(left_gate[id]));
        const float gr = _sigmoid_f(__half2float(right_gate[id]));
        l *= gl;
        r *= gr;
        left[id] = __float2half_rn(l);
        right[id] = __float2half_rn(r);
      }
    }
  }
}

void apply_mask_gate_lr_f16(torch::Tensor left,
                            torch::Tensor right,
                            torch::Tensor left_gate,
                            torch::Tensor right_gate,
                            torch::Tensor mask) {
  _ck_tensor_cuda_contig(left);
  _ck_tensor_cuda_contig(right);
  _ck_tensor_cuda_contig(left_gate);
  _ck_tensor_cuda_contig(right_gate);
  _ck_tensor_cuda_contig(mask);

  _ck(left.dtype() == torch::kFloat16, "left must be float16");
  _ck(right.dtype() == torch::kFloat16, "right must be float16");
  _ck(left_gate.dtype() == torch::kFloat16, "left_gate must be float16");
  _ck(right_gate.dtype() == torch::kFloat16, "right_gate must be float16");
  _ck(mask.dim() == 3, "mask must be 3D");

  const int hidden = (int)left.size(0);
  _ck(hidden == 128, "hidden_dim must be 128");
  _ck(right.numel() == left.numel(), "lr size mismatch");
  _ck(left_gate.numel() == left.numel(), "lg size mismatch");
  _ck(right_gate.numel() == left.numel(), "rg size mismatch");

  const int64_t inner64 = mask.numel();
  _ck(inner64 > 0 && inner64 <= INT_MAX, "mask too large");
  const int inner = (int)inner64;
  _ck((int64_t)hidden * (int64_t)inner == left.numel(), "mask/hidden mismatch");

  const int quads = (inner + 3) >> 2;
  const dim3 block(256, 1, 1);
  const dim3 grid((quads + (int)block.x - 1) / (int)block.x, hidden, 1);

  const auto st = mask.scalar_type();
  if (st == torch::kFloat32) {
    _mask_gate_lr_fuse_f16_vec4<float><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const float*)mask.data_ptr<float>(),
        inner);
  } else if (st == torch::kFloat16) {
    _mask_gate_lr_fuse_f16_vec4<__half><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const __half*)mask.data_ptr<at::Half>(),
        inner);
  } else if (st == torch::kInt64) {
    _mask_gate_lr_fuse_f16_vec4<int64_t><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const int64_t*)mask.data_ptr<int64_t>(),
        inner);
  } else if (st == torch::kInt32) {
    _mask_gate_lr_fuse_f16_vec4<int32_t><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const int32_t*)mask.data_ptr<int32_t>(),
        inner);
  } else if (st == torch::kUInt8) {
    _mask_gate_lr_fuse_f16_vec4<uint8_t><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const uint8_t*)mask.data_ptr<uint8_t>(),
        inner);
  } else if (st == torch::kBool) {
    _mask_gate_lr_fuse_f16_vec4<bool><<<grid, block>>>(
        (__half*)left.data_ptr<at::Half>(),
        (__half*)right.data_ptr<at::Half>(),
        (const __half*)left_gate.data_ptr<at::Half>(),
        (const __half*)right_gate.data_ptr<at::Half>(),
        (const bool*)mask.data_ptr<bool>(),
        inner);
  } else {
    throw std::runtime_error("unsupported mask dtype");
  }
}

// X: [M, K] 行主序(f16)
// W: [N, K] 行主序(f16)
// Y: [M, N] 行主序(f16)
torch::Tensor gemm_f16(torch::Tensor x, torch::Tensor w) {
  _ck_tensor_cuda_contig(x);
  _ck_tensor_cuda_contig(w);
  _ck(x.dtype() == torch::kFloat16, "x must be float16");
  _ck(w.dtype() == torch::kFloat16, "w must be float16");
  _ck(x.dim() == 2, "x must be 2D");
  _ck(w.dim() == 2, "w must be 2D");

  const int64_t M64 = x.size(0);
  const int64_t K64 = x.size(1);
  const int64_t N64 = w.size(0);
  _ck(w.size(1) == K64, "w shape mismatch");
  _ck(M64 > 0 && N64 > 0 && K64 > 0, "empty mat");
  _ck(M64 <= INT_MAX && N64 <= INT_MAX && K64 <= INT_MAX, "mat too large");

  auto y = torch::empty({M64, N64}, x.options());

  const int M = (int)M64;
  const int N = (int)N64;
  const int K = (int)K64;

  cublasHandle_t handle = _get_handle_tc();
  const cublasComputeType_t ct = _get_ct_fast();

  const float alpha = 1.0f;
  const float beta = 0.0f;

  _ck_cublas(
      cublasGemmEx(
          handle,
          CUBLAS_OP_T, CUBLAS_OP_N,
          N, M, K,
          &alpha,
          w.data_ptr<at::Half>(), CUDA_R_16F, K,
          x.data_ptr<at::Half>(), CUDA_R_16F, K,
          &beta,
          y.data_ptr<at::Half>(), CUDA_R_16F, N,
          ct,
          CUBLAS_GEMM_DEFAULT_TENSOR_OP));

  return y;
}

// A: [B, M, K] 行主序(f16)
// B: [B, N, K] 行主序(f16)
// Y: [B, M, N] 行主序(f16,f32 累加)
void gemm_sb_f16_out(torch::Tensor a, torch::Tensor b, torch::Tensor y) {
  _ck_tensor_cuda_contig(a);
  _ck_tensor_cuda_contig(b);
  _ck_tensor_cuda_contig(y);
  _ck(a.dtype() == torch::kFloat16, "a must be float16");
  _ck(b.dtype() == torch::kFloat16, "b must be float16");
  _ck(y.dtype() == torch::kFloat16, "y must be float16");
  _ck(a.dim() == 3, "a must be 3D");
  _ck(b.dim() == 3, "b must be 3D");
  _ck(y.dim() == 3, "y must be 3D");

  const int64_t B64 = a.size(0);
  const int64_t M64 = a.size(1);
  const int64_t K64 = a.size(2);
  _ck(b.size(0) == B64, "batch mismatch");
  _ck(b.size(2) == K64, "k mismatch");
  const int64_t N64 = b.size(1);
  _ck(y.size(0) == B64 && y.size(1) == M64 && y.size(2) == N64, "y shape mismatch");

  _ck(B64 > 0 && M64 > 0 && N64 > 0 && K64 > 0, "empty batched gemm");
  _ck(B64 <= INT_MAX && M64 <= INT_MAX && N64 <= INT_MAX && K64 <= INT_MAX, "batched gemm too large");

  const int Bc = (int)B64;
  const int M = (int)M64;
  const int N = (int)N64;
  const int K = (int)K64;

  cublasHandle_t handle = _get_handle_tc();
  const cublasComputeType_t ct = _get_ct_fast();

  const float alpha = 1.0f;
  const float beta = 0.0f;

  const long long strideA = (long long)N64 * (long long)K64;
  const long long strideB = (long long)M64 * (long long)K64;
  const long long strideC = (long long)M64 * (long long)N64;

  _ck_cublas(
      cublasGemmStridedBatchedEx(
          handle,
          CUBLAS_OP_T, CUBLAS_OP_N,
          N, M, K,
          &alpha,
          b.data_ptr<at::Half>(), CUDA_R_16F, K, strideA,
          a.data_ptr<at::Half>(), CUDA_R_16F, K, strideB,
          &beta,
          y.data_ptr<at::Half>(), CUDA_R_16F, N, strideC,
          Bc,
          ct,
          CUBLAS_GEMM_DEFAULT_TENSOR_OP));
}

__device__ __forceinline__ float _warp_reduce_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);
  return v;
}

template <int D>
__global__ void _ln_fwd_f16_warp4_kernel(
    const float* __restrict__ x,
    const float* __restrict__ w,
    const float* __restrict__ b,
    __half* __restrict__ y,
    int rows) {
  const int tid = (int)threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  const int warps = (int)blockDim.x >> 5;
  const int row = (int)blockIdx.x * warps + warp;
  if (row >= rows) return;

  const int base = row * D;

  const int off0 = lane << 2;
  float4 v0 = *(const float4*)(x + base + off0);
  float sum = (v0.x + v0.y) + (v0.z + v0.w);
  float sumsq = (v0.x * v0.x + v0.y * v0.y) + (v0.z * v0.z + v0.w * v0.w);

  float4 v1, v2, v3, v4, v5;
  if constexpr (D >= 256) {
    v1 = *(const float4*)(x + base + 128 + off0);
    sum += (v1.x + v1.y) + (v1.z + v1.w);
    sumsq += (v1.x * v1.x + v1.y * v1.y) + (v1.z * v1.z + v1.w * v1.w);
  }
  if constexpr (D >= 384) {
    v2 = *(const float4*)(x + base + 256 + off0);
    sum += (v2.x + v2.y) + (v2.z + v2.w);
    sumsq += (v2.x * v2.x + v2.y * v2.y) + (v2.z * v2.z + v2.w * v2.w);
  }
  if constexpr (D >= 512) {
    v3 = *(const float4*)(x + base + 384 + off0);
    sum += (v3.x + v3.y) + (v3.z + v3.w);
    sumsq += (v3.x * v3.x + v3.y * v3.y) + (v3.z * v3.z + v3.w * v3.w);
  }
  if constexpr (D >= 640) {
    v4 = *(const float4*)(x + base + 512 + off0);
    sum += (v4.x + v4.y) + (v4.z + v4.w);
    sumsq += (v4.x * v4.x + v4.y * v4.y) + (v4.z * v4.z + v4.w * v4.w);
  }
  if constexpr (D >= 768) {
    v5 = *(const float4*)(x + base + 640 + off0);
    sum += (v5.x + v5.y) + (v5.z + v5.w);
    sumsq += (v5.x * v5.x + v5.y * v5.y) + (v5.z * v5.z + v5.w * v5.w);
  }

  const float sum_r = _warp_reduce_sum(sum);
  const float sumsq_r = _warp_reduce_sum(sumsq);

  const float inv_d = 1.0f / (float)D;
  const float sum_t = __shfl_sync(0xffffffff, sum_r, 0);
  const float sumsq_t = __shfl_sync(0xffffffff, sumsq_r, 0);
  const float mean = sum_t * inv_d;
  const float var = sumsq_t * inv_d - mean * mean;
  const float inv = rsqrtf(var + 1.0e-5f);

  float4 w0 = *(const float4*)(w + off0);
  float4 b0 = *(const float4*)(b + off0);

  float4 o0;
  o0.x = (v0.x - mean) * inv * w0.x + b0.x;
  o0.y = (v0.y - mean) * inv * w0.y + b0.y;
  o0.z = (v0.z - mean) * inv * w0.z + b0.z;
  o0.w = (v0.w - mean) * inv * w0.w + b0.w;

  *(__half2*)(y + base + off0) = __floats2half2_rn(o0.x, o0.y);
  *(__half2*)(y + base + off0 + 2) = __floats2half2_rn(o0.z, o0.w);

  if constexpr (D >= 256) {
    float4 w1 = *(const float4*)(w + 128 + off0);
    float4 b1 = *(const float4*)(b + 128 + off0);
    float4 o1;
    o1.x = (v1.x - mean) * inv * w1.x + b1.x;
    o1.y = (v1.y - mean) * inv * w1.y + b1.y;
    o1.z = (v1.z - mean) * inv * w1.z + b1.z;
    o1.w = (v1.w - mean) * inv * w1.w + b1.w;
    *(__half2*)(y + base + 128 + off0) = __floats2half2_rn(o1.x, o1.y);
    *(__half2*)(y + base + 128 + off0 + 2) = __floats2half2_rn(o1.z, o1.w);
  }
  if constexpr (D >= 384) {
    float4 w2 = *(const float4*)(w + 256 + off0);
    float4 b2 = *(const float4*)(b + 256 + off0);
    float4 o2;
    o2.x = (v2.x - mean) * inv * w2.x + b2.x;
    o2.y = (v2.y - mean) * inv * w2.y + b2.y;
    o2.z = (v2.z - mean) * inv * w2.z + b2.z;
    o2.w = (v2.w - mean) * inv * w2.w + b2.w;
    *(__half2*)(y + base + 256 + off0) = __floats2half2_rn(o2.x, o2.y);
    *(__half2*)(y + base + 256 + off0 + 2) = __floats2half2_rn(o2.z, o2.w);
  }
  if constexpr (D >= 512) {
    float4 w3 = *(const float4*)(w + 384 + off0);
    float4 b3 = *(const float4*)(b + 384 + off0);
    float4 o3;
    o3.x = (v3.x - mean) * inv * w3.x + b3.x;
    o3.y = (v3.y - mean) * inv * w3.y + b3.y;
    o3.z = (v3.z - mean) * inv * w3.z + b3.z;
    o3.w = (v3.w - mean) * inv * w3.w + b3.w;
    *(__half2*)(y + base + 384 + off0) = __floats2half2_rn(o3.x, o3.y);
    *(__half2*)(y + base + 384 + off0 + 2) = __floats2half2_rn(o3.z, o3.w);
  }
  if constexpr (D >= 640) {
    float4 w4 = *(const float4*)(w + 512 + off0);
    float4 b4 = *(const float4*)(b + 512 + off0);
    float4 o4;
    o4.x = (v4.x - mean) * inv * w4.x + b4.x;
    o4.y = (v4.y - mean) * inv * w4.y + b4.y;
    o4.z = (v4.z - mean) * inv * w4.z + b4.z;
    o4.w = (v4.w - mean) * inv * w4.w + b4.w;
    *(__half2*)(y + base + 512 + off0) = __floats2half2_rn(o4.x, o4.y);
    *(__half2*)(y + base + 512 + off0 + 2) = __floats2half2_rn(o4.z, o4.w);
  }
  if constexpr (D >= 768) {
    float4 w5 = *(const float4*)(w + 640 + off0);
    float4 b5 = *(const float4*)(b + 640 + off0);
    float4 o5;
    o5.x = (v5.x - mean) * inv * w5.x + b5.x;
    o5.y = (v5.y - mean) * inv * w5.y + b5.y;
    o5.z = (v5.z - mean) * inv * w5.z + b5.z;
    o5.w = (v5.w - mean) * inv * w5.w + b5.w;
    *(__half2*)(y + base + 640 + off0) = __floats2half2_rn(o5.x, o5.y);
    *(__half2*)(y + base + 640 + off0 + 2) = __floats2half2_rn(o5.z, o5.w);
  }
}

__global__ void _ln_fwd_f16_kernel(
    const float* __restrict__ x,
    const float* __restrict__ w,
    const float* __restrict__ b,
    __half* __restrict__ y,
    int rows,
    int d) {
  const int row = (int)blockIdx.x;
  if (row >= rows) return;

  const int tid = (int)threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;

  const int base = row * d;

  float v0 = 0.0f, v1 = 0.0f, v2 = 0.0f, v3 = 0.0f;
  const int i0 = tid;
  const int i1 = tid + 128;
  const int i2 = tid + 256;
  const int i3 = tid + 384;
  const bool p0 = (i0 < d);
  const bool p1 = (i1 < d);
  const bool p2 = (i2 < d);
  const bool p3 = (i3 < d);
  if (p0) v0 = x[base + i0];
  if (p1) v1 = x[base + i1];
  if (p2) v2 = x[base + i2];
  if (p3) v3 = x[base + i3];

  float sum = 0.0f;
  float sumsq = 0.0f;
  if (p0) { sum += v0; sumsq += v0 * v0; }
  if (p1) { sum += v1; sumsq += v1 * v1; }
  if (p2) { sum += v2; sumsq += v2 * v2; }
  if (p3) { sum += v3; sumsq += v3 * v3; }

  for (int k = tid + 512; k < d; k += 128) {
    const float v = x[base + k];
    sum += v;
    sumsq += v * v;
  }

  sum = _warp_reduce_sum(sum);
  sumsq = _warp_reduce_sum(sumsq);

  __shared__ float warp_sum[4];
  __shared__ float warp_sumsq[4];
  __shared__ float mean_s;
  __shared__ float inv_s;

  if (lane == 0) {
    warp_sum[warp] = sum;
    warp_sumsq[warp] = sumsq;
  }
  __syncthreads();

  if (warp == 0) {
    float s0 = (lane < 4) ? warp_sum[lane] : 0.0f;
    float s1 = (lane < 4) ? warp_sumsq[lane] : 0.0f;
    s0 = _warp_reduce_sum(s0);
    s1 = _warp_reduce_sum(s1);
    if (lane == 0) {
      const float inv_d = 1.0f / (float)d;
      const float mean = s0 * inv_d;
      const float var = s1 * inv_d - mean * mean;
      mean_s = mean;
      inv_s = rsqrtf(var + 1.0e-5f);
    }
  }
  __syncthreads();

  const float mean = mean_s;
  const float inv = inv_s;

  if (p0) {
    const float o = (v0 - mean) * inv * w[i0] + b[i0];
    y[base + i0] = __float2half_rn(o);
  }
  if (p1) {
    const float o = (v1 - mean) * inv * w[i1] + b[i1];
    y[base + i1] = __float2half_rn(o);
  }
  if (p2) {
    const float o = (v2 - mean) * inv * w[i2] + b[i2];
    y[base + i2] = __float2half_rn(o);
  }
  if (p3) {
    const float o = (v3 - mean) * inv * w[i3] + b[i3];
    y[base + i3] = __float2half_rn(o);
  }
  for (int k = tid + 512; k < d; k += 128) {
    const float v = x[base + k];
    const float o = (v - mean) * inv * w[k] + b[k];
    y[base + k] = __float2half_rn(o);
  }
}

torch::Tensor ln_fwd_f16(torch::Tensor x, torch::Tensor w, torch::Tensor b) {
  _ck_tensor_cuda_contig(x);
  _ck_tensor_cuda_contig(w);
  _ck_tensor_cuda_contig(b);
  _ck(x.dtype() == torch::kFloat32, "x must be float32");
  _ck(w.dtype() == torch::kFloat32, "w must be float32");
  _ck(b.dtype() == torch::kFloat32, "b must be float32");
  _ck(w.dim() == 1, "w must be 1D");
  _ck(b.dim() == 1, "b must be 1D");

  const int64_t d64 = w.numel();
  _ck(d64 == b.numel(), "w/b mismatch");
  _ck(d64 > 0 && d64 <= INT_MAX, "bad d");
  const int d = (int)d64;
  _ck(x.size(-1) == d64, "x last dim mismatch");

  auto y = torch::empty_like(x, x.options().dtype(torch::kFloat16));
  const int64_t rows64 = x.numel() / d64;
  _ck(rows64 > 0 && rows64 <= INT_MAX, "bad rows");
  const int rows = (int)rows64;

  if (d == 128 || d == 256 || d == 384 || d == 512 || d == 768) {
    const dim3 block(256, 1, 1);
    const int warps = (int)block.x >> 5;
    const dim3 grid((rows + warps - 1) / warps, 1, 1);
    if (d == 128) {
      _ln_fwd_f16_warp4_kernel<128><<<grid, block>>>(
          x.data_ptr<float>(),
          w.data_ptr<float>(),
          b.data_ptr<float>(),
          (__half*)y.data_ptr<at::Half>(),
          rows);
    } else if (d == 256) {
      _ln_fwd_f16_warp4_kernel<256><<<grid, block>>>(
          x.data_ptr<float>(),
          w.data_ptr<float>(),
          b.data_ptr<float>(),
          (__half*)y.data_ptr<at::Half>(),
          rows);
    } else if (d == 384) {
      _ln_fwd_f16_warp4_kernel<384><<<grid, block>>>(
          x.data_ptr<float>(),
          w.data_ptr<float>(),
          b.data_ptr<float>(),
          (__half*)y.data_ptr<at::Half>(),
          rows);
    } else if (d == 512) {
      _ln_fwd_f16_warp4_kernel<512><<<grid, block>>>(
          x.data_ptr<float>(),
          w.data_ptr<float>(),
          b.data_ptr<float>(),
          (__half*)y.data_ptr<at::Half>(),
          rows);
    } else {
      _ln_fwd_f16_warp4_kernel<768><<<grid, block>>>(
          x.data_ptr<float>(),
          w.data_ptr<float>(),
          b.data_ptr<float>(),
          (__half*)y.data_ptr<at::Half>(),
          rows);
    }
  } else {
    const dim3 block(128, 1, 1);
    const dim3 grid(rows, 1, 1);
    _ln_fwd_f16_kernel<<<grid, block>>>(
        x.data_ptr<float>(),
        w.data_ptr<float>(),
        b.data_ptr<float>(),
        (__half*)y.data_ptr<at::Half>(),
        rows,
        d);
  }
  return y;
}

__global__ void _pack5_f32_to_f16_vec2_kernel(
    const float* __restrict__ w0,
    const float* __restrict__ w1,
    const float* __restrict__ w2,
    const float* __restrict__ w3,
    const float* __restrict__ w4,
    __half* __restrict__ out,
    int elems_per_mat) {
  const int g = (int)blockIdx.y;
  const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
  const int i = t << 1;
  if (i >= elems_per_mat) return;
  const float* src = nullptr;
  if (g == 0) src = w0;
  else if (g == 1) src = w1;
  else if (g == 2) src = w2;
  else if (g == 3) src = w3;
  else src = w4;

  const int o = g * elems_per_mat + i;
  if (i + 1 < elems_per_mat) {
    const float2 v = *(const float2*)(src + i);
    *(__half2*)(out + o) = __floats2half2_rn(v.x, v.y);
  } else {
    out[o] = __float2half_rn(src[i]);
  }
}

__global__ void _pack5_and_to_out_f32_to_f16_vec4_kernel(
    const float* __restrict__ w0,
    const float* __restrict__ w1,
    const float* __restrict__ w2,
    const float* __restrict__ w3,
    const float* __restrict__ w4,
    const float* __restrict__ w_to_out,
    __half* __restrict__ out_pack5,
    __half* __restrict__ out_to_out,
    int elems_per_mat) {
  const int g = (int)blockIdx.y;
  const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
  const int i = t << 2;
  if (i >= elems_per_mat) return;

  const float* src = nullptr;
  __half* dst = nullptr;
  if (g == 0) { src = w0; dst = out_pack5 + 0 * elems_per_mat; }
  else if (g == 1) { src = w1; dst = out_pack5 + 1 * elems_per_mat; }
  else if (g == 2) { src = w2; dst = out_pack5 + 2 * elems_per_mat; }
  else if (g == 3) { src = w3; dst = out_pack5 + 3 * elems_per_mat; }
  else if (g == 4) { src = w4; dst = out_pack5 + 4 * elems_per_mat; }
  else { src = w_to_out; dst = out_to_out; }

  if (i + 3 < elems_per_mat) {
    const float4 v = *(const float4*)(src + i);
    *(__half2*)(dst + i) = __floats2half2_rn(v.x, v.y);
    *(__half2*)(dst + i + 2) = __floats2half2_rn(v.z, v.w);
  } else {
    #pragma unroll
    for (int off = 0; off < 4; ++off) {
      const int j = i + off;
      if (j < elems_per_mat) {
        dst[j] = __float2half_rn(src[j]);
      }
    }
  }
}

torch::Tensor pack_w5_f16(torch::Tensor w0,
                          torch::Tensor w1,
                          torch::Tensor w2,
                          torch::Tensor w3,
                          torch::Tensor w4) {
  _ck_tensor_cuda_contig(w0);
  _ck_tensor_cuda_contig(w1);
  _ck_tensor_cuda_contig(w2);
  _ck_tensor_cuda_contig(w3);
  _ck_tensor_cuda_contig(w4);
  _ck(w0.dtype() == torch::kFloat32, "w0 must be float32");
  _ck(w1.dtype() == torch::kFloat32, "w1 must be float32");
  _ck(w2.dtype() == torch::kFloat32, "w2 must be float32");
  _ck(w3.dtype() == torch::kFloat32, "w3 must be float32");
  _ck(w4.dtype() == torch::kFloat32, "w4 must be float32");
  _ck(w0.dim() == 2, "w0 must be 2D");
  _ck(w1.dim() == 2, "w1 must be 2D");
  _ck(w2.dim() == 2, "w2 must be 2D");
  _ck(w3.dim() == 2, "w3 must be 2D");
  _ck(w4.dim() == 2, "w4 must be 2D");

  const int64_t h64 = w0.size(0);
  const int64_t d64 = w0.size(1);
  _ck(h64 == 128, "hidden_dim must be 128");
  _ck(w1.sizes() == w0.sizes(), "w1 shape mismatch");
  _ck(w2.sizes() == w0.sizes(), "w2 shape mismatch");
  _ck(w3.sizes() == w0.sizes(), "w3 shape mismatch");
  _ck(w4.sizes() == w0.sizes(), "w4 shape mismatch");

  const int64_t elems64 = h64 * d64;
  _ck(elems64 > 0 && elems64 <= INT_MAX, "weight too large");
  const int elems = (int)elems64;

  auto out = torch::empty({5 * h64, d64}, w0.options().dtype(torch::kFloat16));

  const int pairs = (elems + 1) >> 1;
  const dim3 block(256, 1, 1);
  const dim3 grid((pairs + (int)block.x - 1) / (int)block.x, 5, 1);
  _pack5_f32_to_f16_vec2_kernel<<<grid, block>>>(
      w0.data_ptr<float>(),
      w1.data_ptr<float>(),
      w2.data_ptr<float>(),
      w3.data_ptr<float>(),
      w4.data_ptr<float>(),
      (__half*)out.data_ptr<at::Half>(),
      elems);
  return out;
}

__global__ void _cast_f32_to_f16_vec2_kernel(
    const float* __restrict__ src,
    __half* __restrict__ dst,
    int n) {
  const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
  const int i = t << 1;
  if (i >= n) return;
  if (i + 1 < n) {
    const float2 v = *(const float2*)(src + i);
    *(__half2*)(dst + i) = __floats2half2_rn(v.x, v.y);
  } else {
    dst[i] = __float2half_rn(src[i]);
  }
}

torch::Tensor cast_f32_to_f16(torch::Tensor x) {
  _ck_tensor_cuda_contig(x);
  _ck(x.dtype() == torch::kFloat32, "x must be float32");
  const int64_t n64 = x.numel();
  _ck(n64 > 0 && n64 <= INT_MAX, "x too large");
  const int n = (int)n64;
  auto y = torch::empty_like(x, x.options().dtype(torch::kFloat16));
  const int pairs = (n + 1) >> 1;
  const dim3 block(256, 1, 1);
  const dim3 grid((pairs + (int)block.x - 1) / (int)block.x, 1, 1);
  _cast_f32_to_f16_vec2_kernel<<<grid, block>>>(
      x.data_ptr<float>(),
      (__half*)y.data_ptr<at::Half>(),
      n);
  return y;
}

__global__ void _ln_gate_transpose_f16_kernel(
    const __half* __restrict__ x,
    const __half* __restrict__ g,
    const float* __restrict__ w,
    const float* __restrict__ b,
    __half* __restrict__ y,
    int inner) {
  const int tx = (int)threadIdx.x;
  const int ty = (int)threadIdx.y;
  const int col0 = (int)blockIdx.x * 32;
  const int col = col0 + tx;

  const int tid = ty * 32 + tx;

  __shared__ float sw[128];
  __shared__ float sb[128];
  if (tid < 128) {
    sw[tid] = w[tid];
    sb[tid] = b[tid];
  }

  __shared__ __half sx[128][33];
  __shared__ __half sg[128][33];
  __shared__ __half so[128][33];

  float psum = 0.0f;
  float psumsq = 0.0f;

  #pragma unroll
  for (int k = 0; k < 32; ++k) {
    const int d = ty + (k << 2);
    __half xh = __float2half_rn(0.0f);
    __half gh = __float2half_rn(0.0f);
    float xv = 0.0f;
    if (col < inner) {
      xh = x[d * inner + col];
      gh = g[d * inner + col];
      xv = __half2float(xh);
    }
    sx[d][tx] = xh;
    sg[d][tx] = gh;
    psum += xv;
    psumsq += xv * xv;
  }

  __shared__ float ssum[4][32];
  __shared__ float ssumsq[4][32];
  ssum[ty][tx] = psum;
  ssumsq[ty][tx] = psumsq;
  __syncthreads();

  __shared__ float smean[32];
  __shared__ float sinv[32];
  if (ty == 0) {
    const float sum = ssum[0][tx] + ssum[1][tx] + ssum[2][tx] + ssum[3][tx];
    const float sumsq = ssumsq[0][tx] + ssumsq[1][tx] + ssumsq[2][tx] + ssumsq[3][tx];
    const float inv_d = 1.0f / 128.0f;
    const float mean = sum * inv_d;
    const float var = sumsq * inv_d - mean * mean;
    smean[tx] = mean;
    sinv[tx] = rsqrtf(var + 1.0e-5f);
  }
  __syncthreads();

  const float mean = smean[tx];
  const float inv = sinv[tx];

  #pragma unroll
  for (int k = 0; k < 32; ++k) {
    const int d = ty + (k << 2);
    const float xv = __half2float(sx[d][tx]);
    const float gv = __half2float(sg[d][tx]);
    const float go = _sigmoid_f(gv);
    const float o = ((xv - mean) * inv * sw[d] + sb[d]) * go;
    so[d][tx] = __float2half_rn(o);
  }
  __syncthreads();

  const int d0 = tid;
  if (d0 < 128) {
    #pragma unroll
    for (int c = 0; c < 32; ++c) {
      const int cc = col0 + c;
      if (cc < inner) {
        y[cc * 128 + d0] = so[d0][c];
      }
    }
  }
}

void ln_gate_transpose_f16_out(torch::Tensor x,
                               torch::Tensor w,
                               torch::Tensor b,
                               torch::Tensor g,
                               torch::Tensor y) {
  _ck_tensor_cuda_contig(x);
  _ck_tensor_cuda_contig(w);
  _ck_tensor_cuda_contig(b);
  _ck_tensor_cuda_contig(g);
  _ck_tensor_cuda_contig(y);
  _ck(x.dtype() == torch::kFloat16, "x must be float16");
  _ck(g.dtype() == torch::kFloat16, "g must be float16");
  _ck(y.dtype() == torch::kFloat16, "y must be float16");
  _ck(w.dtype() == torch::kFloat32, "w must be float32");
  _ck(b.dtype() == torch::kFloat32, "b must be float32");
  _ck(x.dim() == 2, "x must be 2D");
  _ck(g.dim() == 2, "g must be 2D");
  _ck(y.dim() == 2, "y must be 2D");
  _ck(w.dim() == 1, "w must be 1D");
  _ck(b.dim() == 1, "b must be 1D");

  const int64_t h64 = x.size(0);
  const int64_t inner64 = x.size(1);
  _ck(h64 == 128, "hidden_dim must be 128");
  _ck(g.sizes() == x.sizes(), "g shape mismatch");
  _ck(w.numel() == h64 && b.numel() == h64, "w/b mismatch");
  _ck(inner64 > 0 && inner64 <= INT_MAX, "inner too large");
  _ck(y.size(0) == inner64 && y.size(1) == h64, "y shape mismatch");
  const int inner = (int)inner64;

  const dim3 block(32, 4, 1);
  const dim3 grid((inner + 31) / 32, 1, 1);
  _ln_gate_transpose_f16_kernel<<<grid, block>>>(
      (const __half*)x.data_ptr<at::Half>(),
      (const __half*)g.data_ptr<at::Half>(),
      w.data_ptr<float>(),
      b.data_ptr<float>(),
      (__half*)y.data_ptr<at::Half>(),
      inner);
}

torch::Tensor trimul_fwd_f16(torch::Tensor x,
                             torch::Tensor mask,
                             torch::Tensor w_norm,
                             torch::Tensor b_norm,
                             torch::Tensor w_out_norm,
                             torch::Tensor b_out_norm,
                             torch::Tensor w0,
                             torch::Tensor w1,
                             torch::Tensor w2,
                             torch::Tensor w3,
                             torch::Tensor w4,
                             torch::Tensor w_to_out) {
  _ck_tensor_cuda_contig(x);
  _ck_tensor_cuda_contig(mask);
  _ck_tensor_cuda_contig(w_norm);
  _ck_tensor_cuda_contig(b_norm);
  _ck_tensor_cuda_contig(w_out_norm);
  _ck_tensor_cuda_contig(b_out_norm);
  _ck_tensor_cuda_contig(w0);
  _ck_tensor_cuda_contig(w1);
  _ck_tensor_cuda_contig(w2);
  _ck_tensor_cuda_contig(w3);
  _ck_tensor_cuda_contig(w4);
  _ck_tensor_cuda_contig(w_to_out);

  _ck(x.dtype() == torch::kFloat32, "x must be float32");
  _ck(mask.dim() == 3, "mask must be 3D");

  _ck(w_norm.dtype() == torch::kFloat32 && b_norm.dtype() == torch::kFloat32, "norm must be f32");
  _ck(w_out_norm.dtype() == torch::kFloat32 && b_out_norm.dtype() == torch::kFloat32, "out norm must be f32");
  _ck(w0.dtype() == torch::kFloat32, "w0 must be float32");
  _ck(w1.dtype() == torch::kFloat32, "w1 must be float32");
  _ck(w2.dtype() == torch::kFloat32, "w2 must be float32");
  _ck(w3.dtype() == torch::kFloat32, "w3 must be float32");
  _ck(w4.dtype() == torch::kFloat32, "w4 must be float32");
  _ck(w_to_out.dtype() == torch::kFloat32, "w_to_out must be float32");

  _ck(x.dim() == 4, "x must be 4D");
  const int64_t bs = x.size(0);
  const int64_t n = x.size(1);
  _ck(x.size(2) == n, "x must be square");
  const int64_t dim = x.size(3);
  _ck(dim > 0 && dim <= INT_MAX, "bad dim");
  _ck(bs > 0 && bs <= INT_MAX, "bad bs");
  _ck(n > 0 && n <= INT_MAX, "bad n");

  _ck(mask.size(0) == bs && mask.size(1) == n && mask.size(2) == n, "mask shape mismatch");
  _ck(w_norm.numel() == dim && b_norm.numel() == dim, "norm param mismatch");

  const int64_t hidden = 128;
  _ck(w_out_norm.numel() == hidden && b_out_norm.numel() == hidden, "out norm param mismatch");

  _ck(w0.dim() == 2 && w0.size(0) == hidden && w0.size(1) == dim, "w0 shape mismatch");
  _ck(w1.dim() == 2 && w1.size(0) == hidden && w1.size(1) == dim, "w1 shape mismatch");
  _ck(w2.dim() == 2 && w2.size(0) == hidden && w2.size(1) == dim, "w2 shape mismatch");
  _ck(w3.dim() == 2 && w3.size(0) == hidden && w3.size(1) == dim, "w3 shape mismatch");
  _ck(w4.dim() == 2 && w4.size(0) == hidden && w4.size(1) == dim, "w4 shape mismatch");
  _ck(w_to_out.dim() == 2 && w_to_out.size(0) == dim && w_to_out.size(1) == hidden, "w_to_out shape mismatch");

  auto x16 = ln_fwd_f16(x, w_norm, b_norm);
  const int64_t m = bs * n * n;
  auto x2 = x16.view({m, dim});

  const int64_t elems64 = hidden * dim;
  _ck(elems64 > 0 && elems64 <= INT_MAX, "weight too large");
  const int elems = (int)elems64;

  auto w_cat16 = torch::empty({5 * hidden, dim}, x.options().dtype(torch::kFloat16));
  auto w_to_out16 = torch::empty({dim, hidden}, x.options().dtype(torch::kFloat16));

  const int quads = (elems + 3) >> 2;
  const dim3 block_w(256, 1, 1);
  const dim3 grid_w((quads + (int)block_w.x - 1) / (int)block_w.x, 6, 1);
  _pack5_and_to_out_f32_to_f16_vec4_kernel<<<grid_w, block_w>>>(
      w0.data_ptr<float>(),
      w1.data_ptr<float>(),
      w2.data_ptr<float>(),
      w3.data_ptr<float>(),
      w4.data_ptr<float>(),
      w_to_out.data_ptr<float>(),
      (__half*)w_cat16.data_ptr<at::Half>(),
      (__half*)w_to_out16.data_ptr<at::Half>(),
      elems);

  auto proj_all = gemm_f16(w_cat16, x2);
  proj_all = proj_all.view({5, hidden, bs, n, n});

  auto left = proj_all.select(0, 0);
  auto right = proj_all.select(0, 1);
  auto left_gate = proj_all.select(0, 2);
  auto right_gate = proj_all.select(0, 3);
  auto out_gate = proj_all.select(0, 4);

  apply_mask_gate_lr_f16(left, right, left_gate, right_gate, mask);

  const int64_t batch = bs * hidden;
  auto a = left.reshape({batch, n, n});
  auto bb = right.reshape({batch, n, n});
  auto c_buf = left_gate.reshape({batch, n, n});
  gemm_sb_f16_out(a, bb, c_buf);

  auto out_flat = c_buf.view({hidden, m});
  auto gate_flat = out_gate.view({hidden, m});

  auto out2 = right_gate.view({m, hidden});
  ln_gate_transpose_f16_out(out_flat, w_out_norm, b_out_norm, gate_flat, out2);

  auto y16 = gemm_f16(out2, w_to_out16);
  return y16.view({bs, n, n, dim});
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("gemm_f16", &gemm_f16, "矩阵乘(f16 输出)");
  m.def("gemm_sb_f16_out", &gemm_sb_f16_out, "批量矩阵乘(写入输出)");
  m.def("apply_mask_gate_lr_f16", &apply_mask_gate_lr_f16, "mask+gate 融合(不处理 out_gate)");
  m.def("ln_fwd_f16", &ln_fwd_f16, "LayerNorm 前向(f16 输出)");
  m.def("pack_w5_f16", &pack_w5_f16, "5 组权重打包与转换(f16)");
  m.def("cast_f32_to_f16", &cast_f32_to_f16, "f32->f16 转换");
  m.def("ln_gate_transpose_f16_out", &ln_gate_transpose_f16_out, "LN+gate+转置(写入输出)");
  m.def("trimul_fwd_f16", &trimul_fwd_f16, "TriMul Outgoing 前向(f16 输出)");
}
"""

    _EXT = load_inline(
        name="trimul_ext_f16_v11",
        cpp_sources="",
        cuda_sources=cuda_src,
        functions=None,
        with_cuda=True,
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        extra_cflags=["-O3"],
        verbose=False,
    )
    return _EXT


def _t_contig_f32(t: torch.Tensor) -> torch.Tensor:
    if t.dtype != torch.float32:
        raise RuntimeError("weight must be float32")
    if not t.is_cuda:
        raise RuntimeError("weight must be CUDA")
    return t.contiguous() if not t.is_contiguous() else t


@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
    _ = config

    if not x.is_cuda:
        raise RuntimeError("CUDA only")
    if x.dtype != torch.float32:
        raise RuntimeError("x must be float32")
    if not x.is_contiguous():
        x = x.contiguous()

    if not mask.is_cuda:
        raise RuntimeError("mask must be CUDA")
    if not mask.is_contiguous():
        mask = mask.contiguous()

    w_norm = _t_contig_f32(weights["norm.weight"])
    b_norm = _t_contig_f32(weights["norm.bias"])

    w_out_norm = _t_contig_f32(weights["to_out_norm.weight"])
    b_out_norm = _t_contig_f32(weights["to_out_norm.bias"])

    w0 = _t_contig_f32(weights["left_proj.weight"])
    w1 = _t_contig_f32(weights["right_proj.weight"])
    w2 = _t_contig_f32(weights["left_gate.weight"])
    w3 = _t_contig_f32(weights["right_gate.weight"])
    w4 = _t_contig_f32(weights["out_gate.weight"])
    w_to_out = _t_contig_f32(weights["to_out.weight"])

    ext = _get_ext()
    return ext.trimul_fwd_f16(
        x,
        mask,
        w_norm,
        b_norm,
        w_out_norm,
        b_out_norm,
        w0,
        w1,
        w2,
        w3,
        w4,
        w_to_out,
    )
scrolls · 1131 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 409251.

⋯ 2 unchanged lines
from typing import Any, Dict, Tuple
import torch
- import torch.nn.functional as F
+ _EXT = None
- _FUSED_W_5X: torch.Tensor | None = None
- _FUSED_W_5X_META: tuple[int, int, int, int, int, int] | None = None
+ def _get_ext():
+ global _EXT
+ if _EXT is not None:
+ return _EXT
- def _get_fused_w_5x(weights: Dict[str, torch.Tensor], *, device: torch.device) -> torch.Tensor:
- global _FUSED_W_5X, _FUSED_W_5X_META
+ from torch.utils.cpp_extension import load_inline
- w0 = weights["left_proj.weight"]
- w1 = weights["right_proj.weight"]
- w2 = weights["left_gate.weight"]
- w3 = weights["right_gate.weight"]
- w4 = weights["out_gate.weight"]
+ cuda_src = r"""
+ #include <torch/extension.h>
+ #include <ATen/cuda/CUDABlas.h>
+ #include <cublas_v2.h>
+ #include <cuda.h>
+ #include <cuda_fp16.h>
+ #include <cuda_runtime.h>
- meta = (
- int(device.index) if device.type == "cuda" else -1,
- int(w0.data_ptr()),
- int(w1.data_ptr()),
- int(w2.data_ptr()),
- int(w3.data_ptr()),
- int(w4.data_ptr()),
- )
- if _FUSED_W_5X is not None and _FUSED_W_5X_META == meta:
- return _FUSED_W_5X
+ #include <type_traits>
- fused = torch.cat((w0, w1, w2, w3, w4), dim=0).contiguous()
- _FUSED_W_5X = fused
- _FUSED_W_5X_META = meta
- return fused
+ static inline void _ck(bool ok, const char* msg) {
+ if (!ok) { throw std::runtime_error(msg); }
+ }
+ static inline void _ck_tensor_cuda_contig(const torch::Tensor& t) {
+ _ck(t.is_cuda(), "tensor must be CUDA");
+ _ck(t.is_contiguous(), "tensor must be contiguous");
+ }
- def _contract_outgoing_bmm(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
-
- bs, n, _, hidden = left.shape
+ static inline void _ck_cublas(cublasStatus_t st) {
+ if (st != CUBLAS_STATUS_SUCCESS) {
+ throw std::runtime_error("cublas call failed");
+ }
+ }
-
- a = left.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)
- b = right.permute(0, 3, 1, 2).contiguous().reshape(bs * hidden, n, n)
- out = torch.bmm(a, b.transpose(1, 2))
- return out.reshape(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
+ static inline cublasHandle_t _get_handle_tc() {
+ cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
+ static thread_local cublasHandle_t last = nullptr;
+ if (h != last) {
+ _ck_cublas(cublasSetMathMode(h, CUBLAS_TENSOR_OP_MATH));
+ last = h;
+ }
+ return h;
+ }
+ static inline cublasComputeType_t _get_ct_fast() {
+ #if defined(CUBLAS_COMPUTE_32F_FAST_16F)
+ return CUBLAS_COMPUTE_32F_FAST_16F;
+ #else
+ return CUBLAS_COMPUTE_32F;
+ #endif
+ }
- @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
+ // Sigmoid:保持与参考实现一致的 fast-math 路径
+ __device__ __forceinline__ float _sigmoid_f(float x) {
+ return __fdividef(1.0f, 1.0f + __expf(-x));
+ }
- if not x.is_cuda:
- raise RuntimeError("CUDA tensors required")
- if x.dtype != torch.float32:
- x = x.to(dtype=torch.float32)
+ __device__ __forceinline__ float2 _sigmoid_f2(float2 v) {
+ v.x = _sigmoid_f(v.x);
+ v.y = _sigmoid_f(v.y);
+ return v;
+ }
- dim = int(config["dim"])
- hidden_dim = int(config["hidden_dim"])
+ template <typename MaskT>
+ __device__ __forceinline__ float _mask_to_f32(MaskT v) {
+ return static_cast<float>(v);
+ }
-
- torch.backends.cuda.matmul.allow_tf32 = True
- torch.backends.cudnn.allow_tf32 = True
+ template <>
+ __device__ __forceinline__ float _mask_to_f32<__half>(__half v) {
+ return __half2float(v);
+ }
- x = x.contiguous()
- x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
+ template <>
+ __device__ __forceinline__ float _mask_to_f32<bool>(bool v) {
+ return v ? 1.0f : 0.0f;
+ }
- fused_w = _get_fused_w_5x(weights, device=x.device)
- proj = F.linear(x, fused_w, None)
- left, right, left_gate, right_gate, out_gate = proj.split(hidden_dim, dim=-1)
+ template <typename MaskT>
+ __global__ void _mask_gate_lr_fuse_f16_vec4(
+ __half* __restrict__ left,
+ __half* __restrict__ right,
+ const __half* __restrict__ left_gate,
+ const __half* __restrict__ right_gate,
+ const MaskT* __restrict__ mask,
+ int inner) {
+ const int d = (int)blockIdx.y;
+ const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
+ const int col = t << 2;
+ if (col >= inner) return;
+ const int idx = d * inner + col;
- left_gate.sigmoid_()
- right_gate.sigmoid_()
- out_gate.sigmoid_()
+ if (col + 3 < inner) {
+ float m0, m1, m2, m3;
+ if constexpr (std::is_same<MaskT, float>::value) {
+ const float4 mv = *(const float4*)(mask + col);
+ m0 = mv.x; m1 = mv.y; m2 = mv.z; m3 = mv.w;
+ } else {
+ m0 = _mask_to_f32<MaskT>(mask[col]);
+ m1 = _mask_to_f32<MaskT>(mask[col + 1]);
+ m2 = _mask_to_f32<MaskT>(mask[col + 2]);
+ m3 = _mask_to_f32<MaskT>(mask[col + 3]);
+ }
- mask_f = mask.unsqueeze(-1)
- if mask_f.dtype != left.dtype:
- mask_f = mask_f.to(dtype=left.dtype)
+ const __half2 l2_0 = *(const __half2*)(left + idx);
+ const __half2 l2_1 = *(const __half2*)(left + idx + 2);
+ const __half2 r2_0 = *(const __half2*)(right + idx);
+ const __half2 r2_1 = *(const __half2*)(right + idx + 2);
+ const __half2 lg2_0 = *(const __half2*)(left_gate + idx);
+ const __half2 lg2_1 = *(const __half2*)(left_gate + idx + 2);
+ const __half2 rg2_0 = *(const __half2*)(right_gate + idx);
+ const __half2 rg2_1 = *(const __half2*)(right_gate + idx + 2);
- left.mul_(mask_f).mul_(left_gate)
- right.mul_(mask_f).mul_(right_gate)
+ const float2 gl0 = _sigmoid_f2(__half22float2(lg2_0));
+ const float2 gl1 = _sigmoid_f2(__half22float2(lg2_1));
+ const float2 gr0 = _sigmoid_f2(__half22float2(rg2_0));
+ const float2 gr1 = _sigmoid_f2(__half22float2(rg2_1));
- out = _contract_outgoing_bmm(left, right)
+ float2 lv0 = __half22float2(l2_0);
+ float2 lv1 = __half22float2(l2_1);
+ float2 rv0 = __half22float2(r2_0);
+ float2 rv1 = __half22float2(r2_1);
- out = F.layer_norm(
- out,
- (hidden_dim,),
- weights["to_out_norm.weight"],
- weights["to_out_norm.bias"],
- 1e-5,
+ lv0.x = lv0.x * m0 * gl0.x;
+ lv0.y = lv0.y * m1 * gl0.y;
+ lv1.x = lv1.x * m2 * gl1.x;
+ lv1.y = lv1.y * m3 * gl1.y;
+
+ rv0.x = rv0.x * m0 * gr0.x;
+ rv0.y = rv0.y * m1 * gr0.y;
+ rv1.x = rv1.x * m2 * gr1.x;
+ rv1.y = rv1.y * m3 * gr1.y;
+
+ *(__half2*)(left + idx) = __floats2half2_rn(lv0.x, lv0.y);
+ *(__half2*)(left + idx + 2) = __floats2half2_rn(lv1.x, lv1.y);
+ *(__half2*)(right + idx) = __floats2half2_rn(rv0.x, rv0.y);
+ *(__half2*)(right + idx + 2) = __floats2half2_rn(rv1.x, rv1.y);
+ } else {
+ #pragma unroll
+ for (int off = 0; off < 4; ++off) {
+ const int c = col + off;
+ if (c < inner) {
+ const float m = _mask_to_f32<MaskT>(mask[c]);
+ const int id = idx + off;
+ float l = __half2float(left[id]) * m;
+ float r = __half2float(right[id]) * m;
+ const float gl = _sigmoid_f(__half2float(left_gate[id]));
+ const float gr = _sigmoid_f(__half2float(right_gate[id]));
+ l *= gl;
+ r *= gr;
+ left[id] = __float2half_rn(l);
+ right[id] = __float2half_rn(r);
+ }
+ }
+ }
+ }
+
+ void apply_mask_gate_lr_f16(torch::Tensor left,
+ torch::Tensor right,
+ torch::Tensor left_gate,
+ torch::Tensor right_gate,
+ torch::Tensor mask) {
+ _ck_tensor_cuda_contig(left);
+ _ck_tensor_cuda_contig(right);
+ _ck_tensor_cuda_contig(left_gate);
+ _ck_tensor_cuda_contig(right_gate);
+ _ck_tensor_cuda_contig(mask);
+
+ _ck(left.dtype() == torch::kFloat16, "left must be float16");
+ _ck(right.dtype() == torch::kFloat16, "right must be float16");
+ _ck(left_gate.dtype() == torch::kFloat16, "left_gate must be float16");
+ _ck(right_gate.dtype() == torch::kFloat16, "right_gate must be float16");
+ _ck(mask.dim() == 3, "mask must be 3D");
+
+ const int hidden = (int)left.size(0);
+ _ck(hidden == 128, "hidden_dim must be 128");
+ _ck(right.numel() == left.numel(), "lr size mismatch");
+ _ck(left_gate.numel() == left.numel(), "lg size mismatch");
+ _ck(right_gate.numel() == left.numel(), "rg size mismatch");
+
+ const int64_t inner64 = mask.numel();
+ _ck(inner64 > 0 && inner64 <= INT_MAX, "mask too large");
+ const int inner = (int)inner64;
+ _ck((int64_t)hidden * (int64_t)inner == left.numel(), "mask/hidden mismatch");
+
+ const int quads = (inner + 3) >> 2;
+ const dim3 block(256, 1, 1);
+ const dim3 grid((quads + (int)block.x - 1) / (int)block.x, hidden, 1);
+
+ const auto st = mask.scalar_type();
+ if (st == torch::kFloat32) {
+ _mask_gate_lr_fuse_f16_vec4<float><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const float*)mask.data_ptr<float>(),
+ inner);
+ } else if (st == torch::kFloat16) {
+ _mask_gate_lr_fuse_f16_vec4<__half><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const __half*)mask.data_ptr<at::Half>(),
+ inner);
+ } else if (st == torch::kInt64) {
+ _mask_gate_lr_fuse_f16_vec4<int64_t><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const int64_t*)mask.data_ptr<int64_t>(),
+ inner);
+ } else if (st == torch::kInt32) {
+ _mask_gate_lr_fuse_f16_vec4<int32_t><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const int32_t*)mask.data_ptr<int32_t>(),
+ inner);
+ } else if (st == torch::kUInt8) {
+ _mask_gate_lr_fuse_f16_vec4<uint8_t><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const uint8_t*)mask.data_ptr<uint8_t>(),
+ inner);
+ } else if (st == torch::kBool) {
+ _mask_gate_lr_fuse_f16_vec4<bool><<<grid, block>>>(
+ (__half*)left.data_ptr<at::Half>(),
+ (__half*)right.data_ptr<at::Half>(),
+ (const __half*)left_gate.data_ptr<at::Half>(),
+ (const __half*)right_gate.data_ptr<at::Half>(),
+ (const bool*)mask.data_ptr<bool>(),
+ inner);
+ } else {
+ throw std::runtime_error("unsupported mask dtype");
+ }
+ }
+
+ // X: [M, K] 行主序(f16)
+ // W: [N, K] 行主序(f16)
+ // Y: [M, N] 行主序(f16)
+ torch::Tensor gemm_f16(torch::Tensor x, torch::Tensor w) {
+ _ck_tensor_cuda_contig(x);
+ _ck_tensor_cuda_contig(w);
+ _ck(x.dtype() == torch::kFloat16, "x must be float16");
+ _ck(w.dtype() == torch::kFloat16, "w must be float16");
+ _ck(x.dim() == 2, "x must be 2D");
+ _ck(w.dim() == 2, "w must be 2D");
+
+ const int64_t M64 = x.size(0);
+ const int64_t K64 = x.size(1);
+ const int64_t N64 = w.size(0);
+ _ck(w.size(1) == K64, "w shape mismatch");
+ _ck(M64 > 0 && N64 > 0 && K64 > 0, "empty mat");
+ _ck(M64 <= INT_MAX && N64 <= INT_MAX && K64 <= INT_MAX, "mat too large");
+
+ auto y = torch::empty({M64, N64}, x.options());
+
+ const int M = (int)M64;
+ const int N = (int)N64;
+ const int K = (int)K64;
+
+ cublasHandle_t handle = _get_handle_tc();
+ const cublasComputeType_t ct = _get_ct_fast();
+
+ const float alpha = 1.0f;
+ const float beta = 0.0f;
+
+ _ck_cublas(
+ cublasGemmEx(
+ handle,
+ CUBLAS_OP_T, CUBLAS_OP_N,
+ N, M, K,
+ &alpha,
+ w.data_ptr<at::Half>(), CUDA_R_16F, K,
+ x.data_ptr<at::Half>(), CUDA_R_16F, K,
+ &beta,
+ y.data_ptr<at::Half>(), CUDA_R_16F, N,
+ ct,
+ CUBLAS_GEMM_DEFAULT_TENSOR_OP));
+
+ return y;
+ }
+
+ // A: [B, M, K] 行主序(f16)
+ // B: [B, N, K] 行主序(f16)
+ // Y: [B, M, N] 行主序(f16,f32 累加)
+ void gemm_sb_f16_out(torch::Tensor a, torch::Tensor b, torch::Tensor y) {
+ _ck_tensor_cuda_contig(a);
+ _ck_tensor_cuda_contig(b);
+ _ck_tensor_cuda_contig(y);
+ _ck(a.dtype() == torch::kFloat16, "a must be float16");
+ _ck(b.dtype() == torch::kFloat16, "b must be float16");
+ _ck(y.dtype() == torch::kFloat16, "y must be float16");
+ _ck(a.dim() == 3, "a must be 3D");
+ _ck(b.dim() == 3, "b must be 3D");
+ _ck(y.dim() == 3, "y must be 3D");
+
+ const int64_t B64 = a.size(0);
+ const int64_t M64 = a.size(1);
+ const int64_t K64 = a.size(2);
+ _ck(b.size(0) == B64, "batch mismatch");
+ _ck(b.size(2) == K64, "k mismatch");
+ const int64_t N64 = b.size(1);
+ _ck(y.size(0) == B64 && y.size(1) == M64 && y.size(2) == N64, "y shape mismatch");
+
+ _ck(B64 > 0 && M64 > 0 && N64 > 0 && K64 > 0, "empty batched gemm");
+ _ck(B64 <= INT_MAX && M64 <= INT_MAX && N64 <= INT_MAX && K64 <= INT_MAX, "batched gemm too large");
+
+ const int Bc = (int)B64;
+ const int M = (int)M64;
+ const int N = (int)N64;
+ const int K = (int)K64;
+
+ cublasHandle_t handle = _get_handle_tc();
+ const cublasComputeType_t ct = _get_ct_fast();
+
+ const float alpha = 1.0f;
+ const float beta = 0.0f;
+
+ const long long strideA = (long long)N64 * (long long)K64;
+ const long long strideB = (long long)M64 * (long long)K64;
+ const long long strideC = (long long)M64 * (long long)N64;
+
+ _ck_cublas(
+ cublasGemmStridedBatchedEx(
+ handle,
+ CUBLAS_OP_T, CUBLAS_OP_N,
+ N, M, K,
+ &alpha,
+ b.data_ptr<at::Half>(), CUDA_R_16F, K, strideA,
+ a.data_ptr<at::Half>(), CUDA_R_16F, K, strideB,
+ &beta,
+ y.data_ptr<at::Half>(), CUDA_R_16F, N, strideC,
+ Bc,
+ ct,
+ CUBLAS_GEMM_DEFAULT_TENSOR_OP));
+ }
+
+ __device__ __forceinline__ float _warp_reduce_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);
+ return v;
+ }
+
+ template <int D>
+ __global__ void _ln_fwd_f16_warp4_kernel(
+ const float* __restrict__ x,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ __half* __restrict__ y,
+ int rows) {
+ const int tid = (int)threadIdx.x;
+ const int lane = tid & 31;
+ const int warp = tid >> 5;
+ const int warps = (int)blockDim.x >> 5;
+ const int row = (int)blockIdx.x * warps + warp;
+ if (row >= rows) return;
+
+ const int base = row * D;
+
+ const int off0 = lane << 2;
+ float4 v0 = *(const float4*)(x + base + off0);
+ float sum = (v0.x + v0.y) + (v0.z + v0.w);
+ float sumsq = (v0.x * v0.x + v0.y * v0.y) + (v0.z * v0.z + v0.w * v0.w);
+
+ float4 v1, v2, v3, v4, v5;
+ if constexpr (D >= 256) {
+ v1 = *(const float4*)(x + base + 128 + off0);
+ sum += (v1.x + v1.y) + (v1.z + v1.w);
+ sumsq += (v1.x * v1.x + v1.y * v1.y) + (v1.z * v1.z + v1.w * v1.w);
+ }
+ if constexpr (D >= 384) {
+ v2 = *(const float4*)(x + base + 256 + off0);
+ sum += (v2.x + v2.y) + (v2.z + v2.w);
+ sumsq += (v2.x * v2.x + v2.y * v2.y) + (v2.z * v2.z + v2.w * v2.w);
+ }
+ if constexpr (D >= 512) {
+ v3 = *(const float4*)(x + base + 384 + off0);
+ sum += (v3.x + v3.y) + (v3.z + v3.w);
+ sumsq += (v3.x * v3.x + v3.y * v3.y) + (v3.z * v3.z + v3.w * v3.w);
+ }
+ if constexpr (D >= 640) {
+ v4 = *(const float4*)(x + base + 512 + off0);
+ sum += (v4.x + v4.y) + (v4.z + v4.w);
+ sumsq += (v4.x * v4.x + v4.y * v4.y) + (v4.z * v4.z + v4.w * v4.w);
+ }
+ if constexpr (D >= 768) {
+ v5 = *(const float4*)(x + base + 640 + off0);
+ sum += (v5.x + v5.y) + (v5.z + v5.w);
+ sumsq += (v5.x * v5.x + v5.y * v5.y) + (v5.z * v5.z + v5.w * v5.w);
+ }
+
+ const float sum_r = _warp_reduce_sum(sum);
+ const float sumsq_r = _warp_reduce_sum(sumsq);
+
+ const float inv_d = 1.0f / (float)D;
+ const float sum_t = __shfl_sync(0xffffffff, sum_r, 0);
+ const float sumsq_t = __shfl_sync(0xffffffff, sumsq_r, 0);
+ const float mean = sum_t * inv_d;
+ const float var = sumsq_t * inv_d - mean * mean;
+ const float inv = rsqrtf(var + 1.0e-5f);
+
+ float4 w0 = *(const float4*)(w + off0);
+ float4 b0 = *(const float4*)(b + off0);
+
+ float4 o0;
+ o0.x = (v0.x - mean) * inv * w0.x + b0.x;
+ o0.y = (v0.y - mean) * inv * w0.y + b0.y;
+ o0.z = (v0.z - mean) * inv * w0.z + b0.z;
+ o0.w = (v0.w - mean) * inv * w0.w + b0.w;
+
+ *(__half2*)(y + base + off0) = __floats2half2_rn(o0.x, o0.y);
+ *(__half2*)(y + base + off0 + 2) = __floats2half2_rn(o0.z, o0.w);
+
+ if constexpr (D >= 256) {
+ float4 w1 = *(const float4*)(w + 128 + off0);
+ float4 b1 = *(const float4*)(b + 128 + off0);
+ float4 o1;
+ o1.x = (v1.x - mean) * inv * w1.x + b1.x;
+ o1.y = (v1.y - mean) * inv * w1.y + b1.y;
+ o1.z = (v1.z - mean) * inv * w1.z + b1.z;
+ o1.w = (v1.w - mean) * inv * w1.w + b1.w;
+ *(__half2*)(y + base + 128 + off0) = __floats2half2_rn(o1.x, o1.y);
+ *(__half2*)(y + base + 128 + off0 + 2) = __floats2half2_rn(o1.z, o1.w);
+ }
+ if constexpr (D >= 384) {
+ float4 w2 = *(const float4*)(w + 256 + off0);
+ float4 b2 = *(const float4*)(b + 256 + off0);
+ float4 o2;
+ o2.x = (v2.x - mean) * inv * w2.x + b2.x;
+ o2.y = (v2.y - mean) * inv * w2.y + b2.y;
+ o2.z = (v2.z - mean) * inv * w2.z + b2.z;
+ o2.w = (v2.w - mean) * inv * w2.w + b2.w;
+ *(__half2*)(y + base + 256 + off0) = __floats2half2_rn(o2.x, o2.y);
+ *(__half2*)(y + base + 256 + off0 + 2) = __floats2half2_rn(o2.z, o2.w);
+ }
+ if constexpr (D >= 512) {
+ float4 w3 = *(const float4*)(w + 384 + off0);
+ float4 b3 = *(const float4*)(b + 384 + off0);
+ float4 o3;
+ o3.x = (v3.x - mean) * inv * w3.x + b3.x;
+ o3.y = (v3.y - mean) * inv * w3.y + b3.y;
+ o3.z = (v3.z - mean) * inv * w3.z + b3.z;
+ o3.w = (v3.w - mean) * inv * w3.w + b3.w;
+ *(__half2*)(y + base + 384 + off0) = __floats2half2_rn(o3.x, o3.y);
+ *(__half2*)(y + base + 384 + off0 + 2) = __floats2half2_rn(o3.z, o3.w);
+ }
+ if constexpr (D >= 640) {
+ float4 w4 = *(const float4*)(w + 512 + off0);
+ float4 b4 = *(const float4*)(b + 512 + off0);
+ float4 o4;
+ o4.x = (v4.x - mean) * inv * w4.x + b4.x;
+ o4.y = (v4.y - mean) * inv * w4.y + b4.y;
+ o4.z = (v4.z - mean) * inv * w4.z + b4.z;
+ o4.w = (v4.w - mean) * inv * w4.w + b4.w;
+ *(__half2*)(y + base + 512 + off0) = __floats2half2_rn(o4.x, o4.y);
+ *(__half2*)(y + base + 512 + off0 + 2) = __floats2half2_rn(o4.z, o4.w);
+ }
+ if constexpr (D >= 768) {
+ float4 w5 = *(const float4*)(w + 640 + off0);
+ float4 b5 = *(const float4*)(b + 640 + off0);
+ float4 o5;
+ o5.x = (v5.x - mean) * inv * w5.x + b5.x;
+ o5.y = (v5.y - mean) * inv * w5.y + b5.y;
+ o5.z = (v5.z - mean) * inv * w5.z + b5.z;
+ o5.w = (v5.w - mean) * inv * w5.w + b5.w;
+ *(__half2*)(y + base + 640 + off0) = __floats2half2_rn(o5.x, o5.y);
+ *(__half2*)(y + base + 640 + off0 + 2) = __floats2half2_rn(o5.z, o5.w);
+ }
+ }
+
+ __global__ void _ln_fwd_f16_kernel(
+ const float* __restrict__ x,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ __half* __restrict__ y,
+ int rows,
+ int d) {
+ const int row = (int)blockIdx.x;
+ if (row >= rows) return;
+
+ const int tid = (int)threadIdx.x;
+ const int lane = tid & 31;
+ const int warp = tid >> 5;
+
+ const int base = row * d;
+
+ float v0 = 0.0f, v1 = 0.0f, v2 = 0.0f, v3 = 0.0f;
+ const int i0 = tid;
+ const int i1 = tid + 128;
+ const int i2 = tid + 256;
+ const int i3 = tid + 384;
+ const bool p0 = (i0 < d);
+ const bool p1 = (i1 < d);
+ const bool p2 = (i2 < d);
+ const bool p3 = (i3 < d);
+ if (p0) v0 = x[base + i0];
+ if (p1) v1 = x[base + i1];
+ if (p2) v2 = x[base + i2];
+ if (p3) v3 = x[base + i3];
+
+ float sum = 0.0f;
+ float sumsq = 0.0f;
+ if (p0) { sum += v0; sumsq += v0 * v0; }
+ if (p1) { sum += v1; sumsq += v1 * v1; }
+ if (p2) { sum += v2; sumsq += v2 * v2; }
+ if (p3) { sum += v3; sumsq += v3 * v3; }
+
+ for (int k = tid + 512; k < d; k += 128) {
+ const float v = x[base + k];
+ sum += v;
+ sumsq += v * v;
+ }
+
+ sum = _warp_reduce_sum(sum);
+ sumsq = _warp_reduce_sum(sumsq);
+
+ __shared__ float warp_sum[4];
+ __shared__ float warp_sumsq[4];
+ __shared__ float mean_s;
+ __shared__ float inv_s;
+
+ if (lane == 0) {
+ warp_sum[warp] = sum;
+ warp_sumsq[warp] = sumsq;
+ }
+ __syncthreads();
+
+ if (warp == 0) {
+ float s0 = (lane < 4) ? warp_sum[lane] : 0.0f;
+ float s1 = (lane < 4) ? warp_sumsq[lane] : 0.0f;
+ s0 = _warp_reduce_sum(s0);
+ s1 = _warp_reduce_sum(s1);
+ if (lane == 0) {
+ const float inv_d = 1.0f / (float)d;
+ const float mean = s0 * inv_d;
+ const float var = s1 * inv_d - mean * mean;
+ mean_s = mean;
+ inv_s = rsqrtf(var + 1.0e-5f);
+ }
+ }
+ __syncthreads();
+
+ const float mean = mean_s;
+ const float inv = inv_s;
+
+ if (p0) {
+ const float o = (v0 - mean) * inv * w[i0] + b[i0];
+ y[base + i0] = __float2half_rn(o);
+ }
+ if (p1) {
+ const float o = (v1 - mean) * inv * w[i1] + b[i1];
+ y[base + i1] = __float2half_rn(o);
+ }
+ if (p2) {
+ const float o = (v2 - mean) * inv * w[i2] + b[i2];
+ y[base + i2] = __float2half_rn(o);
+ }
+ if (p3) {
+ const float o = (v3 - mean) * inv * w[i3] + b[i3];
+ y[base + i3] = __float2half_rn(o);
+ }
+ for (int k = tid + 512; k < d; k += 128) {
+ const float v = x[base + k];
+ const float o = (v - mean) * inv * w[k] + b[k];
+ y[base + k] = __float2half_rn(o);
+ }
+ }
+
+ torch::Tensor ln_fwd_f16(torch::Tensor x, torch::Tensor w, torch::Tensor b) {
+ _ck_tensor_cuda_contig(x);
+ _ck_tensor_cuda_contig(w);
+ _ck_tensor_cuda_contig(b);
+ _ck(x.dtype() == torch::kFloat32, "x must be float32");
+ _ck(w.dtype() == torch::kFloat32, "w must be float32");
+ _ck(b.dtype() == torch::kFloat32, "b must be float32");
+ _ck(w.dim() == 1, "w must be 1D");
+ _ck(b.dim() == 1, "b must be 1D");
+
+ const int64_t d64 = w.numel();
+ _ck(d64 == b.numel(), "w/b mismatch");
+ _ck(d64 > 0 && d64 <= INT_MAX, "bad d");
+ const int d = (int)d64;
+ _ck(x.size(-1) == d64, "x last dim mismatch");
+
+ auto y = torch::empty_like(x, x.options().dtype(torch::kFloat16));
+ const int64_t rows64 = x.numel() / d64;
+ _ck(rows64 > 0 && rows64 <= INT_MAX, "bad rows");
+ const int rows = (int)rows64;
+
+ if (d == 128 || d == 256 || d == 384 || d == 512 || d == 768) {
+ const dim3 block(256, 1, 1);
+ const int warps = (int)block.x >> 5;
+ const dim3 grid((rows + warps - 1) / warps, 1, 1);
+ if (d == 128) {
+ _ln_fwd_f16_warp4_kernel<128><<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows);
+ } else if (d == 256) {
+ _ln_fwd_f16_warp4_kernel<256><<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows);
+ } else if (d == 384) {
+ _ln_fwd_f16_warp4_kernel<384><<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows);
+ } else if (d == 512) {
+ _ln_fwd_f16_warp4_kernel<512><<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows);
+ } else {
+ _ln_fwd_f16_warp4_kernel<768><<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows);
+ }
+ } else {
+ const dim3 block(128, 1, 1);
+ const dim3 grid(rows, 1, 1);
+ _ln_fwd_f16_kernel<<<grid, block>>>(
+ x.data_ptr<float>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ rows,
+ d);
+ }
+ return y;
+ }
+
+ __global__ void _pack5_f32_to_f16_vec2_kernel(
+ const float* __restrict__ w0,
+ const float* __restrict__ w1,
+ const float* __restrict__ w2,
+ const float* __restrict__ w3,
+ const float* __restrict__ w4,
+ __half* __restrict__ out,
+ int elems_per_mat) {
+ const int g = (int)blockIdx.y;
+ const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
+ const int i = t << 1;
+ if (i >= elems_per_mat) return;
+ const float* src = nullptr;
+ if (g == 0) src = w0;
+ else if (g == 1) src = w1;
+ else if (g == 2) src = w2;
+ else if (g == 3) src = w3;
+ else src = w4;
+
+ const int o = g * elems_per_mat + i;
+ if (i + 1 < elems_per_mat) {
+ const float2 v = *(const float2*)(src + i);
+ *(__half2*)(out + o) = __floats2half2_rn(v.x, v.y);
+ } else {
+ out[o] = __float2half_rn(src[i]);
+ }
+ }
+
+ __global__ void _pack5_and_to_out_f32_to_f16_vec4_kernel(
+ const float* __restrict__ w0,
+ const float* __restrict__ w1,
+ const float* __restrict__ w2,
+ const float* __restrict__ w3,
+ const float* __restrict__ w4,
+ const float* __restrict__ w_to_out,
+ __half* __restrict__ out_pack5,
+ __half* __restrict__ out_to_out,
+ int elems_per_mat) {
+ const int g = (int)blockIdx.y;
+ const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
+ const int i = t << 2;
+ if (i >= elems_per_mat) return;
+
+ const float* src = nullptr;
+ __half* dst = nullptr;
+ if (g == 0) { src = w0; dst = out_pack5 + 0 * elems_per_mat; }
+ else if (g == 1) { src = w1; dst = out_pack5 + 1 * elems_per_mat; }
+ else if (g == 2) { src = w2; dst = out_pack5 + 2 * elems_per_mat; }
+ else if (g == 3) { src = w3; dst = out_pack5 + 3 * elems_per_mat; }
+ else if (g == 4) { src = w4; dst = out_pack5 + 4 * elems_per_mat; }
+ else { src = w_to_out; dst = out_to_out; }
+
+ if (i + 3 < elems_per_mat) {
+ const float4 v = *(const float4*)(src + i);
+ *(__half2*)(dst + i) = __floats2half2_rn(v.x, v.y);
+ *(__half2*)(dst + i + 2) = __floats2half2_rn(v.z, v.w);
+ } else {
+ #pragma unroll
+ for (int off = 0; off < 4; ++off) {
+ const int j = i + off;
+ if (j < elems_per_mat) {
+ dst[j] = __float2half_rn(src[j]);
+ }
+ }
+ }
+ }
+
+ torch::Tensor pack_w5_f16(torch::Tensor w0,
+ torch::Tensor w1,
+ torch::Tensor w2,
+ torch::Tensor w3,
+ torch::Tensor w4) {
+ _ck_tensor_cuda_contig(w0);
+ _ck_tensor_cuda_contig(w1);
+ _ck_tensor_cuda_contig(w2);
+ _ck_tensor_cuda_contig(w3);
+ _ck_tensor_cuda_contig(w4);
+ _ck(w0.dtype() == torch::kFloat32, "w0 must be float32");
+ _ck(w1.dtype() == torch::kFloat32, "w1 must be float32");
+ _ck(w2.dtype() == torch::kFloat32, "w2 must be float32");
+ _ck(w3.dtype() == torch::kFloat32, "w3 must be float32");
+ _ck(w4.dtype() == torch::kFloat32, "w4 must be float32");
+ _ck(w0.dim() == 2, "w0 must be 2D");
+ _ck(w1.dim() == 2, "w1 must be 2D");
+ _ck(w2.dim() == 2, "w2 must be 2D");
+ _ck(w3.dim() == 2, "w3 must be 2D");
+ _ck(w4.dim() == 2, "w4 must be 2D");
+
+ const int64_t h64 = w0.size(0);
+ const int64_t d64 = w0.size(1);
+ _ck(h64 == 128, "hidden_dim must be 128");
+ _ck(w1.sizes() == w0.sizes(), "w1 shape mismatch");
+ _ck(w2.sizes() == w0.sizes(), "w2 shape mismatch");
+ _ck(w3.sizes() == w0.sizes(), "w3 shape mismatch");
+ _ck(w4.sizes() == w0.sizes(), "w4 shape mismatch");
+
+ const int64_t elems64 = h64 * d64;
+ _ck(elems64 > 0 && elems64 <= INT_MAX, "weight too large");
+ const int elems = (int)elems64;
+
+ auto out = torch::empty({5 * h64, d64}, w0.options().dtype(torch::kFloat16));
+
+ const int pairs = (elems + 1) >> 1;
+ const dim3 block(256, 1, 1);
+ const dim3 grid((pairs + (int)block.x - 1) / (int)block.x, 5, 1);
+ _pack5_f32_to_f16_vec2_kernel<<<grid, block>>>(
+ w0.data_ptr<float>(),
+ w1.data_ptr<float>(),
+ w2.data_ptr<float>(),
+ w3.data_ptr<float>(),
+ w4.data_ptr<float>(),
+ (__half*)out.data_ptr<at::Half>(),
+ elems);
+ return out;
+ }
+
+ __global__ void _cast_f32_to_f16_vec2_kernel(
+ const float* __restrict__ src,
+ __half* __restrict__ dst,
+ int n) {
+ const int t = (int)blockIdx.x * (int)blockDim.x + (int)threadIdx.x;
+ const int i = t << 1;
+ if (i >= n) return;
+ if (i + 1 < n) {
+ const float2 v = *(const float2*)(src + i);
+ *(__half2*)(dst + i) = __floats2half2_rn(v.x, v.y);
+ } else {
+ dst[i] = __float2half_rn(src[i]);
+ }
+ }
+
+ torch::Tensor cast_f32_to_f16(torch::Tensor x) {
+ _ck_tensor_cuda_contig(x);
+ _ck(x.dtype() == torch::kFloat32, "x must be float32");
+ const int64_t n64 = x.numel();
+ _ck(n64 > 0 && n64 <= INT_MAX, "x too large");
+ const int n = (int)n64;
+ auto y = torch::empty_like(x, x.options().dtype(torch::kFloat16));
+ const int pairs = (n + 1) >> 1;
+ const dim3 block(256, 1, 1);
+ const dim3 grid((pairs + (int)block.x - 1) / (int)block.x, 1, 1);
+ _cast_f32_to_f16_vec2_kernel<<<grid, block>>>(
+ x.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ n);
+ return y;
+ }
+
+ __global__ void _ln_gate_transpose_f16_kernel(
+ const __half* __restrict__ x,
+ const __half* __restrict__ g,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ __half* __restrict__ y,
+ int inner) {
+ const int tx = (int)threadIdx.x;
+ const int ty = (int)threadIdx.y;
+ const int col0 = (int)blockIdx.x * 32;
+ const int col = col0 + tx;
+
+ const int tid = ty * 32 + tx;
+
+ __shared__ float sw[128];
+ __shared__ float sb[128];
+ if (tid < 128) {
+ sw[tid] = w[tid];
+ sb[tid] = b[tid];
+ }
+
+ __shared__ __half sx[128][33];
+ __shared__ __half sg[128][33];
+ __shared__ __half so[128][33];
+
+ float psum = 0.0f;
+ float psumsq = 0.0f;
+
+ #pragma unroll
+ for (int k = 0; k < 32; ++k) {
+ const int d = ty + (k << 2);
+ __half xh = __float2half_rn(0.0f);
+ __half gh = __float2half_rn(0.0f);
+ float xv = 0.0f;
+ if (col < inner) {
+ xh = x[d * inner + col];
+ gh = g[d * inner + col];
+ xv = __half2float(xh);
+ }
+ sx[d][tx] = xh;
+ sg[d][tx] = gh;
+ psum += xv;
+ psumsq += xv * xv;
+ }
+
+ __shared__ float ssum[4][32];
+ __shared__ float ssumsq[4][32];
+ ssum[ty][tx] = psum;
+ ssumsq[ty][tx] = psumsq;
+ __syncthreads();
+
+ __shared__ float smean[32];
+ __shared__ float sinv[32];
+ if (ty == 0) {
+ const float sum = ssum[0][tx] + ssum[1][tx] + ssum[2][tx] + ssum[3][tx];
+ const float sumsq = ssumsq[0][tx] + ssumsq[1][tx] + ssumsq[2][tx] + ssumsq[3][tx];
+ const float inv_d = 1.0f / 128.0f;
+ const float mean = sum * inv_d;
+ const float var = sumsq * inv_d - mean * mean;
+ smean[tx] = mean;
+ sinv[tx] = rsqrtf(var + 1.0e-5f);
+ }
+ __syncthreads();
+
+ const float mean = smean[tx];
+ const float inv = sinv[tx];
+
+ #pragma unroll
+ for (int k = 0; k < 32; ++k) {
+ const int d = ty + (k << 2);
+ const float xv = __half2float(sx[d][tx]);
+ const float gv = __half2float(sg[d][tx]);
+ const float go = _sigmoid_f(gv);
+ const float o = ((xv - mean) * inv * sw[d] + sb[d]) * go;
+ so[d][tx] = __float2half_rn(o);
+ }
+ __syncthreads();
+
+ const int d0 = tid;
+ if (d0 < 128) {
+ #pragma unroll
+ for (int c = 0; c < 32; ++c) {
+ const int cc = col0 + c;
+ if (cc < inner) {
+ y[cc * 128 + d0] = so[d0][c];
+ }
+ }
+ }
+ }
+
+ void ln_gate_transpose_f16_out(torch::Tensor x,
+ torch::Tensor w,
+ torch::Tensor b,
+ torch::Tensor g,
+ torch::Tensor y) {
+ _ck_tensor_cuda_contig(x);
+ _ck_tensor_cuda_contig(w);
+ _ck_tensor_cuda_contig(b);
+ _ck_tensor_cuda_contig(g);
+ _ck_tensor_cuda_contig(y);
+ _ck(x.dtype() == torch::kFloat16, "x must be float16");
+ _ck(g.dtype() == torch::kFloat16, "g must be float16");
+ _ck(y.dtype() == torch::kFloat16, "y must be float16");
+ _ck(w.dtype() == torch::kFloat32, "w must be float32");
+ _ck(b.dtype() == torch::kFloat32, "b must be float32");
+ _ck(x.dim() == 2, "x must be 2D");
+ _ck(g.dim() == 2, "g must be 2D");
+ _ck(y.dim() == 2, "y must be 2D");
+ _ck(w.dim() == 1, "w must be 1D");
+ _ck(b.dim() == 1, "b must be 1D");
+
+ const int64_t h64 = x.size(0);
+ const int64_t inner64 = x.size(1);
+ _ck(h64 == 128, "hidden_dim must be 128");
+ _ck(g.sizes() == x.sizes(), "g shape mismatch");
+ _ck(w.numel() == h64 && b.numel() == h64, "w/b mismatch");
+ _ck(inner64 > 0 && inner64 <= INT_MAX, "inner too large");
+ _ck(y.size(0) == inner64 && y.size(1) == h64, "y shape mismatch");
+ const int inner = (int)inner64;
+
+ const dim3 block(32, 4, 1);
+ const dim3 grid((inner + 31) / 32, 1, 1);
+ _ln_gate_transpose_f16_kernel<<<grid, block>>>(
+ (const __half*)x.data_ptr<at::Half>(),
+ (const __half*)g.data_ptr<at::Half>(),
+ w.data_ptr<float>(),
+ b.data_ptr<float>(),
+ (__half*)y.data_ptr<at::Half>(),
+ inner);
+ }
+
+ torch::Tensor trimul_fwd_f16(torch::Tensor x,
+ torch::Tensor mask,
+ torch::Tensor w_norm,
+ torch::Tensor b_norm,
+ torch::Tensor w_out_norm,
+ torch::Tensor b_out_norm,
+ torch::Tensor w0,
+ torch::Tensor w1,
+ torch::Tensor w2,
+ torch::Tensor w3,
+ torch::Tensor w4,
+ torch::Tensor w_to_out) {
+ _ck_tensor_cuda_contig(x);
+ _ck_tensor_cuda_contig(mask);
+ _ck_tensor_cuda_contig(w_norm);
+ _ck_tensor_cuda_contig(b_norm);
+ _ck_tensor_cuda_contig(w_out_norm);
+ _ck_tensor_cuda_contig(b_out_norm);
+ _ck_tensor_cuda_contig(w0);
+ _ck_tensor_cuda_contig(w1);
+ _ck_tensor_cuda_contig(w2);
+ _ck_tensor_cuda_contig(w3);
+ _ck_tensor_cuda_contig(w4);
+ _ck_tensor_cuda_contig(w_to_out);
+
+ _ck(x.dtype() == torch::kFloat32, "x must be float32");
+ _ck(mask.dim() == 3, "mask must be 3D");
+
+ _ck(w_norm.dtype() == torch::kFloat32 && b_norm.dtype() == torch::kFloat32, "norm must be f32");
+ _ck(w_out_norm.dtype() == torch::kFloat32 && b_out_norm.dtype() == torch::kFloat32, "out norm must be f32");
+ _ck(w0.dtype() == torch::kFloat32, "w0 must be float32");
+ _ck(w1.dtype() == torch::kFloat32, "w1 must be float32");
+ _ck(w2.dtype() == torch::kFloat32, "w2 must be float32");
+ _ck(w3.dtype() == torch::kFloat32, "w3 must be float32");
+ _ck(w4.dtype() == torch::kFloat32, "w4 must be float32");
+ _ck(w_to_out.dtype() == torch::kFloat32, "w_to_out must be float32");
+
+ _ck(x.dim() == 4, "x must be 4D");
+ const int64_t bs = x.size(0);
+ const int64_t n = x.size(1);
+ _ck(x.size(2) == n, "x must be square");
+ const int64_t dim = x.size(3);
+ _ck(dim > 0 && dim <= INT_MAX, "bad dim");
+ _ck(bs > 0 && bs <= INT_MAX, "bad bs");
+ _ck(n > 0 && n <= INT_MAX, "bad n");
+
+ _ck(mask.size(0) == bs && mask.size(1) == n && mask.size(2) == n, "mask shape mismatch");
+ _ck(w_norm.numel() == dim && b_norm.numel() == dim, "norm param mismatch");
+
+ const int64_t hidden = 128;
+ _ck(w_out_norm.numel() == hidden && b_out_norm.numel() == hidden, "out norm param mismatch");
+
+ _ck(w0.dim() == 2 && w0.size(0) == hidden && w0.size(1) == dim, "w0 shape mismatch");
+ _ck(w1.dim() == 2 && w1.size(0) == hidden && w1.size(1) == dim, "w1 shape mismatch");
+ _ck(w2.dim() == 2 && w2.size(0) == hidden && w2.size(1) == dim, "w2 shape mismatch");
+ _ck(w3.dim() == 2 && w3.size(0) == hidden && w3.size(1) == dim, "w3 shape mismatch");
+ _ck(w4.dim() == 2 && w4.size(0) == hidden && w4.size(1) == dim, "w4 shape mismatch");
+ _ck(w_to_out.dim() == 2 && w_to_out.size(0) == dim && w_to_out.size(1) == hidden, "w_to_out shape mismatch");
+
+ auto x16 = ln_fwd_f16(x, w_norm, b_norm);
+ const int64_t m = bs * n * n;
+ auto x2 = x16.view({m, dim});
+
+ const int64_t elems64 = hidden * dim;
+ _ck(elems64 > 0 && elems64 <= INT_MAX, "weight too large");
+ const int elems = (int)elems64;
+
+ auto w_cat16 = torch::empty({5 * hidden, dim}, x.options().dtype(torch::kFloat16));
+ auto w_to_out16 = torch::empty({dim, hidden}, x.options().dtype(torch::kFloat16));
+
+ const int quads = (elems + 3) >> 2;
+ const dim3 block_w(256, 1, 1);
+ const dim3 grid_w((quads + (int)block_w.x - 1) / (int)block_w.x, 6, 1);
+ _pack5_and_to_out_f32_to_f16_vec4_kernel<<<grid_w, block_w>>>(
+ w0.data_ptr<float>(),
+ w1.data_ptr<float>(),
+ w2.data_ptr<float>(),
+ w3.data_ptr<float>(),
+ w4.data_ptr<float>(),
+ w_to_out.data_ptr<float>(),
+ (__half*)w_cat16.data_ptr<at::Half>(),
+ (__half*)w_to_out16.data_ptr<at::Half>(),
+ elems);
+
+ auto proj_all = gemm_f16(w_cat16, x2);
+ proj_all = proj_all.view({5, hidden, bs, n, n});
+
+ auto left = proj_all.select(0, 0);
+ auto right = proj_all.select(0, 1);
+ auto left_gate = proj_all.select(0, 2);
+ auto right_gate = proj_all.select(0, 3);
+ auto out_gate = proj_all.select(0, 4);
+
+ apply_mask_gate_lr_f16(left, right, left_gate, right_gate, mask);
+
+ const int64_t batch = bs * hidden;
+ auto a = left.reshape({batch, n, n});
+ auto bb = right.reshape({batch, n, n});
+ auto c_buf = left_gate.reshape({batch, n, n});
+ gemm_sb_f16_out(a, bb, c_buf);
+
+ auto out_flat = c_buf.view({hidden, m});
+ auto gate_flat = out_gate.view({hidden, m});
+
+ auto out2 = right_gate.view({m, hidden});
+ ln_gate_transpose_f16_out(out_flat, w_out_norm, b_out_norm, gate_flat, out2);
+
+ auto y16 = gemm_f16(out2, w_to_out16);
+ return y16.view({bs, n, n, dim});
+ }
+
+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
+ m.def("gemm_f16", &gemm_f16, "矩阵乘(f16 输出)");
+ m.def("gemm_sb_f16_out", &gemm_sb_f16_out, "批量矩阵乘(写入输出)");
+ m.def("apply_mask_gate_lr_f16", &apply_mask_gate_lr_f16, "mask+gate 融合(不处理 out_gate)");
+ m.def("ln_fwd_f16", &ln_fwd_f16, "LayerNorm 前向(f16 输出)");
+ m.def("pack_w5_f16", &pack_w5_f16, "5 组权重打包与转换(f16)");
+ m.def("cast_f32_to_f16", &cast_f32_to_f16, "f32->f16 转换");
+ m.def("ln_gate_transpose_f16_out", &ln_gate_transpose_f16_out, "LN+gate+转置(写入输出)");
+ m.def("trimul_fwd_f16", &trimul_fwd_f16, "TriMul Outgoing 前向(f16 输出)");
+ }
+ """
+
+ _EXT = load_inline(
+ name="trimul_ext_f16_v11",
+ cpp_sources="",
+ cuda_sources=cuda_src,
+ functions=None,
+ with_cuda=True,
+ extra_cuda_cflags=["-O3", "--use_fast_math"],
+ extra_cflags=["-O3"],
+ verbose=False,
)
- out.mul_(out_gate)
- out = F.linear(out, weights["to_out.weight"], None)
- return out
+ return _EXT
- __all__ = ["custom_kernel"]
+ def _t_contig_f32(t: torch.Tensor) -> torch.Tensor:
+ if t.dtype != torch.float32:
+ raise RuntimeError("weight must be float32")
+ if not t.is_cuda:
+ raise RuntimeError("weight must be CUDA")
+ return t.contiguous() if not t.is_contiguous() else t
+
+ @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
+ _ = config
+
+ if not x.is_cuda:
+ raise RuntimeError("CUDA only")
+ if x.dtype != torch.float32:
+ raise RuntimeError("x must be float32")
+ if not x.is_contiguous():
+ x = x.contiguous()
+
+ if not mask.is_cuda:
+ raise RuntimeError("mask must be CUDA")
+ if not mask.is_contiguous():
+ mask = mask.contiguous()
+
+ w_norm = _t_contig_f32(weights["norm.weight"])
+ b_norm = _t_contig_f32(weights["norm.bias"])
+
+ w_out_norm = _t_contig_f32(weights["to_out_norm.weight"])
+ b_out_norm = _t_contig_f32(weights["to_out_norm.bias"])
+
+ w0 = _t_contig_f32(weights["left_proj.weight"])
+ w1 = _t_contig_f32(weights["right_proj.weight"])
+ w2 = _t_contig_f32(weights["left_gate.weight"])
+ w3 = _t_contig_f32(weights["right_gate.weight"])
+ w4 = _t_contig_f32(weights["out_gate.weight"])
+ w_to_out = _t_contig_f32(weights["to_out.weight"])
+
+ ext = _get_ext()
+ return ext.trimul_fwd_f16(
+ x,
+ mask,
+ w_norm,
+ b_norm,
+ w_out_norm,
+ b_out_norm,
+ w0,
+ w1,
+ w2,
+ w3,
+ w4,
+ w_to_out,
+ )
scrolls · 1197 diff lines total

Best evidence level for this revision: reported

JSON