use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::check_arity;
use crate::dtype::{
output_direct_write_eligible, slice_byte_range, to_dense_f32_widen, write_dense_f32_narrow,
};
use crate::kernels::activations::silu_f32_slice;
#[derive(Clone, Copy, PartialEq, Eq)]
enum ConvActivation {
None,
Silu,
}
pub struct CausalConvWithStateKernel {
activation: ConvActivation,
}
pub struct CausalConvWithStateFactory;
impl KernelFactory for CausalConvWithStateFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let ndim = node.attr("ndim").and_then(|a| a.as_int()).unwrap_or(1);
if ndim != 1 {
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: only ndim=1 (channels-first [B, C, L]) is supported by the \
CPU kernel, got ndim={ndim}"
)));
}
let activation = match node.attr("activation").and_then(|a| a.as_str()) {
None | Some("none") => ConvActivation::None,
Some("silu") | Some("swish") => ConvActivation::Silu,
Some(other) => {
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: activation must be one of none, silu, swish; got {other:?}"
)));
}
};
Ok(Box::new(CausalConvWithStateKernel { activation }))
}
}
impl Kernel for CausalConvWithStateKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("CausalConvWithState", inputs, outputs, 2, 4, 1)?;
let x_shape = inputs[0].shape;
if x_shape.len() != 3 {
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: input must be rank 3 [B, C, L] for ndim=1, got shape {x_shape:?}"
)));
}
let (batch, channels, length) = (x_shape[0], x_shape[1], x_shape[2]);
let w_shape = inputs[1].shape;
if w_shape.len() != 3 || w_shape[0] != channels || w_shape[1] != 1 {
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: weight must be depthwise [C, 1, K] with C={channels}, got \
shape {w_shape:?}"
)));
}
let kernel_size = w_shape[2];
if kernel_size == 0 {
return Err(EpError::KernelFailed(
"CausalConvWithState: kernel size K must be >= 1".into(),
));
}
let pad = kernel_size - 1;
let has_bias = inputs.len() >= 3;
let has_state = inputs.len() >= 4;
let x = to_dense_f32_widen("CausalConvWithState", &inputs[0])?;
let weight = to_dense_f32_widen("CausalConvWithState", &inputs[1])?;
let bias = if has_bias {
let b = to_dense_f32_widen("CausalConvWithState", &inputs[2])?;
if b.len() != channels {
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: bias must be 1D with size C={channels}, got {} elements",
b.len()
)));
}
Some(b)
} else {
None
};
let state = if has_state {
let s_shape = inputs[3].shape;
if s_shape.len() != 3
|| s_shape[0] != batch
|| s_shape[1] != channels
|| s_shape[2] != pad
{
return Err(EpError::KernelFailed(format!(
"CausalConvWithState: past_state must be [B={batch}, C={channels}, K-1={pad}], \
got shape {s_shape:?}"
)));
}
Some(to_dense_f32_widen("CausalConvWithState", &inputs[3])?)
} else {
None
};
let (primary_output, present_outputs) = outputs.split_at_mut(1);
let out_len = batch * channels * length;
let read_ranges: Vec<_> = [Some(&*x), Some(&*weight), bias.as_deref(), state.as_deref()]
.into_iter()
.flatten()
.map(slice_byte_range)
.collect();
let direct_out =
output_direct_write_eligible(&mut primary_output[0], out_len, &read_ranges);
let mut owned_out;
let out: &mut [f32] = if direct_out {
unsafe {
std::slice::from_raw_parts_mut(primary_output[0].data_ptr_mut::<f32>(), out_len)
}
} else {
owned_out = vec![0.0f32; out_len];
&mut owned_out
};
let present_len = batch * channels * pad;
let direct_present = present_outputs.first_mut().is_some_and(|present| {
output_direct_write_eligible(present, present_len, &read_ranges)
});
let mut owned_present;
let mut present = if let Some(output) = present_outputs.first_mut() {
if direct_present {
Some(unsafe {
std::slice::from_raw_parts_mut(output.data_ptr_mut::<f32>(), present_len)
})
} else {
owned_present = vec![0.0f32; present_len];
Some(owned_present.as_mut_slice())
}
} else {
None
};
for b in 0..batch {
for c in 0..channels {
let x_row = &x[(b * channels + c) * length..(b * channels + c) * length + length];
let w_row = &weight[c * kernel_size..c * kernel_size + kernel_size];
let bias_c = bias.as_ref().map_or(0.0, |bv| bv[c]);
let state_row = state
.as_ref()
.map(|sv| &sv[(b * channels + c) * pad..(b * channels + c) * pad + pad]);
let out_row =
&mut out[(b * channels + c) * length..(b * channels + c) * length + length];
if length == 1 {
let mut acc = bias_c;
for (k, &w) in w_row[..pad].iter().enumerate() {
acc += w * state_row.map_or(0.0, |s| s[k]);
}
acc += w_row[pad] * x_row[0];
out_row[0] = acc;
} else {
for (t, out_t) in out_row.iter_mut().enumerate() {
let mut acc = bias_c;
for (k, &w) in w_row.iter().enumerate() {
let pos = t + k;
let val = if pos < pad {
state_row.map_or(0.0, |s| s[pos])
} else {
x_row[pos - pad]
};
acc += w * val;
}
*out_t = acc;
}
}
if let Some(present) = present.as_deref_mut() {
let present_row =
&mut present[(b * channels + c) * pad..(b * channels + c) * pad + pad];
if length == 1 && pad > 0 {
for (slot, source) in present_row[..pad - 1].iter_mut().zip(1..pad) {
*slot = state_row.map_or(0.0, |s| s[source]);
}
present_row[pad - 1] = x_row[0];
} else {
for (p, slot) in present_row.iter_mut().enumerate() {
let pos = length + p;
*slot = if pos < pad {
state_row.map_or(0.0, |s| s[pos])
} else {
x_row[pos - pad]
};
}
}
}
}
}
if self.activation == ConvActivation::Silu {
let mut activated = vec![0.0f32; out.len()];
silu_f32_slice(out, &mut activated);
out.copy_from_slice(&activated);
}
if !direct_out {
write_dense_f32_narrow("CausalConvWithState", &mut primary_output[0], out)?;
}
if !direct_present
&& let (Some(output), Some(present)) = (present_outputs.first_mut(), present.as_deref())
{
write_dense_f32_narrow("CausalConvWithState", output, present)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
use onnx_runtime_ir::{Attribute, DataType, Node, NodeId};
fn kernel(silu: bool) -> Box<dyn Kernel> {
let mut node = Node::new(NodeId(0), "CausalConvWithState", vec![], vec![]);
node.domain = "com.microsoft".to_string();
node.attributes
.insert("ndim".to_string(), Attribute::Int(1));
node.attributes.insert(
"activation".to_string(),
Attribute::String(if silu {
b"silu".to_vec()
} else {
b"none".to_vec()
}),
);
CausalConvWithStateFactory.create(&node, &[]).unwrap()
}
fn assert_close(got: &[f32], want: &[f32], tag: &str) {
assert_eq!(got.len(), want.len(), "{tag} length");
for (i, (&a, &b)) in got.iter().zip(want).enumerate() {
let diff = (a - b).abs();
let rel = diff / b.abs().max(1e-6);
assert!(
diff <= 1e-4 || rel <= 1e-4,
"{tag}[{i}]: got {a}, want {b} (abs {diff}, rel {rel})"
);
}
}
#[test]
fn optional_bias_and_state_default_to_zero() {
let mut node = Node::new(NodeId(0), "CausalConvWithState", vec![], vec![]);
node.domain = "com.microsoft".to_string();
node.attributes
.insert("ndim".to_string(), Attribute::Int(1));
node.attributes.insert(
"activation".to_string(),
Attribute::String(b"none".to_vec()),
);
let kern = CausalConvWithStateFactory.create(&node, &[]).unwrap();
let x = Owned::f32(&[1, 1, 2], &[1.0, 2.0]);
let w = Owned::f32(&[1, 1, 2], &[3.0, 5.0]);
let mut y = Owned::zeros_f32(&[1, 1, 2]);
let mut present = Owned::zeros_f32(&[1, 1, 1]);
let ins = [x.view(), w.view()];
let mut outs = [y.view_mut(), present.view_mut()];
kern.execute(&ins, &mut outs).unwrap();
assert_eq!(y.to_f32(), vec![5.0, 13.0]);
assert_eq!(present.to_f32(), vec![2.0]);
}
#[test]
fn rejects_unknown_activation() {
let mut node = Node::new(NodeId(0), "CausalConvWithState", vec![], vec![]);
node.domain = "com.microsoft".to_string();
node.attributes.insert(
"activation".to_string(),
Attribute::String(b"gelu".to_vec()),
);
assert!(CausalConvWithStateFactory.create(&node, &[]).is_err());
}
#[test]
fn aliased_bindings_match_disjoint_result() {
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut};
use onnx_runtime_ir::{DeviceId, compute_contiguous_strides};
let (batch, c, k, s) = (1usize, 2usize, 3usize, 1usize);
let pad = k - 1;
let x_vals = [0.5f32, -1.5];
let w_vals = [0.1f32, 0.2, 0.3, -0.4, 0.5, -0.6];
let b_vals = [0.05f32, -0.05];
let st_vals = [1.0f32, 2.0, -3.0, 4.0];
let owned_x = || Owned::f32(&[batch, c, s], &x_vals);
let owned_w = || Owned::f32(&[c, 1, k], &w_vals);
let owned_bias = || Owned::f32(&[c], &b_vals);
let owned_state = || Owned::f32(&[batch, c, pad], &st_vals);
let (y_ref, present_ref) = {
let (x, w, bias, state) = (owned_x(), owned_w(), owned_bias(), owned_state());
let mut y = Owned::zeros_f32(&[batch, c, s]);
let mut present = Owned::zeros_f32(&[batch, c, pad]);
kernel(true)
.execute(
&[x.view(), w.view(), bias.view(), state.view()],
&mut [y.view_mut(), present.view_mut()],
)
.unwrap();
(y.to_f32(), present.to_f32())
};
let f32c = DataType::Float32;
let cpu = DeviceId::cpu();
{
let mut shared_state = st_vals.to_vec();
let state_ptr = shared_state.as_ptr() as *const std::ffi::c_void;
let present_ptr = shared_state.as_mut_ptr() as *mut std::ffi::c_void;
let sshape = [batch, c, pad];
let sstrides = compute_contiguous_strides(&sshape);
let (x, w, bias) = (owned_x(), owned_w(), owned_bias());
let mut y = Owned::zeros_f32(&[batch, c, s]);
let state_view = TensorView::new(DevicePtr(state_ptr), f32c, &sshape, &sstrides, cpu);
let present_mut =
TensorMut::new(DevicePtrMut(present_ptr), f32c, &sshape, &sstrides, cpu);
kernel(true)
.execute(
&[x.view(), w.view(), bias.view(), state_view],
&mut [y.view_mut(), present_mut],
)
.unwrap();
assert_close(&y.to_f32(), &y_ref, "present-aliases-state y");
assert_close(&shared_state, &present_ref, "present-aliases-state present");
}
{
let mut shared_x = x_vals.to_vec();
let x_ptr = shared_x.as_ptr() as *const std::ffi::c_void;
let y_ptr = shared_x.as_mut_ptr() as *mut std::ffi::c_void;
let xshape = [batch, c, s];
let xstrides = compute_contiguous_strides(&xshape);
let (w, bias, state) = (owned_w(), owned_bias(), owned_state());
let mut present = Owned::zeros_f32(&[batch, c, pad]);
let x_view = TensorView::new(DevicePtr(x_ptr), f32c, &xshape, &xstrides, cpu);
let y_mut = TensorMut::new(DevicePtrMut(y_ptr), f32c, &xshape, &xstrides, cpu);
kernel(true)
.execute(
&[x_view, w.view(), bias.view(), state.view()],
&mut [y_mut, present.view_mut()],
)
.unwrap();
assert_close(&shared_x, &y_ref, "y-aliases-x y");
assert_close(&present.to_f32(), &present_ref, "y-aliases-x present");
}
}
}