1use super::*;
2
3#[inline]
8fn entropy_log_plus_one(p: f64) -> f64 {
9 if p > 0.0 { p.ln() + 1.0 } else { 0.0 }
10}
11
12#[derive(Debug, Clone, Copy)]
27pub enum SparsityKind {
28 SmoothedL1 { eps: f64 },
29 Hoyer,
30 Log { delta: f64 },
31}
32
33#[derive(Debug, Clone)]
47pub struct SparsityPenalty {
48 pub target_tier: PenaltyTier,
49 pub kind: SparsityKind,
50 pub weight: f64,
51 pub weight_schedule: Option<ScalarWeightSchedule>,
52 pub strength_rho_index: usize,
54 pub eps_rho_index: Option<usize>,
58}
59
60#[derive(Debug, Clone)]
76pub struct SoftmaxAssignmentSparsityPenalty {
77 pub k_atoms: usize,
78 pub temperature: f64,
79 pub weight: f64,
80 pub weight_schedule: Option<ScalarWeightSchedule>,
81 pub row_weights: Option<std::sync::Arc<[f64]>>,
92}
93
94impl SoftmaxAssignmentSparsityPenalty {
95 #[must_use]
96 pub fn new(k_atoms: usize, temperature: f64) -> Self {
97 assert!(k_atoms > 0);
98 assert!(temperature > 0.0);
99 Self {
100 k_atoms,
101 temperature,
102 weight: 1.0,
103 weight_schedule: None,
104 row_weights: None,
105 }
106 }
107
108 #[must_use]
112 pub fn with_row_weights(mut self, weights: Option<&[f64]>) -> Self {
113 self.row_weights = weights.map(|w| std::sync::Arc::from(w.to_vec()));
114 self
115 }
116
117 #[must_use]
121 pub fn row_weight(&self, row: usize) -> f64 {
122 self.row_weights.as_ref().map_or(1.0, |w| w[row])
123 }
124
125 impl_with_weight_schedule!(weight);
126
127 fn softmax_row(&self, row: &[f64]) -> Vec<f64> {
128 let inv_tau = 1.0 / self.temperature;
129 let mut max_logit = f64::NEG_INFINITY;
130 for (idx, &v) in row.iter().enumerate() {
131 assert!(
132 v.is_finite(),
133 "SoftmaxAssignmentSparsityPenalty: non-finite logit at atom {idx}: {v}"
134 );
135 max_logit = max_logit.max(v);
136 }
137 let mut out = vec![0.0; self.k_atoms];
138 let mut sum = 0.0;
139 for i in 0..self.k_atoms {
140 let v = ((row[i] - max_logit) * inv_tau).exp();
141 out[i] = v;
142 sum += v;
143 }
144 assert!(
145 sum.is_finite() && sum > 0.0,
146 "SoftmaxAssignmentSparsityPenalty: non-finite softmax normalizer"
147 );
148 for v in out.iter_mut() {
149 *v /= sum;
150 }
151 out
152 }
153
154 pub fn psd_majorizer_abs_row_sums(&self, row: &[f64], scale: f64) -> Vec<f64> {
174 let a = self.softmax_row(row);
175 let k = self.k_atoms;
176 let l: Vec<f64> = (0..k).map(|i| entropy_log_plus_one(a[i])).collect();
177 let m: f64 = (0..k).map(|i| a[i] * l[i]).sum();
178 let mut d = vec![0.0_f64; k];
179 for kk in 0..k {
180 let h_kk = scale * a[kk] * ((m - l[kk] - 1.0) + a[kk] * (2.0 * l[kk] + 1.0 - 2.0 * m));
182 let mut acc = h_kk.abs();
183 for jj in 0..k {
185 if jj == kk {
186 continue;
187 }
188 let h_kj = scale * a[kk] * a[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m);
189 acc += h_kj.abs();
190 }
191 d[kk] = acc;
192 }
193 d
194 }
195
196 #[must_use]
212 pub fn row_dense_hessian(&self, row_logits: &[f64], scale: f64) -> Array2<f64> {
213 let k = self.k_atoms;
214 let a = self.softmax_row(row_logits);
215 let l: Vec<f64> = (0..k).map(|i| entropy_log_plus_one(a[i])).collect();
216 let m: f64 = (0..k).map(|i| a[i] * l[i]).sum();
217 let mut h = Array2::<f64>::zeros((k, k));
218 for kk in 0..k {
219 for jj in 0..k {
220 let indicator = if kk == jj { 1.0 } else { 0.0 };
221 h[[kk, jj]] = scale
222 * a[kk]
223 * (indicator * (m - l[kk] - 1.0) + a[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m));
224 }
225 }
226 h
227 }
228
229 #[must_use]
237 pub fn row_dense_hessian_logit_derivative(
238 &self,
239 row_logits: &[f64],
240 scale: f64,
241 w: usize,
242 ) -> Array2<f64> {
243 let k = self.k_atoms;
244 let inv_tau = 1.0 / self.temperature;
245 let a = self.softmax_row(row_logits);
246 let l: Vec<f64> = (0..k).map(|i| entropy_log_plus_one(a[i])).collect();
247 let m: f64 = (0..k).map(|i| a[i] * l[i]).sum();
248 let da: Vec<f64> = (0..k)
250 .map(|r| a[r] * (if r == w { 1.0 } else { 0.0 } - a[w]) * inv_tau)
251 .collect();
252 let dl: Vec<f64> = (0..k)
253 .map(|r| if a[r] > 0.0 { da[r] / a[r] } else { 0.0 })
254 .collect();
255 let dm: f64 = (0..k).map(|r| da[r] * l[r] + a[r] * dl[r]).sum();
256 let mut dh = Array2::<f64>::zeros((k, k));
257 for kk in 0..k {
258 for jj in 0..k {
259 let indicator = if kk == jj { 1.0 } else { 0.0 };
260 let bracket =
262 indicator * (m - l[kk] - 1.0) + a[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m);
263 let dbracket = indicator * (dm - dl[kk])
264 + da[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m)
265 + a[jj] * (dl[kk] + dl[jj] - 2.0 * dm);
266 dh[[kk, jj]] = scale * (da[kk] * bracket + a[kk] * dbracket);
267 }
268 }
269 dh
270 }
271
272 #[must_use]
290 pub fn row_psd_majorizer(&self, row_logits: &[f64], scale: f64) -> Array2<f64> {
291 let k = self.k_atoms;
292 let d = self.psd_majorizer_abs_row_sums(row_logits, scale);
293 let mut out = Array2::<f64>::zeros((k, k));
294 for kk in 0..k {
295 out[[kk, kk]] = d[kk];
296 }
297 out
298 }
299
300 #[must_use]
310 pub fn row_psd_majorizer_logit_derivative(
311 &self,
312 row_logits: &[f64],
313 scale: f64,
314 w: usize,
315 ) -> Array2<f64> {
316 let k = self.k_atoms;
317 let h = self.row_dense_hessian(row_logits, scale);
318 let dh = self.row_dense_hessian_logit_derivative(row_logits, scale, w);
319 let mut out = Array2::<f64>::zeros((k, k));
320 for kk in 0..k {
321 let mut acc = 0.0_f64;
322 for jj in 0..k {
323 let s = h[[kk, jj]].signum();
324 if h[[kk, jj]] != 0.0 {
325 acc += s * dh[[kk, jj]];
326 }
327 }
328 out[[kk, kk]] = acc;
329 }
330 out
331 }
332
333 #[must_use]
352 pub fn row_fisher_metric(&self, row_logits: &[f64], scale: f64) -> Array2<f64> {
353 let k = self.k_atoms;
354 let a = self.softmax_row(row_logits);
355 let mut g = Array2::<f64>::zeros((k, k));
356 for kk in 0..k {
357 for jj in 0..k {
358 let indicator = if kk == jj { 1.0 } else { 0.0 };
359 g[[kk, jj]] = scale * a[kk] * (indicator - a[jj]);
360 }
361 }
362 g
363 }
364
365 #[must_use]
377 pub fn row_fisher_metric_logit_derivative(
378 &self,
379 row_logits: &[f64],
380 scale: f64,
381 w: usize,
382 ) -> Array2<f64> {
383 let k = self.k_atoms;
384 let inv_tau = 1.0 / self.temperature;
385 let a = self.softmax_row(row_logits);
386 let da: Vec<f64> = (0..k)
389 .map(|r| a[r] * (if r == w { 1.0 } else { 0.0 } - a[w]) * inv_tau)
390 .collect();
391 let mut dg = Array2::<f64>::zeros((k, k));
392 for kk in 0..k {
393 for jj in 0..k {
394 let indicator = if kk == jj { 1.0 } else { 0.0 };
395 dg[[kk, jj]] = scale * (da[kk] * (indicator - a[jj]) - a[kk] * da[jj]);
396 }
397 }
398 dg
399 }
400}
401
402impl AnalyticPenalty for SoftmaxAssignmentSparsityPenalty {
403 fn tier(&self) -> PenaltyTier {
404 PenaltyTier::Psi
405 }
406
407 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
408 let lambda = resolve_learnable_weight(self.weight, rho[0]);
409 let n = target.len() / self.k_atoms;
410 let values: Vec<f64> = target.iter().copied().collect();
411 let mut acc = 0.0;
412 for row in 0..n {
413 let start = row * self.k_atoms;
414 let a = self.softmax_row(&values[start..start + self.k_atoms]);
415 let w_row = self.row_weight(row);
416 for v in a {
417 if v > 0.0 {
418 acc += -w_row * v * v.ln();
419 }
420 }
421 }
422 lambda * acc
423 }
424
425 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
426 let lambda = resolve_learnable_weight(self.weight, rho[0]);
427 let n = target.len() / self.k_atoms;
428 let values: Vec<f64> = target.iter().copied().collect();
429 let mut out = Array1::<f64>::zeros(target.len());
430 let inv_tau = 1.0 / self.temperature;
431 for row in 0..n {
432 let start = row * self.k_atoms;
433 let a = self.softmax_row(&values[start..start + self.k_atoms]);
434 let w_row = self.row_weight(row);
435 let mut d_h_da = vec![0.0; self.k_atoms];
436 let mut mean = 0.0;
437 for k in 0..self.k_atoms {
438 d_h_da[k] = -lambda * entropy_log_plus_one(a[k]);
439 mean += a[k] * d_h_da[k];
440 }
441 for k in 0..self.k_atoms {
442 out[start + k] = w_row * a[k] * (d_h_da[k] - mean) * inv_tau;
443 }
444 }
445 out
446 }
447
448 fn hessian_diag(
449 &self,
450 target: ArrayView1<'_, f64>,
451 rho: ArrayView1<'_, f64>,
452 ) -> Option<Array1<f64>> {
453 assert_eq!(rho.len(), 1, "softmax entropy expects one rho parameter");
454 assert!(
455 rho.iter().all(|value| value.is_finite()),
456 "softmax entropy rho must be finite"
457 );
458 assert_eq!(
459 target.len() % self.k_atoms,
460 0,
461 "softmax entropy target length must be divisible by k_atoms"
462 );
463 let lambda = resolve_learnable_weight(self.weight, rho[0]);
472 let inv_tau = 1.0 / self.temperature;
473 let scale = lambda * inv_tau * inv_tau;
474 let n = target.len() / self.k_atoms;
475 let values: Vec<f64> = target.iter().copied().collect();
476 let mut out = Array1::<f64>::zeros(target.len());
477 for row in 0..n {
478 let start = row * self.k_atoms;
479 let a = self.softmax_row(&values[start..start + self.k_atoms]);
480 let w_row = self.row_weight(row);
481 let mut mean_log_plus_one = 0.0;
482 for k in 0..self.k_atoms {
483 mean_log_plus_one += a[k] * entropy_log_plus_one(a[k]);
484 }
485 for k in 0..self.k_atoms {
486 let log_plus_one = entropy_log_plus_one(a[k]);
487 let term = (1.0 - 2.0 * a[k]) * (mean_log_plus_one - log_plus_one) + a[k] - 1.0;
488 out[start + k] = w_row * scale * a[k] * term;
489 }
490 }
491 Some(out)
492 }
493
494 fn hvp(
495 &self,
496 target: ArrayView1<'_, f64>,
497 rho: ArrayView1<'_, f64>,
498 v: ArrayView1<'_, f64>,
499 ) -> Array1<f64> {
500 let lambda = resolve_learnable_weight(self.weight, rho[0]);
511 assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
512 let n = target.len() / self.k_atoms;
513 let values: Vec<f64> = target.iter().copied().collect();
514 let mut out = Array1::<f64>::zeros(target.len());
515 let inv_tau = 1.0 / self.temperature;
516 let scale = lambda * inv_tau * inv_tau;
517 for row in 0..n {
518 let start = row * self.k_atoms;
519 let a = self.softmax_row(&values[start..start + self.k_atoms]);
520 let w_row = self.row_weight(row);
521 let mut mean_log_plus_one = 0.0;
522 let mut mean_v = 0.0;
523 for k in 0..self.k_atoms {
524 mean_log_plus_one += a[k] * entropy_log_plus_one(a[k]);
525 mean_v += a[k] * v[start + k];
526 }
527 let mut mean_centered_v_log_plus_one = 0.0;
528 for k in 0..self.k_atoms {
529 let centered_v = v[start + k] - mean_v;
530 mean_centered_v_log_plus_one += a[k] * centered_v * entropy_log_plus_one(a[k]);
531 }
532 for k in 0..self.k_atoms {
533 let log_plus_one = entropy_log_plus_one(a[k]);
534 let centered_v = v[start + k] - mean_v;
535 out[start + k] = w_row
536 * scale
537 * a[k]
538 * (centered_v * (mean_log_plus_one - log_plus_one - 1.0)
539 + mean_centered_v_log_plus_one);
540 }
541 }
542 out
543 }
544
545 fn psd_majorizer_diag(
546 &self,
547 target: ArrayView1<'_, f64>,
548 rho: ArrayView1<'_, f64>,
549 ) -> Option<Array1<f64>> {
550 assert_eq!(rho.len(), 1, "softmax entropy expects one rho parameter");
551 assert_eq!(
552 target.len() % self.k_atoms,
553 0,
554 "softmax entropy target length must be divisible by k_atoms"
555 );
556 let lambda = resolve_learnable_weight(self.weight, rho[0]);
565 let inv_tau = 1.0 / self.temperature;
566 let scale = lambda * inv_tau * inv_tau;
567 let n = target.len() / self.k_atoms;
568 let values: Vec<f64> = target.iter().copied().collect();
569 let mut out = Array1::<f64>::zeros(target.len());
570 for row in 0..n {
571 let start = row * self.k_atoms;
572 let w_row = self.row_weight(row);
573 let d = self.psd_majorizer_abs_row_sums(&values[start..start + self.k_atoms], scale);
574 for k in 0..self.k_atoms {
575 out[start + k] = w_row * d[k];
576 }
577 }
578 Some(out)
579 }
580
581 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
582 Array1::from_vec(vec![self.value(target, rho)])
583 }
584
585 fn rho_count(&self) -> usize {
586 1
587 }
588
589 fn name(&self) -> &str {
590 "softmax_assignment_sparsity"
591 }
592
593 impl_scalar_apply_schedule!(weight);
594}
595
596impl SparsityPenalty {
597 #[must_use = "build error must be handled"]
598 pub fn smoothed_l1(target_tier: PenaltyTier, eps: f64) -> Result<Self, String> {
599 if !(eps.is_finite() && eps > 0.0) {
600 return Err(format!(
601 "SparsityPenalty::smoothed_l1 requires eps > 0 \
602 (Hessian / gradient have a `1/sqrt(x² + eps²)` factor that needs eps > 0 \
603 for differentiability at x = 0); got eps = {eps}"
604 ));
605 }
606 Ok(Self {
607 target_tier,
608 kind: SparsityKind::SmoothedL1 { eps },
609 weight: 1.0,
610 weight_schedule: None,
611 strength_rho_index: 0,
612 eps_rho_index: None,
613 })
614 }
615
616 #[must_use = "build error must be handled"]
617 pub fn log(target_tier: PenaltyTier, delta: f64) -> Result<Self, String> {
618 if !(delta.is_finite() && delta > 0.0) {
619 return Err(format!(
620 "SparsityPenalty::log requires delta > 0 \
621 (the log-sparsifier is log(1 + x²/δ²), undefined at δ = 0); \
622 got delta = {delta}"
623 ));
624 }
625 Ok(Self {
626 target_tier,
627 kind: SparsityKind::Log { delta },
628 weight: 1.0,
629 weight_schedule: None,
630 strength_rho_index: 0,
631 eps_rho_index: None,
632 })
633 }
634
635 #[must_use]
638 pub fn hoyer(target_tier: PenaltyTier) -> Self {
639 Self {
640 target_tier,
641 kind: SparsityKind::Hoyer,
642 weight: 1.0,
643 weight_schedule: None,
644 strength_rho_index: 0,
645 eps_rho_index: None,
646 }
647 }
648
649 impl_with_weight_schedule!(weight);
650
651 #[must_use]
652 pub fn with_eps_reml(mut self, eps_rho_index: usize) -> Self {
653 self.eps_rho_index = Some(eps_rho_index);
654 self
655 }
656
657 fn resolved(&self, rho: ArrayView1<'_, f64>) -> (f64, f64) {
659 let strength = resolve_learnable_weight(self.weight, rho[self.strength_rho_index]);
660 let smoothing = match (self.eps_rho_index, self.kind) {
661 (Some(idx), _) => rho[idx].exp().max(f64::MIN_POSITIVE),
667 (None, SparsityKind::SmoothedL1 { eps }) => eps,
668 (None, SparsityKind::Log { delta }) => delta,
669 (None, SparsityKind::Hoyer) => 0.0,
670 };
671 (strength, smoothing)
672 }
673}
674
675impl AnalyticPenalty for SparsityPenalty {
676 fn tier(&self) -> PenaltyTier {
677 self.target_tier
678 }
679
680 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
681 let (lam, smooth) = self.resolved(rho);
682 match self.kind {
683 SparsityKind::SmoothedL1 { .. } => {
684 let mut acc = 0.0;
685 for &x in target.iter() {
686 acc += (x * x + smooth * smooth).sqrt();
687 }
688 lam * acc
689 }
690 SparsityKind::Hoyer => {
691 let n = target.len() as f64;
698 assert!(n > 1.0, "Hoyer requires n > 1");
699 let l1: f64 = target.iter().map(|x| x.abs()).sum();
700 let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
701 if l2 == 0.0 {
702 return 0.0;
703 }
704 let h = (l1 / l2 - 1.0) / (n.sqrt() - 1.0);
705 lam * h
706 }
707 SparsityKind::Log { .. } => {
708 let mut acc = 0.0;
709 let d2 = smooth * smooth;
710 for &x in target.iter() {
711 acc += (1.0 + x * x / d2).ln();
712 }
713 lam * acc
714 }
715 }
716 }
717
718 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
719 let (lam, smooth) = self.resolved(rho);
720 let mut g = Array1::<f64>::zeros(target.len());
721 match self.kind {
722 SparsityKind::SmoothedL1 { .. } => {
723 let eps2 = smooth * smooth;
724 for (i, &x) in target.iter().enumerate() {
725 g[i] = lam * x / (x * x + eps2).sqrt();
726 }
727 }
728 SparsityKind::Hoyer => {
729 let n = target.len() as f64;
732 assert!(n > 1.0, "Hoyer requires n > 1");
733 let l1: f64 = target.iter().map(|x| x.abs()).sum();
734 let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
735 if l2 == 0.0 {
736 return g;
737 }
738 let denom = n.sqrt() - 1.0;
739 let a = lam / denom;
740 let inv_l2 = 1.0 / l2;
741 let inv_l2_cubed = inv_l2 * inv_l2 * inv_l2;
742 for (i, &x) in target.iter().enumerate() {
743 let sgn = if x > 0.0 {
744 1.0
745 } else if x < 0.0 {
746 -1.0
747 } else {
748 0.0
749 };
750 g[i] = a * (sgn * inv_l2 - l1 * x * inv_l2_cubed);
751 }
752 }
753 SparsityKind::Log { .. } => {
754 let d2 = smooth * smooth;
755 for (i, &x) in target.iter().enumerate() {
756 g[i] = lam * 2.0 * x / (d2 + x * x);
757 }
758 }
759 }
760 g
761 }
762
763 fn hessian_diag(
764 &self,
765 target: ArrayView1<'_, f64>,
766 rho: ArrayView1<'_, f64>,
767 ) -> Option<Array1<f64>> {
768 let (lam, smooth) = self.resolved(rho);
769 match self.kind {
770 SparsityKind::SmoothedL1 { .. } => {
771 let mut d = Array1::<f64>::zeros(target.len());
772 let eps2 = smooth * smooth;
773 for (i, &x) in target.iter().enumerate() {
774 let r = (x * x + eps2).sqrt();
775 d[i] = lam * eps2 / (r * r * r);
776 }
777 Some(d)
778 }
779 SparsityKind::Log { .. } => {
780 let mut d = Array1::<f64>::zeros(target.len());
781 let d2 = smooth * smooth;
790 for (i, &x) in target.iter().enumerate() {
791 let denom = d2 + x * x;
792 d[i] = lam * 2.0 * (d2 - x * x) / (denom * denom);
793 }
794 Some(d)
795 }
796 SparsityKind::Hoyer => None,
803 }
804 }
805
806 fn hvp(
807 &self,
808 target: ArrayView1<'_, f64>,
809 rho: ArrayView1<'_, f64>,
810 v: ArrayView1<'_, f64>,
811 ) -> Array1<f64> {
812 let (lam, smooth) = self.resolved(rho);
818 let n_target = target.len();
819 assert_eq!(v.len(), n_target, "hvp dimension mismatch");
820 match self.kind {
821 SparsityKind::SmoothedL1 { .. } => {
822 let mut out = Array1::<f64>::zeros(n_target);
823 let eps2 = smooth * smooth;
824 for (i, &x) in target.iter().enumerate() {
825 let r = (x * x + eps2).sqrt();
826 out[i] = lam * eps2 / (r * r * r) * v[i];
827 }
828 out
829 }
830 SparsityKind::Log { .. } => {
831 let mut out = Array1::<f64>::zeros(n_target);
837 let d2 = smooth * smooth;
838 for (i, &x) in target.iter().enumerate() {
839 let denom = d2 + x * x;
840 out[i] = lam * 2.0 * (d2 - x * x) / (denom * denom) * v[i];
841 }
842 out
843 }
844 SparsityKind::Hoyer => {
845 let n = n_target as f64;
851 assert!(n > 1.0, "Hoyer requires n > 1");
852 let l1: f64 = target.iter().map(|x| x.abs()).sum();
853 let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
854 let mut out = Array1::<f64>::zeros(n_target);
855 if l2 == 0.0 {
856 return out;
857 }
858 let a = lam / (n.sqrt() - 1.0);
859 let inv_l2_cubed = 1.0 / (l2 * l2 * l2);
860 let inv_l2_5 = inv_l2_cubed / (l2 * l2);
861 let mut x_dot_v = 0.0;
862 let mut s_dot_v = 0.0;
863 for i in 0..n_target {
864 let xi = target[i];
865 let si = if xi > 0.0 {
866 1.0
867 } else if xi < 0.0 {
868 -1.0
869 } else {
870 0.0
871 };
872 x_dot_v += xi * v[i];
873 s_dot_v += si * v[i];
874 }
875 for i in 0..n_target {
876 let xi = target[i];
877 let si = if xi > 0.0 {
878 1.0
879 } else if xi < 0.0 {
880 -1.0
881 } else {
882 0.0
883 };
884 out[i] = a
885 * (-si * x_dot_v * inv_l2_cubed
886 - xi * s_dot_v * inv_l2_cubed
887 - l1 * v[i] * inv_l2_cubed
888 + 3.0 * l1 * xi * x_dot_v * inv_l2_5);
889 }
890 out
891 }
892 }
893 }
894
895 fn psd_majorizer_diag(
896 &self,
897 target: ArrayView1<'_, f64>,
898 rho: ArrayView1<'_, f64>,
899 ) -> Option<Array1<f64>> {
900 let (lam, smooth) = self.resolved(rho);
901 match self.kind {
902 SparsityKind::SmoothedL1 { .. } => self.hessian_diag(target, rho),
904 SparsityKind::Log { .. } => {
908 let mut d = Array1::<f64>::zeros(target.len());
909 let d2 = smooth * smooth;
910 for (i, &x) in target.iter().enumerate() {
911 d[i] = lam * 2.0 / (d2 + x * x);
912 }
913 Some(d)
914 }
915 SparsityKind::Hoyer => None,
918 }
919 }
920
921 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
922 let n_rho = self.rho_count();
925 let mut out = Array1::<f64>::zeros(n_rho);
926 let p_val = self.value(target, rho);
927 out[self.strength_rho_index] = p_val;
928 if let Some(eps_idx) = self.eps_rho_index {
929 let (lam, smooth) = self.resolved(rho);
930 let mut dp_deps = 0.0;
931 match self.kind {
932 SparsityKind::SmoothedL1 { .. } => {
933 for &x in target.iter() {
934 dp_deps += smooth / (x * x + smooth * smooth).sqrt();
935 }
936 dp_deps *= lam;
937 }
938 SparsityKind::Log { .. } => {
939 let d2 = smooth * smooth;
941 for &x in target.iter() {
942 dp_deps += -2.0 * x * x / (smooth * (d2 + x * x));
943 }
944 dp_deps *= lam;
945 }
946 SparsityKind::Hoyer => {}
947 }
948 out[eps_idx] = smooth * dp_deps;
950 }
951 out
952 }
953
954 fn rho_count(&self) -> usize {
955 1 + if self.eps_rho_index.is_some() { 1 } else { 0 }
956 }
957
958 fn name(&self) -> &str {
959 "sparsity"
960 }
961
962 impl_scalar_apply_schedule!(weight);
963}
964
965#[derive(Debug, Clone)]
970pub struct TopKActivationPenalty {
971 pub target: PsiSlice,
972 pub k: usize,
973 pub latent_dim: usize,
974 pub weight: f64,
975 pub weight_schedule: Option<ScalarWeightSchedule>,
976}
977
978impl TopKActivationPenalty {
979 #[must_use = "build error must be handled"]
980 pub fn new(target: PsiSlice, k: usize, weight: f64) -> Result<Self, String> {
981 let latent_dim = target
982 .latent_dim
983 .ok_or_else(|| "TopKActivationPenalty::new requires target.latent_dim".to_string())?;
984 if latent_dim == 0 {
985 return Err("TopKActivationPenalty::new requires latent_dim > 0".to_string());
986 }
987 if k == 0 || k > latent_dim {
988 return Err(format!(
989 "TopKActivationPenalty::new requires 0 < k <= latent_dim; got k={k}, latent_dim={latent_dim}"
990 ));
991 }
992 if !(weight.is_finite() && weight > 0.0) {
993 return Err(format!(
994 "TopKActivationPenalty::new requires finite weight > 0, got {weight}"
995 ));
996 }
997 Ok(Self {
998 target,
999 k,
1000 latent_dim,
1001 weight,
1002 weight_schedule: None,
1003 })
1004 }
1005
1006 impl_with_weight_schedule!(weight);
1007
1008 fn topk_mask_row(&self, target: ArrayView1<'_, f64>, row: usize, mask: &mut [bool]) {
1009 mask.fill(false);
1010 let d = self.latent_dim;
1011 let base = row * d;
1012 let mut order = (0..d).collect::<Vec<_>>();
1013 order.sort_by(|&a, &b| {
1014 target[base + b]
1015 .abs()
1016 .total_cmp(&target[base + a].abs())
1017 .then_with(|| a.cmp(&b))
1018 });
1019 for &axis in order.iter().take(self.k) {
1020 mask[axis] = true;
1021 }
1022 }
1023}
1024
1025impl AnalyticPenalty for TopKActivationPenalty {
1026 fn tier(&self) -> PenaltyTier {
1027 PenaltyTier::Psi
1028 }
1029
1030 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
1031 assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
1032 let d = self.latent_dim;
1033 let n_obs = target.len() / d;
1034 let mut mask = vec![false; d];
1035 let mut acc = 0.0;
1036 for row in 0..n_obs {
1037 self.topk_mask_row(target, row, &mut mask);
1038 let base = row * d;
1039 for axis in 0..d {
1040 if mask[axis] {
1041 let v = target[base + axis];
1042 acc += 0.5 * self.weight * v * v;
1043 }
1044 }
1045 }
1046 acc
1047 }
1048
1049 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1050 assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
1051 let d = self.latent_dim;
1052 let n_obs = target.len() / d;
1053 let mut mask = vec![false; d];
1054 let mut grad = Array1::<f64>::zeros(target.len());
1055 for row in 0..n_obs {
1056 self.topk_mask_row(target, row, &mut mask);
1057 let base = row * d;
1058 for axis in 0..d {
1059 if mask[axis] {
1060 grad[base + axis] = self.weight * target[base + axis];
1061 }
1062 }
1063 }
1064 grad
1065 }
1066
1067 fn hessian_diag(
1068 &self,
1069 target: ArrayView1<'_, f64>,
1070 rho: ArrayView1<'_, f64>,
1071 ) -> Option<Array1<f64>> {
1072 assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
1073 let d = self.latent_dim;
1074 let n_obs = target.len() / d;
1075 let mut mask = vec![false; d];
1076 let mut diag = Array1::<f64>::zeros(target.len());
1077 for row in 0..n_obs {
1078 self.topk_mask_row(target, row, &mut mask);
1079 let base = row * d;
1080 for axis in 0..d {
1081 if mask[axis] {
1082 diag[base + axis] = self.weight;
1083 }
1084 }
1085 }
1086 Some(diag)
1087 }
1088
1089 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1090 assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
1091 assert_eq!(
1092 target.len() % self.latent_dim,
1093 0,
1094 "TopKActivationPenalty target length must be a multiple of latent_dim"
1095 );
1096 Array1::<f64>::zeros(0)
1097 }
1098
1099 fn rho_count(&self) -> usize {
1100 0
1101 }
1102
1103 fn name(&self) -> &str {
1104 "topk_activation"
1105 }
1106
1107 impl_scalar_apply_schedule!(weight);
1108}
1109
1110#[derive(Debug, Clone)]
1115pub struct JumpReLUPenalty {
1116 pub target: PsiSlice,
1117 pub latent_dim: usize,
1118 pub thresholds: Array1<f64>,
1119 pub weight: f64,
1120 pub smoothing_eps: f64,
1121 pub weight_schedule: Option<ScalarWeightSchedule>,
1122}
1123
1124impl JumpReLUPenalty {
1125 #[must_use = "build error must be handled"]
1126 pub fn new(
1127 target: PsiSlice,
1128 thresholds: Array1<f64>,
1129 weight: f64,
1130 smoothing_eps: f64,
1131 ) -> Result<Self, String> {
1132 let latent_dim = target
1133 .latent_dim
1134 .ok_or_else(|| "JumpReLUPenalty::new requires target.latent_dim".to_string())?;
1135 if latent_dim == 0 {
1136 return Err("JumpReLUPenalty::new requires latent_dim > 0".to_string());
1137 }
1138 if thresholds.len() != latent_dim {
1139 return Err(format!(
1140 "JumpReLUPenalty::new thresholds length {} does not match latent_dim {latent_dim}",
1141 thresholds.len()
1142 ));
1143 }
1144 for (idx, &tau) in thresholds.iter().enumerate() {
1145 if !(tau.is_finite() && tau > 0.0) {
1146 return Err(format!(
1147 "JumpReLUPenalty::new thresholds[{idx}] must be finite and > 0, got {tau}"
1148 ));
1149 }
1150 }
1151 if !(weight.is_finite() && weight > 0.0) {
1152 return Err(format!(
1153 "JumpReLUPenalty::new requires finite weight > 0, got {weight}"
1154 ));
1155 }
1156 if !(smoothing_eps.is_finite() && smoothing_eps > 0.0) {
1157 return Err(format!(
1158 "JumpReLUPenalty::new requires finite smoothing_eps > 0, got {smoothing_eps}"
1159 ));
1160 }
1161 Ok(Self {
1162 target,
1163 latent_dim,
1164 thresholds,
1165 weight,
1166 smoothing_eps,
1167 weight_schedule: None,
1168 })
1169 }
1170
1171 impl_with_weight_schedule!(weight);
1172
1173 fn threshold(&self, axis: usize, rho: ArrayView1<'_, f64>) -> f64 {
1174 resolve_learnable_weight(self.thresholds[axis], rho[axis])
1178 }
1179
1180 pub(crate) fn sigmoid_gate(&self, x: f64) -> f64 {
1181 if x >= 0.0 {
1182 1.0 / (1.0 + (-x).exp())
1183 } else {
1184 let ex = x.exp();
1185 ex / (1.0 + ex)
1186 }
1187 }
1188
1189 fn true_hessian_diag_entry(&self, tau: f64, gate: f64) -> f64 {
1190 self.weight * tau * gate * (1.0 - gate) * (1.0 - 2.0 * gate)
1191 / (self.smoothing_eps * self.smoothing_eps)
1192 }
1193
1194 fn psd_hessian_diag_entry(&self, tau: f64, gate: f64) -> f64 {
1195 let slope = gate * (1.0 - gate);
1211 let reweighted_l2 = slope * slope;
1212 let abs_exact = slope * (1.0 - 2.0 * gate).abs();
1213 self.weight * tau * reweighted_l2.max(abs_exact) / (self.smoothing_eps * self.smoothing_eps)
1214 }
1215}
1216
1217#[must_use]
1232pub fn jumprelu_gate_value_grad(z: f64, tau: f64, smoothing_eps: f64) -> (f64, f64, f64) {
1233 let g = gam_linalg::utils::stable_logistic((z - tau) / smoothing_eps);
1234 let value = if z > tau { z } else { 0.0 };
1235 let slope = z * g * (1.0 - g) / smoothing_eps;
1236 let dphi_dz = g + slope;
1237 let dphi_dtau = -slope;
1238 (value, dphi_dz, dphi_dtau)
1239}
1240
1241impl AnalyticPenalty for JumpReLUPenalty {
1242 fn tier(&self) -> PenaltyTier {
1243 PenaltyTier::Psi
1244 }
1245
1246 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
1247 let d = self.latent_dim;
1248 let n_obs = target.len() / d;
1249 let mut acc = 0.0;
1250 for row in 0..n_obs {
1251 let base = row * d;
1252 for axis in 0..d {
1253 let tau = self.threshold(axis, rho);
1254 let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
1255 acc += self.weight * tau * gate;
1256 }
1257 }
1258 acc
1259 }
1260
1261 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1262 let d = self.latent_dim;
1263 let n_obs = target.len() / d;
1264 let mut grad = Array1::<f64>::zeros(target.len());
1265 for row in 0..n_obs {
1266 let base = row * d;
1267 for axis in 0..d {
1268 let tau = self.threshold(axis, rho);
1269 let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
1270 grad[base + axis] = self.weight * tau * gate * (1.0 - gate) / self.smoothing_eps;
1271 }
1272 }
1273 grad
1274 }
1275
1276 fn hessian_diag(
1277 &self,
1278 target: ArrayView1<'_, f64>,
1279 rho: ArrayView1<'_, f64>,
1280 ) -> Option<Array1<f64>> {
1281 let d = self.latent_dim;
1282 let n_obs = target.len() / d;
1283 let mut diag = Array1::<f64>::zeros(target.len());
1284 for row in 0..n_obs {
1285 let base = row * d;
1286 for axis in 0..d {
1287 let tau = self.threshold(axis, rho);
1288 let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
1289 diag[base + axis] = self.true_hessian_diag_entry(tau, gate);
1290 }
1291 }
1292 Some(diag)
1293 }
1294
1295 fn hvp(
1296 &self,
1297 target: ArrayView1<'_, f64>,
1298 rho: ArrayView1<'_, f64>,
1299 v: ArrayView1<'_, f64>,
1300 ) -> Array1<f64> {
1301 assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
1302 let d = self.latent_dim;
1303 let n_obs = target.len() / d;
1304 let mut out = Array1::<f64>::zeros(target.len());
1305 for row in 0..n_obs {
1306 let base = row * d;
1307 for axis in 0..d {
1308 let tau = self.threshold(axis, rho);
1309 let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
1310 out[base + axis] = self.true_hessian_diag_entry(tau, gate) * v[base + axis];
1311 }
1312 }
1313 out
1314 }
1315
1316 fn psd_majorizer_diag(
1317 &self,
1318 target: ArrayView1<'_, f64>,
1319 rho: ArrayView1<'_, f64>,
1320 ) -> Option<Array1<f64>> {
1321 let d = self.latent_dim;
1329 let n_obs = target.len() / d;
1330 let mut diag = Array1::<f64>::zeros(target.len());
1331 for row in 0..n_obs {
1332 let base = row * d;
1333 for axis in 0..d {
1334 let tau = self.threshold(axis, rho);
1335 let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
1336 diag[base + axis] = self.psd_hessian_diag_entry(tau, gate);
1337 }
1338 }
1339 Some(diag)
1340 }
1341
1342 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1343 let d = self.latent_dim;
1344 let n_obs = target.len() / d;
1345 let mut out = Array1::<f64>::zeros(d);
1346 for axis in 0..d {
1347 let tau = self.threshold(axis, rho);
1348 let mut g_tau = 0.0;
1349 for row in 0..n_obs {
1350 let x = target[row * d + axis];
1351 let gate = self.sigmoid_gate((x - tau) / self.smoothing_eps);
1352 g_tau += gate - tau * gate * (1.0 - gate) / self.smoothing_eps;
1353 }
1354 out[axis] = self.weight * tau * g_tau;
1355 }
1356 out
1357 }
1358
1359 fn rho_count(&self) -> usize {
1360 self.latent_dim
1361 }
1362
1363 fn name(&self) -> &str {
1364 "jumprelu"
1365 }
1366
1367 impl_scalar_apply_schedule!(weight);
1368}
1369
1370#[cfg(test)]
1371mod fisher_majorizer_1419_tests {
1372 use super::*;
1373 use approx::assert_abs_diff_eq;
1374 use gam_linalg::faer_ndarray::FaerEigh;
1375 use ndarray::Array2;
1376
1377 #[test]
1391 fn gershgorin_majorizes_entropy_where_fisher_does_not_1419() {
1392 let temperature = 1.0_f64;
1395 let scale = 1.0_f64; let pen = SoftmaxAssignmentSparsityPenalty::new(2, temperature);
1397 let z1 = 0.0_f64;
1398 let z0 = z1 + (0.95_f64 / 0.05_f64).ln();
1399 let row = [z0, z1];
1400
1401 let a = pen.softmax_row(&row);
1403 assert_abs_diff_eq!(a[0], 0.95, epsilon = 1e-12);
1404 assert_abs_diff_eq!(a[1], 0.05, epsilon = 1e-12);
1405
1406 let h = pen.row_dense_hessian(&row, scale);
1408 let g = pen.row_fisher_metric(&row, scale);
1409 let m = pen.row_psd_majorizer(&row, scale);
1410
1411 assert_abs_diff_eq!(h[[0, 0]], 0.0783747664, epsilon = 1e-9);
1414 assert_abs_diff_eq!(g[[0, 0]], 0.95 * 0.05, epsilon = 1e-12);
1415
1416 for kk in 0..2 {
1418 let row_sum: f64 = (0..2).map(|jj| h[[kk, jj]].abs()).sum();
1419 assert_abs_diff_eq!(m[[kk, kk]], row_sum, epsilon = 1e-12);
1420 }
1421 assert_abs_diff_eq!(m[[0, 1]], 0.0, epsilon = 1e-15);
1423 assert_abs_diff_eq!(m[[1, 0]], 0.0, epsilon = 1e-15);
1424 assert!(m[[0, 0]] >= 0.0 && m[[1, 1]] >= 0.0);
1425
1426 let fisher_free = g[[0, 0]] - h[[0, 0]];
1430 let major_free = m[[0, 0]] - h[[0, 0]];
1431 assert!(
1432 fisher_free < -1e-3,
1433 "Fisher must FAIL the majorizer bound in the free direction (#1419); \
1434 G_11 − H_11 = {fisher_free}"
1435 );
1436 assert!(
1437 major_free >= -1e-12,
1438 "Gershgorin majorizer must SATISFY the bound in the free direction (#1419); \
1439 D_11 − H_11 = {major_free}"
1440 );
1441
1442 let mut m_minus_h = Array2::<f64>::zeros((2, 2));
1446 let mut g_minus_h = Array2::<f64>::zeros((2, 2));
1447 for i in 0..2 {
1448 for j in 0..2 {
1449 m_minus_h[[i, j]] = m[[i, j]] - h[[i, j]];
1450 g_minus_h[[i, j]] = g[[i, j]] - h[[i, j]];
1451 }
1452 }
1453 let (m_evals, _) = m_minus_h.eigh(faer::Side::Lower).expect("eigh(M−H)");
1454 let (g_evals, _) = g_minus_h.eigh(faer::Side::Lower).expect("eigh(G−H)");
1455 let m_min = m_evals.iter().cloned().fold(f64::INFINITY, f64::min);
1456 let g_min = g_evals.iter().cloned().fold(f64::INFINITY, f64::min);
1457 assert!(
1458 m_min >= -1e-12,
1459 "Gershgorin majorizer must be a Loewner majorizer (M − H ⪰ 0, #1419); \
1460 smallest eigenvalue of M−H = {m_min}"
1461 );
1462 assert!(
1463 g_min < -1e-9,
1464 "the OLD Fisher metric must FAIL the Loewner majorizer test (#1419); \
1465 smallest eigenvalue of G−H = {g_min} (expected strictly negative)"
1466 );
1467 }
1468
1469 #[test]
1476 fn gershgorin_majorizer_logit_derivative_matches_fd_1419() {
1477 let pen = SoftmaxAssignmentSparsityPenalty::new(4, 0.8);
1478 let row = [0.3_f64, -0.6, 0.9, 0.2];
1479 let scale = 1.1_f64 * (1.0 / 0.8_f64) * (1.0 / 0.8_f64);
1480 let eps = 1e-6;
1481 for w in 0..4 {
1482 let dd = pen.row_psd_majorizer_logit_derivative(&row, scale, w);
1483 let mut rp = row;
1484 let mut rm = row;
1485 rp[w] += eps;
1486 rm[w] -= eps;
1487 let mp = pen.row_psd_majorizer(&rp, scale);
1488 let mm = pen.row_psd_majorizer(&rm, scale);
1489 for k in 0..4 {
1490 let fd = (mp[[k, k]] - mm[[k, k]]) / (2.0 * eps);
1491 assert_abs_diff_eq!(dd[[k, k]], fd, epsilon = 1e-6);
1492 }
1493 for i in 0..4 {
1495 for j in 0..4 {
1496 if i != j {
1497 assert_abs_diff_eq!(dd[[i, j]], 0.0, epsilon = 1e-15);
1498 }
1499 }
1500 }
1501 }
1502 }
1503}
1504
1505#[cfg(test)]
1506mod row_weighted_prior_991_tests {
1507 use super::AnalyticPenalty;
1514 use super::*;
1515 use approx::assert_abs_diff_eq;
1516 use ndarray::{Array1, s};
1517
1518 fn logits(n: usize, k: usize) -> Array1<f64> {
1519 let mut v = Array1::<f64>::zeros(n * k);
1522 for r in 0..n {
1523 for a in 0..k {
1524 v[r * k + a] =
1525 0.35 * (r as f64) - 0.6 * (a as f64) + 0.11 * ((r * k + a) as f64).sin();
1526 }
1527 }
1528 v
1529 }
1530
1531 #[test]
1535 fn weighted_value_is_per_row_reweight_of_unweighted() {
1536 let (n, k) = (5usize, 3usize);
1537 let temperature = 0.7_f64;
1538 let rho = Array1::from_vec(vec![0.2_f64]);
1539 let target = logits(n, k);
1540 let base = SoftmaxAssignmentSparsityPenalty::new(k, temperature);
1541 let mut per_row = vec![0.0_f64; n];
1543 for r in 0..n {
1544 let row = target.slice(s![r * k..r * k + k]).to_owned();
1545 per_row[r] = base.value(row.view(), rho.view());
1546 }
1547 let unweighted: f64 = per_row.iter().sum();
1548 assert_abs_diff_eq!(
1549 base.value(target.view(), rho.view()),
1550 unweighted,
1551 epsilon = 1e-12
1552 );
1553
1554 let w = vec![1.7_f64, 0.3, 1.1, 0.5, 1.4]; let weighted = base.clone().with_row_weights(Some(&w));
1556 let expect: f64 = (0..n).map(|r| w[r] * per_row[r]).sum();
1557 assert_abs_diff_eq!(
1558 weighted.value(target.view(), rho.view()),
1559 expect,
1560 epsilon = 1e-12
1561 );
1562 assert_abs_diff_eq!(
1565 weighted.value(target.view(), rho.view()),
1566 (0..n).map(|r| w[r] * per_row[r]).sum::<f64>(),
1567 epsilon = 1e-12
1568 );
1569 }
1570
1571 #[test]
1576 fn weighted_value_grad_are_fd_consistent() {
1577 let (n, k) = (4usize, 3usize);
1578 let temperature = 0.9_f64;
1579 let rho = Array1::from_vec(vec![-0.1_f64]);
1580 let target = logits(n, k);
1581 let w = vec![1.9_f64, 0.4, 0.8, 0.9];
1582 let pen = SoftmaxAssignmentSparsityPenalty::new(k, temperature).with_row_weights(Some(&w));
1583 let grad = pen.grad_target(target.view(), rho.view());
1584 let eps = 1e-6;
1585 for idx in 0..n * k {
1586 let mut plus = target.clone();
1587 let mut minus = target.clone();
1588 plus[idx] += eps;
1589 minus[idx] -= eps;
1590 let fd = (pen.value(plus.view(), rho.view()) - pen.value(minus.view(), rho.view()))
1591 / (2.0 * eps);
1592 assert_abs_diff_eq!(grad[idx], fd, epsilon = 1e-7);
1593 }
1594 }
1595
1596 #[test]
1600 fn every_channel_scales_by_w_row_identically() {
1601 let (n, k) = (4usize, 3usize);
1602 let temperature = 0.8_f64;
1603 let rho = Array1::from_vec(vec![0.15_f64]);
1604 let target = logits(n, k);
1605 let v = logits(n, k); let w = vec![1.6_f64, 0.25, 1.05, 1.1];
1607 let base = SoftmaxAssignmentSparsityPenalty::new(k, temperature);
1608 let wtd = base.clone().with_row_weights(Some(&w));
1609
1610 let g0 = base.grad_target(target.view(), rho.view());
1611 let g1 = wtd.grad_target(target.view(), rho.view());
1612 let d0 = base.hessian_diag(target.view(), rho.view()).unwrap();
1613 let d1 = wtd.hessian_diag(target.view(), rho.view()).unwrap();
1614 let m0 = base.psd_majorizer_diag(target.view(), rho.view()).unwrap();
1615 let m1 = wtd.psd_majorizer_diag(target.view(), rho.view()).unwrap();
1616 let h0 = base.hvp(target.view(), rho.view(), v.view());
1617 let h1 = wtd.hvp(target.view(), rho.view(), v.view());
1618 for r in 0..n {
1619 for a in 0..k {
1620 let i = r * k + a;
1621 assert_abs_diff_eq!(g1[i], w[r] * g0[i], epsilon = 1e-12);
1622 assert_abs_diff_eq!(d1[i], w[r] * d0[i], epsilon = 1e-12);
1623 assert_abs_diff_eq!(m1[i], w[r] * m0[i], epsilon = 1e-12);
1624 assert_abs_diff_eq!(h1[i], w[r] * h0[i], epsilon = 1e-12);
1625 }
1626 }
1627 let r0 = base.grad_rho(target.view(), rho.view())[0];
1629 let r1 = wtd.grad_rho(target.view(), rho.view())[0];
1630 let expect: f64 = (0..n)
1631 .map(|r| {
1632 let row = target.slice(s![r * k..r * k + k]).to_owned();
1633 w[r] * base.value(row.view(), rho.view())
1634 })
1635 .sum();
1636 assert_abs_diff_eq!(r1, expect, epsilon = 1e-12);
1637 assert!(r0.is_finite());
1638 }
1639
1640 #[test]
1642 fn none_weights_are_bit_for_bit_unweighted() {
1643 let (n, k) = (3usize, 4usize);
1644 let rho = Array1::from_vec(vec![0.0_f64]);
1645 let target = logits(n, k);
1646 let base = SoftmaxAssignmentSparsityPenalty::new(k, 1.0);
1647 let none = base.clone().with_row_weights(None);
1648 assert_eq!(
1649 base.value(target.view(), rho.view()).to_bits(),
1650 none.value(target.view(), rho.view()).to_bits()
1651 );
1652 let g0 = base.grad_target(target.view(), rho.view());
1653 let g1 = none.grad_target(target.view(), rho.view());
1654 for i in 0..n * k {
1655 assert_eq!(g0[i].to_bits(), g1[i].to_bits());
1656 }
1657 }
1658}