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)] #[allow(clippy::derived_hash_with_manual_eq)] pub 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 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, (Move(_, _), Reshape(_, _, _)) => None, (Reshape(_, _, _), Reshape(_, _, _)) => None, _ => 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 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 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
679fn 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 return Some(expr.clone());
695 }
696
697 if let AxisOp::Reshape(at, from_dims, to_dims) = axis_op.canonical().as_ref() {
698 if from_dims.iter().all(|d| d.is_one()) && to_dims.iter().all(|d| d.is_one()) {
700 return Some(expr.clone());
701 }
702 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 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 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 }
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 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
984fn 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
1051pub 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 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 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 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 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
1129pub 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 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 for i in common..current_input.len() {
1167 let i_group = ¤t_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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 assert!(compute_shape_with_onnx_rules(s![3, 0], s!(0), false).is_err())
1753 }
1754
1755 #[test]
1756 fn compute_allowzero_reordered() {
1757 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 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 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 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}