use nalgebra::{ArrayStorage, Const, SMatrix};
use std::fmt::Debug;
use super::modular_def::{ModularError, ModularForm, ModularTransformationGroup};
use crate::arithmetic_utils::Field;
#[allow(dead_code)]
pub struct SumModularForm<
const TWICE_WEIGHT: usize,
R: Field,
TRANSFORM: ModularTransformationGroup<R>,
> {
pub(crate) summands: Vec<(
Box<dyn ModularForm<TWICE_WEIGHT, R, TransformationGroup = TRANSFORM>>,
R,
)>,
}
impl<const TWICE_WEIGHT: usize, R, TRANSFORM: ModularTransformationGroup<R>>
SumModularForm<TWICE_WEIGHT, R, TRANSFORM>
where
R: Debug + Field + 'static,
{
#[allow(clippy::type_complexity)]
pub fn new_from_some_coeffs<const DIM_SPACE: usize>(
basis_of_space: [Box<dyn ModularForm<TWICE_WEIGHT, R, TransformationGroup = TRANSFORM>>;
DIM_SPACE],
constrained_coeffs_values: &[Result<(usize, R), (R, R)>],
) -> Result<Self, ModularError>
where
R: nalgebra::ComplexField,
{
if constrained_coeffs_values.len() < DIM_SPACE {
return Err(ModularError::NotEnoughConstraints);
}
let mut matrix = SMatrix::<R, DIM_SPACE, DIM_SPACE>::zeros();
let mut b_col_vector = SMatrix::<R, DIM_SPACE, 1>::zeros();
for idx in 0..DIM_SPACE {
for jdx in 0..DIM_SPACE {
matrix[(idx, jdx)] = match &constrained_coeffs_values[idx] {
Ok((which_coeff, _)) => basis_of_space[jdx].extract_coeffs(*which_coeff)?,
Err((q, _)) => basis_of_space[jdx].evaluate_at(q)?,
};
}
}
for idx in 0..DIM_SPACE {
b_col_vector[(idx, 0)] = match &constrained_coeffs_values[idx] {
Ok((_, value)) | Err((_, value)) => value.clone(),
};
}
let mut out = matrix.clone();
let succeeded =
nalgebra::try_invert_to::<R, Const<DIM_SPACE>, ArrayStorage<R, DIM_SPACE, DIM_SPACE>>(
matrix, &mut out,
);
if !succeeded {
return Err(ModularError::SingularSystem);
}
let w = out * b_col_vector;
let mut summands = Vec::with_capacity(DIM_SPACE);
for (idx, cur_summand) in basis_of_space.into_iter().enumerate() {
if w[idx].is_zero() {
continue;
}
if !w[idx].is_finite() {
return Err(ModularError::NonFiniteCoefficient);
}
let value = (cur_summand, w[idx].clone());
summands.push(value);
}
Ok(Self { summands })
}
#[allow(dead_code)]
pub fn extract_coeffs(&self, which_coeff: usize) -> Result<R, ModularError> {
let mut to_return = R::zero();
for summand in &self.summands {
to_return += summand.0.extract_coeffs(which_coeff)? * summand.1.clone();
}
Ok(to_return)
}
#[allow(dead_code)]
pub fn evaluate_at(&self, q: &R) -> Result<R, ModularError> {
let mut to_return = R::zero();
for summand in &self.summands {
to_return += summand.0.evaluate_at(q)? * summand.1.clone();
}
Ok(to_return)
}
}