use crate::{BasisError, PenaltyMatrix};
use ndarray::{Array1, Array2, ArrayView1, ArrayViewMut1};
use std::any::Any;
use std::ops::Range;
use std::sync::Arc;
#[derive(Clone)]
pub struct CustomFamilyBlockPsiDerivative {
pub penalty_index: Option<usize>,
pub x_psi: Array2<f64>,
pub s_psi: Array2<f64>,
pub s_psi_components: Option<Vec<(usize, Array2<f64>)>>,
pub s_psi_penalty_components: Option<Vec<(usize, PenaltyMatrix)>>,
pub x_psi_psi: Option<Vec<Array2<f64>>>,
pub s_psi_psi: Option<Vec<Array2<f64>>>,
pub s_psi_psi_components: Option<Vec<Vec<(usize, Array2<f64>)>>>,
pub s_psi_psi_penalty_components: Option<Vec<Vec<(usize, PenaltyMatrix)>>>,
pub implicit_operator: Option<Arc<dyn CustomFamilyPsiDerivativeOperator>>,
pub implicit_axis: usize,
pub implicit_group_id: Option<usize>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CustomFamilyHyperAxis {
DesignPenalty {
block: usize,
derivative_index: usize,
},
Family {
family_axis: usize,
},
}
#[derive(Clone)]
pub struct CustomFamilyHyperLayout {
design_derivative_blocks: Vec<Vec<CustomFamilyBlockPsiDerivative>>,
family_axes: Vec<usize>,
values: Array1<f64>,
design_axis_count: usize,
axis_count: usize,
}
impl CustomFamilyHyperLayout {
pub fn new(
design_derivative_blocks: Vec<Vec<CustomFamilyBlockPsiDerivative>>,
family_axes: Vec<usize>,
values: Array1<f64>,
) -> Result<Self, String> {
for (expected, &actual) in family_axes.iter().enumerate() {
if actual != expected {
return Err(format!(
"custom-family hyper layout family axes must be contiguous and ordered: \
position {expected} carries family axis {actual}"
));
}
}
let design_axis_count =
design_derivative_blocks
.iter()
.try_fold(0usize, |count, derivatives| {
count.checked_add(derivatives.len()).ok_or_else(|| {
"custom-family hyper layout design-axis count exceeds usize".to_string()
})
})?;
let axis_count = design_axis_count
.checked_add(family_axes.len())
.ok_or_else(|| "custom-family hyper layout axis count exceeds usize".to_string())?;
if values.len() != axis_count {
return Err(format!(
"custom-family hyper layout value length mismatch: got {}, expected {axis_count}",
values.len()
));
}
if let Some((axis, value)) = values
.iter()
.copied()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
return Err(format!(
"custom-family hyper layout axis {axis} has non-finite value {value}"
));
}
Ok(Self {
design_derivative_blocks,
family_axes,
values,
design_axis_count,
axis_count,
})
}
pub fn block_count(&self) -> usize {
self.design_derivative_blocks.len()
}
pub fn design_axis_count(&self) -> usize {
self.design_axis_count
}
pub fn family_axis_count(&self) -> usize {
self.family_axes.len()
}
pub fn len(&self) -> usize {
self.axis_count
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn design_derivative_blocks(&self) -> &[Vec<CustomFamilyBlockPsiDerivative>] {
&self.design_derivative_blocks
}
pub fn values(&self) -> &Array1<f64> {
&self.values
}
pub fn axis(&self, global_index: usize) -> Option<CustomFamilyHyperAxis> {
if global_index < self.design_axis_count {
let mut remaining = global_index;
return self.design_derivative_blocks.iter().enumerate().find_map(
|(block, derivatives)| {
if remaining < derivatives.len() {
Some(CustomFamilyHyperAxis::DesignPenalty {
block,
derivative_index: remaining,
})
} else {
remaining -= derivatives.len();
None
}
},
);
}
let family_offset = global_index.checked_sub(self.design_axis_count)?;
self.family_axes
.get(family_offset)
.copied()
.map(|family_axis| CustomFamilyHyperAxis::Family { family_axis })
}
pub fn design_derivative(
&self,
global_index: usize,
) -> Option<(usize, usize, &CustomFamilyBlockPsiDerivative)> {
match self.axis(global_index)? {
CustomFamilyHyperAxis::DesignPenalty {
block,
derivative_index,
} => self
.design_derivative_blocks
.get(block)?
.get(derivative_index)
.map(|derivative| (block, derivative_index, derivative)),
CustomFamilyHyperAxis::Family { .. } => None,
}
}
pub fn family_axis(&self, global_index: usize) -> Option<usize> {
match self.axis(global_index)? {
CustomFamilyHyperAxis::Family { family_axis } => Some(family_axis),
CustomFamilyHyperAxis::DesignPenalty { .. } => None,
}
}
}
pub type SharedCustomFamilyHyperLayout = Arc<CustomFamilyHyperLayout>;
impl CustomFamilyBlockPsiDerivative {
pub fn new(
penalty_index: Option<usize>,
x_psi: Array2<f64>,
s_psi: Array2<f64>,
s_psi_components: Option<Vec<(usize, Array2<f64>)>>,
x_psi_psi: Option<Vec<Array2<f64>>>,
s_psi_psi: Option<Vec<Array2<f64>>>,
s_psi_psi_components: Option<Vec<Vec<(usize, Array2<f64>)>>>,
) -> Self {
Self {
penalty_index,
x_psi,
s_psi,
s_psi_components,
s_psi_penalty_components: None,
x_psi_psi,
s_psi_psi,
s_psi_psi_components,
s_psi_psi_penalty_components: None,
implicit_operator: None,
implicit_axis: 0,
implicit_group_id: None,
}
}
}
pub trait CustomFamilyPsiDerivativeOperator: Send + Sync + Any {
fn as_any(&self) -> &dyn Any;
fn n_data(&self) -> usize;
fn p_out(&self) -> usize;
fn transpose_mul(
&self,
axis: usize,
v: &ArrayView1<'_, f64>,
) -> Result<Array1<f64>, BasisError>;
fn forward_mul(&self, axis: usize, u: &ArrayView1<'_, f64>) -> Result<Array1<f64>, BasisError>;
fn transpose_mul_second_diag(
&self,
axis: usize,
v: &ArrayView1<'_, f64>,
) -> Result<Array1<f64>, BasisError>;
fn transpose_mul_second_cross(
&self,
axis_d: usize,
axis_e: usize,
v: &ArrayView1<'_, f64>,
) -> Result<Array1<f64>, BasisError>;
fn forward_mul_second_diag(
&self,
axis: usize,
u: &ArrayView1<'_, f64>,
) -> Result<Array1<f64>, BasisError>;
fn forward_mul_second_cross(
&self,
axis_d: usize,
axis_e: usize,
u: &ArrayView1<'_, f64>,
) -> Result<Array1<f64>, BasisError>;
fn row_chunk_first(&self, axis: usize, rows: Range<usize>) -> Result<Array2<f64>, BasisError>;
fn row_vector_first_into(
&self,
axis: usize,
row: usize,
mut out: ArrayViewMut1<'_, f64>,
) -> Result<(), BasisError> {
let chunk = self.row_chunk_first(axis, row..row + 1)?;
out.assign(&chunk.row(0));
Ok(())
}
fn row_chunk_second_diag(
&self,
axis: usize,
rows: Range<usize>,
) -> Result<Array2<f64>, BasisError>;
fn row_chunk_second_cross(
&self,
axis_d: usize,
axis_e: usize,
rows: Range<usize>,
) -> Result<Array2<f64>, BasisError>;
fn as_materializable(&self) -> Option<&dyn MaterializablePsiDerivativeOperator> {
None
}
}
pub trait MaterializablePsiDerivativeOperator: CustomFamilyPsiDerivativeOperator {
fn materialize_first(&self, axis: usize) -> Result<Array2<f64>, BasisError>;
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum JointHessianSourcePreference {
Dense,
Operator,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MaterializationIntent {
InnerSolve,
LogdetFactorization,
OuterEvaluation,
OuterGradient,
}