1use std::path::Path;
2
3use serde::{Deserialize, Serialize};
4
5use crate::{
6 Result,
7 callbacks::Callback,
8 error::Error,
9 experiment::{Illumination, Optics},
10 illumination_calibration::{
11 CalibrationConditioning, CalibrationConvergenceReason, CalibrationLossHistoryEntry,
12 CalibrationParameterHistoryEntry, IlluminationCalibration, IlluminationCalibrationState,
13 PlanarArrayParameterValues,
14 },
15 measurements::MeasurementRead,
16 model::{ImagePlaneModel, ReconstructionShape},
17 reconstruction::{
18 AlgorithmMetricRecord, Batch, ReconstructionCheckpoint, ReconstructionProblem,
19 ReconstructionResult, ReconstructionState, RunOptions, Runner,
20 },
21};
22
23use super::{AlgorithmIterationMetrics, ReconstructionAlgorithm, StepOutput, StepSummary};
24
25#[derive(Clone, Debug, Default)]
27pub struct JointIterationMetrics {
28 data_loss: Option<f64>,
29 regularization_loss: Option<f64>,
30 total_loss: Option<f64>,
31 geometry_recompilations: usize,
32 multiplicative_updates: usize,
33 rejected_steps: usize,
34 object_metrics: Vec<(String, String, f64)>,
35}
36
37impl AlgorithmIterationMetrics for JointIterationMetrics {
38 fn merge(&mut self, other: Self) {
39 self.data_loss = other.data_loss.or(self.data_loss);
40 self.regularization_loss = other.regularization_loss.or(self.regularization_loss);
41 self.total_loss = other.total_loss.or(self.total_loss);
42 self.geometry_recompilations += other.geometry_recompilations;
43 self.multiplicative_updates += other.multiplicative_updates;
44 self.rejected_steps += other.rejected_steps;
45 self.object_metrics.extend(other.object_metrics);
46 }
47
48 fn append_records(&self, iteration: usize, output: &mut Vec<AlgorithmMetricRecord>) {
49 for (metric, value) in [
50 ("data_loss", self.data_loss),
51 ("regularization_loss", self.regularization_loss),
52 ("total_loss", self.total_loss),
53 ] {
54 if let Some(value) = value {
55 output.push(AlgorithmMetricRecord {
56 iteration,
57 namespace: "physical_illumination".into(),
58 metric: metric.into(),
59 value,
60 });
61 }
62 }
63 for (metric, value) in [
64 (
65 "geometry_recompilations",
66 self.geometry_recompilations as f64,
67 ),
68 ("multiplicative_updates", self.multiplicative_updates as f64),
69 ("rejected_steps", self.rejected_steps as f64),
70 ] {
71 output.push(AlgorithmMetricRecord {
72 iteration,
73 namespace: "physical_illumination".into(),
74 metric: metric.into(),
75 value,
76 });
77 }
78 output.extend(
79 self.object_metrics
80 .iter()
81 .map(|(namespace, metric, value)| AlgorithmMetricRecord {
82 iteration,
83 namespace: format!("object_update.{namespace}"),
84 metric: metric.clone(),
85 value: *value,
86 }),
87 );
88 }
89}
90
91#[derive(Clone, Debug)]
108pub struct JointReconstruction<A> {
109 pub object_algorithm: A,
111 pub optics: Optics,
113 pub initial_illumination: Illumination,
115 pub illumination_calibration: IlluminationCalibration,
117 pub outer_iterations: usize,
119 pub object_iterations_per_outer: usize,
121 pub illumination_steps_per_outer: usize,
123}
124
125impl<A> JointReconstruction<A> {
126 pub fn new(
128 object_algorithm: A,
129 optics: Optics,
130 initial_illumination: Illumination,
131 illumination_calibration: IlluminationCalibration,
132 outer_iterations: usize,
133 ) -> Self {
134 Self {
135 object_algorithm,
136 optics,
137 initial_illumination,
138 illumination_calibration,
139 outer_iterations,
140 object_iterations_per_outer: 1,
141 illumination_steps_per_outer: 1,
142 }
143 }
144
145 pub fn object_iterations_per_outer(mut self, iterations: usize) -> Self {
147 self.object_iterations_per_outer = iterations;
148 self
149 }
150
151 pub fn illumination_steps_per_outer(mut self, steps: usize) -> Self {
153 self.illumination_steps_per_outer = steps;
154 self
155 }
156}
157
158impl<A: ReconstructionAlgorithm> JointReconstruction<A> {
159 pub fn run<M: MeasurementRead>(
161 self,
162 problem: &ReconstructionProblem<M>,
163 ) -> Result<JointReconstructionResult> {
164 let options = RunOptions {
165 max_iterations: self.outer_iterations,
166 batch_size: usize::MAX,
167 ..RunOptions::default()
168 };
169 JointReconstructionResult::from_reconstruction(Runner::new(self, options).run(problem)?)
170 }
171
172 pub fn run_with_callbacks<M: MeasurementRead>(
178 self,
179 problem: &ReconstructionProblem<M>,
180 callbacks: Vec<Box<dyn Callback>>,
181 ) -> Result<JointReconstructionResult> {
182 let options = RunOptions {
183 max_iterations: self.outer_iterations,
184 batch_size: usize::MAX,
185 ..RunOptions::default()
186 };
187 JointReconstructionResult::from_reconstruction(
188 Runner::new(self, options)
189 .with_callbacks(callbacks)
190 .run(problem)?,
191 )
192 }
193
194 pub fn run_from_checkpoint<M: MeasurementRead>(
196 self,
197 problem: &ReconstructionProblem<M>,
198 checkpoint: ReconstructionCheckpoint,
199 ) -> Result<JointReconstructionResult> {
200 let options = RunOptions {
201 max_iterations: self.outer_iterations,
202 batch_size: usize::MAX,
203 ..RunOptions::default()
204 };
205 JointReconstructionResult::from_reconstruction(
206 Runner::new(self, options)
207 .resume_from(checkpoint)
208 .run(problem)?,
209 )
210 }
211}
212
213impl<A: ReconstructionAlgorithm> ReconstructionAlgorithm for JointReconstruction<A> {
214 type IterationMetrics = JointIterationMetrics;
215
216 fn validate(&self) -> Result<()> {
217 self.object_algorithm.validate()?;
218 if !self.object_algorithm.supports_joint_reconstruction() {
219 return Err(Error::InvalidParameter {
220 name: "object_algorithm",
221 reason: "the selected algorithm carries state that cannot be transported across physical model recompilation"
222 .into(),
223 });
224 }
225 self.illumination_calibration
226 .validate_for(&self.initial_illumination)?;
227 self.optics.validate()?;
228 if self.outer_iterations == 0
229 || self.object_iterations_per_outer == 0
230 || self.illumination_steps_per_outer == 0
231 {
232 return Err(Error::InvalidParameter {
233 name: "joint iteration counts",
234 reason: "outer, object, and illumination iteration counts must be positive".into(),
235 });
236 }
237 Ok(())
238 }
239
240 fn validate_problem<M: MeasurementRead>(
241 &self,
242 problem: &ReconstructionProblem<M>,
243 ) -> Result<()> {
244 self.validate()?;
245 self.object_algorithm.validate_problem(problem)?;
246 let expected = ImagePlaneModel::from_experiment(
247 &self.optics,
248 &self.initial_illumination,
249 problem.model.image_shape(),
250 ReconstructionShape::Exact(problem.model.reconstruction_shape()),
251 )?;
252 if expected.source_count() != problem.model.source_count()
253 || expected.frame_count() != problem.model.frame_count()
254 || expected.k_vectors() != problem.model.k_vectors()
255 {
256 return Err(Error::InvalidModel(
257 "joint reconstruction problem model was not compiled from the supplied initial illumination"
258 .into(),
259 ));
260 }
261 Ok(())
262 }
263
264 fn initialize<M: MeasurementRead>(
265 &self,
266 problem: &ReconstructionProblem<M>,
267 ) -> Result<ReconstructionState> {
268 let mut state = self.object_algorithm.initialize(problem)?;
269 let calibration_state = self
270 .illumination_calibration
271 .initialize(&self.initial_illumination, &self.optics)?;
272 let mut calibrated_model = problem.model.clone();
273 self.illumination_calibration.synchronize_model(
274 &self.optics,
275 &mut calibrated_model,
276 &self.initial_illumination,
277 &calibration_state,
278 )?;
279 state.frame_gains = calibrated_model.frame_gains().map(<[f64]>::to_vec);
280 state.physical_illumination_calibration = Some(calibration_state);
281 state.calibrated_model = Some(calibrated_model);
282 Ok(state)
283 }
284
285 fn initialize_with_backend<M: MeasurementRead>(
286 &self,
287 problem: &ReconstructionProblem<M>,
288 backend: std::sync::Arc<dyn crate::backend::Backend>,
289 ) -> Result<ReconstructionState> {
290 let mut state = self
291 .object_algorithm
292 .initialize_with_backend(problem, backend)?;
293 let calibration_state = self
294 .illumination_calibration
295 .initialize(&self.initial_illumination, &self.optics)?;
296 let mut calibrated_model = problem.model.clone();
297 self.illumination_calibration.synchronize_model(
298 &self.optics,
299 &mut calibrated_model,
300 &self.initial_illumination,
301 &calibration_state,
302 )?;
303 state.frame_gains = calibrated_model.frame_gains().map(<[f64]>::to_vec);
304 state.physical_illumination_calibration = Some(calibration_state);
305 state.calibrated_model = Some(calibrated_model);
306 Ok(state)
307 }
308
309 fn step<M: MeasurementRead>(
310 &mut self,
311 problem: &ReconstructionProblem<M>,
312 state: &mut ReconstructionState,
313 batch: &Batch,
314 iteration: usize,
315 ) -> Result<StepOutput<Self::IterationMetrics>> {
316 let mut effective_model = state
317 .calibrated_model
318 .clone()
319 .ok_or_else(|| Error::InvalidModel("joint calibrated model is missing".into()))?;
320 let local_problem = ReconstructionProblem {
321 measurements: &problem.measurements,
322 model: effective_model.clone(),
323 name: problem.name.clone(),
324 };
325 let mut summary = StepSummary::default();
326 let mut object_metrics = Vec::new();
327 for _ in 0..self.object_iterations_per_outer {
328 let output = self
329 .object_algorithm
330 .step(&local_problem, state, batch, iteration)?;
331 self.object_algorithm
332 .canonicalize_state(&local_problem, state)?;
333 summary.merge(output.summary);
334 let mut records = Vec::new();
335 output.metrics.append_records(iteration + 1, &mut records);
336 object_metrics.extend(
337 records
338 .into_iter()
339 .map(|record| (record.namespace, record.metric, record.value)),
340 );
341 }
342
343 let geometry_before = state
344 .physical_illumination_calibration
345 .as_ref()
346 .ok_or_else(|| {
347 Error::InvalidModel("joint physical calibration state is missing".into())
348 })?
349 .geometry_recompilations;
350 let multiplicative_before = state
351 .physical_illumination_calibration
352 .as_ref()
353 .expect("checked above")
354 .multiplicative_updates;
355 let rejected_before = state
356 .physical_illumination_calibration
357 .as_ref()
358 .expect("checked above")
359 .rejected_steps;
360 let mut objective = None;
361 for _ in 0..self.illumination_steps_per_outer {
362 objective = Some(
363 self.illumination_calibration.optimize(
364 &problem.measurements,
365 &self.optics,
366 state.object_spectrum.ndarray_view(),
367 &state.pupil,
368 &mut effective_model,
369 state
370 .physical_illumination_calibration
371 .as_mut()
372 .expect("checked above"),
373 iteration + 1,
374 )?,
375 );
376 }
377 state.frame_gains = effective_model.frame_gains().map(<[f64]>::to_vec);
378 state.calibrated_model = Some(effective_model);
379 let calibration = state
380 .physical_illumination_calibration
381 .as_ref()
382 .expect("checked above");
383 let objective = objective.expect("positive illumination phase count was validated");
384 Ok(StepOutput {
385 summary,
386 metrics: JointIterationMetrics {
387 data_loss: Some(objective.data_loss),
388 regularization_loss: Some(objective.regularization_loss),
389 total_loss: Some(objective.total_loss),
390 geometry_recompilations: calibration.geometry_recompilations - geometry_before,
391 multiplicative_updates: calibration.multiplicative_updates - multiplicative_before,
392 rejected_steps: calibration.rejected_steps - rejected_before,
393 object_metrics,
394 },
395 })
396 }
397
398 fn canonicalize_state<M: MeasurementRead>(
399 &self,
400 problem: &ReconstructionProblem<M>,
401 state: &mut ReconstructionState,
402 ) -> Result<()> {
403 let model = state
404 .calibrated_model
405 .clone()
406 .unwrap_or_else(|| problem.model.clone());
407 let local_problem = ReconstructionProblem {
408 measurements: &problem.measurements,
409 model,
410 name: problem.name.clone(),
411 };
412 self.object_algorithm
413 .canonicalize_state(&local_problem, state)
414 }
415
416 fn iterations(&self) -> usize {
417 self.outer_iterations
418 }
419
420 fn batch_size(&self) -> usize {
421 usize::MAX
422 }
423}
424
425#[derive(Clone, Debug, Serialize, Deserialize)]
427#[serde(deny_unknown_fields)]
428pub struct JointReconstructionResult {
429 pub reconstruction: ReconstructionResult,
431 pub initial_illumination: Illumination,
433 pub calibrated_illumination: Illumination,
435 pub calibrated_model: ImagePlaneModel,
437 pub initial_parameters: PlanarArrayParameterValues,
439 pub final_parameters: PlanarArrayParameterValues,
441 pub parameter_history: Vec<CalibrationParameterHistoryEntry>,
443 pub loss_history: Vec<CalibrationLossHistoryEntry>,
445 pub convergence_reason: Option<CalibrationConvergenceReason>,
447 pub conditioning: CalibrationConditioning,
449 pub diagnostics: IlluminationCalibrationState,
451}
452
453impl JointReconstructionResult {
454 pub fn from_reconstruction(reconstruction: ReconstructionResult) -> Result<Self> {
456 let diagnostics = reconstruction
457 .physical_illumination_calibration
458 .clone()
459 .ok_or_else(|| {
460 Error::InvalidModel("joint result has no physical calibration".into())
461 })?;
462 let calibrated_model = reconstruction
463 .calibrated_model
464 .clone()
465 .ok_or_else(|| Error::InvalidModel("joint result has no calibrated model".into()))?;
466 Ok(Self {
467 initial_illumination: diagnostics.initial_illumination.clone(),
468 calibrated_illumination: diagnostics.current_illumination.clone(),
469 calibrated_model,
470 initial_parameters: diagnostics.initial_parameters.clone(),
471 final_parameters: diagnostics.current_parameters.clone(),
472 parameter_history: diagnostics.parameter_history.clone(),
473 loss_history: diagnostics.loss_history.clone(),
474 convergence_reason: diagnostics.convergence_reason,
475 conditioning: diagnostics.conditioning.clone(),
476 diagnostics,
477 reconstruction,
478 })
479 }
480
481 pub fn save_json(&self, path: impl AsRef<Path>) -> Result<()> {
483 let writer = std::io::BufWriter::new(std::fs::File::create(path)?);
484 serde_json::to_writer(writer, self)?;
485 Ok(())
486 }
487
488 pub fn load_json(path: impl AsRef<Path>) -> Result<Self> {
490 let reader = std::io::BufReader::new(std::fs::File::open(path)?);
491 let result: Self = serde_json::from_reader(reader)?;
492 result.reconstruction.validate()?;
493 result.calibrated_model.validate()?;
494 Ok(result)
495 }
496
497 #[cfg(feature = "parquet")]
499 pub fn write_bundle(
500 &self,
501 path: impl AsRef<Path>,
502 options: crate::reconstruction::BundleExportOptions,
503 ) -> Result<crate::reconstruction::ResultBundle> {
504 self.reconstruction.write_bundle(path, options)
505 }
506}