Skip to main content

laddu_compile/
normalization.rs

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/// Compiler-selected family of accepted-normalization implementation.
14#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
15pub enum NormalizationStrategy {
16    /// Coherent groups represented by packed Hermitian statistics.
17    Hermitian,
18    /// A general exact sum of parameter coefficients times event bases.
19    LinearStatistics,
20    /// Exact sufficient statistics plus a nonseparable additive residual.
21    Hybrid,
22    /// Ordinary event reduction.
23    General,
24}
25
26/// Stable reason that compiler-native normalization was not selected.
27#[derive(Clone, Debug, PartialEq, Eq)]
28pub enum NormalizationFallbackReason {
29    /// The model root is not a real scalar.
30    NonScalarIntensity,
31    /// An operation mixed parameter and event dependence without an exact rule.
32    UnsupportedMixedOperation {
33        /// Offending optimized graph node.
34        node: ExprId,
35        /// Operation category.
36        operation: &'static str,
37    },
38    /// Exact symbolic distribution exceeded the compiler expansion budget.
39    ExpansionBudgetExceeded {
40        /// Maximum permitted number of sufficient-statistic terms.
41        budget: usize,
42    },
43}
44
45/// Stable summary of compiler-native normalization analysis.
46#[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    /// Returns the compiler-selected candidate family.
57    pub fn strategy(&self) -> NormalizationStrategy {
58        self.strategy
59    }
60
61    /// Returns the number of event-basis statistics in the candidate.
62    pub fn basis_count(&self) -> usize {
63        self.basis_count
64    }
65
66    /// Returns the number of recognized coherent squared-norm groups.
67    pub fn coherent_group_count(&self) -> usize {
68        self.coherent_group_count
69    }
70
71    /// Returns whether evaluation retains a general additive residual.
72    pub fn has_residual(&self) -> bool {
73        self.has_residual
74    }
75
76    /// Returns the structured general-path reason, when present.
77    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/// Exact compiler artifact consumed by execution backends.
89///
90/// This is a workspace-internal contract shared with `laddu-runtime`. It is
91/// public only because Rust crate boundaries do not provide workspace-scoped
92/// visibility. Downstream users should use [`NormalizationDiagnostics`]
93/// instead of depending on this artifact's representation or methods.
94#[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    /// Returns stable compiler diagnostics.
199    pub fn diagnostics(&self) -> &NormalizationDiagnostics {
200        &self.diagnostics
201    }
202
203    /// Returns whether the source graph proves every event value nonnegative.
204    pub fn proven_nonnegative(&self) -> bool {
205        self.proven_nonnegative
206    }
207
208    /// Builds one event-only compiled model per statistic.
209    ///
210    /// # Errors
211    ///
212    /// Returns a compiler error if an extracted basis graph cannot be lowered.
213    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    /// Builds the parameter-only scalar contraction for accumulated statistics.
226    ///
227    /// # Errors
228    ///
229    /// Returns a compiler error if the statistic count is incompatible with
230    /// this plan or the coefficient graph cannot be lowered.
231    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    /// Builds the nonseparable residual model, when present.
253    ///
254    /// # Errors
255    ///
256    /// Returns a compiler error if the residual graph cannot be lowered.
257    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    /// Recognize a real, event-independent square without taking a square root.
664    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}