use crate::context::InferenceContext;
use crate::dim_expr::DimExpr;
use crate::error::ShapeInferError;
use crate::registry::InferenceRegistry;
fn const_ints(ctx: &InferenceContext, input: usize) -> Option<Vec<i64>> {
ctx.input_shape_data(input)?
.elems
.iter()
.map(DimExpr::as_const)
.collect()
}
fn const_float_scalar(ctx: &InferenceContext, input: usize) -> Option<f64> {
ctx.input_shape_data(input)?.as_float_scalar()
}
fn tile(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let Some(input) = ctx.input_type(0).cloned() else {
return Ok(());
};
if let Some(rank) = ctx.input_rank(1)
&& rank != 1
{
return Err(ShapeInferError::InvalidRank {
op: "Tile".into(),
index: 1,
rank,
detail: "repeats must be a 1-D tensor".into(),
});
}
let repeats = const_ints(ctx, 1);
if let Some(repeats) = &repeats
&& repeats.len() != input.rank()
{
return Err(ShapeInferError::Invalid {
op: "Tile".into(),
detail: format!(
"repeats has {} values but input rank is {}",
repeats.len(),
input.rank()
),
});
}
let mut output = Vec::with_capacity(input.rank());
for (i, dim) in input.shape.iter().enumerate() {
let output_dim = match repeats.as_ref().and_then(|values| values.get(i)).copied() {
Some(repeat) if repeat >= 0 => {
if let Some(extent) = dim.as_const() {
let product = i128::from(extent) * i128::from(repeat);
if product > isize::MAX as i128 {
return Err(ShapeInferError::Invalid {
op: "Tile".into(),
detail: format!("output extent {extent} * {repeat} exceeds isize::MAX"),
});
}
}
dim.mul(&DimExpr::constant(repeat))
}
Some(_) | None => ctx.fresh_dim(),
};
output.push(output_dim);
}
ctx.set_output(0, input.dtype, output);
Ok(())
}
enum RangeLength {
Known(i64),
Unknown,
TooLarge,
}
fn range_len(start: i64, limit: i64, delta: i64) -> RangeLength {
if delta == 0 {
return RangeLength::Unknown;
}
let start = i128::from(start);
let limit = i128::from(limit);
let delta = i128::from(delta);
let distance = limit - start;
let count = if (distance > 0 && delta > 0) || (distance < 0 && delta < 0) {
let distance = distance.abs();
let step = delta.abs();
(distance + step - 1) / step
} else {
0
};
if count > isize::MAX as i128 {
RangeLength::TooLarge
} else {
RangeLength::Known(count as i64)
}
}
fn range_len_f32(start: f64, limit: f64, delta: f64) -> RangeLength {
let start = start as f32;
let limit = limit as f32;
let delta = delta as f32;
if delta == 0.0 || !start.is_finite() || !limit.is_finite() || !delta.is_finite() {
return RangeLength::Unknown;
}
let count = ((limit - start) / delta).ceil().max(0.0);
if !count.is_finite() || count >= isize::MAX as f32 {
return RangeLength::TooLarge;
}
RangeLength::Known(count as i64)
}
fn range_len_f64(start: f64, limit: f64, delta: f64) -> RangeLength {
if delta == 0.0 || !start.is_finite() || !limit.is_finite() || !delta.is_finite() {
return RangeLength::Unknown;
}
let count = ((limit - start) / delta).ceil().max(0.0);
if !count.is_finite() || count >= isize::MAX as f64 {
return RangeLength::TooLarge;
}
RangeLength::Known(count as i64)
}
fn range(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
let Some(dtype) = ctx.input_dtype(0) else {
return Ok(());
};
let length = range_int_len(ctx)
.or_else(|| range_float_len(ctx, dtype))
.unwrap_or(RangeLength::Unknown);
let dim = match length {
RangeLength::Known(length) => DimExpr::constant(length),
RangeLength::Unknown => ctx.fresh_dim(),
RangeLength::TooLarge => {
return Err(ShapeInferError::Invalid {
op: "Range".into(),
detail: "output length exceeds isize::MAX".into(),
});
}
};
ctx.set_output(0, dtype, vec![dim]);
Ok(())
}
fn range_int_len(ctx: &InferenceContext) -> Option<RangeLength> {
match (const_ints(ctx, 0), const_ints(ctx, 1), const_ints(ctx, 2)) {
(Some(start), Some(limit), Some(delta))
if start.len() == 1 && limit.len() == 1 && delta.len() == 1 =>
{
Some(range_len(start[0], limit[0], delta[0]))
}
_ => None,
}
}
fn range_float_len(
ctx: &InferenceContext,
dtype: onnx_runtime_ir::DataType,
) -> Option<RangeLength> {
let start = const_float_scalar(ctx, 0)?;
let limit = const_float_scalar(ctx, 1)?;
let delta = const_float_scalar(ctx, 2)?;
match dtype {
onnx_runtime_ir::DataType::Float32 => Some(range_len_f32(start, limit, delta)),
onnx_runtime_ir::DataType::Float64 => Some(range_len_f64(start, limit, delta)),
_ => None,
}
}
fn cum_sum(ctx: &mut InferenceContext) -> Result<(), ShapeInferError> {
if let Some(input) = ctx.input_type(0).cloned() {
ctx.set_output_type(0, input);
}
Ok(())
}
pub fn register(reg: &mut InferenceRegistry) {
reg.register("", "Tile", 6, tile);
reg.register("", "Range", 11, range);
reg.register("", "CumSum", 11, cum_sum);
reg.register("", "CumSum", 14, cum_sum);
}