1use std::hash::{Hash, Hasher};
2
3use laddu_expr::{
4 BinaryOp, ExprGraph, ExprGraphRebuilder, ExprId, ExprMetadata, ExprNode, ExprSourceKind,
5 UnaryOp, ValueKind,
6};
7use num::complex::Complex64;
8
9use crate::{CompileError, CompileResult, CompiledModel, GraphFacts, graph_utils::compact_to_root};
10
11const DEFAULT_EXPANSION_BUDGET: usize = 4_096;
12
13#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
15pub enum NormalizationStrategy {
16 Hermitian,
18 LinearStatistics,
20 Hybrid,
22 General,
24}
25
26#[derive(Clone, Debug, PartialEq, Eq)]
28pub enum NormalizationFallbackReason {
29 NonScalarIntensity,
31 UnsupportedMixedOperation {
33 node: ExprId,
35 operation: &'static str,
37 },
38 ExpansionBudgetExceeded {
40 budget: usize,
42 },
43}
44
45#[derive(Clone, Debug, PartialEq, Eq)]
47pub struct NormalizationDiagnostics {
48 strategy: NormalizationStrategy,
49 basis_count: usize,
50 coherent_group_count: usize,
51 has_residual: bool,
52 fallback_reason: Option<NormalizationFallbackReason>,
53}
54
55impl NormalizationDiagnostics {
56 pub fn strategy(&self) -> NormalizationStrategy {
58 self.strategy
59 }
60
61 pub fn basis_count(&self) -> usize {
63 self.basis_count
64 }
65
66 pub fn coherent_group_count(&self) -> usize {
68 self.coherent_group_count
69 }
70
71 pub fn has_residual(&self) -> bool {
73 self.has_residual
74 }
75
76 pub fn fallback_reason(&self) -> Option<&NormalizationFallbackReason> {
78 self.fallback_reason.as_ref()
79 }
80}
81
82#[derive(Copy, Clone, Debug)]
83struct SeparableTerm {
84 coefficient: ExprId,
85 basis: ExprId,
86}
87
88#[doc(hidden)]
95#[derive(Clone, Debug)]
96pub struct NormalizationPlan {
97 graph: ExprGraph,
98 terms: Vec<SeparableTerm>,
99 residual: Option<ExprId>,
100 diagnostics: NormalizationDiagnostics,
101 proven_nonnegative: bool,
102}
103
104impl NormalizationPlan {
105 pub(crate) fn hash_structure<H: Hasher>(&self, state: &mut H) {
106 self.graph.root().hash(state);
107 self.graph.nodes().len().hash(state);
108 for node in self.graph.nodes() {
109 node.structural_key().hash(state);
110 }
111 self.terms.len().hash(state);
112 for term in &self.terms {
113 term.coefficient.hash(state);
114 term.basis.hash(state);
115 }
116 self.residual.hash(state);
117 self.diagnostics.strategy.hash(state);
118 self.proven_nonnegative.hash(state);
119 }
120
121 pub(crate) fn analyze_disabled(graph: &ExprGraph) -> Self {
122 Self::general(
123 graph,
124 NormalizationFallbackReason::UnsupportedMixedOperation {
125 node: graph.root(),
126 operation: "normalization analysis disabled",
127 },
128 )
129 }
130
131 pub(crate) fn analyze(graph: &ExprGraph, facts: &GraphFacts) -> Self {
132 if !matches!(
133 facts.get(graph.root()).map(|facts| facts.value_kind),
134 Some(ValueKind::Real | ValueKind::Complex)
135 ) {
136 return Self::general(graph, NormalizationFallbackReason::NonScalarIntensity);
137 }
138
139 let decomposition = NormalizationAnalyzer::new(graph, facts, DEFAULT_EXPANSION_BUDGET)
140 .analyze(graph.root());
141
142 if decomposition.terms.is_empty() {
143 return Self::general(
144 graph,
145 decomposition.last_failure().cloned().unwrap_or(
146 NormalizationFallbackReason::UnsupportedMixedOperation {
147 node: graph.root(),
148 operation: "root",
149 },
150 ),
151 );
152 }
153
154 let built = decomposition
155 .build(graph)
156 .expect("normalization decomposition emits a valid graph");
157 let terms = built.terms;
158 let residual = built.residual;
159 let strategy = if residual.is_some() {
160 NormalizationStrategy::Hybrid
161 } else if built.coherent_groups > 0 {
162 NormalizationStrategy::Hermitian
163 } else {
164 NormalizationStrategy::LinearStatistics
165 };
166 let diagnostics = NormalizationDiagnostics {
167 strategy,
168 basis_count: terms.len(),
169 coherent_group_count: built.coherent_groups,
170 has_residual: residual.is_some(),
171 fallback_reason: built.last_failure,
172 };
173 Self {
174 graph: built.graph,
175 terms,
176 residual,
177 diagnostics,
178 proven_nonnegative: proves_nonnegative(graph, facts, graph.root()),
179 }
180 }
181
182 fn general(graph: &ExprGraph, reason: NormalizationFallbackReason) -> Self {
183 Self {
184 graph: graph.clone(),
185 terms: Vec::new(),
186 residual: None,
187 diagnostics: NormalizationDiagnostics {
188 strategy: NormalizationStrategy::General,
189 basis_count: 0,
190 coherent_group_count: 0,
191 has_residual: false,
192 fallback_reason: Some(reason),
193 },
194 proven_nonnegative: false,
195 }
196 }
197
198 pub fn diagnostics(&self) -> &NormalizationDiagnostics {
200 &self.diagnostics
201 }
202
203 pub fn proven_nonnegative(&self) -> bool {
205 self.proven_nonnegative
206 }
207
208 pub fn basis_models(&self) -> CompileResult<Vec<CompiledModel>> {
214 self.terms
215 .iter()
216 .map(|term| {
217 CompiledModel::from_graph_without_normalization(compact_to_root(
218 &self.graph,
219 term.basis,
220 )?)
221 })
222 .collect()
223 }
224
225 pub fn evaluator_model(&self, statistics: &[Complex64]) -> CompileResult<CompiledModel> {
232 if statistics.len() != self.terms.len() {
233 return Err(CompileError::InvalidExecutablePlan(format!(
234 "normalization expected {} statistics, got {}",
235 self.terms.len(),
236 statistics.len()
237 )));
238 }
239 let mut builder = NormalizationGraphBuilder::new(&self.graph);
240 let mut products = Vec::with_capacity(self.terms.len());
241 for (term, statistic) in self.terms.iter().zip(statistics) {
242 let constant = builder.constant(*statistic);
243 let coefficient = builder.source(term.coefficient);
244 products.push(builder.product(&[constant, coefficient]));
245 }
246 let sum = builder.sum(&products);
247 let root = builder.unary(UnaryOp::Real, sum);
248 let graph = builder.finish(root)?;
249 CompiledModel::from_graph_without_normalization(compact_to_root(&graph, root)?)
250 }
251
252 pub fn residual_model(&self) -> CompileResult<Option<CompiledModel>> {
258 self.residual
259 .map(|root| {
260 CompiledModel::from_graph_without_normalization(compact_to_root(&self.graph, root)?)
261 })
262 .transpose()
263 }
264}
265
266fn proves_nonnegative(graph: &ExprGraph, facts: &GraphFacts, id: ExprId) -> bool {
267 match graph.node(id).expect("normalization node exists") {
268 ExprNode::RealConst(value) => value.is_finite() && *value >= 0.0,
269 ExprNode::ComplexConst(value) => value.im == 0.0 && value.re.is_finite() && value.re >= 0.0,
270 ExprNode::Unary {
271 op: UnaryOp::NormSqr,
272 ..
273 } => true,
274 ExprNode::Unary {
275 op: UnaryOp::Real,
276 input,
277 } => proves_nonnegative(graph, facts, *input),
278 ExprNode::Unary {
279 op: UnaryOp::PowI(power),
280 input,
281 } => {
282 *power >= 0
283 && power % 2 == 0
284 && facts
285 .get(*input)
286 .is_some_and(|facts| facts.value_kind == ValueKind::Real)
287 }
288 ExprNode::Binary {
289 op: BinaryOp::Add,
290 lhs,
291 rhs,
292 } => proves_nonnegative(graph, facts, *lhs) && proves_nonnegative(graph, facts, *rhs),
293 ExprNode::Binary {
294 op: BinaryOp::Mul,
295 lhs,
296 rhs,
297 } => proves_nonnegative(graph, facts, *lhs) && proves_nonnegative(graph, facts, *rhs),
298 ExprNode::NaryAdd { terms } => terms
299 .iter()
300 .all(|term| proves_nonnegative(graph, facts, *term)),
301 ExprNode::NaryMul { factors } => factors
302 .iter()
303 .all(|factor| proves_nonnegative(graph, facts, *factor)),
304 _ => false,
305 }
306}
307
308#[derive(Copy, Clone, Debug, PartialEq, Eq)]
309enum DecompositionNode {
310 Source(ExprId),
311 Generated(usize),
312}
313
314#[derive(Clone, Debug)]
315enum GraphOperation {
316 Constant(Complex64),
317 Unary {
318 op: UnaryOp,
319 input: DecompositionNode,
320 },
321 Binary {
322 op: BinaryOp,
323 lhs: DecompositionNode,
324 rhs: DecompositionNode,
325 },
326 Product(Vec<DecompositionNode>),
327 Sum(Vec<DecompositionNode>),
328}
329
330#[derive(Clone, Debug)]
331struct AnalyzedTerm {
332 coefficient: DecompositionNode,
333 basis: DecompositionNode,
334}
335
336#[derive(Clone, Debug)]
337struct Decomposition {
338 operations: Vec<GraphOperation>,
339 terms: Vec<AnalyzedTerm>,
340 residual: Option<DecompositionNode>,
341 coherent_groups: usize,
342 failures: Vec<NormalizationFallbackReason>,
343}
344
345impl Decomposition {
346 fn last_failure(&self) -> Option<&NormalizationFallbackReason> {
347 self.failures.last()
348 }
349
350 fn build(self, graph: &ExprGraph) -> CompileResult<BuiltDecomposition> {
351 let mut builder = NormalizationGraphBuilder::new(graph);
352 let mut generated = Vec::with_capacity(self.operations.len());
353 for operation in self.operations {
354 let resolve = |node: DecompositionNode| match node {
355 DecompositionNode::Source(id) => builder.source(id),
356 DecompositionNode::Generated(index) => generated[index],
357 };
358 let id = match operation {
359 GraphOperation::Constant(value) => builder.constant(value),
360 GraphOperation::Unary { op, input } => {
361 let input = resolve(input);
362 builder.unary(op, input)
363 }
364 GraphOperation::Binary { op, lhs, rhs } => {
365 let lhs = resolve(lhs);
366 let rhs = resolve(rhs);
367 builder.binary(op, lhs, rhs)
368 }
369 GraphOperation::Product(factors) => {
370 let factors = factors.into_iter().map(resolve).collect::<Vec<_>>();
371 builder.product(&factors)
372 }
373 GraphOperation::Sum(terms) => {
374 let terms = terms.into_iter().map(resolve).collect::<Vec<_>>();
375 builder.sum(&terms)
376 }
377 };
378 generated.push(id);
379 }
380 let resolve = |node: DecompositionNode| match node {
381 DecompositionNode::Source(id) => builder.source(id),
382 DecompositionNode::Generated(index) => generated[index],
383 };
384 let terms = self
385 .terms
386 .into_iter()
387 .map(|term| SeparableTerm {
388 coefficient: resolve(term.coefficient),
389 basis: resolve(term.basis),
390 })
391 .collect();
392 let residual = self.residual.map(resolve);
393 let root = builder.source(ExprId::from_index(0));
394 Ok(BuiltDecomposition {
395 graph: builder.finish(root)?,
396 terms,
397 residual,
398 coherent_groups: self.coherent_groups,
399 last_failure: self.failures.last().cloned(),
400 })
401 }
402}
403
404struct BuiltDecomposition {
405 graph: ExprGraph,
406 terms: Vec<SeparableTerm>,
407 residual: Option<ExprId>,
408 coherent_groups: usize,
409 last_failure: Option<NormalizationFallbackReason>,
410}
411
412#[derive(Copy, Clone, Debug)]
413struct ExpansionBudget(usize);
414
415impl ExpansionBudget {
416 fn ensure(self, count: usize) -> Result<(), NormalizationFallbackReason> {
417 if count <= self.0 {
418 Ok(())
419 } else {
420 Err(self.exceeded())
421 }
422 }
423
424 fn product_count(self, lhs: usize, rhs: usize) -> Result<usize, NormalizationFallbackReason> {
425 let count = lhs.checked_mul(rhs).ok_or_else(|| self.exceeded())?;
426 self.ensure(count)?;
427 Ok(count)
428 }
429
430 fn packed_triangle_count(self, count: usize) -> Result<usize, NormalizationFallbackReason> {
431 let next = count.checked_add(1).ok_or_else(|| self.exceeded())?;
432 let packed = count.checked_mul(next).ok_or_else(|| self.exceeded())? / 2;
433 self.ensure(packed)?;
434 Ok(packed)
435 }
436
437 fn exceeded(self) -> NormalizationFallbackReason {
438 NormalizationFallbackReason::ExpansionBudgetExceeded { budget: self.0 }
439 }
440}
441
442struct NormalizationAnalyzer<'a> {
443 graph: &'a ExprGraph,
444 facts: &'a GraphFacts,
445 operations: Vec<GraphOperation>,
446 one: DecompositionNode,
447 budget: ExpansionBudget,
448 coherent_groups: usize,
449}
450
451impl<'a> NormalizationAnalyzer<'a> {
452 fn new(graph: &'a ExprGraph, facts: &'a GraphFacts, budget: usize) -> Self {
453 let mut analyzer = Self {
454 graph,
455 facts,
456 operations: Vec::new(),
457 one: DecompositionNode::Source(graph.root()),
458 budget: ExpansionBudget(budget),
459 coherent_groups: 0,
460 };
461 analyzer.one = analyzer.constant(Complex64::new(1.0, 0.0));
462 analyzer
463 }
464
465 fn analyze(mut self, root: ExprId) -> Decomposition {
466 let mut terms = Vec::new();
467 let mut residuals = Vec::new();
468 let mut failures = Vec::new();
469 for root in self.additive_roots(root) {
470 match self.decompose(root) {
471 Ok(mut extracted) => terms.append(&mut extracted),
472 Err(reason) => {
473 failures.push(reason);
474 residuals.push(DecompositionNode::Source(root));
475 }
476 }
477 }
478 let residual = self.sum(&residuals);
479 Decomposition {
480 operations: self.operations,
481 terms,
482 residual,
483 coherent_groups: self.coherent_groups,
484 failures,
485 }
486 }
487
488 fn dependency(&self, id: ExprId) -> crate::DependencyFacts {
489 self.facts.get(id).expect("facts are complete").dependency
490 }
491
492 fn additive_roots(&self, root: ExprId) -> Vec<ExprId> {
493 match self.graph.node(root).expect("normalization node exists") {
494 ExprNode::NaryAdd { terms } => terms.clone(),
495 ExprNode::Binary {
496 op: BinaryOp::Add,
497 lhs,
498 rhs,
499 } => {
500 let mut roots = self.additive_roots(*lhs);
501 roots.extend(self.additive_roots(*rhs));
502 roots
503 }
504 _ => vec![root],
505 }
506 }
507
508 fn decompose(&mut self, id: ExprId) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
509 let dependency = self.dependency(id);
510 if !dependency.depends_on_event {
511 return Ok(vec![AnalyzedTerm {
512 coefficient: DecompositionNode::Source(id),
513 basis: self.one,
514 }]);
515 }
516 if !dependency.depends_on_free_params {
517 return Ok(vec![AnalyzedTerm {
518 coefficient: self.one,
519 basis: DecompositionNode::Source(id),
520 }]);
521 }
522
523 match self
524 .graph
525 .node(id)
526 .expect("normalization node exists")
527 .clone()
528 {
529 ExprNode::Binary { op, lhs, rhs } => self.decompose_binary(id, op, lhs, rhs),
530 ExprNode::NaryAdd { terms } => {
531 let mut result = Vec::new();
532 for term in terms {
533 result.extend(self.decompose(term)?);
534 self.ensure_budget(&result)?;
535 }
536 Ok(result)
537 }
538 ExprNode::NaryMul { factors } => {
539 if let [lhs, rhs] = factors.as_slice()
540 && let Some((scale, amplitude)) = self.scaled_coherent_parts(*lhs, *rhs)
541 {
542 return self.decompose_scaled_coherent(scale, amplitude);
543 }
544 let mut result = vec![AnalyzedTerm {
545 coefficient: self.one,
546 basis: self.one,
547 }];
548 for factor in factors {
549 let factor_terms = self.decompose(factor)?;
550 result = self.multiply_terms(&result, &factor_terms)?;
551 }
552 Ok(result)
553 }
554 ExprNode::Unary { op, input } => self.decompose_unary(id, op, input),
555 _ => Err(self.unsupported(id, "structured mixed operation")),
556 }
557 }
558
559 fn decompose_binary(
560 &mut self,
561 id: ExprId,
562 op: BinaryOp,
563 lhs: ExprId,
564 rhs: ExprId,
565 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
566 match op {
567 BinaryOp::Add => {
568 let mut terms = self.decompose(lhs)?;
569 terms.extend(self.decompose(rhs)?);
570 self.ensure_budget(&terms)?;
571 Ok(terms)
572 }
573 BinaryOp::Sub => {
574 let mut terms = self.decompose(lhs)?;
575 for mut term in self.decompose(rhs)? {
576 term.coefficient = self.unary(UnaryOp::Neg, term.coefficient);
577 terms.push(term);
578 }
579 self.ensure_budget(&terms)?;
580 Ok(terms)
581 }
582 BinaryOp::Mul => {
583 if let Some((scale, amplitude)) = self.scaled_coherent_parts(lhs, rhs) {
584 return self.decompose_scaled_coherent(scale, amplitude);
585 }
586 let left = self.decompose(lhs)?;
587 let right = self.decompose(rhs)?;
588 self.multiply_terms(&left, &right)
589 }
590 BinaryOp::Div => {
591 let denominator = self.dependency(rhs);
592 let mut terms = self.decompose(lhs)?;
593 if !denominator.depends_on_event {
594 for term in &mut terms {
595 term.coefficient = self.binary(
596 BinaryOp::Div,
597 term.coefficient,
598 DecompositionNode::Source(rhs),
599 );
600 }
601 Ok(terms)
602 } else if !denominator.depends_on_free_params {
603 for term in &mut terms {
604 term.basis =
605 self.binary(BinaryOp::Div, term.basis, DecompositionNode::Source(rhs));
606 }
607 Ok(terms)
608 } else {
609 Err(self.unsupported(id, "mixed division"))
610 }
611 }
612 BinaryOp::Atan2 => Err(self.unsupported(id, "atan2")),
613 }
614 }
615
616 fn decompose_unary(
617 &mut self,
618 id: ExprId,
619 op: UnaryOp,
620 input: ExprId,
621 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
622 match op {
623 UnaryOp::Neg => {
624 let mut terms = self.decompose(input)?;
625 for term in &mut terms {
626 term.coefficient = self.unary(UnaryOp::Neg, term.coefficient);
627 }
628 Ok(terms)
629 }
630 UnaryOp::Conj => {
631 let mut terms = self.decompose(input)?;
632 for term in &mut terms {
633 term.coefficient = self.unary(UnaryOp::Conj, term.coefficient);
634 term.basis = self.unary(UnaryOp::Conj, term.basis);
635 }
636 Ok(terms)
637 }
638 UnaryOp::NormSqr => {
639 let terms = self.decompose(input)?;
640 self.pack_norm_sqr(&terms)
641 }
642 UnaryOp::PowI(power) if power >= 0 => {
643 let base = self.decompose(input)?;
644 let mut result = vec![AnalyzedTerm {
645 coefficient: self.one,
646 basis: self.one,
647 }];
648 for _ in 0..power {
649 result = self.multiply_terms(&result, &base)?;
650 }
651 Ok(result)
652 }
653 UnaryOp::Real | UnaryOp::Imag => self.decompose_projection(op, input),
654 UnaryOp::Sqrt
655 | UnaryOp::Exp
656 | UnaryOp::Sin
657 | UnaryOp::Cos
658 | UnaryOp::Log
659 | UnaryOp::PowI(_) => Err(self.unsupported(id, "nonlinear unary operation")),
660 }
661 }
662
663 fn coherent_scale_root(&self, id: ExprId) -> Option<ExprId> {
665 let root = match self.graph.node(id)? {
666 ExprNode::Unary {
667 op: UnaryOp::PowI(2),
668 input,
669 } => *input,
670 ExprNode::Binary {
671 op: BinaryOp::Mul,
672 lhs,
673 rhs,
674 } if lhs == rhs => *lhs,
675 ExprNode::NaryMul { factors } if factors.len() == 2 && factors[0] == factors[1] => {
676 factors[0]
677 }
678 _ => return None,
679 };
680 let facts = self.facts.get(root)?;
681 (facts.value_kind == ValueKind::Real && !facts.dependency.depends_on_event).then_some(root)
682 }
683
684 fn scaled_coherent_parts(&self, lhs: ExprId, rhs: ExprId) -> Option<(ExprId, ExprId)> {
685 for (scale, norm) in [(lhs, rhs), (rhs, lhs)] {
686 if let Some(root) = self.coherent_scale_root(scale)
687 && let Some(ExprNode::Unary {
688 op: UnaryOp::NormSqr,
689 input,
690 }) = self.graph.node(norm)
691 {
692 return Some((root, *input));
693 }
694 }
695 None
696 }
697
698 fn decompose_scaled_coherent(
699 &mut self,
700 scale: ExprId,
701 amplitude: ExprId,
702 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
703 let mut terms = self.decompose(amplitude)?;
704 for term in &mut terms {
705 term.coefficient = self.product(&[DecompositionNode::Source(scale), term.coefficient]);
706 }
707 self.pack_norm_sqr(&terms)
708 }
709
710 fn pack_norm_sqr(
711 &mut self,
712 terms: &[AnalyzedTerm],
713 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
714 self.coherent_groups += 1;
715 let packed_len = self.budget.packed_triangle_count(terms.len())?;
716 let two = self.constant(Complex64::new(2.0, 0.0));
717 let mut packed = Vec::with_capacity(packed_len);
718 for (row, left) in terms.iter().enumerate() {
719 for (column, right) in terms.iter().enumerate().skip(row) {
720 let right_coefficient = self.unary(UnaryOp::Conj, right.coefficient);
721 let right_basis = self.unary(UnaryOp::Conj, right.basis);
722 let mut coefficient = self.product(&[left.coefficient, right_coefficient]);
723 if column != row {
724 coefficient = self.product(&[two, coefficient]);
725 }
726 packed.push(AnalyzedTerm {
727 coefficient,
728 basis: self.product(&[left.basis, right_basis]),
729 });
730 }
731 }
732 Ok(packed)
733 }
734
735 fn decompose_projection(
736 &mut self,
737 op: UnaryOp,
738 input: ExprId,
739 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
740 let terms = self.decompose(input)?;
741 let mut result = Vec::with_capacity(terms.len() * 2);
742 let factor = if op == UnaryOp::Real {
743 Complex64::new(0.5, 0.0)
744 } else {
745 Complex64::new(0.0, -0.5)
746 };
747 let conjugate_factor = if op == UnaryOp::Real { factor } else { -factor };
748 let factor = self.constant(factor);
749 let conjugate_factor = self.constant(conjugate_factor);
750 for term in terms {
751 let coefficient = self.product(&[factor, term.coefficient]);
752 result.push(AnalyzedTerm {
753 coefficient,
754 basis: term.basis,
755 });
756 let conjugated_coefficient = self.unary(UnaryOp::Conj, term.coefficient);
757 let conjugated_basis = self.unary(UnaryOp::Conj, term.basis);
758 let coefficient = self.product(&[conjugate_factor, conjugated_coefficient]);
759 result.push(AnalyzedTerm {
760 coefficient,
761 basis: conjugated_basis,
762 });
763 }
764 self.ensure_budget(&result)?;
765 Ok(result)
766 }
767
768 fn multiply_terms(
769 &mut self,
770 lhs: &[AnalyzedTerm],
771 rhs: &[AnalyzedTerm],
772 ) -> Result<Vec<AnalyzedTerm>, NormalizationFallbackReason> {
773 let count = self.budget.product_count(lhs.len(), rhs.len())?;
774 let mut result = Vec::with_capacity(count);
775 for lhs in lhs {
776 for rhs in rhs {
777 let coefficient = self.product(&[lhs.coefficient, rhs.coefficient]);
778 let basis = self.product(&[lhs.basis, rhs.basis]);
779 result.push(AnalyzedTerm { coefficient, basis });
780 }
781 }
782 Ok(result)
783 }
784
785 fn ensure_budget(&self, terms: &[AnalyzedTerm]) -> Result<(), NormalizationFallbackReason> {
786 self.budget.ensure(terms.len())
787 }
788
789 fn unsupported(&self, node: ExprId, operation: &'static str) -> NormalizationFallbackReason {
790 NormalizationFallbackReason::UnsupportedMixedOperation { node, operation }
791 }
792
793 fn constant(&mut self, value: Complex64) -> DecompositionNode {
794 self.push(GraphOperation::Constant(value))
795 }
796
797 fn unary(&mut self, op: UnaryOp, input: DecompositionNode) -> DecompositionNode {
798 self.push(GraphOperation::Unary { op, input })
799 }
800
801 fn binary(
802 &mut self,
803 op: BinaryOp,
804 lhs: DecompositionNode,
805 rhs: DecompositionNode,
806 ) -> DecompositionNode {
807 self.push(GraphOperation::Binary { op, lhs, rhs })
808 }
809
810 fn product(&mut self, factors: &[DecompositionNode]) -> DecompositionNode {
811 match factors {
812 [] => self.one,
813 [only] => *only,
814 _ => self.push(GraphOperation::Product(factors.to_vec())),
815 }
816 }
817
818 fn sum(&mut self, roots: &[DecompositionNode]) -> Option<DecompositionNode> {
819 match roots {
820 [] => None,
821 [only] => Some(*only),
822 _ => Some(self.push(GraphOperation::Sum(roots.to_vec()))),
823 }
824 }
825
826 fn push(&mut self, operation: GraphOperation) -> DecompositionNode {
827 let id = DecompositionNode::Generated(self.operations.len());
828 self.operations.push(operation);
829 id
830 }
831}
832
833struct NormalizationGraphBuilder {
834 rebuild: ExprGraphRebuilder<ExprId>,
835}
836
837impl NormalizationGraphBuilder {
838 fn new(graph: &ExprGraph) -> Self {
839 let mut rebuild = ExprGraphRebuilder::with_capacity(graph.nodes().len());
840 for index in 0..graph.nodes().len() {
841 let old_id = ExprId::from_index(index);
842 let node = graph
843 .node(old_id)
844 .expect("normalization graph node exists")
845 .map_children(|child| {
846 rebuild
847 .remapped(&child)
848 .expect("validated expression graphs emit children before parents")
849 });
850 let metadata = graph
851 .metadata(old_id)
852 .expect("normalization graph metadata is complete")
853 .clone();
854 rebuild.emit(old_id, node, metadata);
855 }
856 Self { rebuild }
857 }
858
859 fn source(&self, id: ExprId) -> ExprId {
860 self.rebuild
861 .remapped(&id)
862 .expect("normalization source node was copied")
863 }
864
865 fn constant(&mut self, value: Complex64) -> ExprId {
866 self.emit(ExprNode::from_folded_const(value), ExprSourceKind::Const)
867 }
868
869 fn unary(&mut self, op: UnaryOp, input: ExprId) -> ExprId {
870 self.emit(ExprNode::Unary { op, input }, ExprSourceKind::Unary)
871 }
872
873 fn binary(&mut self, op: BinaryOp, lhs: ExprId, rhs: ExprId) -> ExprId {
874 self.emit(ExprNode::Binary { op, lhs, rhs }, ExprSourceKind::Binary)
875 }
876
877 fn product(&mut self, factors: &[ExprId]) -> ExprId {
878 match factors {
879 [] => self.constant(Complex64::new(1.0, 0.0)),
880 [only] => *only,
881 _ => self.emit(
882 ExprNode::NaryMul {
883 factors: factors.to_vec(),
884 },
885 ExprSourceKind::Binary,
886 ),
887 }
888 }
889
890 fn sum(&mut self, terms: &[ExprId]) -> ExprId {
891 match terms {
892 [] => self.constant(Complex64::new(0.0, 0.0)),
893 [only] => *only,
894 _ => self.emit(
895 ExprNode::NaryAdd {
896 terms: terms.to_vec(),
897 },
898 ExprSourceKind::Binary,
899 ),
900 }
901 }
902
903 fn emit(&mut self, node: ExprNode, source: ExprSourceKind) -> ExprId {
904 self.rebuild.emit_anonymous(node, ExprMetadata::new(source))
905 }
906
907 fn finish(self, root: ExprId) -> CompileResult<ExprGraph> {
908 Ok(self.rebuild.finish(root)?)
909 }
910}
911
912#[cfg(test)]
913mod tests {
914 use laddu_expr::{Expr, complex, event_scalar, parameter, polar_complex};
915
916 use super::*;
917
918 fn diagnostics(expression: &Expr) -> NormalizationDiagnostics {
919 CompiledModel::from_expr(expression)
920 .unwrap()
921 .normalization_diagnostics()
922 .clone()
923 }
924
925 fn decomposition(expression: &Expr, budget: usize) -> Decomposition {
926 let graph = expression.to_graph();
927 let facts = GraphFacts::analyze(&graph);
928 NormalizationAnalyzer::new(&graph, &facts, budget).analyze(graph.root())
929 }
930
931 #[test]
932 fn extracts_rectangular_and_polar_coherent_models() {
933 let basis = complex(event_scalar("x"), event_scalar("y"));
934 let rectangular = (complex(parameter!("re"), parameter!("im")) * basis.clone()).norm_sqr();
935 let polar = (polar_complex(parameter!("mag"), parameter!("phase")) * basis).norm_sqr();
936 for model in [&rectangular, &polar] {
937 let diagnostics = diagnostics(model);
938 assert_eq!(
939 diagnostics.strategy(),
940 NormalizationStrategy::Hermitian,
941 "{diagnostics:?}"
942 );
943 assert!(!diagnostics.has_residual());
944 assert!(diagnostics.basis_count() >= 1);
945 }
946 }
947
948 #[test]
949 fn scaled_coherent_intensity_remains_hermitian() {
950 let scale = Expr::from(parameter!("scale"));
951 let wave = complex(event_scalar("x"), event_scalar("y"))
952 + Expr::from(parameter!("mix")) * complex(event_scalar("y"), 0.5);
953 for intensity in [
954 scale.clone().powi(2) * wave.clone().norm_sqr(),
955 (scale.clone() * scale) * wave.norm_sqr(),
956 ] {
957 let diagnostics = diagnostics(&intensity);
958 assert_eq!(diagnostics.strategy(), NormalizationStrategy::Hermitian);
959 assert!(!diagnostics.has_residual());
960 assert_eq!(diagnostics.basis_count(), 3);
961 }
962 }
963
964 #[test]
965 fn decomposes_separable_and_nonseparable_additive_parts() {
966 let x = event_scalar("x");
967 let scale = Expr::from(parameter!("scale"));
968 let expression = scale.clone() * x.clone() + (scale * x).sin();
969 let diagnostics = diagnostics(&expression);
970 assert_eq!(diagnostics.strategy(), NormalizationStrategy::Hybrid);
971 assert!(diagnostics.has_residual());
972 }
973
974 #[test]
975 fn classifies_binary_operations_before_building_a_graph() {
976 let parameter = Expr::from(parameter!("scale"));
977 let event = event_scalar("x");
978 let mixed = parameter.clone() * event.clone();
979 let cases = [
980 ("add", parameter.clone() + event.clone(), 2, false),
981 ("sub", parameter.clone() - event.clone(), 2, false),
982 ("mul", mixed.clone() * mixed.clone(), 1, false),
983 (
984 "parameter divisor",
985 mixed.clone() / parameter.clone(),
986 1,
987 false,
988 ),
989 ("event divisor", mixed.clone() / event.clone(), 1, false),
990 (
991 "mixed divisor",
992 mixed.clone() / (parameter + event),
993 0,
994 true,
995 ),
996 ];
997
998 for (name, expression, expected_terms, has_failure) in cases {
999 let decomposition = decomposition(&expression, DEFAULT_EXPANSION_BUDGET);
1000 assert_eq!(decomposition.terms.len(), expected_terms, "{name}");
1001 assert_eq!(!decomposition.failures.is_empty(), has_failure, "{name}");
1002 }
1003 }
1004
1005 #[test]
1006 fn classifies_unary_operations_before_building_a_graph() {
1007 let mixed = Expr::from(parameter!("scale")) * event_scalar("x");
1008 let cases = [
1009 ("neg", -mixed.clone(), 1, 0),
1010 ("conj", mixed.clone().conj(), 1, 0),
1011 ("norm_sqr", mixed.clone().norm_sqr(), 1, 1),
1012 ("powi", mixed.clone().powi(2), 1, 0),
1013 ("real", mixed.clone().real(), 2, 0),
1014 ("imag", mixed.clone().imag(), 2, 0),
1015 ("sin", mixed.sin(), 0, 0),
1016 ];
1017
1018 for (name, expression, expected_terms, coherent_groups) in cases {
1019 let decomposition = decomposition(&expression, DEFAULT_EXPANSION_BUDGET);
1020 assert_eq!(decomposition.terms.len(), expected_terms, "{name}");
1021 assert_eq!(decomposition.coherent_groups, coherent_groups, "{name}");
1022 assert_eq!(
1023 decomposition.failures.is_empty(),
1024 expected_terms > 0,
1025 "{name}"
1026 );
1027 }
1028 }
1029
1030 #[test]
1031 fn expansion_budget_checks_exact_boundaries_and_overflow() {
1032 let budget = ExpansionBudget(6);
1033 assert_eq!(budget.product_count(2, 3).unwrap(), 6);
1034 assert_eq!(budget.packed_triangle_count(3).unwrap(), 6);
1035 assert_eq!(budget.ensure(6), Ok(()));
1036
1037 for result in [
1038 budget.product_count(2, 4),
1039 budget.packed_triangle_count(4),
1040 budget.product_count(usize::MAX, 2),
1041 budget.packed_triangle_count(usize::MAX),
1042 ] {
1043 assert_eq!(
1044 result,
1045 Err(NormalizationFallbackReason::ExpansionBudgetExceeded { budget: 6 })
1046 );
1047 }
1048 }
1049
1050 #[test]
1051 fn fallback_diagnostics_preserve_the_last_unsupported_root() {
1052 let mixed = Expr::from(parameter!("scale")) * event_scalar("x");
1053 let expression = mixed.clone().sin() + mixed.cos();
1054 let graph = expression.to_graph();
1055 let facts = GraphFacts::analyze(&graph);
1056 let roots = NormalizationAnalyzer::new(&graph, &facts, DEFAULT_EXPANSION_BUDGET)
1057 .additive_roots(graph.root());
1058 let decomposition = NormalizationAnalyzer::new(&graph, &facts, DEFAULT_EXPANSION_BUDGET)
1059 .analyze(graph.root());
1060
1061 assert_eq!(decomposition.failures.len(), 2);
1062 assert_eq!(
1063 decomposition.last_failure(),
1064 Some(&NormalizationFallbackReason::UnsupportedMixedOperation {
1065 node: roots[1],
1066 operation: "nonlinear unary operation",
1067 })
1068 );
1069 }
1070}