Skip to main content

tract_core/ops/
change_axes.rs

1use std::borrow::Borrow;
2use std::fmt::Debug;
3
4use crate::internal::*;
5use crate::model::{TypedModel, TypedNode};
6use crate::ops::identity::Identity;
7use AxisOp::*;
8use tract_itertools::Itertools;
9use tract_linalg::block_quant::{BlockQuantFact, BlockQuantStorage};
10use tract_ndarray::{ArrayViewD, ArrayViewMutD};
11
12#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
13pub enum InOut {
14    Out(usize),
15    In(usize),
16}
17
18impl InOut {
19    pub fn as_outlet<F: Clone + Fact, O: Clone>(&self, node: &Node<F, O>) -> OutletId {
20        match self {
21            InOut::In(ix) => node.inputs[*ix],
22            InOut::Out(ix) => OutletId::new(node.id, *ix),
23        }
24    }
25
26    pub fn is_input(&self) -> bool {
27        matches!(self, InOut::In(_))
28    }
29
30    pub fn is_output(&self) -> bool {
31        matches!(self, InOut::Out(_))
32    }
33
34    pub fn slot(&self) -> usize {
35        match self {
36            InOut::Out(o) => *o,
37            InOut::In(i) => *i,
38        }
39    }
40}
41
42#[derive(Clone, Hash, Eq)]
43#[allow(clippy::large_enum_variant)] // FIXME ?
44#[allow(clippy::derived_hash_with_manual_eq)] // FIXME. this one may be pretty bad. how about a.canonical() == b.canonical() ? need proper canonicalizeation of Reshape
45pub enum AxisOp {
46    Add(usize),
47    Rm(usize),
48    Move(usize, usize),
49    Reshape(usize, TVec<TDim>, TVec<TDim>),
50}
51
52impl Debug for AxisOp {
53    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54        match self {
55            AxisOp::Add(a) => write!(f, "Add({a})"),
56            AxisOp::Rm(a) => write!(f, "Rm({a})"),
57            AxisOp::Move(from, to) => write!(f, "Move({from},{to})"),
58            AxisOp::Reshape(at, from, to) => {
59                write!(f, "Reshape({at}, [{}], [{}])", from.iter().join(","), to.iter().join(","))
60            }
61        }
62    }
63}
64
65impl PartialEq for AxisOp {
66    fn eq(&self, other: &AxisOp) -> bool {
67        if self.is_noop() && other.is_noop() {
68            true
69        } else if self.is_noop() != other.is_noop() {
70            false
71        } else {
72            match (self, other) {
73                (Add(a), Add(b)) | (Rm(a), Rm(b)) => a == b,
74                (Move(f1, t1), Move(f2, t2)) => {
75                    (f1 == f2 && t1 == t2)
76                        || ((*t1 == f1 + 1 || *f1 == t1 + 1) && t2 == f1 && t1 == f2)
77                }
78                (Reshape(at1, f1, t1), Reshape(at2, f2, t2)) => at1 == at2 && f1 == f2 && t1 == t2,
79                _ => false,
80            }
81        }
82    }
83}
84
85impl AxisOp {
86    pub fn canonical(&self) -> Cow<'_, AxisOp> {
87        match self {
88            Move(from, to) if *from == to + 1 => Cow::Owned(Move(*to, *from)),
89            Reshape(at, from, to)
90                if from.len() == 1 && to.len() == 2 && from[0] == to[0] && to[1].is_one() =>
91            {
92                Cow::Owned(Add(*at + 1))
93            }
94            Reshape(at, from, to)
95                if from.len() == 1 && to.len() == 2 && from[0] == to[1] && to[0].is_one() =>
96            {
97                Cow::Owned(Add(*at))
98            }
99            Reshape(at, from, to)
100                if from.len() == 2 && to.len() == 1 && from[0] == to[0] && from[1].is_one() =>
101            {
102                Cow::Owned(Rm(*at + 1))
103            }
104            Reshape(at, from, to)
105                if from.len() == 2 && to.len() == 1 && from[1] == to[0] && from[0].is_one() =>
106            {
107                Cow::Owned(Rm(*at))
108            }
109            other => Cow::Borrowed(other),
110        }
111    }
112
113    pub fn simplify(&self) -> TVec<AxisOp> {
114        match self.canonical().borrow() {
115            Reshape(_, from, to) if from == to => tvec!(),
116            Reshape(at, from, to) if to.len() == 0 => tvec!(Rm(*at); from.len()),
117            Reshape(at, from, to) if from.len() == 0 => tvec!(Add(*at); to.len()),
118            Reshape(at, from, to) if from[0] == to[0] => {
119                Reshape(at + 1, from[1..].into(), to[1..].into()).simplify()
120            }
121            Reshape(at, from, to) if from[from.len() - 1] == to[to.len() - 1] => {
122                Reshape(*at, from[..from.len() - 1].into(), to[..to.len() - 1].into()).simplify()
123            }
124            Reshape(at, from, to) if from[0] == 1.to_dim() => std::iter::once(Rm(*at))
125                .chain(Reshape(*at, from[1..].into(), to.clone()).simplify())
126                .collect(),
127            Reshape(at, from, to) if to[0] == 1.to_dim() => {
128                Reshape(*at, from.clone(), to[1..].into())
129                    .simplify()
130                    .into_iter()
131                    .chain(std::iter::once(Add(*at)))
132                    .collect()
133            }
134            Reshape(at, from, to) if from[from.len() - 1] == 1.to_dim() => {
135                std::iter::once(Rm(at + from.len() - 1))
136                    .chain(Reshape(*at, from[..from.len() - 1].into(), to.clone()).simplify())
137                    .collect()
138            }
139            Reshape(at, from, to) if to[to.len() - 1] == 1.to_dim() => {
140                std::iter::once(Add(at + from.len()))
141                    .chain(Reshape(*at, from.clone(), to[..to.len() - 1].into()).simplify())
142                    .collect()
143            }
144            other => tvec!(other.clone()),
145        }
146    }
147
148    pub fn transform_axis(&self, axis: usize) -> Option<usize> {
149        match self.canonical().as_ref() {
150            Add(ix) => Some(axis + (axis >= *ix) as usize),
151            Rm(ix) => {
152                if axis == *ix {
153                    None
154                } else {
155                    Some(axis - (axis > *ix) as usize)
156                }
157            }
158            Move(from, to) if from < to => {
159                if axis < *from || axis > *to {
160                    Some(axis)
161                } else if axis == *from {
162                    Some(*to)
163                } else {
164                    Some(axis - 1)
165                }
166            }
167            Move(from, to) => {
168                if axis < *to || axis > *from {
169                    Some(axis)
170                } else if axis == *from {
171                    Some(*to)
172                } else {
173                    Some(axis + 1)
174                }
175            }
176            Reshape(at, _, _) if axis < *at => Some(axis),
177            Reshape(at, from, to) if axis >= at + from.len() => Some(axis + to.len() - from.len()),
178            Reshape(_, _, _) => None,
179        }
180    }
181
182    // if sucessful return Some()
183    // first item is the Op we want to be replaced by. if none, we are now identity.
184    // second item is the change to propagate. if none, the output is not
185    // changed
186    pub fn merge_incoming_change(
187        &self,
188        change: &AxisOp,
189    ) -> Option<(Option<AxisOp>, Option<AxisOp>)> {
190        match (self.canonical().as_ref(), change.canonical().as_ref()) {
191            (Add(op), Add(c)) => {
192                Some((Some(Add(op + (c < op) as usize)), Some(Add(c + (c >= op) as usize))))
193            }
194            (Add(op), Rm(c)) => {
195                Some((Some(Add(op - (c < op) as usize)), Some(Rm(c + (c >= op) as usize))))
196            }
197            (Rm(op), Add(c)) => {
198                Some((Some(Rm(op + (c <= op) as usize)), Some(Add(c - (op < c) as usize))))
199            }
200            (Rm(op), Rm(c)) => {
201                Some((Some(Rm(op - (c < op) as usize)), Some(Rm(c - (op <= c) as usize))))
202            }
203
204            (Add(x), Move(from, to)) => {
205                if x <= from.min(to) {
206                    Some((Some(self.clone()), Some(Move(from + 1, to + 1))))
207                } else if x > from.max(to) {
208                    Some((Some(self.clone()), Some(change.clone())))
209                } else {
210                    None
211                }
212            }
213
214            (Move(from, to), Add(x)) => {
215                if x <= from.min(to) {
216                    Some((Some(Move(from + 1, to + 1)), Some(Add(*x))))
217                } else if x > from.max(to) {
218                    Some((Some(Move(*from, *to)), Some(Add(*x))))
219                } else {
220                    None
221                }
222            }
223
224            (Rm(x), Move(from, to)) => {
225                if x == from {
226                    Some((Some(Rm(*to)), None))
227                } else if x < from.min(to) {
228                    Some((Some(self.clone()), Some(Move(from - 1, to - 1))))
229                } else if x > from.max(to) {
230                    Some((Some(self.clone()), Some(change.clone())))
231                } else if from + 1 == *to && x == to {
232                    Some((Some(Rm(*from)), None))
233                } else if from < to && x <= to {
234                    Some((Some(Rm(x - 1)), Some(Move(*from, *to - 1))))
235                } else {
236                    Some((Some(Rm(x + 1)), Some(Move(*from - 1, *to))))
237                }
238            }
239
240            (Move(from, to), Rm(x)) => {
241                if x < from.min(to) {
242                    Some((Some(Move(from - 1, to - 1)), Some(Rm(*x))))
243                } else if x > from.max(to) {
244                    Some((Some(Move(*from, *to)), Some(Rm(*x))))
245                } else {
246                    None
247                }
248            }
249
250            (Add(op), Reshape(at, from, to)) => {
251                if op <= at {
252                    Some((Some(Add(*op)), Some(Reshape(at + 1, from.clone(), to.clone()))))
253                } else if *op > at + from.len() {
254                    Some((
255                        Some(Add(*op + to.len() - from.len())),
256                        Some(Reshape(*at, from.clone(), to.clone())),
257                    ))
258                } else {
259                    None
260                }
261            }
262            (Rm(op), Reshape(at, from, to)) => {
263                if op < at {
264                    Some((Some(Rm(*op)), Some(Reshape(at - 1, from.clone(), to.clone()))))
265                } else if *op > at + from.len() {
266                    Some((
267                        Some(Rm(*op + to.len() - from.len())),
268                        Some(Reshape(*at, from.clone(), to.clone())),
269                    ))
270                } else {
271                    None
272                }
273            }
274            (Reshape(at, from, to), Add(change)) => {
275                if change < at {
276                    Some((Some(Reshape(at + 1, from.clone(), to.clone())), Some(Add(*change))))
277                } else if *change > *at + from.len() {
278                    Some((
279                        Some(Reshape(*at, from.clone(), to.clone())),
280                        Some(Add(change + to.len() - from.len())),
281                    ))
282                } else {
283                    None
284                }
285            }
286            (Reshape(at, from, to), Rm(change)) => {
287                if change < at {
288                    Some((Some(Reshape(at - 1, from.clone(), to.clone())), Some(Rm(*change))))
289                } else if *change > *at + from.len() {
290                    Some((
291                        Some(Reshape(*at, from.clone(), to.clone())),
292                        Some(Rm(change + to.len() - from.len())),
293                    ))
294                } else {
295                    None
296                }
297            }
298            (Reshape(_, _, _), Move(_, _)) => None, // todo, some are manageable
299            (Move(_, _), Reshape(_, _, _)) => None, // todo, some are manageable
300            (Reshape(_, _, _), Reshape(_, _, _)) => None, // todo, some are manageable
301            _ => None,
302        }
303    }
304
305    pub fn change_shape_array<D: DimLike>(
306        &self,
307        shape: &mut TVec<D>,
308        broadcasting: bool,
309    ) -> TractResult<()> {
310        match self.canonical().as_ref() {
311            Add(ix) => {
312                ensure!(*ix <= shape.len());
313                shape.insert(*ix, D::one());
314            }
315            Rm(ix) => {
316                ensure!(*ix < shape.len());
317                shape.remove(*ix);
318            }
319            Move(from, to) => {
320                ensure!(*from < shape.len());
321                ensure!(*to < shape.len());
322                let axis = shape.remove(*from);
323                shape.insert(*to, axis);
324            }
325            Reshape(at, from, to) => {
326                let from_volume = from.iter().product::<TDim>();
327                let to_volume = to.iter().product::<TDim>();
328                // Two algebraically equal volumes can land in different
329                // factored forms when the same dimension is built two ways
330                // (e.g. (B+2BY)·(1+Y) vs B·(1+Y)·(1+2Y) on Conformer-style
331                // streaming attention).  Compare polynomial expansions so
332                // structural mismatch on factor ordering doesn't fail the
333                // check.
334                ensure!(
335                    from_volume.clone().expand_polynomial()
336                        == to_volume.clone().expand_polynomial(),
337                    "{from_volume} should be equal to {to_volume}"
338                );
339                ensure!(*at + from.len() <= shape.len());
340                if shape.len() >= from.len() + *at
341                    && tract_itertools::izip!(shape.iter().skip(*at), from)
342                        .all(|(shape, spec)| shape.to_dim() == *spec)
343                {
344                    for _ in from {
345                        shape.remove(*at);
346                    }
347                    for d in to.iter().rev() {
348                        shape.insert(*at, d.try_into()?);
349                    }
350                } else if broadcasting
351                    && shape.iter().skip(*at).take(from.len()).all(|d| d.to_dim() == 1.to_dim())
352                {
353                    for _ in from {
354                        shape.remove(*at);
355                    }
356                    for _ in to.iter().rev() {
357                        shape.insert(*at, 1.into());
358                    }
359                } else {
360                    bail!("Incompatible reshape for shape {:?} and {:?}", shape, self);
361                }
362            }
363        }
364        Ok(())
365    }
366
367    pub fn change_shape(&self, shape: &mut ShapeFact, broadcasting: bool) -> TractResult<()> {
368        match self.canonical().as_ref() {
369            Add(ix) => {
370                if *ix > shape.rank() {
371                    bail!("Attempt to insert axis #{} on shape {:?}", ix, shape);
372                }
373                shape.insert_axis(*ix)
374            }
375            Rm(ix) => {
376                if shape.rank() <= *ix {
377                    bail!("Attempt to remove axis #{} on shape {:?}", ix, shape);
378                }
379                if shape[*ix] != 1.to_dim() {
380                    bail!("Removing non-trivial axis #{} of dim: {:?}", ix, shape);
381                }
382                shape.remove_axis(*ix)
383            }
384            _ => {
385                let mut array = shape.to_tvec();
386                self.change_shape_array(&mut array, broadcasting)?;
387                let mut new_shape = ShapeFact::from_dims(array);
388                std::mem::swap(shape, &mut new_shape);
389                Ok(())
390            }
391        }
392    }
393
394    pub fn change_tensor(&self, tensor: &mut Tensor, broadcasting: bool) -> TractResult<()> {
395        if tensor.storage_as::<BlockQuantStorage>().is_some() {
396            let bqs = tensor.try_storage_as::<BlockQuantStorage>()?.clone();
397            let mut new_shape: TVec<usize> = tensor.shape().into();
398            self.change_shape_array(&mut new_shape, false)?;
399            let mut new_tensor = bqs.into_tensor_with_shape(tensor.datum_type(), &new_shape);
400            std::mem::swap(tensor, &mut new_tensor);
401            return Ok(());
402        }
403        ensure!(self.required_rank() <= tensor.rank());
404        match self.canonical().as_ref() {
405            Add(ix) => tensor.insert_axis(*ix),
406            Rm(ix) => tensor.remove_axis(*ix),
407            Move(from, to) => {
408                let mut tmp = tensor.clone().move_axis(*from, *to)?;
409                std::mem::swap(tensor, &mut tmp);
410                Ok(())
411            }
412            Reshape(at, from, to) => {
413                let mut shape: TVec<usize> = tensor.shape().into();
414                self.change_shape_array(&mut shape, true)?;
415                if tensor.set_shape(&shape).is_ok() {
416                    Ok(())
417                } else if broadcasting
418                    && tensor.shape().iter().skip(*at).take(from.len()).all(|d| *d == 1)
419                {
420                    if from.len() > to.len() {
421                        for _ in to.len()..from.len() {
422                            tensor.remove_axis(*at)?;
423                        }
424                    }
425                    if to.len() > from.len() {
426                        for _ in from.len()..to.len() {
427                            tensor.insert_axis(*at)?;
428                        }
429                    }
430                    Ok(())
431                } else {
432                    bail!(
433                        "Invalid reshaping: {:?} on tensor {:?} (broadcasting allowed: {:?})",
434                        self,
435                        tensor,
436                        broadcasting
437                    )
438                }
439            }
440        }
441    }
442
443    pub fn change_view<D>(&self, view: &mut ArrayViewD<D>) -> TractResult<()> {
444        use tract_ndarray::Axis;
445        match *self {
446            AxisOp::Rm(axis) => view.index_axis_inplace(Axis(axis), 0),
447            AxisOp::Add(axis) => view.insert_axis_inplace(Axis(axis)),
448            AxisOp::Move(from, to) if from < to => {
449                for left in from..to {
450                    view.swap_axes(left, left + 1);
451                }
452            }
453            AxisOp::Move(from, to) => {
454                for left in (to..from).rev() {
455                    view.swap_axes(left, left + 1);
456                }
457            }
458            AxisOp::Reshape(_, _, _) => bail!("Reshape can not change views in place"),
459        }
460        Ok(())
461    }
462
463    pub fn change_view_mut<D>(&self, view: &mut ArrayViewMutD<D>) -> TractResult<()> {
464        use tract_ndarray::Axis;
465        match *self {
466            AxisOp::Rm(axis) => view.index_axis_inplace(Axis(axis), 0),
467            AxisOp::Add(axis) => view.insert_axis_inplace(Axis(axis)),
468            AxisOp::Move(from, to) if from < to => {
469                for left in from..to {
470                    view.swap_axes(left, left + 1);
471                }
472            }
473            AxisOp::Move(from, to) => {
474                for left in (to..from).rev() {
475                    view.swap_axes(left, left + 1);
476                }
477            }
478            AxisOp::Reshape(_, _, _) => bail!("Reshape can not change views in place"),
479        }
480        Ok(())
481    }
482
483    pub fn recip(&self) -> AxisOp {
484        match self.canonical().as_ref() {
485            Add(ix) => Rm(*ix),
486            Rm(ix) => Add(*ix),
487            Move(from, to) if from == to => self.clone(),
488            Move(from, to) if *from + 1 == *to => self.clone(),
489            Move(from, to) if *from == *to + 1 => {
490                unreachable!();
491            }
492            Move(from, to) => Move(*to, *from),
493            Reshape(at, from, to) => Reshape(*at, to.clone(), from.clone()),
494        }
495    }
496
497    pub fn is_noop(&self) -> bool {
498        match self {
499            Move(f, t) if f == t => true,
500            Reshape(_, f, t) if f == t => true,
501            _ => false,
502        }
503    }
504
505    pub fn only_shape(&self) -> bool {
506        if self.is_noop() {
507            return true;
508        }
509        !matches!(self, Move(_, _))
510    }
511
512    pub fn wire_split_axis(
513        model: &mut TypedModel,
514        name: impl ToString,
515        outlet: OutletId,
516        axis: usize,
517        outer_dim: usize,
518    ) -> TractResult<TVec<OutletId>> {
519        let fact = model.outlet_fact(outlet)?;
520        let dim: TDim = fact.shape[axis].clone();
521        let inner_dim = dim.clone() / outer_dim;
522        let op = Self::Reshape(axis, tvec!(dim.clone()), tvec!(outer_dim.to_dim(), inner_dim));
523        model.wire_node(name.to_string(), op, &[outlet])
524    }
525
526    pub fn wire_collapse_axis(
527        model: &mut TypedModel,
528        name: impl ToString,
529        outlet: OutletId,
530        axis: usize,
531    ) -> TractResult<TVec<OutletId>> {
532        let fact = model.outlet_fact(outlet)?;
533        let dim: TDim = fact.shape[axis].clone();
534        let next_dim: TDim = fact.shape[axis + 1].clone();
535        let op = Self::Reshape(axis, tvec!(dim.clone(), next_dim.clone()), tvec!(dim * next_dim));
536        model.wire_node(name.to_string(), op, &[outlet])
537    }
538
539    #[inline]
540    pub fn required_rank(&self) -> usize {
541        match self {
542            Rm(r) => r + 1,
543            Add(a) => *a,
544            Reshape(at, from, _to) => at + from.len(),
545            Move(from, to) => *from.max(to),
546        }
547    }
548
549    /// The same op on wires that gained `prefix` axes in front: every axis it
550    /// names moves right by as much, so none of them can land ahead of the new
551    /// axes. Inverse of [`AxisOp::trim_left`], and infallible where that one can
552    /// refuse.
553    pub fn pad_left(&self, prefix: usize) -> AxisOp {
554        match self {
555            Rm(r) => Rm(r + prefix),
556            Add(a) => Add(a + prefix),
557            Reshape(at, from, to) => Reshape(at + prefix, from.clone(), to.clone()),
558            Move(from, to) => Move(from + prefix, to + prefix),
559        }
560    }
561
562    pub fn trim_left(&self, prefix: usize) -> TractResult<AxisOp> {
563        Ok(match self {
564            Rm(r) if *r >= prefix => Rm(r - prefix),
565            Add(a) if *a >= prefix => Add(a - prefix),
566            Reshape(at, from, to) if *at >= prefix => {
567                Reshape(at - prefix, from.clone(), to.clone())
568            }
569            Move(from, to) if *from >= prefix && *to >= prefix => Move(from - prefix, to - prefix),
570            _ => bail!("Can no trim left {self:?} by {prefix}"),
571        })
572    }
573}
574
575pub fn wire_rank_broadcast(
576    prefix: impl AsRef<str>,
577    target: &mut TypedModel,
578    inputs: &[OutletId],
579) -> TractResult<TVec<OutletId>> {
580    let facts =
581        inputs.iter().map(|o| target.outlet_fact(*o).cloned()).collect::<TractResult<TVec<_>>>()?;
582    let max_rank = facts.iter().map(|f| f.rank()).max().unwrap();
583    let mut wires = tvec!();
584    for i in 0..inputs.len() {
585        let mut wire = inputs[i];
586        for _ in facts[i].rank()..max_rank {
587            let name = target.unique_name(prefix.as_ref().to_string() + ".fix-rank");
588            wire = target.wire_node(name, AxisOp::Add(0), &[wire])?[0];
589        }
590        wires.push(wire);
591    }
592    Ok(wires)
593}
594
595pub fn wire_with_rank_broadcast(
596    prefix: impl AsRef<str>,
597    target: &mut TypedModel,
598    op: impl Into<Box<dyn TypedOp>>,
599    inputs: &[OutletId],
600) -> TractResult<TVec<OutletId>> {
601    let prefix = prefix.as_ref();
602    let wires = wire_rank_broadcast(prefix, target, inputs)?;
603    target.wire_node(prefix, op.into(), &wires)
604}
605
606#[derive(Clone, Debug, PartialEq, Eq, Hash)]
607pub struct AxisChange {
608    pub outlet: OutletId,
609    pub op: AxisOp,
610}
611
612#[derive(Clone, Default, Debug)]
613pub struct AxisChangeConsequence {
614    pub substitute_op: Option<Box<dyn TypedOp>>,
615    pub wire_changes: TVec<(InOut, AxisOp)>,
616}
617
618impl AxisChangeConsequence {
619    pub fn new(
620        _model: &TypedModel,
621        node: &TypedNode,
622        op: Option<Box<dyn TypedOp>>,
623        axis_op: &AxisOp,
624    ) -> AxisChangeConsequence {
625        let mut wire_changes = tvec!();
626        for i in 0..node.inputs.len() {
627            wire_changes.push((InOut::In(i), axis_op.clone()));
628        }
629        for i in 0..node.outputs.len() {
630            wire_changes.push((InOut::Out(i), axis_op.clone()));
631        }
632        AxisChangeConsequence { wire_changes, substitute_op: op }
633    }
634}
635
636impl Op for AxisOp {
637    fn name(&self) -> StaticName {
638        match self {
639            Add(_) => "AddAxis".into(),
640            Rm(_) => "RmAxis".into(),
641            Move(_, _) => "MoveAxis".into(),
642            Reshape(_, _, _) => "Reshape".into(),
643        }
644    }
645
646    fn info(&self) -> TractResult<Vec<String>> {
647        match self {
648            Add(axis) | Rm(axis) => Ok(vec![format!("Axis: {axis}")]),
649            Move(from, to) => Ok(vec![format!("Axis {from} to {to}")]),
650            Reshape(at, from, to) => Ok(vec![format!(
651                "Axes starting at {}: {:?} to {:?}",
652                at,
653                from.iter().join(","),
654                to.iter().join(",")
655            )]),
656        }
657    }
658
659    op_as_typed_op!();
660}
661
662impl EvalOp for AxisOp {
663    op_out_of_plan!();
664
665    fn eval(&self, ctx: &EvalContext, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
666        let mut input = args_1!(inputs).into_tensor();
667        match self {
668            AxisOp::Reshape(skip, from, to) => {
669                let from = from.iter().map(|d| d.eval(ctx.symbols)).collect();
670                let to = to.iter().map(|d| d.eval(ctx.symbols)).collect();
671                AxisOp::Reshape(*skip, from, to).change_tensor(&mut input, false)?
672            }
673            _ => self.change_tensor(&mut input, false)?,
674        }
675        Ok(tvec!(input.into_tvalue()))
676    }
677}
678
679/// Remap coordinate symbols in a TDim expression according to an AxisOp.
680/// Returns None if the remapping cannot be determined (e.g. general reshape
681/// with both ends > 1).
682fn remap_uniform_tdim(expr: &TDim, axis_op: &AxisOp) -> Option<TDim> {
683    let syms = expr.symbols();
684    let coord_syms: Vec<(usize, Symbol)> = syms
685        .into_iter()
686        .filter_map(|s| {
687            let name = format!("{s}");
688            name.strip_prefix("🎯").and_then(|rest| rest.parse::<usize>().ok()).map(|k| (k, s))
689        })
690        .collect();
691
692    if coord_syms.is_empty() {
693        // No coordinate symbols – the value is uniform across all positions; propagate as-is.
694        return Some(expr.clone());
695    }
696
697    if let AxisOp::Reshape(at, from_dims, to_dims) = axis_op.canonical().as_ref() {
698        // Trivial all-ones case: shape change is purely cosmetic, value is unaffected.
699        if from_dims.iter().all(|d| d.is_one()) && to_dims.iter().all(|d| d.is_one()) {
700            return Some(expr.clone());
701        }
702        // Pure split: from = [D], to = [d_0, …, d_{k-1}], Π = D.  The input
703        // axis-`at` position decomposes as
704        //     pos[at] = Σ_i pos[at+i]_new · stride_i
705        // with `stride_i = Π_{j>i} to_dims[j]` (last stride is 1).  Other
706        // input axes shift right by `k-1` (the net rank change).
707        if from_dims.len() == 1 {
708            let from_dim = from_dims[0].clone();
709            let to_product: TDim = to_dims.iter().fold(TDim::Val(1), |acc, d| acc * d.clone());
710            if to_product == from_dim {
711                let k_to = to_dims.len();
712                let mut map: HashMap<Symbol, TDim> = HashMap::default();
713                for (k, sym) in &coord_syms {
714                    let scope = sym.scope()?;
715                    let new_expr = if *k < *at {
716                        TDim::Sym(sym.clone())
717                    } else if *k == *at {
718                        let mut sum = TDim::Val(0);
719                        let mut stride = TDim::Val(1);
720                        for i in (0..k_to).rev() {
721                            let new_sym = scope.coord_sym(*at + i);
722                            sum += TDim::Sym(new_sym) * stride.clone();
723                            stride *= to_dims[i].clone();
724                        }
725                        sum
726                    } else {
727                        TDim::Sym(scope.coord_sym(*k + k_to - 1))
728                    };
729                    map.insert(sym.clone(), new_expr);
730                }
731                return expr.substitute_all(&map).ok().map(|e| e.reduce());
732            }
733        }
734        // Pure merge: from = [d_0, …, d_{k-1}], to = [D].  We can express
735        // `pos[at+i]_old` from `pos[at]_new` only via integer division and
736        // modulo, which TDim doesn't carry.  Special-case the easy form
737        // where all but one of the merged dims is 1 — then the lone
738        // non-trivial sub-axis just maps to the new merged axis.
739        if to_dims.len() == 1 {
740            let to_dim = to_dims[0].clone();
741            let from_product: TDim = from_dims.iter().fold(TDim::Val(1), |acc, d| acc * d.clone());
742            if from_product == to_dim {
743                let k_from = from_dims.len();
744                let mut map: HashMap<Symbol, TDim> = HashMap::default();
745                for (k, sym) in &coord_syms {
746                    let scope = sym.scope()?;
747                    let new_expr = if *k < *at {
748                        TDim::Sym(sym.clone())
749                    } else if *k < *at + k_from {
750                        let i = *k - *at;
751                        let only_nontrivial =
752                            from_dims.iter().enumerate().all(|(j, d)| j == i || d.is_one());
753                        if only_nontrivial {
754                            TDim::Sym(scope.coord_sym(*at))
755                        } else {
756                            return None;
757                        }
758                    } else {
759                        TDim::Sym(scope.coord_sym(*k - (k_from - 1)))
760                    };
761                    map.insert(sym.clone(), new_expr);
762                }
763                return expr.substitute_all(&map).ok().map(|e| e.reduce());
764            }
765        }
766        return None;
767    }
768
769    // For Add/Rm/Move: use transform_axis and substitute all at once to avoid
770    // double-substitution when two axes swap positions (e.g. Move).
771    let map: HashMap<Symbol, TDim> = coord_syms
772        .into_iter()
773        .filter_map(|(k, sym)| {
774            let new_k = axis_op.transform_axis(k)?;
775            if new_k == k {
776                return None;
777            }
778            let scope = sym.scope()?;
779            Some((sym, TDim::Sym(scope.coord_sym(new_k))))
780        })
781        .collect();
782    if map.is_empty() {
783        return Some(expr.clone());
784    }
785    expr.substitute_all(&map).ok()
786}
787
788impl TypedOp for AxisOp {
789    as_op!();
790
791    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
792        if let Some(bqf) =
793            inputs[0].exotic_fact().and_then(|of| of.downcast_ref::<BlockQuantFact>())
794        {
795            let mut new_shape: TVec<usize> = bqf.shape().into();
796            self.change_shape_array(&mut new_shape, false)?;
797            let new_bqf = BlockQuantFact::new(bqf.format.clone(), new_shape.clone());
798            let shape: TVec<TDim> = new_shape.iter().map(|d| d.to_dim()).collect();
799            let mut new_fact = inputs[0].datum_type.fact(&*shape).with_exotic_fact(new_bqf);
800            if let Some(k) = &inputs[0].konst {
801                let mut new = k.clone().into_tensor();
802                self.change_tensor(&mut new, false)?;
803                new_fact.konst = Some(new.into());
804            }
805            return Ok(tvec!(new_fact));
806        }
807        let mut shape = inputs[0].shape.clone();
808        self.change_shape(&mut shape, false)?;
809        let mut fact = inputs[0].datum_type.fact(shape);
810        fact.exotic_fact.clone_from(&inputs[0].exotic_fact);
811        if let Some(tdim) = &inputs[0].uniform_tdim {
812            fact.uniform_tdim = remap_uniform_tdim(tdim, self);
813        }
814        Ok(tvec!(fact))
815    }
816
817    fn input_roi(
818        &self,
819        model: &TypedModel,
820        node: &TypedNode,
821    ) -> TractResult<Option<TVec<Option<TDim>>>> {
822        crate::optim::propagate_roi::bubble_roi(model, node)
823    }
824
825    fn axes_mapping(
826        &self,
827        inputs: &[&TypedFact],
828        outputs: &[&TypedFact],
829    ) -> TractResult<AxesMapping> {
830        let mut axes: Vec<Axis> = (0..inputs[0].rank())
831            .zip('a'..)
832            .map(|(axis_id, repr)| {
833                let mut axis = Axis::new(repr, inputs.len(), outputs.len()).input(0, axis_id);
834                if let Some(out) = self.transform_axis(axis_id) {
835                    axis = axis.output(0, out);
836                }
837                axis
838            })
839            .collect();
840        for (axis, letter) in (0..outputs[0].rank()).zip('A'..) {
841            if self.recip().transform_axis(axis).is_none() {
842                axes.push(Axis::new(letter, inputs.len(), outputs.len()).output(0, axis));
843            }
844        }
845        AxesMapping::new(inputs.len(), outputs.len(), axes)
846    }
847
848    fn declutter(
849        &self,
850        model: &TypedModel,
851        node: &TypedNode,
852    ) -> TractResult<Option<TypedModelPatch>> {
853        if self.is_noop()
854            && let Some(p) = TypedModelPatch::shunt_one_op(model, node)?
855        {
856            return Ok(Some(p));
857        }
858        let simplified = self.simplify();
859        if simplified.len() != 1 || &simplified[0] != self {
860            let mut patch = TypedModelPatch::default();
861            let mut wire = patch.tap_model(model, node.inputs[0])?;
862            for (ix, op) in simplified.into_iter().enumerate() {
863                wire = patch.wire_node(format!("{}.{}", node.name, ix), op, &[wire])?[0];
864            }
865            patch.shunt_outside(model, node.id.into(), wire)?;
866            Ok(Some(patch))
867        } else {
868            Ok(None)
869        }
870    }
871
872    fn suggested_axis_changes(&self) -> TractResult<TVec<(InOut, AxisOp)>> {
873        Ok(tvec!((InOut::Out(0), self.recip()), (InOut::In(0), self.clone())))
874    }
875
876    fn change_axes(
877        &self,
878        _model: &TypedModel,
879        _node: &TypedNode,
880        io: InOut,
881        change: &AxisOp,
882    ) -> TractResult<Option<AxisChangeConsequence>> {
883        let op = if let InOut::Out(0) = io {
884            rule_if_some!(more = self.recip().change_axes(_model, _node, InOut::In(0), change)?);
885            AxisChangeConsequence {
886                substitute_op: more.substitute_op.map(|op| {
887                    if let Some(op) = op.as_op().downcast_ref::<AxisOp>() {
888                        Box::new(op.recip())
889                    } else {
890                        op // have to be identity
891                    }
892                }),
893                wire_changes: more
894                    .wire_changes
895                    .into_iter()
896                    .map(|wc| {
897                        (if wc.0 == InOut::In(0) { InOut::Out(0) } else { InOut::In(0) }, wc.1)
898                    })
899                    .collect(),
900            }
901        } else if change == self {
902            AxisChangeConsequence { substitute_op: Some(Box::new(Identity)), wire_changes: tvec!() }
903        } else {
904            rule_if_some!((new_op, new_change) = self.merge_incoming_change(change));
905            trace!("  Change:{change:?} self:{self:?} -> change:{new_change:?} op:{new_op:?}");
906            let substitute_op: Box<dyn TypedOp> =
907                if let Some(o) = new_op { Box::new(o) as _ } else { Box::new(Identity) };
908            let mut wire_changes = tvec!();
909            if !change.is_noop() {
910                wire_changes.push((InOut::In(0), change.clone()))
911            }
912            if let Some(new_change) = new_change {
913                wire_changes.push((InOut::Out(0), new_change))
914            }
915            AxisChangeConsequence { substitute_op: Some(substitute_op), wire_changes }
916        };
917        Ok(Some(op))
918    }
919
920    fn set_symbols(
921        &self,
922        _source: &TypedModel,
923        node: &TypedNode,
924        target: &mut TypedModel,
925        mapping: &HashMap<OutletId, OutletId>,
926        subs: &HashMap<Symbol, TDim>,
927    ) -> TractResult<TVec<OutletId>> {
928        let op = if let AxisOp::Reshape(axis, from, to) = self {
929            AxisOp::Reshape(
930                *axis,
931                from.iter().map(|d| d.substitute_all(subs)).collect::<TractResult<_>>()?,
932                to.iter().map(|d| d.substitute_all(subs)).collect::<TractResult<_>>()?,
933            )
934        } else {
935            self.clone()
936        };
937        target.wire_node(&node.name, op, &[mapping[&node.inputs[0]]])
938    }
939
940    fn slice(
941        &self,
942        patch: &mut TypedModelPatch,
943        _model: &TypedModel,
944        node: &TypedNode,
945        _prefix: &str,
946        inputs: &[OutletId],
947        output_axis: usize,
948        _start: &TDim,
949        _end: &TDim,
950    ) -> TractResult<Option<TVec<OutletId>>> {
951        // is this test really useful ? or axis mapping preempt this ?
952        if let Reshape(pos, _from, to) = self
953            && output_axis >= *pos
954            && output_axis < pos + to.len()
955        {
956            return Ok(None);
957        }
958        patch.wire_node(&node.name, &node.op, inputs).map(Some)
959    }
960
961    fn codegen(
962        &self,
963        model: &TypedModel,
964        node: &TypedNode,
965    ) -> TractResult<Option<TypedModelPatch>> {
966        rule_if!(node.outputs[0].fact.exotic_fact.is_none());
967        if let Some(shape) = node.outputs[0].fact.shape.as_concrete()
968            && !matches!(self, AxisOp::Move(_, _))
969        {
970            let (inputs, outputs) = model.node_facts(node.id)?;
971            let mapping = self.axes_mapping(&inputs, &outputs)?;
972            let op = IntoShape {
973                mapping,
974                len: shape.iter().product(),
975                strides: Tensor::natural_strides(shape),
976                dims: shape.into(),
977            };
978            return Ok(Some(TypedModelPatch::replace_single_op(model, node, &node.inputs, op)?));
979        }
980        Ok(None)
981    }
982}
983
984// a, b, c is a <- b, b <- c, c <- a
985fn perm_to_cycles(perm: &[usize]) -> TVec<TVec<usize>> {
986    let mut cycles: TVec<TVec<usize>> = tvec!();
987    let mut done = 0;
988    while done < perm.len() {
989        if perm[done] == done || cycles.iter().any(|c| c.contains(&done)) {
990            done += 1;
991            continue;
992        }
993        let mut cycle = tvec!();
994        let mut current = done;
995        loop {
996            cycle.push(current);
997            current = perm[current];
998            if current == done {
999                break;
1000            }
1001        }
1002        cycles.push(cycle)
1003    }
1004    cycles
1005}
1006
1007fn is_rotation_cycle(cycle: &[usize]) -> Option<(usize, usize)> {
1008    if cycle.windows(2).all(|w| w[0] + 1 == w[1]) {
1009        Some((cycle[0], cycle[cycle.len() - 1]))
1010    } else if cycle[1..cycle.len()].windows(2).all(|w| w[0] - 1 == w[1])
1011        && cycle[cycle.len() - 1] - 1 == cycle[0]
1012    {
1013        Some((cycle[1], cycle[0]))
1014    } else {
1015        None
1016    }
1017}
1018
1019fn perm_to_atoms(input: &[usize]) -> TVec<(usize, usize)> {
1020    let mut changes: TVec<(usize, usize)> = tvec!();
1021    'top: loop {
1022        let mut reached: TVec<usize> = (0..input.len()).collect();
1023        changes.iter().for_each(|(f, t)| {
1024            let axis = reached.remove(*f);
1025            reached.insert(*t, axis);
1026        });
1027        if &*reached == input {
1028            return changes;
1029        }
1030        let remaining: TVec<usize> =
1031            input.iter().map(|x| reached.iter().position(|y| y == x).unwrap()).collect();
1032        let cycles = perm_to_cycles(&remaining);
1033        for cycle in &cycles {
1034            if let Some(rot) = is_rotation_cycle(cycle) {
1035                changes.push(rot);
1036                continue 'top;
1037            }
1038        }
1039        changes.push((cycles[0][1], cycles[0][0]));
1040    }
1041}
1042
1043pub fn perm_to_ops(input: &[usize]) -> TVec<AxisOp> {
1044    perm_to_atoms(input).into_iter().map(|pair| AxisOp::Move(pair.0, pair.1)).collect()
1045}
1046
1047pub fn compute_shape_with_tf_rules(input: &[TDim], shape_spec: &[TDim]) -> TractResult<TVec<TDim>> {
1048    compute_shape_with_onnx_rules(input, shape_spec, false)
1049}
1050
1051/// Resolve a `Reshape` target shape, honouring ONNX's `allowzero`.
1052///
1053/// With `allowzero` clear, a 0 in the target copies the input dimension at the same position.
1054/// This is the TensorFlow rule and the ONNX default, and is what `compute_shape_with_tf_rules`
1055/// asks for. With `allowzero` set, ONNX gives a 0 its literal meaning: a zero-length dimension,
1056/// "not taken from input tensor". `allowzero` was added in ONNX opset 14.
1057pub fn compute_shape_with_onnx_rules(
1058    input: &[TDim],
1059    shape_spec: &[TDim],
1060    allowzero: bool,
1061) -> TractResult<TVec<TDim>> {
1062    let mut shape: TVec<TDim> = shape_spec.into();
1063    if !allowzero {
1064        // Replace 0s with corresponding input dims (positional, per TF and ONNX allowzero=0)
1065        for (i, s) in shape.iter_mut().enumerate() {
1066            if *s == 0.into() {
1067                *s = input
1068                    .get(i)
1069                    .with_context(|| {
1070                        format!(
1071                            "Reshape: 0 at position {i} but input only has {} dims",
1072                            input.len()
1073                        )
1074                    })?
1075                    .clone();
1076            }
1077        }
1078    }
1079    let input_vol: TDim = input.iter().product();
1080    if let Some(pos) = shape.iter().position(|d| *d == (-1).into()) {
1081        // ONNX: "If the attribute 'allowzero' is set, it is invalid for the specified shape to
1082        // contain both a zero value and -1, as the value of the dimension corresponding to -1
1083        // cannot be determined uniquely."
1084        if allowzero && shape.iter().any(|d| *d == 0.into()) {
1085            bail!(
1086                "Reshape: allowzero is set, so the target shape {shape_spec:?} may not contain both 0 and -1"
1087            )
1088        }
1089        let shape_vol: TDim = shape.iter().filter(|d| **d != (-1).into()).product();
1090        // Dividing by a zero remainder leaves -1 undetermined rather than wrong: every value
1091        // satisfies the element count.
1092        if shape_vol == 0.into() {
1093            bail!(
1094                "Reshape: -1 in target shape {shape_spec:?} cannot be inferred, because the other \
1095                 dimensions of {shape:?} leave a volume of zero"
1096            )
1097        }
1098        let div = input_vol.maybe_div(&shape_vol)?;
1099        if div.1 != 1 {
1100            bail!("invalid")
1101        }
1102        shape[pos] = div.0;
1103    } else {
1104        let shape_vol: TDim = shape.iter().product();
1105        if input_vol != shape_vol {
1106            bail!(
1107                "Reshape volume mismatch: input {input:?} (vol={input_vol}) vs shape {shape:?} (vol={shape_vol})"
1108            );
1109        }
1110    }
1111    // Neither branch above rules a negative out: -1 resolves only the first, and comparing
1112    // volumes accepts one when a dimension is zero, as -2 * 0 == 3 * 0.
1113    if let Some(bad) = shape.iter().position(|d| d.as_i64().map(|d| d < 0).unwrap_or(false)) {
1114        bail!(
1115            "Reshape: dimension {bad} of target shape {shape_spec:?} is still negative after \
1116             inference (resolved to {shape:?}); at most one dimension may be -1"
1117        )
1118    }
1119    Ok(shape)
1120}
1121
1122pub fn to_axis_ops_with_tf_rules(
1123    input_orig: &[TDim],
1124    output_spec: &[TDim],
1125) -> TractResult<TVec<AxisOp>> {
1126    to_axis_ops_with_onnx_rules(input_orig, output_spec, false)
1127}
1128
1129/// As `to_axis_ops_with_tf_rules`, but honouring ONNX's `allowzero`.
1130pub fn to_axis_ops_with_onnx_rules(
1131    input_orig: &[TDim],
1132    output_spec: &[TDim],
1133    allowzero: bool,
1134) -> TractResult<TVec<AxisOp>> {
1135    let final_output = compute_shape_with_onnx_rules(input_orig, output_spec, allowzero)?;
1136    // A zero-length dimension makes every group volume zero, so the greedy volume matching below
1137    // finds spurious partial matches and builds an incoherent stack. There are no elements to
1138    // move, so the whole reshape is one unambiguous operation.
1139    if final_output.iter().chain(input_orig.iter()).any(|d| *d == 0.into()) {
1140        return Ok(if *input_orig == *final_output {
1141            tvec!()
1142        } else {
1143            tvec!(AxisOp::Reshape(0, input_orig.into(), final_output.clone()))
1144        });
1145    }
1146    let mut stack: TVec<AxisOp> = tvec!();
1147    'top: loop {
1148        let current_input =
1149            stack.iter().try_fold(TVec::from(input_orig), |mut shape, op| -> TractResult<_> {
1150                op.change_shape_array(&mut shape, false)?;
1151                Ok(shape)
1152            })?;
1153        if current_input == final_output {
1154            return Ok(stack);
1155        }
1156        if let Some(common) =
1157            current_input.iter().zip(final_output.iter()).position(|(a, b)| a != b)
1158        {
1159            if current_input[common].is_one() {
1160                stack.push(AxisOp::Rm(common));
1161            } else if final_output[common].is_one() {
1162                stack.push(AxisOp::Add(common));
1163            } else {
1164                // actual regrouping. search for a match. this is quadratic, but
1165                // rank is expected to be somewhat reasonable
1166                for i in common..current_input.len() {
1167                    let i_group = &current_input[common..i + 1];
1168                    let i_volume: TDim = i_group.iter().product();
1169                    for o in common..final_output.len() {
1170                        let o_group = &final_output[common..o + 1];
1171                        let o_volume: TDim = o_group.iter().product();
1172                        if i_volume == o_volume {
1173                            stack.push(AxisOp::Reshape(common, i_group.into(), o_group.into()));
1174                            continue 'top;
1175                        }
1176                    }
1177                }
1178                bail!(
1179                    "Could not find matching reshape grouping: current_input={current_input:?} final_output={final_output:?} common={common}"
1180                )
1181            }
1182        } else if final_output.len() > current_input.len() {
1183            stack.push(AxisOp::Add(current_input.len()));
1184        } else {
1185            stack.push(AxisOp::Rm(current_input.len() - 1));
1186        }
1187    }
1188}
1189
1190#[derive(Clone, Debug, PartialEq, Eq, Hash)]
1191pub struct IntoShape {
1192    pub mapping: AxesMapping,
1193    pub len: usize,
1194    pub dims: TVec<usize>,
1195    pub strides: TVec<isize>,
1196}
1197
1198impl Op for IntoShape {
1199    fn name(&self) -> StaticName {
1200        "IntoShape".into()
1201    }
1202
1203    fn info(&self) -> TractResult<Vec<String>> {
1204        Ok(vec![format!("{}", self.mapping)])
1205    }
1206
1207    op_as_typed_op!();
1208}
1209
1210impl EvalOp for IntoShape {
1211    op_out_of_plan!();
1212
1213    fn eval(&self, _ctx: &EvalContext, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
1214        let mut input = args_1!(inputs).into_tensor();
1215        ensure!(input.len() == self.len);
1216        unsafe { input.set_geometry_unchecked(&self.dims, &self.strides) };
1217        Ok(tvec!(input.into_tvalue()))
1218    }
1219}
1220
1221impl TypedOp for IntoShape {
1222    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
1223        let mut fact = inputs[0].datum_type.fact(&self.dims);
1224        if let Some(of) = &inputs[0].exotic_fact {
1225            fact = fact.with_exotic_fact(of.clone());
1226        }
1227        Ok(tvec!(fact))
1228    }
1229
1230    fn declutter(
1231        &self,
1232        model: &TypedModel,
1233        node: &TypedNode,
1234    ) -> TractResult<Option<TypedModelPatch>> {
1235        let input = model.outlet_fact(node.inputs[0])?;
1236        if input.shape.as_concrete().is_some_and(|shape| shape == &*self.dims) {
1237            return TypedModelPatch::shunt_one_op(model, node);
1238        }
1239        if let Some(succ) = model.single_succ(node.id)?
1240            && let Some(into_shape) = succ.op_as::<IntoShape>()
1241        {
1242            let op =
1243                Self { mapping: self.mapping.compose(&into_shape.mapping)?, ..into_shape.clone() };
1244            return Ok(Some(TypedModelPatch::fuse_with_next(model, node, op)?));
1245        }
1246        Ok(None)
1247    }
1248
1249    as_op!();
1250}
1251
1252#[cfg(test)]
1253mod test {
1254    use super::*;
1255
1256    #[test]
1257    fn test_perm_to_cycles() {
1258        assert_eq!(perm_to_cycles(&[1, 2, 0]), tvec!(tvec!(0, 1, 2)));
1259        assert_eq!(perm_to_cycles(&[2, 0, 1]), tvec!(tvec!(0, 2, 1)));
1260        assert_eq!(perm_to_cycles(&[1, 2, 3, 0]), tvec!(tvec!(0, 1, 2, 3)));
1261        assert_eq!(perm_to_cycles(&[3, 0, 1, 2]), tvec!(tvec!(0, 3, 2, 1)));
1262        assert_eq!(perm_to_cycles(&[3, 1, 2, 0, 4]), tvec!(tvec!(0, 3)));
1263    }
1264
1265    #[test]
1266    fn is_rotation() {
1267        assert_eq!(is_rotation_cycle(&[0, 1, 2]), Some((0, 2)));
1268        assert_eq!(is_rotation_cycle(&[0, 2, 1]), Some((2, 0)));
1269    }
1270
1271    #[test]
1272    fn test_perm_one_rotation() {
1273        assert_eq!(perm_to_atoms(&[1, 2, 0, 3, 4]), tvec!((0, 2)));
1274    }
1275
1276    #[test]
1277    fn test_perm_two_rotations() {
1278        assert_eq!(perm_to_atoms(&[1, 2, 0, 4, 3]), tvec!((0, 2), (3, 4)));
1279    }
1280
1281    #[test]
1282    fn test_perm_complex() {
1283        assert_eq!(perm_to_atoms(&[3, 1, 2, 0, 4]), tvec!((3, 0), (1, 3)));
1284    }
1285
1286    // ADD-ADD
1287
1288    //                          Op
1289    //           b,c   ------|Add(0)|----->        n,b,c
1290    //   Add(0)                                            Add(1)
1291    //         a,b,c   ------|Add(0)|----->        a,n,b,c
1292    #[test]
1293    pub fn transform_op_add_0_add_0() {
1294        let change = Add(0);
1295        let op = Add(0);
1296        assert_eq!(op.merge_incoming_change(&change), Some((Some(Add(0)), Some(Add(1)))));
1297    }
1298
1299    //                          Op
1300    //           b,c   ------|Add(1)|----->        b,n,c
1301    //   Add(0)                                                 Add(0)
1302    //         a,b,c   ------|Add(2)|----->        a,b,n,c
1303    #[test]
1304    pub fn transform_op_add_0_add_1() {
1305        let change = Add(0);
1306        let op = Add(1);
1307        assert_eq!(op.merge_incoming_change(&change), Some((Some(Add(2)), Some(Add(0)))));
1308    }
1309
1310    //                          Op
1311    //           a,c   ------|Add(0)|----->        n,a,c
1312    //   Add(1)                                                 Add(2)
1313    //         a,b,c   ------|Add(0)|----->        n,a,b,c
1314    #[test]
1315    pub fn transform_op_add_1_add_0() {
1316        let change = Add(1);
1317        let op = Add(0);
1318        assert_eq!(op.merge_incoming_change(&change), Some((Some(Add(0)), Some(Add(2)))));
1319    }
1320
1321    //                          Op
1322    //         a,b,c   ------|Rm(1)|----->         a,c
1323    //   Rm(0)                                             Rm(0)
1324    //           b,c   ------|Rm(0)|----->         c
1325    #[test]
1326    pub fn transform_op_rm_0_rm_1() {
1327        let change = Rm(0);
1328        let op = Rm(1);
1329        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(0)), Some(Rm(0)))));
1330    }
1331
1332    //                          Op
1333    //         a,b,c   ------|Rm(0)|----->         b,c
1334    //   Rm(1)                                             Rm(0)
1335    //           a,c   ------|Rm(0)|----->         c
1336    #[test]
1337    pub fn transform_op_rm_1_rm_0() {
1338        let change = Rm(1);
1339        let op = Rm(0);
1340        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(0)), Some(Rm(0)))));
1341    }
1342
1343    // ADD - RM
1344
1345    //                          Op
1346    //          b,c     ------|Rm(0)|------>        c
1347    //   Add(0)                                                 Add(0)
1348    //          a,b,c   ------|Rm(1)|----->         a,c
1349    #[test]
1350    pub fn transform_op_add_0_rm_0() {
1351        let change = Add(0);
1352        let op = Rm(0);
1353        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(1)), Some(Add(0)))));
1354    }
1355
1356    //                          Op
1357    //          b,c     ------|Rm(1)|------>        b
1358    //   Add(0)                                                 Add(0)
1359    //          a,b,c   ------|Rm(2)|----->         a,b
1360    #[test]
1361    pub fn transform_op_add_0_rm_1() {
1362        let change = Add(0);
1363        let op = Rm(1);
1364        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(2)), Some(Add(0)))));
1365    }
1366
1367    //                          Op
1368    //          a,c     ------|Rm(0)|------>        c
1369    //   Add(1)                                                 Add(0)
1370    //          a,b,c   ------|Rm(0)|----->         b,c
1371    #[test]
1372    pub fn transform_op_add_1_rm_0() {
1373        let change = Add(1);
1374        let op = Rm(0);
1375        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(0)), Some(Add(0)))));
1376    }
1377
1378    // RM - ADD
1379
1380    //                          Op
1381    //         a,b,c   ------|Add(0)|----->        X,a,b,c
1382    //   Rm(1)                                                 Rm(2)
1383    //           a,c   ------|Add(0)|----->        X,a,c
1384    #[test]
1385    pub fn transform_op_rm_1_add_0() {
1386        let change = Rm(1);
1387        let op = Add(0);
1388        assert_eq!(op.merge_incoming_change(&change), Some((Some(Add(0)), Some(Rm(2)))));
1389    }
1390
1391    //                          Op
1392    //         a,b,c   ------|Add(1)|----->        a,X,b,c
1393    //   Rm(0)                                                 Rm(0)
1394    //           b,c   ------|Add(0)|----->        X,b,c
1395    #[test]
1396    pub fn transform_op_rm_0_add_1() {
1397        let change = Rm(0);
1398        let op = Add(1);
1399        assert_eq!(op.merge_incoming_change(&change), Some((Some(Add(0)), Some(Rm(0)))));
1400    }
1401
1402    //                          Op
1403    //         a,b,c   ------|Rm(2)|----->        a,b
1404    //   Move(0, 2)                                           Move(0,1)
1405    //         b,c,a   ------|Rm(1)|----->        b,a
1406    #[test]
1407    pub fn transform_op_mv_02_rm_2() {
1408        let change = Move(0, 2);
1409        let op = Rm(2);
1410        assert_eq!(op.merge_incoming_change(&change), Some((Some(Rm(1)), Some(Move(0, 1)))));
1411    }
1412}
1413
1414#[cfg(test)]
1415mod proptests {
1416    use super::*;
1417    use proptest::prelude::*;
1418
1419    #[derive(Debug)]
1420    struct ComposeProblem {
1421        input: TVec<usize>,
1422        ops: TVec<AxisOp>,
1423    }
1424
1425    impl Arbitrary for AxisOp {
1426        type Parameters = TVec<usize>;
1427        type Strategy = BoxedStrategy<AxisOp>;
1428        fn arbitrary_with(shape: TVec<usize>) -> Self::Strategy {
1429            let mut ops: BoxedStrategy<AxisOp> = (0usize..shape.len() + 1).prop_map(Add).boxed();
1430            if shape.len() > 1 {
1431                ops = ops
1432                    .prop_union(
1433                        (0..shape.len(), 0..shape.len() - 1)
1434                            .prop_map(|(a, b)| Move(a, b + (b >= a) as usize))
1435                            .boxed(),
1436                    )
1437                    .boxed()
1438            }
1439            let rms = (0..shape.len()).filter(|&ax| shape[ax] == 1).map(Rm).collect::<Vec<_>>();
1440            if rms.len() > 0 {
1441                ops = ops
1442                    .prop_union((0..rms.len()).prop_map(move |rm| rms[rm].clone()).boxed())
1443                    .boxed()
1444            }
1445            let mergeable: Vec<AxisOp> = shape
1446                .windows(2)
1447                .enumerate()
1448                .filter(|(_, w)| w[0] > 1 && w[1] > 1)
1449                .map(|(ix, w)| {
1450                    Reshape(ix, tvec!(w[0].to_dim(), w[1].to_dim()), tvec!((w[0] * w[1]).to_dim()))
1451                })
1452                .collect();
1453            if mergeable.len() > 1 {
1454                ops = ops
1455                    .prop_union(
1456                        (0..mergeable.len()).prop_map(move |ix| mergeable[ix].clone()).boxed(),
1457                    )
1458                    .boxed()
1459            }
1460            ops
1461        }
1462    }
1463
1464    impl Arbitrary for ComposeProblem {
1465        type Parameters = ();
1466        type Strategy = BoxedStrategy<ComposeProblem>;
1467        fn arbitrary_with(_args: ()) -> Self::Strategy {
1468            let input = proptest::collection::vec(1usize..4, 1usize..4);
1469            fn tail(len: usize, shape: TVec<usize>) -> BoxedStrategy<TVec<AxisOp>> {
1470                if len == 0 {
1471                    Just(tvec!()).boxed()
1472                } else {
1473                    AxisOp::arbitrary_with(shape.clone())
1474                        .prop_flat_map(move |op| {
1475                            let mut shape = shape.clone();
1476                            op.change_shape_array(&mut shape, false).unwrap();
1477                            tail(len - 1, shape.clone()).prop_map(move |mut t| {
1478                                t.insert(0, op.clone());
1479                                t
1480                            })
1481                        })
1482                        .boxed()
1483                }
1484            }
1485            (input, 1usize..=5)
1486                .prop_flat_map(|(input, len)| (Just(input.clone()), tail(len, input.into())))
1487                .prop_map(|(input, ops)| ComposeProblem { input: input.into(), ops })
1488                .boxed()
1489        }
1490    }
1491
1492    impl ComposeProblem {
1493        pub fn model(&self) -> TractResult<TypedModel> {
1494            let mut model = TypedModel::default();
1495            let mut wire = model.add_source("source", i64::fact(&self.input))?;
1496            for (ix, op) in self.ops.iter().enumerate() {
1497                wire = model.wire_node(format!("op_{ix}"), op.clone(), &[wire])?[0];
1498            }
1499            model.select_output_outlets(&[wire])?;
1500            Ok(model)
1501        }
1502
1503        fn input(&self) -> TractResult<Tensor> {
1504            unsafe {
1505                let mut t = Tensor::uninitialized::<i64>(&self.input)?;
1506                for i in 0..t.len() {
1507                    t.try_as_plain_mut().unwrap().as_slice_mut().unwrap()[i] = i as i64;
1508                }
1509                Ok(t)
1510            }
1511        }
1512
1513        fn check(&self) -> TractResult<()> {
1514            crate::setup_test_logger();
1515            let input = self.input()?;
1516            let model = self.model()?;
1517            let raw = model.into_runnable()?.run(tvec!(input.clone().into_tvalue()))?;
1518            let optimized = self.model()?.into_decluttered()?;
1519            let opt = optimized.into_runnable()?.run(tvec!(input.into_tvalue()))?;
1520            opt[0].close_enough(&raw[0], false)
1521        }
1522    }
1523
1524    proptest! {
1525        #[test]
1526        fn recip(pb in any::<AxisOp>()) {
1527            assert_eq!(pb.recip().recip(), pb);
1528        }
1529
1530        #[test]
1531        fn axis_ops(pb in any::<ComposeProblem>()) {
1532            pb.check().unwrap()
1533        }
1534    }
1535
1536    #[test]
1537    fn add_0_rm_0() {
1538        let pb = ComposeProblem { input: tvec![1], ops: tvec![Add(0), Rm(0)] };
1539        pb.check().unwrap();
1540    }
1541
1542    #[test]
1543    fn add_0_move_01() {
1544        let pb = ComposeProblem { input: tvec![2], ops: tvec![Add(0), Move(0, 1)] };
1545        pb.check().unwrap();
1546    }
1547
1548    #[test]
1549    fn add_0_move_01_add_1() {
1550        let pb = ComposeProblem { input: tvec![2], ops: tvec![Add(0), Move(0, 1), Add(1)] };
1551        pb.check().unwrap();
1552    }
1553
1554    #[test]
1555    fn recip_move_01() {
1556        let op = Move(1, 0);
1557        assert_eq!(op.recip().recip(), op);
1558    }
1559
1560    #[test]
1561    fn recip_move_20() {
1562        let op = Move(2, 0);
1563        assert_eq!(op.recip().recip(), op);
1564    }
1565
1566    #[test]
1567    fn recip_move_02() {
1568        let op = Move(0, 2);
1569        assert_eq!(op.recip().recip(), op);
1570    }
1571
1572    #[test]
1573    fn add_0_add_1_move_02() {
1574        let pb = ComposeProblem { input: tvec![2], ops: tvec![Add(0), Add(1), Move(0, 2)] };
1575        pb.check().unwrap();
1576    }
1577
1578    #[test]
1579    fn add_0_add_0() {
1580        let pb = ComposeProblem { input: tvec![1], ops: tvec![Add(0), Add(0)] };
1581        pb.check().unwrap();
1582    }
1583
1584    #[test]
1585    fn add_0_add_0_move_02() {
1586        let pb = ComposeProblem { input: tvec![2], ops: tvec![Add(0), Add(0), Move(0, 2)] };
1587        pb.check().unwrap();
1588    }
1589
1590    #[test]
1591    fn add_0_add_2_move_12() {
1592        let pb = ComposeProblem { input: tvec![2], ops: tvec![Add(0), Add(2), Move(1, 2)] };
1593        pb.check().unwrap();
1594    }
1595
1596    #[test]
1597    fn add_0_add_0_move_02_rm_0() {
1598        let pb = ComposeProblem { input: tvec![1], ops: tvec![Add(0), Add(0), Move(0, 2), Rm(0)] };
1599        pb.check().unwrap();
1600    }
1601
1602    #[test]
1603    fn add_0_add_0_move_20_move_20() {
1604        let pb =
1605            ComposeProblem { input: tvec![2], ops: tvec![Add(0), Add(0), Move(2, 0), Move(2, 0)] };
1606        pb.check().unwrap();
1607    }
1608
1609    #[test]
1610    fn move_01_add_0() {
1611        let pb = ComposeProblem { input: tvec![1, 1], ops: tvec![Move(0, 1), Add(0)] };
1612        pb.check().unwrap();
1613    }
1614
1615    #[test]
1616    fn add_0_move_02_move_02() {
1617        let pb = ComposeProblem { input: tvec![1, 1], ops: tvec![Add(0), Move(0, 2), Move(0, 2),] };
1618        pb.check().unwrap();
1619    }
1620
1621    #[test]
1622    fn add_0_add_2_move_20_move_12_rm_2() {
1623        let pb = ComposeProblem {
1624            input: tvec![3],
1625            ops: tvec![Add(0), Add(2), Move(2, 0), Move(1, 2), Rm(2)],
1626        };
1627        pb.check().unwrap();
1628    }
1629
1630    #[test]
1631    fn move_02_move_02() {
1632        let pb = ComposeProblem { input: tvec![2, 1, 1], ops: tvec![Move(0, 2), Move(0, 2)] };
1633        pb.check().unwrap();
1634    }
1635
1636    #[test]
1637    fn rm_1_perm_10_add_0() {
1638        let pb = ComposeProblem { input: tvec![1, 1, 2], ops: tvec![Rm(1), Move(0, 1), Add(0)] };
1639        pb.check().unwrap();
1640    }
1641
1642    #[test]
1643    fn add_2_move_02_move_02() {
1644        let pb = ComposeProblem { input: tvec![3, 2], ops: tvec![Add(2), Move(0, 2), Move(0, 2)] };
1645        pb.check().unwrap();
1646    }
1647
1648    #[test]
1649    fn move_01_move_20_move_20() {
1650        let pb = ComposeProblem {
1651            input: tvec![2, 3, 2],
1652            ops: tvec![Move(0, 1), Move(2, 0), Move(2, 0)],
1653        };
1654        pb.check().unwrap();
1655    }
1656
1657    #[test]
1658    fn reshape_axes_tracking() {
1659        let pb = ComposeProblem {
1660            input: tvec![2, 2, 2],
1661            ops: tvec![Reshape(0, tvec!(2.to_dim(), 2.to_dim()), tvec!(4.to_dim()))],
1662        };
1663        pb.check().unwrap();
1664    }
1665
1666    #[test]
1667    fn simplify_reshape() {
1668        macro_rules! d {
1669            ($($dim: expr),*) =>  { tvec!($($dim.to_dim()),*) }
1670        }
1671        assert_eq!(Reshape(3, d!(), d!()).simplify(), tvec!());
1672        assert_eq!(Reshape(3, d!(2, 3), d!(2, 3)).simplify(), tvec!());
1673        assert_eq!(Reshape(3, d!(1), d!()).simplify(), tvec!(Rm(3)));
1674        assert_eq!(Reshape(3, d!(), d!(1)).simplify(), tvec!(Add(3)));
1675        assert_eq!(
1676            Reshape(3, d!(2, 3, 4), d!(2, 4, 3)).simplify(),
1677            tvec!(Reshape(4, d!(3, 4), d!(4, 3)))
1678        );
1679        assert_eq!(
1680            Reshape(3, d!(3, 4, 2), d!(4, 3, 2)).simplify(),
1681            tvec!(Reshape(3, d!(3, 4), d!(4, 3)))
1682        );
1683        assert_eq!(
1684            Reshape(3, d!(1, 2, 3), d!(3, 2)).simplify(),
1685            tvec!(Rm(3), Reshape(3, d!(2, 3), d!(3, 2)))
1686        );
1687        assert_eq!(
1688            Reshape(3, d!(2, 3), d!(1, 3, 2)).simplify(),
1689            tvec!(Reshape(3, d!(2, 3), d!(3, 2)), Add(3))
1690        );
1691        assert_eq!(
1692            Reshape(3, d!(2, 3, 1), d!(3, 2)).simplify(),
1693            tvec!(Rm(5), Reshape(3, d!(2, 3), d!(3, 2)))
1694        );
1695        assert_eq!(
1696            Reshape(3, d!(2, 3), d!(3, 2, 1)).simplify(),
1697            tvec!(Add(5), Reshape(3, d!(2, 3), d!(3, 2)))
1698        );
1699        assert_eq!(
1700            Reshape(2, d!(2, 2, 1), d!(4)).simplify(),
1701            tvec!(Rm(4), Reshape(2, d!(2, 2), d!(4)))
1702        );
1703        assert_eq!(Reshape(1, d!(1, 2), d!(2)).simplify(), tvec!(Rm(1)));
1704    }
1705
1706    macro_rules! s {
1707        ($($a:expr),*) => {&[ $($a.clone().into()),* ]}
1708    }
1709
1710    macro_rules! r {
1711        ($at: expr ; $($from:expr),* => $($to:expr),*) => {
1712            AxisOp::Reshape($at, tvec!($($from.into()),*),  tvec!($($to.into()),*))
1713        }
1714    }
1715
1716    #[test]
1717    fn compute_invalid() {
1718        assert!(compute_shape_with_tf_rules(s![3, 4, 5], s!(100)).is_err());
1719    }
1720
1721    #[test]
1722    fn compute_with_leading_zero() {
1723        assert_eq!(&*compute_shape_with_tf_rules(s![3, 4, 5], s!(0, 0, 5)).unwrap(), s![3, 4, 5])
1724    }
1725
1726    #[test]
1727    fn compute_with_leading_zero_with_flatten() {
1728        assert_eq!(
1729            &*compute_shape_with_tf_rules(s![2, 3, 5, 7], s!(2, 0, 35)).unwrap(),
1730            s![2, 3, 35]
1731        )
1732    }
1733
1734    #[test]
1735    fn compute_with_trailing_zero() {
1736        assert_eq!(&*compute_shape_with_tf_rules(s![3, 4, 5], s!(3, -1, 0)).unwrap(), s![3, 4, 5])
1737    }
1738
1739    // ONNX allowzero. A 0 in the target is literal rather than copied from the input, which is
1740    // only reachable through the ONNX frontend: TensorFlow has no such attribute and always
1741    // takes the copy rule above.
1742
1743    #[test]
1744    fn compute_allowzero_keeps_a_literal_zero() {
1745        assert_eq!(&*compute_shape_with_onnx_rules(s![3, 0], s!(0), true).unwrap(), s![0])
1746    }
1747
1748    #[test]
1749    fn compute_without_allowzero_copies_the_input_dim() {
1750        // Same model, attribute clear: the 0 becomes the input's dim 0, so 3 elements are asked
1751        // of a tensor holding none.
1752        assert!(compute_shape_with_onnx_rules(s![3, 0], s!(0), false).is_err())
1753    }
1754
1755    #[test]
1756    fn compute_allowzero_reordered() {
1757        // ONNX's own test_reshape_allowzero_reordered case.
1758        assert_eq!(
1759            &*compute_shape_with_onnx_rules(s![0, 3, 4], s!(3, 4, 0), true).unwrap(),
1760            s![3, 4, 0]
1761        )
1762    }
1763
1764    #[test]
1765    fn compute_rejects_two_inferred_dims() {
1766        assert!(compute_shape_with_tf_rules(s![2, 3, 4], s!(-1, -1)).is_err())
1767    }
1768
1769    #[test]
1770    fn compute_rejects_three_inferred_dims() {
1771        assert!(compute_shape_with_tf_rules(s![2, 3, 4], s!(-1, -1, -1)).is_err())
1772    }
1773
1774    #[test]
1775    fn compute_rejects_inferred_dim_beside_other_negative() {
1776        assert!(compute_shape_with_tf_rules(s![2, 3, 4], s!(-1, -2)).is_err());
1777        assert!(compute_shape_with_tf_rules(s![2, 3, 4], s!(-2, -1)).is_err())
1778    }
1779
1780    #[test]
1781    fn compute_rejects_negative_that_survives_the_volume_check() {
1782        assert!(compute_shape_with_tf_rules(s![3, 0], s!(-2, 0)).is_err())
1783    }
1784
1785    #[test]
1786    fn compute_rejects_inferred_dim_with_zero_volume_remainder() {
1787        assert!(compute_shape_with_tf_rules(s![3, 0], s!(-1, 0)).is_err());
1788        assert!(compute_shape_with_tf_rules(s![0, 3], s!(0, -1)).is_err())
1789    }
1790
1791    #[test]
1792    fn compute_still_infers_a_single_dim() {
1793        assert_eq!(&*compute_shape_with_tf_rules(s![2, 3, 4], s!(-1)).unwrap(), s![24]);
1794        assert_eq!(&*compute_shape_with_tf_rules(s![2, 3, 4], s!(2, -1)).unwrap(), s![2, 12]);
1795        assert_eq!(&*compute_shape_with_tf_rules(s![3, 0], s!(-1)).unwrap(), s![0]);
1796    }
1797
1798    #[test]
1799    fn compute_leaves_symbolic_dims_alone() {
1800        let table = SymbolScope::default();
1801        let s = table.new_with_prefix("S");
1802        assert_eq!(&*compute_shape_with_tf_rules(s![s, 2, 128], s!(0, -1)).unwrap(), s![s, 256]);
1803    }
1804
1805    #[test]
1806    fn compute_allowzero_rejects_zero_beside_minus_one() {
1807        // "it is invalid for the specified shape to contain both a zero value and -1, as the
1808        // value of the dimension corresponding to -1 cannot be determined uniquely."
1809        assert!(compute_shape_with_onnx_rules(s![2, 3], s!(0, -1), true).is_err())
1810    }
1811
1812    #[test]
1813    fn to_axis_ops_allowzero_keeps_a_literal_zero() {
1814        // The wiring path resolves the shape the same way, so it must agree with the rules path.
1815        assert!(to_axis_ops_with_onnx_rules(s![3, 0], s!(0), true).is_ok())
1816    }
1817
1818    #[test]
1819    fn compute_bug_1() {
1820        let table = SymbolScope::default();
1821        let s = table.new_with_prefix("S");
1822        assert_eq!(
1823            &*compute_shape_with_tf_rules(s![s, 1, 2, 128], s!(0, 0, -1)).unwrap(),
1824            s![s, 1, 256]
1825        )
1826    }
1827
1828    #[test]
1829    fn compute_bug_2() {
1830        let table = SymbolScope::default();
1831        let b = table.new_with_prefix("B");
1832        let s = table.new_with_prefix("S");
1833        assert_eq!(
1834            &*compute_shape_with_tf_rules(s![s, b, 2, 128], s!(0, 0, -1)).unwrap(),
1835            s![s, b, 256]
1836        )
1837    }
1838
1839    #[test]
1840    fn compute_zero_with_rank_change() {
1841        // Moonshine RoPE: input rank 4, output rank 5, two leading 0s
1842        assert_eq!(
1843            &*compute_shape_with_tf_rules(s![1, 52, 8, 32], s!(0, 0, 8, 16, 2)).unwrap(),
1844            s![1, 52, 8, 16, 2]
1845        )
1846    }
1847
1848    #[test]
1849    fn axis_op_rm_begin() {
1850        assert_eq!(&*to_axis_ops_with_tf_rules(s![1, 2, 3], s!(2, 3)).unwrap(), &[Rm(0)])
1851    }
1852
1853    #[test]
1854    fn axis_op_rm_end() {
1855        assert_eq!(&*to_axis_ops_with_tf_rules(s![2, 3, 1], s!(2, 3)).unwrap(), &[Rm(2)])
1856    }
1857
1858    #[test]
1859    fn axis_op_insert_begin() {
1860        assert_eq!(&*to_axis_ops_with_tf_rules(s![2, 3], s!(1, 2, 3)).unwrap(), &[Add(0)])
1861    }
1862
1863    #[test]
1864    fn axis_op_insert_end() {
1865        assert_eq!(&*to_axis_ops_with_tf_rules(s![2, 3], s!(2, 3, 1)).unwrap(), &[Add(2)])
1866    }
1867
1868    #[test]
1869    fn axis_op_merge() {
1870        assert_eq!(
1871            &*to_axis_ops_with_tf_rules(s![2, 3, 5, 7], s!(2, 0, 35)).unwrap(),
1872            &[r!(2 ; 5,7 => 35 )]
1873        )
1874    }
1875
1876    #[test]
1877    fn axis_op_complex() {
1878        assert_eq!(
1879            &*to_axis_ops_with_tf_rules(s![1, 2, 3, 5, 7], s!(2, 1, 3, 35, 1)).unwrap(),
1880            &[Rm(0), Add(1), r!(3 ; 5,7 => 35 ), Add(4)]
1881        )
1882    }
1883}