1use phasesmith_core::{
4 ConstantWavelengthInstrument, CwContributionsView, FcjGeometry, SupportPolicy,
5};
6use phasesmith_crystallography::{
7 IntegratedIntensityCorrectionModel, PreparedNeutronScattering, PreparedXrayScattering,
8 SpaceGroup, UnitCell,
9};
10use phasesmith_execution::ExecutionContext;
11
12use crate::structural_pattern::{
13 BuiltInScatteringModel, MonochromaticPositionCorrection, StructuralPatternDenseResult,
14 StructuralPatternError, StructuralPatternInputView, StructuralPatternJvpResult,
15 StructuralPatternResult, StructuralPatternVjpResult,
16 calculate_structural_pattern_dense_with_context, calculate_structural_pattern_jvp_with_context,
17 calculate_structural_pattern_vjp_with_context, calculate_structural_pattern_with_context,
18};
19
20#[derive(Clone, Debug, PartialEq)]
22pub struct StructuralPhaseDefinition {
23 pub cell: UnitCell,
25 pub space_group: SpaceGroup,
27 pub hkl: Vec<[i32; 3]>,
29 pub multiplicity: Vec<usize>,
31 pub fractional_xyz: Vec<[f64; 3]>,
33 pub occupancy: Vec<f64>,
35 pub u_iso_angstrom2: Vec<f64>,
37 pub anisotropic_mask: Vec<bool>,
39 pub u_aniso_cif_angstrom2: Vec<[f64; 6]>,
41 pub scattering_species: Vec<String>,
43 pub scattering_real_offset: Vec<f64>,
45 pub scattering_imag_offset: Vec<f64>,
47 pub scale: f64,
49 pub coordinate_tolerance: f64,
51 pub scattering_model: BuiltInScatteringModel,
53 pub correction_model: IntegratedIntensityCorrectionModel,
55}
56
57impl StructuralPhaseDefinition {
58 pub fn validate(&self) -> Result<(), StructuralPatternError> {
65 validate_definition(self)
66 }
67}
68
69#[derive(Clone, Copy, Debug)]
71pub struct PreparedStructuralPatternInputView<'a> {
72 pub x_deg: &'a [f64],
74 pub instrument: ConstantWavelengthInstrument,
76 pub axial_geometry: Option<FcjGeometry>,
78 pub position_correction: MonochromaticPositionCorrection,
80 pub contributions: CwContributionsView<'a>,
82 pub support: SupportPolicy,
84}
85
86#[derive(Clone)]
88pub struct PreparedStructuralPhase {
89 definition: StructuralPhaseDefinition,
90 execution: ExecutionContext,
91}
92
93impl PreparedStructuralPhase {
94 pub fn new(
101 definition: StructuralPhaseDefinition,
102 execution: ExecutionContext,
103 ) -> Result<Self, StructuralPatternError> {
104 definition.validate()?;
105 Ok(Self {
106 definition,
107 execution,
108 })
109 }
110
111 #[must_use]
113 pub fn reflection_count(&self) -> usize {
114 self.definition.hkl.len()
115 }
116
117 #[must_use]
119 pub fn structural_parameter_count(&self) -> usize {
120 6 + 5 * self.definition.fractional_xyz.len() + 1
121 }
122
123 #[must_use]
125 pub fn execution_threads(&self) -> usize {
126 self.execution.threads()
127 }
128
129 #[must_use]
131 pub const fn definition(&self) -> &StructuralPhaseDefinition {
132 &self.definition
133 }
134
135 pub fn calculate(
142 &self,
143 input: &PreparedStructuralPatternInputView<'_>,
144 ) -> Result<StructuralPatternResult, StructuralPatternError> {
145 self.with_input(input, calculate_structural_pattern_with_context)
146 }
147
148 pub fn linearize(
155 &self,
156 input: &PreparedStructuralPatternInputView<'_>,
157 ) -> Result<StructuralPatternDenseResult, StructuralPatternError> {
158 self.with_input(input, calculate_structural_pattern_dense_with_context)
159 }
160
161 pub fn jvp(
167 &self,
168 input: &PreparedStructuralPatternInputView<'_>,
169 tangent: &[f64],
170 ) -> Result<StructuralPatternJvpResult, StructuralPatternError> {
171 self.with_input(input, |cell, group, structural_input, execution| {
172 calculate_structural_pattern_jvp_with_context(
173 cell,
174 group,
175 structural_input,
176 tangent,
177 execution,
178 )
179 })
180 }
181
182 pub fn vjp(
188 &self,
189 input: &PreparedStructuralPatternInputView<'_>,
190 sample_weights: &[f64],
191 ) -> Result<StructuralPatternVjpResult, StructuralPatternError> {
192 self.with_input(input, |cell, group, structural_input, execution| {
193 calculate_structural_pattern_vjp_with_context(
194 cell,
195 group,
196 structural_input,
197 sample_weights,
198 execution,
199 )
200 })
201 }
202
203 fn with_input<R>(
204 &self,
205 input: &PreparedStructuralPatternInputView<'_>,
206 operation: impl FnOnce(
207 UnitCell,
208 &SpaceGroup,
209 &StructuralPatternInputView<'_>,
210 &ExecutionContext,
211 ) -> Result<R, StructuralPatternError>,
212 ) -> Result<R, StructuralPatternError> {
213 let definition = &self.definition;
214 let species = definition
215 .scattering_species
216 .iter()
217 .map(String::as_str)
218 .collect::<Vec<_>>();
219 let structural_input = StructuralPatternInputView {
220 x_deg: input.x_deg,
221 hkl: &definition.hkl,
222 multiplicity: &definition.multiplicity,
223 fractional_xyz: &definition.fractional_xyz,
224 occupancy: &definition.occupancy,
225 u_iso_angstrom2: &definition.u_iso_angstrom2,
226 anisotropic_mask: &definition.anisotropic_mask,
227 u_aniso_cif_angstrom2: &definition.u_aniso_cif_angstrom2,
228 scattering_species: &species,
229 scattering_real_offset: &definition.scattering_real_offset,
230 scattering_imag_offset: &definition.scattering_imag_offset,
231 scale: definition.scale,
232 coordinate_tolerance: definition.coordinate_tolerance,
233 instrument: input.instrument,
234 axial_geometry: input.axial_geometry,
235 position_correction: input.position_correction,
236 correction_model: correction_for_wavelength(
237 definition.correction_model,
238 input.instrument.wavelength_angstrom,
239 ),
240 scattering_model: definition.scattering_model,
241 contributions: input.contributions,
242 support: input.support,
243 };
244 operation(
245 definition.cell,
246 &definition.space_group,
247 &structural_input,
248 &self.execution,
249 )
250 }
251}
252
253fn validate_definition(
254 definition: &StructuralPhaseDefinition,
255) -> Result<(), StructuralPatternError> {
256 definition
257 .cell
258 .geometry()
259 .map(|_| ())
260 .map_err(StructuralPatternError::InvalidCell)?;
261 if definition.hkl.len() != definition.multiplicity.len() {
262 return Err(StructuralPatternError::ReflectionLengthMismatch);
263 }
264 let site_count = definition.fractional_xyz.len();
265 if site_count != definition.occupancy.len()
266 || site_count != definition.u_iso_angstrom2.len()
267 || site_count != definition.anisotropic_mask.len()
268 || site_count != definition.u_aniso_cif_angstrom2.len()
269 || site_count != definition.scattering_species.len()
270 || (!definition.scattering_real_offset.is_empty()
271 && site_count != definition.scattering_real_offset.len())
272 || (!definition.scattering_imag_offset.is_empty()
273 && site_count != definition.scattering_imag_offset.len())
274 || definition.scattering_real_offset.is_empty()
275 != definition.scattering_imag_offset.is_empty()
276 {
277 return Err(StructuralPatternError::SiteLengthMismatch);
278 }
279 if definition
280 .scattering_real_offset
281 .iter()
282 .chain(&definition.scattering_imag_offset)
283 .any(|value| !value.is_finite())
284 {
285 return Err(StructuralPatternError::NonFiniteScatteringOffset);
286 }
287 match definition.scattering_model {
288 BuiltInScatteringModel::XrayNonResonant => {
289 PreparedXrayScattering::new(definition.scattering_species.iter().map(String::as_str))
290 .map(|_| ())
291 }
292 BuiltInScatteringModel::NeutronNuclear => {
293 PreparedNeutronScattering::new(definition.scattering_species.iter().map(String::as_str))
294 .map(|_| ())
295 }
296 }
297 .map_err(StructuralPatternError::Scattering)
298}
299
300fn correction_for_wavelength(
301 correction: IntegratedIntensityCorrectionModel,
302 wavelength_angstrom: f64,
303) -> IntegratedIntensityCorrectionModel {
304 match correction {
305 IntegratedIntensityCorrectionModel::Neutral => IntegratedIntensityCorrectionModel::Neutral,
306 IntegratedIntensityCorrectionModel::BraggBrentanoUnpolarizedLp { .. } => {
307 IntegratedIntensityCorrectionModel::BraggBrentanoUnpolarizedLp {
308 wavelength_angstrom,
309 }
310 }
311 IntegratedIntensityCorrectionModel::BraggBrentanoPolarizedLp { polarization, .. } => {
312 IntegratedIntensityCorrectionModel::BraggBrentanoPolarizedLp {
313 wavelength_angstrom,
314 polarization,
315 }
316 }
317 IntegratedIntensityCorrectionModel::ConstantWavelengthNeutronLorentz { .. } => {
318 IntegratedIntensityCorrectionModel::ConstantWavelengthNeutronLorentz {
319 wavelength_angstrom,
320 }
321 }
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328 use phasesmith_crystallography::SymmetryOperation;
329
330 fn definition() -> StructuralPhaseDefinition {
331 StructuralPhaseDefinition {
332 cell: UnitCell {
333 a_angstrom: 5.0,
334 b_angstrom: 5.0,
335 c_angstrom: 5.0,
336 alpha_deg: 90.0,
337 beta_deg: 90.0,
338 gamma_deg: 90.0,
339 },
340 space_group: SpaceGroup::new(vec![SymmetryOperation::identity()]).expect("P1"),
341 hkl: vec![[1, 0, 0]],
342 multiplicity: vec![2],
343 fractional_xyz: vec![[0.0, 0.0, 0.0]],
344 occupancy: vec![1.0],
345 u_iso_angstrom2: vec![0.01],
346 anisotropic_mask: vec![false],
347 u_aniso_cif_angstrom2: vec![[0.0; 6]],
348 scattering_species: vec!["Si".to_owned()],
349 scattering_real_offset: Vec::new(),
350 scattering_imag_offset: Vec::new(),
351 scale: 1.0,
352 coordinate_tolerance: 1.0e-10,
353 scattering_model: BuiltInScatteringModel::XrayNonResonant,
354 correction_model: IntegratedIntensityCorrectionModel::Neutral,
355 }
356 }
357
358 #[test]
359 fn prepared_phase_owns_validated_data_and_execution() {
360 let phase = PreparedStructuralPhase::new(
361 definition(),
362 ExecutionContext::new(2).expect("execution context"),
363 )
364 .expect("prepared phase");
365 assert_eq!(phase.reflection_count(), 1);
366 assert_eq!(phase.structural_parameter_count(), 12);
367 assert_eq!(phase.execution_threads(), 2);
368 }
369
370 #[test]
371 fn prepared_phase_rejects_inconsistent_owned_shapes() {
372 let mut invalid_reflections = definition();
373 invalid_reflections.multiplicity.clear();
374 assert!(matches!(
375 PreparedStructuralPhase::new(invalid_reflections, ExecutionContext::serial()),
376 Err(StructuralPatternError::ReflectionLengthMismatch)
377 ));
378
379 let mut invalid_sites = definition();
380 invalid_sites.occupancy.clear();
381 assert!(matches!(
382 PreparedStructuralPhase::new(invalid_sites, ExecutionContext::serial()),
383 Err(StructuralPatternError::SiteLengthMismatch)
384 ));
385 }
386
387 #[test]
388 fn prepared_phase_rejects_unknown_scattering_species() {
389 let mut invalid = definition();
390 invalid.scattering_species[0] = "not-an-element".to_owned();
391 assert!(matches!(
392 PreparedStructuralPhase::new(invalid, ExecutionContext::serial()),
393 Err(StructuralPatternError::Scattering(_))
394 ));
395 }
396
397 #[test]
398 fn dynamic_wavelength_replaces_the_stored_correction_wavelength() {
399 let correction = correction_for_wavelength(
400 IntegratedIntensityCorrectionModel::BraggBrentanoPolarizedLp {
401 wavelength_angstrom: 0.5,
402 polarization: 0.7,
403 },
404 1.5406,
405 );
406 assert_eq!(
407 correction,
408 IntegratedIntensityCorrectionModel::BraggBrentanoPolarizedLp {
409 wavelength_angstrom: 1.5406,
410 polarization: 0.7,
411 }
412 );
413 }
414}