use std::ffi::c_void;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use cudarc::driver::sys::CUdeviceptr;
use cudarc::driver::{LaunchConfig, PushKernelArg};
use onnx_runtime_ep_api::{
DeviceGraphResource, DevicePtr, DevicePtrMut, EpError, Kernel, KernelFactory, Result,
TensorMut, TensorView,
};
use onnx_runtime_ir::{DataType, Node};
use crate::error::{driver_err, not_implemented};
use crate::runtime::{CudaRuntime, GraphDeviceAllocation, cuptr, raw_ptr};
use super::softmax::resolve_axis;
const LAYERNORM_SRC: &str = r#"
#include <cuda_fp16.h>
#include <cuda_bf16.h>
__device__ __forceinline__ float load_layernorm_param(
const void* values, const int is_half, const int index) {
return is_half
? __half2float(((const __half*)values)[index])
: ((const float*)values)[index];
}
extern "C" __global__ void layernorm_f32(
const float* x,
const float* scale,
const float* bias, // null when absent
float* y,
float* mean_out, // null when not requested
float* invstd_out, // null when not requested
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
// Pass 1: mean.
float s = 0.0f;
for (int j = tid; j < norm_size; j += nt) s += x[base + j];
red[tid] = s;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float mean = red[0] / (float)norm_size;
__syncthreads();
// Pass 2: population variance.
float v = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float d = x[base + j] - mean;
v += d * d;
}
red[tid] = v;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float var = red[0] / (float)norm_size;
const float inv_std = 1.0f / sqrtf(var + epsilon);
if (tid == 0) {
if (mean_out) mean_out[g] = mean;
if (invstd_out) invstd_out[g] = inv_std;
}
// Pass 3: normalize + affine.
for (int j = tid; j < norm_size; j += nt) {
const float xhat = (x[base + j] - mean) * inv_std;
float o = xhat * scale[j];
if (has_bias) o += bias[j];
y[base + j] = o;
}
}
extern "C" __global__ void layernorm_f16(
const __half* x,
const void* scale,
const void* bias,
__half* y,
float* mean_out,
float* invstd_out,
const int num_groups,
const int norm_size,
const int scale_is_half,
const int bias_is_half,
const int has_bias,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
float s = 0.0f;
for (int j = tid; j < norm_size; j += nt)
s += __half2float(x[base + j]);
red[tid] = s;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float mean = red[0] / (float)norm_size;
__syncthreads();
float v = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float d = __half2float(x[base + j]) - mean;
v += d * d;
}
red[tid] = v;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float inv_std =
1.0f / sqrtf(red[0] / (float)norm_size + epsilon);
if (tid == 0) {
if (mean_out) mean_out[g] = mean;
if (invstd_out) invstd_out[g] = inv_std;
}
for (int j = tid; j < norm_size; j += nt) {
const float xhat = (__half2float(x[base + j]) - mean) * inv_std;
float o = xhat * load_layernorm_param(scale, scale_is_half, j);
if (has_bias)
o += load_layernorm_param(bias, bias_is_half, j);
y[base + j] = __float2half_rn(o);
}
}
__device__ __forceinline__ float load_layernorm_bf16_param(
const void* values, const int is_bf16, const int index) {
return is_bf16
? __bfloat162float(((const __nv_bfloat16*)values)[index])
: ((const float*)values)[index];
}
extern "C" __global__ void layernorm_bf16(
const __nv_bfloat16* x,
const void* scale,
const void* bias,
__nv_bfloat16* y,
float* mean_out,
float* invstd_out,
const int num_groups,
const int norm_size,
const int scale_is_bf16,
const int bias_is_bf16,
const int has_bias,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
float s = 0.0f;
for (int j = tid; j < norm_size; j += nt)
s += __bfloat162float(x[base + j]);
red[tid] = s;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float mean = red[0] / (float)norm_size;
__syncthreads();
float v = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float d = __bfloat162float(x[base + j]) - mean;
v += d * d;
}
red[tid] = v;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float inv_std =
1.0f / sqrtf(red[0] / (float)norm_size + epsilon);
if (tid == 0) {
if (mean_out) mean_out[g] = mean;
if (invstd_out) invstd_out[g] = inv_std;
}
for (int j = tid; j < norm_size; j += nt) {
const float xhat = (__bfloat162float(x[base + j]) - mean) * inv_std;
float o = xhat * load_layernorm_bf16_param(scale, scale_is_bf16, j);
if (has_bias)
o += load_layernorm_bf16_param(bias, bias_is_bf16, j);
y[base + j] = __float2bfloat16_rn(o);
}
}
"#;
const RMSNORM_SRC: &str = r#"
#include <cuda_fp16.h>
#include <cuda_bf16.h>
__device__ __forceinline__ float load_rmsnorm_scale(
const void* values, const int is_half, const int index) {
return is_half
? __half2float(((const __half*)values)[index])
: ((const float*)values)[index];
}
extern "C" __global__ void rmsnorm_f32(
const float* x,
const float* scale,
float* y,
float* invstd_out, // null when not requested
const int num_groups,
const int norm_size,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
// Keep the correctness path in the CPU kernel's left-to-right f32 order.
// Accuracy-level-4 MatMulNBits quantizes activations, so even a one-ulp
// normalization difference can cross an int8 rounding boundary in decode.
if (tid == 0) {
float ss = 0.0f;
for (int j = 0; j < norm_size; ++j) {
const float xv = x[base + j];
// Match the CPU kernel's separate multiply then add. NVRTC otherwise
// contracts this expression to FMA and changes recurrent decode state.
ss = __fadd_rn(ss, __fmul_rn(xv, xv));
}
red[0] = ss;
}
__syncthreads();
const float ms = red[0] / (float)norm_size;
const float inv_std = 1.0f / sqrtf(ms + epsilon);
if (tid == 0 && invstd_out) invstd_out[g] = inv_std;
for (int j = tid; j < norm_size; j += nt)
y[base + j] = x[base + j] * inv_std * scale[j];
}
extern "C" __global__ void rmsnorm_f16(
const __half* x,
const void* scale,
__half* y,
float* invstd_out,
const int num_groups,
const int norm_size,
const int scale_is_half,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
float ss = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float xv = __half2float(x[base + j]);
ss += xv * xv;
}
red[tid] = ss;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float inv_std =
1.0f / sqrtf(red[0] / (float)norm_size + epsilon);
if (tid == 0 && invstd_out) invstd_out[g] = inv_std;
for (int j = tid; j < norm_size; j += nt) {
const float o = __half2float(x[base + j]) * inv_std
* load_rmsnorm_scale(scale, scale_is_half, j);
y[base + j] = __float2half_rn(o);
}
}
__device__ __forceinline__ float load_rmsnorm_bf16_scale(
const void* values, const int is_bf16, const int index) {
return is_bf16
? __bfloat162float(((const __nv_bfloat16*)values)[index])
: ((const float*)values)[index];
}
extern "C" __global__ void rmsnorm_bf16(
const __nv_bfloat16* x,
const void* scale,
__nv_bfloat16* y,
float* invstd_out,
const int num_groups,
const int norm_size,
const int scale_is_bf16,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
// Parallel f32 tree reduction of the mean-square. The bf16 activations
// upcast to f32 losslessly, then accumulate in f32 (matmul-free: just x*x),
// so precision is full-f32 throughout; only the *summation order* differs
// from the serial `rmsnorm_f32` reference (tree vs strict left-to-right).
// A pairwise tree is at least as accurate as sequential accumulation (lower
// error growth, O(log n) vs O(n)), and a per-element f64 oracle confirms the
// tree result is within a couple of ulp of the f64 ground-truth RMS. Because
// decode feeds an accuracy-level-4 int4 MatMulNBits (quantized activations),
// a sub-ulp normalization difference can still flip a downstream int8
// rounding boundary, so greedy token ids stay byte-exact for the first ~38
// steps then exhibit expected sub-ulp greedy sensitivity. The strict
// CPU-order serial path remains available via
// ONNX_GENAI_CUDA_DISABLE_NORM_CAST_FOLD=1 (routes back to rmsnorm_f32).
float ss = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float xv = __bfloat162float(x[base + j]);
ss += xv * xv;
}
red[tid] = ss;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float inv_std =
1.0f / sqrtf(red[0] / (float)norm_size + epsilon);
if (tid == 0 && invstd_out) invstd_out[g] = inv_std;
for (int j = tid; j < norm_size; j += nt) {
const float o = __bfloat162float(x[base + j]) * inv_std
* load_rmsnorm_bf16_scale(scale, scale_is_bf16, j);
y[base + j] = __float2bfloat16_rn(o);
}
}
"#;
const SKIP_RMSNORM_SRC: &str = r#"
#include <cuda_fp16.h>
#include <cuda_bf16.h>
__device__ __forceinline__ float load_skip_val(
const void* values, const int is_half, const int index) {
return is_half
? __half2float(((const __half*)values)[index])
: ((const float*)values)[index];
}
// Warp-shuffle tail of the launch-invariant `red[tid]` power-of-two tree.
//
// Precondition: `red[tid]` already holds this thread's fp32 partial and the
// caller has executed a `__syncthreads()`. `nt` is a power of two (the block
// size chosen by `reduction_launch_config`).
//
// The inter-warp offsets (>= 32) combine threads living in DIFFERENT warps, so
// they must stay in shared memory with a `__syncthreads()` barrier. Once the
// tree has collapsed to `offset == 32`, `red[0..31]` hold the 32 per-warp
// partial sums. The remaining offsets 16,8,4,2,1 pair lane `tid` with lane
// `tid+offset` ENTIRELY inside warp 0 — exactly the pairing and order that
// `__shfl_down_sync(0xffffffff, v, offset)` produces — so warp 0 finishes the
// reduction in registers, dropping 5 `__syncthreads` + 5 shared read/write
// rounds. The accumulation order (and `__fadd_rn` round-to-nearest fp32 add) is
// BIT-IDENTICAL to the full shared-memory tree; see the `warp_reduce` unit tests
// which assert byte-for-byte equality of the `y` output flag-off vs flag-on.
//
// After this returns, `red[0]` holds the total for every thread to read.
__device__ __forceinline__ float skip_rmsnorm_warp_tail(
float* red, const int tid, const int nt) {
for (int offset = nt >> 1; offset >= 32; offset >>= 1) {
if (tid < offset) red[tid] = __fadd_rn(red[tid], red[tid + offset]);
__syncthreads();
}
if (tid < 32) {
float v = (tid < nt) ? red[tid] : 0.0f;
#pragma unroll
for (int o = 16; o > 0; o >>= 1) {
v = __fadd_rn(v, __shfl_down_sync(0xffffffffu, v, o));
}
if (tid == 0) red[0] = v;
}
__syncthreads();
return red[0];
}
template <bool DenseSkip, bool WarpTail>
__device__ __forceinline__ void skip_rmsnorm_f32_tpl(
const float* input,
const float* skip,
const float* gamma,
const float* bias, // null when absent
float* y,
float* sum_out, // null when not requested
float* mean_out, // null when not requested (always zero)
float* invstd_out, // null when not requested
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
const unsigned long long* shape = metadata;
const unsigned long long* skip_strides = metadata + rank;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
float sum_squares = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
unsigned long long skip_index = (unsigned long long)base + j;
if (!DenseSkip) {
unsigned long long linear = skip_index;
skip_index = 0;
for (int d = rank - 1; d >= 0; --d) {
const unsigned long long coord = linear % shape[d];
linear /= shape[d];
skip_index += coord * skip_strides[d];
}
}
float sv = input[base + j] + skip[skip_index];
if (has_bias) sv += bias[j];
y[base + j] = sv;
if (sum_out) sum_out[base + j] = sv;
sum_squares = __fadd_rn(sum_squares, __fmul_rn(sv, sv));
}
// Fixed block tree: every thread owns the same strided subsequence, then
// power-of-two offsets combine partials in a launch-invariant order. With
// WarpTail, the intra-warp offsets (<= 16) finish via __shfl_down_sync in
// registers — bit-identical pairing/order to the full shared tree.
red[tid] = sum_squares;
__syncthreads();
float reduced;
if (WarpTail) {
reduced = skip_rmsnorm_warp_tail(red, tid, nt);
} else {
for (int offset = nt >> 1; offset > 0; offset >>= 1) {
if (tid < offset) {
red[tid] = __fadd_rn(red[tid], red[tid + offset]);
}
__syncthreads();
}
reduced = red[0];
}
const float inv_std = 1.0f / sqrtf(reduced / (float)norm_size + epsilon);
if (tid == 0) {
if (mean_out) mean_out[g] = 0.0f;
if (invstd_out) invstd_out[g] = inv_std;
}
__syncthreads();
for (int j = tid; j < norm_size; j += nt)
y[base + j] = (y[base + j] * inv_std) * gamma[j];
}
extern "C" __global__ void skip_rmsnorm_f32_dense(
const float* input,
const float* skip,
const float* gamma,
const float* bias,
float* y,
float* sum_out,
float* mean_out,
float* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
skip_rmsnorm_f32_tpl<true, false>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, epsilon);
}
extern "C" __global__ void skip_rmsnorm_f32_dense_warp(
const float* input,
const float* skip,
const float* gamma,
const float* bias,
float* y,
float* sum_out,
float* mean_out,
float* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
skip_rmsnorm_f32_tpl<true, true>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, epsilon);
}
extern "C" __global__ void skip_rmsnorm_f32(
const float* input,
const float* skip,
const float* gamma,
const float* bias,
float* y,
float* sum_out,
float* mean_out,
float* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
skip_rmsnorm_f32_tpl<false, false>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, epsilon);
}
extern "C" __global__ void skip_rmsnorm_f32_warp(
const float* input,
const float* skip,
const float* gamma,
const float* bias,
float* y,
float* sum_out,
float* mean_out,
float* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const float epsilon)
{
skip_rmsnorm_f32_tpl<false, true>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, epsilon);
}
// ── bf16 fused Add→RMSNorm, byte-exact with standalone `Add(bf16)` +
// `rmsnorm_bf16`. Two invariants make it byte-exact:
// (1) the residual sum is rounded to bf16 BEFORE the RMS reduction — exactly
// what the standalone bf16 `Add` op writes (`__float2bfloat16_rn(
// f32(input) + f32(skip))`), so the value stored to `y`/`sum_out` (the
// residual reused by the next layer) is bit-identical, and the fp32
// mean-square accumulates over the SAME bf16-rounded operand rmsnorm_bf16
// would read back from DRAM;
// (2) the reduction is the identical fixed block tree rmsnorm_bf16 uses
// (strided per-thread partials over `blockDim.x` threads, then power-of-two
// shared-memory combine), launched with the same NORM_BLOCK config, so the
// summation ORDER matches bit-for-bit.
// gamma is only a final multiplicand (never in the fp32 variance), so an fp32 or
// bf16 gamma are both loaded at full precision.
__device__ __forceinline__ float load_skip_bf16_param(
const void* values, const int is_bf16, const int index) {
return is_bf16
? __bfloat162float(((const __nv_bfloat16*)values)[index])
: ((const float*)values)[index];
}
template <bool WarpTail>
__device__ __forceinline__ void skip_rmsnorm_bf16_tpl(
const __nv_bfloat16* input,
const __nv_bfloat16* skip,
const void* gamma,
const void* bias, // null when absent
__nv_bfloat16* y,
__nv_bfloat16* sum_out, // null when not requested
void* mean_out, // null when not requested (always zero)
void* invstd_out, // null when not requested
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_bf16,
const int bias_is_bf16,
const int stat_is_bf16,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
const unsigned long long* shape = metadata;
const unsigned long long* skip_strides = metadata + rank;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
// Pass 1: residual sum (rounded to bf16, stored) + fp32 mean-square over the
// rounded value. `ss += xv*xv` / `red[tid] += red[tid+off]` match rmsnorm_bf16
// exactly (both compile to fadd.rn/fmul.rn).
float ss = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
unsigned long long skip_index = (unsigned long long)base + j;
if (!dense_skip) {
unsigned long long linear = skip_index;
skip_index = 0;
for (int d = rank - 1; d >= 0; --d) {
const unsigned long long coord = linear % shape[d];
linear /= shape[d];
skip_index += coord * skip_strides[d];
}
}
float sv = __bfloat162float(input[base + j]) + __bfloat162float(skip[skip_index]);
if (has_bias) sv += load_skip_bf16_param(bias, bias_is_bf16, j);
const __nv_bfloat16 svb = __float2bfloat16_rn(sv);
y[base + j] = svb;
if (sum_out) sum_out[base + j] = svb;
const float rounded = __bfloat162float(svb);
ss += rounded * rounded;
}
red[tid] = ss;
__syncthreads();
float reduced;
if (WarpTail) {
reduced = skip_rmsnorm_warp_tail(red, tid, nt);
} else {
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
reduced = red[0];
}
const float inv_std = 1.0f / sqrtf(reduced / (float)norm_size + epsilon);
if (tid == 0) {
if (mean_out) {
if (stat_is_bf16) ((__nv_bfloat16*)mean_out)[g] = __float2bfloat16_rn(0.0f);
else ((float*)mean_out)[g] = 0.0f;
}
if (invstd_out) {
if (stat_is_bf16) ((__nv_bfloat16*)invstd_out)[g] = __float2bfloat16_rn(inv_std);
else ((float*)invstd_out)[g] = inv_std;
}
}
__syncthreads();
// Pass 2: scale by inv_std · gamma, round to bf16 — identical to rmsnorm_bf16.
for (int j = tid; j < norm_size; j += nt) {
const float o = __bfloat162float(y[base + j]) * inv_std
* load_skip_bf16_param(gamma, gamma_is_bf16, j);
y[base + j] = __float2bfloat16_rn(o);
}
}
extern "C" __global__ void skip_rmsnorm_bf16(
const __nv_bfloat16* input,
const __nv_bfloat16* skip,
const void* gamma,
const void* bias, // null when absent
__nv_bfloat16* y,
__nv_bfloat16* sum_out, // null when not requested
void* mean_out, // null when not requested (always zero)
void* invstd_out, // null when not requested
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_bf16,
const int bias_is_bf16,
const int stat_is_bf16,
const float epsilon)
{
skip_rmsnorm_bf16_tpl<false>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, dense_skip,
gamma_is_bf16, bias_is_bf16, stat_is_bf16, epsilon);
}
extern "C" __global__ void skip_rmsnorm_bf16_warp(
const __nv_bfloat16* input,
const __nv_bfloat16* skip,
const void* gamma,
const void* bias, // null when absent
__nv_bfloat16* y,
__nv_bfloat16* sum_out, // null when not requested
void* mean_out, // null when not requested (always zero)
void* invstd_out, // null when not requested
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_bf16,
const int bias_is_bf16,
const int stat_is_bf16,
const float epsilon)
{
skip_rmsnorm_bf16_tpl<true>(input, skip, gamma, bias, y, sum_out, mean_out,
invstd_out, metadata, rank, num_groups, norm_size, has_bias, dense_skip,
gamma_is_bf16, bias_is_bf16, stat_is_bf16, epsilon);
}
union SkipHalf4 {
unsigned long long raw;
__half2 pair[2];
};
// One warp covers aligned half4 chunks. The launch predicate guarantees that
// norm_size is divisible by 32 lanes * 4 halves, so every lane owns the same
// number of complete chunks and no tail handling is needed.
extern "C" __global__ void skip_rmsnorm_f16_warp_half4(
const __half* input,
const __half* skip,
const void* gamma,
const void* bias,
__half* y,
__half* sum_out,
void* mean_out,
void* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_half,
const int bias_is_half,
const int stat_is_half,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
const int lane = threadIdx.x;
const int chunks_per_lane = norm_size / (32 * 4);
const unsigned long long* input4 =
(const unsigned long long*)(input + base);
const unsigned long long* skip4 =
(const unsigned long long*)(skip + base);
const unsigned long long* gamma4 =
(const unsigned long long*)gamma;
unsigned long long* y4 = (unsigned long long*)(y + base);
unsigned long long* sum4 =
sum_out ? (unsigned long long*)(sum_out + base) : 0;
float ss0 = 0.0f;
float ss1 = 0.0f;
float ss2 = 0.0f;
float ss3 = 0.0f;
for (int item = 0; item < chunks_per_lane; ++item) {
const int chunk = lane + item * 32;
SkipHalf4 input_v;
SkipHalf4 skip_v;
SkipHalf4 residual;
input_v.raw = input4[chunk];
skip_v.raw = skip4[chunk];
residual.pair[0] = __hadd2(input_v.pair[0], skip_v.pair[0]);
residual.pair[1] = __hadd2(input_v.pair[1], skip_v.pair[1]);
y4[chunk] = residual.raw;
if (sum4) sum4[chunk] = residual.raw;
const float2 rounded0 = __half22float2(residual.pair[0]);
const float2 rounded1 = __half22float2(residual.pair[1]);
ss0 += rounded0.x * rounded0.x;
ss1 += rounded0.y * rounded0.y;
ss2 += rounded1.x * rounded1.x;
ss3 += rounded1.y * rounded1.y;
}
float ss = (ss0 + ss1) + (ss2 + ss3);
for (int off = 16; off > 0; off >>= 1) {
ss += __shfl_down_sync(0xffffffffu, ss, off);
}
float inv_std = 0.0f;
if (lane == 0) {
inv_std = 1.0f / sqrtf(ss / (float)norm_size + epsilon);
if (mean_out) {
if (stat_is_half) ((__half*)mean_out)[g] = __float2half_rn(0.0f);
else ((float*)mean_out)[g] = 0.0f;
}
if (invstd_out) {
if (stat_is_half) ((__half*)invstd_out)[g] = __float2half_rn(inv_std);
else ((float*)invstd_out)[g] = inv_std;
}
}
inv_std = __shfl_sync(0xffffffffu, inv_std, 0);
const float* gamma_f = (const float*)gamma;
for (int item = 0; item < chunks_per_lane; ++item) {
const int chunk = lane + item * 32;
SkipHalf4 residual;
SkipHalf4 output;
residual.raw = y4[chunk];
const float2 value0 = __half22float2(residual.pair[0]);
const float2 value1 = __half22float2(residual.pair[1]);
// gamma is only ever a final multiplicand (never part of the fp32
// variance accumulation), so an fp32 gamma is loaded at full precision
// while an fp16 gamma keeps the wide half4 load. This lets decoders that
// export gamma in fp32 (e.g. Phi) still take the vectorized warp path.
float scale0x, scale0y, scale1x, scale1y;
if (gamma_is_half) {
SkipHalf4 scale;
scale.raw = gamma4[chunk];
const float2 scale0 = __half22float2(scale.pair[0]);
const float2 scale1 = __half22float2(scale.pair[1]);
scale0x = scale0.x;
scale0y = scale0.y;
scale1x = scale1.x;
scale1y = scale1.y;
} else {
const int j = chunk << 2;
scale0x = gamma_f[j];
scale0y = gamma_f[j + 1];
scale1x = gamma_f[j + 2];
scale1y = gamma_f[j + 3];
}
output.pair[0] = __floats2half2_rn(
value0.x * inv_std * scale0x,
value0.y * inv_std * scale0y);
output.pair[1] = __floats2half2_rn(
value1.x * inv_std * scale1x,
value1.y * inv_std * scale1y);
y4[chunk] = output.raw;
}
}
// Decode-shaped variant of `skip_rmsnorm_f16_warp_half4`. The warp path runs one
// 32-lane warp per row, which saturates the machine only when there are many
// rows (prefill). At decode there is a single row (num_groups == 1), so a lone
// warp leaves the GPU almost entirely idle (measured 1.56% achieved occupancy,
// Grid 1 x Block 32) and stalls ~92% of cycles on Long-Scoreboard global-load
// latency with too few resident warps to hide it. This variant spreads the same
// half4 chunks of one row across a full multi-warp block and reduces the
// sum-of-squares through the file's launch-invariant `red[tid]` block tree, so
// many warps are resident to hide the load latency. Same launch predicate as the
// warp path (norm_size % 128 == 0, dense skip, no bias). The residual (`y` /
// `sum_out`) is written per-chunk by exactly one thread with the identical
// __hadd2 rounding, so it is byte-identical to the warp path; only the fp32
// sum-of-squares reduction order differs, perturbing the shared 1/rms by ULPs.
template <bool WarpTail>
__device__ __forceinline__ void skip_rmsnorm_f16_block_half4_tpl(
const __half* input,
const __half* skip,
const void* gamma,
const void* bias,
__half* y,
__half* sum_out,
void* mean_out,
void* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_half,
const int bias_is_half,
const int stat_is_half,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
const int tid = threadIdx.x;
const int nt = blockDim.x;
const int chunks = norm_size >> 2;
const unsigned long long* input4 =
(const unsigned long long*)(input + base);
const unsigned long long* skip4 =
(const unsigned long long*)(skip + base);
const unsigned long long* gamma4 =
(const unsigned long long*)gamma;
unsigned long long* y4 = (unsigned long long*)(y + base);
unsigned long long* sum4 =
sum_out ? (unsigned long long*)(sum_out + base) : 0;
extern __shared__ float red[];
float ss = 0.0f;
for (int chunk = tid; chunk < chunks; chunk += nt) {
SkipHalf4 input_v;
SkipHalf4 skip_v;
SkipHalf4 residual;
input_v.raw = input4[chunk];
skip_v.raw = skip4[chunk];
residual.pair[0] = __hadd2(input_v.pair[0], skip_v.pair[0]);
residual.pair[1] = __hadd2(input_v.pair[1], skip_v.pair[1]);
y4[chunk] = residual.raw;
if (sum4) sum4[chunk] = residual.raw;
const float2 rounded0 = __half22float2(residual.pair[0]);
const float2 rounded1 = __half22float2(residual.pair[1]);
ss += rounded0.x * rounded0.x;
ss += rounded0.y * rounded0.y;
ss += rounded1.x * rounded1.x;
ss += rounded1.y * rounded1.y;
}
red[tid] = ss;
__syncthreads();
float reduced;
if (WarpTail) {
reduced = skip_rmsnorm_warp_tail(red, tid, nt);
} else {
for (int offset = nt >> 1; offset > 0; offset >>= 1) {
if (tid < offset) red[tid] += red[tid + offset];
__syncthreads();
}
reduced = red[0];
}
const float inv_std = 1.0f / sqrtf(reduced / (float)norm_size + epsilon);
if (tid == 0) {
if (mean_out) {
if (stat_is_half) ((__half*)mean_out)[g] = __float2half_rn(0.0f);
else ((float*)mean_out)[g] = 0.0f;
}
if (invstd_out) {
if (stat_is_half) ((__half*)invstd_out)[g] = __float2half_rn(inv_std);
else ((float*)invstd_out)[g] = inv_std;
}
}
const float* gamma_f = (const float*)gamma;
for (int chunk = tid; chunk < chunks; chunk += nt) {
SkipHalf4 residual;
SkipHalf4 output;
residual.raw = y4[chunk];
const float2 value0 = __half22float2(residual.pair[0]);
const float2 value1 = __half22float2(residual.pair[1]);
float scale0x, scale0y, scale1x, scale1y;
if (gamma_is_half) {
SkipHalf4 scale;
scale.raw = gamma4[chunk];
const float2 scale0 = __half22float2(scale.pair[0]);
const float2 scale1 = __half22float2(scale.pair[1]);
scale0x = scale0.x;
scale0y = scale0.y;
scale1x = scale1.x;
scale1y = scale1.y;
} else {
const int j = chunk << 2;
scale0x = gamma_f[j];
scale0y = gamma_f[j + 1];
scale1x = gamma_f[j + 2];
scale1y = gamma_f[j + 3];
}
output.pair[0] = __floats2half2_rn(
value0.x * inv_std * scale0x,
value0.y * inv_std * scale0y);
output.pair[1] = __floats2half2_rn(
value1.x * inv_std * scale1x,
value1.y * inv_std * scale1y);
y4[chunk] = output.raw;
}
}
extern "C" __global__ void skip_rmsnorm_f16_block_half4(
const __half* input,
const __half* skip,
const void* gamma,
const void* bias,
__half* y,
__half* sum_out,
void* mean_out,
void* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_half,
const int bias_is_half,
const int stat_is_half,
const float epsilon)
{
skip_rmsnorm_f16_block_half4_tpl<false>(input, skip, gamma, bias, y, sum_out,
mean_out, invstd_out, metadata, rank, num_groups, norm_size, has_bias,
dense_skip, gamma_is_half, bias_is_half, stat_is_half, epsilon);
}
extern "C" __global__ void skip_rmsnorm_f16_block_half4_warp(
const __half* input,
const __half* skip,
const void* gamma,
const void* bias,
__half* y,
__half* sum_out,
void* mean_out,
void* invstd_out,
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_half,
const int bias_is_half,
const int stat_is_half,
const float epsilon)
{
skip_rmsnorm_f16_block_half4_tpl<true>(input, skip, gamma, bias, y, sum_out,
mean_out, invstd_out, metadata, rank, num_groups, norm_size, has_bias,
dense_skip, gamma_is_half, bias_is_half, stat_is_half, epsilon);
}
extern "C" __global__ void skip_rmsnorm_f16(
const __half* input,
const __half* skip,
const void* gamma,
const void* bias, // null when absent
__half* y,
__half* sum_out, // null when not requested
void* mean_out, // null when not requested (always zero)
void* invstd_out, // null when not requested
const unsigned long long* metadata,
const int rank,
const int num_groups,
const int norm_size,
const int has_bias,
const int dense_skip,
const int gamma_is_half,
const int bias_is_half,
const int stat_is_half,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
const unsigned long long* shape = metadata;
const unsigned long long* skip_strides = metadata + rank;
const int lane = threadIdx.x;
// fp32 accumulate over the fp16-rounded residual so the RMS matches the
// residual value stored into `sum_out` and reused by the next layer.
float ss = 0.0f;
const bool vectorized = dense_skip && ((base & 1) == 0);
if (vectorized) {
const int pairs = norm_size >> 1;
const __half2* input2 = (const __half2*)(input + base);
const __half2* skip2 = (const __half2*)(skip + base);
__half2* y2 = (__half2*)(y + base);
__half2* sum2 = sum_out ? (__half2*)(sum_out + base) : 0;
for (int pair = lane; pair < pairs; pair += 32) {
const float2 input_v = __half22float2(input2[pair]);
const float2 skip_v = __half22float2(skip2[pair]);
const int j = pair << 1;
float sv0 = input_v.x + skip_v.x;
float sv1 = input_v.y + skip_v.y;
if (has_bias) {
sv0 += load_skip_val(bias, bias_is_half, j);
sv1 += load_skip_val(bias, bias_is_half, j + 1);
}
const __half svh0 = __float2half_rn(sv0);
const __half svh1 = __float2half_rn(sv1);
const __half2 svh = __halves2half2(svh0, svh1);
y2[pair] = svh;
if (sum2) sum2[pair] = svh;
const float2 rounded = __half22float2(svh);
ss += rounded.x * rounded.x;
ss += rounded.y * rounded.y;
}
if ((norm_size & 1) && lane == 0) {
const int j = norm_size - 1;
float sv = __half2float(input[base + j]) + __half2float(skip[base + j]);
if (has_bias) sv += load_skip_val(bias, bias_is_half, j);
const __half svh = __float2half_rn(sv);
y[base + j] = svh;
if (sum_out) sum_out[base + j] = svh;
const float rounded = __half2float(svh);
ss += rounded * rounded;
}
} else {
for (int j = lane; j < norm_size; j += 32) {
unsigned long long linear = (unsigned long long)base + j;
unsigned long long skip_index = 0;
for (int d = rank - 1; d >= 0; --d) {
const unsigned long long coord = linear % shape[d];
linear /= shape[d];
skip_index += coord * skip_strides[d];
}
float sv = __half2float(input[base + j]) + __half2float(skip[skip_index]);
if (has_bias) sv += load_skip_val(bias, bias_is_half, j);
const __half svh = __float2half_rn(sv);
y[base + j] = svh;
if (sum_out) sum_out[base + j] = svh;
const float rounded = __half2float(svh);
ss += rounded * rounded;
}
}
__syncwarp();
for (int off = 16; off > 0; off >>= 1) {
ss += __shfl_down_sync(0xffffffffu, ss, off);
}
float inv_std = 0.0f;
if (lane == 0) {
inv_std = 1.0f / sqrtf(ss / (float)norm_size + epsilon);
if (mean_out) {
if (stat_is_half) ((__half*)mean_out)[g] = __float2half_rn(0.0f);
else ((float*)mean_out)[g] = 0.0f;
}
if (invstd_out) {
if (stat_is_half) ((__half*)invstd_out)[g] = __float2half_rn(inv_std);
else ((float*)invstd_out)[g] = inv_std;
}
}
inv_std = __shfl_sync(0xffffffffu, inv_std, 0);
if (vectorized) {
const int pairs = norm_size >> 1;
__half2* y2 = (__half2*)(y + base);
for (int pair = lane; pair < pairs; pair += 32) {
const float2 residual = __half22float2(y2[pair]);
const int j = pair << 1;
const float out0 = residual.x * inv_std
* load_skip_val(gamma, gamma_is_half, j);
const float out1 = residual.y * inv_std
* load_skip_val(gamma, gamma_is_half, j + 1);
y2[pair] = __floats2half2_rn(out0, out1);
}
if ((norm_size & 1) && lane == 0) {
const int j = norm_size - 1;
const float v = __half2float(y[base + j]) * inv_std
* load_skip_val(gamma, gamma_is_half, j);
y[base + j] = __float2half_rn(v);
}
} else {
for (int j = lane; j < norm_size; j += 32) {
const float v = __half2float(y[base + j]) * inv_std
* load_skip_val(gamma, gamma_is_half, j);
y[base + j] = __float2half_rn(v);
}
}
}
"#;
const SKIP_LAYERNORM_SRC: &str = r#"
#include <cuda_fp16.h>
#include <cuda_bf16.h>
__device__ __forceinline__ float skip_ln_load(
const void* data, size_t index, int dtype) {
if (dtype == 0) return ((const float*)data)[index];
if (dtype == 1) return __half2float(((const __half*)data)[index]);
return __bfloat162float(((const __nv_bfloat16*)data)[index]);
}
__device__ __forceinline__ void skip_ln_store(
void* data, size_t index, float value, int dtype) {
if (dtype == 0) ((float*)data)[index] = value;
else if (dtype == 1) ((__half*)data)[index] = __float2half_rn(value);
else ((__nv_bfloat16*)data)[index] = __float2bfloat16_rn(value);
}
extern "C" __global__ void skip_layernorm(
const void* input,
const void* skip,
const void* gamma,
const void* beta, // null when absent
const void* bias, // null when absent (per-channel, length norm_size)
void* y,
void* sum_out, // null when not requested
void* mean_out, // null when not requested
void* invstd_out, // null when not requested
const int num_groups,
const int norm_size,
const int dtype,
const int gamma_dtype,
const int beta_dtype,
const int bias_dtype,
const int stat_dtype,
const int has_beta,
const int has_bias,
const float epsilon)
{
const int g = blockIdx.x;
if (g >= num_groups) return;
const size_t base = (size_t)g * norm_size;
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
// Residual sum s = input + skip (+ bias); stash in y and optionally sum_out.
for (int j = tid; j < norm_size; j += nt) {
float sv = skip_ln_load(input, base + j, dtype)
+ skip_ln_load(skip, base + j, dtype);
if (has_bias) sv += skip_ln_load(bias, j, bias_dtype);
skip_ln_store(y, base + j, sv, dtype);
if (sum_out) skip_ln_store(sum_out, base + j, sv, dtype);
}
__syncthreads();
// Pass 1: mean of s.
float s = 0.0f;
for (int j = tid; j < norm_size; j += nt)
s += skip_ln_load(y, base + j, dtype);
red[tid] = s;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float mean = red[0] / (float)norm_size;
__syncthreads();
// Pass 2: population variance of s.
float v = 0.0f;
for (int j = tid; j < norm_size; j += nt) {
const float d = skip_ln_load(y, base + j, dtype) - mean;
v += d * d;
}
red[tid] = v;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red[tid] += red[tid + off];
__syncthreads();
}
const float var = red[0] / (float)norm_size;
const float inv_std = 1.0f / sqrtf(var + epsilon);
if (tid == 0) {
if (mean_out) skip_ln_store(mean_out, g, mean, stat_dtype);
if (invstd_out) skip_ln_store(invstd_out, g, inv_std, stat_dtype);
}
__syncthreads();
// Pass 3: normalize + affine (gamma / optional beta).
for (int j = tid; j < norm_size; j += nt) {
const float xhat = (skip_ln_load(y, base + j, dtype) - mean) * inv_std;
float o = xhat * skip_ln_load(gamma, j, gamma_dtype);
if (has_beta) o += skip_ln_load(beta, j, beta_dtype);
skip_ln_store(y, base + j, o, dtype);
}
}
"#;
const LAYERNORM_MODULE: &str = "layernorm_bf16_v2";
const RMSNORM_MODULE: &str = "rmsnorm_bf16_v4";
const SKIP_RMSNORM_MODULE: &str = "skip_rmsnorm_f16_warp_v6_block";
pub(crate) const SKIP_RMSNORM_BLOCK_MAX_GROUPS: u32 = 8;
const SKIP_RMSNORM_BLOCK_THREADS: u32 = 1024;
const SKIP_RMSNORM_BLOCK_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_SKIP_RMSNORM_BLOCK";
fn skip_rmsnorm_block_disabled() -> bool {
std::env::var_os(SKIP_RMSNORM_BLOCK_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
const SKIP_LAYERNORM_MODULE: &str = "skip_layernorm_typed_v2";
const NORM_BLOCK: u32 = 256;
const SKIP_RMSNORM_WARP_HALF4_MULTIPLE: usize = 32 * 4;
fn preferred_norm_block_threads(norm_size: usize, max_threads_per_block: u32) -> u32 {
let reported_limit = max_threads_per_block.clamp(32, NORM_BLOCK);
let device_limit = 1 << (31 - reported_limit.leading_zeros());
let useful_threads = norm_size
.max(32)
.next_power_of_two()
.min(NORM_BLOCK as usize) as u32;
useful_threads.min(device_limit)
}
fn norm_block_threads(
norm_size: usize,
num_groups: usize,
multiprocessor_count: u32,
max_threads_per_block: u32,
) -> u32 {
let cap = if num_groups < multiprocessor_count.max(1) as usize {
max_threads_per_block.max(NORM_BLOCK)
} else {
NORM_BLOCK
};
let reported_limit = max_threads_per_block.clamp(32, cap);
let device_limit = 1 << (31 - reported_limit.leading_zeros());
let useful_threads = norm_size.max(32).next_power_of_two().min(cap as usize) as u32;
useful_threads.min(device_limit)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum SkipRmsnormVariant {
F32Dense,
F32,
F16Generic,
F16WarpHalf4,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct SkipRmsnormSelection {
variant: SkipRmsnormVariant,
entry: &'static str,
reason: &'static str,
}
const FP32_GAMMA_WARP_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_FP32_GAMMA_WARP_NORM";
fn fp32_gamma_warp_disabled() -> bool {
std::env::var_os(FP32_GAMMA_WARP_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
fn select_skip_rmsnorm_variant(
is_half: bool,
dense_skip: bool,
norm_size: usize,
has_bias: bool,
gamma_is_half: bool,
) -> SkipRmsnormSelection {
let gamma_ok = gamma_is_half || !fp32_gamma_warp_disabled();
if is_half
&& dense_skip
&& norm_size.is_multiple_of(SKIP_RMSNORM_WARP_HALF4_MULTIPLE)
&& !has_bias
&& gamma_ok
{
SkipRmsnormSelection {
variant: SkipRmsnormVariant::F16WarpHalf4,
entry: "skip_rmsnorm_f16_warp_half4",
reason: if gamma_is_half {
"variant=warp_half4;dtype=fp16;dense_skip;bias=none;gamma=fp16;\
hidden%128==0;one_warp"
} else {
"variant=warp_half4;dtype=fp16;dense_skip;bias=none;gamma=fp32;\
hidden%128==0;one_warp"
},
}
} else if is_half {
SkipRmsnormSelection {
variant: SkipRmsnormVariant::F16Generic,
entry: "skip_rmsnorm_f16",
reason: "variant=generic;dtype=fp16;not(dense_skip & bias=none & \
hidden%128==0)",
}
} else if dense_skip {
SkipRmsnormSelection {
variant: SkipRmsnormVariant::F32Dense,
entry: "skip_rmsnorm_f32_dense",
reason: "variant=parallel_block_tree;dtype=fp32;dense_skip;fixed_reduction_order",
}
} else {
SkipRmsnormSelection {
variant: SkipRmsnormVariant::F32,
entry: "skip_rmsnorm_f32",
reason: "variant=generic;dtype=fp32",
}
}
}
fn require_f32(op: &str, name: &str, dt: DataType) -> Result<()> {
if dt != DataType::Float32 {
return Err(not_implemented(format!(
"{op} with {name} dtype {dt:?} (this slice is f32-only; f16/bf16 pending)"
)));
}
Ok(())
}
fn require_float_storage(op: &str, name: &str, dt: DataType) -> Result<()> {
if !matches!(
dt,
DataType::Float16 | DataType::BFloat16 | DataType::Float32
) {
return Err(not_implemented(format!(
"{op} with {name} dtype {dt:?} (expected f16, bf16, or f32)"
)));
}
Ok(())
}
fn require_param_for_activation(
op: &str,
name: &str,
activation: DataType,
parameter: DataType,
) -> Result<()> {
if parameter != DataType::Float32 && parameter != activation {
return Err(not_implemented(format!(
"{op} with {name} dtype {parameter:?} for {activation:?} activations \
(expected matching storage dtype or f32)"
)));
}
Ok(())
}
fn require_f16_or_f32(op: &str, name: &str, dt: DataType) -> Result<()> {
if !matches!(dt, DataType::Float16 | DataType::Float32) {
return Err(not_implemented(format!(
"{op} with {name} dtype {dt:?} (expected f16 or f32)"
)));
}
Ok(())
}
fn layernorm_entry(dtype: DataType) -> &'static str {
match dtype {
DataType::Float16 => "layernorm_f16",
DataType::BFloat16 => "layernorm_bf16",
DataType::Float32 => "layernorm_f32",
_ => unreachable!("LayerNormalization dtype must be validated before dispatch"),
}
}
fn rmsnorm_entry(dtype: DataType) -> &'static str {
match dtype {
DataType::Float16 => "rmsnorm_f16",
DataType::BFloat16 => "rmsnorm_bf16",
DataType::Float32 => "rmsnorm_f32",
_ => unreachable!("RMSNormalization dtype must be validated before dispatch"),
}
}
fn require_contiguous(op: &str, name: &str, contiguous: bool) -> Result<()> {
if !contiguous {
return Err(not_implemented(format!(
"{op} with a non-contiguous (strided) {name}; \
insert an explicit copy to materialise it before the op"
)));
}
Ok(())
}
fn dim_overflow(op: &str, name: &str, v: usize) -> EpError {
EpError::KernelFailed(format!(
"cuda_ep {op}: {name} ({v}) exceeds the i32 kernel bound"
))
}
pub struct LayerNormFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for LayerNormFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let axis = node.attr("axis").and_then(|a| a.as_int()).unwrap_or(-1);
let epsilon = node
.attr("epsilon")
.and_then(|a| a.as_float())
.unwrap_or(1e-5);
Ok(Box::new(LayerNormKernel {
axis,
epsilon,
runtime: self.runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}))
}
}
#[derive(Debug)]
pub struct LayerNormKernel {
axis: i64,
epsilon: f32,
runtime: Arc<CudaRuntime>,
warmed_signature: Mutex<Option<NormCaptureSignature>>,
last_call_capture_safe: AtomicBool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct NormCaptureSignature {
activation_dtype: DataType,
scale_dtype: DataType,
bias_dtype: Option<DataType>,
input_shape: Vec<usize>,
output_dtypes_and_shapes: Vec<(DataType, Vec<usize>)>,
}
impl NormCaptureSignature {
fn build(
activation_dtype: DataType,
scale_dtype: DataType,
bias_dtype: Option<DataType>,
input_shape: &[usize],
outputs: &[TensorMut],
) -> Self {
Self {
activation_dtype,
scale_dtype,
bias_dtype,
input_shape: input_shape.to_vec(),
output_dtypes_and_shapes: outputs
.iter()
.map(|output| (output.dtype, output.shape.to_vec()))
.collect(),
}
}
}
impl LayerNormKernel {
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.last_call_capture_safe.store(false, Ordering::Relaxed);
if !(2..=3).contains(&inputs.len()) || outputs.is_empty() || outputs.len() > 3 {
return Err(EpError::KernelFailed(format!(
"cuda_ep LayerNormalization: expected 2-3 inputs (X, Scale[, B]) and \
1-3 outputs (Y[, Mean, InvStdDev]), got {} and {}",
inputs.len(),
outputs.len()
)));
}
let x = &inputs[0];
let scale = &inputs[1];
let bias = inputs.get(2);
require_float_storage("LayerNormalization", "X", x.dtype)?;
if x.dtype == DataType::Float32 {
require_f32("LayerNormalization", "Scale", scale.dtype)?;
} else {
require_param_for_activation("LayerNormalization", "Scale", x.dtype, scale.dtype)?;
}
if outputs[0].dtype != x.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep LayerNormalization: Y dtype {:?} must match X dtype {:?}",
outputs[0].dtype, x.dtype
)));
}
require_contiguous("LayerNormalization", "X", x.is_contiguous())?;
require_contiguous("LayerNormalization", "Scale", scale.is_contiguous())?;
require_contiguous("LayerNormalization", "Y", outputs[0].is_contiguous())?;
let rank = x.shape.len();
let axis = resolve_axis("LayerNormalization", self.axis, rank)?;
let norm_size: usize = x.shape[axis..].iter().product();
let num_groups: usize = x.shape[..axis].iter().product();
if norm_size == 0 {
return Err(EpError::KernelFailed(
"cuda_ep LayerNormalization: empty normalization axis".into(),
));
}
if scale.numel() != norm_size {
return Err(EpError::KernelFailed(format!(
"cuda_ep LayerNormalization: Scale has {} elements, expected {norm_size} \
(= prod(shape[axis..]))",
scale.numel()
)));
}
let bias_ptr = match bias {
None => 0u64,
Some(b) => {
if x.dtype == DataType::Float32 {
require_f32("LayerNormalization", "B", b.dtype)?;
} else {
require_param_for_activation("LayerNormalization", "B", x.dtype, b.dtype)?;
}
require_contiguous("LayerNormalization", "B", b.is_contiguous())?;
if b.numel() != norm_size {
return Err(EpError::KernelFailed(format!(
"cuda_ep LayerNormalization: B has {} elements, expected {norm_size}",
b.numel()
)));
}
cuptr(b.data_ptr::<u8>() as *const c_void)
}
};
if outputs[0].shape != x.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep LayerNormalization: Y shape {:?} must equal X shape {:?}",
outputs[0].shape, x.shape
)));
}
if num_groups == 0 {
return Ok(());
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let elements = x.numel() as u64;
let groups = num_groups as u64;
let mut flops = elements
.saturating_mul(7)
.saturating_add(groups.saturating_mul(5));
if bias.is_some() {
flops = flops.saturating_add(elements);
}
flops
});
let x_ptr = cuptr(x.data_ptr::<u8>() as *const c_void);
let scale_ptr = cuptr(scale.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let (mean_ptr, invstd_ptr) = optional_stat_ptrs("LayerNormalization", outputs, num_groups)?;
let (groups_u, norm_i) = (
u32::try_from(num_groups)
.map_err(|_| dim_overflow("LayerNormalization", "num_groups", num_groups))?,
i32::try_from(norm_size)
.map_err(|_| dim_overflow("LayerNormalization", "norm_size", norm_size))?,
);
let has_bias: i32 = i32::from(bias_ptr != 0);
let eps = self.epsilon;
let groups_i = groups_u_i32(groups_u);
let signature = NormCaptureSignature::build(
x.dtype,
scale.dtype,
bias.map(|bias| bias.dtype),
x.shape,
outputs,
);
let capturing = self.runtime.is_capturing()?;
let mut warmed_signature = self
.warmed_signature
.lock()
.expect("cuda_ep LayerNormalization capture signature poisoned");
if capturing && warmed_signature.as_ref() != Some(&signature) {
return Err(EpError::KernelFailed(
"cuda_ep LayerNormalization: dtype or shape changed during CUDA graph capture; warm the exact fixed-shape signature before capture"
.into(),
));
}
let entry = layernorm_entry(x.dtype);
let func = self
.runtime
.nvrtc_function(LAYERNORM_MODULE, LAYERNORM_SRC, entry)?;
let cfg = self.runtime.reduction_launch_config(
&func,
groups_u,
NORM_BLOCK,
std::mem::size_of::<f32>() as u32,
)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
let scale_is_half = i32::from(scale.dtype == DataType::Float16);
let bias_is_half = i32::from(bias.is_some_and(|bias| bias.dtype == DataType::Float16));
let scale_is_bf16 = i32::from(scale.dtype == DataType::BFloat16);
let bias_is_bf16 = i32::from(bias.is_some_and(|bias| bias.dtype == DataType::BFloat16));
builder
.arg(&x_ptr)
.arg(&scale_ptr)
.arg(&bias_ptr)
.arg(&y_ptr)
.arg(&mean_ptr)
.arg(&invstd_ptr)
.arg(&groups_i)
.arg(&norm_i);
match x.dtype {
DataType::Float16 => {
builder
.arg(&scale_is_half)
.arg(&bias_is_half)
.arg(&has_bias)
.arg(&eps);
}
DataType::BFloat16 => {
builder
.arg(&scale_is_bf16)
.arg(&bias_is_bf16)
.arg(&has_bias)
.arg(&eps);
}
DataType::Float32 => {
builder.arg(&has_bias).arg(&eps);
}
_ => unreachable!("LayerNormalization dtype validated above"),
}
unsafe { builder.launch(cfg) }.map_err(|e| driver_err(&format!("launch {entry}"), e))?;
if !capturing {
*warmed_signature = Some(signature.clone());
}
self.last_call_capture_safe.store(true, Ordering::Relaxed);
Ok(())
}
}
impl Kernel for LayerNormKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.run(inputs, outputs)
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
fn capture_support(&self) -> onnx_runtime_ep_api::CaptureSupport {
if self.last_call_capture_safe.load(Ordering::Relaxed) {
onnx_runtime_ep_api::CaptureSupport::Supported
} else {
onnx_runtime_ep_api::CaptureSupport::unsupported(
"LayerNormalization shape/dtype signature does not match the warmed fixed-shape capture signature",
)
}
}
}
pub struct RmsNormFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for RmsNormFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
if node
.attr("stash_type")
.is_some_and(|attribute| attribute.as_int() != Some(1))
{
return Err(EpError::KernelFailed(
"RMSNormalization: stash_type must be 1 (float)".into(),
));
}
let axis = node.attr("axis").and_then(|a| a.as_int()).unwrap_or(-1);
let epsilon = node
.attr("epsilon")
.and_then(|a| a.as_float())
.unwrap_or(1e-5);
Ok(Box::new(RmsNormKernel {
axis,
epsilon,
runtime: self.runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}))
}
}
#[derive(Debug)]
pub struct RmsNormKernel {
axis: i64,
epsilon: f32,
runtime: Arc<CudaRuntime>,
warmed_signature: Mutex<Option<NormCaptureSignature>>,
last_call_capture_safe: AtomicBool,
}
impl RmsNormKernel {
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.last_call_capture_safe.store(false, Ordering::Relaxed);
let op = "RMSNormalization";
if inputs.len() != 2 || outputs.is_empty() || outputs.len() > 2 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: expected 2 inputs (X, Scale) and 1-2 outputs \
(Y[, InvStdDev]), got {} and {}",
inputs.len(),
outputs.len()
)));
}
let x = &inputs[0];
let scale = &inputs[1];
require_float_storage(op, "X", x.dtype)?;
if x.dtype == DataType::Float32 {
require_f32(op, "Scale", scale.dtype)?;
} else {
require_param_for_activation(op, "Scale", x.dtype, scale.dtype)?;
}
if outputs[0].dtype != x.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: Y dtype {:?} must match X dtype {:?}",
outputs[0].dtype, x.dtype
)));
}
require_contiguous(op, "X", x.is_contiguous())?;
require_contiguous(op, "Scale", scale.is_contiguous())?;
require_contiguous(op, "Y", outputs[0].is_contiguous())?;
let rank = x.shape.len();
let axis = resolve_axis(op, self.axis, rank)?;
let norm_size: usize = x.shape[axis..].iter().product();
let num_groups: usize = x.shape[..axis].iter().product();
if norm_size == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: empty normalization axis"
)));
}
if scale.numel() != norm_size {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: Scale has {} elements, expected {norm_size}",
scale.numel()
)));
}
if outputs[0].shape != x.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: Y shape {:?} must equal X shape {:?}",
outputs[0].shape, x.shape
)));
}
if num_groups == 0 {
return Ok(());
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let elements = x.numel() as u64;
elements
.saturating_mul(4)
.saturating_add((num_groups as u64).saturating_mul(4))
});
let x_ptr = cuptr(x.data_ptr::<u8>() as *const c_void);
let scale_ptr = cuptr(scale.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let invstd_ptr = match outputs.get_mut(1) {
None => 0u64,
Some(t) => {
require_f32(op, "InvStdDev", t.dtype)?;
if t.numel() != num_groups {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: InvStdDev has {} elements, expected {num_groups}",
t.numel()
)));
}
cuptr(t.data_ptr_mut::<u8>() as *const c_void)
}
};
let (groups_u, norm_i) = (
u32::try_from(num_groups).map_err(|_| dim_overflow(op, "num_groups", num_groups))?,
i32::try_from(norm_size).map_err(|_| dim_overflow(op, "norm_size", norm_size))?,
);
let eps = self.epsilon;
let signature = NormCaptureSignature::build(x.dtype, scale.dtype, None, x.shape, outputs);
let capturing = self.runtime.is_capturing()?;
let mut warmed_signature = self
.warmed_signature
.lock()
.expect("cuda_ep RMSNormalization capture signature poisoned");
if capturing && warmed_signature.as_ref() != Some(&signature) {
return Err(EpError::KernelFailed(
"cuda_ep RMSNormalization: dtype or shape changed during CUDA graph capture; warm the exact fixed-shape signature before capture"
.into(),
));
}
let entry = rmsnorm_entry(x.dtype);
let func = self
.runtime
.nvrtc_function(RMSNORM_MODULE, RMSNORM_SRC, entry)?;
let caps = self.runtime.capabilities();
let cfg = self.runtime.reduction_launch_config(
&func,
groups_u,
norm_block_threads(
norm_size,
groups_u as usize,
caps.multiprocessor_count(),
caps.max_threads_per_block(),
),
std::mem::size_of::<f32>() as u32,
)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
let groups_i = groups_u_i32(groups_u);
let scale_is_half = i32::from(scale.dtype == DataType::Float16);
let scale_is_bf16 = i32::from(scale.dtype == DataType::BFloat16);
builder
.arg(&x_ptr)
.arg(&scale_ptr)
.arg(&y_ptr)
.arg(&invstd_ptr)
.arg(&groups_i)
.arg(&norm_i);
match x.dtype {
DataType::Float16 => {
builder.arg(&scale_is_half).arg(&eps);
}
DataType::BFloat16 => {
builder.arg(&scale_is_bf16).arg(&eps);
}
DataType::Float32 => {
builder.arg(&eps);
}
_ => unreachable!("RMSNormalization dtype validated above"),
}
unsafe { builder.launch(cfg) }.map_err(|e| driver_err(&format!("launch {entry}"), e))?;
if !capturing {
*warmed_signature = Some(signature.clone());
}
self.last_call_capture_safe.store(true, Ordering::Relaxed);
Ok(())
}
}
impl Kernel for RmsNormKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.run(inputs, outputs)
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
fn capture_support(&self) -> onnx_runtime_ep_api::CaptureSupport {
if self.last_call_capture_safe.load(Ordering::Relaxed) {
onnx_runtime_ep_api::CaptureSupport::Supported
} else {
onnx_runtime_ep_api::CaptureSupport::unsupported(
"SimplifiedLayerNormalization/RMSNorm shape/dtype signature does not match the warmed fixed-shape capture signature",
)
}
}
}
pub struct SkipSimplifiedLayerNormFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for SkipSimplifiedLayerNormFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let epsilon = node
.attr("epsilon")
.and_then(|a| a.as_float())
.unwrap_or(1e-5);
Ok(Box::new(SkipSimplifiedLayerNormKernel {
epsilon,
runtime: self.runtime.clone(),
metadata: Mutex::new(SkipBroadcastMetadataCache::new(self.runtime.clone())),
bf16_scratch: Mutex::new(NormBf16Scratch::new(self.runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
last_call_used_bf16_scratch: AtomicBool::new(false),
}))
}
}
#[derive(Debug)]
pub struct SkipSimplifiedLayerNormKernel {
epsilon: f32,
runtime: Arc<CudaRuntime>,
metadata: Mutex<SkipBroadcastMetadataCache>,
bf16_scratch: Mutex<NormBf16Scratch>,
last_call_capture_safe: AtomicBool,
last_call_used_bf16_scratch: AtomicBool,
}
struct SkipNormWarmRollback<'a> {
kernel: &'a SkipSimplifiedLayerNormKernel,
metadata: Option<SkipBroadcastMetadataCache>,
scratch: Option<NormBf16Scratch>,
capture_safe: bool,
used_bf16_scratch: bool,
committed: bool,
}
impl<'a> SkipNormWarmRollback<'a> {
fn new(kernel: &'a SkipSimplifiedLayerNormKernel) -> Result<Self> {
Ok(Self {
metadata: Some(
kernel
.metadata
.lock()
.map_err(|_| {
EpError::KernelFailed(
"cuda_ep SkipSimplifiedLayerNormalization: metadata lock poisoned"
.into(),
)
})?
.clone(),
),
scratch: Some(
kernel
.bf16_scratch
.lock()
.map_err(|_| {
EpError::KernelFailed(
"cuda_ep SkipSimplifiedLayerNormalization: scratch lock poisoned"
.into(),
)
})?
.clone(),
),
capture_safe: kernel.last_call_capture_safe.load(Ordering::Relaxed),
used_bf16_scratch: kernel.last_call_used_bf16_scratch.load(Ordering::Relaxed),
kernel,
committed: false,
})
}
fn finish(mut self, result: Result<()>) -> Result<()> {
self.committed = result.is_ok();
result
}
}
impl Drop for SkipNormWarmRollback<'_> {
fn drop(&mut self) {
if self.committed {
return;
}
if !self.kernel.runtime.is_capturing().unwrap_or(true) {
let _ = self.kernel.runtime.drain_for_unmap();
}
if let Some(metadata) = self.metadata.take()
&& let Ok(mut current) = self.kernel.metadata.lock()
{
*current = metadata;
}
if let Some(scratch) = self.scratch.take()
&& let Ok(mut current) = self.kernel.bf16_scratch.lock()
{
*current = scratch;
}
self.kernel
.last_call_capture_safe
.store(self.capture_safe, Ordering::Relaxed);
self.kernel
.last_call_used_bf16_scratch
.store(self.used_bf16_scratch, Ordering::Relaxed);
}
}
#[derive(Clone, Debug)]
struct NormBf16Scratch {
runtime: Arc<CudaRuntime>,
allocation: Option<Arc<GraphDeviceAllocation>>,
cap: usize,
}
impl NormBf16Scratch {
fn new(runtime: Arc<CudaRuntime>) -> Self {
Self {
runtime,
allocation: None,
cap: 0,
}
}
fn ensure(&mut self, bytes: usize) -> Result<(CUdeviceptr, bool)> {
if bytes <= self.cap
&& let Some(allocation) = self.allocation.as_ref()
{
return Ok((allocation.ptr(), false));
}
if self.allocation.is_some() {
self.runtime.drain_for_unmap()?;
}
let allocation = GraphDeviceAllocation::allocate(&self.runtime, bytes.max(1))?;
self.runtime
.staged_warm_cache_mutation("SkipSimplifiedLayerNorm bf16 scratch allocation")?;
let ptr = allocation.ptr();
self.allocation = Some(allocation);
self.cap = bytes;
Ok((ptr, true))
}
fn device_graph_resource(&self) -> Option<DeviceGraphResource> {
self.allocation
.as_ref()
.map(GraphDeviceAllocation::device_graph_resource)
}
}
#[derive(Clone, Debug)]
struct SkipBroadcastMetadataCache {
runtime: Arc<CudaRuntime>,
allocation: Option<Arc<GraphDeviceAllocation>>,
input_shape: Vec<usize>,
skip_shape: Vec<usize>,
}
impl SkipBroadcastMetadataCache {
fn new(runtime: Arc<CudaRuntime>) -> Self {
Self {
runtime,
allocation: None,
input_shape: Vec::new(),
skip_shape: Vec::new(),
}
}
fn reserve(&mut self, input_shape: &[usize], skip_shape: &[usize]) -> Result<CUdeviceptr> {
if self.input_shape == input_shape
&& self.skip_shape == skip_shape
&& let Some(allocation) = self.allocation.as_ref()
{
return Ok(allocation.ptr());
}
if self.runtime.is_capturing()? {
return Err(EpError::KernelFailed(
"cuda_ep SkipSimplifiedLayerNormalization: broadcast metadata shape changed \
during CUDA graph capture; warm the fixed decode shape before capture"
.into(),
));
}
let metadata = skip_broadcast_metadata(input_shape, skip_shape);
let metadata_bytes = u64_bytes(&metadata);
let allocation = GraphDeviceAllocation::allocate(&self.runtime, metadata_bytes.len())?;
unsafe { self.runtime.htod(metadata_bytes, allocation.ptr()) }?;
self.runtime.staged_warm_cache_mutation(
"SkipSimplifiedLayerNorm broadcast metadata allocation/upload",
)?;
if self.allocation.is_some() {
self.runtime.synchronize()?;
}
let ptr = allocation.ptr();
self.allocation = Some(allocation);
self.input_shape.clear();
self.input_shape.extend_from_slice(input_shape);
self.skip_shape.clear();
self.skip_shape.extend_from_slice(skip_shape);
Ok(ptr)
}
fn device_graph_resource(&self) -> Option<DeviceGraphResource> {
self.allocation
.as_ref()
.map(GraphDeviceAllocation::device_graph_resource)
}
}
impl SkipSimplifiedLayerNormKernel {
fn run_bf16_via_f32(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
let align = |bytes: usize| bytes.div_ceil(16) * 16;
let mut total = 0usize;
let mut input_offsets: Vec<Option<usize>> = Vec::with_capacity(inputs.len());
for input in inputs {
if input.dtype == DataType::BFloat16 && !input.is_absent() {
input_offsets.push(Some(total));
total += align(input.numel().max(1) * 4);
} else {
input_offsets.push(None);
}
}
let mut output_offsets: Vec<usize> = Vec::with_capacity(outputs.len());
for output in outputs.iter() {
output_offsets.push(total);
total += align(output.numel().max(1) * 4);
}
let mut arena = self
.bf16_scratch
.lock()
.expect("cuda_ep SkipSimplifiedLayerNormalization bf16 scratch mutex poisoned");
let (base, grew) = arena.ensure(total)?;
let staged = (|| -> Result<()> {
let mut f32_inputs: Vec<TensorView> = Vec::with_capacity(inputs.len());
for (input, offset) in inputs.iter().zip(input_offsets.iter()) {
match offset {
Some(off) => {
let ptr = base + *off as CUdeviceptr;
super::cast::launch_cast_raw(
&self.runtime,
cuptr(input.data_ptr::<u8>() as *const c_void),
DataType::BFloat16,
ptr,
DataType::Float32,
input.numel(),
)?;
f32_inputs.push(TensorView::new(
DevicePtr(raw_ptr(ptr) as *const c_void),
DataType::Float32,
input.shape,
input.strides,
input.device,
));
}
None => f32_inputs.push(*input),
}
}
let mut f32_outputs: Vec<TensorMut> = Vec::with_capacity(outputs.len());
for (output, off) in outputs.iter().zip(output_offsets.iter()) {
let ptr = base + *off as CUdeviceptr;
f32_outputs.push(TensorMut::new(
DevicePtrMut(raw_ptr(ptr)),
DataType::Float32,
output.shape,
output.strides,
output.device,
));
}
self.run(&f32_inputs, &mut f32_outputs)?;
for (output, off) in outputs.iter_mut().zip(output_offsets.iter()) {
if output.is_absent() || output.numel() == 0 {
continue;
}
let n = output.numel();
super::cast::launch_cast_raw(
&self.runtime,
base + *off as CUdeviceptr,
DataType::Float32,
cuptr(output.data_ptr_mut::<u8>() as *const c_void),
output.dtype,
n,
)?;
}
Ok(())
})();
if staged.is_ok() && grew && self.runtime.is_capturing().unwrap_or(true) {
self.last_call_capture_safe.store(false, Ordering::Relaxed);
}
staged
}
fn run_bf16_native(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.last_call_capture_safe.store(false, Ordering::Relaxed);
let op = "SkipSimplifiedLayerNormalization";
let input = &inputs[0];
let skip = &inputs[1];
let gamma = &inputs[2];
if skip.dtype != DataType::BFloat16 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip dtype {:?} must match input dtype BFloat16",
skip.dtype
)));
}
if outputs[0].dtype != DataType::BFloat16 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output dtype {:?} must match input dtype BFloat16",
outputs[0].dtype
)));
}
require_param_for_activation(op, "gamma", DataType::BFloat16, gamma.dtype)?;
require_contiguous(op, "input", input.is_contiguous())?;
require_contiguous(op, "skip", skip.is_contiguous())?;
require_contiguous(op, "gamma", gamma.is_contiguous())?;
require_contiguous(op, "output", outputs[0].is_contiguous())?;
self.runtime.require_nvrtc_half_headers(op)?;
let rank = input.shape.len();
if rank == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input must have rank >= 1"
)));
}
let norm_size = input.shape[rank - 1];
let num_groups: usize = input.shape[..rank - 1].iter().product();
if norm_size == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: empty hidden (last) dimension"
)));
}
if gamma.shape != [norm_size] {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: gamma shape {:?} must equal [{norm_size}]",
gamma.shape
)));
}
let broadcast =
onnx_runtime_ir::broadcast_shapes(input.shape, skip.shape).map_err(EpError::Ir)?;
if broadcast != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip shape {:?} is not broadcastable to input shape {:?}",
skip.shape, input.shape
)));
}
if outputs[0].shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output shape {:?} must equal input shape {:?}",
outputs[0].shape, input.shape
)));
}
if num_groups == 0 {
return Ok(());
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let elements = input.numel() as u64;
let groups = num_groups as u64;
elements
.saturating_mul(5)
.saturating_add(groups.saturating_mul(4))
});
let mean_ptr = optional_bf16_stat_ptr(op, "Mean", outputs, 1, num_groups)?;
let invstd_ptr = optional_bf16_stat_ptr(op, "InvStdDev", outputs, 2, num_groups)?;
let stat_is_bf16 = i32::from(
outputs
.get(1)
.is_some_and(|t| t.dtype == DataType::BFloat16)
|| outputs
.get(2)
.is_some_and(|t| t.dtype == DataType::BFloat16),
);
let sum_ptr = match outputs.get_mut(3) {
None => 0u64,
Some(t) => {
if t.dtype != DataType::BFloat16 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum dtype {:?} must match input dtype BFloat16",
t.dtype
)));
}
if t.shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum shape {:?} must equal input shape {:?}",
t.shape, input.shape
)));
}
cuptr(t.data_ptr_mut::<u8>() as *const c_void)
}
};
let (groups_u, norm_i) = (
u32::try_from(num_groups).map_err(|_| dim_overflow(op, "num_groups", num_groups))?,
i32::try_from(norm_size).map_err(|_| dim_overflow(op, "norm_size", norm_size))?,
);
let rank_i = i32::try_from(rank).map_err(|_| dim_overflow(op, "rank", rank))?;
let input_ptr = cuptr(input.data_ptr::<u8>() as *const c_void);
let skip_ptr = cuptr(skip.data_ptr::<u8>() as *const c_void);
let gamma_ptr = cuptr(gamma.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let mut metadata = self
.metadata
.lock()
.expect("cuda_ep skip normalization metadata cache poisoned");
let metadata_ptr = metadata.reserve(input.shape, skip.shape)?;
let dense_skip = i32::from(skip.numel() == input.numel());
let bias_ptr = 0u64;
let has_bias = 0i32;
let gamma_is_bf16 = i32::from(gamma.dtype == DataType::BFloat16);
let bias_is_bf16 = 0i32;
let bf16_entry = "skip_rmsnorm_bf16_warp";
onnx_runtime_ep_api::record_kernel_variant!(
"skip_rmsnorm_bf16",
"SkipSimplifiedLayerNormalization hidden={norm_size}: native byte-exact bf16 \
(bf16-rounded residual sum, rmsnorm_bf16 block-tree reduction)"
);
let func =
self.runtime
.nvrtc_function(SKIP_RMSNORM_MODULE, SKIP_RMSNORM_SRC, bf16_entry)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
let groups_i = groups_u_i32(groups_u);
builder
.arg(&input_ptr)
.arg(&skip_ptr)
.arg(&gamma_ptr)
.arg(&bias_ptr)
.arg(&y_ptr)
.arg(&sum_ptr)
.arg(&mean_ptr)
.arg(&invstd_ptr)
.arg(&metadata_ptr)
.arg(&rank_i)
.arg(&groups_i)
.arg(&norm_i)
.arg(&has_bias)
.arg(&dense_skip)
.arg(&gamma_is_bf16)
.arg(&bias_is_bf16)
.arg(&stat_is_bf16)
.arg(&self.epsilon);
let caps = self.runtime.capabilities();
let cfg = self.runtime.reduction_launch_config(
&func,
groups_u,
norm_block_threads(
norm_size,
groups_u as usize,
caps.multiprocessor_count(),
caps.max_threads_per_block(),
),
std::mem::size_of::<f32>() as u32,
)?;
unsafe { builder.launch(cfg) }.map_err(|e| driver_err("launch skip_rmsnorm_bf16", e))?;
self.last_call_capture_safe
.store(num_groups == 1, Ordering::Relaxed);
Ok(())
}
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
let rollback = SkipNormWarmRollback::new(self)?;
self.last_call_capture_safe.store(false, Ordering::Relaxed);
self.last_call_used_bf16_scratch
.store(false, Ordering::Relaxed);
let op = "SkipSimplifiedLayerNormalization";
if !(3..=4).contains(&inputs.len()) || outputs.is_empty() || outputs.len() > 4 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: expected 3-4 inputs (input, skip, gamma[, bias]) and 1-4 outputs, got {} and {}",
inputs.len(),
outputs.len()
)));
}
let input = &inputs[0];
let skip = &inputs[1];
let gamma = &inputs[2];
let bias = inputs.get(3).filter(|bias| !bias.is_absent());
if input.dtype == DataType::BFloat16 {
if self.runtime.require_nvrtc_half_headers(op).is_ok()
&& skip.numel() == input.numel()
&& bias.is_none()
{
return rollback.finish(self.run_bf16_native(inputs, outputs));
}
self.last_call_used_bf16_scratch
.store(true, Ordering::Relaxed);
return rollback.finish(self.run_bf16_via_f32(inputs, outputs));
}
require_f16_or_f32(op, "input", input.dtype)?;
let is_half = input.dtype == DataType::Float16;
if skip.dtype != input.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip dtype {:?} must match input dtype {:?}",
skip.dtype, input.dtype
)));
}
if outputs[0].dtype != input.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output dtype {:?} must match input dtype {:?}",
outputs[0].dtype, input.dtype
)));
}
if is_half {
require_f16_or_f32(op, "gamma", gamma.dtype)?;
} else {
require_f32(op, "gamma", gamma.dtype)?;
}
require_contiguous(op, "input", input.is_contiguous())?;
require_contiguous(op, "skip", skip.is_contiguous())?;
require_contiguous(op, "gamma", gamma.is_contiguous())?;
require_contiguous(op, "output", outputs[0].is_contiguous())?;
if is_half {
self.runtime.require_nvrtc_half_headers(op)?;
}
let rank = input.shape.len();
if rank == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input must have rank >= 1"
)));
}
let norm_size = input.shape[rank - 1];
let num_groups: usize = input.shape[..rank - 1].iter().product();
if norm_size == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: empty hidden (last) dimension"
)));
}
if gamma.shape != [norm_size] {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: gamma shape {:?} must equal [{norm_size}]",
gamma.shape
)));
}
let bias_ptr = optional_norm_vec_ptr(op, "bias", bias, norm_size, is_half)?;
let broadcast =
onnx_runtime_ir::broadcast_shapes(input.shape, skip.shape).map_err(EpError::Ir)?;
if broadcast != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip shape {:?} is not broadcastable to input shape {:?}",
skip.shape, input.shape
)));
}
if outputs[0].shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output shape {:?} must equal input shape {:?}",
outputs[0].shape, input.shape
)));
}
if num_groups == 0 {
return Ok(());
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let elements = input.numel() as u64;
let groups = num_groups as u64;
let mut flops = elements
.saturating_mul(5)
.saturating_add(groups.saturating_mul(4));
if bias.is_some() {
flops = flops.saturating_add(elements);
}
flops
});
let gamma_is_half = i32::from(gamma.dtype == DataType::Float16);
let bias_is_half = i32::from(bias.is_some_and(|b| b.dtype == DataType::Float16));
let (mean_ptr, invstd_ptr, stat_is_half) = if is_half {
let mean = optional_half_stat_ptr(op, "Mean", outputs, 1, num_groups)?;
let invstd = optional_half_stat_ptr(op, "InvStdDev", outputs, 2, num_groups)?;
let stat_half = i32::from(
outputs.get(1).is_some_and(|t| t.dtype == DataType::Float16)
|| outputs.get(2).is_some_and(|t| t.dtype == DataType::Float16),
);
(mean, invstd, stat_half)
} else {
let (mean, invstd) = optional_stat_ptrs(op, outputs, num_groups)?;
(mean, invstd, 0)
};
let sum_ptr = match outputs.get_mut(3) {
None => 0u64,
Some(t) => {
if t.dtype != input.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum dtype {:?} must match input dtype {:?}",
t.dtype, input.dtype
)));
}
if t.shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum shape {:?} must equal input shape {:?}",
t.shape, input.shape
)));
}
cuptr(t.data_ptr_mut::<u8>() as *const c_void)
}
};
let (groups_u, norm_i) = (
u32::try_from(num_groups).map_err(|_| dim_overflow(op, "num_groups", num_groups))?,
i32::try_from(norm_size).map_err(|_| dim_overflow(op, "norm_size", norm_size))?,
);
let rank_i = i32::try_from(rank).map_err(|_| dim_overflow(op, "rank", rank))?;
let has_bias = i32::from(bias_ptr != 0);
let input_ptr = cuptr(input.data_ptr::<u8>() as *const c_void);
let skip_ptr = cuptr(skip.data_ptr::<u8>() as *const c_void);
let gamma_ptr = cuptr(gamma.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let mut metadata = self
.metadata
.lock()
.expect("cuda_ep skip normalization metadata cache poisoned");
let metadata_ptr = metadata.reserve(input.shape, skip.shape)?;
let dense_skip = i32::from(skip.numel() == input.numel());
let selection = select_skip_rmsnorm_variant(
is_half,
dense_skip != 0,
norm_size,
bias_ptr != 0,
gamma_is_half != 0,
);
let variant_name = match selection.variant {
SkipRmsnormVariant::F32Dense => "skip_rmsnorm_f32_dense",
SkipRmsnormVariant::F32 => "skip_rmsnorm_f32",
SkipRmsnormVariant::F16Generic => "skip_rmsnorm_f16_generic",
SkipRmsnormVariant::F16WarpHalf4 => "skip_rmsnorm_f16_warp_half4",
};
let use_skip_block = matches!(selection.variant, SkipRmsnormVariant::F16WarpHalf4)
&& groups_u <= SKIP_RMSNORM_BLOCK_MAX_GROUPS
&& !skip_rmsnorm_block_disabled();
let entry = if use_skip_block {
"skip_rmsnorm_f16_block_half4_warp"
} else {
match selection.variant {
SkipRmsnormVariant::F32Dense => "skip_rmsnorm_f32_dense_warp",
SkipRmsnormVariant::F32 => "skip_rmsnorm_f32_warp",
_ => selection.entry,
}
};
onnx_runtime_ep_api::record_kernel_variant!(
variant_name,
"SkipSimplifiedLayerNormalization hidden={norm_size}: {}",
selection.reason
);
let func = self
.runtime
.nvrtc_function(SKIP_RMSNORM_MODULE, SKIP_RMSNORM_SRC, entry)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
let groups_i = groups_u_i32(groups_u);
builder
.arg(&input_ptr)
.arg(&skip_ptr)
.arg(&gamma_ptr)
.arg(&bias_ptr)
.arg(&y_ptr)
.arg(&sum_ptr)
.arg(&mean_ptr)
.arg(&invstd_ptr)
.arg(&metadata_ptr)
.arg(&rank_i)
.arg(&groups_i)
.arg(&norm_i)
.arg(&has_bias);
if is_half {
builder
.arg(&dense_skip)
.arg(&gamma_is_half)
.arg(&bias_is_half)
.arg(&stat_is_half)
.arg(&self.epsilon);
} else {
builder.arg(&self.epsilon);
}
let cfg = if use_skip_block {
self.runtime.reduction_launch_config(
&func,
groups_u,
SKIP_RMSNORM_BLOCK_THREADS,
std::mem::size_of::<f32>() as u32,
)?
} else if is_half {
LaunchConfig {
grid_dim: (groups_u, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
}
} else {
let preferred_threads = preferred_norm_block_threads(
norm_size,
self.runtime.capabilities().max_threads_per_block(),
);
self.runtime.reduction_launch_config(
&func,
groups_u,
preferred_threads,
std::mem::size_of::<f32>() as u32,
)?
};
unsafe { builder.launch(cfg) }.map_err(|e| driver_err(&format!("launch {entry}"), e))?;
self.last_call_capture_safe.store(true, Ordering::Relaxed);
rollback.finish(Ok(()))
}
}
impl Kernel for SkipSimplifiedLayerNormKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.run(inputs, outputs)
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
fn device_graph_resources(&self) -> Vec<DeviceGraphResource> {
let mut resources = Vec::with_capacity(2);
if let Ok(metadata) = self.metadata.lock()
&& let Some(resource) = metadata.device_graph_resource()
{
resources.push(resource);
}
if self.last_call_used_bf16_scratch.load(Ordering::Relaxed)
&& let Ok(scratch) = self.bf16_scratch.lock()
&& let Some(resource) = scratch.device_graph_resource()
{
resources.push(resource);
}
resources
}
fn capture_support(&self) -> onnx_runtime_ep_api::CaptureSupport {
if self.last_call_capture_safe.load(Ordering::Relaxed) {
onnx_runtime_ep_api::CaptureSupport::Supported
} else {
onnx_runtime_ep_api::CaptureSupport::unsupported(
"SkipSimplifiedLayerNormalization shape/dtype signature does not match the warmed capture signature (pre-warm the fixed M=K shape before capture)",
)
}
}
}
pub struct SkipLayerNormFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for SkipLayerNormFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let epsilon = node
.attr("epsilon")
.and_then(|a| a.as_float())
.unwrap_or(1e-5);
Ok(Box::new(SkipLayerNormKernel {
epsilon,
runtime: self.runtime.clone(),
last_call_capture_safe: AtomicBool::new(false),
}))
}
}
#[derive(Debug)]
pub struct SkipLayerNormKernel {
epsilon: f32,
runtime: Arc<CudaRuntime>,
last_call_capture_safe: AtomicBool,
}
impl SkipLayerNormKernel {
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.last_call_capture_safe.store(false, Ordering::Relaxed);
let op = "SkipLayerNormalization";
if !(3..=5).contains(&inputs.len()) || outputs.is_empty() || outputs.len() > 4 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: expected 3-5 inputs (input, skip, gamma[, beta][, bias]) \
and 1-4 outputs, got {} and {}",
inputs.len(),
outputs.len()
)));
}
let input = &inputs[0];
let skip = &inputs[1];
let gamma = &inputs[2];
let beta = inputs.get(3);
let bias = inputs.get(4);
require_float_storage(op, "input", input.dtype)?;
if skip.dtype != input.dtype || outputs[0].dtype != input.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip/output dtypes ({:?}/{:?}) must match input dtype {:?}",
skip.dtype, outputs[0].dtype, input.dtype
)));
}
require_param_for_activation(op, "gamma", input.dtype, gamma.dtype)?;
require_contiguous(op, "input", input.is_contiguous())?;
require_contiguous(op, "skip", skip.is_contiguous())?;
require_contiguous(op, "gamma", gamma.is_contiguous())?;
require_contiguous(op, "output", outputs[0].is_contiguous())?;
let rank = input.shape.len();
if rank == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input must have rank >= 1"
)));
}
if skip.shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: skip shape {:?} must equal input shape {:?}",
skip.shape, input.shape
)));
}
let norm_size = input.shape[rank - 1];
let num_groups: usize = input.shape[..rank - 1].iter().product();
if norm_size == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: empty hidden (last) dimension"
)));
}
if gamma.numel() != norm_size {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: gamma has {} elements, expected {norm_size} (hidden size)",
gamma.numel()
)));
}
let beta_ptr = optional_param_ptr(op, "beta", beta, norm_size, input.dtype)?;
let bias_ptr = optional_param_ptr(op, "bias", bias, norm_size, input.dtype)?;
if outputs[0].shape != input.shape {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output shape {:?} must equal input shape {:?}",
outputs[0].shape, input.shape
)));
}
if num_groups == 0 {
return Ok(());
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let elements = input.numel() as u64;
let groups = num_groups as u64;
let mut flops = elements
.saturating_mul(8)
.saturating_add(groups.saturating_mul(5));
if beta_ptr != 0 {
flops = flops.saturating_add(elements);
}
if bias_ptr != 0 {
flops = flops.saturating_add(elements);
}
flops
});
let input_ptr = cuptr(input.data_ptr::<u8>() as *const c_void);
let skip_ptr = cuptr(skip.data_ptr::<u8>() as *const c_void);
let gamma_ptr = cuptr(gamma.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let (mean_ptr, invstd_ptr, stat_dtype) =
optional_stat_ptrs_typed(op, outputs, num_groups, input.dtype)?;
let sum_ptr = match outputs.get_mut(3) {
None => 0u64,
Some(t) => {
if t.dtype != input.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum dtype {:?} must match input dtype {:?}",
t.dtype, input.dtype
)));
}
if t.numel() != input.numel() {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: input_skip_bias_sum has {} elements, expected {}",
t.numel(),
input.numel()
)));
}
cuptr(t.data_ptr_mut::<u8>() as *const c_void)
}
};
let (groups_u, norm_i) = (
u32::try_from(num_groups).map_err(|_| dim_overflow(op, "num_groups", num_groups))?,
i32::try_from(norm_size).map_err(|_| dim_overflow(op, "norm_size", norm_size))?,
);
let has_beta: i32 = i32::from(beta_ptr != 0);
let has_bias: i32 = i32::from(bias_ptr != 0);
let eps = self.epsilon;
let dtype = storage_kind(input.dtype);
let gamma_dtype = storage_kind(gamma.dtype);
let beta_dtype = beta.map_or(dtype, |tensor| storage_kind(tensor.dtype));
let bias_dtype = bias.map_or(dtype, |tensor| storage_kind(tensor.dtype));
let stat_dtype_kind = storage_kind(stat_dtype);
let func = self.runtime.nvrtc_function(
SKIP_LAYERNORM_MODULE,
SKIP_LAYERNORM_SRC,
"skip_layernorm",
)?;
let cfg = self.runtime.reduction_launch_config(
&func,
groups_u,
NORM_BLOCK,
std::mem::size_of::<f32>() as u32,
)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
let groups_i = groups_u_i32(groups_u);
builder
.arg(&input_ptr)
.arg(&skip_ptr)
.arg(&gamma_ptr)
.arg(&beta_ptr)
.arg(&bias_ptr)
.arg(&y_ptr)
.arg(&sum_ptr)
.arg(&mean_ptr)
.arg(&invstd_ptr)
.arg(&groups_i)
.arg(&norm_i)
.arg(&dtype)
.arg(&gamma_dtype)
.arg(&beta_dtype)
.arg(&bias_dtype)
.arg(&stat_dtype_kind)
.arg(&has_beta)
.arg(&has_bias)
.arg(&eps);
unsafe { builder.launch(cfg) }.map_err(|e| driver_err("launch skip_layernorm", e))?;
self.last_call_capture_safe
.store(num_groups == 1, Ordering::Relaxed);
Ok(())
}
}
impl Kernel for SkipLayerNormKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.run(inputs, outputs)
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
fn capture_support(&self) -> onnx_runtime_ep_api::CaptureSupport {
if self.last_call_capture_safe.load(Ordering::Relaxed) {
onnx_runtime_ep_api::CaptureSupport::Supported
} else {
onnx_runtime_ep_api::CaptureSupport::unsupported(
"SkipLayerNormalization shape/dtype signature does not match the warmed single-group capture signature",
)
}
}
}
fn groups_u_i32(groups: u32) -> i32 {
groups as i32
}
fn skip_broadcast_metadata(input: &[usize], skip: &[usize]) -> Vec<u64> {
let mut metadata = input.iter().map(|&dim| dim as u64).collect::<Vec<_>>();
let contiguous = onnx_runtime_ir::compute_contiguous_strides(skip);
let leading = input.len() - skip.len();
metadata.extend((0..input.len()).map(|axis| {
if axis < leading || skip[axis - leading] == 1 {
0
} else {
contiguous[axis - leading] as u64
}
}));
metadata
}
fn u64_bytes(values: &[u64]) -> &[u8] {
unsafe {
std::slice::from_raw_parts(values.as_ptr().cast::<u8>(), std::mem::size_of_val(values))
}
}
fn optional_stat_ptrs(
op: &str,
outputs: &mut [TensorMut],
num_groups: usize,
) -> Result<(CUdeviceptr, CUdeviceptr)> {
let mean = optional_out_ptr(op, "Mean", outputs, 1, num_groups)?;
let invstd = optional_out_ptr(op, "InvStdDev", outputs, 2, num_groups)?;
Ok((mean, invstd))
}
fn storage_kind(dtype: DataType) -> i32 {
match dtype {
DataType::Float32 => 0,
DataType::Float16 => 1,
DataType::BFloat16 => 2,
_ => unreachable!("normalization storage dtype must be validated before dispatch"),
}
}
fn optional_stat_ptrs_typed(
op: &str,
outputs: &mut [TensorMut],
num_groups: usize,
dtype: DataType,
) -> Result<(CUdeviceptr, CUdeviceptr, DataType)> {
let mut stat_dtype = DataType::Float32;
for idx in [1, 2] {
if let Some(t) = outputs.get(idx) {
if t.dtype != DataType::Float32 && t.dtype != dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: stat output {idx} dtype {:?} must be Float32 or the input dtype {dtype:?}",
t.dtype
)));
}
stat_dtype = t.dtype;
}
}
if let (Some(mean), Some(invstd)) = (outputs.get(1), outputs.get(2))
&& mean.dtype != invstd.dtype
{
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: Mean dtype {:?} and InvStdDev dtype {:?} must match",
mean.dtype, invstd.dtype
)));
}
let mean = optional_out_ptr_typed(op, "Mean", outputs, 1, num_groups, stat_dtype)?;
let invstd = optional_out_ptr_typed(op, "InvStdDev", outputs, 2, num_groups, stat_dtype)?;
Ok((mean, invstd, stat_dtype))
}
fn optional_out_ptr_typed(
op: &str,
name: &str,
outputs: &mut [TensorMut],
idx: usize,
expect: usize,
dtype: DataType,
) -> Result<CUdeviceptr> {
match outputs.get_mut(idx) {
None => Ok(0),
Some(t) => {
if t.dtype != dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} dtype {:?} must match input dtype {dtype:?}",
t.dtype
)));
}
if t.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
t.numel()
)));
}
Ok(cuptr(t.data_ptr_mut::<u8>() as *const c_void))
}
}
}
fn optional_param_ptr(
op: &str,
name: &str,
tensor: Option<&TensorView>,
expect: usize,
activation_dtype: DataType,
) -> Result<CUdeviceptr> {
match tensor {
None => Ok(0),
Some(tensor) => {
require_param_for_activation(op, name, activation_dtype, tensor.dtype)?;
require_contiguous(op, name, tensor.is_contiguous())?;
if tensor.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
tensor.numel()
)));
}
Ok(cuptr(tensor.data_ptr::<u8>() as *const c_void))
}
}
}
fn optional_out_ptr(
op: &str,
name: &str,
outputs: &mut [TensorMut],
idx: usize,
expect: usize,
) -> Result<CUdeviceptr> {
match outputs.get_mut(idx) {
None => Ok(0),
Some(t) => {
require_f32(op, name, t.dtype)?;
if t.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
t.numel()
)));
}
Ok(cuptr(t.data_ptr_mut::<u8>() as *const c_void))
}
}
}
fn optional_norm_vec_ptr(
op: &str,
name: &str,
t: Option<&TensorView>,
expect: usize,
allow_half: bool,
) -> Result<CUdeviceptr> {
match t {
None => Ok(0),
Some(v) => {
if allow_half {
require_f16_or_f32(op, name, v.dtype)?;
} else {
require_f32(op, name, v.dtype)?;
}
require_contiguous(op, name, v.is_contiguous())?;
if v.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
v.numel()
)));
}
Ok(cuptr(v.data_ptr::<u8>() as *const c_void))
}
}
}
fn optional_half_stat_ptr(
op: &str,
name: &str,
outputs: &mut [TensorMut],
idx: usize,
expect: usize,
) -> Result<CUdeviceptr> {
match outputs.get_mut(idx) {
None => Ok(0),
Some(t) => {
require_f16_or_f32(op, name, t.dtype)?;
if t.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
t.numel()
)));
}
Ok(cuptr(t.data_ptr_mut::<u8>() as *const c_void))
}
}
}
fn optional_bf16_stat_ptr(
op: &str,
name: &str,
outputs: &mut [TensorMut],
idx: usize,
expect: usize,
) -> Result<CUdeviceptr> {
match outputs.get_mut(idx) {
None => Ok(0),
Some(t) => {
if !matches!(t.dtype, DataType::BFloat16 | DataType::Float32) {
return Err(not_implemented(format!(
"{op} with {name} dtype {:?} (expected bf16 or f32)",
t.dtype
)));
}
if t.numel() != expect {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: {name} has {} elements, expected {expect}",
t.numel()
)));
}
Ok(cuptr(t.data_ptr_mut::<u8>() as *const c_void))
}
}
}
#[cfg(test)]
mod tests {
use half::{bf16, f16};
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut, ExecutionProvider};
use onnx_runtime_ir::compute_contiguous_strides;
use super::*;
use crate::CudaExecutionProvider;
#[test]
fn sources_expose_their_entry_points() {
assert!(LAYERNORM_SRC.contains("layernorm_f32"));
assert!(LAYERNORM_SRC.contains("layernorm_f16"));
assert!(LAYERNORM_SRC.contains("layernorm_bf16"));
assert!(RMSNORM_SRC.contains("rmsnorm_f32"));
assert!(RMSNORM_SRC.contains("rmsnorm_f16"));
assert!(RMSNORM_SRC.contains("rmsnorm_bf16"));
assert!(SKIP_LAYERNORM_SRC.contains("skip_layernorm"));
assert!(SKIP_LAYERNORM_SRC.contains("__half"));
assert!(SKIP_LAYERNORM_SRC.contains("__nv_bfloat16"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f32"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_bf16"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f32_dense"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f16"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f16_warp_half4"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f16_block_half4"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f16_block_half4_warp"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f32_warp"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_f32_dense_warp"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_bf16_warp"));
assert!(SKIP_RMSNORM_SRC.contains("skip_rmsnorm_warp_tail"));
}
#[test]
fn norm_dispatch_preserves_existing_entries_and_adds_bf16() {
assert_eq!(layernorm_entry(DataType::Float16), "layernorm_f16");
assert_eq!(layernorm_entry(DataType::Float32), "layernorm_f32");
assert_eq!(layernorm_entry(DataType::BFloat16), "layernorm_bf16");
assert_eq!(rmsnorm_entry(DataType::Float16), "rmsnorm_f16");
assert_eq!(rmsnorm_entry(DataType::Float32), "rmsnorm_f32");
assert_eq!(rmsnorm_entry(DataType::BFloat16), "rmsnorm_bf16");
}
#[test]
fn norm_block_width_respects_shape_and_device_limit() {
assert_eq!(preferred_norm_block_threads(3584, 1024), 256);
assert_eq!(preferred_norm_block_threads(96, 1024), 128);
assert_eq!(preferred_norm_block_threads(3584, 128), 128);
assert_eq!(preferred_norm_block_threads(17, 1024), 32);
}
#[test]
fn norm_block_widens_only_when_the_grid_cannot_fill_the_device() {
assert_eq!(norm_block_threads(6656, 1, 108, 1024), 1024);
assert_eq!(norm_block_threads(6656, 108, 108, 1024), NORM_BLOCK);
assert_eq!(norm_block_threads(6656, 4096, 108, 1024), NORM_BLOCK);
assert_eq!(norm_block_threads(6656, 107, 108, 1024), 1024);
}
#[test]
fn norm_block_never_exceeds_what_the_shape_or_the_device_can_use() {
assert_eq!(norm_block_threads(128, 32, 108, 1024), 128);
assert_eq!(norm_block_threads(17, 1, 108, 1024), 32);
assert_eq!(norm_block_threads(6656, 1, 108, 512), 512);
assert_eq!(norm_block_threads(6656, 1, 108, 768), 512);
assert_eq!(norm_block_threads(6656, 1, 0, 1024), NORM_BLOCK);
}
fn skip_rmsnorm_residuals(hidden: usize) -> (Vec<f16>, Vec<f16>) {
let residual = (0..hidden)
.map(|index| {
let input = f16::from_f32(((index * 37 % 101) as f32 - 50.0) / 31.0);
let skip = f16::from_f32(((index * 17 % 67) as f32 - 33.0) / 47.0);
let bias = f16::from_f32(((index * 11 % 29) as f32 - 14.0) / 113.0);
f16::from_f32(input.to_f32() + skip.to_f32() + bias.to_f32())
})
.collect();
let gamma = (0..hidden)
.map(|index| f16::from_f32(0.75 + (index * 13 % 41) as f32 / 64.0))
.collect();
(residual, gamma)
}
fn normalize_f16(residual: &[f16], gamma: &[f16], sum_squares: f32) -> Vec<f16> {
let inv_std = 1.0 / (sum_squares / residual.len() as f32 + 1e-5).sqrt();
residual
.iter()
.zip(gamma)
.map(|(residual, gamma)| f16::from_f32(residual.to_f32() * inv_std * gamma.to_f32()))
.collect()
}
fn previous_shared_tree_skip_rmsnorm(residual: &[f16], gamma: &[f16]) -> Vec<f16> {
let mut lanes = [0.0f32; NORM_BLOCK as usize];
for (lane, sum) in lanes.iter_mut().enumerate() {
for value in residual.iter().skip(lane).step_by(NORM_BLOCK as usize) {
let value = value.to_f32();
*sum += value * value;
}
}
let mut offset = lanes.len() / 2;
while offset > 0 {
for lane in 0..offset {
lanes[lane] += lanes[lane + offset];
}
offset /= 2;
}
normalize_f16(residual, gamma, lanes[0])
}
fn generic_warp_shuffle_skip_rmsnorm(residual: &[f16], gamma: &[f16]) -> Vec<f16> {
let mut lanes = [0.0f32; 32];
let pairs = residual.len() / 2;
for (lane, sum) in lanes.iter_mut().enumerate() {
for pair in (lane..pairs).step_by(32) {
let first = residual[pair * 2].to_f32();
let second = residual[pair * 2 + 1].to_f32();
*sum += first * first;
*sum += second * second;
}
}
if !residual.len().is_multiple_of(2) {
let tail = residual[residual.len() - 1].to_f32();
lanes[0] += tail * tail;
}
let mut offset = 16;
while offset > 0 {
let previous = lanes;
for lane in 0..(32 - offset) {
lanes[lane] += previous[lane + offset];
}
offset /= 2;
}
normalize_f16(residual, gamma, lanes[0])
}
fn half4_warp_skip_rmsnorm(residual: &[f16], gamma: &[f16]) -> Vec<f16> {
assert!(
residual
.len()
.is_multiple_of(SKIP_RMSNORM_WARP_HALF4_MULTIPLE)
);
let mut lanes = [0.0f32; 32];
let chunks_per_lane = residual.len() / SKIP_RMSNORM_WARP_HALF4_MULTIPLE;
for (lane, sum) in lanes.iter_mut().enumerate() {
let mut ss0 = 0.0f32;
let mut ss1 = 0.0f32;
let mut ss2 = 0.0f32;
let mut ss3 = 0.0f32;
for item in 0..chunks_per_lane {
let base = (lane + item * 32) * 4;
let value0 = residual[base].to_f32();
let value1 = residual[base + 1].to_f32();
let value2 = residual[base + 2].to_f32();
let value3 = residual[base + 3].to_f32();
ss0 += value0 * value0;
ss1 += value1 * value1;
ss2 += value2 * value2;
ss3 += value3 * value3;
}
*sum = (ss0 + ss1) + (ss2 + ss3);
}
let mut offset = 16;
while offset > 0 {
let previous = lanes;
for lane in 0..(32 - offset) {
lanes[lane] += previous[lane + offset];
}
offset /= 2;
}
normalize_f16(residual, gamma, lanes[0])
}
fn fixed_seven_half4_warp_skip_rmsnorm(residual: &[f16; 896], gamma: &[f16; 896]) -> Vec<f16> {
let mut lanes = [0.0f32; 32];
for (lane, sum) in lanes.iter_mut().enumerate() {
let mut ss0 = 0.0f32;
let mut ss1 = 0.0f32;
let mut ss2 = 0.0f32;
let mut ss3 = 0.0f32;
for item in 0..7 {
let base = (lane + item * 32) * 4;
let value0 = residual[base].to_f32();
let value1 = residual[base + 1].to_f32();
let value2 = residual[base + 2].to_f32();
let value3 = residual[base + 3].to_f32();
ss0 += value0 * value0;
ss1 += value1 * value1;
ss2 += value2 * value2;
ss3 += value3 * value3;
}
*sum = (ss0 + ss1) + (ss2 + ss3);
}
let mut offset = 16;
while offset > 0 {
let previous = lanes;
for lane in 0..(32 - offset) {
lanes[lane] += previous[lane + offset];
}
offset /= 2;
}
normalize_f16(residual, gamma, lanes[0])
}
#[test]
fn warp_shuffle_skip_rmsnorm_matches_shared_tree_for_hidden_and_tail_sizes() {
for hidden in [896, 1024, 2048, 4096, 5120] {
let (residual, gamma) = skip_rmsnorm_residuals(hidden);
let previous = previous_shared_tree_skip_rmsnorm(&residual, &gamma);
let warp = half4_warp_skip_rmsnorm(&residual, &gamma);
let max_error = previous
.iter()
.zip(&warp)
.map(|(previous, warp)| (previous.to_f32() - warp.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 2.0e-3,
"hidden={hidden} shared-tree/warp max fp16 error {max_error}"
);
}
let hidden = 900;
let (residual, gamma) = skip_rmsnorm_residuals(hidden);
let previous = previous_shared_tree_skip_rmsnorm(&residual, &gamma);
let generic = generic_warp_shuffle_skip_rmsnorm(&residual, &gamma);
let max_error = previous
.iter()
.zip(&generic)
.map(|(previous, generic)| (previous.to_f32() - generic.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 2.0e-3,
"hidden={hidden} shared-tree/generic max fp16 error {max_error}"
);
}
#[test]
fn fp16_skip_rmsnorm_warp_selection_is_structural() {
for hidden in [128, 256, 512, 896, 1024, 2048, 4096, 5120] {
let selection = select_skip_rmsnorm_variant(true, true, hidden, false, true);
assert_eq!(
selection.variant,
SkipRmsnormVariant::F16WarpHalf4,
"hidden={hidden}: {}",
selection.reason
);
assert!(selection.reason.contains("hidden%128==0"));
}
let tail = select_skip_rmsnorm_variant(true, true, 900, false, true);
assert_eq!(tail.variant, SkipRmsnormVariant::F16Generic);
assert!(tail.reason.contains("hidden%128==0"));
}
#[test]
fn generalized_half4_warp_is_bit_identical_for_hidden_896() {
let (residual, gamma) = skip_rmsnorm_residuals(896);
let residual: [f16; 896] = residual.try_into().unwrap();
let gamma: [f16; 896] = gamma.try_into().unwrap();
let fixed = fixed_seven_half4_warp_skip_rmsnorm(&residual, &gamma);
let generalized = half4_warp_skip_rmsnorm(&residual, &gamma);
assert_eq!(
fixed
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
generalized
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
}
fn f16_bytes(values: &[f16]) -> &[u8] {
unsafe {
std::slice::from_raw_parts(values.as_ptr().cast::<u8>(), std::mem::size_of_val(values))
}
}
fn run_fp16_skip_rmsnorm_gpu(
ep: &CudaExecutionProvider,
hidden: usize,
) -> (Vec<f16>, Vec<f16>, Vec<f16>) {
let shape = [1, hidden];
let strides = compute_contiguous_strides(&shape);
let gamma_shape = [hidden];
let gamma_strides = compute_contiguous_strides(&gamma_shape);
let input = (0..hidden)
.map(|index| f16::from_f32(((index * 37 % 101) as f32 - 50.0) / 31.0))
.collect::<Vec<_>>();
let skip = (0..hidden)
.map(|index| f16::from_f32(((index * 17 % 67) as f32 - 33.0) / 47.0))
.collect::<Vec<_>>();
let gamma = (0..hidden)
.map(|index| f16::from_f32(0.75 + (index * 13 % 41) as f32 / 64.0))
.collect::<Vec<_>>();
let residual = input
.iter()
.zip(&skip)
.map(|(input, skip)| f16::from_f32(input.to_f32() + skip.to_f32()))
.collect::<Vec<_>>();
let input_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let skip_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let gamma_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let mut output_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let runtime = ep.runtime();
unsafe {
runtime
.htod(f16_bytes(&input), cuptr(input_buffer.as_ptr()))
.unwrap();
runtime
.htod(f16_bytes(&skip), cuptr(skip_buffer.as_ptr()))
.unwrap();
runtime
.htod(f16_bytes(&gamma), cuptr(gamma_buffer.as_ptr()))
.unwrap();
}
{
let inputs = [
TensorView::new(
DevicePtr(input_buffer.as_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(skip_buffer.as_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(gamma_buffer.as_ptr()),
DataType::Float16,
&gamma_shape,
&gamma_strides,
ep.device_id(),
),
];
let output = TensorMut::new(
DevicePtrMut(output_buffer.as_mut_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
);
let kernel = SkipSimplifiedLayerNormKernel {
epsilon: 1e-5,
runtime: runtime.clone(),
metadata: Mutex::new(SkipBroadcastMetadataCache::new(runtime.clone())),
bf16_scratch: Mutex::new(NormBf16Scratch::new(runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
last_call_used_bf16_scratch: AtomicBool::new(false),
};
kernel.run(&inputs, &mut [output]).unwrap();
}
let mut output_bytes = vec![0u8; hidden * std::mem::size_of::<f16>()];
unsafe {
runtime
.dtoh(&mut output_bytes, cuptr(output_buffer.as_ptr()))
.unwrap();
}
let output = output_bytes
.chunks_exact(2)
.map(|raw| f16::from_bits(u16::from_ne_bytes(raw.try_into().unwrap())))
.collect();
ep.deallocate(input_buffer).unwrap();
ep.deallocate(skip_buffer).unwrap();
ep.deallocate(gamma_buffer).unwrap();
ep.deallocate(output_buffer).unwrap();
(output, residual, gamma)
}
fn f32_bytes(values: &[f32]) -> &[u8] {
unsafe {
std::slice::from_raw_parts(values.as_ptr().cast::<u8>(), std::mem::size_of_val(values))
}
}
#[test]
fn fp32_dense_skip_rmsnorm_matches_reference_and_optional_outputs() {
let Ok(ep) = CudaExecutionProvider::new(0) else {
eprintln!("skipping fp32 dense skip RMSNorm test: CUDA unavailable");
return;
};
let hidden = 3584usize;
let shape = [1, hidden];
let strides = compute_contiguous_strides(&shape);
let gamma_shape = [hidden];
let gamma_strides = compute_contiguous_strides(&gamma_shape);
let stat_shape = [1usize];
let stat_strides = [1i64];
let input = (0..hidden)
.map(|index| ((index * 37 % 101) as f32 - 50.0) / 31.0)
.collect::<Vec<_>>();
let skip = (0..hidden)
.map(|index| ((index * 17 % 67) as f32 - 33.0) / 47.0)
.collect::<Vec<_>>();
let gamma = (0..hidden)
.map(|index| 0.75 + (index * 13 % 41) as f32 / 64.0)
.collect::<Vec<_>>();
let residual = input
.iter()
.zip(&skip)
.map(|(input, skip)| input + skip)
.collect::<Vec<_>>();
let sum_squares = residual.iter().fold(0.0f64, |sum, value| {
sum + f64::from(*value) * f64::from(*value)
});
let inverse_standard_deviation =
(sum_squares / hidden as f64 + 1e-5f64).sqrt().recip() as f32;
let allocate = |elements: usize| {
ep.allocate(elements * std::mem::size_of::<f32>(), 256)
.unwrap()
};
let input_buffer = allocate(hidden);
let skip_buffer = allocate(hidden);
let gamma_buffer = allocate(hidden);
let mut output_buffer = allocate(hidden);
let mut mean_buffer = allocate(1);
let mut inverse_standard_deviation_buffer = allocate(1);
let mut sum_buffer = allocate(hidden);
let runtime = ep.runtime();
unsafe {
runtime
.htod(f32_bytes(&input), cuptr(input_buffer.as_ptr()))
.unwrap();
runtime
.htod(f32_bytes(&skip), cuptr(skip_buffer.as_ptr()))
.unwrap();
runtime
.htod(f32_bytes(&gamma), cuptr(gamma_buffer.as_ptr()))
.unwrap();
}
let inputs = [
TensorView::new(
DevicePtr(input_buffer.as_ptr()),
DataType::Float32,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(skip_buffer.as_ptr()),
DataType::Float32,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(gamma_buffer.as_ptr()),
DataType::Float32,
&gamma_shape,
&gamma_strides,
ep.device_id(),
),
];
let mut outputs = [
TensorMut::new(
DevicePtrMut(output_buffer.as_mut_ptr()),
DataType::Float32,
&shape,
&strides,
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(mean_buffer.as_mut_ptr()),
DataType::Float32,
&stat_shape,
&stat_strides,
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(inverse_standard_deviation_buffer.as_mut_ptr()),
DataType::Float32,
&stat_shape,
&stat_strides,
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(sum_buffer.as_mut_ptr()),
DataType::Float32,
&shape,
&strides,
ep.device_id(),
),
];
let kernel = SkipSimplifiedLayerNormKernel {
epsilon: 1e-5,
runtime: runtime.clone(),
metadata: Mutex::new(SkipBroadcastMetadataCache::new(runtime.clone())),
bf16_scratch: Mutex::new(NormBf16Scratch::new(runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
last_call_used_bf16_scratch: AtomicBool::new(false),
};
kernel.run(&inputs, &mut outputs).unwrap();
runtime.synchronize().unwrap();
let mut output = vec![0.0f32; hidden];
let mut mean = [f32::NAN];
let mut got_inverse_standard_deviation = [f32::NAN];
let mut sum = vec![0.0f32; hidden];
unsafe {
runtime
.dtoh(f32_bytes_mut(&mut output), cuptr(output_buffer.as_ptr()))
.unwrap();
runtime
.dtoh(f32_bytes_mut(&mut mean), cuptr(mean_buffer.as_ptr()))
.unwrap();
runtime
.dtoh(
f32_bytes_mut(&mut got_inverse_standard_deviation),
cuptr(inverse_standard_deviation_buffer.as_ptr()),
)
.unwrap();
runtime
.dtoh(f32_bytes_mut(&mut sum), cuptr(sum_buffer.as_ptr()))
.unwrap();
}
assert_eq!(mean[0], 0.0);
assert_eq!(sum, residual);
assert!((got_inverse_standard_deviation[0] - inverse_standard_deviation).abs() < 2e-6);
for index in 0..hidden {
let reference = residual[index] * inverse_standard_deviation * gamma[index];
assert!(
(output[index] - reference).abs() < 2e-5,
"output mismatch at {index}: {} vs {reference}",
output[index]
);
}
for buffer in [
input_buffer,
skip_buffer,
gamma_buffer,
output_buffer,
mean_buffer,
inverse_standard_deviation_buffer,
sum_buffer,
] {
ep.deallocate(buffer).unwrap();
}
}
fn assert_bf16_skip_byte_exact(gamma_bf16: bool) {
let Ok(ep) = CudaExecutionProvider::new(0) else {
eprintln!("skipping bf16 skip RMSNorm byte-exactness test: CUDA unavailable");
return;
};
let hidden = 6656usize; let shape = [1usize, hidden];
let strides = compute_contiguous_strides(&shape);
let gamma_shape = [hidden];
let gamma_strides = compute_contiguous_strides(&gamma_shape);
let runtime = ep.runtime();
let input: Vec<bf16> = (0..hidden)
.map(|i| bf16::from_f32(((i * 37 % 101) as f32 - 50.0) / 31.0))
.collect();
let skip: Vec<bf16> = (0..hidden)
.map(|i| bf16::from_f32(((i * 17 % 67) as f32 - 33.0) / 47.0))
.collect();
let gamma_f32: Vec<f32> = (0..hidden)
.map(|i| 0.75 + (i * 13 % 41) as f32 / 64.0)
.collect();
let residual_ref: Vec<bf16> = input
.iter()
.zip(&skip)
.map(|(a, b)| bf16::from_f32(a.to_f32() + b.to_f32()))
.collect();
let bytes = std::mem::size_of_val(input.as_slice());
let gamma_bytes = if gamma_bf16 {
hidden * std::mem::size_of::<bf16>()
} else {
hidden * std::mem::size_of::<f32>()
};
let input_buffer = ep.allocate(bytes, 256).unwrap();
let skip_buffer = ep.allocate(bytes, 256).unwrap();
let residual_buffer = ep.allocate(bytes, 256).unwrap();
let gamma_buffer = ep.allocate(gamma_bytes, 256).unwrap();
let mut ref_out_buffer = ep.allocate(bytes, 256).unwrap();
let mut y_buffer = ep.allocate(bytes, 256).unwrap();
let mut sum_buffer = ep.allocate(bytes, 256).unwrap();
let mut mean_buffer = ep.allocate(std::mem::size_of::<f32>(), 256).unwrap();
let mut invstd_buffer = ep.allocate(std::mem::size_of::<f32>(), 256).unwrap();
let gamma_bf16_vec: Vec<bf16> = gamma_f32.iter().map(|v| bf16::from_f32(*v)).collect();
unsafe {
runtime
.htod(bf16_bytes(&input), cuptr(input_buffer.as_ptr()))
.unwrap();
runtime
.htod(bf16_bytes(&skip), cuptr(skip_buffer.as_ptr()))
.unwrap();
runtime
.htod(bf16_bytes(&residual_ref), cuptr(residual_buffer.as_ptr()))
.unwrap();
if gamma_bf16 {
runtime
.htod(bf16_bytes(&gamma_bf16_vec), cuptr(gamma_buffer.as_ptr()))
.unwrap();
} else {
runtime
.htod(f32_bytes(&gamma_f32), cuptr(gamma_buffer.as_ptr()))
.unwrap();
}
}
let gamma_dtype = if gamma_bf16 {
DataType::BFloat16
} else {
DataType::Float32
};
let gamma_view = TensorView::new(
DevicePtr(gamma_buffer.as_ptr()),
gamma_dtype,
&gamma_shape,
&gamma_strides,
ep.device_id(),
);
RmsNormKernel {
axis: -1,
epsilon: 1e-5,
runtime: runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}
.run(
&[
TensorView::new(
DevicePtr(residual_buffer.as_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
),
gamma_view,
],
&mut [TensorMut::new(
DevicePtrMut(ref_out_buffer.as_mut_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
)],
)
.unwrap();
let kernel = SkipSimplifiedLayerNormKernel {
epsilon: 1e-5,
runtime: runtime.clone(),
metadata: Mutex::new(SkipBroadcastMetadataCache::new(runtime.clone())),
bf16_scratch: Mutex::new(NormBf16Scratch::new(runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
last_call_used_bf16_scratch: AtomicBool::new(false),
};
kernel
.run(
&[
TensorView::new(
DevicePtr(input_buffer.as_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(skip_buffer.as_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
),
gamma_view,
],
&mut [
TensorMut::new(
DevicePtrMut(y_buffer.as_mut_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(mean_buffer.as_mut_ptr()),
DataType::Float32,
&[1usize],
&[1i64],
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(invstd_buffer.as_mut_ptr()),
DataType::Float32,
&[1usize],
&[1i64],
ep.device_id(),
),
TensorMut::new(
DevicePtrMut(sum_buffer.as_mut_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
),
],
)
.unwrap();
runtime.synchronize().unwrap();
assert!(
kernel.last_call_capture_safe.load(Ordering::Relaxed),
"native bf16 skip must stay capture-safe at num_groups==1"
);
let read_bf16 = |buffer: &onnx_runtime_ep_api::DeviceBuffer| -> Vec<u16> {
let mut raw = vec![0u8; bytes];
unsafe {
runtime.dtoh(&mut raw, cuptr(buffer.as_ptr())).unwrap();
}
raw.chunks_exact(2)
.map(|c| u16::from_ne_bytes(c.try_into().unwrap()))
.collect()
};
let ref_bits = read_bf16(&ref_out_buffer);
let y_bits = read_bf16(&y_buffer);
let sum_bits = read_bf16(&sum_buffer);
let residual_ref_bits: Vec<u16> = residual_ref.iter().map(|v| v.to_bits()).collect();
assert_eq!(
sum_bits, residual_ref_bits,
"fused residual sum must be bit-identical to standalone Add(bf16) (gamma_bf16={gamma_bf16})"
);
assert_eq!(
y_bits, ref_bits,
"fused norm output must be bit-identical to standalone Add(bf16)+rmsnorm_bf16 (gamma_bf16={gamma_bf16})"
);
for buffer in [
input_buffer,
skip_buffer,
residual_buffer,
gamma_buffer,
ref_out_buffer,
y_buffer,
sum_buffer,
mean_buffer,
invstd_buffer,
] {
ep.deallocate(buffer).unwrap();
}
}
#[test]
fn bf16_native_skip_rmsnorm_is_byte_exact_with_bf16_gamma() {
assert_bf16_skip_byte_exact(true);
}
#[test]
fn bf16_native_skip_rmsnorm_is_byte_exact_with_f32_gamma() {
assert_bf16_skip_byte_exact(false);
}
fn f32_bytes_mut(values: &mut [f32]) -> &mut [u8] {
unsafe {
std::slice::from_raw_parts_mut(
values.as_mut_ptr().cast::<u8>(),
std::mem::size_of_val(values),
)
}
}
fn bf16_bytes(values: &[bf16]) -> &[u8] {
unsafe {
std::slice::from_raw_parts(values.as_ptr().cast::<u8>(), std::mem::size_of_val(values))
}
}
fn run_bf16_norm_gpu(ep: &CudaExecutionProvider, layer_norm: bool) -> Vec<bf16> {
let shape = [2, 5];
let strides = compute_contiguous_strides(&shape);
let param_shape = [shape[1]];
let param_strides = compute_contiguous_strides(¶m_shape);
let input = (0..shape.iter().product())
.map(|index| bf16::from_f32((index as f32 - 4.5) / 3.0))
.collect::<Vec<_>>();
let scale = (0..shape[1])
.map(|index| bf16::from_f32(0.75 + index as f32 * 0.125))
.collect::<Vec<_>>();
let bias = (0..shape[1])
.map(|index| bf16::from_f32((index as f32 - 2.0) / 16.0))
.collect::<Vec<_>>();
let bytes = std::mem::size_of_val(input.as_slice());
let param_bytes = std::mem::size_of_val(scale.as_slice());
let input_buffer = ep.allocate(bytes, 256).unwrap();
let scale_buffer = ep.allocate(param_bytes, 256).unwrap();
let bias_buffer = ep.allocate(param_bytes, 256).unwrap();
let mut output_buffer = ep.allocate(bytes, 256).unwrap();
let runtime = ep.runtime();
unsafe {
runtime
.htod(bf16_bytes(&input), cuptr(input_buffer.as_ptr()))
.unwrap();
runtime
.htod(bf16_bytes(&scale), cuptr(scale_buffer.as_ptr()))
.unwrap();
runtime
.htod(bf16_bytes(&bias), cuptr(bias_buffer.as_ptr()))
.unwrap();
}
let x = TensorView::new(
DevicePtr(input_buffer.as_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
);
let scale_view = TensorView::new(
DevicePtr(scale_buffer.as_ptr()),
DataType::BFloat16,
¶m_shape,
¶m_strides,
ep.device_id(),
);
let output = TensorMut::new(
DevicePtrMut(output_buffer.as_mut_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
);
if layer_norm {
let bias_view = TensorView::new(
DevicePtr(bias_buffer.as_ptr()),
DataType::BFloat16,
¶m_shape,
¶m_strides,
ep.device_id(),
);
LayerNormKernel {
axis: -1,
epsilon: 1e-5,
runtime: runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}
.run(&[x, scale_view, bias_view], &mut [output])
.unwrap();
} else {
RmsNormKernel {
axis: -1,
epsilon: 1e-5,
runtime: runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}
.run(&[x, scale_view], &mut [output])
.unwrap();
}
let mut output_bytes = vec![0u8; bytes];
unsafe {
runtime
.dtoh(&mut output_bytes, cuptr(output_buffer.as_ptr()))
.unwrap();
}
let output = output_bytes
.chunks_exact(2)
.map(|raw| bf16::from_bits(u16::from_ne_bytes(raw.try_into().unwrap())))
.collect();
ep.deallocate(input_buffer).unwrap();
ep.deallocate(scale_buffer).unwrap();
ep.deallocate(bias_buffer).unwrap();
ep.deallocate(output_buffer).unwrap();
output
}
fn bf16_norm_reference(layer_norm: bool) -> Vec<bf16> {
let groups = 2;
let hidden = 5;
let input = (0..groups * hidden)
.map(|index| bf16::from_f32((index as f32 - 4.5) / 3.0).to_f32())
.collect::<Vec<_>>();
let scale = (0..hidden)
.map(|index| bf16::from_f32(0.75 + index as f32 * 0.125).to_f32())
.collect::<Vec<_>>();
let bias = (0..hidden)
.map(|index| bf16::from_f32((index as f32 - 2.0) / 16.0).to_f32())
.collect::<Vec<_>>();
let mut output = Vec::with_capacity(input.len());
for group in input.chunks_exact(hidden) {
if layer_norm {
let mean = group.iter().sum::<f32>() / hidden as f32;
let variance = group
.iter()
.map(|value| (value - mean) * (value - mean))
.sum::<f32>()
/ hidden as f32;
let inv_std = 1.0 / (variance + 1e-5).sqrt();
output.extend((0..hidden).map(|index| {
bf16::from_f32((group[index] - mean) * inv_std * scale[index] + bias[index])
}));
} else {
let mean_square =
group.iter().map(|value| value * value).sum::<f32>() / hidden as f32;
let inv_std = 1.0 / (mean_square + 1e-5).sqrt();
output.extend(
(0..hidden).map(|index| bf16::from_f32(group[index] * inv_std * scale[index])),
);
}
}
output
}
fn run_skip_rmsnorm_gpu_f32_gamma(
ep: &CudaExecutionProvider,
hidden: usize,
) -> (Vec<f16>, Vec<f16>, Vec<f32>) {
let shape = [1, hidden];
let strides = compute_contiguous_strides(&shape);
let gamma_shape = [hidden];
let gamma_strides = compute_contiguous_strides(&gamma_shape);
let input = (0..hidden)
.map(|index| f16::from_f32(((index * 37 % 101) as f32 - 50.0) / 31.0))
.collect::<Vec<_>>();
let skip = (0..hidden)
.map(|index| f16::from_f32(((index * 17 % 67) as f32 - 33.0) / 47.0))
.collect::<Vec<_>>();
let gamma = (0..hidden)
.map(|index| 0.7501 + (index % 41) as f32 * 0.012_345)
.collect::<Vec<f32>>();
let residual = input
.iter()
.zip(&skip)
.map(|(input, skip)| f16::from_f32(input.to_f32() + skip.to_f32()))
.collect::<Vec<_>>();
let input_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let skip_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let gamma_buffer = ep
.allocate(hidden * std::mem::size_of::<f32>(), 256)
.unwrap();
let mut output_buffer = ep
.allocate(hidden * std::mem::size_of::<f16>(), 256)
.unwrap();
let runtime = ep.runtime();
unsafe {
runtime
.htod(f16_bytes(&input), cuptr(input_buffer.as_ptr()))
.unwrap();
runtime
.htod(f16_bytes(&skip), cuptr(skip_buffer.as_ptr()))
.unwrap();
runtime
.htod(f32_bytes(&gamma), cuptr(gamma_buffer.as_ptr()))
.unwrap();
}
{
let inputs = [
TensorView::new(
DevicePtr(input_buffer.as_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(skip_buffer.as_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
),
TensorView::new(
DevicePtr(gamma_buffer.as_ptr()),
DataType::Float32,
&gamma_shape,
&gamma_strides,
ep.device_id(),
),
];
let output = TensorMut::new(
DevicePtrMut(output_buffer.as_mut_ptr()),
DataType::Float16,
&shape,
&strides,
ep.device_id(),
);
let kernel = SkipSimplifiedLayerNormKernel {
epsilon: 1e-5,
runtime: runtime.clone(),
metadata: Mutex::new(SkipBroadcastMetadataCache::new(runtime.clone())),
bf16_scratch: Mutex::new(NormBf16Scratch::new(runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
last_call_used_bf16_scratch: AtomicBool::new(false),
};
kernel.run(&inputs, &mut [output]).unwrap();
}
let mut output_bytes = vec![0u8; hidden * std::mem::size_of::<f16>()];
unsafe {
runtime
.dtoh(&mut output_bytes, cuptr(output_buffer.as_ptr()))
.unwrap();
}
let output = output_bytes
.chunks_exact(2)
.map(|raw| f16::from_bits(u16::from_ne_bytes(raw.try_into().unwrap())))
.collect();
ep.deallocate(input_buffer).unwrap();
ep.deallocate(skip_buffer).unwrap();
ep.deallocate(gamma_buffer).unwrap();
ep.deallocate(output_buffer).unwrap();
(output, residual, gamma)
}
fn half4_warp_skip_rmsnorm_f32_gamma(residual: &[f16], gamma: &[f32]) -> Vec<f16> {
let mut lanes = [0.0f32; 32];
let chunks_per_lane = residual.len() / SKIP_RMSNORM_WARP_HALF4_MULTIPLE;
for (lane, sum) in lanes.iter_mut().enumerate() {
let (mut ss0, mut ss1, mut ss2, mut ss3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
for item in 0..chunks_per_lane {
let base = (lane + item * 32) * 4;
let v0 = residual[base].to_f32();
let v1 = residual[base + 1].to_f32();
let v2 = residual[base + 2].to_f32();
let v3 = residual[base + 3].to_f32();
ss0 += v0 * v0;
ss1 += v1 * v1;
ss2 += v2 * v2;
ss3 += v3 * v3;
}
*sum = (ss0 + ss1) + (ss2 + ss3);
}
let mut offset = 16;
while offset > 0 {
let previous = lanes;
for lane in 0..(32 - offset) {
lanes[lane] += previous[lane + offset];
}
offset /= 2;
}
let inv_std = 1.0 / (lanes[0] / residual.len() as f32 + 1e-5).sqrt();
residual
.iter()
.zip(gamma)
.map(|(r, g)| f16::from_f32(r.to_f32() * inv_std * g))
.collect()
}
fn f16_accumulation_skip_rmsnorm_f32_gamma(residual: &[f16], gamma: &[f32]) -> Vec<f16> {
let mut ss = f16::from_f32(0.0);
for r in residual {
ss = f16::from_f32(ss.to_f32() + (r.to_f32() * r.to_f32()));
}
let inv_std = 1.0 / (ss.to_f32() / residual.len() as f32 + 1e-5).sqrt();
residual
.iter()
.zip(gamma)
.map(|(r, g)| f16::from_f32(r.to_f32() * inv_std * g))
.collect()
}
#[test]
fn f32_gamma_warp_selection_is_structural_and_gated() {
for hidden in [128usize, 3072, 4096] {
let sel = select_skip_rmsnorm_variant(true, true, hidden, false, false);
assert_eq!(
sel.variant,
SkipRmsnormVariant::F16WarpHalf4,
"hidden={hidden} fp32-gamma should take warp_half4"
);
assert!(sel.reason.contains("gamma=fp32"));
}
let half = select_skip_rmsnorm_variant(true, true, 3072, false, true);
assert_eq!(half.variant, SkipRmsnormVariant::F16WarpHalf4);
assert!(half.reason.contains("gamma=fp16"));
}
#[test]
fn fp32_gamma_gpu_skip_rmsnorm_matches_warp_reference_at_phi_and_qwen_dims() {
let ep = match CudaExecutionProvider::new_default() {
Ok(ep) => ep,
Err(error) => {
eprintln!("skip: no CUDA GPU/runtime available ({error})");
return;
}
};
for hidden in [128usize, 3072] {
let (output, residual, gamma) = run_skip_rmsnorm_gpu_f32_gamma(&ep, hidden);
let reference = half4_warp_skip_rmsnorm_f32_gamma(&residual, &gamma);
let max_error = output
.iter()
.zip(&reference)
.map(|(got, want)| (got.to_f32() - want.to_f32()).abs())
.fold(0.0f32, f32::max);
let parity_tol = 1.0e-3f32;
assert!(
max_error <= parity_tol,
"hidden={hidden} fp32-gamma warp GPU max error {max_error}"
);
let broken = f16_accumulation_skip_rmsnorm_f32_gamma(&residual, &gamma);
let broken_error = reference
.iter()
.zip(&broken)
.map(|(want, bad)| (want.to_f32() - bad.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
broken_error > parity_tol,
"hidden={hidden} fp16-accumulation guard too weak ({broken_error}); \
test cannot detect a broken accumulation dtype"
);
}
}
#[test]
fn fp16_skip_rmsnorm_gpu_is_generic_across_structural_hidden_sizes() {
let ep = match CudaExecutionProvider::new_default() {
Ok(ep) => ep,
Err(error) => {
eprintln!("skip: no CUDA GPU/runtime available ({error})");
return;
}
};
for hidden in [896, 1024, 2048, 4096, 5120] {
let selection = select_skip_rmsnorm_variant(true, true, hidden, false, true);
assert_eq!(selection.variant, SkipRmsnormVariant::F16WarpHalf4);
let (output, residual, gamma) = run_fp16_skip_rmsnorm_gpu(&ep, hidden);
let reference = previous_shared_tree_skip_rmsnorm(&residual, &gamma);
let max_error = output
.iter()
.zip(&reference)
.map(|(output, reference)| (output.to_f32() - reference.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 2.0e-3,
"hidden={hidden} GPU half4/shared-tree max fp16 error {max_error}"
);
}
let hidden = 900;
let selection = select_skip_rmsnorm_variant(true, true, hidden, false, true);
assert_eq!(selection.variant, SkipRmsnormVariant::F16Generic);
let (output, residual, gamma) = run_fp16_skip_rmsnorm_gpu(&ep, hidden);
let reference = previous_shared_tree_skip_rmsnorm(&residual, &gamma);
let max_error = output
.iter()
.zip(&reference)
.map(|(output, reference)| (output.to_f32() - reference.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 2.0e-3,
"hidden={hidden} GPU generic/shared-tree max fp16 error {max_error}"
);
}
#[test]
fn fp16_skip_rmsnorm_source_uses_one_warp_without_shared_reduction() {
let warp_start = SKIP_RMSNORM_SRC
.find("extern \"C\" __global__ void skip_rmsnorm_f16_warp_half4")
.unwrap();
let block_start = SKIP_RMSNORM_SRC
.find("__device__ __forceinline__ void skip_rmsnorm_f16_block_half4_tpl")
.unwrap();
assert!(block_start > warp_start);
let warp_body = &SKIP_RMSNORM_SRC[warp_start..block_start];
assert!(warp_body.contains("__half2"));
assert!(warp_body.contains("__shfl_down_sync"));
assert!(!warp_body.contains("extern __shared__"));
assert!(!warp_body.contains("__syncthreads"));
let block_body = &SKIP_RMSNORM_SRC[block_start..];
assert!(block_body.contains("extern __shared__ float red[]"));
assert!(block_body.contains("__syncthreads"));
}
#[test]
fn require_f32_names_op_and_dtype() {
let e = require_f32("LayerNormalization", "Scale", DataType::Float16).unwrap_err();
let msg = format!("{e}");
assert!(msg.contains("LayerNormalization"), "{msg}");
assert!(msg.contains("Float16"), "{msg}");
}
#[test]
fn require_contiguous_is_actionable() {
let e = require_contiguous("RMSNormalization", "X", false).unwrap_err();
let msg = format!("{e}");
assert!(msg.contains("non-contiguous"), "{msg}");
assert!(msg.contains("materialise"), "{msg}");
}
#[test]
fn norm_group_split_matches_axis() {
let shape = [4usize, 8];
let axis = resolve_axis("LayerNormalization", -1, shape.len()).unwrap();
let norm_size: usize = shape[axis..].iter().product();
let groups: usize = shape[..axis].iter().product();
assert_eq!((groups, norm_size), (4, 8));
}
#[test]
fn bf16_layernorm_and_rmsnorm_match_fp32_references() {
let ep = match CudaExecutionProvider::new_default() {
Ok(ep) => ep,
Err(error) => {
eprintln!("skip: no CUDA GPU/runtime available ({error})");
return;
}
};
for layer_norm in [false, true] {
let output = run_bf16_norm_gpu(&ep, layer_norm);
let reference = bf16_norm_reference(layer_norm);
let max_error = output
.iter()
.zip(&reference)
.map(|(actual, expected)| (actual.to_f32() - expected.to_f32()).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 0.015625,
"{} max bf16 error {max_error}",
if layer_norm {
"LayerNormalization"
} else {
"RMSNormalization"
}
);
}
}
#[test]
fn bf16_rmsnorm_tree_reduction_matches_f64_oracle_at_muse_glimmer_width() {
let ep = match CudaExecutionProvider::new_default() {
Ok(ep) => ep,
Err(error) => {
eprintln!("skip: no CUDA GPU/runtime available ({error})");
return;
}
};
let hidden = 6656usize;
let epsilon = 1e-6f32;
let x_f32: Vec<f32> = (0..hidden)
.map(|i| {
let t = (i as f32) * 0.017_f32;
let sign = if i % 2 == 0 { 1.0 } else { -1.0 };
sign * (0.001 + (t.sin() * t.cos()).abs() * 3.0)
})
.collect();
let x_bf16: Vec<bf16> = x_f32.iter().map(|v| bf16::from_f32(*v)).collect();
let scale_f32: Vec<f32> = (0..hidden).map(|i| 1.0 + ((i % 7) as f32) * 0.03).collect();
let shape = [1usize, hidden];
let strides = compute_contiguous_strides(&shape);
let param_shape = [hidden];
let param_strides = compute_contiguous_strides(¶m_shape);
let x_buffer = ep
.allocate(std::mem::size_of_val(x_bf16.as_slice()), 256)
.unwrap();
let scale_buffer = ep
.allocate(std::mem::size_of_val(scale_f32.as_slice()), 256)
.unwrap();
let mut out_buffer = ep
.allocate(std::mem::size_of_val(x_bf16.as_slice()), 256)
.unwrap();
let runtime = ep.runtime();
unsafe {
runtime
.htod(bf16_bytes(&x_bf16), cuptr(x_buffer.as_ptr()))
.unwrap();
runtime
.htod(f32_bytes(&scale_f32), cuptr(scale_buffer.as_ptr()))
.unwrap();
}
let x = TensorView::new(
DevicePtr(x_buffer.as_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
);
let scale_view = TensorView::new(
DevicePtr(scale_buffer.as_ptr()),
DataType::Float32,
¶m_shape,
¶m_strides,
ep.device_id(),
);
let output = TensorMut::new(
DevicePtrMut(out_buffer.as_mut_ptr()),
DataType::BFloat16,
&shape,
&strides,
ep.device_id(),
);
RmsNormKernel {
axis: -1,
epsilon,
runtime: runtime.clone(),
warmed_signature: Mutex::new(None),
last_call_capture_safe: AtomicBool::new(false),
}
.run(&[x, scale_view], &mut [output])
.unwrap();
runtime.synchronize().unwrap();
let mut out_bytes = vec![0u8; std::mem::size_of_val(x_bf16.as_slice())];
unsafe {
runtime
.dtoh(&mut out_bytes, cuptr(out_buffer.as_ptr()))
.unwrap();
}
let out_bf16: Vec<bf16> = out_bytes
.chunks_exact(2)
.map(|raw| bf16::from_bits(u16::from_ne_bytes(raw.try_into().unwrap())))
.collect();
let x_ref: Vec<f64> = x_bf16.iter().map(|v| f64::from(v.to_f32())).collect();
let ms_f64: f64 = x_ref.iter().map(|v| v * v).sum::<f64>() / hidden as f64;
let inv_std_f64 = 1.0 / (ms_f64 + f64::from(epsilon)).sqrt();
let mut max_ulp = 0i32;
for i in 0..hidden {
let expect = x_ref[i] * inv_std_f64 * f64::from(scale_f32[i]);
let expect_bf16 = bf16::from_f32(expect as f32);
let ulp = (i32::from(out_bf16[i].to_bits()) - i32::from(expect_bf16.to_bits())).abs();
max_ulp = max_ulp.max(ulp);
}
assert!(
max_ulp <= 1,
"bf16 rmsnorm output diverges from f64 oracle by {max_ulp} bf16 ulp (want <= 1)"
);
let serial_ms: f32 = {
let mut ss = 0.0f32;
for v in &x_bf16 {
let f = v.to_f32();
ss += f * f;
}
ss / hidden as f32
};
let tree_ms: f32 = {
let mut level: Vec<f32> = x_bf16
.iter()
.map(|v| {
let f = v.to_f32();
f * f
})
.collect();
while level.len() > 1 {
let mut next = Vec::with_capacity(level.len().div_ceil(2));
let mut i = 0;
while i + 1 < level.len() {
next.push(level[i] + level[i + 1]);
i += 2;
}
if i < level.len() {
next.push(level[i]);
}
level = next;
}
level[0] / hidden as f32
};
let serial_err = (f64::from(serial_ms) - ms_f64).abs();
let tree_err = (f64::from(tree_ms) - ms_f64).abs();
assert!(
tree_err <= serial_err * 1.000_001 + 1e-9,
"tree mean-square error {tree_err:e} exceeds serial error {serial_err:e}"
);
}
}
#[cfg(test)]
mod claim_probes {
use std::ffi::c_void;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use half::{bf16, f16};
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut, TensorMut, TensorView};
use onnx_runtime_ir::{DataType, DeviceId};
use super::SkipLayerNormKernel;
use crate::runtime::CudaRuntime;
fn maybe_runtime() -> Option<Arc<CudaRuntime>> {
crate::test_support::maybe_runtime()
}
fn reference(input: &[f32], skip: &[f32], gamma: &[f32], eps: f32) -> Vec<f32> {
let n = input.len();
let s: Vec<f32> = (0..n).map(|i| input[i] + skip[i]).collect();
let mean = s.iter().sum::<f32>() / n as f32;
let var = s.iter().map(|&v| (v - mean) * (v - mean)).sum::<f32>() / n as f32;
let inv = 1.0 / (var + eps).sqrt();
(0..n).map(|i| (s[i] - mean) * inv * gamma[i]).collect()
}
#[test]
fn typed_skip_layernorm_f16_bf16_match_reference_on_gpu() {
let Some(runtime) = maybe_runtime() else {
eprintln!("skipping typed SkipLayerNorm GPU probe: CUDA runtime unavailable");
return;
};
if runtime
.require_nvrtc_half_headers("SkipLayerNormalization")
.is_err()
{
eprintln!("skipping typed SkipLayerNorm GPU probe: fp16 headers unavailable");
return;
}
let input = [1.0f32, 2.0, 3.0, 4.0];
let skip = [0.5f32, 0.5, 0.5, 0.5];
let gamma = [1.0f32, 0.5, 2.0, 1.5];
let eps = 1e-5f32;
let expect = reference(&input, &skip, &gamma, eps);
run_half::<f16>(
&runtime,
DataType::Float16,
f16::from_f32,
f16::to_f32,
&input,
&skip,
&gamma,
eps,
&expect,
3.0e-2,
);
run_half::<bf16>(
&runtime,
DataType::BFloat16,
bf16::from_f32,
bf16::to_f32,
&input,
&skip,
&gamma,
eps,
&expect,
1.5e-1,
);
}
#[allow(clippy::too_many_arguments)]
fn run_half<T: Copy>(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
to_h: impl Fn(f32) -> T,
from_h: impl Fn(T) -> f32,
input: &[f32],
skip: &[f32],
gamma: &[f32],
eps: f32,
expect: &[f32],
tol: f32,
) {
let n = input.len();
let hin: Vec<T> = input.iter().map(|&v| to_h(v)).collect();
let hskip: Vec<T> = skip.iter().map(|&v| to_h(v)).collect();
let hgamma: Vec<T> = gamma.iter().map(|&v| to_h(v)).collect();
let elem = std::mem::size_of::<T>();
let bytes = elem * n;
let in_dev = runtime.alloc_raw(bytes).unwrap();
let skip_dev = runtime.alloc_raw(bytes).unwrap();
let gamma_dev = runtime.alloc_raw(bytes).unwrap();
let out_dev = runtime.alloc_raw(bytes).unwrap();
let as_bytes = |v: &[T]| unsafe {
std::slice::from_raw_parts(v.as_ptr().cast::<u8>(), std::mem::size_of_val(v))
};
unsafe {
runtime.htod(as_bytes(&hin), in_dev).unwrap();
runtime.htod(as_bytes(&hskip), skip_dev).unwrap();
runtime.htod(as_bytes(&hgamma), gamma_dev).unwrap();
}
let device = DeviceId::cuda(0);
let shape = [1usize, n];
let strides = [n as i64, 1];
let mk = |ptr: u64| {
TensorView::new(
DevicePtr(ptr as usize as *const c_void),
dtype,
&shape,
&strides,
device,
)
};
let inputs = [mk(in_dev), mk(skip_dev), mk(gamma_dev)];
let mut outputs = [TensorMut::new(
DevicePtrMut(out_dev as usize as *mut c_void),
dtype,
&shape,
&strides,
device,
)];
let kernel = SkipLayerNormKernel {
epsilon: eps,
runtime: runtime.clone(),
last_call_capture_safe: AtomicBool::new(false),
};
kernel.run(&inputs, &mut outputs).unwrap();
runtime.synchronize().unwrap();
let mut out = vec![to_h(0.0); n];
let out_bytes =
unsafe { std::slice::from_raw_parts_mut(out.as_mut_ptr().cast::<u8>(), bytes) };
unsafe { runtime.dtoh(out_bytes, out_dev).unwrap() };
unsafe {
runtime.free_raw(in_dev).unwrap();
runtime.free_raw(skip_dev).unwrap();
runtime.free_raw(gamma_dev).unwrap();
runtime.free_raw(out_dev).unwrap();
}
for (i, (&o, &e)) in out.iter().zip(expect).enumerate() {
let got = from_h(o);
assert!(
(got - e).abs() <= tol,
"{dtype:?} SkipLayerNorm index {i}: expected {e}, got {got}"
);
}
}
}