use std::collections::BTreeSet;
use std::error::Error;
use std::fmt::{Display, Formatter};
use phasesmith_core::{
TOF_GLOBAL_PARAMETER_COUNT, TofBankGeometry, TofError, TofInstrument, TofInstrumentParameter,
};
use phasesmith_crystallography::{IntegratedIntensityCorrectionModel, P1ParameterLayout};
use phasesmith_engine::{
BuiltInScatteringModel, StructuralTofError, StructuralTofInputView, StructuralTofResult,
calculate_structural_tof_pattern_jvp_with_context,
calculate_structural_tof_pattern_vjp_with_context,
calculate_structural_tof_pattern_with_context,
};
use phasesmith_execution::ExecutionPolicy;
use phasesmith_model::{DomainError, RecordId, TofPatternRecord};
use crate::{
LatticeBounds, ParameterBounds, ParameterError, ParameterKey, ParameterSet, ParameterSpec,
ResidualError, ResidualEvaluation, ResidualOptions, RietveldError, RietveldParameterError,
RietveldPhase, RietveldStructuralLayout, RietveldStructuralSelection, TofChebyshevBackground,
TofInstrumentParameterBound, TofLeBailError, evaluate_tof_residuals,
};
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofBank {
pub bank_id: RecordId,
pub pattern: TofPatternRecord,
pub instrument: TofInstrument,
pub geometry: TofBankGeometry,
pub correction_model: IntegratedIntensityCorrectionModel,
pub scale: f64,
pub scale_bounds: ParameterBounds,
pub refine_scale: bool,
pub background: Option<TofChebyshevBackground>,
pub refine_background: bool,
pub instrument_bounds: Vec<TofInstrumentParameterBound>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofMultiBankInput {
pub phase: RietveldPhase,
pub structural_selection: RietveldStructuralSelection,
pub lattice_bounds: Option<LatticeBounds>,
pub banks: Vec<StructuralTofBank>,
pub support_fwhm: f64,
pub tail_log: f64,
pub use_uncertainty: bool,
pub execution: ExecutionPolicy,
}
impl StructuralTofMultiBankInput {
pub fn validate(&self) -> Result<(), StructuralTofMultiBankError> {
self.phase.validate()?;
let definition = self.phase.definition();
if definition.scattering_model != BuiltInScatteringModel::NeutronNuclear
|| !definition.scattering_real_offset.is_empty()
|| !definition.scattering_imag_offset.is_empty()
{
return Err(StructuralTofMultiBankError::InvalidPhaseContract(
"structural TOF requires built-in neutron scattering without X-ray offsets",
));
}
if definition.correction_model != IntegratedIntensityCorrectionModel::Neutral
|| definition.scale.to_bits() != 1.0_f64.to_bits()
{
return Err(StructuralTofMultiBankError::InvalidPhaseContract(
"the shared phase must use neutral correction and unit placeholder scale",
));
}
if self.phase.sample_physics().is_some()
|| self.phase.reflection_domain().is_some()
|| self.phase.contributions()
!= &phasesmith_core::OwnedCwContributions::neutral(definition.hkl.len())
{
return Err(StructuralTofMultiBankError::InvalidPhaseContract(
"CW sample physics and dynamic topology are not part of structural TOF",
));
}
if self.structural_selection.phase_scale {
return Err(StructuralTofMultiBankError::InvalidPhaseContract(
"structural TOF phase scale is bank-local and cannot be shared",
));
}
if !self.support_fwhm.is_finite()
|| self.support_fwhm <= 0.0
|| !self.tail_log.is_finite()
|| self.tail_log <= 0.0
{
return Err(StructuralTofMultiBankError::InvalidSupport);
}
if self.banks.is_empty() {
return Err(StructuralTofMultiBankError::TooFewBanks);
}
if self
.banks
.iter()
.map(|bank| &bank.bank_id)
.collect::<BTreeSet<_>>()
.len()
!= self.banks.len()
{
return Err(StructuralTofMultiBankError::DuplicateBankId);
}
for bank in &self.banks {
validate_bank(bank)?;
}
RietveldStructuralLayout::new(
std::slice::from_ref(&self.phase),
self.structural_selection,
std::slice::from_ref(&self.lattice_bounds),
)?;
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
struct BankParameterMapping {
bank_id: RecordId,
pattern: TofPatternRecord,
geometry: TofBankGeometry,
correction_model: IntegratedIntensityCorrectionModel,
scale_bounds: ParameterBounds,
scale: Option<usize>,
instrument_bounds: Vec<TofInstrumentParameterBound>,
instrument: Vec<(TofInstrumentParameter, usize)>,
background_contract: Option<(RecordId, [f64; 2], usize)>,
background: Vec<usize>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofMultiBankLayout {
parameters: ParameterSet,
structural: RietveldStructuralLayout,
structural_count: usize,
structural_selection: RietveldStructuralSelection,
lattice_bounds: Option<LatticeBounds>,
reflection_ids: Vec<String>,
support_fwhm: f64,
tail_log: f64,
use_uncertainty: bool,
execution: ExecutionPolicy,
banks: Vec<BankParameterMapping>,
}
impl StructuralTofMultiBankLayout {
pub fn new(input: &StructuralTofMultiBankInput) -> Result<Self, StructuralTofMultiBankError> {
input.validate()?;
let structural = RietveldStructuralLayout::new(
std::slice::from_ref(&input.phase),
input.structural_selection,
std::slice::from_ref(&input.lattice_bounds),
)?;
let mut specs = structural.parameters().specs().to_vec();
let structural_count = specs.len();
let mut banks = Vec::with_capacity(input.banks.len());
for bank in &input.banks {
let owner = bank.bank_id.as_str();
let scale = if bank.refine_scale {
let index = specs.len();
specs.push(ParameterSpec::new(
ParameterKey::new("tof_scale", owner, "scale")?,
bank.scale,
"relative",
bank.scale_bounds,
bank.scale.abs().max(1.0),
true,
)?);
Some(index)
} else {
None
};
let instrument_values = bank.instrument.values();
let mut instrument = Vec::with_capacity(bank.instrument_bounds.len());
for bound in &bank.instrument_bounds {
let index = specs.len();
let value = instrument_values[bound.parameter.index()];
let half_span = 0.5 * (bound.upper - bound.lower);
specs.push(ParameterSpec::new(
ParameterKey::new("tof_instrument", owner, bound.parameter.name())?,
value,
instrument_unit(bound.parameter),
ParameterBounds::new(bound.lower, bound.upper)?,
value.abs().max(half_span).max(f64::EPSILON.sqrt()),
true,
)?);
instrument.push((bound.parameter, index));
}
let mut background = Vec::new();
if bank.refine_background {
let model = bank.background.as_ref().ok_or(
StructuralTofMultiBankError::InvalidBankContract(
"refine_background requires a background model",
),
)?;
for (order, value) in model.coefficients().iter().copied().enumerate() {
let index = specs.len();
specs.push(ParameterSpec::new(
ParameterKey::new("tof_background", owner, format!("coefficient_{order}"))?,
value,
"intensity",
ParameterBounds::default(),
value.abs().max(1.0),
true,
)?);
background.push(index);
}
}
banks.push(BankParameterMapping {
bank_id: bank.bank_id.clone(),
pattern: bank.pattern.clone(),
geometry: bank.geometry,
correction_model: bank.correction_model,
scale_bounds: bank.scale_bounds,
scale,
instrument_bounds: bank.instrument_bounds.clone(),
instrument,
background_contract: bank.background.as_ref().map(|model| {
(
model.background_id().clone(),
model.domain_us(),
model.coefficients().len(),
)
}),
background,
});
}
Ok(Self {
parameters: ParameterSet::new(specs)?,
structural,
structural_count,
structural_selection: input.structural_selection,
lattice_bounds: input.lattice_bounds.clone(),
reflection_ids: input.phase.reflection_ids().to_vec(),
support_fwhm: input.support_fwhm,
tail_log: input.tail_log,
use_uncertainty: input.use_uncertainty,
execution: input.execution.clone(),
banks,
})
}
#[must_use]
pub const fn parameters(&self) -> &ParameterSet {
&self.parameters
}
pub fn apply_values(
&self,
input: &StructuralTofMultiBankInput,
values: &[f64],
) -> Result<StructuralTofMultiBankInput, StructuralTofMultiBankError> {
let current = self
.parameters
.specs()
.iter()
.map(ParameterSpec::value)
.collect::<Vec<_>>();
self.apply_value_change(input, ¤t, values)
}
pub fn apply_value_change(
&self,
input: &StructuralTofMultiBankInput,
current_values: &[f64],
values: &[f64],
) -> Result<StructuralTofMultiBankInput, StructuralTofMultiBankError> {
self.validate_contract(input)?;
if current_values.len() != self.parameters.specs().len()
|| values.len() != self.parameters.specs().len()
|| current_values.iter().any(|value| !value.is_finite())
{
return Err(StructuralTofMultiBankError::ParameterLengthMismatch);
}
for (spec, value) in self.parameters.specs().iter().zip(values) {
if !value.is_finite() || !spec.bounds().contains(*value) {
return Err(StructuralTofMultiBankError::Parameter(
ParameterError::ValueOutsideBounds {
key: spec.key().clone(),
value: *value,
},
));
}
}
let phase = self
.structural
.apply_value_change(
std::slice::from_ref(&input.phase),
¤t_values[..self.structural_count],
&values[..self.structural_count],
)?
.into_iter()
.next()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
let mut result = input.clone();
result.phase = phase;
for ((bank, mapping), original) in
result.banks.iter_mut().zip(&self.banks).zip(&input.banks)
{
if let Some(index) = mapping.scale {
bank.scale = values[index];
}
let mut instrument_values = original.instrument.values();
for &(parameter, index) in &mapping.instrument {
instrument_values[parameter.index()] = values[index];
}
bank.instrument = TofInstrument::from_values(instrument_values)?;
if !mapping.background.is_empty() {
let coefficients = mapping
.background
.iter()
.map(|index| values[*index])
.collect();
bank.background = Some(
original
.background
.as_ref()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?
.with_coefficients(coefficients)?,
);
}
}
result.validate()?;
Ok(result)
}
fn validate_contract(
&self,
input: &StructuralTofMultiBankInput,
) -> Result<(), StructuralTofMultiBankError> {
input.validate()?;
self.structural
.validate_phases(std::slice::from_ref(&input.phase))?;
if input.structural_selection != self.structural_selection
|| input.lattice_bounds != self.lattice_bounds
|| input.phase.reflection_ids() != self.reflection_ids
|| input.support_fwhm.to_bits() != self.support_fwhm.to_bits()
|| input.tail_log.to_bits() != self.tail_log.to_bits()
|| input.use_uncertainty != self.use_uncertainty
|| input.execution != self.execution
|| input.banks.len() != self.banks.len()
{
return Err(StructuralTofMultiBankError::BankContractMismatch);
}
for (bank, mapping) in input.banks.iter().zip(&self.banks) {
let background_contract = bank.background.as_ref().map(|model| {
(
model.background_id().clone(),
model.domain_us(),
model.coefficients().len(),
)
});
if bank.bank_id != mapping.bank_id
|| bank.pattern != mapping.pattern
|| bank.geometry != mapping.geometry
|| bank.correction_model != mapping.correction_model
|| bank.scale_bounds != mapping.scale_bounds
|| bank.refine_scale != mapping.scale.is_some()
|| bank.instrument_bounds != mapping.instrument_bounds
|| bank.refine_background == mapping.background.is_empty()
|| background_contract != mapping.background_contract
{
return Err(StructuralTofMultiBankError::BankContractMismatch);
}
}
Ok(())
}
fn native_tangent(
&self,
bank_index: usize,
direction: &[f64],
native_count: usize,
) -> Result<Vec<f64>, StructuralTofMultiBankError> {
if direction.len() != self.parameters.specs().len() {
return Err(StructuralTofMultiBankError::ParameterLengthMismatch);
}
let mut tangent = self
.structural
.native_tangents(&direction[..self.structural_count])?
.into_iter()
.next()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
if tangent.len() != native_count {
return Err(StructuralTofMultiBankError::InternalInvariant);
}
if let Some(index) = self.banks[bank_index].scale {
tangent[native_count - 1] = direction[index];
}
Ok(tangent)
}
fn scatter_native_gradient(
&self,
bank_index: usize,
native: &[f64],
output: &mut [f64],
) -> Result<(), StructuralTofMultiBankError> {
let shared = self.structural.project_native_gradients(&[native])?;
for (target, value) in output[..self.structural_count].iter_mut().zip(shared) {
*target += value;
}
if let Some(index) = self.banks[bank_index].scale {
output[index] += native
.last()
.copied()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
}
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofBankCalculation {
pub bank_id: RecordId,
pub y: Vec<f64>,
pub profile_y: Vec<f64>,
pub background_y: Vec<f64>,
pub structural: StructuralTofResult,
pub metrics: ResidualEvaluation,
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofMultiBankCalculation {
pub banks: Vec<StructuralTofBankCalculation>,
pub objective: f64,
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofMultiBankProduct {
pub bank_id: RecordId,
pub y: Vec<f64>,
pub derivative: Vec<f64>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct StructuralTofMultiBankGradient {
pub calculation: StructuralTofMultiBankCalculation,
pub gradient: Vec<f64>,
}
pub struct PreparedStructuralTofMultiBankObjective {
input: StructuralTofMultiBankInput,
layout: StructuralTofMultiBankLayout,
background_bases: Vec<Option<crate::TofChebyshevBasis>>,
}
impl PreparedStructuralTofMultiBankObjective {
pub fn new(input: StructuralTofMultiBankInput) -> Result<Self, StructuralTofMultiBankError> {
let layout = StructuralTofMultiBankLayout::new(&input)?;
let background_bases = input
.banks
.iter()
.map(|bank| {
bank.background
.as_ref()
.map(|model| model.basis(&bank.pattern.tof_us))
.transpose()
})
.collect::<Result<Vec<_>, _>>()?;
Ok(Self {
input,
layout,
background_bases,
})
}
#[must_use]
pub const fn input(&self) -> &StructuralTofMultiBankInput {
&self.input
}
#[must_use]
pub const fn layout(&self) -> &StructuralTofMultiBankLayout {
&self.layout
}
pub fn calculate(
&self,
) -> Result<StructuralTofMultiBankCalculation, StructuralTofMultiBankError> {
let mut banks = Vec::with_capacity(self.input.banks.len());
let mut objective = 0.0;
for (index, bank) in self.input.banks.iter().enumerate() {
let structural = calculate_bank(&self.input, bank)?;
let profile_y = structural.accumulation.y.clone();
let background_y = background_values(bank, self.background_bases[index].as_ref())?;
let y = profile_y
.iter()
.zip(&background_y)
.map(|(profile, background)| profile + background)
.collect::<Vec<_>>();
let metrics = evaluate_tof_residuals(
&bank.pattern,
&y,
ResidualOptions {
use_uncertainty: self.input.use_uncertainty,
parameter_count: self.layout.parameters.specs().len(),
},
)?;
objective += 0.5 * metrics.chi_square;
banks.push(StructuralTofBankCalculation {
bank_id: bank.bank_id.clone(),
y,
profile_y,
background_y,
structural,
metrics,
});
}
Ok(StructuralTofMultiBankCalculation { banks, objective })
}
pub fn jvp(
&self,
direction: &[f64],
) -> Result<Vec<StructuralTofMultiBankProduct>, StructuralTofMultiBankError> {
let native_count = P1ParameterLayout {
site_count: self.input.phase.definition().fractional_xyz.len(),
}
.parameter_count();
let mut products = Vec::with_capacity(self.input.banks.len());
for (bank_index, bank) in self.input.banks.iter().enumerate() {
let tangent = self
.layout
.native_tangent(bank_index, direction, native_count)?;
let species = species(self.input.phase.definition());
let view = bank_view(&self.input, bank, &species);
let forward = calculate_structural_tof_pattern_jvp_with_context(
self.input.phase.definition().cell,
&self.input.phase.definition().space_group,
&view,
&tangent,
self.input.execution.context(),
)?;
let mut derivative = forward.d_y;
add_instrument_jvp(
&mut derivative,
&forward.result,
&self.layout.banks[bank_index],
direction,
)?;
add_background_jvp(
&mut derivative,
self.background_bases[bank_index].as_ref(),
&self.layout.banks[bank_index],
direction,
)?;
let background = background_values(bank, self.background_bases[bank_index].as_ref())?;
let y = forward
.result
.accumulation
.y
.iter()
.zip(background)
.map(|(profile, background)| profile + background)
.collect();
products.push(StructuralTofMultiBankProduct {
bank_id: bank.bank_id.clone(),
y,
derivative,
});
}
Ok(products)
}
pub fn vjp(&self, weights: &[Vec<f64>]) -> Result<Vec<f64>, StructuralTofMultiBankError> {
if weights.len() != self.input.banks.len() {
return Err(StructuralTofMultiBankError::BankWeightCountMismatch);
}
let mut result = vec![0.0; self.layout.parameters.specs().len()];
for (bank_index, (bank, weights)) in self.input.banks.iter().zip(weights).enumerate() {
let species = species(self.input.phase.definition());
let view = bank_view(&self.input, bank, &species);
let reverse = calculate_structural_tof_pattern_vjp_with_context(
self.input.phase.definition().cell,
&self.input.phase.definition().space_group,
&view,
weights,
self.input.execution.context(),
)?;
self.layout
.scatter_native_gradient(bank_index, &reverse.gradient, &mut result)?;
add_instrument_vjp(
&mut result,
&reverse.result,
&self.layout.banks[bank_index],
weights,
)?;
add_background_vjp(
&mut result,
self.background_bases[bank_index].as_ref(),
&self.layout.banks[bank_index],
weights,
)?;
}
Ok(result)
}
pub fn gradient(&self) -> Result<StructuralTofMultiBankGradient, StructuralTofMultiBankError> {
let calculation = self.calculate()?;
let weights = self
.input
.banks
.iter()
.zip(&calculation.banks)
.map(|(bank, calculation)| {
objective_weights(bank, &calculation.metrics, self.input.use_uncertainty)
})
.collect::<Vec<_>>();
let gradient = self.vjp(&weights)?;
Ok(StructuralTofMultiBankGradient {
calculation,
gradient,
})
}
pub fn normal_product(
&self,
direction: &[f64],
damping: f64,
) -> Result<Vec<f64>, StructuralTofMultiBankError> {
if !damping.is_finite() || damping < 0.0 {
return Err(StructuralTofMultiBankError::InvalidDamping);
}
let products = self.jvp(direction)?;
let weights = self
.input
.banks
.iter()
.zip(products)
.map(|(bank, product)| {
weighted_direction(bank, product.derivative, self.input.use_uncertainty)
})
.collect::<Vec<_>>();
let mut result = self.vjp(&weights)?;
for (value, direction) in result.iter_mut().zip(direction) {
*value += damping * direction;
}
Ok(result)
}
}
fn validate_bank(bank: &StructuralTofBank) -> Result<(), StructuralTofMultiBankError> {
bank.pattern.validate()?;
if bank.pattern.observed_y.is_none() {
return Err(StructuralTofMultiBankError::MissingObservations);
}
bank.instrument.validate()?;
bank.geometry.validate()?;
if !bank.scale.is_finite() || bank.scale < 0.0 || !bank.scale_bounds.contains(bank.scale) {
return Err(StructuralTofMultiBankError::InvalidBankContract(
"bank scale must be finite, non-negative, and inside its bounds",
));
}
match bank.correction_model {
IntegratedIntensityCorrectionModel::Neutral => {}
IntegratedIntensityCorrectionModel::TimeOfFlightNeutronLorentz { two_theta_deg }
if two_theta_deg.to_bits() == bank.geometry.two_theta_deg.to_bits() => {}
_ => {
return Err(StructuralTofMultiBankError::InvalidBankContract(
"bank correction must be neutral or match the bank angle exactly",
));
}
}
let mut selected = BTreeSet::new();
let values = bank.instrument.values();
for bound in &bank.instrument_bounds {
if !bound.lower.is_finite()
|| !bound.upper.is_finite()
|| bound.lower >= bound.upper
|| !selected.insert(bound.parameter)
|| !(bound.lower..=bound.upper).contains(&values[bound.parameter.index()])
{
return Err(StructuralTofMultiBankError::InvalidBankContract(
"instrument selections require unique finite bounds containing the current value",
));
}
}
if let Some(background) = &bank.background {
background.validate()?;
background.basis(&bank.pattern.tof_us)?;
} else if bank.refine_background {
return Err(StructuralTofMultiBankError::InvalidBankContract(
"refine_background requires a background model",
));
}
Ok(())
}
fn species(definition: &phasesmith_engine::StructuralPhaseDefinition) -> Vec<&str> {
definition
.scattering_species
.iter()
.map(String::as_str)
.collect()
}
fn bank_view<'a>(
input: &'a StructuralTofMultiBankInput,
bank: &'a StructuralTofBank,
species: &'a [&'a str],
) -> StructuralTofInputView<'a> {
let definition = input.phase.definition();
StructuralTofInputView {
tof_us: &bank.pattern.tof_us,
hkl: &definition.hkl,
multiplicity: &definition.multiplicity,
fractional_xyz: &definition.fractional_xyz,
occupancy: &definition.occupancy,
u_iso_angstrom2: &definition.u_iso_angstrom2,
anisotropic_mask: &definition.anisotropic_mask,
u_aniso_cif_angstrom2: &definition.u_aniso_cif_angstrom2,
scattering_species: species,
scale: bank.scale,
coordinate_tolerance: definition.coordinate_tolerance,
correction_model: bank.correction_model,
bank_geometry: bank.geometry,
instrument: bank.instrument,
support_fwhm: input.support_fwhm,
tail_log: input.tail_log,
}
}
fn calculate_bank(
input: &StructuralTofMultiBankInput,
bank: &StructuralTofBank,
) -> Result<StructuralTofResult, StructuralTofMultiBankError> {
let species = species(input.phase.definition());
Ok(calculate_structural_tof_pattern_with_context(
input.phase.definition().cell,
&input.phase.definition().space_group,
&bank_view(input, bank, &species),
input.execution.context(),
)?)
}
fn background_values(
bank: &StructuralTofBank,
basis: Option<&crate::TofChebyshevBasis>,
) -> Result<Vec<f64>, StructuralTofMultiBankError> {
let mut values = if let (Some(model), Some(basis)) = (&bank.background, basis) {
model.calculate_from_basis(basis)?
} else {
vec![0.0; bank.pattern.sample_count()]
};
for (value, fixed) in values.iter_mut().zip(&bank.pattern.background_y) {
*value += fixed;
}
Ok(values)
}
fn add_instrument_jvp(
derivative: &mut [f64],
result: &StructuralTofResult,
mapping: &BankParameterMapping,
direction: &[f64],
) -> Result<(), StructuralTofMultiBankError> {
let global = result
.accumulation
.derivatives
.global
.as_ref()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
if global.parameter_count != TOF_GLOBAL_PARAMETER_COUNT
|| global.values.len() != TOF_GLOBAL_PARAMETER_COUNT * derivative.len()
{
return Err(StructuralTofMultiBankError::InternalInvariant);
}
for &(parameter, index) in &mapping.instrument {
let row = &global.values
[parameter.index() * derivative.len()..(parameter.index() + 1) * derivative.len()];
for (target, value) in derivative.iter_mut().zip(row) {
*target += direction[index] * value;
}
}
Ok(())
}
fn add_instrument_vjp(
output: &mut [f64],
result: &StructuralTofResult,
mapping: &BankParameterMapping,
weights: &[f64],
) -> Result<(), StructuralTofMultiBankError> {
let global = result
.accumulation
.derivatives
.global
.as_ref()
.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
if global.values.len() != TOF_GLOBAL_PARAMETER_COUNT * weights.len() {
return Err(StructuralTofMultiBankError::InternalInvariant);
}
for &(parameter, index) in &mapping.instrument {
let row = &global.values
[parameter.index() * weights.len()..(parameter.index() + 1) * weights.len()];
output[index] += row.iter().zip(weights).map(|(a, b)| a * b).sum::<f64>();
}
Ok(())
}
fn add_background_jvp(
derivative: &mut [f64],
basis: Option<&crate::TofChebyshevBasis>,
mapping: &BankParameterMapping,
direction: &[f64],
) -> Result<(), StructuralTofMultiBankError> {
if mapping.background.is_empty() {
return Ok(());
}
let basis = basis.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
if basis.rows != derivative.len() || basis.columns != mapping.background.len() {
return Err(StructuralTofMultiBankError::InternalInvariant);
}
for (sample, row) in basis.values.chunks_exact(basis.columns).enumerate() {
derivative[sample] += row
.iter()
.zip(&mapping.background)
.map(|(value, index)| value * direction[*index])
.sum::<f64>();
}
Ok(())
}
fn add_background_vjp(
output: &mut [f64],
basis: Option<&crate::TofChebyshevBasis>,
mapping: &BankParameterMapping,
weights: &[f64],
) -> Result<(), StructuralTofMultiBankError> {
if mapping.background.is_empty() {
return Ok(());
}
let basis = basis.ok_or(StructuralTofMultiBankError::InternalInvariant)?;
if basis.rows != weights.len() || basis.columns != mapping.background.len() {
return Err(StructuralTofMultiBankError::InternalInvariant);
}
for (sample, row) in basis.values.chunks_exact(basis.columns).enumerate() {
for (value, index) in row.iter().zip(&mapping.background) {
output[*index] += weights[sample] * value;
}
}
Ok(())
}
fn objective_weights(
bank: &StructuralTofBank,
metrics: &ResidualEvaluation,
use_uncertainty: bool,
) -> Vec<f64> {
(0..bank.pattern.sample_count())
.map(|sample| {
if !metrics.included[sample] {
0.0
} else if use_uncertainty {
let sigma = bank.pattern.uncertainty.as_ref().map_or(1.0, |v| v[sample]);
metrics.residual[sample] / (sigma * sigma)
} else {
metrics.residual[sample]
}
})
.collect()
}
fn weighted_direction(
bank: &StructuralTofBank,
direction: Vec<f64>,
use_uncertainty: bool,
) -> Vec<f64> {
direction
.into_iter()
.enumerate()
.map(|(sample, value)| {
if bank.pattern.mask.as_ref().is_some_and(|mask| !mask[sample]) {
0.0
} else if use_uncertainty {
let sigma = bank.pattern.uncertainty.as_ref().map_or(1.0, |v| v[sample]);
value / (sigma * sigma)
} else {
value
}
})
.collect()
}
const fn instrument_unit(parameter: TofInstrumentParameter) -> &'static str {
match parameter {
TofInstrumentParameter::Zero | TofInstrumentParameter::Z => "microsecond",
TofInstrumentParameter::Difc | TofInstrumentParameter::X => "microsecond/angstrom",
TofInstrumentParameter::Difa | TofInstrumentParameter::Y => "microsecond/angstrom^2",
TofInstrumentParameter::Difb => "microsecond*angstrom",
TofInstrumentParameter::Alpha => "microsecond^-1*angstrom",
TofInstrumentParameter::Beta0 => "microsecond^-1",
TofInstrumentParameter::Beta1 => "angstrom^4/microsecond",
TofInstrumentParameter::Betaq => "angstrom^2/microsecond",
TofInstrumentParameter::Sigma0 => "microsecond^2",
TofInstrumentParameter::Sigma1 => "microsecond^2/angstrom^2",
TofInstrumentParameter::Sigma2 => "microsecond^2/angstrom^4",
TofInstrumentParameter::Sigmaq => "microsecond^2/angstrom",
}
}
#[derive(Debug)]
pub enum StructuralTofMultiBankError {
TooFewBanks,
DuplicateBankId,
MissingObservations,
InvalidPhaseContract(&'static str),
InvalidBankContract(&'static str),
InvalidSupport,
BankContractMismatch,
ParameterLengthMismatch,
BankWeightCountMismatch,
InvalidDamping,
InternalInvariant,
Pattern(DomainError),
Rietveld(RietveldError),
StructuralParameters(RietveldParameterError),
Parameter(ParameterError),
Background(TofLeBailError),
Structural(StructuralTofError),
Tof(TofError),
Residual(ResidualError),
}
impl Display for StructuralTofMultiBankError {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
match self {
Self::TooFewBanks => formatter.write_str("structural TOF requires at least one bank"),
Self::DuplicateBankId => formatter.write_str("structural TOF bank IDs must be unique"),
Self::MissingObservations => {
formatter.write_str("every structural TOF bank requires observations")
}
Self::InvalidPhaseContract(reason) | Self::InvalidBankContract(reason) => {
formatter.write_str(reason)
}
Self::InvalidSupport => {
formatter.write_str("structural TOF support controls must be finite and positive")
}
Self::BankContractMismatch => formatter
.write_str("structural TOF bank contract changed under the prepared layout"),
Self::ParameterLengthMismatch => {
formatter.write_str("structural TOF parameter value/direction length is wrong")
}
Self::BankWeightCountMismatch => formatter
.write_str("structural TOF reverse products require one weight vector per bank"),
Self::InvalidDamping => {
formatter.write_str("structural TOF damping must be finite and non-negative")
}
Self::InternalInvariant => {
formatter.write_str("structural TOF internal shape invariant failed")
}
Self::Pattern(error) => Display::fmt(error, formatter),
Self::Rietveld(error) => Display::fmt(error, formatter),
Self::StructuralParameters(error) => Display::fmt(error, formatter),
Self::Parameter(error) => Display::fmt(error, formatter),
Self::Background(error) => Display::fmt(error, formatter),
Self::Structural(error) => Display::fmt(error, formatter),
Self::Tof(error) => Display::fmt(error, formatter),
Self::Residual(error) => Display::fmt(error, formatter),
}
}
}
impl Error for StructuralTofMultiBankError {}
impl From<DomainError> for StructuralTofMultiBankError {
fn from(value: DomainError) -> Self {
Self::Pattern(value)
}
}
impl From<RietveldError> for StructuralTofMultiBankError {
fn from(value: RietveldError) -> Self {
Self::Rietveld(value)
}
}
impl From<RietveldParameterError> for StructuralTofMultiBankError {
fn from(value: RietveldParameterError) -> Self {
Self::StructuralParameters(value)
}
}
impl From<ParameterError> for StructuralTofMultiBankError {
fn from(value: ParameterError) -> Self {
Self::Parameter(value)
}
}
impl From<TofLeBailError> for StructuralTofMultiBankError {
fn from(value: TofLeBailError) -> Self {
Self::Background(value)
}
}
impl From<StructuralTofError> for StructuralTofMultiBankError {
fn from(value: StructuralTofError) -> Self {
Self::Structural(value)
}
}
impl From<TofError> for StructuralTofMultiBankError {
fn from(value: TofError) -> Self {
Self::Tof(value)
}
}
impl From<ResidualError> for StructuralTofMultiBankError {
fn from(value: ResidualError) -> Self {
Self::Residual(value)
}
}