use ndarray::{Array1, Array2, Array3};
use crate::families::compiler::{RowHessian, RowJacobianOperator, scale_jacobian_by_sqrt_h_with};
use gam_problem::FamilyChannelHessian;
pub struct BernoulliRowHessian {
w: Array1<f64>,
}
impl BernoulliRowHessian {
pub fn from_row_weights(w: Array1<f64>) -> Self {
Self { w }
}
pub fn row_weights(&self) -> &Array1<f64> {
&self.w
}
}
impl RowHessian for BernoulliRowHessian {
fn k(&self) -> usize {
1
}
fn nrows(&self) -> usize {
self.w.len()
}
fn fill_row(&self, row: usize, out: &mut [f64]) {
assert_eq!(out.len(), 1, "BernoulliRowHessian::fill_row expects K=1");
out[0] = self.w[row];
}
fn evaluate_full(&self) -> Array3<f64> {
let n = self.w.len();
let mut out = Array3::<f64>::zeros((n, 1, 1));
for i in 0..n {
out[[i, 0, 0]] = self.w[i];
}
out
}
}
impl FamilyChannelHessian for BernoulliRowHessian {
fn fill_subject(&self, i: usize, out: &mut [f64]) {
assert_eq!(
out.len(),
1,
"BernoulliRowHessian::fill_subject expects K=1"
);
out[0] = self.w[i];
}
fn n_subjects(&self) -> usize {
self.w.len()
}
fn n_outputs(&self) -> usize {
1
}
fn evaluate_full(&self) -> ndarray::Array3<f64> {
let n = self.w.len();
let mut out = ndarray::Array3::<f64>::zeros((n, 1, 1));
for i in 0..n {
out[[i, 0, 0]] = self.w[i];
}
out
}
}
pub struct BernoulliDenseDesignOperator {
design: Array2<f64>,
}
impl BernoulliDenseDesignOperator {
pub fn new(design: Array2<f64>) -> Self {
Self { design }
}
}
impl RowJacobianOperator for BernoulliDenseDesignOperator {
fn k(&self) -> usize {
1
}
fn ncols(&self) -> usize {
self.design.ncols()
}
fn nrows(&self) -> usize {
self.design.nrows()
}
fn apply_row(&self, row: usize, delta_beta: &[f64], out: &mut [f64]) {
assert_eq!(out.len(), 1);
assert_eq!(delta_beta.len(), self.design.ncols());
let mut acc = 0.0;
for (j, &b) in delta_beta.iter().enumerate() {
acc += self.design[[row, j]] * b;
}
out[0] = acc;
}
fn evaluate_full(&self) -> Array3<f64> {
let n = self.design.nrows();
let p = self.design.ncols();
let mut out = Array3::<f64>::zeros((n, p, 1));
for i in 0..n {
for j in 0..p {
out[[i, j, 0]] = self.design[[i, j]];
}
}
out
}
fn scaled_design_by_sqrt_h(&self, h_full: &Array3<f64>) -> Array2<f64> {
let n = self.design.nrows();
let p = self.design.ncols();
scale_jacobian_by_sqrt_h_with(n, p, 1, h_full, |i, a, c| {
assert_eq!(c, 0, "K=1 operator has only channel 0");
self.design[[i, a]]
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dense_design_operator_evaluate_full_shape() {
let design = Array2::from_shape_fn((5, 3), |(i, j)| (i as f64) * 0.1 + (j as f64));
let op = BernoulliDenseDesignOperator::new(design.clone());
let full = op.evaluate_full();
assert_eq!(full.shape(), &[5, 3, 1]);
for i in 0..5 {
for j in 0..3 {
assert_eq!(full[[i, j, 0]], design[[i, j]]);
}
}
}
}