Skip to main content

laddu_runtime/
normalization.rs

1use laddu_compile::{
2    CompiledModel, NormalizationDiagnostics, NormalizationStrategy, ReductionPlan,
3};
4use laddu_data::{
5    data::{CacheStorage, Dataset},
6    io::ReadPlan,
7};
8use laddu_expr::parameters::{ParamError, ParamLayout, ParamProjection, ParamValues};
9use laddu_memory::{MemoryFitRequest, MemoryFootprint};
10use num::complex::{Complex32, Complex64};
11use std::sync::{
12    Arc,
13    atomic::{AtomicBool, Ordering},
14};
15
16use crate::{
17    CpuBackend, CpuPlan, Execution, MemoryLease, NormalizationMode, PreparedDataset,
18    PreparedDatasetStats, PreparedModel, RuntimeError, RuntimeResult,
19};
20
21/// Runtime diagnostics for one compiler-native accepted normalization.
22#[derive(Clone, Debug, PartialEq, Eq)]
23pub struct PreparedNormalizationDiagnostics {
24    strategy: NormalizationStrategy,
25    compiler: NormalizationDiagnostics,
26    retained_bytes: usize,
27    preparation_passes: usize,
28    cache_hit: bool,
29    tag_projection_reused_parent: bool,
30}
31
32impl PreparedNormalizationDiagnostics {
33    /// Returns the runtime-selected normalization strategy.
34    pub fn strategy(&self) -> NormalizationStrategy {
35        self.strategy
36    }
37
38    /// Returns compiler analysis diagnostics.
39    pub fn compiler(&self) -> &NormalizationDiagnostics {
40        &self.compiler
41    }
42
43    /// Returns retained sufficient-statistic bytes on this rank.
44    pub fn retained_bytes(&self) -> usize {
45        self.retained_bytes
46    }
47
48    /// Returns the number of accepted-source passes used during preparation.
49    pub fn preparation_passes(&self) -> usize {
50        self.preparation_passes
51    }
52
53    /// Returns whether an execution-scoped prepared artifact was reused.
54    pub fn cache_hit(&self) -> bool {
55        self.cache_hit
56    }
57
58    /// Returns whether a tag projection reused parent statistics.
59    pub fn tag_projection_reused_parent(&self) -> bool {
60        self.tag_projection_reused_parent
61    }
62
63    /// Constructs diagnostics for an ordinary prepared event reduction.
64    #[doc(hidden)]
65    pub fn general(compiler: NormalizationDiagnostics) -> Self {
66        Self {
67            strategy: NormalizationStrategy::General,
68            compiler,
69            retained_bytes: 0,
70            preparation_passes: 1,
71            cache_hit: false,
72            tag_projection_reused_parent: false,
73        }
74    }
75}
76
77#[derive(Clone, Debug)]
78struct GeneralResidual {
79    plan: PreparedModel,
80    dataset: PreparedDataset,
81    parameters: ParamProjection,
82}
83
84#[derive(Debug)]
85enum StoredStatistics {
86    F32(Vec<Complex32>),
87    F64(Vec<Complex64>),
88}
89
90impl StoredStatistics {
91    fn from_f64(values: Vec<Complex64>, precision: crate::Precision) -> Self {
92        if precision == crate::Precision::F32 {
93            Self::F32(
94                values
95                    .into_iter()
96                    .map(|value| Complex32::new(value.re as f32, value.im as f32))
97                    .collect(),
98            )
99        } else {
100            Self::F64(values)
101        }
102    }
103
104    fn resident_bytes(&self) -> usize {
105        match self {
106            Self::F32(values) => values.capacity() * std::mem::size_of::<Complex32>(),
107            Self::F64(values) => values.capacity() * std::mem::size_of::<Complex64>(),
108        }
109    }
110
111    fn evaluator_values(&self) -> Vec<Complex64> {
112        match self {
113            Self::F32(values) => values
114                .iter()
115                .map(|value| Complex64::new(value.re as f64, value.im as f64))
116                .collect(),
117            Self::F64(values) => values.clone(),
118        }
119    }
120}
121
122fn normalization_projection(
123    child: &ParamLayout,
124    parent: &ParamLayout,
125) -> RuntimeResult<ParamProjection> {
126    child.projection_from(parent).map_err(|error| match error {
127        ParamError::UnknownName(name) => RuntimeError::Data(format!(
128            "normalization parameter `{name}` is absent from the source model"
129        )),
130        ParamError::ParameterConflict { name, .. } => RuntimeError::Data(format!(
131            "normalization parameter `{name}` is unexpectedly fixed in the source model"
132        )),
133        error => RuntimeError::Parameter(error.to_string()),
134    })
135}
136
137fn project_normalization(
138    projection: &ParamProjection,
139    params: &ParamValues,
140) -> RuntimeResult<ParamValues> {
141    projection.project(params).map_err(|error| match error {
142        ParamError::UnknownName(name) => RuntimeError::Data(format!(
143            "normalization parameter `{name}` is absent from supplied values"
144        )),
145        error => RuntimeError::Parameter(error.to_string()),
146    })
147}
148
149/// Prepared sufficient statistics and their parameter-only contraction.
150#[derive(Debug)]
151pub struct PreparedNormalization {
152    evaluator: CpuPlan,
153    evaluator_parameters: ParamProjection,
154    statistics: StoredStatistics,
155    residual: Option<GeneralResidual>,
156    verification: Option<GeneralResidual>,
157    stats: PreparedDatasetStats,
158    diagnostics: PreparedNormalizationDiagnostics,
159    cache_reused: AtomicBool,
160    _memory_lease: MemoryLease,
161}
162
163#[derive(Debug)]
164struct NormalizationEvaluation {
165    value: f64,
166    gradient: Option<Vec<f64>>,
167}
168
169impl PreparedNormalization {
170    /// Prepares compiler-native normalization when selected by execution policy.
171    ///
172    /// # Errors
173    ///
174    /// Returns a runtime error when basis compilation, dataset traversal,
175    /// memory reservation, or backend preparation fails.
176    pub fn prepare(
177        model: &CompiledModel,
178        general_plan: &PreparedModel,
179        dataset: &Dataset,
180        execution: &Execution,
181    ) -> RuntimeResult<Option<Arc<Self>>> {
182        if execution.normalization_mode() == NormalizationMode::General
183            || model.normalization_diagnostics().strategy() == NormalizationStrategy::General
184            || (execution.normalization_mode() == NormalizationMode::Auto
185                && !model.normalization_plan().proven_nonnegative())
186        {
187            return Ok(None);
188        }
189
190        let key = (
191            model.optimized_digest(),
192            dataset.identity(),
193            execution.normalization_mode(),
194        );
195        let mut cache = execution
196            .normalization_cache()
197            .lock()
198            .unwrap_or_else(|error| error.into_inner());
199        cache.retain(|_, prepared| prepared.strong_count() > 0);
200        if let Some(prepared) = cache.get(&key).and_then(std::sync::Weak::upgrade) {
201            prepared.cache_reused.store(true, Ordering::Relaxed);
202            return Ok(Some(prepared));
203        }
204        let Some(prepared) = Self::prepare_uncached(model, general_plan, dataset, execution)?
205        else {
206            return Ok(None);
207        };
208        let prepared = Arc::new(prepared);
209        cache.insert(key, Arc::downgrade(&prepared));
210        Ok(Some(prepared))
211    }
212
213    fn prepare_uncached(
214        model: &CompiledModel,
215        general_plan: &PreparedModel,
216        dataset: &Dataset,
217        execution: &Execution,
218    ) -> RuntimeResult<Option<Self>> {
219        let basis_models = model
220            .normalization_plan()
221            .basis_models()
222            .map_err(|error| RuntimeError::Data(error.to_string()))?;
223        let statistic_bytes = if execution.precision() == crate::Precision::F32 {
224            std::mem::size_of::<Complex32>()
225        } else {
226            std::mem::size_of::<Complex64>()
227        };
228        let retained_bytes = basis_models.len().saturating_mul(statistic_bytes);
229        let memory_lease = match execution
230            .host_memory()
231            .reserve(u64::try_from(retained_bytes).unwrap_or(u64::MAX))
232        {
233            Ok(lease) => lease,
234            Err(_) if execution.normalization_mode() == NormalizationMode::Auto => return Ok(None),
235            Err(error) => return Err(error.into()),
236        };
237        let basis_plans = basis_models
238            .iter()
239            .map(|basis| {
240                CpuBackend.prepare_shared_with_autodiff_mode(basis, execution.autodiff_mode())
241            })
242            .collect::<RuntimeResult<Vec<_>>>()?;
243        let basis_params = basis_models
244            .iter()
245            .map(|basis| basis.params().default_values())
246            .collect::<Vec<_>>();
247        let (statistics, stats) =
248            accumulate_statistics(&basis_plans, &basis_params, dataset, execution)?;
249        let statistics = StoredStatistics::from_f64(statistics, execution.precision());
250        let evaluator_statistics = statistics.evaluator_values();
251        let evaluator_model = model
252            .normalization_plan()
253            .evaluator_model(&evaluator_statistics)
254            .map_err(|error| RuntimeError::Data(error.to_string()))?;
255        // Sufficient statistics may be prepared for an f32 accelerator, but
256        // their tiny parameter-only contraction stays on the CPU in f64 so
257        // value/gradient evaluation remains available and numerically stable.
258        let evaluator = CpuBackend
259            .prepare_with_autodiff_mode(&evaluator_model, execution.autodiff_mode())
260            .map_err(|error| RuntimeError::Data(error.to_string()))?;
261        let evaluator_parameters =
262            normalization_projection(evaluator_model.params(), model.params())?;
263
264        let residual_model = model
265            .normalization_plan()
266            .residual_model()
267            .map_err(|error| RuntimeError::Data(error.to_string()))?;
268        let residual = if let Some(residual_model) = residual_model {
269            let parameters = normalization_projection(residual_model.params(), model.params())?;
270            let plan = PreparedModel::prepare(&residual_model, execution)?;
271            let dataset = plan.prepare_dataset(execution, dataset)?;
272            Some(GeneralResidual {
273                plan,
274                dataset,
275                parameters,
276            })
277        } else {
278            None
279        };
280        let verification = if execution.normalization_mode() == NormalizationMode::Verify {
281            Some(GeneralResidual {
282                plan: general_plan.clone(),
283                dataset: general_plan.prepare_dataset(execution, dataset)?,
284                parameters: normalization_projection(model.params(), model.params())?,
285            })
286        } else {
287            None
288        };
289        let preparation_passes =
290            1 + usize::from(residual.is_some()) + usize::from(verification.is_some());
291        Ok(Some(Self {
292            evaluator,
293            evaluator_parameters,
294            statistics,
295            residual,
296            verification,
297            stats,
298            diagnostics: PreparedNormalizationDiagnostics {
299                strategy: model.normalization_diagnostics().strategy(),
300                compiler: model.normalization_diagnostics().clone(),
301                retained_bytes,
302                preparation_passes,
303                cache_hit: false,
304                tag_projection_reused_parent: false,
305            },
306            cache_reused: AtomicBool::new(false),
307            _memory_lease: memory_lease,
308        }))
309    }
310
311    /// Returns accepted-dataset statistics collected during preparation.
312    pub fn stats(&self) -> &PreparedDatasetStats {
313        &self.stats
314    }
315
316    /// Returns normalization preparation diagnostics.
317    pub fn diagnostics(&self) -> PreparedNormalizationDiagnostics {
318        let mut diagnostics = self.diagnostics.clone();
319        diagnostics.cache_hit = self.cache_reused.load(Ordering::Relaxed);
320        diagnostics
321    }
322
323    /// Returns retained sufficient-statistic storage in bytes.
324    pub fn resident_bytes(&self) -> usize {
325        self.statistics.resident_bytes()
326    }
327
328    /// Evaluates the accepted normalization without constructing a gradient.
329    ///
330    /// # Errors
331    ///
332    /// Returns a runtime error for incompatible parameters, residual backend
333    /// failures, or a verification mismatch.
334    pub fn value(&self, params: &ParamValues, execution: &Execution) -> RuntimeResult<f64> {
335        Ok(self.evaluate_composed(params, execution, false)?.value)
336    }
337
338    /// Evaluates the accepted normalization and its local free-parameter gradient.
339    ///
340    /// # Errors
341    ///
342    /// Returns a runtime error for incompatible parameters, autodiff/backend
343    /// failures, or a verification mismatch.
344    pub fn value_gradient(
345        &self,
346        params: &ParamValues,
347        execution: &Execution,
348    ) -> RuntimeResult<(f64, Vec<f64>)> {
349        let evaluation = self.evaluate_composed(params, execution, true)?;
350        Ok((
351            evaluation.value,
352            evaluation.gradient.ok_or_else(|| {
353                RuntimeError::Data("normalization gradient composition produced no gradient".into())
354            })?,
355        ))
356    }
357
358    fn evaluate_composed(
359        &self,
360        params: &ParamValues,
361        execution: &Execution,
362        with_gradient: bool,
363    ) -> RuntimeResult<NormalizationEvaluation> {
364        let evaluator_params = project_normalization(&self.evaluator_parameters, params)?;
365        let (mut value, mut gradient) = if with_gradient {
366            let evaluation = self.evaluator.evaluate_with_gradient(&evaluator_params)?;
367            let mut gradient = vec![0.0; params.layout().n_free()];
368            let evaluator_gradient = evaluation
369                .gradient()
370                .iter()
371                .map(|value| value.re)
372                .collect::<Vec<_>>();
373            self.evaluator_parameters
374                .scatter_add(&evaluator_gradient, &mut gradient)
375                .map_err(|_| incompatible_gradient_layout())?;
376            (evaluation.value().re, Some(gradient))
377        } else {
378            (self.evaluator.evaluate(&evaluator_params)?.re, None)
379        };
380        if let Some(residual) = &self.residual {
381            let residual_params = project_normalization(&residual.parameters, params)?;
382            if let Some(gradient) = &mut gradient {
383                let residual_evaluation = residual.plan.reduce_with_gradient(
384                    execution,
385                    &residual_params,
386                    &residual.dataset,
387                    ReductionPlan::weighted_real(),
388                )?;
389                value += residual_evaluation.value();
390                residual
391                    .parameters
392                    .scatter_add(residual_evaluation.gradient(), gradient)
393                    .map_err(|_| incompatible_gradient_layout())?;
394            } else {
395                value += residual.plan.reduce(
396                    execution,
397                    &residual_params,
398                    &residual.dataset,
399                    ReductionPlan::weighted_real(),
400                )?;
401            }
402        }
403        if let Some(general) = &self.verification {
404            let general_params = project_normalization(&general.parameters, params)?;
405            if let Some(gradient) = &gradient {
406                let expected = general.plan.reduce_with_gradient(
407                    execution,
408                    &general_params,
409                    &general.dataset,
410                    ReductionPlan::weighted_real(),
411                )?;
412                verify_close("normalization value", value, expected.value(), execution)?;
413                for (index, (actual, expected)) in
414                    gradient.iter().zip(expected.gradient()).enumerate()
415                {
416                    verify_close(
417                        &format!("normalization gradient[{index}]"),
418                        *actual,
419                        *expected,
420                        execution,
421                    )?;
422                }
423            } else {
424                let expected = general.plan.reduce(
425                    execution,
426                    &general_params,
427                    &general.dataset,
428                    ReductionPlan::weighted_real(),
429                )?;
430                verify_close("normalization value", value, expected, execution)?;
431            }
432        }
433        Ok(NormalizationEvaluation { value, gradient })
434    }
435}
436
437fn incompatible_gradient_layout() -> RuntimeError {
438    RuntimeError::Data("normalization gradient has an incompatible parameter layout".into())
439}
440
441fn accumulate_statistics(
442    plans: &[Arc<CpuPlan>],
443    params: &[ParamValues],
444    dataset: &Dataset,
445    execution: &Execution,
446) -> RuntimeResult<(Vec<Complex64>, PreparedDatasetStats)> {
447    let mut sums = vec![Complex64::ZERO; plans.len()];
448    let mut corrections = vec![Complex64::ZERO; plans.len()];
449    let mut read_plan: ReadPlan = execution.read_plan(dataset.read_plan());
450    let local_limit = dataset
451        .num_events()
452        .map_err(|error| RuntimeError::Data(error.to_string()))?
453        .and_then(|events| usize::try_from(events).ok())
454        .unwrap_or(usize::MAX);
455    let statistic_bytes = plans.len().saturating_mul(std::mem::size_of::<Complex64>());
456    let decision = MemoryFitRequest {
457        label: "normalization statistics".into(),
458        footprint: MemoryFootprint::from_usize(statistic_bytes, statistic_bytes),
459        available_bytes: execution.host_memory().remaining(),
460        event_limit: local_limit,
461        strategy: "single-pass sufficient statistics".into(),
462    }
463    .evaluate()?;
464    read_plan.chunk_size = Some(
465        read_plan
466            .chunk_size
467            .map_or(decision.chunk_events, |manual| {
468                manual.min(decision.chunk_events)
469            })
470            .max(1),
471    );
472    execution.record_memory_decision(decision);
473    let local = (|| {
474        let mut events = 0usize;
475        let mut batches = 0usize;
476        let mut weight_sum = 0.0;
477        let mut weight_correction = 0.0;
478        for batch in dataset
479            .stream_with_plan(read_plan)
480            .map_err(|error| RuntimeError::Data(error.to_string()))?
481        {
482            let batch = batch.map_err(|error| RuntimeError::Data(error.to_string()))?;
483            events += batch.len();
484            batches += 1;
485            for row in 0..batch.len() {
486                let weight = batch.weights_at(row);
487                let corrected = weight - weight_correction;
488                let next = weight_sum + corrected;
489                weight_correction = (next - weight_sum) - corrected;
490                weight_sum = next;
491            }
492            for (index, (plan, params)) in plans.iter().zip(params).enumerate() {
493                for (row, value) in plan.evaluate_batch(params, &batch)?.into_iter().enumerate() {
494                    let value = value * batch.weights_at(row);
495                    let corrected = value - corrections[index];
496                    let next = sums[index] + corrected;
497                    corrections[index] = (next - sums[index]) - corrected;
498                    sums[index] = next;
499                }
500            }
501        }
502        Ok::<_, RuntimeError>((events, batches, weight_sum))
503    })();
504    if !execution.all_succeeded(local.is_ok()) {
505        return local.and(Err(RuntimeError::DistributedPeerFailure));
506    }
507    let (events, batches, weight_sum) = local?;
508    for sum in &mut sums {
509        sum.re = execution.sum_f64(sum.re);
510        sum.im = execution.sum_f64(sum.im);
511    }
512    let stats = PreparedDatasetStats::new(
513        events,
514        execution.sum_usize(events),
515        batches,
516        execution.sum_f64(weight_sum),
517        sums.len() * std::mem::size_of::<Complex64>(),
518        CacheStorage::Resident,
519    );
520    Ok((sums, stats))
521}
522
523fn verify_close(
524    label: &str,
525    actual: f64,
526    expected: f64,
527    execution: &Execution,
528) -> RuntimeResult<()> {
529    let tolerance = match execution.precision() {
530        crate::Precision::F32 => 5.0e-4,
531        crate::Precision::Auto | crate::Precision::F64 => 1.0e-10,
532    } * expected.abs().max(1.0);
533    if (actual - expected).abs() <= tolerance {
534        Ok(())
535    } else {
536        Err(RuntimeError::Data(format!(
537            "{label} verification failed: compiler-native={actual}, general={expected}, tolerance={tolerance}"
538        )))
539    }
540}
541
542#[cfg(test)]
543mod tests {
544    use std::sync::Arc;
545
546    use laddu_data::{
547        data::{EventBatch, OwnedEvent},
548        schema::Schema,
549    };
550    use laddu_expr::{complex, event_scalar, parameter};
551
552    use super::*;
553    use crate::{ExecutionOptions, MemoryBudget, MemoryPlan};
554
555    #[test]
556    fn auto_falls_back_when_statistics_exceed_host_budget() {
557        let amplitude = complex(event_scalar("x"), 0.5)
558            + parameter!("mix", initial: 0.3) * complex(event_scalar("x").powi(2), 0.25);
559        let model = CompiledModel::from_expr(&amplitude.norm_sqr()).unwrap();
560        assert_eq!(
561            model.normalization_diagnostics().strategy(),
562            NormalizationStrategy::Hermitian
563        );
564        let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
565        let batch = EventBatch::from_events(schema, [OwnedEvent::weighted(vec![], vec![0.5], 1.0)])
566            .unwrap();
567        let dataset = Dataset::from_batches(vec![batch]).unwrap();
568        let execution = Execution::local(ExecutionOptions {
569            normalization: NormalizationMode::Auto,
570            memory: MemoryPlan::host_device(MemoryBudget::Bytes(1), MemoryBudget::Auto),
571            ..ExecutionOptions::default()
572        })
573        .unwrap();
574        let plan = PreparedModel::prepare(&model, &execution).unwrap();
575        assert!(
576            PreparedNormalization::prepare(&model, &plan, &dataset, &execution)
577                .unwrap()
578                .is_none()
579        );
580    }
581}