use std::ffi::c_void;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use cudarc::driver::PushKernelArg;
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::{DataType, Node};
use crate::cudnn::{CudnnBufferPair, CudnnReduceOp, TensorDescriptorSpec};
use crate::error::{driver_err, not_implemented};
use crate::runtime::{CudaRuntime, cuptr};
const REDUCE_SRC: &str = r#"
extern "C" __global__ void validate_reduce_axes_i64(
const long long* actual,
const long long* expected,
const int count,
unsigned int* capture_error)
{
for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < count;
i += blockDim.x * gridDim.x) {
if (actual[i] != expected[i]) atomicOr(capture_error, 128u);
}
}
extern "C" __global__ void reduce_f32(
const float* x,
float* y,
const long long* base_off, // [out_count]
const long long* delta_off, // [reduce_count]
const int out_count,
const int reduce_count,
const int op, // 0 sum, 1 max, 2 min
const int is_mean,
const unsigned int* capture_error)
{
if (capture_error && *capture_error) return;
const int o = blockIdx.x;
if (o >= out_count) return;
const float NEG_INF = __int_as_float(0xff800000);
const float POS_INF = __int_as_float(0x7f800000);
const float QNAN = __int_as_float(0x7fc00000);
extern __shared__ float red[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
const size_t base = (size_t)base_off[o];
float acc = (op == 1) ? NEG_INF : (op == 2) ? POS_INF : 0.0f;
for (int r = tid; r < reduce_count; r += nt) {
const float v = x[base + (size_t)delta_off[r]];
if (op == 1) acc = (isnan(acc) || isnan(v)) ? QNAN : fmaxf(acc, v);
else if (op == 2) acc = (isnan(acc) || isnan(v)) ? QNAN : fminf(acc, v);
else acc += v;
}
red[tid] = acc;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) {
const float a = red[tid], b = red[tid + off];
if (op == 1) red[tid] = (isnan(a) || isnan(b)) ? QNAN : fmaxf(a, b);
else if (op == 2) red[tid] = (isnan(a) || isnan(b)) ? QNAN : fminf(a, b);
else red[tid] = a + b;
}
__syncthreads();
}
if (tid == 0) {
float out = red[0];
if (is_mean) out /= (float)reduce_count;
y[o] = out;
}
}
extern "C" __global__ void reduce_i64_sum(
const long long* x,
long long* y,
const long long* base_off,
const long long* delta_off,
const int out_count,
const int reduce_count,
const unsigned int* capture_error)
{
if (capture_error && *capture_error) return;
const int o = blockIdx.x;
if (o >= out_count) return;
extern __shared__ long long red_i64[];
const int tid = threadIdx.x;
const int nt = blockDim.x;
const size_t base = (size_t)base_off[o];
long long acc = 0;
for (int r = tid; r < reduce_count; r += nt) {
acc += x[base + (size_t)delta_off[r]];
}
red_i64[tid] = acc;
__syncthreads();
for (int off = nt >> 1; off > 0; off >>= 1) {
if (tid < off) red_i64[tid] += red_i64[tid + off];
__syncthreads();
}
if (tid == 0) y[o] = red_i64[0];
}
"#;
const REDUCE_MODULE: &str = "reduce_f32";
const REDUCE_ENTRY: &str = "reduce_f32";
const REDUCE_I64_SUM_ENTRY: &str = "reduce_i64_sum";
const REDUCE_VALIDATE_AXES_ENTRY: &str = "validate_reduce_axes_i64";
pub const REDUCE_CAPTURE_ERROR_AXES: u32 = 128;
const REDUCE_BLOCK: u32 = 256;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReduceOp {
Sum,
Mean,
Max,
Min,
}
impl ReduceOp {
fn name(self) -> &'static str {
match self {
ReduceOp::Sum => "ReduceSum",
ReduceOp::Mean => "ReduceMean",
ReduceOp::Max => "ReduceMax",
ReduceOp::Min => "ReduceMin",
}
}
fn kernel_tags(self) -> (i32, i32) {
match self {
ReduceOp::Sum => (0, 0),
ReduceOp::Mean => (0, 1),
ReduceOp::Max => (1, 0),
ReduceOp::Min => (2, 0),
}
}
fn cudnn_op(self) -> Option<CudnnReduceOp> {
match self {
ReduceOp::Sum => Some(CudnnReduceOp::Add),
ReduceOp::Mean => Some(CudnnReduceOp::Average),
ReduceOp::Max | ReduceOp::Min => None,
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct ReductionPlan {
pub base: Vec<i64>,
pub delta: Vec<i64>,
pub out_shape: Vec<usize>,
}
fn contiguous_strides(shape: &[usize]) -> Vec<i64> {
let mut strides = vec![0i64; shape.len()];
let mut acc = 1i64;
for d in (0..shape.len()).rev() {
strides[d] = acc;
acc *= shape[d] as i64;
}
strides
}
fn contiguous_strides_usize(shape: &[usize]) -> Vec<usize> {
let mut strides = vec![0usize; shape.len()];
let mut acc = 1usize;
for d in (0..shape.len()).rev() {
strides[d] = acc;
acc *= shape[d];
}
strides
}
fn reduced_output_shape(in_shape: &[usize], reduce: &[bool], keepdims: bool) -> Vec<usize> {
let mut out_shape = Vec::with_capacity(in_shape.len());
for (dim, &is_reduced) in in_shape.iter().zip(reduce) {
if is_reduced {
if keepdims {
out_shape.push(1);
}
} else {
out_shape.push(*dim);
}
}
out_shape
}
pub(crate) fn cudnn_reduce_specs(
dtype: DataType,
in_shape: &[usize],
reduce: &[bool],
) -> Result<(TensorDescriptorSpec, TensorDescriptorSpec)> {
let cudnn_out_shape: Vec<usize> = in_shape
.iter()
.zip(reduce)
.map(|(&dim, &is_reduced)| if is_reduced { 1 } else { dim })
.collect();
let input = TensorDescriptorSpec::new(dtype, in_shape, &contiguous_strides_usize(in_shape))?;
let output = TensorDescriptorSpec::new(
dtype,
&cudnn_out_shape,
&contiguous_strides_usize(&cudnn_out_shape),
)?;
Ok((input, output))
}
pub(crate) fn build_plan(in_shape: &[usize], reduce: &[bool], keepdims: bool) -> ReductionPlan {
let rank = in_shape.len();
let strides = contiguous_strides(in_shape);
let kept_axes: Vec<usize> = (0..rank).filter(|&d| !reduce[d]).collect();
let red_axes: Vec<usize> = (0..rank).filter(|&d| reduce[d]).collect();
let kept_dims: Vec<usize> = kept_axes.iter().map(|&d| in_shape[d]).collect();
let red_dims: Vec<usize> = red_axes.iter().map(|&d| in_shape[d]).collect();
let base = enumerate_offsets(&kept_dims, &kept_axes, &strides);
let delta = enumerate_offsets(&red_dims, &red_axes, &strides);
let out_shape = reduced_output_shape(in_shape, reduce, keepdims);
ReductionPlan {
base,
delta,
out_shape,
}
}
fn enumerate_offsets(dims: &[usize], axes: &[usize], strides: &[i64]) -> Vec<i64> {
let total: usize = dims.iter().product::<usize>().max(1);
let mut out = Vec::with_capacity(total);
let mut idx = vec![0usize; dims.len()];
loop {
let mut off = 0i64;
for k in 0..dims.len() {
off += idx[k] as i64 * strides[axes[k]];
}
out.push(off);
if !next_index(dims, &mut idx) {
break;
}
}
out
}
fn next_index(dims: &[usize], idx: &mut [usize]) -> bool {
for d in (0..dims.len()).rev() {
idx[d] += 1;
if idx[d] < dims[d] {
return true;
}
idx[d] = 0;
}
false
}
macro_rules! reduce_factory {
($factory:ident, $variant:expr) => {
pub struct $factory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for $factory {
fn create(&self, node: &Node, _shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let axes_attr = node
.attr("axes")
.and_then(|a| a.as_ints())
.map(<[i64]>::to_vec);
let keepdims = node.attr("keepdims").and_then(|a| a.as_int()).unwrap_or(1) != 0;
let noop_with_empty_axes = node
.attr("noop_with_empty_axes")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0;
Ok(Box::new(ReduceKernel {
op: $variant,
axes_attr,
keepdims,
noop_with_empty_axes,
runtime: self.runtime.clone(),
int64_metadata: Mutex::new(ReductionMetadataCache::new(self.runtime.clone())),
last_call_capture_safe: AtomicBool::new(false),
}))
}
}
};
}
reduce_factory!(ReduceSumFactory, ReduceOp::Sum);
reduce_factory!(ReduceMeanFactory, ReduceOp::Mean);
reduce_factory!(ReduceMaxFactory, ReduceOp::Max);
reduce_factory!(ReduceMinFactory, ReduceOp::Min);
#[derive(Debug)]
pub struct ReduceKernel {
op: ReduceOp,
axes_attr: Option<Vec<i64>>,
keepdims: bool,
noop_with_empty_axes: bool,
runtime: Arc<CudaRuntime>,
int64_metadata: Mutex<ReductionMetadataCache>,
last_call_capture_safe: AtomicBool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct ReductionMetadataKey {
input_shape: Vec<usize>,
reduce: Vec<bool>,
keepdims: bool,
axes: Vec<i64>,
}
#[derive(Debug)]
struct ReductionMetadataCache {
runtime: Arc<CudaRuntime>,
key: Option<ReductionMetadataKey>,
base: CUdeviceptr,
delta: CUdeviceptr,
axes: CUdeviceptr,
}
impl ReductionMetadataCache {
fn new(runtime: Arc<CudaRuntime>) -> Self {
Self {
runtime,
key: None,
base: 0,
delta: 0,
axes: 0,
}
}
fn prepare(
&mut self,
input_shape: &[usize],
reduce: &[bool],
keepdims: bool,
axes: &[i64],
plan: &ReductionPlan,
) -> Result<(CUdeviceptr, CUdeviceptr, CUdeviceptr)> {
let key = ReductionMetadataKey {
input_shape: input_shape.to_vec(),
reduce: reduce.to_vec(),
keepdims,
axes: axes.to_vec(),
};
if self.key.as_ref() == Some(&key) {
return Ok((self.base, self.delta, self.axes));
}
if self.runtime.is_capturing()? {
return Err(EpError::KernelFailed(
"cuda_ep ReduceSum: int64 reduction metadata changed during CUDA graph capture; warm the fixed decode shape before capture".into(),
));
}
if self.base != 0 || self.delta != 0 || self.axes != 0 {
self.runtime.synchronize()?;
}
let base_bytes = as_i64_bytes(&plan.base);
let delta_bytes = as_i64_bytes(&plan.delta);
let axes_bytes = as_i64_bytes(axes);
let base = self.runtime.alloc_raw(base_bytes.len().max(1))?;
let delta = match self.runtime.alloc_raw(delta_bytes.len().max(1)) {
Ok(delta) => delta,
Err(error) => {
let _ = unsafe { self.runtime.free_raw(base) };
return Err(error);
}
};
let axes_ptr = match self.runtime.alloc_raw(axes_bytes.len().max(1)) {
Ok(axes_ptr) => axes_ptr,
Err(error) => {
let _ = unsafe { self.runtime.free_raw(base) };
let _ = unsafe { self.runtime.free_raw(delta) };
return Err(error);
}
};
let upload = (|| {
unsafe { self.runtime.htod(&base_bytes, base) }?;
unsafe { self.runtime.htod(&delta_bytes, delta) }?;
unsafe { self.runtime.htod(&axes_bytes, axes_ptr) }
})();
if let Err(error) = upload {
let _ = unsafe { self.runtime.free_raw(base) };
let _ = unsafe { self.runtime.free_raw(delta) };
let _ = unsafe { self.runtime.free_raw(axes_ptr) };
return Err(error);
}
if self.base != 0 {
unsafe { self.runtime.free_raw(self.base) }?;
}
if self.delta != 0 {
unsafe { self.runtime.free_raw(self.delta) }?;
}
if self.axes != 0 {
unsafe { self.runtime.free_raw(self.axes) }?;
}
self.key = Some(key);
self.base = base;
self.delta = delta;
self.axes = axes_ptr;
Ok((base, delta, axes_ptr))
}
}
impl Drop for ReductionMetadataCache {
fn drop(&mut self) {
if self.base != 0 {
let _ = unsafe { self.runtime.free_raw(self.base) };
self.base = 0;
}
if self.delta != 0 {
let _ = unsafe { self.runtime.free_raw(self.delta) };
self.delta = 0;
}
if self.axes != 0 {
let _ = unsafe { self.runtime.free_raw(self.axes) };
self.axes = 0;
}
}
}
pub(crate) fn resolve_reduce_mask(
op: &str,
axes_raw: &Option<Vec<i64>>,
rank: usize,
noop_with_empty_axes: bool,
) -> Result<Vec<bool>> {
let mut reduce = vec![false; rank];
match axes_raw {
Some(a) if a.is_empty() => {
if !noop_with_empty_axes {
reduce.iter_mut().for_each(|r| *r = true);
}
}
Some(axes) => {
for &a in axes {
let ax = if a < 0 { a + rank as i64 } else { a };
if ax < 0 || ax as usize >= rank {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: axis {a} is out of range for a rank-{rank} input; \
axis must lie in [-{rank}, {rank})"
)));
}
reduce[ax as usize] = true;
}
}
None => {
if !noop_with_empty_axes {
reduce.iter_mut().for_each(|r| *r = true);
}
}
}
Ok(reduce)
}
impl ReduceKernel {
fn read_axes_input(&self, op: &str, axes: &TensorView) -> Result<Vec<i64>> {
if !axes.is_contiguous() {
return Err(not_implemented(format!(
"{op} with a non-contiguous (strided) axes input; materialise it first"
)));
}
let n = axes.numel();
let src = cuptr(axes.data_ptr::<u8>() as *const c_void);
match axes.dtype {
DataType::Int64 => {
let mut bytes = vec![0u8; n * std::mem::size_of::<i64>()];
unsafe { self.runtime.dtoh(&mut bytes, src) }?;
Ok(bytes
.chunks_exact(8)
.map(|c| i64::from_ne_bytes(c.try_into().unwrap()))
.collect())
}
DataType::Int32 => {
let mut bytes = vec![0u8; n * std::mem::size_of::<i32>()];
unsafe { self.runtime.dtoh(&mut bytes, src) }?;
Ok(bytes
.chunks_exact(4)
.map(|c| i32::from_ne_bytes(c.try_into().unwrap()) as i64)
.collect())
}
other => Err(not_implemented(format!(
"{op} with axes input dtype {other:?} (expected int32 or int64)"
))),
}
}
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
let op = self.op.name();
if !(1..=2).contains(&inputs.len()) || outputs.len() != 1 {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: expected 1-2 inputs (data[, axes]) and 1 output, got {} and {}",
inputs.len(),
outputs.len()
)));
}
let x = &inputs[0];
let cudnn_op = self.op.cudnn_op();
let supported_dtype = if self.op == ReduceOp::Sum && x.dtype == DataType::Int64 {
true
} else if cudnn_op.is_some() {
matches!(
x.dtype,
DataType::Float32 | DataType::Float16 | DataType::BFloat16
)
} else {
x.dtype == DataType::Float32
};
if !supported_dtype {
return Err(not_implemented(format!(
"{op} with input dtype {:?} (sum supports i64/f32/f16/bf16; mean supports \
f32/f16/bf16; max/min are f32)",
x.dtype
)));
}
if outputs[0].dtype != x.dtype {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output dtype {:?} must equal input dtype {:?}",
outputs[0].dtype, x.dtype
)));
}
if !x.is_contiguous() || !outputs[0].is_contiguous() {
return Err(not_implemented(format!(
"{op} with a non-contiguous (strided) input/output; materialise it first"
)));
}
let rank = x.shape.len();
let capturing = self.runtime.is_capturing()?;
let axes_raw: Option<Vec<i64>> = if inputs.len() == 2 && capturing {
if inputs[1].dtype != DataType::Int64 {
return Err(EpError::KernelFailed(
"cuda_ep ReduceSum: captured axes input must be Int64".into(),
));
}
Some(
self.int64_metadata
.lock()
.map_err(|_| {
EpError::KernelFailed(
"cuda_ep ReduceSum: metadata cache lock was poisoned".into(),
)
})?
.key
.as_ref()
.ok_or_else(|| {
EpError::KernelFailed(
"cuda_ep ReduceSum: axes were not warmed before CUDA graph capture"
.into(),
)
})?
.axes
.clone(),
)
} else if inputs.len() == 2 {
Some(self.read_axes_input(op, &inputs[1])?)
} else {
self.axes_attr.clone()
};
let reduce = resolve_reduce_mask(op, &axes_raw, rank, self.noop_with_empty_axes)?;
let expected_shape = reduced_output_shape(x.shape, &reduce, self.keepdims);
if outputs[0].shape != expected_shape.as_slice() {
return Err(EpError::KernelFailed(format!(
"cuda_ep {op}: output shape {:?} does not match the reduced shape {:?} \
(axes {:?}, keepdims {})",
outputs[0].shape, expected_shape, axes_raw, self.keepdims
)));
}
if x.numel() == 0 || outputs[0].numel() == 0 {
return Ok(());
}
if !reduce.iter().any(|&axis| axis) || rank == 0 {
let src = cuptr(x.data_ptr::<u8>() as *const c_void);
let dst = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
if src != dst {
unsafe { self.runtime.dtod(src, dst, x.byte_size()) }?;
}
return Ok(());
}
if x.dtype != DataType::Int64
&& let Some(cudnn_op) = cudnn_op
{
if self.runtime.cudnn().is_available() {
let (input_spec, output_spec) = cudnn_reduce_specs(x.dtype, x.shape, &reduce)?;
let x_ptr = cuptr(x.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
self.runtime.cudnn().with_handle(|handle| {
handle.reduce(
&input_spec,
&output_spec,
cudnn_op,
CudnnBufferPair {
input: x_ptr,
output: y_ptr,
input_numel: x.numel(),
output_numel: outputs[0].numel(),
},
)
})?;
return self.runtime.synchronize();
}
if x.dtype != DataType::Float32 {
return self.runtime.cudnn().with_handle(|_| Ok(()));
}
}
let plan = build_plan(x.shape, &reduce, self.keepdims);
let out_count = plan.base.len();
let reduce_count = plan.delta.len();
if out_count == 0 || reduce_count == 0 {
return Ok(());
}
if x.dtype == DataType::Int64 && (inputs.len() == 1 || inputs[1].dtype == DataType::Int64) {
let axes = axes_raw.as_deref().unwrap_or(&[]);
let mut metadata = self.int64_metadata.lock().map_err(|_| {
EpError::KernelFailed("cuda_ep ReduceSum: metadata cache lock was poisoned".into())
})?;
let (base_buf, delta_buf, expected_axes) =
metadata.prepare(x.shape, &reduce, self.keepdims, axes, &plan)?;
if capturing && inputs.len() == 2 {
self.validate_captured_axes(&inputs[1], expected_axes)?;
}
self.launch(
x,
outputs,
base_buf,
delta_buf,
out_count,
reduce_count,
capturing,
)?;
self.last_call_capture_safe.store(true, Ordering::Relaxed);
return Ok(());
}
let base_bytes = as_i64_bytes(&plan.base);
let delta_bytes = as_i64_bytes(&plan.delta);
let base_buf = self.runtime.alloc_raw(base_bytes.len())?;
let delta_buf = self.runtime.alloc_raw(delta_bytes.len())?;
let result = (|| {
unsafe { self.runtime.htod(&base_bytes, base_buf) }?;
unsafe { self.runtime.htod(&delta_bytes, delta_buf) }?;
self.launch(
x,
outputs,
base_buf,
delta_buf,
out_count,
reduce_count,
false,
)
})();
let free_base = unsafe { self.runtime.free_raw(base_buf) };
let free_delta = unsafe { self.runtime.free_raw(delta_buf) };
result.and(free_base).and(free_delta)
}
fn validate_captured_axes(&self, actual: &TensorView, expected: CUdeviceptr) -> Result<()> {
let count = i32::try_from(actual.numel()).map_err(|_| {
EpError::KernelFailed("cuda_ep ReduceSum: axes count exceeds i32".into())
})?;
let actual = cuptr(actual.data_ptr::<u8>() as *const c_void);
let capture_error = self.runtime.capture_error_ptr();
let func =
self.runtime
.nvrtc_function(REDUCE_MODULE, REDUCE_SRC, REDUCE_VALIDATE_AXES_ENTRY)?;
let mut builder = self.runtime.stream().launch_builder(&func);
builder
.arg(&actual)
.arg(&expected)
.arg(&count)
.arg(&capture_error);
unsafe {
builder.launch(cudarc::driver::LaunchConfig {
grid_dim: ((count as u32).div_ceil(REDUCE_BLOCK).max(1), 1, 1),
block_dim: (REDUCE_BLOCK, 1, 1),
shared_mem_bytes: 0,
})
}
.map_err(|error| driver_err("launch validate_reduce_axes_i64", error))?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn launch(
&self,
x: &TensorView,
outputs: &mut [TensorMut],
base_buf: CUdeviceptr,
delta_buf: CUdeviceptr,
out_count: usize,
reduce_count: usize,
capturing: bool,
) -> Result<()> {
let op = self.op.name();
let out_i = i32::try_from(out_count).map_err(|_| {
EpError::KernelFailed(format!("cuda_ep {op}: {out_count} outputs exceed i32"))
})?;
let red_i = i32::try_from(reduce_count).map_err(|_| {
EpError::KernelFailed(format!(
"cuda_ep {op}: reduction group {reduce_count} exceeds i32"
))
})?;
let grid = u32::try_from(out_count).map_err(|_| {
EpError::KernelFailed(format!("cuda_ep {op}: {out_count} blocks exceed u32"))
})?;
let (op_tag, is_mean) = self.op.kernel_tags();
let x_ptr = cuptr(x.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let capture_error = if capturing {
self.runtime.capture_error_ptr()
} else {
0
};
let entry = if x.dtype == DataType::Int64 {
REDUCE_I64_SUM_ENTRY
} else {
REDUCE_ENTRY
};
let func = self
.runtime
.nvrtc_function(REDUCE_MODULE, REDUCE_SRC, entry)?;
let bytes_per_thread = if x.dtype == DataType::Int64 {
std::mem::size_of::<i64>() as u32
} else {
std::mem::size_of::<f32>() as u32
};
let cfg =
self.runtime
.reduction_launch_config(&func, grid, REDUCE_BLOCK, bytes_per_thread)?;
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
builder
.arg(&x_ptr)
.arg(&y_ptr)
.arg(&base_buf)
.arg(&delta_buf)
.arg(&out_i)
.arg(&red_i);
if x.dtype != DataType::Int64 {
builder.arg(&op_tag).arg(&is_mean).arg(&capture_error);
} else {
builder.arg(&capture_error);
}
unsafe { builder.launch(cfg) }.map_err(|e| driver_err(&format!("launch {entry}"), e))?;
if capturing {
Ok(())
} else {
self.runtime.synchronize()
}
}
}
fn as_i64_bytes(v: &[i64]) -> Vec<u8> {
let mut out = Vec::with_capacity(v.len() * 8);
for &x in v {
out.extend_from_slice(&x.to_ne_bytes());
}
out
}
impl Kernel for ReduceKernel {
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(
"requires a warmed fixed-shape ReduceSum path with stable device-resident axes metadata",
)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn entry_point_present_in_source() {
assert!(REDUCE_SRC.contains(REDUCE_ENTRY));
}
#[test]
fn strides_are_row_major() {
assert_eq!(contiguous_strides(&[2, 3, 4]), vec![12, 4, 1]);
}
#[test]
fn plan_reduce_last_axis_keepdims() {
let reduce = [false, true];
let plan = build_plan(&[2, 3], &reduce, true);
assert_eq!(plan.out_shape, vec![2, 1]);
assert_eq!(plan.base, vec![0, 3]); assert_eq!(plan.delta, vec![0, 1, 2]); }
#[test]
fn plan_reduce_axis0_no_keepdims() {
let reduce = [true, false];
let plan = build_plan(&[2, 3], &reduce, false);
assert_eq!(plan.out_shape, vec![3]);
assert_eq!(plan.base, vec![0, 1, 2]); assert_eq!(plan.delta, vec![0, 3]); }
#[test]
fn plan_reduce_all_axes() {
let reduce = [true, true];
let plan = build_plan(&[2, 3], &reduce, true);
assert_eq!(plan.out_shape, vec![1, 1]);
assert_eq!(plan.base, vec![0]);
assert_eq!(plan.delta, vec![0, 1, 2, 3, 4, 5]);
}
#[test]
fn resolve_mask_negative_axis_and_empty_noop() {
let m = resolve_reduce_mask("ReduceSum", &Some(vec![-1]), 3, false).unwrap();
assert_eq!(m, vec![false, false, true]);
let m = resolve_reduce_mask("ReduceSum", &Some(vec![]), 3, true).unwrap();
assert_eq!(m, vec![false, false, false]);
let m = resolve_reduce_mask("ReduceSum", &Some(vec![]), 3, false).unwrap();
assert_eq!(m, vec![true, true, true]);
let m = resolve_reduce_mask("ReduceSum", &None, 2, false).unwrap();
assert_eq!(m, vec![true, true]);
}
#[test]
fn resolve_mask_rejects_out_of_range_axis() {
let e = resolve_reduce_mask("ReduceMax", &Some(vec![5]), 2, false).unwrap_err();
let msg = format!("{e}");
assert!(msg.contains("out of range"), "{msg}");
assert!(msg.contains("axis 5"), "{msg}");
}
#[test]
fn kernel_tags_map_ops() {
assert_eq!(ReduceOp::Sum.kernel_tags(), (0, 0));
assert_eq!(ReduceOp::Mean.kernel_tags(), (0, 1));
assert_eq!(ReduceOp::Max.kernel_tags(), (1, 0));
assert_eq!(ReduceOp::Min.kernel_tags(), (2, 0));
}
#[test]
fn cudnn_op_mapping_only_ports_sum_and_mean() {
assert_eq!(ReduceOp::Sum.cudnn_op(), Some(CudnnReduceOp::Add));
assert_eq!(ReduceOp::Mean.cudnn_op(), Some(CudnnReduceOp::Average));
assert_eq!(ReduceOp::Max.cudnn_op(), None);
assert_eq!(ReduceOp::Min.cudnn_op(), None);
}
#[test]
fn cudnn_specs_keep_reduced_axes_as_size_one() {
let (input, output) =
cudnn_reduce_specs(DataType::BFloat16, &[2, 3, 4], &[true, false, true]).unwrap();
assert_eq!(input.dims(), &[1, 2, 3, 4]);
assert_eq!(input.strides(), &[24, 12, 4, 1]);
assert_eq!(output.dims(), &[1, 1, 3, 1]);
assert_eq!(output.strides(), &[3, 3, 1, 1]);
}
}