1#[cfg(feature = "parquet")]
8use std::fs;
9use std::{
10 collections::BTreeMap,
11 fs::File,
12 io::BufWriter,
13 path::{Path, PathBuf},
14 time::Instant,
15};
16
17use ndarray::ArrayView2;
18use serde::{Deserialize, Serialize};
19use uuid::Uuid;
20
21use crate::{
22 Complex64, Result,
23 algorithms::ReconstructionAlgorithm,
24 evaluation::{evaluate_frame_intensity, evaluate_reconstruction_with_problem},
25 measurements::MeasurementRead,
26 model::ImagePlaneModel,
27 reconstruction::{ReconstructionProblem, ReconstructionResult},
28};
29
30pub const BENCHMARK_RECORD_FORMAT_VERSION: u32 = 1;
32pub const SMOKE_BENCHMARK_PROFILE: &str = "smoke";
34pub const CPU_BENCHMARK_PROFILE: &str = "cpu";
36
37#[derive(Clone, Copy, Debug, Eq, PartialEq)]
43pub struct BenchmarkProfile {
44 pub name: &'static str,
46 pub description: &'static str,
48 pub expected_runtime: &'static str,
50 pub output_directory: &'static str,
52 pub algorithms: &'static [&'static str],
54}
55
56impl BenchmarkProfile {
57 pub fn output_path(&self) -> PathBuf {
59 PathBuf::from(self.output_directory)
60 }
61}
62
63pub const BENCHMARK_PROFILES: &[BenchmarkProfile] = &[
65 BenchmarkProfile {
66 name: SMOKE_BENCHMARK_PROFILE,
67 description: "offline synthetic sanity profile for all implemented CPU algorithms",
68 expected_runtime: "under 1 minute on a typical laptop CPU",
69 output_directory: "target/benchmark-results/smoke",
70 algorithms: &[
71 "AlternatingProjection",
72 "AdaptiveAlternatingProjection",
73 "Fpie",
74 "Mpie",
75 "Epry",
76 "Admm",
77 "GlobalGaussNewton",
78 "GradientDescent",
79 ],
80 },
81 BenchmarkProfile {
82 name: CPU_BENCHMARK_PROFILE,
83 description: "offline synthetic CPU comparison with longer iteration counts",
84 expected_runtime: "1-5 minutes on a typical laptop CPU",
85 output_directory: "target/benchmark-results/cpu",
86 algorithms: &[
87 "AlternatingProjection",
88 "AdaptiveAlternatingProjection",
89 "Fpie",
90 "Mpie",
91 "Epry",
92 "Admm",
93 "GlobalGaussNewton",
94 "GradientDescent",
95 ],
96 },
97];
98
99pub fn benchmark_profile(name: &str) -> Option<&'static BenchmarkProfile> {
101 BENCHMARK_PROFILES
102 .iter()
103 .find(|profile| profile.name == name)
104}
105
106pub fn annotate_benchmark_profile(record: &mut BenchmarkRecord, profile: &BenchmarkProfile) {
108 record
109 .metadata
110 .insert("benchmark_profile".into(), profile.name.into());
111 record.metadata.insert(
112 "benchmark_profile_expected_runtime".into(),
113 profile.expected_runtime.into(),
114 );
115 record.metadata.insert(
116 "benchmark_profile_output_directory".into(),
117 profile.output_directory.into(),
118 );
119}
120
121#[derive(Clone, Debug, Serialize, Deserialize)]
123#[serde(deny_unknown_fields)]
124pub struct BenchmarkRecord {
125 pub format_version: u32,
127 pub case_id: String,
129 pub run_id: String,
131 pub dataset_name: String,
133 pub dataset_version: Option<String>,
135 pub preset_name: Option<String>,
137 pub crate_version: String,
139 pub random_seed: Option<u64>,
141 pub spatial_crop: Option<[usize; 4]>,
143 pub algorithm: String,
145 pub algorithm_configuration: String,
147 pub success: bool,
149 pub error: Option<String>,
151 pub frame_count: usize,
153 pub image_shape: [usize; 2],
155 pub reconstruction_shape: [usize; 2],
157 pub completed_iterations: usize,
159 pub elapsed_seconds: f64,
161 pub initial_objective: Option<f64>,
163 pub final_objective: Option<f64>,
165 pub final_to_initial_objective_ratio: Option<f64>,
167 pub amplitude_rmse: Option<f64>,
169 pub phase_rmse: Option<f64>,
171 pub complex_field_relative_error: Option<f64>,
173 pub fourier_domain_relative_error: Option<f64>,
175 pub pupil_amplitude_rmse: Option<f64>,
177 pub pupil_phase_rmse: Option<f64>,
179 pub illumination_position_rmse: Option<f64>,
181 pub per_frame_residual_mean: Option<f64>,
183 pub per_frame_residual_max: Option<f64>,
185 pub frames: Vec<BenchmarkFrameRecord>,
187 pub output_paths: Vec<PathBuf>,
189 pub metadata: BTreeMap<String, String>,
191}
192
193#[derive(Clone, Debug, Serialize, Deserialize)]
195#[serde(deny_unknown_fields)]
196pub struct BenchmarkFrameRecord {
197 pub frame_index: usize,
199 pub original_frame_index: usize,
201 pub original_illumination_index: Option<usize>,
203 pub normalized_l2: Option<f64>,
205}
206
207impl BenchmarkRecord {
208 pub fn from_result(
213 case_id: impl Into<String>,
214 dataset_name: impl Into<String>,
215 algorithm_configuration: impl Into<String>,
216 result: &ReconstructionResult,
217 ) -> Self {
218 let image_shape = result.recovered_pupil.shape();
219 let reconstruction_shape = result.object.dim();
220 let frame_count = result
221 .metadata
222 .get("frame_count")
223 .and_then(|value| value.parse().ok())
224 .or_else(|| result.recovered_frame_gains.as_ref().map(Vec::len))
225 .unwrap_or(0);
226 let initial_objective = result
227 .trace
228 .iterations
229 .first()
230 .map(|record| record.objective);
231 let final_objective = result.trace.final_objective();
232 Self {
233 format_version: BENCHMARK_RECORD_FORMAT_VERSION,
234 case_id: case_id.into(),
235 run_id: Uuid::new_v4().to_string(),
236 dataset_name: dataset_name.into(),
237 dataset_version: result.metadata.get("dataset_version").cloned(),
238 preset_name: result.metadata.get("preset_name").cloned(),
239 crate_version: env!("CARGO_PKG_VERSION").into(),
240 random_seed: result
241 .metadata
242 .get("random_seed")
243 .and_then(|value| value.parse().ok()),
244 spatial_crop: None,
245 algorithm: result.runtime.algorithm.clone(),
246 algorithm_configuration: algorithm_configuration.into(),
247 success: true,
248 error: None,
249 frame_count,
250 image_shape: [image_shape.0, image_shape.1],
251 reconstruction_shape: [reconstruction_shape.0, reconstruction_shape.1],
252 completed_iterations: result.runtime.completed_iterations,
253 elapsed_seconds: result.runtime.elapsed_seconds,
254 initial_objective,
255 final_objective,
256 final_to_initial_objective_ratio: initial_objective.and_then(|initial| {
257 final_objective
258 .filter(|_| initial.abs() > f64::EPSILON)
259 .map(|final_value| final_value / initial)
260 }),
261 amplitude_rmse: None,
262 phase_rmse: None,
263 complex_field_relative_error: None,
264 fourier_domain_relative_error: None,
265 pupil_amplitude_rmse: None,
266 pupil_phase_rmse: None,
267 illumination_position_rmse: None,
268 per_frame_residual_mean: None,
269 per_frame_residual_max: None,
270 frames: (0..frame_count)
271 .map(|frame_index| BenchmarkFrameRecord {
272 frame_index,
273 original_frame_index: frame_index,
274 original_illumination_index: None,
275 normalized_l2: None,
276 })
277 .collect(),
278 output_paths: Vec::new(),
279 metadata: BTreeMap::new(),
280 }
281 }
282}
283
284pub fn run_benchmark_case<A, M>(
288 dataset_name: impl Into<String>,
289 algorithm_configuration: impl Into<String>,
290 algorithm: A,
291 problem: &ReconstructionProblem<M>,
292 ground_truth: Option<ArrayView2<'_, Complex64>>,
293 true_model: Option<&ImagePlaneModel>,
294 valid_object_mask: Option<ArrayView2<'_, u8>>,
295) -> (BenchmarkRecord, Option<ReconstructionResult>)
296where
297 A: ReconstructionAlgorithm,
298 M: MeasurementRead,
299{
300 let algorithm_name = short_type_name::<A>().to_owned();
301 let image_shape = problem.measurements.image_shape();
302 let reconstruction_shape = problem.model.reconstruction_shape;
303 let mut record = BenchmarkRecord {
304 format_version: BENCHMARK_RECORD_FORMAT_VERSION,
305 case_id: String::new(),
306 run_id: Uuid::new_v4().to_string(),
307 dataset_name: dataset_name.into(),
308 dataset_version: None,
309 preset_name: None,
310 crate_version: env!("CARGO_PKG_VERSION").into(),
311 random_seed: None,
312 spatial_crop: None,
313 algorithm: algorithm_name,
314 algorithm_configuration: algorithm_configuration.into(),
315 success: false,
316 error: None,
317 frame_count: problem.measurements.frame_count(),
318 image_shape: [image_shape.0, image_shape.1],
319 reconstruction_shape: [reconstruction_shape.0, reconstruction_shape.1],
320 completed_iterations: 0,
321 elapsed_seconds: 0.0,
322 initial_objective: None,
323 final_objective: None,
324 final_to_initial_objective_ratio: None,
325 amplitude_rmse: None,
326 phase_rmse: None,
327 complex_field_relative_error: None,
328 fourier_domain_relative_error: None,
329 pupil_amplitude_rmse: None,
330 pupil_phase_rmse: None,
331 illumination_position_rmse: None,
332 per_frame_residual_mean: None,
333 per_frame_residual_max: None,
334 frames: problem
335 .measurements
336 .frame_metadata()
337 .iter()
338 .enumerate()
339 .map(|(index, metadata)| BenchmarkFrameRecord {
340 frame_index: index,
341 original_frame_index: metadata.original_frame_index.unwrap_or(index),
342 original_illumination_index: metadata
343 .original_illumination_index
344 .or(metadata.illumination_index),
345 normalized_l2: None,
346 })
347 .collect(),
348 output_paths: Vec::new(),
349 metadata: BTreeMap::new(),
350 };
351 record.case_id = format!("{:016x}", case_hash(&record));
352
353 let started = Instant::now();
354 let result = match algorithm.run(problem) {
355 Ok(result) => result,
356 Err(error) => {
357 record.elapsed_seconds = started.elapsed().as_secs_f64();
358 record.error = Some(error.to_string());
359 return (record, None);
360 }
361 };
362 record.elapsed_seconds = started.elapsed().as_secs_f64();
363 record.algorithm = result.runtime.algorithm.clone();
364 record.completed_iterations = result.runtime.completed_iterations;
365 record.initial_objective = result.trace.iterations.first().map(|entry| entry.objective);
366 record.final_objective = result.trace.final_objective();
367 record.final_to_initial_objective_ratio = record.initial_objective.and_then(|initial| {
368 record
369 .final_objective
370 .filter(|_| initial.abs() > f64::EPSILON)
371 .map(|final_objective| final_objective / initial)
372 });
373
374 let residuals: Vec<f64> = if let Some(truth) = ground_truth {
375 let metrics = evaluate_reconstruction_with_problem(
376 &result,
377 problem,
378 truth,
379 true_model,
380 valid_object_mask,
381 );
382 match metrics {
383 Ok(metrics) => {
384 record.amplitude_rmse = Some(metrics.object.amplitude_rmse);
385 record.phase_rmse = Some(metrics.object.phase_rmse);
386 record.complex_field_relative_error = Some(metrics.object.complex_nrmse);
387 record.fourier_domain_relative_error = Some(metrics.object.fourier_nrmse);
388 record.pupil_amplitude_rmse =
389 metrics.pupil.as_ref().map(|value| value.amplitude_rmse);
390 record.pupil_phase_rmse = metrics.pupil.as_ref().map(|value| value.phase_rmse);
391 record.illumination_position_rmse = metrics
392 .illumination
393 .as_ref()
394 .map(|value| value.position_rmse);
395 metrics
396 .intensity
397 .map(|value| {
398 value
399 .per_frame
400 .into_iter()
401 .map(|frame| frame.normalized_l2)
402 .collect()
403 })
404 .unwrap_or_default()
405 }
406 Err(error) => {
407 record.error = Some(format!("benchmark metric calculation failed: {error}"));
408 return (record, Some(result));
409 }
410 }
411 } else {
412 match evaluate_frame_intensity(&result, problem) {
413 Ok(metrics) => metrics
414 .per_frame
415 .into_iter()
416 .map(|frame| frame.normalized_l2)
417 .collect(),
418 Err(error) => {
419 record.error = Some(format!("benchmark residual calculation failed: {error}"));
420 return (record, Some(result));
421 }
422 }
423 };
424 if !residuals.is_empty() {
425 record.per_frame_residual_mean =
426 Some(residuals.iter().sum::<f64>() / residuals.len() as f64);
427 record.per_frame_residual_max = residuals.iter().copied().reduce(f64::max);
428 }
429 for (frame, residual) in record.frames.iter_mut().zip(residuals) {
430 frame.normalized_l2 = Some(residual);
431 }
432 let mut result = result;
433 result
434 .metadata
435 .insert("case_id".into(), record.case_id.clone());
436 result
437 .metadata
438 .insert("dataset_name".into(), record.dataset_name.clone());
439 if let Some(version) = &record.dataset_version {
440 result
441 .metadata
442 .insert("dataset_version".into(), version.clone());
443 }
444 result.metadata.insert(
445 "algorithm_configuration".into(),
446 record.algorithm_configuration.clone(),
447 );
448 result
449 .metadata
450 .insert("frame_count".into(), record.frame_count.to_string());
451 if let Some(seed) = record.random_seed {
452 result
453 .metadata
454 .insert("random_seed".into(), seed.to_string());
455 }
456 record.success = true;
457 (record, Some(result))
458}
459
460#[cfg(feature = "parquet")]
463pub fn save_benchmark_outputs(
464 record: &mut BenchmarkRecord,
465 result: &ReconstructionResult,
466 directory: impl AsRef<Path>,
467) -> Result<()> {
468 let directory = directory.as_ref();
469 fs::create_dir_all(directory)?;
470 let stem = format!(
471 "{}-{}-{:016x}",
472 safe_stem(&record.dataset_name),
473 safe_stem(&record.algorithm),
474 case_hash(record),
475 );
476 let outputs = [
477 directory.join(format!("{stem}-amplitude.png")),
478 directory.join(format!("{stem}-phase.png")),
479 directory.join(format!("{stem}-result")),
480 directory.join(format!("{stem}-trace.csv")),
481 ];
482 result.save_amplitude(&outputs[0])?;
483 result.save_phase(&outputs[1])?;
484 result.write_bundle(
485 &outputs[2],
486 crate::reconstruction::BundleExportOptions {
487 run_id: Some(record.run_id.clone()),
488 label: None,
489 include_previews: true,
490 },
491 )?;
492 result.save_trace_csv(&outputs[3])?;
493 record.output_paths.extend(outputs);
494 Ok(())
495}
496
497#[cfg(not(feature = "parquet"))]
498pub fn save_benchmark_outputs(
509 _record: &mut BenchmarkRecord,
510 _result: &ReconstructionResult,
511 _directory: impl AsRef<Path>,
512) -> Result<()> {
513 Err(crate::Error::Unsupported(
514 "benchmark result bundles require the `parquet` feature".into(),
515 ))
516}
517
518pub fn write_benchmark_json(records: &[BenchmarkRecord], path: impl AsRef<Path>) -> Result<()> {
520 #[derive(Serialize)]
521 struct Report<'a> {
522 format_version: u32,
523 records: &'a [BenchmarkRecord],
524 }
525
526 let writer = BufWriter::new(File::create(path)?);
527 serde_json::to_writer_pretty(
528 writer,
529 &Report {
530 format_version: BENCHMARK_RECORD_FORMAT_VERSION,
531 records,
532 },
533 )?;
534 Ok(())
535}
536
537pub fn write_benchmark_csv(records: &[BenchmarkRecord], path: impl AsRef<Path>) -> Result<()> {
539 let mut writer = csv::Writer::from_path(path)?;
540 writer.write_record([
541 "format_version",
542 "case_id",
543 "run_id",
544 "dataset_name",
545 "dataset_version",
546 "preset_name",
547 "crate_version",
548 "random_seed",
549 "spatial_crop",
550 "algorithm",
551 "algorithm_configuration",
552 "success",
553 "error",
554 "frame_count",
555 "image_height",
556 "image_width",
557 "reconstruction_height",
558 "reconstruction_width",
559 "completed_iterations",
560 "elapsed_seconds",
561 "initial_objective",
562 "final_objective",
563 "final_to_initial_objective_ratio",
564 "amplitude_rmse",
565 "phase_rmse",
566 "complex_field_relative_error",
567 "fourier_domain_relative_error",
568 "pupil_amplitude_rmse",
569 "pupil_phase_rmse",
570 "illumination_position_rmse",
571 "per_frame_residual_mean",
572 "per_frame_residual_max",
573 "frames_json",
574 "output_paths",
575 "metadata_json",
576 ])?;
577 for record in records {
578 writer.write_record([
579 record.format_version.to_string(),
580 record.case_id.clone(),
581 record.run_id.clone(),
582 record.dataset_name.clone(),
583 record.dataset_version.clone().unwrap_or_default(),
584 record.preset_name.clone().unwrap_or_default(),
585 record.crate_version.clone(),
586 record
587 .random_seed
588 .map_or_else(String::new, |seed| seed.to_string()),
589 record.spatial_crop.map_or_else(String::new, |crop| {
590 crop.iter()
591 .map(usize::to_string)
592 .collect::<Vec<_>>()
593 .join(";")
594 }),
595 record.algorithm.clone(),
596 record.algorithm_configuration.clone(),
597 record.success.to_string(),
598 record.error.clone().unwrap_or_default(),
599 record.frame_count.to_string(),
600 record.image_shape[0].to_string(),
601 record.image_shape[1].to_string(),
602 record.reconstruction_shape[0].to_string(),
603 record.reconstruction_shape[1].to_string(),
604 record.completed_iterations.to_string(),
605 record.elapsed_seconds.to_string(),
606 optional_number(record.initial_objective),
607 optional_number(record.final_objective),
608 optional_number(record.final_to_initial_objective_ratio),
609 optional_number(record.amplitude_rmse),
610 optional_number(record.phase_rmse),
611 optional_number(record.complex_field_relative_error),
612 optional_number(record.fourier_domain_relative_error),
613 optional_number(record.pupil_amplitude_rmse),
614 optional_number(record.pupil_phase_rmse),
615 optional_number(record.illumination_position_rmse),
616 optional_number(record.per_frame_residual_mean),
617 optional_number(record.per_frame_residual_max),
618 serde_json::to_string(&record.frames)?,
619 record
620 .output_paths
621 .iter()
622 .map(|path| path.display().to_string())
623 .collect::<Vec<_>>()
624 .join(";"),
625 serde_json::to_string(&record.metadata)?,
626 ])?;
627 }
628 writer.flush()?;
629 Ok(())
630}
631
632fn short_type_name<T>() -> &'static str {
633 std::any::type_name::<T>()
634 .rsplit("::")
635 .next()
636 .unwrap_or("reconstruction algorithm")
637}
638
639fn optional_number(value: Option<f64>) -> String {
640 value.map_or_else(String::new, |value| value.to_string())
641}
642
643#[cfg(feature = "parquet")]
644fn safe_stem(value: &str) -> String {
645 let stem: String = value
646 .chars()
647 .map(|character| {
648 if character.is_ascii_alphanumeric() || matches!(character, '-' | '_') {
649 character
650 } else {
651 '_'
652 }
653 })
654 .collect();
655 if stem.is_empty() {
656 "benchmark".into()
657 } else {
658 stem
659 }
660}
661
662fn case_hash(record: &BenchmarkRecord) -> u64 {
663 let mut hash = 0xcbf29ce484222325_u64;
666 let mut update = |bytes: &[u8]| {
667 for &byte in bytes {
668 hash ^= u64::from(byte);
669 hash = hash.wrapping_mul(0x100000001b3);
670 }
671 hash ^= 0xff;
672 hash = hash.wrapping_mul(0x100000001b3);
673 };
674 update(record.dataset_name.as_bytes());
675 update(
676 record
677 .dataset_version
678 .as_deref()
679 .unwrap_or_default()
680 .as_bytes(),
681 );
682 update(record.preset_name.as_deref().unwrap_or_default().as_bytes());
683 update(record.algorithm.as_bytes());
684 update(record.algorithm_configuration.as_bytes());
685 for frame in &record.frames {
686 update(&frame.original_frame_index.to_le_bytes());
687 update(
688 &frame
689 .original_illumination_index
690 .unwrap_or(usize::MAX)
691 .to_le_bytes(),
692 );
693 }
694 if let Some(crop) = record.spatial_crop {
695 for value in crop {
696 update(&value.to_le_bytes());
697 }
698 }
699 hash
700}