use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::{check_arity, to_dense_bytes, to_dense_i64, write_dense_bytes};
pub struct UnsqueezeKernel {
axes: Option<Vec<i64>>,
}
pub struct UnsqueezeFactory;
impl KernelFactory for UnsqueezeFactory {
fn create(&self, node: &Node, _shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let axes = node
.attr("axes")
.and_then(|a| a.as_ints())
.map(|v| v.to_vec());
Ok(Box::new(UnsqueezeKernel { axes }))
}
}
impl Kernel for UnsqueezeKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("Unsqueeze", inputs, outputs, 1, 2, 1)?;
let axes_len = if inputs.len() >= 2 && !inputs[1].is_absent() {
to_dense_i64(&inputs[1])?.len()
} else {
self.axes
.as_ref()
.ok_or_else(|| {
EpError::KernelFailed(
"Unsqueeze: `axes` not supplied — provide it as the opset-13 second \
input or as the opset-12 `axes` attribute"
.into(),
)
})?
.len()
};
let expected_rank = inputs[0].shape.len() + axes_len;
if outputs[0].shape.len() != expected_rank {
return Err(EpError::KernelFailed(format!(
"Unsqueeze: output rank {} != input rank {} + {} axes",
outputs[0].shape.len(),
inputs[0].shape.len(),
axes_len
)));
}
let data = to_dense_bytes(&inputs[0])?;
write_dense_bytes(&mut outputs[0], &data)
}
fn supports_strided_input(&self, input_idx: usize) -> bool {
input_idx == 0
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
#[test]
fn unsqueeze_inserts_axis() {
let x = Owned::f32(&[3], &[1., 2., 3.]);
let mut out = Owned::zeros_f32(&[1, 3]);
UnsqueezeKernel {
axes: Some(vec![0]),
}
.execute(&[x.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3.]);
}
#[test]
fn unsqueeze_multiple_axes_int64() {
let x = Owned::i64(&[2], &[5, 9]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::Int64, &[1, 2, 1]);
UnsqueezeKernel {
axes: Some(vec![0, 2]),
}
.execute(&[x.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_i64(), vec![5, 9]);
}
#[test]
fn unsqueeze_missing_axes_errors() {
let x = Owned::f32(&[3], &[1., 2., 3.]);
let mut out = Owned::zeros_f32(&[1, 3]);
assert!(
UnsqueezeKernel { axes: None }
.execute(&[x.view()], &mut [out.view_mut()])
.is_err()
);
}
#[test]
fn unsqueeze_axes_as_input_opset13() {
let x = Owned::f32(&[3], &[1., 2., 3.]);
let axes = Owned::i64(&[1], &[0]);
let mut out = Owned::zeros_f32(&[1, 3]);
UnsqueezeKernel { axes: None }
.execute(&[x.view(), axes.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3.]);
}
#[test]
fn unsqueeze_axes_input_multiple() {
let x = Owned::i64(&[2], &[5, 9]);
let axes = Owned::i64(&[2], &[0, 2]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::Int64, &[1, 2, 1]);
UnsqueezeKernel { axes: None }
.execute(&[x.view(), axes.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_i64(), vec![5, 9]);
}
#[test]
fn unsqueeze_axes_input_rank_mismatch_errors() {
let x = Owned::f32(&[3], &[1., 2., 3.]);
let axes = Owned::i64(&[2], &[0, 1]);
let mut out = Owned::zeros_f32(&[1, 3]);
assert!(
UnsqueezeKernel { axes: None }
.execute(&[x.view(), axes.view()], &mut [out.view_mut()])
.is_err()
);
}
#[test]
fn unsqueeze_bf16_preserves_element_bits() {
let x = Owned::bf16(&[2], &[1., -2.]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::BFloat16, &[1, 2]);
UnsqueezeKernel {
axes: Some(vec![0]),
}
.execute(&[x.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_u16_bits(), x.to_u16_bits());
}
}