1use std::error::Error;
4use std::fmt::{Display, Formatter};
5
6use phasesmith_crystallography::P1ParameterLayout;
7use phasesmith_model::RecordId;
8
9use crate::{
10 LatticeBounds, LatticeError, LatticeParameterization, ParameterBounds, ParameterError,
11 ParameterKey, ParameterSet, ParameterSpec, RietveldPhase,
12};
13
14#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
16#[allow(clippy::struct_excessive_bools)]
17pub struct RietveldStructuralSelection {
18 pub lattice: bool,
20 pub coordinates: bool,
22 pub occupancy: bool,
24 pub u_iso: bool,
26 pub phase_scale: bool,
28}
29
30#[derive(Clone, Debug, PartialEq)]
32pub struct SiteCoordinateModel {
33 parameter_names: Vec<String>,
34 basis: Vec<f64>,
35 special_position: bool,
36}
37
38impl SiteCoordinateModel {
39 #[allow(clippy::too_many_lines)]
46 pub fn new(
47 space_group: &phasesmith_crystallography::SpaceGroup,
48 coordinate: [f64; 3],
49 tolerance: f64,
50 ) -> Result<Self, RietveldParameterError> {
51 if coordinate.iter().any(|value| !value.is_finite())
52 || !tolerance.is_finite()
53 || tolerance <= 0.0
54 {
55 return Err(RietveldParameterError::InvalidCoordinateModel);
56 }
57 let mut equations = Vec::<[f64; 3]>::new();
58 for operation in space_group.operations() {
59 let rotation = operation.rotation();
60 let translation = operation.translation();
61 let mut difference = [0.0; 3];
62 for row in 0..3 {
63 difference[row] = translation[row].as_f64() - coordinate[row];
64 for column in 0..3 {
65 difference[row] += f64::from(rotation[row][column]) * coordinate[column];
66 }
67 }
68 if difference
69 .iter()
70 .all(|value| (value - value.round()).abs() <= tolerance)
71 {
72 for row in 0..3 {
73 let mut equation = rotation[row].map(f64::from);
74 equation[row] -= 1.0;
75 equations.push(equation);
76 }
77 }
78 }
79 let (column_count, mut basis) = deterministic_null_space(equations);
80 let special_position = column_count != 3;
81 if !special_position {
82 basis = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0];
83 }
84 for value in &mut basis {
85 if value.abs() < 1.0e-14 {
86 *value = 0.0;
87 }
88 }
89 let parameter_names = if special_position {
90 (0..column_count).map(|index| format!("q{index}")).collect()
91 } else {
92 ["x", "y", "z"].map(str::to_owned).to_vec()
93 };
94 Ok(Self {
95 parameter_names,
96 basis,
97 special_position,
98 })
99 }
100
101 #[must_use]
103 pub fn parameter_names(&self) -> &[String] {
104 &self.parameter_names
105 }
106
107 #[must_use]
109 pub fn basis(&self) -> &[f64] {
110 &self.basis
111 }
112
113 #[must_use]
115 pub const fn is_special_position(&self) -> bool {
116 self.special_position
117 }
118}
119
120#[derive(Clone, Debug, PartialEq)]
121struct ParameterMapping {
122 global_index: usize,
123 native_terms: Vec<(usize, f64)>,
124}
125
126#[derive(Clone, Debug, PartialEq)]
127struct PhaseDerivativeLayout {
128 native_count: usize,
129 mappings: Vec<ParameterMapping>,
130}
131
132#[derive(Clone, Debug, PartialEq)]
134pub struct RietveldStructuralLayout {
135 parameters: ParameterSet,
136 phases: Vec<PhaseDerivativeLayout>,
137 coordinate_models: Vec<Vec<SiteCoordinateModel>>,
138 phase_ids: Vec<RecordId>,
139 site_ids: Vec<Vec<RecordId>>,
140 space_groups: Vec<phasesmith_crystallography::SpaceGroup>,
141 anisotropic_masks: Vec<Vec<bool>>,
142}
143
144impl RietveldStructuralLayout {
145 #[allow(clippy::too_many_lines)]
155 pub fn new(
156 phases: &[RietveldPhase],
157 selection: RietveldStructuralSelection,
158 lattice_bounds: &[Option<LatticeBounds>],
159 ) -> Result<Self, RietveldParameterError> {
160 if phases.len() != lattice_bounds.len() {
161 return Err(RietveldParameterError::PhaseCountMismatch);
162 }
163 let mut specs = Vec::new();
164 let mut layouts = Vec::with_capacity(phases.len());
165 let mut all_coordinate_models = Vec::with_capacity(phases.len());
166 for (phase, supplied_bounds) in phases.iter().zip(lattice_bounds) {
167 let definition = phase.definition();
168 let native = P1ParameterLayout {
169 site_count: definition.fractional_xyz.len(),
170 };
171 let mut mappings = Vec::new();
172 if selection.lattice {
173 let bounds = supplied_bounds
174 .as_ref()
175 .ok_or(RietveldParameterError::MissingLatticeBounds)?;
176 let parameterization =
177 LatticeParameterization::new(definition.space_group.clone(), definition.cell)?;
178 if bounds.parameter_names() != parameterization.parameter_names() {
179 return Err(RietveldParameterError::Lattice(LatticeError::InvalidBounds));
180 }
181 let values = parameterization.values_from_cell(definition.cell)?;
182 let jacobian = parameterization.cell_jacobian(&values)?;
183 let columns = values.len();
184 for column in 0..columns {
185 let name = ¶meterization.parameter_names()[column];
186 let unit = if name.ends_with("_angstrom") {
187 "angstrom"
188 } else {
189 "degree"
190 };
191 push_mapping(
192 &mut specs,
193 &mut mappings,
194 ParameterKey::new("lattice", phase.phase_id().as_str(), name)?,
195 values[column],
196 unit,
197 ParameterBounds::new(bounds.lower()[column], bounds.upper()[column])?,
198 values[column].abs().max(1.0),
199 (0..6)
200 .filter_map(|row| {
201 let coefficient = jacobian[row * columns + column];
202 (coefficient != 0.0).then_some((row, coefficient))
203 })
204 .collect(),
205 )?;
206 }
207 }
208 let coordinate_models = definition
209 .fractional_xyz
210 .iter()
211 .map(|coordinate| {
212 SiteCoordinateModel::new(
213 &definition.space_group,
214 *coordinate,
215 definition.coordinate_tolerance,
216 )
217 })
218 .collect::<Result<Vec<_>, _>>()?;
219 for (site, (site_id, model)) in
220 phase.site_ids().iter().zip(&coordinate_models).enumerate()
221 {
222 let owner = format!("{}/{}", phase.phase_id(), site_id);
223 if selection.coordinates {
224 let columns = model.parameter_names.len();
225 for column in 0..columns {
226 let value = if model.special_position {
227 0.0
228 } else {
229 definition.fractional_xyz[site][column]
230 };
231 push_mapping(
232 &mut specs,
233 &mut mappings,
234 ParameterKey::new("site", &owner, &model.parameter_names[column])?,
235 value,
236 "fractional",
237 if model.special_position {
238 ParameterBounds::new(-0.5, 0.5)?
239 } else {
240 ParameterBounds::default()
241 },
242 1.0,
243 (0..3)
244 .filter_map(|row| {
245 let coefficient = model.basis[row * columns + column];
246 (coefficient != 0.0)
247 .then_some((native.coordinate(site, row), coefficient))
248 })
249 .collect(),
250 )?;
251 }
252 }
253 if selection.occupancy {
254 let value = definition.occupancy[site];
255 push_mapping(
256 &mut specs,
257 &mut mappings,
258 ParameterKey::new("site", &owner, "occupancy")?,
259 value,
260 "fraction",
261 ParameterBounds::new(0.0, (2.0 * value + 0.1).max(1.0))?,
262 value.max(1.0),
263 vec![(native.occupancy(site), 1.0)],
264 )?;
265 }
266 if selection.u_iso && !definition.anisotropic_mask[site] {
267 let value = definition.u_iso_angstrom2[site];
268 push_mapping(
269 &mut specs,
270 &mut mappings,
271 ParameterKey::new("site", &owner, "u_iso_angstrom2")?,
272 value,
273 "angstrom^2",
274 ParameterBounds::new(0.0, (2.0 * value + 0.05).max(0.5))?,
275 value.max(0.01),
276 vec![(native.u_iso(site), 1.0)],
277 )?;
278 }
279 }
280 if selection.phase_scale {
281 let numerical_scale = if definition.scale == 0.0 {
282 1.0
283 } else {
284 definition.scale.abs()
285 };
286 push_mapping(
287 &mut specs,
288 &mut mappings,
289 ParameterKey::new("phase", phase.phase_id().as_str(), "scale")?,
290 definition.scale,
291 "relative",
292 ParameterBounds::new(0.0, f64::INFINITY)?,
293 numerical_scale,
294 vec![(native.scale(), 1.0)],
295 )?;
296 }
297 layouts.push(PhaseDerivativeLayout {
298 native_count: native.parameter_count(),
299 mappings,
300 });
301 all_coordinate_models.push(coordinate_models);
302 }
303 Ok(Self {
304 parameters: ParameterSet::new(specs)?,
305 phases: layouts,
306 coordinate_models: all_coordinate_models,
307 phase_ids: phases
308 .iter()
309 .map(|phase| phase.phase_id().clone())
310 .collect(),
311 site_ids: phases
312 .iter()
313 .map(|phase| phase.site_ids().to_vec())
314 .collect(),
315 space_groups: phases
316 .iter()
317 .map(|phase| phase.definition().space_group.clone())
318 .collect(),
319 anisotropic_masks: phases
320 .iter()
321 .map(|phase| phase.definition().anisotropic_mask.clone())
322 .collect(),
323 })
324 }
325
326 #[must_use]
328 pub const fn parameters(&self) -> &ParameterSet {
329 &self.parameters
330 }
331
332 #[must_use]
334 pub fn coordinate_models(&self) -> &[Vec<SiteCoordinateModel>] {
335 &self.coordinate_models
336 }
337
338 pub(crate) fn validate_phases(
339 &self,
340 phases: &[RietveldPhase],
341 ) -> Result<(), RietveldParameterError> {
342 if phases.len() != self.phase_ids.len()
343 || phases.iter().enumerate().any(|(index, phase)| {
344 phase.phase_id() != &self.phase_ids[index]
345 || phase.site_ids() != self.site_ids[index]
346 || phase.definition().space_group != self.space_groups[index]
347 || phase.definition().anisotropic_mask != self.anisotropic_masks[index]
348 || 6 + 5 * phase.definition().fractional_xyz.len() + 1
349 != self.phases[index].native_count
350 })
351 {
352 return Err(RietveldParameterError::PhaseIdentityMismatch);
353 }
354 Ok(())
355 }
356
357 pub fn apply_values(
364 &self,
365 phases: &[RietveldPhase],
366 values: &[f64],
367 ) -> Result<Vec<RietveldPhase>, RietveldParameterError> {
368 let current = self
369 .parameters
370 .specs()
371 .iter()
372 .map(ParameterSpec::value)
373 .collect::<Vec<_>>();
374 self.apply_value_change(phases, ¤t, values)
375 }
376
377 pub fn apply_value_change(
388 &self,
389 phases: &[RietveldPhase],
390 current_values: &[f64],
391 values: &[f64],
392 ) -> Result<Vec<RietveldPhase>, RietveldParameterError> {
393 self.validate_phases(phases)?;
394 if current_values.len() != self.parameters.specs().len()
395 || values.len() != self.parameters.specs().len()
396 || current_values.iter().any(|value| !value.is_finite())
397 || values.iter().any(|value| !value.is_finite())
398 {
399 return Err(RietveldParameterError::ValueLengthMismatch);
400 }
401 if let Some((spec, value)) = self
402 .parameters
403 .specs()
404 .iter()
405 .zip(values)
406 .find(|(spec, value)| !spec.bounds().contains(**value))
407 {
408 return Err(RietveldParameterError::Parameter(
409 ParameterError::ValueOutsideBounds {
410 key: spec.key().clone(),
411 value: *value,
412 },
413 ));
414 }
415 let mut updated = Vec::with_capacity(phases.len());
416 for (phase_index, phase) in phases.iter().enumerate() {
417 let mut definition = phase.definition().clone();
418 let phase_id = phase.phase_id().as_str();
419 let parameterization =
420 LatticeParameterization::new(definition.space_group.clone(), definition.cell)?;
421 let mut lattice_values = parameterization.values_from_cell(definition.cell)?;
422 let mut lattice_changed = false;
423 for (index, name) in parameterization.parameter_names().iter().enumerate() {
424 if let Some(value) = self.value_for("lattice", phase_id, name, values)? {
425 lattice_values[index] = value;
426 lattice_changed = true;
427 }
428 }
429 if lattice_changed {
430 definition.cell = parameterization.to_cell(&lattice_values)?;
431 }
432 for (site, site_id) in phase.site_ids().iter().enumerate() {
433 let owner = format!("{phase_id}/{site_id}");
434 let model = &self.coordinate_models[phase_index][site];
435 if model.special_position {
436 let columns = model.parameter_names.len();
437 for column in 0..columns {
438 if let Some((before, after)) = self.value_change_for(
439 "site",
440 &owner,
441 &model.parameter_names[column],
442 current_values,
443 values,
444 )? {
445 for row in 0..3 {
446 definition.fractional_xyz[site][row] +=
447 model.basis[row * columns + column] * (after - before);
448 }
449 }
450 }
451 } else {
452 for (component, name) in ["x", "y", "z"].iter().enumerate() {
453 if let Some(value) = self.value_for("site", &owner, name, values)? {
454 definition.fractional_xyz[site][component] = value;
455 }
456 }
457 }
458 if let Some(value) = self.value_for("site", &owner, "occupancy", values)? {
459 definition.occupancy[site] = value;
460 }
461 if let Some(value) = self.value_for("site", &owner, "u_iso_angstrom2", values)? {
462 definition.u_iso_angstrom2[site] = value;
463 }
464 }
465 if let Some(value) = self.value_for("phase", phase_id, "scale", values)? {
466 definition.scale = value;
467 }
468 updated.push(
469 phase
470 .with_definition(definition)
471 .map_err(RietveldParameterError::Rietveld)?,
472 );
473 }
474 Ok(updated)
475 }
476
477 fn value_change_for(
478 &self,
479 module: &str,
480 owner: &str,
481 name: &str,
482 current_values: &[f64],
483 values: &[f64],
484 ) -> Result<Option<(f64, f64)>, RietveldParameterError> {
485 let key = ParameterKey::new(module, owner, name)?;
486 Ok(self
487 .parameters
488 .index_of(&key)
489 .map(|index| (current_values[index], values[index])))
490 }
491
492 fn value_for(
493 &self,
494 module: &str,
495 owner: &str,
496 name: &str,
497 values: &[f64],
498 ) -> Result<Option<f64>, RietveldParameterError> {
499 let key = ParameterKey::new(module, owner, name)?;
500 Ok(self.parameters.index_of(&key).map(|index| values[index]))
501 }
502
503 pub fn native_tangents(
510 &self,
511 direction: &[f64],
512 ) -> Result<Vec<Vec<f64>>, RietveldParameterError> {
513 if direction.len() != self.parameters.specs().len() {
514 return Err(RietveldParameterError::DirectionLengthMismatch);
515 }
516 Ok(self
517 .phases
518 .iter()
519 .map(|phase| {
520 let mut tangent = vec![0.0; phase.native_count];
521 for mapping in &phase.mappings {
522 for &(native_index, coefficient) in &mapping.native_terms {
523 tangent[native_index] += coefficient * direction[mapping.global_index];
524 }
525 }
526 tangent
527 })
528 .collect())
529 }
530
531 pub fn project_native_gradients(
538 &self,
539 gradients: &[&[f64]],
540 ) -> Result<Vec<f64>, RietveldParameterError> {
541 if gradients.len() != self.phases.len() {
542 return Err(RietveldParameterError::PhaseCountMismatch);
543 }
544 let mut projected = vec![0.0; self.parameters.specs().len()];
545 for (phase, gradient) in self.phases.iter().zip(gradients) {
546 if gradient.len() != phase.native_count {
547 return Err(RietveldParameterError::NativeGradientLengthMismatch);
548 }
549 for mapping in &phase.mappings {
550 projected[mapping.global_index] += mapping
551 .native_terms
552 .iter()
553 .map(|(native_index, coefficient)| coefficient * gradient[*native_index])
554 .sum::<f64>();
555 }
556 }
557 Ok(projected)
558 }
559
560 pub fn project_native_jacobians(
567 &self,
568 jacobians: &[(&[f64], usize)],
569 sample_count: usize,
570 ) -> Result<Vec<f64>, RietveldParameterError> {
571 if jacobians.len() != self.phases.len() {
572 return Err(RietveldParameterError::PhaseCountMismatch);
573 }
574 let element_count = self
575 .parameters
576 .specs()
577 .len()
578 .checked_mul(sample_count)
579 .ok_or(RietveldParameterError::AllocationOverflow)?;
580 let mut projected = vec![0.0; element_count];
581 for (phase, (jacobian, native_count)) in self.phases.iter().zip(jacobians) {
582 let native_elements = native_count
583 .checked_mul(sample_count)
584 .ok_or(RietveldParameterError::AllocationOverflow)?;
585 if *native_count != phase.native_count || jacobian.len() != native_elements {
586 return Err(RietveldParameterError::NativeGradientLengthMismatch);
587 }
588 for mapping in &phase.mappings {
589 let target_start = mapping
590 .global_index
591 .checked_mul(sample_count)
592 .ok_or(RietveldParameterError::AllocationOverflow)?;
593 for &(native_index, coefficient) in &mapping.native_terms {
594 let source_start = native_index
595 .checked_mul(sample_count)
596 .ok_or(RietveldParameterError::AllocationOverflow)?;
597 for sample in 0..sample_count {
598 projected[target_start + sample] +=
599 coefficient * jacobian[source_start + sample];
600 }
601 }
602 }
603 }
604 Ok(projected)
605 }
606}
607
608#[allow(clippy::too_many_arguments)]
609fn push_mapping(
610 specs: &mut Vec<ParameterSpec>,
611 mappings: &mut Vec<ParameterMapping>,
612 key: ParameterKey,
613 value: f64,
614 unit: &str,
615 bounds: ParameterBounds,
616 scale: f64,
617 native_terms: Vec<(usize, f64)>,
618) -> Result<(), ParameterError> {
619 let global_index = specs.len();
620 specs.push(ParameterSpec::new(key, value, unit, bounds, scale, true)?);
621 mappings.push(ParameterMapping {
622 global_index,
623 native_terms,
624 });
625 Ok(())
626}
627
628fn deterministic_null_space(mut equations: Vec<[f64; 3]>) -> (usize, Vec<f64>) {
629 let mut pivot_columns = Vec::new();
630 let mut pivot_row = 0;
631 for column in 0..3 {
632 let Some(row) =
633 (pivot_row..equations.len()).find(|row| equations[*row][column].abs() > 1.0e-12)
634 else {
635 continue;
636 };
637 equations.swap(pivot_row, row);
638 let pivot = equations[pivot_row][column];
639 for value in &mut equations[pivot_row] {
640 *value /= pivot;
641 }
642 let pivot_values = equations[pivot_row];
643 for (row, equation) in equations.iter_mut().enumerate() {
644 if row != pivot_row {
645 let factor = equation[column];
646 for index in 0..3 {
647 equation[index] -= factor * pivot_values[index];
648 }
649 }
650 }
651 pivot_columns.push(column);
652 pivot_row += 1;
653 }
654 let free_columns = (0..3)
655 .filter(|column| !pivot_columns.contains(column))
656 .collect::<Vec<_>>();
657 let column_count = free_columns.len();
658 let mut basis = vec![0.0; 3 * column_count];
659 for (basis_column, free_column) in free_columns.iter().copied().enumerate() {
660 let mut vector = [0.0; 3];
661 vector[free_column] = 1.0;
662 for (row, pivot_column) in pivot_columns.iter().copied().enumerate() {
663 vector[pivot_column] = -equations[row][free_column];
664 }
665 let norm = vector.iter().map(|value| value * value).sum::<f64>().sqrt();
666 for row in 0..3 {
667 basis[row * column_count + basis_column] = vector[row] / norm;
668 }
669 }
670 (column_count, basis)
671}
672
673#[derive(Debug)]
675pub enum RietveldParameterError {
676 PhaseCountMismatch,
678 MissingLatticeBounds,
680 DirectionLengthMismatch,
682 NativeGradientLengthMismatch,
684 AllocationOverflow,
686 ValueLengthMismatch,
688 PhaseIdentityMismatch,
690 InvalidCoordinateModel,
692 Parameter(ParameterError),
694 Lattice(LatticeError),
696 Rietveld(crate::RietveldError),
698}
699
700impl Display for RietveldParameterError {
701 fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
702 match self {
703 Self::PhaseCountMismatch => formatter.write_str("Rietveld phase counts must match"),
704 Self::MissingLatticeBounds => {
705 formatter.write_str("selected Rietveld lattices require finite bounds")
706 }
707 Self::DirectionLengthMismatch => {
708 formatter.write_str("Rietveld parameter direction length mismatch")
709 }
710 Self::NativeGradientLengthMismatch => {
711 formatter.write_str("Rietveld native gradient length mismatch")
712 }
713 Self::AllocationOverflow => {
714 formatter.write_str("Rietveld derivative allocation overflow")
715 }
716 Self::ValueLengthMismatch => {
717 formatter.write_str("Rietveld physical value length mismatch")
718 }
719 Self::PhaseIdentityMismatch => {
720 formatter.write_str("Rietveld parameter layout phase identities differ")
721 }
722 Self::InvalidCoordinateModel => {
723 formatter.write_str("Rietveld site coordinate model is invalid")
724 }
725 Self::Parameter(error) => Display::fmt(error, formatter),
726 Self::Lattice(error) => Display::fmt(error, formatter),
727 Self::Rietveld(error) => Display::fmt(error, formatter),
728 }
729 }
730}
731
732impl Error for RietveldParameterError {
733 fn source(&self) -> Option<&(dyn Error + 'static)> {
734 match self {
735 Self::Parameter(error) => Some(error),
736 Self::Lattice(error) => Some(error),
737 Self::Rietveld(error) => Some(error),
738 Self::PhaseCountMismatch
739 | Self::MissingLatticeBounds
740 | Self::DirectionLengthMismatch
741 | Self::NativeGradientLengthMismatch
742 | Self::AllocationOverflow
743 | Self::ValueLengthMismatch
744 | Self::PhaseIdentityMismatch
745 | Self::InvalidCoordinateModel => None,
746 }
747 }
748}
749
750impl From<ParameterError> for RietveldParameterError {
751 fn from(value: ParameterError) -> Self {
752 Self::Parameter(value)
753 }
754}
755
756impl From<LatticeError> for RietveldParameterError {
757 fn from(value: LatticeError) -> Self {
758 Self::Lattice(value)
759 }
760}