use std::collections::BTreeMap;
use onnx_runtime_ir::SymbolId;
type Monomial = Vec<u32>;
#[derive(Clone, PartialEq, Eq, Hash, Debug, Default)]
pub struct DimExpr {
terms: BTreeMap<Monomial, i64>,
overflow: bool,
}
impl DimExpr {
pub fn constant(n: i64) -> Self {
let mut terms = BTreeMap::new();
if n != 0 {
terms.insert(Vec::new(), n);
}
Self {
terms,
overflow: false,
}
}
pub fn symbol(s: SymbolId) -> Self {
let mut terms = BTreeMap::new();
terms.insert(vec![s.0], 1);
Self {
terms,
overflow: false,
}
}
pub fn overflow() -> Self {
Self {
terms: BTreeMap::new(),
overflow: true,
}
}
pub fn is_overflow(&self) -> bool {
self.overflow
}
pub fn as_const(&self) -> Option<i64> {
if self.overflow {
return None;
}
match self.terms.len() {
0 => Some(0),
1 => self.terms.get(&Vec::new()).copied(),
_ => None,
}
}
pub fn as_symbol(&self) -> Option<SymbolId> {
if self.overflow || self.terms.len() != 1 {
return None;
}
let (mono, &coeff) = self.terms.iter().next()?;
if coeff == 1 && mono.len() == 1 {
Some(SymbolId(mono[0]))
} else {
None
}
}
pub fn is_const(&self) -> bool {
self.as_const().is_some()
}
fn prune(mut self) -> Self {
self.terms.retain(|_, c| *c != 0);
self
}
pub fn add(&self, other: &DimExpr) -> DimExpr {
if self.overflow || other.overflow {
return DimExpr::overflow();
}
let mut terms = self.terms.clone();
for (mono, &coeff) in &other.terms {
let slot = terms.entry(mono.clone()).or_insert(0);
match slot.checked_add(coeff) {
Some(v) => *slot = v,
None => return DimExpr::overflow(),
}
}
DimExpr {
terms,
overflow: false,
}
.prune()
}
pub fn sub(&self, other: &DimExpr) -> DimExpr {
if self.overflow || other.overflow {
return DimExpr::overflow();
}
let mut terms = self.terms.clone();
for (mono, &coeff) in &other.terms {
let slot = terms.entry(mono.clone()).or_insert(0);
match slot.checked_sub(coeff) {
Some(v) => *slot = v,
None => return DimExpr::overflow(),
}
}
DimExpr {
terms,
overflow: false,
}
.prune()
}
pub fn mul(&self, other: &DimExpr) -> DimExpr {
if self.overflow || other.overflow {
return DimExpr::overflow();
}
let mut terms: BTreeMap<Monomial, i64> = BTreeMap::new();
for (a_mono, &a_c) in &self.terms {
for (b_mono, &b_c) in &other.terms {
let Some(prod) = a_c.checked_mul(b_c) else {
return DimExpr::overflow();
};
let mut mono = a_mono.clone();
mono.extend_from_slice(b_mono);
mono.sort_unstable();
let slot = terms.entry(mono).or_insert(0);
match slot.checked_add(prod) {
Some(v) => *slot = v,
None => return DimExpr::overflow(),
}
}
}
DimExpr {
terms,
overflow: false,
}
.prune()
}
pub fn checked_div(&self, other: &DimExpr) -> Option<DimExpr> {
if self.overflow || other.overflow {
return None;
}
if self.terms.is_empty() {
return Some(DimExpr::constant(0));
}
if other.terms.len() != 1 {
return None;
}
let (div_mono, &div_coeff) = other.terms.iter().next()?;
if div_coeff == 0 {
return None;
}
let mut out: BTreeMap<Monomial, i64> = BTreeMap::new();
for (mono, &coeff) in &self.terms {
if coeff.checked_rem(div_coeff)? != 0 {
return None;
}
let mut remaining = mono.clone();
for sym in div_mono {
let pos = remaining.iter().position(|s| s == sym)?;
remaining.remove(pos);
}
out.insert(remaining, coeff.checked_div(div_coeff)?);
}
Some(
DimExpr {
terms: out,
overflow: false,
}
.prune(),
)
}
pub fn product(exprs: &[DimExpr]) -> DimExpr {
let mut acc = DimExpr::constant(1);
for e in exprs {
acc = acc.mul(e);
}
acc
}
}
impl From<onnx_runtime_ir::Dim> for DimExpr {
fn from(d: onnx_runtime_ir::Dim) -> Self {
match d {
onnx_runtime_ir::Dim::Static(n) => DimExpr::constant(n as i64),
onnx_runtime_ir::Dim::Symbolic(s) => DimExpr::symbol(s),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(n: u32) -> DimExpr {
DimExpr::symbol(SymbolId(n))
}
#[test]
fn constant_folding() {
let e = DimExpr::constant(3).add(&DimExpr::constant(4));
assert_eq!(e.as_const(), Some(7));
assert!(e.is_const());
}
#[test]
fn zero_is_canonical() {
assert_eq!(DimExpr::constant(0).as_const(), Some(0));
assert_eq!(
DimExpr::constant(5).sub(&DimExpr::constant(5)).as_const(),
Some(0)
);
}
#[test]
fn symbol_roundtrip() {
let e = sym(2);
assert_eq!(e.as_symbol(), Some(SymbolId(2)));
assert_eq!(e.add(&sym(2)).as_symbol(), None);
}
#[test]
fn affine_expression() {
let e = sym(0).add(&DimExpr::constant(5));
assert_eq!(e.as_const(), None);
assert_eq!(e.as_symbol(), None);
assert_eq!(e.sub(&DimExpr::constant(5)).as_symbol(), Some(SymbolId(0)));
}
#[test]
fn product_of_symbols() {
let e = sym(0).mul(&sym(1));
assert_eq!(e, sym(1).mul(&sym(0)));
}
#[test]
fn exact_constant_division() {
let e = DimExpr::constant(48);
assert_eq!(
e.checked_div(&DimExpr::constant(6)).unwrap().as_const(),
Some(8)
);
assert!(
DimExpr::constant(7)
.checked_div(&DimExpr::constant(2))
.is_none()
);
}
#[test]
fn reshape_minus_one_cancellation() {
let b = sym(0);
let s = sym(1);
let total = DimExpr::product(&[b.clone(), s.clone(), DimExpr::constant(768)]);
let known = DimExpr::product(&[b, s, DimExpr::constant(12)]);
let missing = total.checked_div(&known).unwrap();
assert_eq!(missing.as_const(), Some(64));
}
#[test]
fn division_by_multiterm_is_none() {
let total = sym(0).mul(&sym(1));
let divisor = sym(0).add(&DimExpr::constant(1));
assert!(total.checked_div(&divisor).is_none());
}
#[test]
fn symbolic_product_division_keeps_symbol() {
let e = sym(0).mul(&DimExpr::constant(768));
let q = e.checked_div(&DimExpr::constant(768)).unwrap();
assert_eq!(q.as_symbol(), Some(SymbolId(0)));
}
#[test]
fn from_ir_dim() {
use onnx_runtime_ir::Dim;
assert_eq!(DimExpr::from(Dim::Static(4)).as_const(), Some(4));
assert_eq!(
DimExpr::from(Dim::Symbolic(SymbolId(9))).as_symbol(),
Some(SymbolId(9))
);
}
#[test]
fn mul_overflow_degrades_to_unknown() {
let big = DimExpr::constant(1 << 20);
let total = DimExpr::product(&[big.clone(), big.clone(), big.clone(), big]);
assert!(total.is_overflow());
assert_eq!(total.as_const(), None); assert_eq!(total.as_symbol(), None);
}
#[test]
fn add_and_sub_overflow_degrade() {
let max = DimExpr::constant(i64::MAX);
assert!(max.add(&DimExpr::constant(1)).is_overflow());
let min = DimExpr::constant(i64::MIN);
assert!(min.sub(&DimExpr::constant(1)).is_overflow());
}
#[test]
fn overflow_poisons_further_arithmetic() {
let poisoned = DimExpr::overflow();
assert!(poisoned.add(&DimExpr::constant(1)).is_overflow());
assert!(poisoned.mul(&DimExpr::constant(2)).is_overflow());
assert!(poisoned.sub(&DimExpr::constant(3)).is_overflow());
assert!(poisoned.checked_div(&DimExpr::constant(4)).is_none());
assert!(DimExpr::constant(8).checked_div(&poisoned).is_none());
}
#[test]
fn checked_div_guards_i64_min_over_neg_one() {
let num = DimExpr::constant(i64::MIN);
let div = DimExpr::constant(-1);
assert!(num.checked_div(&div).is_none());
}
}