1use std::{
2 collections::{HashMap, HashSet},
3 fmt,
4 hash::Hash,
5 sync::{Arc, OnceLock},
6};
7
8use num::complex::Complex64;
9use serde::{Deserialize, Serialize};
10
11use crate::{
12 ExprGraphError, ExprShapeError, ParamError, ParamResult,
13 parameters::{InitialSpec, ParamState, Parameter, ParameterUpdate},
14};
15
16#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
18pub struct ExprId(u64);
19
20impl ExprId {
21 pub fn from_index(index: usize) -> Self {
23 Self(index as u64)
24 }
25
26 pub fn index(self) -> usize {
28 self.0 as usize
29 }
30}
31
32#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
34pub enum ValueKind {
35 Real,
37 Complex,
39 Vector {
41 len: usize,
43 },
44 Matrix {
46 rows: usize,
48 cols: usize,
50 },
51}
52
53#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
55pub enum NumberClass {
56 Unknown,
58 Real,
60 Imaginary,
62 Complex,
64}
65
66#[derive(Copy, Clone, Debug, PartialEq, Eq)]
68pub struct ExprNodeSemantics {
69 pub value_kind: ValueKind,
71 pub number_class: NumberClass,
73}
74
75fn add_number_class(lhs: NumberClass, rhs: NumberClass) -> NumberClass {
76 use NumberClass::{Complex, Imaginary, Real, Unknown};
77 match (lhs, rhs) {
78 (Real, Real) => Real,
79 (Imaginary, Imaginary) => Imaginary,
80 (Complex, _) | (_, Complex) => Complex,
81 (Unknown, _) | (_, Unknown) => Unknown,
82 _ => Complex,
83 }
84}
85
86fn mul_number_class(lhs: NumberClass, rhs: NumberClass) -> NumberClass {
87 use NumberClass::{Complex, Imaginary, Real, Unknown};
88 match (lhs, rhs) {
89 (Real, Real) | (Imaginary, Imaginary) => Real,
90 (Real, Imaginary) | (Imaginary, Real) => Imaginary,
91 (Complex, _) | (_, Complex) => Complex,
92 (Unknown, _) | (_, Unknown) => Unknown,
93 }
94}
95
96#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
98pub enum ExprDependencyKind {
99 Constant,
101 Parameter,
103 Event,
105 Children,
107}
108
109#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
111pub enum ExprShape {
112 Scalar,
114 Vector {
116 len: usize,
118 },
119 Matrix {
121 rows: usize,
123 cols: usize,
125 },
126}
127
128impl fmt::Display for ExprShape {
129 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
130 match self {
131 Self::Scalar => write!(f, "scalar"),
132 Self::Vector { len } => write!(f, "vector[{len}]"),
133 Self::Matrix { rows, cols } => write!(f, "matrix[{rows}x{cols}]"),
134 }
135 }
136}
137
138pub trait ComponentIndex {
140 fn component_index(self) -> usize;
142}
143
144impl ComponentIndex for usize {
145 fn component_index(self) -> usize {
146 self
147 }
148}
149
150impl ComponentIndex for i32 {
151 fn component_index(self) -> usize {
152 usize::try_from(self).expect("component index must be nonnegative")
153 }
154}
155
156#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
157pub enum P4Component {
159 E,
161 Px,
163 Py,
165 Pz,
167}
168
169impl P4Component {
170 pub fn label(self) -> &'static str {
172 match self {
173 Self::E => "e",
174 Self::Px => "px",
175 Self::Py => "py",
176 Self::Pz => "pz",
177 }
178 }
179
180 pub fn index(self) -> usize {
182 match self {
183 Self::E => 0,
184 Self::Px => 1,
185 Self::Py => 2,
186 Self::Pz => 3,
187 }
188 }
189}
190
191#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
193pub enum UnaryOp {
194 Neg,
196 Real,
198 Imag,
200 Conj,
202 NormSqr,
204 Sqrt,
206 Exp,
208 Sin,
210 Cos,
212 Log,
214 PowI(i32),
216}
217
218impl UnaryOp {
219 pub fn evaluate(&self, value: Complex64) -> Complex64 {
221 match self {
222 Self::Neg => -value,
223 Self::Real => Complex64::from(value.re),
224 Self::Imag => Complex64::from(value.im),
225 Self::Conj => value.conj(),
226 Self::NormSqr => Complex64::from(value.norm_sqr()),
227 Self::Sqrt => value.sqrt(),
228 Self::Exp => value.exp(),
229 Self::Sin => value.sin(),
230 Self::Cos => value.cos(),
231 Self::Log => value.ln(),
232 Self::PowI(power) => value.powi(*power),
233 }
234 }
235}
236
237#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
239pub enum BinaryOp {
240 Add,
242 Sub,
244 Mul,
246 Div,
248 Atan2,
250}
251
252impl BinaryOp {
253 pub fn evaluate(&self, a: Complex64, b: Complex64) -> Complex64 {
255 match self {
256 Self::Add => a + b,
257 Self::Sub => a - b,
258 Self::Mul => a * b,
259 Self::Div => a / b,
260 Self::Atan2 => Complex64::from(a.re.atan2(b.re)),
261 }
262 }
263}
264
265#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
267pub enum ExprNode {
268 RealConst(f64),
270 ComplexConst(Complex64),
272 ScalarParam(Parameter),
274 EventScalar(Arc<str>),
276 EventP4Component {
278 name: Arc<str>,
280 component: P4Component,
282 },
283 Unary {
285 op: UnaryOp,
287 input: ExprId,
289 },
290 Binary {
292 op: BinaryOp,
294 lhs: ExprId,
296 rhs: ExprId,
298 },
299 NaryAdd {
301 terms: Vec<ExprId>,
303 },
304 NaryMul {
306 factors: Vec<ExprId>,
308 },
309 Complex {
311 re: ExprId,
313 im: ExprId,
315 },
316 Vector {
318 elements: Vec<ExprId>,
320 },
321 Matrix {
323 rows: usize,
325 cols: usize,
327 elements: Vec<ExprId>,
329 },
330 Component {
332 input: ExprId,
334 index: usize,
336 },
337 MatrixElement {
339 input: ExprId,
341 row: usize,
343 col: usize,
345 },
346 MatMul {
348 lhs: ExprId,
350 rhs: ExprId,
352 },
353 MatVec {
355 matrix: ExprId,
357 vector: ExprId,
359 },
360 Dot {
362 lhs: ExprId,
364 rhs: ExprId,
366 },
367 Solve {
369 matrix: ExprId,
371 rhs: ExprId,
373 },
374}
375
376#[doc(hidden)]
382#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
383pub struct ParameterStructuralKey {
384 name: Arc<str>,
385 state: ParameterStateStructuralKey,
386 initial: InitialStructuralKey,
387 bounds: (Option<u64>, Option<u64>),
388 periodic: bool,
389 scale: Option<u64>,
390 unit: Option<Arc<str>>,
391 latex: Option<Arc<str>>,
392 description: Option<Arc<str>>,
393}
394
395#[doc(hidden)]
396#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
397pub enum ParameterStateStructuralKey {
398 Free,
399 Fixed(u64),
400}
401
402#[doc(hidden)]
403#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
404pub enum InitialStructuralKey {
405 Default,
406 Value(u64),
407 Uniform { min: u64, max: u64 },
408}
409
410#[doc(hidden)]
417#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
418pub enum ExprNodeStructuralKey {
419 RealConst(u64),
420 ComplexConst {
421 re: u64,
422 im: u64,
423 },
424 ScalarParam(ParameterStructuralKey),
425 EventScalar(Arc<str>),
426 EventP4Component {
427 name: Arc<str>,
428 component: P4Component,
429 },
430 Unary {
431 op: UnaryOp,
432 input: ExprId,
433 },
434 Binary {
435 op: BinaryOp,
436 lhs: ExprId,
437 rhs: ExprId,
438 },
439 NaryAdd {
440 terms: Vec<ExprId>,
441 },
442 NaryMul {
443 factors: Vec<ExprId>,
444 },
445 Complex {
446 re: ExprId,
447 im: ExprId,
448 },
449 Vector {
450 elements: Vec<ExprId>,
451 },
452 Matrix {
453 rows: usize,
454 cols: usize,
455 elements: Vec<ExprId>,
456 },
457 Component {
458 input: ExprId,
459 index: usize,
460 },
461 MatrixElement {
462 input: ExprId,
463 row: usize,
464 col: usize,
465 },
466 MatMul {
467 lhs: ExprId,
468 rhs: ExprId,
469 },
470 MatVec {
471 matrix: ExprId,
472 vector: ExprId,
473 },
474 Dot {
475 lhs: ExprId,
476 rhs: ExprId,
477 },
478 Solve {
479 matrix: ExprId,
480 rhs: ExprId,
481 },
482}
483
484impl From<&Parameter> for ParameterStructuralKey {
485 fn from(parameter: &Parameter) -> Self {
486 let state = match parameter.state() {
487 ParamState::Free => ParameterStateStructuralKey::Free,
488 ParamState::Fixed(value) => ParameterStateStructuralKey::Fixed(value.to_bits()),
489 };
490 let initial = match parameter.initial_spec() {
491 InitialSpec::Default => InitialStructuralKey::Default,
492 InitialSpec::Value(value) => InitialStructuralKey::Value(value.to_bits()),
493 InitialSpec::Uniform { min, max } => InitialStructuralKey::Uniform {
494 min: min.to_bits(),
495 max: max.to_bits(),
496 },
497 };
498 Self {
499 name: Arc::from(parameter.name()),
500 state,
501 initial,
502 bounds: (
503 parameter.bounds_spec().min.map(f64::to_bits),
504 parameter.bounds_spec().max.map(f64::to_bits),
505 ),
506 periodic: parameter.is_periodic(),
507 scale: parameter.scale().map(f64::to_bits),
508 unit: parameter.unit_label().map(Arc::from),
509 latex: parameter.latex_label().map(Arc::from),
510 description: parameter.description_text().map(Arc::from),
511 }
512 }
513}
514
515impl From<Complex64> for ExprNode {
516 fn from(value: Complex64) -> Self {
517 if value.im == 0.0 {
518 Self::RealConst(value.re)
519 } else {
520 Self::ComplexConst(value)
521 }
522 }
523}
524
525impl ExprNode {
526 pub fn semantics(&self, children: &[ExprNodeSemantics]) -> ExprNodeSemantics {
529 ExprNodeSemantics {
530 value_kind: self.infer_value_kind(children),
531 number_class: self.infer_number_class(children),
532 }
533 }
534
535 fn infer_value_kind(&self, children: &[ExprNodeSemantics]) -> ValueKind {
536 match self {
537 Self::RealConst(_) | Self::ScalarParam(_) => ValueKind::Real,
538 Self::ComplexConst(value) => {
539 if value.im == 0.0 {
540 ValueKind::Real
541 } else {
542 ValueKind::Complex
543 }
544 }
545 Self::EventScalar(_) | Self::EventP4Component { .. } => ValueKind::Real,
546 Self::Unary { op, input } => match op {
547 UnaryOp::Real | UnaryOp::Imag | UnaryOp::NormSqr => ValueKind::Real,
548 UnaryOp::Neg
549 | UnaryOp::Conj
550 | UnaryOp::Sqrt
551 | UnaryOp::Exp
552 | UnaryOp::Sin
553 | UnaryOp::Cos
554 | UnaryOp::Log
555 | UnaryOp::PowI(_) => children[input.index()].value_kind,
556 },
557 Self::Binary { op, lhs, rhs } => {
558 if *op == BinaryOp::Atan2 {
559 return ValueKind::Real;
560 }
561 if children[lhs.index()].value_kind == ValueKind::Real
562 && children[rhs.index()].value_kind == ValueKind::Real
563 {
564 ValueKind::Real
565 } else {
566 ValueKind::Complex
567 }
568 }
569 Self::NaryAdd { terms } => {
570 if terms
571 .iter()
572 .all(|id| children[id.index()].value_kind == ValueKind::Real)
573 {
574 ValueKind::Real
575 } else {
576 ValueKind::Complex
577 }
578 }
579 Self::NaryMul { factors } => {
580 if factors
581 .iter()
582 .all(|id| children[id.index()].value_kind == ValueKind::Real)
583 {
584 ValueKind::Real
585 } else {
586 ValueKind::Complex
587 }
588 }
589 Self::Complex { .. } => ValueKind::Complex,
590 Self::Vector { elements } => ValueKind::Vector {
591 len: elements.len(),
592 },
593 Self::Matrix { rows, cols, .. } => ValueKind::Matrix {
594 rows: *rows,
595 cols: *cols,
596 },
597 Self::Component { input, .. } => match children[input.index()].value_kind {
598 ValueKind::Vector { .. } => ValueKind::Complex,
599 kind => kind,
600 },
601 Self::MatrixElement { .. } | Self::Dot { .. } => ValueKind::Complex,
602 Self::MatMul { lhs, rhs } => {
603 let ValueKind::Matrix { rows, .. } = children[lhs.index()].value_kind else {
604 return ValueKind::Complex;
605 };
606 let ValueKind::Matrix { cols, .. } = children[rhs.index()].value_kind else {
607 return ValueKind::Complex;
608 };
609 ValueKind::Matrix { rows, cols }
610 }
611 Self::MatVec { matrix, .. } => {
612 let ValueKind::Matrix { rows, .. } = children[matrix.index()].value_kind else {
613 return ValueKind::Complex;
614 };
615 ValueKind::Vector { len: rows }
616 }
617 Self::Solve { rhs, .. } => children[rhs.index()].value_kind,
618 }
619 }
620
621 fn infer_number_class(&self, children: &[ExprNodeSemantics]) -> NumberClass {
622 match self {
623 Self::RealConst(_) | Self::ScalarParam(_) => NumberClass::Real,
624 Self::ComplexConst(value) => match (value.re == 0.0, value.im == 0.0) {
625 (_, true) => NumberClass::Real,
626 (true, false) => NumberClass::Imaginary,
627 (false, false) => NumberClass::Complex,
628 },
629 Self::EventScalar(_) | Self::EventP4Component { .. } => NumberClass::Real,
630 Self::Unary { op, input } => match op {
631 UnaryOp::Neg | UnaryOp::Conj => children[input.index()].number_class,
632 UnaryOp::Real | UnaryOp::Imag | UnaryOp::NormSqr => NumberClass::Real,
633 UnaryOp::Exp | UnaryOp::Sin | UnaryOp::Cos | UnaryOp::PowI(_) => {
634 let input = children[input.index()].number_class;
635 if input == NumberClass::Real {
636 NumberClass::Real
637 } else {
638 NumberClass::Unknown
639 }
640 }
641 UnaryOp::Sqrt | UnaryOp::Log => NumberClass::Unknown,
642 },
643 Self::Binary { op, lhs, rhs } => {
644 let lhs = children[lhs.index()].number_class;
645 let rhs = children[rhs.index()].number_class;
646 match op {
647 BinaryOp::Add | BinaryOp::Sub => add_number_class(lhs, rhs),
648 BinaryOp::Mul | BinaryOp::Div => mul_number_class(lhs, rhs),
649 BinaryOp::Atan2 => NumberClass::Real,
650 }
651 }
652 Self::NaryAdd { terms } => {
653 let mut classes = terms.iter().map(|id| children[id.index()].number_class);
654 let Some(first) = classes.next() else {
655 return NumberClass::Real;
656 };
657 classes.fold(first, add_number_class)
658 }
659 Self::NaryMul { factors } => {
660 let mut classes = factors.iter().map(|id| children[id.index()].number_class);
661 let Some(first) = classes.next() else {
662 return NumberClass::Real;
663 };
664 classes.fold(first, mul_number_class)
665 }
666 Self::Complex { .. } => NumberClass::Complex,
667 Self::Vector { .. }
668 | Self::Matrix { .. }
669 | Self::Component { .. }
670 | Self::MatrixElement { .. }
671 | Self::MatMul { .. }
672 | Self::MatVec { .. }
673 | Self::Dot { .. }
674 | Self::Solve { .. } => NumberClass::Unknown,
675 }
676 }
677
678 pub fn dependency_kind(&self) -> ExprDependencyKind {
680 match self {
681 Self::RealConst(_) | Self::ComplexConst(_) => ExprDependencyKind::Constant,
682 Self::ScalarParam(_) => ExprDependencyKind::Parameter,
683 Self::EventScalar(_) | Self::EventP4Component { .. } => ExprDependencyKind::Event,
684 _ => ExprDependencyKind::Children,
685 }
686 }
687
688 #[doc(hidden)]
690 pub fn structural_key(&self) -> ExprNodeStructuralKey {
691 match self {
692 Self::RealConst(value) => ExprNodeStructuralKey::RealConst(value.to_bits()),
693 Self::ComplexConst(value) => ExprNodeStructuralKey::ComplexConst {
694 re: value.re.to_bits(),
695 im: value.im.to_bits(),
696 },
697 Self::ScalarParam(parameter) => {
698 ExprNodeStructuralKey::ScalarParam(ParameterStructuralKey::from(parameter))
699 }
700 Self::EventScalar(name) => ExprNodeStructuralKey::EventScalar(Arc::clone(name)),
701 Self::EventP4Component { name, component } => ExprNodeStructuralKey::EventP4Component {
702 name: Arc::clone(name),
703 component: *component,
704 },
705 Self::Unary { op, input } => ExprNodeStructuralKey::Unary {
706 op: *op,
707 input: *input,
708 },
709 Self::Binary { op, lhs, rhs } => ExprNodeStructuralKey::Binary {
710 op: *op,
711 lhs: *lhs,
712 rhs: *rhs,
713 },
714 Self::NaryAdd { terms } => ExprNodeStructuralKey::NaryAdd {
715 terms: terms.clone(),
716 },
717 Self::NaryMul { factors } => ExprNodeStructuralKey::NaryMul {
718 factors: factors.clone(),
719 },
720 Self::Complex { re, im } => ExprNodeStructuralKey::Complex { re: *re, im: *im },
721 Self::Vector { elements } => ExprNodeStructuralKey::Vector {
722 elements: elements.clone(),
723 },
724 Self::Matrix {
725 rows,
726 cols,
727 elements,
728 } => ExprNodeStructuralKey::Matrix {
729 rows: *rows,
730 cols: *cols,
731 elements: elements.clone(),
732 },
733 Self::Component { input, index } => ExprNodeStructuralKey::Component {
734 input: *input,
735 index: *index,
736 },
737 Self::MatrixElement { input, row, col } => ExprNodeStructuralKey::MatrixElement {
738 input: *input,
739 row: *row,
740 col: *col,
741 },
742 Self::MatMul { lhs, rhs } => ExprNodeStructuralKey::MatMul {
743 lhs: *lhs,
744 rhs: *rhs,
745 },
746 Self::MatVec { matrix, vector } => ExprNodeStructuralKey::MatVec {
747 matrix: *matrix,
748 vector: *vector,
749 },
750 Self::Dot { lhs, rhs } => ExprNodeStructuralKey::Dot {
751 lhs: *lhs,
752 rhs: *rhs,
753 },
754 Self::Solve { matrix, rhs } => ExprNodeStructuralKey::Solve {
755 matrix: *matrix,
756 rhs: *rhs,
757 },
758 }
759 }
760
761 pub fn from_folded_const(value: Complex64) -> Self {
763 if value.im == 0.0 && value.im.is_sign_positive() {
764 Self::RealConst(value.re)
765 } else {
766 Self::ComplexConst(value)
767 }
768 }
769
770 pub fn const_value(&self) -> Option<Complex64> {
772 match self {
773 ExprNode::RealConst(value) => Some(Complex64::from(*value)),
774 ExprNode::ComplexConst(value) => Some(*value),
775 _ => None,
776 }
777 }
778
779 pub fn is_zero(node: &ExprNode) -> bool {
781 node.const_value()
782 .is_some_and(|value| value == Complex64::ZERO)
783 }
784
785 pub fn is_one(node: &ExprNode) -> bool {
787 node.const_value()
788 .is_some_and(|value| value == Complex64::ONE)
789 }
790
791 pub fn children(&self) -> impl ExactSizeIterator<Item = ExprId> + DoubleEndedIterator + '_ {
797 (0..self.child_count()).map(|index| self.child_at(index))
798 }
799
800 pub fn child_ids(&self) -> Vec<ExprId> {
805 self.children().collect()
806 }
807
808 pub fn map_children(&self, mut map: impl FnMut(ExprId) -> ExprId) -> Self {
813 match self {
814 Self::RealConst(_)
815 | Self::ComplexConst(_)
816 | Self::ScalarParam(_)
817 | Self::EventScalar(_)
818 | Self::EventP4Component { .. } => self.clone(),
819 Self::Unary { op, input } => Self::Unary {
820 op: *op,
821 input: map(*input),
822 },
823 Self::Binary { op, lhs, rhs } => Self::Binary {
824 op: *op,
825 lhs: map(*lhs),
826 rhs: map(*rhs),
827 },
828 Self::NaryAdd { terms } => Self::NaryAdd {
829 terms: terms.iter().copied().map(&mut map).collect(),
830 },
831 Self::NaryMul { factors } => Self::NaryMul {
832 factors: factors.iter().copied().map(&mut map).collect(),
833 },
834 Self::Complex { re, im } => Self::Complex {
835 re: map(*re),
836 im: map(*im),
837 },
838 Self::Vector { elements } => Self::Vector {
839 elements: elements.iter().copied().map(&mut map).collect(),
840 },
841 Self::Matrix {
842 rows,
843 cols,
844 elements,
845 } => Self::Matrix {
846 rows: *rows,
847 cols: *cols,
848 elements: elements.iter().copied().map(&mut map).collect(),
849 },
850 Self::Component { input, index } => Self::Component {
851 input: map(*input),
852 index: *index,
853 },
854 Self::MatrixElement { input, row, col } => Self::MatrixElement {
855 input: map(*input),
856 row: *row,
857 col: *col,
858 },
859 Self::MatMul { lhs, rhs } => Self::MatMul {
860 lhs: map(*lhs),
861 rhs: map(*rhs),
862 },
863 Self::MatVec { matrix, vector } => Self::MatVec {
864 matrix: map(*matrix),
865 vector: map(*vector),
866 },
867 Self::Dot { lhs, rhs } => Self::Dot {
868 lhs: map(*lhs),
869 rhs: map(*rhs),
870 },
871 Self::Solve { matrix, rhs } => Self::Solve {
872 matrix: map(*matrix),
873 rhs: map(*rhs),
874 },
875 }
876 }
877
878 fn child_count(&self) -> usize {
879 match self {
880 Self::RealConst(_)
881 | Self::ComplexConst(_)
882 | Self::ScalarParam(_)
883 | Self::EventScalar(_)
884 | Self::EventP4Component { .. } => 0,
885 Self::Unary { .. } | Self::Component { .. } | Self::MatrixElement { .. } => 1,
886 Self::Binary { .. }
887 | Self::Complex { .. }
888 | Self::MatMul { .. }
889 | Self::MatVec { .. }
890 | Self::Dot { .. }
891 | Self::Solve { .. } => 2,
892 Self::NaryAdd { terms } => terms.len(),
893 Self::NaryMul { factors } => factors.len(),
894 Self::Vector { elements } | Self::Matrix { elements, .. } => elements.len(),
895 }
896 }
897
898 fn child_at(&self, index: usize) -> ExprId {
899 match self {
900 Self::Unary { input, .. }
901 | Self::Component { input, .. }
902 | Self::MatrixElement { input, .. } => *input,
903 Self::Binary { lhs, rhs, .. }
904 | Self::Complex { re: lhs, im: rhs }
905 | Self::MatMul { lhs, rhs }
906 | Self::Dot { lhs, rhs } => [*lhs, *rhs][index],
907 Self::MatVec { matrix, vector } => [*matrix, *vector][index],
908 Self::Solve { matrix, rhs } => [*matrix, *rhs][index],
909 Self::NaryAdd { terms } => terms[index],
910 Self::NaryMul { factors } => factors[index],
911 Self::Vector { elements } | Self::Matrix { elements, .. } => elements[index],
912 Self::RealConst(_)
913 | Self::ComplexConst(_)
914 | Self::ScalarParam(_)
915 | Self::EventScalar(_)
916 | Self::EventP4Component { .. } => unreachable!("leaf node has no children"),
917 }
918 }
919}
920
921#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
923pub enum ExprSourceKind {
924 Const,
926 Param,
928 Event,
930 Unary,
932 Binary,
934 Complex,
936 Vector,
938 Matrix,
940 LinearAlgebra,
942}
943
944#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
946pub struct ExprMetadata {
947 source: ExprSourceKind,
948 name: Option<Arc<str>>,
949 tags: Vec<Arc<str>>,
950}
951
952impl ExprMetadata {
953 pub fn new(source: ExprSourceKind) -> Self {
955 Self {
956 source,
957 name: None,
958 tags: Vec::new(),
959 }
960 }
961
962 pub fn source(&self) -> ExprSourceKind {
964 self.source
965 }
966
967 pub fn name(&self) -> Option<&str> {
969 self.name.as_deref()
970 }
971
972 pub fn tags(&self) -> &[Arc<str>] {
974 &self.tags
975 }
976
977 pub fn has_tag(&self, tag: &str) -> bool {
979 self.tags.iter().any(|candidate| candidate.as_ref() == tag)
980 }
981}
982
983#[derive(Clone, Debug)]
989pub struct Expr {
990 node: Arc<DagNode>,
991}
992
993#[derive(Clone, Debug)]
994struct DagNode {
995 kind: DagNodeKind,
996 metadata: ExprMetadata,
997 shape: OnceLock<Result<ExprShape, ExprShapeError>>,
998}
999
1000#[derive(Clone, Debug)]
1001enum DagNodeKind {
1002 RealConst(f64),
1003 ComplexConst(Complex64),
1004 ScalarParam(Parameter),
1005 EventScalar(Arc<str>),
1006 EventP4Component {
1007 name: Arc<str>,
1008 component: P4Component,
1009 },
1010 Unary {
1011 op: UnaryOp,
1012 input: Expr,
1013 },
1014 Binary {
1015 op: BinaryOp,
1016 lhs: Expr,
1017 rhs: Expr,
1018 },
1019 TensorScalar {
1020 op: BinaryOp,
1021 tensor: Expr,
1022 scalar: Expr,
1023 scalar_on_left: bool,
1024 },
1025 Complex {
1026 re: Expr,
1027 im: Expr,
1028 },
1029 Vector {
1030 elements: Vec<Expr>,
1031 },
1032 Matrix {
1033 rows: usize,
1034 cols: usize,
1035 elements: Vec<Expr>,
1036 },
1037 Component {
1038 input: Expr,
1039 index: usize,
1040 },
1041 MatrixElement {
1042 input: Expr,
1043 row: usize,
1044 col: usize,
1045 },
1046 MatMul {
1047 lhs: Expr,
1048 rhs: Expr,
1049 },
1050 MatVec {
1051 matrix: Expr,
1052 vector: Expr,
1053 },
1054 Dot {
1055 lhs: Expr,
1056 rhs: Expr,
1057 },
1058 Solve {
1059 matrix: Expr,
1060 rhs: Expr,
1061 },
1062}
1063
1064impl DagNodeKind {
1065 fn child_count(&self) -> usize {
1066 match self {
1067 Self::RealConst(_)
1068 | Self::ComplexConst(_)
1069 | Self::ScalarParam(_)
1070 | Self::EventScalar(_)
1071 | Self::EventP4Component { .. } => 0,
1072 Self::Unary { .. } | Self::Component { .. } | Self::MatrixElement { .. } => 1,
1073 Self::Binary { .. }
1074 | Self::TensorScalar { .. }
1075 | Self::Complex { .. }
1076 | Self::MatMul { .. }
1077 | Self::MatVec { .. }
1078 | Self::Dot { .. }
1079 | Self::Solve { .. } => 2,
1080 Self::Vector { elements } | Self::Matrix { elements, .. } => elements.len(),
1081 }
1082 }
1083
1084 fn child_at(&self, index: usize) -> &Expr {
1085 match self {
1086 Self::Unary { input, .. }
1087 | Self::Component { input, .. }
1088 | Self::MatrixElement { input, .. } => input,
1089 Self::Binary { lhs, rhs, .. } | Self::MatMul { lhs, rhs } | Self::Dot { lhs, rhs } => {
1090 [lhs, rhs][index]
1091 }
1092 Self::TensorScalar { tensor, scalar, .. } => [tensor, scalar][index],
1093 Self::Complex { re, im } => [re, im][index],
1094 Self::MatVec { matrix, vector } => [matrix, vector][index],
1095 Self::Solve { matrix, rhs } => [matrix, rhs][index],
1096 Self::Vector { elements } | Self::Matrix { elements, .. } => &elements[index],
1097 Self::RealConst(_)
1098 | Self::ComplexConst(_)
1099 | Self::ScalarParam(_)
1100 | Self::EventScalar(_)
1101 | Self::EventP4Component { .. } => unreachable!("leaf nodes have no children"),
1102 }
1103 }
1104
1105 fn map_children(&self, mut map: impl FnMut(&Expr) -> Expr) -> Self {
1106 match self {
1107 Self::RealConst(value) => Self::RealConst(*value),
1108 Self::ComplexConst(value) => Self::ComplexConst(*value),
1109 Self::ScalarParam(parameter) => Self::ScalarParam(parameter.clone()),
1110 Self::EventScalar(name) => Self::EventScalar(Arc::clone(name)),
1111 Self::EventP4Component { name, component } => Self::EventP4Component {
1112 name: Arc::clone(name),
1113 component: *component,
1114 },
1115 Self::Unary { op, input } => Self::Unary {
1116 op: *op,
1117 input: map(input),
1118 },
1119 Self::Binary { op, lhs, rhs } => Self::Binary {
1120 op: *op,
1121 lhs: map(lhs),
1122 rhs: map(rhs),
1123 },
1124 Self::TensorScalar {
1125 op,
1126 tensor,
1127 scalar,
1128 scalar_on_left,
1129 } => Self::TensorScalar {
1130 op: *op,
1131 tensor: map(tensor),
1132 scalar: map(scalar),
1133 scalar_on_left: *scalar_on_left,
1134 },
1135 Self::Complex { re, im } => Self::Complex {
1136 re: map(re),
1137 im: map(im),
1138 },
1139 Self::Vector { elements } => Self::Vector {
1140 elements: elements.iter().map(&mut map).collect(),
1141 },
1142 Self::Matrix {
1143 rows,
1144 cols,
1145 elements,
1146 } => Self::Matrix {
1147 rows: *rows,
1148 cols: *cols,
1149 elements: elements.iter().map(&mut map).collect(),
1150 },
1151 Self::Component { input, index } => Self::Component {
1152 input: map(input),
1153 index: *index,
1154 },
1155 Self::MatrixElement { input, row, col } => Self::MatrixElement {
1156 input: map(input),
1157 row: *row,
1158 col: *col,
1159 },
1160 Self::MatMul { lhs, rhs } => Self::MatMul {
1161 lhs: map(lhs),
1162 rhs: map(rhs),
1163 },
1164 Self::MatVec { matrix, vector } => Self::MatVec {
1165 matrix: map(matrix),
1166 vector: map(vector),
1167 },
1168 Self::Dot { lhs, rhs } => Self::Dot {
1169 lhs: map(lhs),
1170 rhs: map(rhs),
1171 },
1172 Self::Solve { matrix, rhs } => Self::Solve {
1173 matrix: map(matrix),
1174 rhs: map(rhs),
1175 },
1176 }
1177 }
1178}
1179
1180impl Expr {
1181 fn new(kind: DagNodeKind) -> Self {
1182 let source = source_kind(&kind);
1183 Self {
1184 node: Arc::new(DagNode {
1185 kind,
1186 metadata: ExprMetadata::new(source),
1187 shape: OnceLock::new(),
1188 }),
1189 }
1190 }
1191
1192 pub fn named(self, name: impl Into<Arc<str>>) -> Self {
1194 self.with_metadata(|metadata| metadata.name = Some(name.into()))
1195 }
1196
1197 pub fn tagged(self, tag: impl Into<Arc<str>>) -> Self {
1199 let tag = tag.into();
1200 self.with_metadata(|metadata| {
1201 if !metadata.tags.iter().any(|existing| existing == &tag) {
1202 metadata.tags.push(tag);
1203 }
1204 })
1205 }
1206
1207 pub fn tagged_with(self, tags: impl IntoIterator<Item = impl Into<Arc<str>>>) -> Self {
1209 tags.into_iter().fold(self, Self::tagged)
1210 }
1211
1212 pub fn project_tags<'a>(&self, tags: impl IntoIterator<Item = &'a str>) -> Self {
1216 let tags: Vec<_> = tags.into_iter().collect();
1217 enum Frame {
1218 Visit(Expr),
1219 Rebuild(Expr, usize),
1220 }
1221
1222 let mut projected = Vec::new();
1223 let mut stack = vec![Frame::Visit(self.clone())];
1224 while let Some(frame) = stack.pop() {
1225 match frame {
1226 Frame::Visit(expr) => {
1227 let has_tags = !expr.node.metadata.tags.is_empty();
1228 if has_tags || expr.node.kind.child_count() == 0 {
1229 projected.push(
1230 if !has_tags
1231 || expr
1232 .node
1233 .metadata
1234 .tags
1235 .iter()
1236 .any(|candidate| tags.contains(&candidate.as_ref()))
1237 {
1238 expr.clone()
1239 } else {
1240 expr.zero_like()
1241 },
1242 );
1243 continue;
1244 }
1245 let child_count = expr.node.kind.child_count();
1246 stack.push(Frame::Rebuild(expr.clone(), child_count));
1247 for index in (0..child_count).rev() {
1248 stack.push(Frame::Visit(expr.node.kind.child_at(index).clone()));
1249 }
1250 }
1251 Frame::Rebuild(expr, child_count) => {
1252 let child_start = projected.len() - child_count;
1253 let children = projected.split_off(child_start);
1254 let mut children = children.into_iter();
1255 let kind = expr.node.kind.map_children(|original| {
1256 children.next().unwrap_or_else(|| original.clone())
1257 });
1258 projected.push(
1259 Expr::new(kind)
1260 .with_metadata(|metadata| *metadata = expr.node.metadata.clone()),
1261 );
1262 }
1263 }
1264 }
1265 projected.pop().unwrap_or_else(|| self.clone())
1266 }
1267
1268 fn zero_like(&self) -> Self {
1269 match self
1270 .shape()
1271 .expect("valid expression shapes are cached eagerly")
1272 {
1273 ExprShape::Scalar => Expr::from(0.0),
1274 ExprShape::Vector { len } => vector((0..len).map(|_| Expr::from(0.0))),
1275 ExprShape::Matrix { rows, cols } => {
1276 matrix_from_flat(rows, cols, (0..rows * cols).map(|_| Expr::from(0.0)))
1277 .expect("zero matrix dimensions match")
1278 }
1279 }
1280 }
1281
1282 pub fn real(&self) -> Self {
1284 unary(UnaryOp::Real, self)
1285 }
1286
1287 pub fn imag(&self) -> Self {
1289 unary(UnaryOp::Imag, self)
1290 }
1291
1292 pub fn conj(&self) -> Self {
1297 match self.shape() {
1298 Ok(ExprShape::Vector { len }) => {
1299 vector((0..len).map(|index| self.component(index).conj()))
1300 }
1301 Ok(ExprShape::Matrix { rows, cols }) => matrix_from_flat(
1302 rows,
1303 cols,
1304 (0..rows)
1305 .flat_map(|row| (0..cols).map(move |col| self.matrix_element(row, col).conj())),
1306 )
1307 .expect("conjugating a valid matrix preserves its dimensions"),
1308 _ => unary(UnaryOp::Conj, self),
1309 }
1310 }
1311
1312 pub fn norm_sqr(&self) -> Self {
1314 unary(UnaryOp::NormSqr, self)
1315 }
1316
1317 pub fn sqrt(&self) -> Self {
1319 unary(UnaryOp::Sqrt, self)
1320 }
1321
1322 pub fn exp(&self) -> Self {
1324 unary(UnaryOp::Exp, self)
1325 }
1326
1327 pub fn sin(&self) -> Self {
1329 unary(UnaryOp::Sin, self)
1330 }
1331
1332 pub fn cos(&self) -> Self {
1334 unary(UnaryOp::Cos, self)
1335 }
1336
1337 pub fn acos(&self) -> Self {
1339 atan2((Expr::from(1.0) - self.powi(2)).sqrt(), self)
1340 }
1341
1342 pub fn log(&self) -> Self {
1344 unary(UnaryOp::Log, self)
1345 }
1346
1347 pub fn powi(&self, power: i32) -> Self {
1349 unary(UnaryOp::PowI(power), self)
1350 }
1351
1352 pub fn component(&self, index: impl ComponentIndex) -> Self {
1354 Expr::new(DagNodeKind::Component {
1355 input: self.clone(),
1356 index: index.component_index(),
1357 })
1358 }
1359
1360 pub fn matrix_element(&self, row: usize, col: usize) -> Self {
1362 Expr::new(DagNodeKind::MatrixElement {
1363 input: self.clone(),
1364 row,
1365 col,
1366 })
1367 }
1368
1369 pub fn to_graph(&self) -> ExprGraph {
1371 if !self.contains_tensor_scalar() {
1372 return GraphBuilder::new().build(self);
1373 }
1374 let (root, mut nodes) = self.lower_tensor_scalars();
1375 let graph = GraphBuilder::new().build(&root);
1376 drop(root);
1377 while let Some(node) = nodes.pop() {
1380 drop(node);
1381 }
1382 graph
1383 }
1384
1385 fn contains_tensor_scalar(&self) -> bool {
1386 let mut seen = HashSet::new();
1387 let mut stack = vec![self];
1388 while let Some(expr) = stack.pop() {
1389 if !seen.insert(Arc::as_ptr(&expr.node) as usize) {
1390 continue;
1391 }
1392 if matches!(expr.node.kind, DagNodeKind::TensorScalar { .. }) {
1393 return true;
1394 }
1395 for index in 0..expr.node.kind.child_count() {
1396 stack.push(expr.node.kind.child_at(index));
1397 }
1398 }
1399 false
1400 }
1401
1402 fn lower_tensor_scalars(&self) -> (Self, Vec<Self>) {
1403 let mut lowered = HashMap::<usize, Expr>::new();
1404 let mut order = Vec::new();
1405 let mut stack = vec![(self.clone(), false)];
1406 while let Some((expr, visited)) = stack.pop() {
1407 let key = Arc::as_ptr(&expr.node) as usize;
1408 if lowered.contains_key(&key) {
1409 continue;
1410 }
1411 if visited {
1412 let kind = expr
1413 .node
1414 .kind
1415 .map_children(|child| lowered[&(Arc::as_ptr(&child.node) as usize)].clone());
1416 let result = match kind {
1417 DagNodeKind::TensorScalar {
1418 op,
1419 tensor,
1420 scalar,
1421 scalar_on_left,
1422 } => lower_tensor_scalar(op, tensor, scalar, scalar_on_left),
1423 other => Expr::new(other),
1424 }
1425 .with_metadata(|metadata| *metadata = expr.node.metadata.clone());
1426 order.push(result.clone());
1427 lowered.insert(key, result);
1428 } else {
1429 stack.push((expr.clone(), true));
1430 for index in (0..expr.node.kind.child_count()).rev() {
1431 stack.push((expr.node.kind.child_at(index).clone(), false));
1432 }
1433 }
1434 }
1435 (lowered[&(Arc::as_ptr(&self.node) as usize)].clone(), order)
1436 }
1437
1438 pub fn from_graph(graph: ExprGraph) -> Result<Self, ExprGraphError> {
1446 let ExprGraph {
1447 root,
1448 nodes,
1449 metadata,
1450 } = graph;
1451 let graph = ExprGraph::from_parts(root, nodes, metadata)?;
1452 let mut expressions: Vec<Expr> = Vec::with_capacity(graph.nodes.len());
1453 for (index, node) in graph.nodes.iter().enumerate() {
1454 let child = |id: ExprId| expressions[id.index()].clone();
1455 let expression = match node {
1456 ExprNode::RealConst(value) => Expr::new(DagNodeKind::RealConst(*value)),
1457 ExprNode::ComplexConst(value) => Expr::new(DagNodeKind::ComplexConst(*value)),
1458 ExprNode::ScalarParam(parameter) => {
1459 Expr::new(DagNodeKind::ScalarParam(parameter.clone()))
1460 }
1461 ExprNode::EventScalar(name) => {
1462 Expr::new(DagNodeKind::EventScalar(Arc::clone(name)))
1463 }
1464 ExprNode::EventP4Component { name, component } => {
1465 Expr::new(DagNodeKind::EventP4Component {
1466 name: Arc::clone(name),
1467 component: *component,
1468 })
1469 }
1470 ExprNode::Unary { op, input } => Expr::new(DagNodeKind::Unary {
1471 op: *op,
1472 input: child(*input),
1473 }),
1474 ExprNode::Binary { op, lhs, rhs } => Expr::new(DagNodeKind::Binary {
1475 op: *op,
1476 lhs: child(*lhs),
1477 rhs: child(*rhs),
1478 }),
1479 ExprNode::NaryAdd { terms } => terms
1480 .iter()
1481 .map(|id| child(*id))
1482 .reduce(|lhs, rhs| binary(BinaryOp::Add, &lhs, &rhs))
1483 .unwrap_or_else(|| Expr::from(0.0)),
1484 ExprNode::NaryMul { factors } => factors
1485 .iter()
1486 .map(|id| child(*id))
1487 .reduce(|lhs, rhs| binary(BinaryOp::Mul, &lhs, &rhs))
1488 .unwrap_or_else(|| Expr::from(1.0)),
1489 ExprNode::Complex { re, im } => Expr::new(DagNodeKind::Complex {
1490 re: child(*re),
1491 im: child(*im),
1492 }),
1493 ExprNode::Vector { elements } => Expr::new(DagNodeKind::Vector {
1494 elements: elements.iter().map(|id| child(*id)).collect(),
1495 }),
1496 ExprNode::Matrix {
1497 rows,
1498 cols,
1499 elements,
1500 } => Expr::new(DagNodeKind::Matrix {
1501 rows: *rows,
1502 cols: *cols,
1503 elements: elements.iter().map(|id| child(*id)).collect(),
1504 }),
1505 ExprNode::Component { input, index } => Expr::new(DagNodeKind::Component {
1506 input: child(*input),
1507 index: *index,
1508 }),
1509 ExprNode::MatrixElement { input, row, col } => {
1510 Expr::new(DagNodeKind::MatrixElement {
1511 input: child(*input),
1512 row: *row,
1513 col: *col,
1514 })
1515 }
1516 ExprNode::MatMul { lhs, rhs } => Expr::new(DagNodeKind::MatMul {
1517 lhs: child(*lhs),
1518 rhs: child(*rhs),
1519 }),
1520 ExprNode::MatVec { matrix, vector } => Expr::new(DagNodeKind::MatVec {
1521 matrix: child(*matrix),
1522 vector: child(*vector),
1523 }),
1524 ExprNode::Dot { lhs, rhs } => Expr::new(DagNodeKind::Dot {
1525 lhs: child(*lhs),
1526 rhs: child(*rhs),
1527 }),
1528 ExprNode::Solve { matrix, rhs } => Expr::new(DagNodeKind::Solve {
1529 matrix: child(*matrix),
1530 rhs: child(*rhs),
1531 }),
1532 };
1533 let mut dag = (*expression.node).clone();
1534 dag.metadata = graph.metadata[index].clone();
1535 expressions.push(Expr {
1536 node: Arc::new(dag),
1537 });
1538 }
1539 Ok(expressions[graph.root.index()].clone())
1540 }
1541
1542 pub fn shape(&self) -> Result<ExprShape, ExprShapeError> {
1549 self.node
1550 .shape
1551 .get_or_init(|| self.node.kind.shape())
1552 .clone()
1553 }
1554
1555 fn with_metadata(self, f: impl FnOnce(&mut ExprMetadata)) -> Self {
1556 let mut node = (*self.node).clone();
1557 f(&mut node.metadata);
1558 Self {
1559 node: Arc::new(node),
1560 }
1561 }
1562}
1563
1564impl Serialize for Expr {
1565 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1566 where
1567 S: serde::Serializer,
1568 {
1569 self.to_graph().serialize(serializer)
1570 }
1571}
1572
1573impl<'de> Deserialize<'de> for Expr {
1574 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1575 where
1576 D: serde::Deserializer<'de>,
1577 {
1578 Expr::from_graph(ExprGraph::deserialize(deserializer)?).map_err(serde::de::Error::custom)
1579 }
1580}
1581
1582impl DagNodeKind {
1583 fn shape(&self) -> Result<ExprShape, ExprShapeError> {
1584 match self {
1585 Self::RealConst(_)
1586 | Self::ComplexConst(_)
1587 | Self::ScalarParam(_)
1588 | Self::EventScalar(_)
1589 | Self::EventP4Component { .. } => Ok(ExprShape::Scalar),
1590 Self::Unary { input, .. } => {
1591 input.expect_shape("unary operation", ExprShape::Scalar)?;
1592 Ok(ExprShape::Scalar)
1593 }
1594 Self::Binary { lhs, rhs, .. } => {
1595 lhs.expect_shape("binary operation", ExprShape::Scalar)?;
1596 rhs.expect_shape("binary operation", ExprShape::Scalar)?;
1597 Ok(ExprShape::Scalar)
1598 }
1599 Self::TensorScalar {
1600 op,
1601 tensor,
1602 scalar,
1603 scalar_on_left,
1604 } => {
1605 scalar.expect_shape("tensor-scalar operation", ExprShape::Scalar)?;
1606 if *scalar_on_left && *op == BinaryOp::Div {
1607 return Err(ExprShapeError::new(
1608 "tensor-scalar operation",
1609 "scalar division by a tensor is unsupported",
1610 ));
1611 }
1612 match tensor.shape()? {
1613 shape @ (ExprShape::Vector { .. } | ExprShape::Matrix { .. }) => Ok(shape),
1614 ExprShape::Scalar => Err(ExprShapeError::new(
1615 "tensor-scalar operation",
1616 "expected a vector or matrix operand",
1617 )),
1618 }
1619 }
1620 Self::Complex { re, im } => {
1621 re.expect_shape("complex constructor", ExprShape::Scalar)?;
1622 im.expect_shape("complex constructor", ExprShape::Scalar)?;
1623 Ok(ExprShape::Scalar)
1624 }
1625 Self::Vector { elements } => {
1626 for element in elements {
1627 element.expect_shape("vector constructor", ExprShape::Scalar)?;
1628 }
1629 Ok(ExprShape::Vector {
1630 len: elements.len(),
1631 })
1632 }
1633 Self::Matrix {
1634 rows,
1635 cols,
1636 elements,
1637 } => {
1638 let expected = rows.checked_mul(*cols).ok_or_else(|| {
1639 ExprShapeError::new("matrix constructor", "row/column product overflowed")
1640 })?;
1641 if elements.len() != expected {
1642 return Err(ExprShapeError::new(
1643 "matrix constructor",
1644 format!(
1645 "shape {rows}x{cols} requires {expected} elements, got {}",
1646 elements.len()
1647 ),
1648 ));
1649 }
1650 for element in elements {
1651 element.expect_shape("matrix constructor", ExprShape::Scalar)?;
1652 }
1653 Ok(ExprShape::Matrix {
1654 rows: *rows,
1655 cols: *cols,
1656 })
1657 }
1658 Self::Component { input, index } => {
1659 let ExprShape::Vector { len } = input.shape()? else {
1660 return Err(ExprShapeError::new(
1661 "component",
1662 format!("expected vector, got {}", input.shape()?),
1663 ));
1664 };
1665 if *index >= len {
1666 return Err(ExprShapeError::new(
1667 "component",
1668 format!("index {index} is out of bounds for vector[{len}]"),
1669 ));
1670 }
1671 Ok(ExprShape::Scalar)
1672 }
1673 Self::MatrixElement { input, row, col } => {
1674 let ExprShape::Matrix { rows, cols } = input.shape()? else {
1675 return Err(ExprShapeError::new(
1676 "matrix element",
1677 format!("expected matrix, got {}", input.shape()?),
1678 ));
1679 };
1680 if *row >= rows || *col >= cols {
1681 return Err(ExprShapeError::new(
1682 "matrix element",
1683 format!("index ({row}, {col}) is out of bounds for matrix[{rows}x{cols}]"),
1684 ));
1685 }
1686 Ok(ExprShape::Scalar)
1687 }
1688 Self::MatMul { lhs, rhs } => {
1689 let ExprShape::Matrix {
1690 rows: lhs_rows,
1691 cols: lhs_cols,
1692 } = lhs.shape()?
1693 else {
1694 return Err(ExprShapeError::new(
1695 "matrix multiplication",
1696 format!("left input must be a matrix, got {}", lhs.shape()?),
1697 ));
1698 };
1699 let ExprShape::Matrix {
1700 rows: rhs_rows,
1701 cols: rhs_cols,
1702 } = rhs.shape()?
1703 else {
1704 return Err(ExprShapeError::new(
1705 "matrix multiplication",
1706 format!("right input must be a matrix, got {}", rhs.shape()?),
1707 ));
1708 };
1709 if lhs_cols != rhs_rows {
1710 return Err(ExprShapeError::new(
1711 "matrix multiplication",
1712 format!("cannot multiply {lhs_rows}x{lhs_cols} by {rhs_rows}x{rhs_cols}"),
1713 ));
1714 }
1715 Ok(ExprShape::Matrix {
1716 rows: lhs_rows,
1717 cols: rhs_cols,
1718 })
1719 }
1720 Self::MatVec { matrix, vector } => {
1721 let ExprShape::Matrix { rows, cols } = matrix.shape()? else {
1722 return Err(ExprShapeError::new(
1723 "matrix-vector multiplication",
1724 format!("left input must be a matrix, got {}", matrix.shape()?),
1725 ));
1726 };
1727 let ExprShape::Vector { len } = vector.shape()? else {
1728 return Err(ExprShapeError::new(
1729 "matrix-vector multiplication",
1730 format!("right input must be a vector, got {}", vector.shape()?),
1731 ));
1732 };
1733 if cols != len {
1734 return Err(ExprShapeError::new(
1735 "matrix-vector multiplication",
1736 format!("cannot multiply {rows}x{cols} matrix by vector[{len}]"),
1737 ));
1738 }
1739 Ok(ExprShape::Vector { len: rows })
1740 }
1741 Self::Dot { lhs, rhs } => {
1742 let ExprShape::Vector { len: lhs_len } = lhs.shape()? else {
1743 return Err(ExprShapeError::new(
1744 "dot product",
1745 format!("left input must be a vector, got {}", lhs.shape()?),
1746 ));
1747 };
1748 let ExprShape::Vector { len: rhs_len } = rhs.shape()? else {
1749 return Err(ExprShapeError::new(
1750 "dot product",
1751 format!("right input must be a vector, got {}", rhs.shape()?),
1752 ));
1753 };
1754 if lhs_len != rhs_len {
1755 return Err(ExprShapeError::new(
1756 "dot product",
1757 format!("vector lengths differ: {lhs_len} and {rhs_len}"),
1758 ));
1759 }
1760 Ok(ExprShape::Scalar)
1761 }
1762 Self::Solve { matrix, rhs } => {
1763 let ExprShape::Matrix { rows, cols } = matrix.shape()? else {
1764 return Err(ExprShapeError::new(
1765 "linear solve",
1766 format!("left input must be a matrix, got {}", matrix.shape()?),
1767 ));
1768 };
1769 let ExprShape::Vector { len } = rhs.shape()? else {
1770 return Err(ExprShapeError::new(
1771 "linear solve",
1772 format!("right input must be a vector, got {}", rhs.shape()?),
1773 ));
1774 };
1775 if rows != cols || rows != len {
1776 return Err(ExprShapeError::new(
1777 "linear solve",
1778 format!("cannot solve matrix[{rows}x{cols}] against vector[{len}]"),
1779 ));
1780 }
1781 Ok(ExprShape::Vector { len })
1782 }
1783 }
1784 }
1785}
1786
1787impl Expr {
1788 fn is_tensor_syntax(&self) -> bool {
1789 matches!(
1790 self.node.kind,
1791 DagNodeKind::Vector { .. }
1792 | DagNodeKind::Matrix { .. }
1793 | DagNodeKind::MatMul { .. }
1794 | DagNodeKind::MatVec { .. }
1795 | DagNodeKind::Solve { .. }
1796 | DagNodeKind::TensorScalar { .. }
1797 )
1798 }
1799
1800 fn expect_shape(
1801 &self,
1802 operation: &'static str,
1803 expected: ExprShape,
1804 ) -> Result<(), ExprShapeError> {
1805 let actual = self.shape()?;
1806 if actual != expected {
1807 return Err(ExprShapeError::new(
1808 operation,
1809 format!("expected {expected}, got {actual}"),
1810 ));
1811 }
1812 Ok(())
1813 }
1814}
1815
1816auto_ops::impl_op_ex!(+ |a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1817
1818auto_ops::impl_op_ex!(+ |a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Add, a, b) });
1819auto_ops::impl_op_ex!(+ |a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1820
1821auto_ops::impl_op_ex!(+ |a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Add, a, b) });
1822auto_ops::impl_op_ex!(+ |a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1823
1824auto_ops::impl_op_ex!(+ |a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1825auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1826
1827auto_ops::impl_op_ex!(+ |a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Add, a, b) });
1828auto_ops::impl_op_ex!(+ |a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1829
1830auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Add, a, b) });
1831auto_ops::impl_op_ex!(+ |a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1832
1833auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1834
1835auto_ops::impl_op_ex!(-|a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1836
1837auto_ops::impl_op_ex!(-|a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Sub, a, b) });
1838
1839auto_ops::impl_op_ex!(-|a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1840
1841auto_ops::impl_op_ex!(-|a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Sub, a, b) });
1842
1843auto_ops::impl_op_ex!(-|a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1844
1845auto_ops::impl_op_ex!(-|a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1846
1847auto_ops::impl_op_ex!(-|a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1848
1849auto_ops::impl_op_ex!(-|a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1850
1851auto_ops::impl_op_ex!(-|a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Sub, a, b) });
1852
1853auto_ops::impl_op_ex!(-|a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1854
1855auto_ops::impl_op_ex!(-|a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Sub, a, b) });
1856
1857auto_ops::impl_op_ex!(-|a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1858
1859auto_ops::impl_op_ex!(*|a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1860
1861auto_ops::impl_op_ex!(*|a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Mul, a, b) });
1862auto_ops::impl_op_ex!(*|a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1863
1864auto_ops::impl_op_ex!(*|a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Mul, a, b) });
1865auto_ops::impl_op_ex!(*|a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1866
1867auto_ops::impl_op_ex!(*|a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1868auto_ops::impl_op_ex!(*|a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1869
1870auto_ops::impl_op_ex!(*|a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1871auto_ops::impl_op_ex!(*|a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Mul, a, b) });
1872
1873auto_ops::impl_op_ex!(*|a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1874auto_ops::impl_op_ex!(*|a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Mul, a, b) });
1875
1876auto_ops::impl_op_ex!(*|a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1877
1878auto_ops::impl_op_ex!(/ |a: &Expr, b: &Expr| -> Expr {
1879 binary(BinaryOp::Div, a, b)
1880});
1881
1882auto_ops::impl_op_ex!(/ |a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Div, a, b) });
1883auto_ops::impl_op_ex!(/ |a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Div, a, b) });
1884
1885auto_ops::impl_op_ex!(/ |a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Div, a, b) });
1886auto_ops::impl_op_ex!(/ |a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Div, a, b) });
1887
1888auto_ops::impl_op_ex!(/|a: &Expr, b: &Parameter| -> Expr {
1889 binary(BinaryOp::Div, a, b)
1890});
1891auto_ops::impl_op_ex!(/|a: &Parameter, b: &Expr| -> Expr {
1892 binary(BinaryOp::Div, a, b)
1893});
1894
1895auto_ops::impl_op_ex!(/|a: &f64, b: &Parameter| -> Expr {
1896 binary(BinaryOp::Div, a, b)
1897});
1898auto_ops::impl_op_ex!(/|a: &Parameter, b: &f64| -> Expr {
1899 binary(BinaryOp::Div, a, b)
1900});
1901
1902auto_ops::impl_op_ex!(/|a: &Complex64, b: &Parameter| -> Expr {
1903 binary(BinaryOp::Div, a, b)
1904});
1905auto_ops::impl_op_ex!(/|a: &Parameter, b: &Complex64| -> Expr {
1906 binary(BinaryOp::Div, a, b)
1907});
1908
1909auto_ops::impl_op_ex!(/|a: &Parameter, b: &Parameter| -> Expr {
1910 binary(BinaryOp::Div, a, b)
1911});
1912
1913auto_ops::impl_op_ex!(-|a: &Expr| -> Expr { unary(UnaryOp::Neg, a) });
1914auto_ops::impl_op_ex!(-|a: &Parameter| -> Expr { unary(UnaryOp::Neg, a) });
1915
1916auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Expr| {
1917 *a = binary(BinaryOp::Add, &*a, b);
1918});
1919auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &f64| {
1920 *a = binary(BinaryOp::Add, &*a, b);
1921});
1922auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Complex64| {
1923 *a = binary(BinaryOp::Add, &*a, b);
1924});
1925auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Parameter| {
1926 *a = binary(BinaryOp::Add, &*a, b);
1927});
1928
1929auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Expr| {
1930 *a = binary(BinaryOp::Sub, &*a, b);
1931});
1932auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &f64| {
1933 *a = binary(BinaryOp::Sub, &*a, b);
1934});
1935auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Complex64| {
1936 *a = binary(BinaryOp::Sub, &*a, b);
1937});
1938auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Parameter| {
1939 *a = binary(BinaryOp::Sub, &*a, b);
1940});
1941
1942auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Expr| {
1943 *a = binary(BinaryOp::Mul, &*a, b);
1944});
1945auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &f64| {
1946 *a = binary(BinaryOp::Mul, &*a, b);
1947});
1948auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Complex64| {
1949 *a = binary(BinaryOp::Mul, &*a, b);
1950});
1951auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Parameter| {
1952 *a = binary(BinaryOp::Mul, &*a, b);
1953});
1954
1955auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Expr| {
1956 *a = binary(BinaryOp::Div, &*a, b);
1957});
1958auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &f64| {
1959 *a = binary(BinaryOp::Div, &*a, b);
1960});
1961auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Complex64| {
1962 *a = binary(BinaryOp::Div, &*a, b);
1963});
1964auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Parameter| {
1965 *a = binary(BinaryOp::Div, &*a, b);
1966});
1967
1968impl From<f64> for Expr {
1969 fn from(value: f64) -> Self {
1970 Self::new(DagNodeKind::RealConst(value))
1971 }
1972}
1973
1974impl From<&f64> for Expr {
1975 fn from(value: &f64) -> Self {
1976 Self::new(DagNodeKind::RealConst(*value))
1977 }
1978}
1979
1980impl From<Complex64> for Expr {
1981 fn from(value: Complex64) -> Self {
1982 Self::new(DagNodeKind::ComplexConst(value))
1983 }
1984}
1985
1986impl From<&Complex64> for Expr {
1987 fn from(value: &Complex64) -> Self {
1988 Self::new(DagNodeKind::ComplexConst(*value))
1989 }
1990}
1991
1992impl From<&Expr> for Expr {
1993 fn from(value: &Expr) -> Self {
1994 value.clone()
1995 }
1996}
1997
1998impl From<Parameter> for Expr {
1999 fn from(parameter: Parameter) -> Self {
2000 Expr::new(DagNodeKind::ScalarParam(parameter))
2001 }
2002}
2003
2004impl From<&Parameter> for Expr {
2005 fn from(parameter: &Parameter) -> Self {
2006 parameter.clone().into()
2007 }
2008}
2009
2010pub fn cis(phase: Expr) -> Expr {
2012 phase.cos() + Complex64::I * phase.sin()
2013}
2014
2015pub fn complex(re: impl Into<Expr>, im: impl Into<Expr>) -> Expr {
2017 Expr::new(DagNodeKind::Complex {
2018 re: re.into(),
2019 im: im.into(),
2020 })
2021}
2022
2023pub fn polar_complex(mag: impl Into<Expr>, phase: impl Into<Expr>) -> Expr {
2025 mag.into() * (Complex64::I * phase.into()).exp()
2026}
2027
2028pub fn event_scalar(name: impl Into<Arc<str>>) -> Expr {
2030 Expr::new(DagNodeKind::EventScalar(name.into()))
2031}
2032
2033pub fn event_p4_component(name: impl Into<Arc<str>>, component: P4Component) -> Expr {
2035 Expr::new(DagNodeKind::EventP4Component {
2036 name: name.into(),
2037 component,
2038 })
2039}
2040
2041pub fn atan2(y: impl Into<Expr>, x: impl Into<Expr>) -> Expr {
2043 binary(BinaryOp::Atan2, y, x)
2044}
2045
2046pub fn acos(value: impl Into<Expr>) -> Expr {
2048 value.into().acos()
2049}
2050
2051pub fn vector<E>(elements: impl IntoIterator<Item = E>) -> Expr
2053where
2054 E: Into<Expr>,
2055 Expr: From<E>,
2056{
2057 Expr::new(DagNodeKind::Vector {
2058 elements: elements.into_iter().map(Expr::from).collect(),
2059 })
2060}
2061
2062pub fn matrix<const R: usize, const C: usize, E>(elements: [[E; C]; R]) -> Expr
2064where
2065 E: Into<Expr>,
2066 Expr: From<E>,
2067{
2068 Expr::new(DagNodeKind::Matrix {
2069 rows: R,
2070 cols: C,
2071 elements: elements.into_iter().flatten().map(Expr::from).collect(),
2072 })
2073}
2074
2075pub fn matrix_from_flat<E>(
2083 rows: usize,
2084 cols: usize,
2085 elements: impl IntoIterator<Item = E>,
2086) -> Result<Expr, ExprShapeError>
2087where
2088 E: Into<Expr>,
2089 Expr: From<E>,
2090{
2091 if rows == 0 || cols == 0 {
2092 return Err(ExprShapeError::new(
2093 "matrix constructor",
2094 format!("matrix dimensions must be nonzero, got {rows}x{cols}"),
2095 ));
2096 }
2097 let expected = rows.checked_mul(cols).ok_or_else(|| {
2098 ExprShapeError::new("matrix constructor", "row/column product overflowed")
2099 })?;
2100 let elements = elements.into_iter().map(Expr::from).collect::<Vec<_>>();
2101 if elements.len() != expected {
2102 return Err(ExprShapeError::new(
2103 "matrix constructor",
2104 format!(
2105 "shape {rows}x{cols} requires {expected} elements, got {}",
2106 elements.len()
2107 ),
2108 ));
2109 }
2110 for element in &elements {
2111 element.expect_shape("matrix constructor", ExprShape::Scalar)?;
2112 }
2113 Ok(Expr::new(DagNodeKind::Matrix {
2114 rows,
2115 cols,
2116 elements,
2117 }))
2118}
2119
2120pub fn matmul(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2122 Expr::new(DagNodeKind::MatMul {
2123 lhs: lhs.into(),
2124 rhs: rhs.into(),
2125 })
2126}
2127
2128pub fn matvec(matrix: impl Into<Expr>, vector: impl Into<Expr>) -> Expr {
2130 Expr::new(DagNodeKind::MatVec {
2131 matrix: matrix.into(),
2132 vector: vector.into(),
2133 })
2134}
2135
2136pub fn dot(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2138 Expr::new(DagNodeKind::Dot {
2139 lhs: lhs.into(),
2140 rhs: rhs.into(),
2141 })
2142}
2143
2144pub fn solve(matrix: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2146 Expr::new(DagNodeKind::Solve {
2147 matrix: matrix.into(),
2148 rhs: rhs.into(),
2149 })
2150}
2151
2152fn unary(op: UnaryOp, expr: impl Into<Expr>) -> Expr {
2153 Expr::new(DagNodeKind::Unary {
2154 op,
2155 input: expr.into(),
2156 })
2157}
2158
2159fn binary(op: BinaryOp, lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2160 let lhs = lhs.into();
2161 let rhs = rhs.into();
2162 match (lhs.is_tensor_syntax(), rhs.is_tensor_syntax()) {
2163 (true, false) if op != BinaryOp::Atan2 => Expr::new(DagNodeKind::TensorScalar {
2164 op,
2165 tensor: lhs,
2166 scalar: rhs,
2167 scalar_on_left: false,
2168 }),
2169 (false, true) if matches!(op, BinaryOp::Add | BinaryOp::Sub | BinaryOp::Mul) => {
2170 Expr::new(DagNodeKind::TensorScalar {
2171 op,
2172 tensor: rhs,
2173 scalar: lhs,
2174 scalar_on_left: true,
2175 })
2176 }
2177 _ => Expr::new(DagNodeKind::Binary { op, lhs, rhs }),
2178 }
2179}
2180
2181fn lower_tensor_scalar(op: BinaryOp, tensor: Expr, scalar: Expr, scalar_on_left: bool) -> Expr {
2182 let element = |value: Expr| {
2183 if scalar_on_left {
2184 binary(op, scalar.clone(), value)
2185 } else {
2186 binary(op, value, scalar.clone())
2187 }
2188 };
2189 let shape = match &tensor.node.kind {
2190 DagNodeKind::Vector { elements } => ExprShape::Vector {
2191 len: elements.len(),
2192 },
2193 DagNodeKind::Matrix { rows, cols, .. } => ExprShape::Matrix {
2194 rows: *rows,
2195 cols: *cols,
2196 },
2197 _ => tensor.shape().expect("tensor-scalar shape was validated"),
2198 };
2199 match shape {
2200 ExprShape::Vector { len } => {
2201 let elements = (0..len).map(|index| {
2202 let value = match &tensor.node.kind {
2203 DagNodeKind::Vector { elements } => elements[index].clone(),
2204 _ => tensor.component(index),
2205 };
2206 element(value)
2207 });
2208 vector(elements)
2209 }
2210 ExprShape::Matrix { rows, cols } => {
2211 let elements = (0..rows).flat_map(|row| {
2212 (0..cols).map({
2213 let element = &element;
2214 let tensor = &tensor;
2215 move |col| {
2216 let value = match &tensor.node.kind {
2217 DagNodeKind::Matrix { elements, .. } => {
2218 elements[row * cols + col].clone()
2219 }
2220 _ => tensor.matrix_element(row, col),
2221 };
2222 element(value)
2223 }
2224 })
2225 });
2226 matrix_from_flat(rows, cols, elements).expect("tensor-scalar matrix shape is valid")
2227 }
2228 ExprShape::Scalar => unreachable!("tensor-scalar input is a tensor"),
2229 }
2230}
2231
2232#[derive(Clone, Debug, Serialize, Deserialize)]
2234pub struct ExprGraph {
2235 root: ExprId,
2236 nodes: Vec<ExprNode>,
2237 metadata: Vec<ExprMetadata>,
2238}
2239
2240#[doc(hidden)]
2246pub struct ExprGraphRebuilder<K> {
2247 nodes: Vec<ExprNode>,
2248 metadata: Vec<ExprMetadata>,
2249 remapped: HashMap<K, ExprId>,
2250}
2251
2252#[doc(hidden)]
2253impl<K> ExprGraphRebuilder<K>
2254where
2255 K: Eq + Hash,
2256{
2257 pub fn with_capacity(capacity: usize) -> Self {
2259 Self {
2260 nodes: Vec::with_capacity(capacity),
2261 metadata: Vec::with_capacity(capacity),
2262 remapped: HashMap::with_capacity(capacity),
2263 }
2264 }
2265
2266 pub fn remapped(&self, key: &K) -> Option<ExprId> {
2268 self.remapped.get(key).copied()
2269 }
2270
2271 pub fn nodes(&self) -> &[ExprNode] {
2273 &self.nodes
2274 }
2275
2276 pub fn metadata(&self) -> &[ExprMetadata] {
2278 &self.metadata
2279 }
2280
2281 pub fn alias(&mut self, key: K, id: ExprId) {
2290 assert!(
2291 !self.remapped.contains_key(&key),
2292 "a rebuild key may only be mapped once"
2293 );
2294 assert!(
2295 id.index() < self.nodes.len(),
2296 "a rebuild alias must reference an emitted node"
2297 );
2298 self.remapped.insert(key, id);
2299 }
2300
2301 pub fn emit_anonymous(&mut self, node: ExprNode, metadata: ExprMetadata) -> ExprId {
2310 let id = ExprId::from_index(self.nodes.len());
2311 assert!(
2312 node.children().all(|child| child.index() < id.index()),
2313 "rebuilt expression children must be emitted before their parent"
2314 );
2315 self.nodes.push(node);
2316 self.metadata.push(metadata);
2317 id
2318 }
2319
2320 pub fn emit(&mut self, key: K, node: ExprNode, metadata: ExprMetadata) -> ExprId {
2327 assert!(
2328 !self.remapped.contains_key(&key),
2329 "a rebuild key may only be mapped once"
2330 );
2331 let id = self.emit_anonymous(node, metadata);
2332 self.remapped.insert(key, id);
2333 id
2334 }
2335
2336 pub fn finish(self, root: ExprId) -> Result<ExprGraph, ExprGraphError> {
2338 ExprGraph::from_parts(root, self.nodes, self.metadata)
2339 }
2340}
2341
2342impl ExprGraph {
2343 #[doc(hidden)]
2349 pub fn reachable_post_order(&self, roots: impl IntoIterator<Item = ExprId>) -> Vec<ExprId> {
2350 let roots = roots.into_iter().collect::<Vec<_>>();
2351 let mut visited = HashSet::with_capacity(self.nodes.len());
2352 let mut stack = Vec::new();
2353 let mut order = Vec::new();
2354
2355 for root in roots.into_iter().rev() {
2356 stack.push((root, false));
2357 }
2358 while let Some((id, expanded)) = stack.pop() {
2359 if expanded {
2360 order.push(id);
2361 continue;
2362 }
2363 if self.node(id).is_none() {
2364 continue;
2365 }
2366 if !visited.insert(id) {
2367 continue;
2368 }
2369 stack.push((id, true));
2370 if let Some(node) = self.node(id) {
2371 for child in node.children().rev() {
2372 stack.push((child, false));
2373 }
2374 }
2375 }
2376 order
2377 }
2378
2379 pub fn with_parameters<N, I>(&self, updates: I) -> ParamResult<Self>
2391 where
2392 N: AsRef<str>,
2393 I: IntoIterator<Item = (N, ParameterUpdate)>,
2394 {
2395 let mut by_name = HashMap::new();
2396 let mut names = Vec::new();
2397 for (name, update) in updates {
2398 let name = name.as_ref().to_owned();
2399 update.validate()?;
2400 if by_name.insert(name.clone(), update).is_some() {
2401 return Err(ParamError::DuplicateName(name));
2402 }
2403 names.push(name);
2404 }
2405
2406 let mut found = HashSet::new();
2407 let mut graph = self.clone();
2408 for node in &mut graph.nodes {
2409 if let ExprNode::ScalarParam(parameter) = node
2410 && let Some(update) = by_name.get(parameter.name())
2411 {
2412 *parameter = parameter.with_update(update)?;
2413 found.insert(parameter.name().to_owned());
2414 }
2415 }
2416 if let Some(name) = names.iter().find(|name| !found.contains(*name)) {
2417 return Err(ParamError::UnknownName(name.clone()));
2418 }
2419
2420 let mut registry = crate::parameters::ParamRegistry::new();
2421 for node in &graph.nodes {
2422 if let ExprNode::ScalarParam(parameter) = node {
2423 registry.register(parameter.clone())?;
2424 }
2425 }
2426 registry.layout()?;
2427 Ok(graph)
2428 }
2429
2430 pub fn project_tags<'a>(&self, tags: impl IntoIterator<Item = &'a str>) -> Self {
2439 let tags: Vec<_> = tags.into_iter().collect();
2440 let mut rebuild = ExprGraphRebuilder::with_capacity(self.nodes.len());
2441 let root_key = (self.root, false);
2442 let mut visited = HashSet::with_capacity(self.nodes.len());
2443 let mut stack = vec![(root_key, false)];
2444 while let Some((key @ (old, retain_all), expanded)) = stack.pop() {
2445 if expanded {
2446 let old_metadata = &self.metadata[old.index()];
2447 let matches = old_metadata
2448 .tags
2449 .iter()
2450 .any(|tag| tags.contains(&tag.as_ref()));
2451 let node = if !retain_all && !old_metadata.tags.is_empty() && !matches {
2452 ExprNode::RealConst(0.0)
2453 } else {
2454 let retain_children = retain_all || matches;
2455 self.nodes[old.index()].map_children(|child| {
2456 rebuild
2457 .remapped(&(child, retain_children))
2458 .expect("tag projection emits children before parents")
2459 })
2460 };
2461 let metadata = if matches || retain_all {
2462 old_metadata.clone()
2463 } else {
2464 ExprMetadata::new(old_metadata.source)
2465 };
2466 rebuild.emit(key, node, metadata);
2467 continue;
2468 }
2469 if !visited.insert(key) {
2470 continue;
2471 }
2472 stack.push((key, true));
2473 let old_metadata = &self.metadata[old.index()];
2474 let matches = old_metadata
2475 .tags
2476 .iter()
2477 .any(|tag| tags.contains(&tag.as_ref()));
2478 if retain_all || old_metadata.tags.is_empty() || matches {
2479 let retain_children = retain_all || matches;
2480 for child in self.nodes[old.index()].children().rev() {
2481 stack.push(((child, retain_children), false));
2482 }
2483 }
2484 }
2485 let root = rebuild
2486 .remapped(&root_key)
2487 .expect("tag projection emits its root");
2488 rebuild
2489 .finish(root)
2490 .expect("tag projection rebuilds a valid expression graph")
2491 }
2492
2493 pub fn from_parts(
2504 root: ExprId,
2505 nodes: Vec<ExprNode>,
2506 metadata: Vec<ExprMetadata>,
2507 ) -> Result<Self, ExprGraphError> {
2508 if nodes.is_empty() {
2509 return Err(ExprGraphError::Empty);
2510 }
2511 if nodes.len() != metadata.len() {
2512 return Err(ExprGraphError::MetadataLength {
2513 node_len: nodes.len(),
2514 metadata_len: metadata.len(),
2515 });
2516 }
2517 if root.index() >= nodes.len() {
2518 return Err(ExprGraphError::InvalidRoot {
2519 root: root.index(),
2520 node_len: nodes.len(),
2521 });
2522 }
2523 for (index, node) in nodes.iter().enumerate() {
2524 for child in node.children() {
2525 if child.index() >= nodes.len() {
2526 return Err(ExprGraphError::InvalidChild {
2527 node: index,
2528 child: child.index(),
2529 });
2530 }
2531 if child.index() >= index {
2532 return Err(ExprGraphError::InvalidChildOrder {
2533 node: index,
2534 child: child.index(),
2535 });
2536 }
2537 }
2538 }
2539 Ok(Self {
2540 root,
2541 nodes,
2542 metadata,
2543 })
2544 }
2545
2546 pub fn root(&self) -> ExprId {
2548 self.root
2549 }
2550
2551 pub fn node(&self, id: ExprId) -> Option<&ExprNode> {
2553 self.nodes.get(id.index())
2554 }
2555
2556 pub fn nodes(&self) -> &[ExprNode] {
2558 &self.nodes
2559 }
2560
2561 pub fn metadata(&self, id: ExprId) -> Option<&ExprMetadata> {
2563 self.metadata.get(id.index())
2564 }
2565}
2566
2567pub(crate) fn node_children(node: &ExprNode) -> Vec<(String, ExprId)> {
2568 node.children()
2569 .enumerate()
2570 .map(|(index, child)| (node_child_label(node, index), child))
2571 .collect()
2572}
2573
2574fn node_child_label(node: &ExprNode, index: usize) -> String {
2575 match node {
2576 ExprNode::Unary { .. } | ExprNode::Component { .. } | ExprNode::MatrixElement { .. } => {
2577 "input".into()
2578 }
2579 ExprNode::Binary { .. } | ExprNode::MatMul { .. } | ExprNode::Dot { .. } => {
2580 if index == 0 { "lhs" } else { "rhs" }.into()
2581 }
2582 ExprNode::NaryAdd { .. } => format!("term[{index}]"),
2583 ExprNode::NaryMul { .. } => format!("factor[{index}]"),
2584 ExprNode::Complex { .. } => if index == 0 { "re" } else { "im" }.into(),
2585 ExprNode::Vector { .. } => format!("element[{index}]"),
2586 ExprNode::Matrix { cols, .. } => {
2587 format!("element[{},{}]", index / cols, index % cols)
2588 }
2589 ExprNode::MatVec { .. } => if index == 0 { "matrix" } else { "vector" }.into(),
2590 ExprNode::Solve { .. } => if index == 0 { "matrix" } else { "rhs" }.into(),
2591 ExprNode::RealConst(_)
2592 | ExprNode::ComplexConst(_)
2593 | ExprNode::ScalarParam(_)
2594 | ExprNode::EventScalar(_)
2595 | ExprNode::EventP4Component { .. } => unreachable!("leaf nodes have no child labels"),
2596 }
2597}
2598
2599#[derive(Default)]
2600struct GraphBuilder {
2601 nodes: Vec<ExprNode>,
2602 metadata: Vec<ExprMetadata>,
2603 ids: HashMap<usize, ExprId>,
2604}
2605
2606impl GraphBuilder {
2607 fn new() -> Self {
2608 Self::default()
2609 }
2610
2611 fn build(mut self, expr: &Expr) -> ExprGraph {
2612 let mut stack = vec![(expr.clone(), false)];
2613 while let Some((expr, expanded)) = stack.pop() {
2614 let key = Arc::as_ptr(&expr.node) as usize;
2615 if self.ids.contains_key(&key) {
2616 continue;
2617 }
2618 if expanded {
2619 let node = self.lower(&expr.node.kind);
2620 let id = ExprId::from_index(self.nodes.len());
2621 self.nodes.push(node);
2622 self.metadata.push(expr.node.metadata.clone());
2623 self.ids.insert(key, id);
2624 continue;
2625 }
2626 stack.push((expr.clone(), true));
2627 for index in (0..expr.node.kind.child_count()).rev() {
2628 stack.push((expr.node.kind.child_at(index).clone(), false));
2629 }
2630 }
2631 let root = self.id(expr);
2632 ExprGraph {
2633 root,
2634 nodes: self.nodes,
2635 metadata: self.metadata,
2636 }
2637 }
2638
2639 fn id(&self, expr: &Expr) -> ExprId {
2640 let key = Arc::as_ptr(&expr.node) as usize;
2641 self.ids[&key]
2642 }
2643
2644 fn lower(&self, kind: &DagNodeKind) -> ExprNode {
2645 match kind {
2646 DagNodeKind::RealConst(value) => ExprNode::RealConst(*value),
2647 DagNodeKind::ComplexConst(value) => ExprNode::ComplexConst(*value),
2648 DagNodeKind::ScalarParam(parameter) => ExprNode::ScalarParam(parameter.clone()),
2649 DagNodeKind::EventScalar(name) => ExprNode::EventScalar(Arc::clone(name)),
2650 DagNodeKind::EventP4Component { name, component } => ExprNode::EventP4Component {
2651 name: Arc::clone(name),
2652 component: *component,
2653 },
2654 DagNodeKind::Unary { op, input } => {
2655 let input = self.id(input);
2656 ExprNode::Unary { op: *op, input }
2657 }
2658 DagNodeKind::Binary { op, lhs, rhs } => {
2659 let lhs = self.id(lhs);
2660 let rhs = self.id(rhs);
2661 ExprNode::Binary { op: *op, lhs, rhs }
2662 }
2663 DagNodeKind::TensorScalar { .. } => {
2664 unreachable!("tensor-scalar nodes are lowered before graph construction")
2665 }
2666 DagNodeKind::Complex { re, im } => {
2667 let re = self.id(re);
2668 let im = self.id(im);
2669 ExprNode::Complex { re, im }
2670 }
2671 DagNodeKind::Vector { elements } => ExprNode::Vector {
2672 elements: elements.iter().map(|expr| self.id(expr)).collect(),
2673 },
2674 DagNodeKind::Matrix {
2675 rows,
2676 cols,
2677 elements,
2678 } => ExprNode::Matrix {
2679 rows: *rows,
2680 cols: *cols,
2681 elements: elements.iter().map(|expr| self.id(expr)).collect(),
2682 },
2683 DagNodeKind::Component { input, index } => {
2684 let input = self.id(input);
2685 ExprNode::Component {
2686 input,
2687 index: *index,
2688 }
2689 }
2690 DagNodeKind::MatrixElement { input, row, col } => {
2691 let input = self.id(input);
2692 ExprNode::MatrixElement {
2693 input,
2694 row: *row,
2695 col: *col,
2696 }
2697 }
2698 DagNodeKind::MatMul { lhs, rhs } => {
2699 let lhs = self.id(lhs);
2700 let rhs = self.id(rhs);
2701 ExprNode::MatMul { lhs, rhs }
2702 }
2703 DagNodeKind::MatVec { matrix, vector } => {
2704 let matrix = self.id(matrix);
2705 let vector = self.id(vector);
2706 ExprNode::MatVec { matrix, vector }
2707 }
2708 DagNodeKind::Dot { lhs, rhs } => {
2709 let lhs = self.id(lhs);
2710 let rhs = self.id(rhs);
2711 ExprNode::Dot { lhs, rhs }
2712 }
2713 DagNodeKind::Solve { matrix, rhs } => {
2714 let matrix = self.id(matrix);
2715 let rhs = self.id(rhs);
2716 ExprNode::Solve { matrix, rhs }
2717 }
2718 }
2719 }
2720}
2721
2722fn source_kind(kind: &DagNodeKind) -> ExprSourceKind {
2723 match kind {
2724 DagNodeKind::RealConst(_) | DagNodeKind::ComplexConst(_) => ExprSourceKind::Const,
2725 DagNodeKind::ScalarParam(_) => ExprSourceKind::Param,
2726 DagNodeKind::EventScalar(_) | DagNodeKind::EventP4Component { .. } => ExprSourceKind::Event,
2727 DagNodeKind::Unary { .. } => ExprSourceKind::Unary,
2728 DagNodeKind::Binary { .. } | DagNodeKind::TensorScalar { .. } => ExprSourceKind::Binary,
2729 DagNodeKind::Complex { .. } => ExprSourceKind::Complex,
2730 DagNodeKind::Vector { .. } | DagNodeKind::Component { .. } | DagNodeKind::Dot { .. } => {
2731 ExprSourceKind::Vector
2732 }
2733 DagNodeKind::Matrix { .. } | DagNodeKind::MatrixElement { .. } => ExprSourceKind::Matrix,
2734 DagNodeKind::MatMul { .. } | DagNodeKind::MatVec { .. } | DagNodeKind::Solve { .. } => {
2735 ExprSourceKind::LinearAlgebra
2736 }
2737 }
2738}
2739
2740#[cfg(test)]
2741mod tests {
2742 use super::*;
2743 use crate::parameter;
2744
2745 #[test]
2746 fn builds_target_syntax_without_layout_or_context() {
2747 let model = (Complex64::I * parameter!("y", initial : 1.0, bounds : (0.0, 2.0))
2748 + parameter!("x"))
2749 .norm_sqr();
2750
2751 let graph = model.to_graph();
2752 assert!(matches!(
2753 graph.node(graph.root()),
2754 Some(ExprNode::Unary {
2755 op: UnaryOp::NormSqr,
2756 ..
2757 })
2758 ));
2759 }
2760
2761 #[test]
2762 fn parameter_nodes_store_specs_but_do_not_make_layouts() {
2763 let graph = Expr::from(parameter!("x", initial: 1.0)).to_graph();
2764 assert!(matches!(
2765 graph.node(graph.root()),
2766 Some(ExprNode::ScalarParam(spec)) if spec.name() == "x"
2767 ));
2768 }
2769
2770 #[test]
2771 fn complex_constructor_builds_expression_node() {
2772 let graph = complex(parameter!("re"), parameter!("im")).to_graph();
2773
2774 assert!(matches!(
2775 graph.node(graph.root()),
2776 Some(ExprNode::Complex { .. })
2777 ));
2778 }
2779
2780 #[test]
2781 fn polar_complex_lowers_to_expression_graph() {
2782 let graph = polar_complex(parameter!("mag"), parameter!("phase")).to_graph();
2783
2784 assert!(graph.nodes().iter().any(|node| matches!(
2785 node,
2786 ExprNode::Unary {
2787 op: UnaryOp::Exp,
2788 ..
2789 }
2790 )));
2791 }
2792
2793 #[test]
2794 fn metadata_survives_graph_construction() {
2795 let graph = event_scalar("mass")
2796 .named("event mass")
2797 .tagged("data")
2798 .tagged("data")
2799 .to_graph();
2800 let metadata = graph.metadata(graph.root()).unwrap();
2801 assert_eq!(metadata.name(), Some("event mass"));
2802 assert_eq!(metadata.tags(), &[Arc::from("data")]);
2803 assert!(metadata.has_tag("data"));
2804 }
2805
2806 #[test]
2807 fn parameter_updates_rewrite_all_same_named_occurrences() {
2808 let expression = Expr::from(parameter!("x")) + Expr::from(parameter!("x"));
2809 let graph = expression
2810 .to_graph()
2811 .with_parameters([(
2812 String::from("x"),
2813 ParameterUpdate {
2814 state: Some(ParamState::Fixed(2.0)),
2815 ..Default::default()
2816 },
2817 )])
2818 .unwrap();
2819
2820 assert_eq!(
2821 graph
2822 .nodes()
2823 .iter()
2824 .filter_map(|node| match node {
2825 ExprNode::ScalarParam(parameter) => Some(parameter),
2826 _ => None,
2827 })
2828 .count(),
2829 2
2830 );
2831 assert!(graph.nodes().iter().all(|node| {
2832 !matches!(node, ExprNode::ScalarParam(parameter) if parameter.name() == "x" && parameter.is_free())
2833 }));
2834 }
2835
2836 #[test]
2837 fn parameter_updates_reject_duplicates_and_unknown_names_without_mutation() {
2838 let graph = Expr::from(parameter!("x")).to_graph();
2839 let original = graph.to_string();
2840 assert!(matches!(
2841 graph.with_parameters([
2842 ("x", ParameterUpdate::default()),
2843 ("x", ParameterUpdate::default()),
2844 ]),
2845 Err(ParamError::DuplicateName(name)) if name == "x"
2846 ));
2847 assert!(matches!(
2848 graph.with_parameters([("missing", ParameterUpdate::default())]),
2849 Err(ParamError::UnknownName(name)) if name == "missing"
2850 ));
2851 assert_eq!(graph.to_string(), original);
2852 }
2853
2854 #[test]
2855 fn expressions_round_trip_through_serde_with_metadata() {
2856 let expression = ((parameter!("x", initial: 1.0) + 2.0).named("offset")
2857 * event_scalar("mass").tagged("data"))
2858 .tagged("model");
2859 let encoded = serde_json::to_string(&expression).unwrap();
2860 let decoded: Expr = serde_json::from_str(&encoded).unwrap();
2861
2862 assert_eq!(
2863 serde_json::to_value(expression.to_graph()).unwrap(),
2864 serde_json::to_value(decoded.to_graph()).unwrap()
2865 );
2866 }
2867
2868 #[test]
2869 fn display_formats_graph_as_labeled_tree() {
2870 let graph = ((parameter!("x") + 1.0).named("offset") * event_scalar("mass").tagged("data"))
2871 .to_graph();
2872 let display = graph.display_tree().to_string();
2873
2874 assert!(display.starts_with("ExprGraph(root=#"));
2875 assert!(display.contains("Binary(Mul)"));
2876 assert!(display.contains("┣ lhs:"));
2877 assert!(display.contains("┗ rhs:"));
2878 assert!(display.contains("Binary(Add) name=\"offset\""));
2879 assert!(display.contains("ScalarParam(x)"));
2880 assert!(display.contains("RealConst(1)"));
2881 assert!(display.contains("EventScalar(mass) tags=[data]"));
2882 }
2883
2884 #[test]
2885 fn display_formats_graph_as_expression() {
2886 let costheta = Expr::from(parameter!("costheta"));
2887 let phi = event_scalar("phi");
2888 let p = Expr::from(parameter!("p"));
2889 let phase = Expr::from(7.0) * Complex64::I;
2890 let graph =
2891 (((costheta.powi(2) * phi.sin()) - 5.2).norm_sqr() * p.conj() - phase.exp()).to_graph();
2892
2893 assert_eq!(
2894 graph.to_string(),
2895 "|costheta^2 * sin(phi) - 5.2|^2 * conj(p) - exp(7 * i)"
2896 );
2897 }
2898
2899 #[test]
2900 fn display_parenthesizes_when_precedence_requires_it() {
2901 let a = Expr::from(parameter!("a"));
2902 let b = Expr::from(parameter!("b"));
2903 let c = Expr::from(parameter!("c"));
2904
2905 assert_eq!(
2906 (a.clone() * (b.clone() + c.clone())).to_graph().to_string(),
2907 "a * (b + c)"
2908 );
2909 assert_eq!(
2910 (a.clone() - (b.clone() - c.clone())).to_graph().to_string(),
2911 "a - (b - c)"
2912 );
2913 assert_eq!(((a / b) / c).to_graph().to_string(), "a / b / c");
2914 }
2915
2916 #[test]
2917 fn display_rounds_tiny_float_representation_noise() {
2918 let metadata = ExprMetadata::new(ExprSourceKind::Const);
2919 let graph = ExprGraph::from_parts(
2920 ExprId::from_index(2),
2921 vec![
2922 ExprNode::RealConst(2.9999999999999996),
2923 ExprNode::ComplexConst(Complex64::new(0.30000000000000004, 1.9999999999999998)),
2924 ExprNode::Binary {
2925 op: BinaryOp::Add,
2926 lhs: ExprId::from_index(0),
2927 rhs: ExprId::from_index(1),
2928 },
2929 ],
2930 vec![metadata.clone(), metadata.clone(), metadata],
2931 )
2932 .unwrap();
2933
2934 assert_eq!(graph.to_string(), "3 + 0.3 + 2i");
2935 assert!(graph.display_tree().to_string().contains("RealConst(3)"));
2936 assert!(
2937 graph
2938 .display_tree()
2939 .to_string()
2940 .contains("ComplexConst(0.3 + 2i)")
2941 );
2942 }
2943
2944 #[test]
2945 fn display_formats_p4_components_and_atan2() {
2946 let expr = atan2(
2947 event_p4_component("ks1", P4Component::Py),
2948 event_p4_component("ks1", P4Component::Px),
2949 );
2950
2951 assert_eq!(expr.to_graph().to_string(), "atan2(ks1.py, ks1.px)");
2952 }
2953
2954 #[test]
2955 fn graph_from_parts_validates_structure() {
2956 let metadata = ExprMetadata::new(ExprSourceKind::Const);
2957 let graph = ExprGraph::from_parts(
2958 ExprId::from_index(1),
2959 vec![
2960 ExprNode::RealConst(1.0),
2961 ExprNode::Unary {
2962 op: UnaryOp::Neg,
2963 input: ExprId::from_index(0),
2964 },
2965 ],
2966 vec![metadata.clone(), metadata.clone()],
2967 )
2968 .unwrap();
2969 assert!(matches!(
2970 graph.node(graph.root()),
2971 Some(ExprNode::Unary {
2972 op: UnaryOp::Neg,
2973 ..
2974 })
2975 ));
2976
2977 let err = ExprGraph::from_parts(
2978 ExprId::from_index(0),
2979 vec![ExprNode::RealConst(1.0)],
2980 Vec::new(),
2981 )
2982 .unwrap_err();
2983 assert!(matches!(err, ExprGraphError::MetadataLength { .. }));
2984
2985 let err = ExprGraph::from_parts(
2986 ExprId::from_index(0),
2987 vec![ExprNode::Unary {
2988 op: UnaryOp::Neg,
2989 input: ExprId::from_index(0),
2990 }],
2991 vec![metadata],
2992 )
2993 .unwrap_err();
2994 assert!(matches!(err, ExprGraphError::InvalidChildOrder { .. }));
2995 }
2996
2997 #[test]
2998 fn graph_preserves_unsimplified_expression_shape() {
2999 let graph = (parameter!("x") + 0.0).to_graph();
3000 assert!(matches!(
3001 graph.node(graph.root()),
3002 Some(ExprNode::Binary {
3003 op: BinaryOp::Add,
3004 ..
3005 })
3006 ));
3007 }
3008
3009 #[test]
3010 fn graph_preserves_written_operand_order_for_commutative_ops() {
3011 let left_param = (parameter!("x") + 1.0).to_graph();
3012 assert!(matches!(
3013 left_param.node(left_param.root()),
3014 Some(ExprNode::Binary {
3015 op: BinaryOp::Add,
3016 lhs,
3017 rhs
3018 }) if matches!(left_param.node(*lhs), Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "x")
3019 && matches!(left_param.node(*rhs), Some(ExprNode::RealConst(1.0)))
3020 ));
3021
3022 let right_param = (1.0 + parameter!("x")).to_graph();
3023 assert!(matches!(
3024 right_param.node(right_param.root()),
3025 Some(ExprNode::Binary {
3026 op: BinaryOp::Add,
3027 lhs,
3028 rhs
3029 }) if matches!(right_param.node(*lhs), Some(ExprNode::RealConst(1.0)))
3030 && matches!(right_param.node(*rhs), Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "x")
3031 ));
3032 }
3033
3034 #[test]
3035 fn represents_kmatrix_style_solve_graph() {
3036 let beta = vector([
3037 complex(parameter!("b0_re"), parameter!("b0_im")),
3038 complex(parameter!("b1_re"), parameter!("b1_im")),
3039 ]);
3040 let a = matrix([
3041 [Complex64::new(1.0, 0.0), Complex64::new(0.0, 1.0)],
3042 [Complex64::new(0.0, -1.0), Complex64::new(1.0, 0.0)],
3043 ]);
3044 let graph = solve(a, beta).component(0).to_graph();
3045
3046 assert!(
3047 graph
3048 .nodes()
3049 .iter()
3050 .any(|node| matches!(node, ExprNode::Solve { .. }))
3051 );
3052 }
3053
3054 #[test]
3055 fn graph_builder_preserves_shared_dag_nodes() {
3056 let shared = event_scalar("x").sin();
3057 let expression = vector((0..1_000).map(|_| shared.clone()));
3058 let graph = expression.to_graph();
3059
3060 assert_eq!(graph.nodes().len(), 3);
3061 let ExprNode::Vector { elements } = graph.node(graph.root()).unwrap() else {
3062 panic!("root should be a vector");
3063 };
3064 assert!(elements.windows(2).all(|pair| pair[0] == pair[1]));
3065 }
3066
3067 #[test]
3068 fn expression_projection_preserves_occurrence_rebuild_behavior() {
3069 let shared = event_scalar("x").sin();
3070 let projected = (shared.clone() + shared).project_tags(["selected"]);
3071 let graph = projected.to_graph();
3072
3073 assert_eq!(
3074 graph
3075 .nodes()
3076 .iter()
3077 .filter(|node| matches!(
3078 node,
3079 ExprNode::Unary {
3080 op: UnaryOp::Sin,
3081 ..
3082 }
3083 ))
3084 .count(),
3085 2
3086 );
3087 }
3088
3089 #[test]
3090 fn iterative_construction_and_projection_handle_deep_expressions() {
3091 let mut expression = event_scalar("x");
3092 for _ in 0..10_000 {
3093 expression = expression.sin();
3094 }
3095
3096 let projected = expression.project_tags(["selected"]);
3097 let graph = projected.to_graph();
3098
3099 assert_eq!(graph.nodes().len(), 10_001);
3100 assert_eq!(
3101 graph.reachable_post_order([graph.root()]).len(),
3102 graph.nodes().len()
3103 );
3104
3105 std::mem::forget(expression);
3108 std::mem::forget(projected);
3109 }
3110
3111 #[test]
3112 fn tensor_scalar_lowering_handles_a_deep_scalar_element() {
3113 let mut scalar = event_scalar("x");
3114 for _ in 0..10_000 {
3115 scalar = scalar.sin();
3116 }
3117 let expression = vector([scalar.clone()]) * 2.0;
3118 let graph = expression.to_graph();
3119 assert!(
3120 matches!(graph.node(graph.root()), Some(ExprNode::Vector { elements }) if elements.len() == 1)
3121 );
3122 std::mem::forget(scalar);
3123 std::mem::forget(expression);
3124 }
3125
3126 #[test]
3127 fn reachable_post_order_preserves_child_order_and_deduplicates_shared_nodes() {
3128 let shared = event_scalar("x").sin();
3129 let graph = (shared.clone() + shared).to_graph();
3130 let order = graph.reachable_post_order([graph.root()]);
3131
3132 assert_eq!(order.len(), graph.nodes().len());
3133 assert_eq!(order.last(), Some(&graph.root()));
3134 for id in order {
3135 for child in graph.node(id).unwrap().children() {
3136 assert!(child.index() < id.index());
3137 }
3138 }
3139 }
3140
3141 #[test]
3142 fn dynamic_matrices_and_shapes_are_checked_eagerly() {
3143 let dynamic = matrix_from_flat(2, 2, [1.0, 2.0, 3.0, 4.0]).unwrap();
3144 assert_eq!(
3145 dynamic.shape().unwrap(),
3146 ExprShape::Matrix { rows: 2, cols: 2 }
3147 );
3148 assert!(matrix_from_flat(2, 2, [1.0, 2.0, 3.0]).is_err());
3149 assert!(matmul(dynamic, matrix([[1.0, 2.0, 3.0]])).shape().is_err());
3150 }
3151
3152 #[test]
3153 fn assignment_operators_build_binary_expression_nodes() {
3154 let mut expr = Expr::from(parameter!("x"));
3155 expr += parameter!("y");
3156 expr -= 1.0;
3157 expr *= Complex64::I;
3158 expr /= Expr::from(parameter!("z"));
3159
3160 let graph = expr.to_graph();
3161 assert!(matches!(
3162 graph.node(graph.root()),
3163 Some(ExprNode::Binary {
3164 op: BinaryOp::Div,
3165 ..
3166 })
3167 ));
3168 assert_eq!(
3169 graph
3170 .nodes()
3171 .iter()
3172 .filter(|node| matches!(node, ExprNode::Binary { .. }))
3173 .count(),
3174 4
3175 );
3176 }
3177
3178 #[test]
3179 fn assignment_operators_accept_borrowed_rhs_values() {
3180 let y = parameter!("y");
3181 let one = 1.0;
3182 let i = Complex64::I;
3183 let z = Expr::from(parameter!("z"));
3184
3185 let mut expr = Expr::from(parameter!("x"));
3186 expr += &y;
3187 expr -= &one;
3188 expr *= &i;
3189 expr /= &z;
3190
3191 let graph = expr.to_graph();
3192 assert!(matches!(
3193 graph.node(graph.root()),
3194 Some(ExprNode::Binary {
3195 op: BinaryOp::Div,
3196 ..
3197 })
3198 ));
3199 }
3200}