use onnx_runtime_ir::{Attribute, DataType};
use crate::context::InferenceContext;
use crate::dim_expr::DimExpr;
use crate::error::ShapeInferError;
use crate::registry::InferenceRegistry;
fn scalar_int(ctx: &InferenceContext, index: usize) -> Option<i64> {
let data = ctx.input_shape_data(index)?;
if !data.is_scalar() {
return None;
}
data.elems.first()?.as_const()
}
fn output_datatype(ctx: &InferenceContext) -> DataType {
ctx.node
.attr("output_datatype")
.and_then(Attribute::as_int)
.and_then(|raw| i32::try_from(raw).ok())
.and_then(DataType::from_onnx)
.unwrap_or(DataType::Float32)
}
fn flag(ctx: &InferenceContext, name: &str) -> bool {
ctx.node.attr(name).and_then(Attribute::as_int).unwrap_or(0) != 0
}
fn onesided_bins(n: i64) -> i64 {
(n >> 1) + 1
}
fn dft_axis_index(axis: i64, rank: usize) -> Option<usize> {
let r = rank as i64;
if axis < -r || axis == -1 || axis >= r - 1 {
return None;
}
let normalized = if axis >= 0 { axis } else { axis + r };
usize::try_from(normalized).ok()
}
fn dft(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let Some(input) = ctx.input_type(0).cloned() else {
return Ok(());
};
let dtype = input.dtype;
let rank = input.shape.len();
if rank < 2 {
return Err(ShapeInferError::InvalidRank {
op: "DFT".into(),
index: 0,
rank,
detail: "input must have rank >= 2 (including the complex dimension)".into(),
});
}
let onesided = flag(ctx, "onesided");
let has_dft_length = ctx.has_input(1);
let last = rank - 1;
let axis_is_input = ctx.opset("") >= 20 && ctx.has_input(2);
let axis_value: Option<i64> = if axis_is_input {
scalar_int(ctx, 2)
} else {
let default_axis = if ctx.opset("") >= 20 { -2 } else { 1 };
Some(
ctx.node
.attr("axis")
.and_then(Attribute::as_int)
.unwrap_or(default_axis),
)
};
let mut out = input.shape.clone();
let Some(axis) = axis_value else {
if onesided || has_dft_length {
out = (0..rank).map(|_| ctx.fresh_dim()).collect();
}
out[last] = DimExpr::constant(2);
ctx.set_output(0, dtype, out);
return Ok(());
};
let Some(axis_idx) = dft_axis_index(axis, rank) else {
return Err(ShapeInferError::Invalid {
op: "DFT".into(),
detail: format!("axis {axis} is invalid for a tensor of rank {rank}"),
});
};
if has_dft_length {
out[axis_idx] = match scalar_int(ctx, 1) {
Some(length) => DimExpr::constant(length),
None => ctx.fresh_dim(),
};
}
if onesided {
out[axis_idx] = match out[axis_idx].as_const() {
Some(n) => DimExpr::constant(onesided_bins(n)),
None => ctx.fresh_dim(),
};
}
out[last] = DimExpr::constant(2);
ctx.set_output(0, dtype, out);
Ok(())
}
fn stft(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let Some(signal) = ctx.input_type(0).cloned() else {
return Ok(());
};
let dtype = signal.dtype;
let rank = signal.shape.len();
if rank != 3 {
return Err(ShapeInferError::InvalidRank {
op: "STFT".into(),
index: 0,
rank,
detail: "signal must have rank 3: [batch, signal_length, 1|2]".into(),
});
}
if let Some(components) = signal.shape[2].as_const()
&& components != 1
&& components != 2
{
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!(
"signal's last dimension must be 1 (real) or 2 (complex), got {components}"
),
});
}
let batch = signal.shape[0].clone();
let signal_length = signal.shape[1].as_const();
if let Some(shape) = ctx.input_shape(1)
&& !shape.is_empty()
{
return Err(ShapeInferError::InvalidRank {
op: "STFT".into(),
index: 1,
rank: shape.len(),
detail: "frame_step must be a scalar".into(),
});
}
let frame_step = scalar_int(ctx, 1);
if let Some(step) = frame_step
&& step <= 0
{
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!("frame_step must be greater than zero, got {step}"),
});
}
let has_window = ctx.has_input(2);
let has_frame_length = ctx.has_input(3);
if !has_window && !has_frame_length {
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: "either optional window or frame_length must be provided".into(),
});
}
let window_length = if has_window {
let shape = ctx.input_shape(2).ok_or_else(|| ShapeInferError::Invalid {
op: "STFT".into(),
detail: "window input has no tensor shape".into(),
})?;
if shape.len() != 1 {
return Err(ShapeInferError::InvalidRank {
op: "STFT".into(),
index: 2,
rank: shape.len(),
detail: "window must have rank 1".into(),
});
}
let length = shape[0].as_const();
if let Some(length) = length
&& length <= 0
{
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!("window length must be greater than zero, got {length}"),
});
}
length
} else {
None
};
let frame_length = if has_frame_length {
if let Some(shape) = ctx.input_shape(3)
&& !shape.is_empty()
{
return Err(ShapeInferError::InvalidRank {
op: "STFT".into(),
index: 3,
rank: shape.len(),
detail: "frame_length must be a scalar".into(),
});
}
let length = scalar_int(ctx, 3);
if let Some(length) = length
&& length <= 0
{
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!("frame_length must be greater than zero, got {length}"),
});
}
length
} else {
None
};
if let (Some(window), Some(frame)) = (window_length, frame_length)
&& window != frame
{
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!("window length {window} must equal frame_length {frame}"),
});
}
let dft_size = if has_frame_length {
frame_length
} else {
window_length
};
let onesided = ctx
.node
.attr("onesided")
.and_then(Attribute::as_int)
.unwrap_or(1)
!= 0;
if onesided && signal.shape[2].as_const() == Some(2) {
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: "onesided=1 requires a real signal (last dimension 1)".into(),
});
}
let bins = dft_size.map(|size| if onesided { onesided_bins(size) } else { size });
let frames = match (signal_length, dft_size, frame_step) {
(Some(length), Some(size), Some(_)) if size > length => {
return Err(ShapeInferError::Invalid {
op: "STFT".into(),
detail: format!(
"frame length {size} exceeds signal length {length}; STFT uses complete unpadded frames"
),
});
}
(Some(length), Some(size), Some(step)) => Some((length - size) / step + 1),
_ => None,
};
let frames_dim = frames.map_or_else(|| ctx.fresh_dim(), DimExpr::constant);
let bins_dim = bins.map_or_else(|| ctx.fresh_dim(), DimExpr::constant);
ctx.set_output(
0,
dtype,
vec![batch, frames_dim, bins_dim, DimExpr::constant(2)],
);
Ok(())
}
fn mel_weight_matrix(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let dtype = output_datatype(ctx);
let rows = match scalar_int(ctx, 1) {
Some(dft_length) if dft_length > 0 => DimExpr::constant(onesided_bins(dft_length)),
_ => ctx.fresh_dim(),
};
let cols = match scalar_int(ctx, 0) {
Some(num_mel_bins) if num_mel_bins > 0 => DimExpr::constant(num_mel_bins),
_ => ctx.fresh_dim(),
};
ctx.set_output(0, dtype, vec![rows, cols]);
Ok(())
}
fn window(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let dtype = output_datatype(ctx);
let length = match scalar_int(ctx, 0) {
Some(size) if size > 0 => DimExpr::constant(size),
_ => ctx.fresh_dim(),
};
ctx.set_output(0, dtype, vec![length]);
Ok(())
}
pub fn register(reg: &mut InferenceRegistry) {
reg.register("", "DFT", 17, dft);
reg.register("", "DFT", 20, dft);
reg.register("", "STFT", 17, stft);
reg.register("", "MelWeightMatrix", 17, mel_weight_matrix);
reg.register("", "HannWindow", 17, window);
reg.register("", "HammingWindow", 17, window);
reg.register("", "BlackmanWindow", 17, window);
}