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#[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 pub fn strategy(&self) -> NormalizationStrategy {
35 self.strategy
36 }
37
38 pub fn compiler(&self) -> &NormalizationDiagnostics {
40 &self.compiler
41 }
42
43 pub fn retained_bytes(&self) -> usize {
45 self.retained_bytes
46 }
47
48 pub fn preparation_passes(&self) -> usize {
50 self.preparation_passes
51 }
52
53 pub fn cache_hit(&self) -> bool {
55 self.cache_hit
56 }
57
58 pub fn tag_projection_reused_parent(&self) -> bool {
60 self.tag_projection_reused_parent
61 }
62
63 #[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#[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 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 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 pub fn stats(&self) -> &PreparedDatasetStats {
313 &self.stats
314 }
315
316 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 pub fn resident_bytes(&self) -> usize {
325 self.statistics.resident_bytes()
326 }
327
328 pub fn value(&self, params: &ParamValues, execution: &Execution) -> RuntimeResult<f64> {
335 Ok(self.evaluate_composed(params, execution, false)?.value)
336 }
337
338 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(&litude.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}