use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AuxOutcomeFamily {
Binomial,
Multinomial { n_classes: usize },
}
impl AuxOutcomeFamily {
pub fn n_eta_channels(&self) -> usize {
match self {
AuxOutcomeFamily::Binomial => 1,
AuxOutcomeFamily::Multinomial { n_classes } => n_classes.saturating_sub(1),
}
}
}
#[derive(Debug, Clone)]
pub struct BehavioralHead {
family: AuxOutcomeFamily,
y: Array1<f64>,
w_row: Array1<f64>,
}
impl BehavioralHead {
pub fn new(
family: AuxOutcomeFamily,
y: Array1<f64>,
w_row: Array1<f64>,
) -> Result<Self, String> {
let n = y.len();
if w_row.len() != n {
return Err(format!(
"BehavioralHead: w_row length {} != labels length {n}",
w_row.len()
));
}
for &v in w_row.iter() {
if !(v.is_finite() && v >= 0.0) {
return Err(format!(
"BehavioralHead: row weights must be finite and ≥ 0, got {v}"
));
}
}
match family {
AuxOutcomeFamily::Binomial => {
for (i, &label) in y.iter().enumerate() {
if label != 0.0 && label != 1.0 {
return Err(format!(
"BehavioralHead(Binomial): label[{i}] = {label} is not 0/1"
));
}
}
}
AuxOutcomeFamily::Multinomial { n_classes } => {
if n_classes < 2 {
return Err(format!(
"BehavioralHead(Multinomial): need ≥ 2 classes, got {n_classes}"
));
}
for (i, &label) in y.iter().enumerate() {
let k = label as usize;
if k as f64 != label || k >= n_classes {
return Err(format!(
"BehavioralHead(Multinomial): label[{i}] = {label} not an \
integer class index in 0..{n_classes}"
));
}
}
}
}
Ok(Self { family, y, w_row })
}
pub fn fully_supervised(family: AuxOutcomeFamily, y: Array1<f64>) -> Result<Self, String> {
let n = y.len();
Self::new(family, y, Array1::from_elem(n, 1.0))
}
pub fn family(&self) -> AuxOutcomeFamily {
self.family
}
pub fn n_obs(&self) -> usize {
self.y.len()
}
pub fn n_coeffs(&self, latent_dim: usize) -> usize {
self.family.n_eta_channels() * (1 + latent_dim)
}
pub fn effective_labeled_count(&self) -> f64 {
self.w_row.iter().sum()
}
fn eta(&self, t: ArrayView2<'_, f64>, coeffs: ArrayView1<'_, f64>) -> Array2<f64> {
let (n, d) = t.dim();
let n_eta = self.family.n_eta_channels();
let mut eta = Array2::<f64>::zeros((n, n_eta));
for c in 0..n_eta {
let base = c * (1 + d);
let a = coeffs[base];
for row in 0..n {
let mut acc = a;
for axis in 0..d {
acc += t[[row, axis]] * coeffs[base + 1 + axis];
}
eta[[row, c]] = acc;
}
}
eta
}
pub fn neg_loglik_and_grad(
&self,
t: ArrayView2<'_, f64>,
coeffs: ArrayView1<'_, f64>,
) -> Result<(f64, Array1<f64>, Array2<f64>), String> {
let (n, d) = t.dim();
if n != self.y.len() {
return Err(format!(
"BehavioralHead: latent rows {n} != labels {}",
self.y.len()
));
}
let n_eta = self.family.n_eta_channels();
if coeffs.len() != n_eta * (1 + d) {
return Err(format!(
"BehavioralHead: coeffs length {} != n_eta·(1+d) = {}",
coeffs.len(),
n_eta * (1 + d)
));
}
let eta = self.eta(t, coeffs);
let mut nll = 0.0_f64;
let mut grad_coeffs = Array1::<f64>::zeros(n_eta * (1 + d));
let mut grad_t = Array2::<f64>::zeros((n, d));
match self.family {
AuxOutcomeFamily::Binomial => {
for row in 0..n {
let w = self.w_row[row];
if w == 0.0 {
continue;
}
let e = eta[[row, 0]];
let log1p = if e > 0.0 {
e + (-e).exp().ln_1p()
} else {
e.exp().ln_1p()
};
let y = self.y[row];
nll += w * (log1p - y * e);
let p = 1.0 / (1.0 + (-e).exp());
let r = w * (p - y);
grad_coeffs[0] += r;
for axis in 0..d {
grad_coeffs[1 + axis] += r * t[[row, axis]];
grad_t[[row, axis]] += r * coeffs[1 + axis];
}
}
}
AuxOutcomeFamily::Multinomial { .. } => {
for row in 0..n {
let w = self.w_row[row];
if w == 0.0 {
continue;
}
let mut max_eta = 0.0_f64;
for c in 0..n_eta {
if eta[[row, c]] > max_eta {
max_eta = eta[[row, c]];
}
}
let mut denom = (0.0 - max_eta).exp();
for c in 0..n_eta {
denom += (eta[[row, c]] - max_eta).exp();
}
let lse = max_eta + denom.ln();
let label = self.y[row] as usize;
let eta_y = if label == 0 {
0.0
} else {
eta[[row, label - 1]]
};
nll += w * (lse - eta_y);
for c in 0..n_eta {
let p_c = (eta[[row, c]] - lse).exp();
let indicator = if label == c + 1 { 1.0 } else { 0.0 };
let r = w * (p_c - indicator);
let base = c * (1 + d);
grad_coeffs[base] += r;
for axis in 0..d {
grad_coeffs[base + 1 + axis] += r * t[[row, axis]];
grad_t[[row, axis]] += r * coeffs[base + 1 + axis];
}
}
}
}
}
Ok((nll, grad_coeffs, grad_t))
}
}
#[derive(Debug, Clone)]
pub struct LeakageAbsorber {
q: Array2<f64>,
}
impl LeakageAbsorber {
pub fn rank(&self) -> usize {
self.q.ncols()
}
pub fn basis(&self) -> ArrayView2<'_, f64> {
self.q.view()
}
}
#[derive(Debug, Clone)]
pub struct HeadFeatureSignificance {
pub statistic: Vec<f64>,
pub p_value: Vec<f64>,
pub fdr_rejected: Vec<usize>,
pub alpha: f64,
}