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
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 linesfrom typing import Any, Dict, Tupleimport 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