1use ndarray::{Array2, ArrayView2, ArrayViewMut2};
2use num_complex::Complex64;
3use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _};
4
5use crate::{
6 Result,
7 array_layout::{StandardArray2, checked_len_2d},
8 array_serde::Array2Data,
9 error::Error,
10 experiment::Optics,
11};
12
13use super::Sampling;
14
15#[derive(Clone, Debug, PartialEq)]
21pub struct Pupil {
22 pub(crate) values: StandardArray2<Complex64>,
23 pub(crate) support: StandardArray2<u8>,
24}
25
26impl Pupil {
27 pub fn new(values: Array2<Complex64>, support: Array2<u8>) -> Result<Self> {
33 let values = StandardArray2::try_from(values)?;
34 let support = StandardArray2::try_from(support)?;
35 if support.dim() != values.dim() {
36 return Err(Error::InvalidShape(format!(
37 "pupil support shape {:?} does not match value shape {:?}",
38 support.dim(),
39 values.dim()
40 )));
41 }
42 if values.dim().0 == 0 || values.dim().1 == 0 {
43 return Err(Error::InvalidShape(format!(
44 "pupil dimensions must be non-zero, got {:?}",
45 values.dim()
46 )));
47 }
48 if values
49 .as_slice()
50 .iter()
51 .any(|value| !value.re.is_finite() || !value.im.is_finite())
52 {
53 return Err(Error::InvalidModel(
54 "pupil contains non-finite complex values".into(),
55 ));
56 }
57 if support.as_slice().iter().any(|&value| value > 1) {
58 return Err(Error::InvalidModel(
59 "pupil support values must be exactly zero or one".into(),
60 ));
61 }
62 Ok(Self { values, support })
63 }
64
65 pub fn circular(shape: (usize, usize), sampling: &Sampling, optics: &Optics) -> Result<Self> {
79 sampling.validate()?;
80 optics.validate()?;
81 let cutoff = std::f64::consts::TAU * optics.objective_na / optics.wavelength_vacuum_m;
82 let medium_k = optics.objective_medium_wavenumber();
83 let length = checked_len_2d(shape)?;
84 if length == 0 {
85 return Err(Error::InvalidShape(format!(
86 "pupil dimensions must be non-zero, got {shape:?}"
87 )));
88 }
89 let mut values = Vec::with_capacity(length);
90 let mut support = Vec::with_capacity(length);
91 let aberration = optics.pupil_aberration.as_ref();
92 for row in 0..shape.0 {
93 let ky = (row as f64 - (shape.0 / 2) as f64) * sampling.dky;
94 for column in 0..shape.1 {
95 let kx = (column as f64 - (shape.1 / 2) as f64) * sampling.dkx;
96 let radius = kx.hypot(ky);
97 let inside = radius <= cutoff;
98 support.push(u8::from(inside));
99 if !inside {
100 values.push(Complex64::new(0.0, 0.0));
101 continue;
102 }
103 let rho = if cutoff > 0.0 { radius / cutoff } else { 0.0 };
104 let theta = ky.atan2(kx);
105 let mut phase = 0.0;
106 if let Some(defocus) = optics.defocus_distance {
107 phase -= defocus * (kx * kx + ky * ky) / (2.0 * medium_k);
108 }
109 let mut amplitude = 1.0;
110 if let Some(aberration) = aberration {
111 phase += aberration.astigmatism * rho * rho * (2.0 * theta).cos();
112 phase += aberration.coma * (3.0 * rho.powi(3) - 2.0 * rho) * theta.cos();
113 phase += aberration.spherical * (6.0 * rho.powi(4) - 6.0 * rho * rho + 1.0);
114 amplitude = (-aberration.edge_apodization * rho * rho).exp();
115 }
116 values.push(Complex64::from_polar(amplitude, phase));
117 }
118 }
119 Ok(Self {
120 values: StandardArray2::from_shape_vec(shape, values)?,
121 support: StandardArray2::from_shape_vec(shape, support)?,
122 })
123 }
124
125 pub fn shape(&self) -> (usize, usize) {
127 self.values.dim()
128 }
129
130 pub fn values(&self) -> ArrayView2<'_, Complex64> {
132 self.values.ndarray_view()
133 }
134
135 pub fn values_mut(&mut self) -> ArrayViewMut2<'_, Complex64> {
137 self.values.ndarray_view_mut()
138 }
139
140 pub fn support(&self) -> ArrayView2<'_, u8> {
142 self.support.ndarray_view()
143 }
144
145 pub fn replace_values(&mut self, values: Array2<Complex64>) -> Result<()> {
149 let values = StandardArray2::try_from(values)?;
150 if values.dim() != self.shape() {
151 return Err(Error::InvalidShape(format!(
152 "replacement pupil shape {:?} does not match {:?}",
153 values.dim(),
154 self.shape()
155 )));
156 }
157 if values
158 .as_slice()
159 .iter()
160 .any(|value| !value.re.is_finite() || !value.im.is_finite())
161 {
162 return Err(Error::InvalidModel(
163 "pupil contains non-finite complex values".into(),
164 ));
165 }
166 self.values = values;
167 Ok(())
168 }
169
170 pub fn apply_support(&mut self) {
172 for (value, &inside) in self
173 .values
174 .as_slice_mut()
175 .iter_mut()
176 .zip(self.support.as_slice())
177 {
178 if inside == 0 {
179 *value = Complex64::new(0.0, 0.0);
180 }
181 }
182 }
183}
184
185impl Serialize for Pupil {
186 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
187 where
188 S: Serializer,
189 {
190 #[derive(Serialize)]
191 struct Representation {
192 values: Array2Data<Complex64>,
193 support: Array2Data<u8>,
194 }
195
196 Representation {
197 values: Array2Data::from_view(self.values()),
198 support: Array2Data::from_view(self.support()),
199 }
200 .serialize(serializer)
201 }
202}
203
204impl<'de> Deserialize<'de> for Pupil {
205 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
206 where
207 D: Deserializer<'de>,
208 {
209 #[derive(Deserialize)]
210 #[serde(deny_unknown_fields)]
211 struct Representation {
212 values: Array2Data<Complex64>,
213 support: Array2Data<u8>,
214 }
215
216 let representation = Representation::deserialize(deserializer)?;
217 let values = representation
218 .values
219 .into_array()
220 .map_err(D::Error::custom)?;
221 let support = representation
222 .support
223 .into_array()
224 .map_err(D::Error::custom)?;
225 Self::new(values, support).map_err(D::Error::custom)
226 }
227}