use std::marker::PhantomData;
use crate::{DesignMatrix, Link, ModelError, RowMultiplier, Softplus, design::scale_active_rows};
const EXPECTED_FINITE: &str = "finite";
pub type SoftplusScalar = TransformedScalar<SoftplusTransform>;
pub type NegativeSoftplusScalar = TransformedScalar<NegativeSoftplusTransform>;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct LinearPredictorBlock<X> {
x: X,
}
impl<X> LinearPredictorBlock<X> {
#[must_use]
#[inline]
pub const fn new(x: X) -> Self {
Self { x }
}
#[must_use]
#[inline]
pub const fn x(&self) -> &X {
&self.x
}
#[must_use]
#[inline]
pub fn into_inner(self) -> X {
self.x
}
}
impl<X> PredictorBlock for LinearPredictorBlock<X>
where
X: DesignMatrix,
{
#[inline]
fn nrows(&self) -> usize {
self.x.nrows()
}
#[inline]
fn nparams(&self) -> usize {
self.x.ncols()
}
#[inline]
fn eta_row(&self, row: usize, beta: &[f64]) -> f64 {
self.x.dot_row(row, beta)
}
#[inline]
fn add_gradient(&self, scores: &[f64], _: &[f64], grad: &mut [f64]) {
self.x.add_t_mul_vec(scores, grad);
}
#[inline]
fn set_constant_start(&self, value: f64, beta: &mut [f64]) -> bool {
self.x.set_constant_start(value, beta)
}
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
_: &[f64],
grad: &mut [f64],
) {
self.x.add_weighted_t_mul_vec(scores, multiplier, grad);
}
#[inline]
fn add_weighted_gradient_by<M>(
&self,
scores: &[f64],
multiplier: &M,
_: &[f64],
grad: &mut [f64],
) where
M: RowMultiplier + ?Sized,
{
self.x.add_weighted_t_mul_vec_by(scores, multiplier, grad);
}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
Some(0.0)
}
}
impl<X: DesignMatrix> HasDesignMatrix for LinearPredictorBlock<X> {
type Matrix = X;
#[inline]
fn design(&self) -> &Self::Matrix {
&self.x
}
}
impl<X: DesignMatrix> LinearPredictorGeometry for LinearPredictorBlock<X> {
#[inline]
fn add_weighted_gram(&self, row_weights: &[f64], out: &mut [f64]) -> Result<(), ModelError> {
validate_geometry_lengths(self.x.nrows(), self.x.ncols(), row_weights, out)?;
self.x.gram_weighted(row_weights, out);
Ok(())
}
#[inline]
fn add_weighted_gram_by<M>(
&self,
row_weights: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
validate_geometry_lengths(self.x.nrows(), self.x.ncols(), row_weights, out)?;
self.x.gram_weighted_by(row_weights, multiplier, out);
Ok(())
}
#[inline]
fn add_t_mul_vec(&self, row_scores: &[f64], out: &mut [f64]) -> Result<(), ModelError> {
validate_vector_geometry_lengths(self.x.nrows(), self.x.ncols(), row_scores, out)?;
self.x.add_t_mul_vec(row_scores, out);
Ok(())
}
#[inline]
fn add_t_mul_vec_by<M>(
&self,
row_scores: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
validate_vector_geometry_lengths(self.x.nrows(), self.x.ncols(), row_scores, out)?;
self.x
.add_weighted_t_mul_vec_by(row_scores, multiplier, out);
Ok(())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct SoftplusTransform;
impl CoefficientTransform for SoftplusTransform {
#[inline]
fn value(beta: f64) -> f64 {
Softplus::inverse(beta)
}
#[inline]
fn derivative(beta: f64) -> f64 {
Softplus::derivative_inverse(beta)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct NegativeSoftplusTransform;
impl CoefficientTransform for NegativeSoftplusTransform {
#[inline]
fn value(beta: f64) -> f64 {
-Softplus::inverse(beta)
}
#[inline]
fn derivative(beta: f64) -> f64 {
-Softplus::derivative_inverse(beta)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TransformedScalar<T> {
nrows: usize,
marker: PhantomData<T>,
}
impl<T> TransformedScalar<T> {
#[must_use]
#[inline]
pub const fn new(nrows: usize) -> Self {
Self {
nrows,
marker: PhantomData,
}
}
#[must_use]
#[inline]
pub const fn nrows(&self) -> usize {
self.nrows
}
}
impl<T> PredictorBlock for TransformedScalar<T>
where
T: CoefficientTransform,
{
#[inline]
fn nrows(&self) -> usize {
self.nrows
}
#[inline]
fn nparams(&self) -> usize {
1
}
#[inline]
fn eta_row(&self, _: usize, beta: &[f64]) -> f64 {
T::value(beta[0])
}
#[inline]
fn add_gradient(&self, scores: &[f64], beta: &[f64], grad: &mut [f64]) {
debug_assert_eq!(scores.len(), self.nrows);
debug_assert_eq!(beta.len(), 1);
debug_assert_eq!(grad.len(), 1);
grad[0] = scores
.iter()
.sum::<f64>()
.mul_add(T::derivative(beta[0]), grad[0]);
}
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
beta: &[f64],
grad: &mut [f64],
) {
debug_assert_eq!(scores.len(), self.nrows);
debug_assert_eq!(multiplier.len(), self.nrows);
debug_assert_eq!(beta.len(), 1);
debug_assert_eq!(grad.len(), 1);
grad[0] = weighted_sum(scores, multiplier).mul_add(T::derivative(beta[0]), grad[0]);
}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
let value = T::value(0.0);
value.is_finite().then_some(value)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct FloorSoftplusScalar {
nrows: usize,
floor: f64,
}
impl FloorSoftplusScalar {
#[must_use]
#[inline]
pub const fn new(nrows: usize, floor: f64) -> Self {
Self { nrows, floor }
}
#[inline]
pub fn try_new(nrows: usize, floor: f64) -> Result<Self, ModelError> {
validate_finite("floor", floor)?;
Ok(Self::new(nrows, floor))
}
#[must_use]
#[inline]
pub const fn nrows(&self) -> usize {
self.nrows
}
#[must_use]
#[inline]
pub const fn floor(&self) -> f64 {
self.floor
}
}
impl PredictorBlock for FloorSoftplusScalar {
#[inline]
fn nrows(&self) -> usize {
self.nrows
}
#[inline]
fn nparams(&self) -> usize {
1
}
#[inline]
fn eta_row(&self, _: usize, beta: &[f64]) -> f64 {
self.floor + Softplus::inverse(beta[0])
}
#[inline]
fn add_gradient(&self, scores: &[f64], beta: &[f64], grad: &mut [f64]) {
debug_assert_eq!(scores.len(), self.nrows);
debug_assert_eq!(beta.len(), 1);
debug_assert_eq!(grad.len(), 1);
grad[0] = scores
.iter()
.sum::<f64>()
.mul_add(Softplus::derivative_inverse(beta[0]), grad[0]);
}
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
beta: &[f64],
grad: &mut [f64],
) {
debug_assert_eq!(scores.len(), self.nrows);
debug_assert_eq!(multiplier.len(), self.nrows);
debug_assert_eq!(beta.len(), 1);
debug_assert_eq!(grad.len(), 1);
grad[0] = weighted_sum(scores, multiplier)
.mul_add(Softplus::derivative_inverse(beta[0]), grad[0]);
}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
let value = self.floor + Softplus::inverse(0.0);
value.is_finite().then_some(value)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct OffsetBlock {
nrows: usize,
value: f64,
}
impl OffsetBlock {
#[must_use]
#[inline]
pub const fn new(nrows: usize, value: f64) -> Self {
Self { nrows, value }
}
#[inline]
pub fn try_new(nrows: usize, value: f64) -> Result<Self, ModelError> {
validate_finite("offset value", value)?;
Ok(Self::new(nrows, value))
}
#[must_use]
#[inline]
pub const fn nrows(&self) -> usize {
self.nrows
}
#[must_use]
#[inline]
pub const fn value(&self) -> f64 {
self.value
}
}
impl PredictorBlock for OffsetBlock {
#[inline]
fn nrows(&self) -> usize {
self.nrows
}
#[inline]
fn nparams(&self) -> usize {
0
}
#[inline]
fn eta_row(&self, _: usize, _: &[f64]) -> f64 {
self.value
}
#[inline]
fn add_gradient(&self, _: &[f64], _: &[f64], _: &mut [f64]) {}
#[inline]
fn add_weighted_gradient(&self, _: &[f64], _: &[f64], _: &[f64], _: &mut [f64]) {}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
self.value.is_finite().then_some(self.value)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ProductBlock<X> {
multiplier: Vec<f64>,
inner: X,
}
impl<X> ProductBlock<X> {
#[must_use]
#[inline]
pub const fn new_unchecked(multiplier: Vec<f64>, inner: X) -> Self {
Self { multiplier, inner }
}
#[must_use]
#[inline]
pub fn multiplier(&self) -> &[f64] {
&self.multiplier
}
#[must_use]
#[inline]
pub const fn inner(&self) -> &X {
&self.inner
}
#[must_use]
#[inline]
pub fn into_inner(self) -> X {
self.inner
}
#[must_use]
#[inline]
pub fn into_parts(self) -> (Vec<f64>, X) {
(self.multiplier, self.inner)
}
}
impl<X> ProductBlock<X>
where
X: PredictorBlock,
{
#[inline]
pub fn try_new(multiplier: Vec<f64>, inner: X) -> Result<Self, ModelError> {
let block = Self::new_unchecked(multiplier, inner);
block.inner.validate()?;
block.validate_multiplier()?;
Ok(block)
}
#[inline]
fn validate_multiplier(&self) -> Result<(), ModelError> {
if self.multiplier.len() != self.inner.nrows() {
return Err(ModelError::DesignRowMismatch {
parameter: "product multiplier",
expected_rows: self.inner.nrows(),
actual_rows: self.multiplier.len(),
});
}
for (index, value) in self.multiplier.iter().copied().enumerate() {
if !value.is_finite() {
return Err(ModelError::InvalidMultiplier { index });
}
}
Ok(())
}
}
impl<X> ProductBlock<X>
where
X: LinearPredictorGeometry,
{
#[inline]
fn validate_geometry_outer_lengths(
&self,
row_weights: &[f64],
out: &[f64],
) -> Result<(), ModelError> {
self.validate_multiplier()?;
validate_geometry_lengths(self.inner.nrows(), self.inner.nparams(), row_weights, out)
}
#[inline]
fn validate_vector_geometry_outer_lengths(
&self,
row_scores: &[f64],
out: &[f64],
) -> Result<(), ModelError> {
self.validate_multiplier()?;
validate_vector_geometry_lengths(self.inner.nrows(), self.inner.nparams(), row_scores, out)
}
}
impl<X> PredictorBlock for ProductBlock<X>
where
X: PredictorBlock,
{
#[inline]
fn nrows(&self) -> usize {
self.inner.nrows()
}
#[inline]
fn nparams(&self) -> usize {
self.inner.nparams()
}
#[inline]
fn eta_row(&self, row: usize, beta: &[f64]) -> f64 {
self.multiplier[row] * self.inner.eta_row(row, beta)
}
#[inline]
fn add_gradient(&self, scores: &[f64], beta: &[f64], grad: &mut [f64]) {
debug_assert_eq!(scores.len(), self.nrows());
debug_assert_eq!(self.multiplier.len(), self.nrows());
self.inner
.add_weighted_gradient(scores, &self.multiplier, beta, grad);
}
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
beta: &[f64],
grad: &mut [f64],
) {
debug_assert_eq!(scores.len(), self.nrows());
debug_assert_eq!(multiplier.len(), self.nrows());
debug_assert_eq!(self.multiplier.len(), self.nrows());
self.add_weighted_gradient_by(scores, multiplier, beta, grad);
}
#[inline]
fn add_weighted_gradient_by<M>(
&self,
scores: &[f64],
multiplier: &M,
beta: &[f64],
grad: &mut [f64],
) where
M: RowMultiplier + ?Sized,
{
debug_assert_eq!(scores.len(), self.nrows());
debug_assert_eq!(self.multiplier.len(), self.nrows());
let product_multiplier = ProductRowMultiplier {
left: self.multiplier.as_slice(),
right: multiplier,
};
self.inner
.add_weighted_gradient_by(scores, &product_multiplier, beta, grad);
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
self.inner.validate()?;
self.validate_multiplier()
}
#[inline]
#[allow(clippy::float_cmp)]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
let inner = self.inner.zero_beta_constant_contribution()?;
if inner == 0.0 {
return Some(0.0);
}
let first = self.multiplier.first().copied()?;
if !first.is_finite() || !self.multiplier.iter().all(|value| *value == first) {
return None;
}
let value = first * inner;
value.is_finite().then_some(value)
}
}
impl<X> LinearPredictorGeometry for ProductBlock<X>
where
X: LinearPredictorGeometry,
{
#[inline]
fn add_weighted_gram(&self, row_weights: &[f64], out: &mut [f64]) -> Result<(), ModelError> {
self.add_weighted_gram_by(row_weights, &UnitRowMultiplier, out)
}
#[inline]
fn add_weighted_gram_by<M>(
&self,
row_weights: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
self.validate_geometry_outer_lengths(row_weights, out)?;
let product_multiplier = ProductSquaredRowMultiplier {
product: self.multiplier.as_slice(),
right: multiplier,
};
self.inner
.add_weighted_gram_by(row_weights, &product_multiplier, out)
}
#[inline]
fn add_t_mul_vec(&self, row_scores: &[f64], out: &mut [f64]) -> Result<(), ModelError> {
self.add_t_mul_vec_by(row_scores, &UnitRowMultiplier, out)
}
#[inline]
fn add_t_mul_vec_by<M>(
&self,
row_scores: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
self.validate_vector_geometry_outer_lengths(row_scores, out)?;
let product_multiplier = ProductRowMultiplier {
left: self.multiplier.as_slice(),
right: multiplier,
};
self.inner
.add_t_mul_vec_by(row_scores, &product_multiplier, out)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct SumBlock<Terms> {
pub terms: Terms,
}
impl<Terms> SumBlock<Terms> {
#[must_use]
#[inline]
pub const fn new(terms: Terms) -> Self {
Self { terms }
}
}
struct ProductRowMultiplier<'a, M>
where
M: RowMultiplier + ?Sized,
{
left: &'a [f64],
right: &'a M,
}
impl<M> RowMultiplier for ProductRowMultiplier<'_, M>
where
M: RowMultiplier + ?Sized,
{
#[inline]
fn multiplier_at(&self, row: usize) -> f64 {
self.left[row] * self.right.multiplier_at(row)
}
}
struct ProductSquaredRowMultiplier<'a, M>
where
M: RowMultiplier + ?Sized,
{
product: &'a [f64],
right: &'a M,
}
impl<M> RowMultiplier for ProductSquaredRowMultiplier<'_, M>
where
M: RowMultiplier + ?Sized,
{
#[inline]
fn multiplier_at(&self, row: usize) -> f64 {
self.product[row] * self.product[row] * self.right.multiplier_at(row)
}
}
struct UnitRowMultiplier;
impl RowMultiplier for UnitRowMultiplier {
#[inline]
fn multiplier_at(&self, _: usize) -> f64 {
1.0
}
}
pub trait PredictorBlock {
fn nrows(&self) -> usize;
fn nparams(&self) -> usize;
fn eta_row(&self, row: usize, beta: &[f64]) -> f64;
#[inline]
fn set_constant_start(&self, _value: f64, _beta: &mut [f64]) -> bool {
false
}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
None
}
fn add_gradient(&self, scores: &[f64], beta: &[f64], grad: &mut [f64]);
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
beta: &[f64],
grad: &mut [f64],
) {
self.add_weighted_gradient_by(scores, multiplier, beta, grad);
}
#[inline]
fn add_weighted_gradient_by<M>(
&self,
scores: &[f64],
multiplier: &M,
beta: &[f64],
grad: &mut [f64],
) where
M: RowMultiplier + ?Sized,
{
debug_assert_eq!(scores.len(), self.nrows());
let scaled_scores = scale_active_rows(scores, multiplier);
self.add_gradient(&scaled_scores, beta, grad);
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
Ok(())
}
}
pub trait HasDesignMatrix: PredictorBlock {
type Matrix: DesignMatrix;
fn design(&self) -> &Self::Matrix;
}
pub trait LinearPredictorGeometry: PredictorBlock {
fn add_weighted_gram(&self, row_weights: &[f64], out: &mut [f64]) -> Result<(), ModelError>;
#[inline]
fn add_weighted_gram_by<M>(
&self,
row_weights: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
let scaled_weights = scale_active_rows(row_weights, multiplier);
self.add_weighted_gram(&scaled_weights, out)
}
fn add_t_mul_vec(&self, row_scores: &[f64], out: &mut [f64]) -> Result<(), ModelError>;
#[inline]
fn add_t_mul_vec_by<M>(
&self,
row_scores: &[f64],
multiplier: &M,
out: &mut [f64],
) -> Result<(), ModelError>
where
M: RowMultiplier + ?Sized,
{
let scaled_scores = scale_active_rows(row_scores, multiplier);
self.add_t_mul_vec(&scaled_scores, out)
}
}
pub trait CoefficientTransform {
fn value(beta: f64) -> f64;
fn derivative(beta: f64) -> f64;
}
#[inline]
fn weighted_sum(scores: &[f64], multiplier: &[f64]) -> f64 {
scores
.iter()
.copied()
.enumerate()
.map(|(row, score)| {
if score == 0.0 {
0.0
} else {
score * multiplier[row]
}
})
.sum()
}
const fn validate_finite(parameter: &'static str, value: f64) -> Result<(), ModelError> {
if value.is_finite() {
Ok(())
} else {
Err(ModelError::InvalidParameter {
parameter,
expected: EXPECTED_FINITE,
})
}
}
fn validate_geometry_lengths(
nrows: usize,
nparams: usize,
row_weights: &[f64],
out: &[f64],
) -> Result<(), ModelError> {
validate_row_values_len(nrows, row_weights)?;
let expected_values = nparams
.checked_mul(nparams)
.ok_or(ModelError::ArithmeticOverflow {
context: "linear predictor geometry Gram value count",
})?;
if out.len() != expected_values {
return Err(ModelError::DesignSize {
expected_values,
actual_values: out.len(),
});
}
Ok(())
}
fn validate_vector_geometry_lengths(
nrows: usize,
nparams: usize,
row_scores: &[f64],
out: &[f64],
) -> Result<(), ModelError> {
validate_row_values_len(nrows, row_scores)?;
if out.len() != nparams {
return Err(ModelError::GradientLength {
expected: nparams,
actual: out.len(),
});
}
Ok(())
}
const fn validate_row_values_len(nrows: usize, row_values: &[f64]) -> Result<(), ModelError> {
if row_values.len() != nrows {
return Err(ModelError::WeightLength {
expected: nrows,
actual: row_values.len(),
});
}
Ok(())
}
macro_rules! impl_sum_block {
(
terms = ($($term:ident),+);
vars = ($($var:ident),+);
indices = ($($idx:tt),+);
names = ($($name:literal),+)
) => {
impl<$($term,)+> PredictorBlock for SumBlock<($($term,)+)>
where
$($term: PredictorBlock,)+
{
#[inline]
fn nrows(&self) -> usize {
self.terms.0.nrows()
}
#[inline]
fn nparams(&self) -> usize {
0 $(+ self.terms.$idx.nparams())+
}
#[inline]
fn eta_row(&self, row: usize, beta: &[f64]) -> f64 {
let mut start = 0;
let mut eta = 0.0;
$(
let $var = &self.terms.$idx;
let end = start + $var.nparams();
eta += $var.eta_row(row, &beta[start..end]);
start = end;
)+
let _ = start;
eta
}
#[inline]
fn add_gradient(&self, scores: &[f64], beta: &[f64], grad: &mut [f64]) {
let mut start = 0;
$(
let $var = &self.terms.$idx;
let end = start + $var.nparams();
$var.add_gradient(scores, &beta[start..end], &mut grad[start..end]);
start = end;
)+
let _ = start;
}
#[inline]
fn set_constant_start(&self, value: f64, beta: &mut [f64]) -> bool {
let baselines = [$(self.terms.$idx.zero_beta_constant_contribution(),)+];
let mut start = 0;
$(
let $var = &self.terms.$idx;
let end = start + $var.nparams();
let other_baseline = baselines
.iter()
.enumerate()
.filter(|(index, _)| *index != $idx)
.try_fold(0.0, |sum, (_, baseline)| baseline.map(|value| sum + value));
if let Some(other_baseline) = other_baseline {
if $var.set_constant_start(value - other_baseline, &mut beta[start..end]) {
return true;
}
}
start = end;
)+
let _ = start;
false
}
#[inline]
fn zero_beta_constant_contribution(&self) -> Option<f64> {
let mut contribution = 0.0;
$(
contribution += self.terms.$idx.zero_beta_constant_contribution()?;
)+
contribution.is_finite().then_some(contribution)
}
#[inline]
fn add_weighted_gradient(
&self,
scores: &[f64],
multiplier: &[f64],
beta: &[f64],
grad: &mut [f64],
) {
debug_assert_eq!(scores.len(), self.nrows());
debug_assert_eq!(multiplier.len(), self.nrows());
debug_assert_eq!(beta.len(), self.nparams());
debug_assert_eq!(grad.len(), self.nparams());
self.add_weighted_gradient_by(scores, multiplier, beta, grad);
}
#[inline]
fn add_weighted_gradient_by<M>(
&self,
scores: &[f64],
multiplier: &M,
beta: &[f64],
grad: &mut [f64],
) where
M: RowMultiplier + ?Sized,
{
debug_assert_eq!(scores.len(), self.nrows());
debug_assert_eq!(beta.len(), self.nparams());
debug_assert_eq!(grad.len(), self.nparams());
let mut start = 0;
$(
let $var = &self.terms.$idx;
let end = start + $var.nparams();
$var.add_weighted_gradient_by(
scores,
multiplier,
&beta[start..end],
&mut grad[start..end],
);
start = end;
)+
let _ = start;
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
let expected_rows = self.terms.0.nrows();
$(
self.terms.$idx.validate()?;
if self.terms.$idx.nrows() != expected_rows {
return Err(ModelError::DesignRowMismatch {
parameter: $name,
expected_rows,
actual_rows: self.terms.$idx.nrows(),
});
}
)+
Ok(())
}
}
};
}
impl_sum_block!(
terms = (T1);
vars = (term1);
indices = (0);
names = ("sum term")
);
impl_sum_block!(
terms = (T1, T2);
vars = (term1, term2);
indices = (0, 1);
names = ("sum first term", "sum second term")
);
impl_sum_block!(
terms = (T1, T2, T3);
vars = (term1, term2, term3);
indices = (0, 1, 2);
names = ("sum first term", "sum second term", "sum third term")
);
impl_sum_block!(
terms = (T1, T2, T3, T4);
vars = (term1, term2, term3, term4);
indices = (0, 1, 2, 3);
names = (
"sum first term",
"sum second term",
"sum third term",
"sum fourth term"
)
);
impl_sum_block!(
terms = (T1, T2, T3, T4, T5);
vars = (term1, term2, term3, term4, term5);
indices = (0, 1, 2, 3, 4);
names = (
"sum first term",
"sum second term",
"sum third term",
"sum fourth term",
"sum fifth term"
)
);
impl_sum_block!(
terms = (T1, T2, T3, T4, T5, T6);
vars = (term1, term2, term3, term4, term5, term6);
indices = (0, 1, 2, 3, 4, 5);
names = (
"sum first term",
"sum second term",
"sum third term",
"sum fourth term",
"sum fifth term",
"sum sixth term"
)
);
impl_sum_block!(
terms = (T1, T2, T3, T4, T5, T6, T7);
vars = (term1, term2, term3, term4, term5, term6, term7);
indices = (0, 1, 2, 3, 4, 5, 6);
names = (
"sum first term",
"sum second term",
"sum third term",
"sum fourth term",
"sum fifth term",
"sum sixth term",
"sum seventh term"
)
);
impl_sum_block!(
terms = (T1, T2, T3, T4, T5, T6, T7, T8);
vars = (term1, term2, term3, term4, term5, term6, term7, term8);
indices = (0, 1, 2, 3, 4, 5, 6, 7);
names = (
"sum first term",
"sum second term",
"sum third term",
"sum fourth term",
"sum fifth term",
"sum sixth term",
"sum seventh term",
"sum eighth term"
)
);
#[cfg(test)]
mod tests {
use approx::assert_relative_eq;
use crate::{
DenseDesign, DesignMatrix, LinearPredictorGeometry, ModelError, PredictorBlock,
RowMultiplier,
};
use super::{
FloorSoftplusScalar, LinearPredictorBlock, NegativeSoftplusScalar, OffsetBlock,
ProductBlock, SoftplusScalar,
};
struct DefaultPaths;
impl PredictorBlock for DefaultPaths {
fn nrows(&self) -> usize {
3
}
fn nparams(&self) -> usize {
1
}
fn eta_row(&self, _: usize, beta: &[f64]) -> f64 {
beta[0]
}
fn add_gradient(&self, scores: &[f64], _: &[f64], grad: &mut [f64]) {
grad[0] += scores.iter().sum::<f64>();
}
}
impl LinearPredictorGeometry for DefaultPaths {
fn add_weighted_gram(
&self,
row_weights: &[f64],
out: &mut [f64],
) -> Result<(), ModelError> {
out[0] += row_weights.iter().sum::<f64>();
Ok(())
}
fn add_t_mul_vec(&self, row_scores: &[f64], out: &mut [f64]) -> Result<(), ModelError> {
out[0] += row_scores.iter().sum::<f64>();
Ok(())
}
}
struct PanicOnMaskedRow;
impl RowMultiplier for PanicOnMaskedRow {
fn multiplier_at(&self, row: usize) -> f64 {
assert_ne!(row, 1, "zero row must not read multiplier");
2.0
}
}
struct LazyGradientBlock;
impl PredictorBlock for LazyGradientBlock {
fn nrows(&self) -> usize {
3
}
fn nparams(&self) -> usize {
1
}
fn eta_row(&self, _: usize, beta: &[f64]) -> f64 {
beta[0]
}
fn add_gradient(&self, _: &[f64], _: &[f64], _: &mut [f64]) {
panic!("composed lazy path must not materialize scores");
}
fn add_weighted_gradient_by<M>(
&self,
scores: &[f64],
multiplier: &M,
_: &[f64],
grad: &mut [f64],
) where
M: RowMultiplier + ?Sized,
{
for (row, score) in scores.iter().copied().enumerate() {
if score != 0.0 {
grad[0] = score.mul_add(multiplier.multiplier_at(row), grad[0]);
}
}
}
}
#[test]
fn linear_predictor_block_matches_design_matrix_operations() {
let design = DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]);
let block = LinearPredictorBlock::new(design);
let beta = [10.0, 1.0];
assert_eq!(block.x().nrows(), 2);
assert_relative_eq!(block.eta_row(1, &beta), 34.0);
let mut grad = vec![0.0, 0.0];
block.add_gradient(&[0.5, 2.0], &beta, &mut grad);
assert_relative_eq!(grad[0], 6.5);
assert_relative_eq!(grad[1], 9.0);
}
#[test]
fn linear_predictor_block_fuses_weighted_gradient() {
let design = DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]);
let block = LinearPredictorBlock::new(design);
let beta = [10.0, 1.0];
let mut grad = vec![1.0, 1.0];
block.add_weighted_gradient(&[0.5, 2.0], &[2.0, -1.0], &beta, &mut grad);
assert_relative_eq!(grad[0], -4.0);
assert_relative_eq!(grad[1], -5.0);
}
#[test]
fn linear_predictor_geometry_delegates_dense_products() {
let block = LinearPredictorBlock::new(DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]));
assert_eq!(block.nrows(), 2);
assert_eq!(block.nparams(), 2);
let mut gram = vec![1.0, 2.0, 3.0, 4.0];
block.add_weighted_gram(&[0.5, 2.0], &mut gram).unwrap();
assert_relative_eq!(gram[0], 19.5);
assert_relative_eq!(gram[1], 27.0);
assert_relative_eq!(gram[2], 28.0);
assert_relative_eq!(gram[3], 38.0);
let mut t_mul = vec![1.0, 1.0];
block.add_t_mul_vec(&[0.5, 2.0], &mut t_mul).unwrap();
assert_relative_eq!(t_mul[0], 7.5);
assert_relative_eq!(t_mul[1], 10.0);
}
#[test]
fn linear_predictor_geometry_validates_lengths() {
let block = LinearPredictorBlock::new(DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]));
assert_eq!(
block.add_weighted_gram(&[1.0], &mut [0.0; 4]).unwrap_err(),
ModelError::WeightLength {
expected: 2,
actual: 1,
}
);
assert_eq!(
block
.add_weighted_gram(&[1.0, 1.0], &mut [0.0; 3])
.unwrap_err(),
ModelError::DesignSize {
expected_values: 4,
actual_values: 3,
}
);
assert_eq!(
block.add_t_mul_vec(&[1.0, 1.0], &mut [0.0]).unwrap_err(),
ModelError::GradientLength {
expected: 2,
actual: 1,
}
);
}
#[test]
fn sum_block_supports_eight_terms() {
let terms = (
LinearPredictorBlock::new(DenseDesign::column(&[1.0, 2.0])),
LinearPredictorBlock::new(DenseDesign::column(&[2.0, 3.0])),
LinearPredictorBlock::new(DenseDesign::column(&[3.0, 4.0])),
LinearPredictorBlock::new(DenseDesign::column(&[4.0, 5.0])),
LinearPredictorBlock::new(DenseDesign::column(&[5.0, 6.0])),
LinearPredictorBlock::new(DenseDesign::column(&[6.0, 7.0])),
LinearPredictorBlock::new(DenseDesign::column(&[7.0, 8.0])),
LinearPredictorBlock::new(DenseDesign::column(&[8.0, 9.0])),
);
let block = crate::SumBlock::new(terms);
let beta = [1.0; 8];
assert_eq!(block.nparams(), 8);
assert_relative_eq!(block.eta_row(1, &beta), 44.0);
let mut grad = vec![0.0; 8];
block.add_gradient(&[0.5, 2.0], &beta, &mut grad);
assert_relative_eq!(grad[0], 4.5);
assert_relative_eq!(grad[7], 22.0);
}
#[test]
fn transformed_scalar_blocks_match_finite_difference() {
let softplus = SoftplusScalar::new(3);
assert_eq!(softplus.nrows(), 3);
assert_scalar_gradient_matches_finite_difference(softplus, &[0.5, 1.0, 2.0]);
assert_scalar_gradient_matches_finite_difference(
NegativeSoftplusScalar::new(3),
&[0.5, 1.0, 2.0],
);
let floored = FloorSoftplusScalar::try_new(3, 10.0).unwrap();
assert_eq!(floored.nrows(), 3);
assert_relative_eq!(floored.floor(), 10.0);
assert_scalar_gradient_matches_finite_difference(floored, &[0.5, 1.0, 2.0]);
}
#[test]
fn floor_softplus_scalar_try_new_validates_floor() {
assert_eq!(
FloorSoftplusScalar::try_new(2, f64::NAN).unwrap_err(),
ModelError::InvalidParameter {
parameter: "floor",
expected: "finite",
}
);
}
#[allow(clippy::needless_pass_by_value)]
fn assert_scalar_gradient_matches_finite_difference(
block: impl PredictorBlock,
scores: &[f64],
) {
let beta = [0.4];
let eps = 1.0e-6;
let mut grad = [0.0];
block.add_gradient(scores, &beta, &mut grad);
let mut finite_difference = 0.0;
for (row, score) in scores.iter().copied().enumerate() {
let plus = block.eta_row(row, &[beta[0] + eps]);
let minus = block.eta_row(row, &[beta[0] - eps]);
finite_difference += score * (plus - minus) / (2.0 * eps);
}
assert_relative_eq!(grad[0], finite_difference, epsilon = 1.0e-6);
}
#[test]
fn offset_block_is_constant_and_has_no_gradient() {
let block = OffsetBlock::try_new(2, 3.5).unwrap();
let mut grad = [];
assert_eq!(block.nrows(), 2);
assert_relative_eq!(block.value(), 3.5);
assert_eq!(block.nparams(), 0);
assert_relative_eq!(block.eta_row(1, &[]), 3.5);
block.add_gradient(&[1.0, 2.0], &[], &mut grad);
}
#[test]
fn offset_block_try_new_validates_value() {
assert_eq!(
OffsetBlock::try_new(2, f64::INFINITY).unwrap_err(),
ModelError::InvalidParameter {
parameter: "offset value",
expected: "finite",
}
);
}
#[test]
fn product_block_scales_eta_and_gradient() {
let inner = LinearPredictorBlock::new(DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]));
let block = ProductBlock::try_new(vec![2.0, -1.0], inner).unwrap();
let beta = [0.5, 1.0];
let scores = [0.25, 2.0];
let mut grad = [0.0, 0.0];
assert_eq!(block.multiplier(), &[2.0, -1.0]);
assert_eq!(block.inner().nparams(), 2);
assert_relative_eq!(block.eta_row(0, &beta), 5.0);
assert_relative_eq!(block.eta_row(1, &beta), -5.5);
block.add_gradient(&scores, &beta, &mut grad);
assert_relative_eq!(grad[0], -5.5);
assert_relative_eq!(grad[1], -7.0);
}
#[test]
fn product_block_geometry_scales_rows_lazily() {
let inner = LinearPredictorBlock::new(DenseDesign::from_rows(&[[1.0, 2.0], [3.0, 4.0]]));
let block = ProductBlock::try_new(vec![2.0, -1.0], inner).unwrap();
let mut gram = vec![0.0; 4];
block.add_weighted_gram(&[0.5, 2.0], &mut gram).unwrap();
assert_relative_eq!(gram[0], 20.0);
assert_relative_eq!(gram[1], 28.0);
assert_relative_eq!(gram[2], 28.0);
assert_relative_eq!(gram[3], 40.0);
let mut t_mul = vec![0.0, 0.0];
block.add_t_mul_vec(&[0.25, 2.0], &mut t_mul).unwrap();
assert_relative_eq!(t_mul[0], -5.5);
assert_relative_eq!(t_mul[1], -7.0);
}
#[test]
fn product_block_validates_multiplier_length() {
let inner = LinearPredictorBlock::new(DenseDesign::intercept(2));
let block = ProductBlock::new_unchecked(vec![1.0], inner);
assert_multiplier_length_error(block.validate().unwrap_err());
}
#[test]
fn product_block_try_new_validates_multiplier_length() {
let inner = LinearPredictorBlock::new(DenseDesign::intercept(2));
assert_multiplier_length_error(ProductBlock::try_new(vec![1.0], inner).unwrap_err());
}
#[test]
fn product_block_validates_multiplier_finiteness() {
let inner = LinearPredictorBlock::new(DenseDesign::intercept(2));
let block = ProductBlock::new_unchecked(vec![1.0, f64::INFINITY], inner);
assert_invalid_multiplier_error(block.validate().unwrap_err());
}
#[test]
fn product_block_try_new_validates_multiplier_finiteness() {
let inner = LinearPredictorBlock::new(DenseDesign::intercept(2));
assert_invalid_multiplier_error(
ProductBlock::try_new(vec![1.0, f64::INFINITY], inner).unwrap_err(),
);
}
#[test]
fn predictor_gradient_default_skips_zero_score_multiplier() {
let block = DefaultPaths;
let values = [1.0, 0.0, 3.0];
let mut gradient = [0.0];
block.add_weighted_gradient_by(&values, &PanicOnMaskedRow, &[], &mut gradient);
assert_relative_eq!(gradient[0], 8.0);
}
#[test]
fn product_of_sum_keeps_weighted_gradient_multiplier_lazy() {
let sum = crate::SumBlock::new((LazyGradientBlock, LazyGradientBlock));
let block = ProductBlock::new_unchecked(vec![2.0, f64::NAN, 4.0], sum);
let mut gradient = [0.0, 0.0];
block.add_weighted_gradient_by(
&[1.0, 0.0, 3.0],
&PanicOnMaskedRow,
&[0.0, 0.0],
&mut gradient,
);
assert_relative_eq!(gradient[0], 28.0);
assert_relative_eq!(gradient[1], 28.0);
}
#[test]
fn product_block_geometry_validates_multiplier_finiteness() {
let inner = LinearPredictorBlock::new(DenseDesign::intercept(2));
let block = ProductBlock::new_unchecked(vec![1.0, f64::NAN], inner);
assert_invalid_multiplier_error(
block
.add_weighted_gram(&[1.0, 0.0], &mut [0.0])
.unwrap_err(),
);
assert_invalid_multiplier_error(block.add_t_mul_vec(&[1.0, 0.0], &mut [0.0]).unwrap_err());
}
#[test]
fn predictor_geometry_defaults_skip_zero_row_multipliers() {
let block = DefaultPaths;
let values = [1.0, 0.0, 3.0];
let mut gram = [0.0];
block
.add_weighted_gram_by(&values, &PanicOnMaskedRow, &mut gram)
.unwrap();
assert_relative_eq!(gram[0], 8.0);
let mut transpose = [0.0];
block
.add_t_mul_vec_by(&values, &PanicOnMaskedRow, &mut transpose)
.unwrap();
assert_relative_eq!(transpose[0], 8.0);
}
#[test]
fn scalar_weighted_gradient_skips_zero_score_nan_multiplier() {
let values = [1.0, 0.0, 3.0];
let scalar = SoftplusScalar::new(3);
let mut scalar_gradient = [0.0];
scalar.add_weighted_gradient(&values, &[2.0, f64::NAN, 2.0], &[0.0], &mut scalar_gradient);
assert!(scalar_gradient[0].is_finite());
}
#[allow(clippy::needless_pass_by_value)]
fn assert_multiplier_length_error(error: ModelError) {
assert_eq!(
error,
ModelError::DesignRowMismatch {
parameter: "product multiplier",
expected_rows: 2,
actual_rows: 1,
}
);
}
#[allow(clippy::needless_pass_by_value)]
fn assert_invalid_multiplier_error(error: ModelError) {
assert_eq!(error, ModelError::InvalidMultiplier { index: 1 });
}
}