1use super::{unflatten_named, PolicyNetwork, RLOptimizationMetrics};
7use crate::error::{OptimError, Result};
8use scirs2_core::ndarray::{Array1, Array2, ScalarOperand};
9use scirs2_core::numeric::Float;
10use std::fmt::Debug;
11
12fn tiny<T: Float>() -> T {
18 T::from(1e-30).unwrap_or_else(T::epsilon)
19}
20
21type SurrogateFn<'a, P, T> = dyn FnMut(&P) -> Result<T> + 'a;
25
26#[derive(Debug, Clone, Copy)]
28pub enum TrustRegionMethod {
29 TRPO,
31
32 CPO,
34
35 Projection,
37
38 NaturalGradient,
40}
41
42#[derive(Debug, Clone)]
44pub struct TrustRegionConfig<T: Float + Debug + Send + Sync + 'static> {
45 pub method: TrustRegionMethod,
47
48 pub max_kl: T,
50
51 pub cg_iters: usize,
53 pub cg_damping: T,
54 pub cg_tolerance: T,
55
56 pub max_backtracks: usize,
58 pub backtrack_coeff: T,
59 pub accept_ratio: T,
60
61 pub fisher_subsample_freq: usize,
63 pub fisher_reg: T,
64}
65
66impl<T: Float + Debug + Send + Sync + 'static> Default for TrustRegionConfig<T> {
67 fn default() -> Self {
68 Self {
69 method: TrustRegionMethod::TRPO,
70 max_kl: T::from(0.01).unwrap_or_else(|| T::zero()),
71 cg_iters: 10,
72 cg_damping: T::from(0.1).unwrap_or_else(|| T::zero()),
73 cg_tolerance: T::from(1e-8).unwrap_or_else(|| T::zero()),
74 max_backtracks: 10,
75 backtrack_coeff: T::from(0.5).unwrap_or_else(|| T::zero()),
76 accept_ratio: T::from(0.1).unwrap_or_else(|| T::zero()),
77 fisher_subsample_freq: 1,
78 fisher_reg: T::from(1e-5).unwrap_or_else(|| T::zero()),
79 }
80 }
81}
82
83pub struct TrustRegionOptimizer<T: Float + Debug + Send + Sync + 'static, P: PolicyNetwork<T>> {
85 config: TrustRegionConfig<T>,
87
88 policy: P,
90
91 score_samples: Option<Array2<T>>,
99
100 cost_constraint: Option<(Array1<T>, T)>,
106
107 natural_grad_state: NaturalGradientState<T>,
109
110 update_count: usize,
112}
113
114#[derive(Debug, Clone)]
116pub struct TrustRegionStepReport<T: Float + Debug + Send + Sync + 'static> {
117 pub accepted: bool,
119
120 pub step_scale: T,
122
123 pub kl: T,
125
126 pub surrogate_improvement: T,
128
129 pub backtracks: usize,
131}
132
133#[derive(Debug, Clone)]
135pub struct NaturalGradientState<T: Float + Debug + Send + Sync + 'static> {
136 pub prev_gradients: Option<Array1<T>>,
138
139 pub momentum: T,
141
142 pub adaptive_lr_state: AdaptiveLRState<T>,
144}
145
146#[derive(Debug, Clone)]
148pub struct AdaptiveLRState<T: Float + Debug + Send + Sync + 'static> {
149 pub learning_rate: T,
151
152 pub adapt_factor: T,
154
155 pub success_count: usize,
157
158 pub failure_count: usize,
160}
161
162impl<
163 T: Float + Debug + Send + Sync + std::iter::Sum + ScalarOperand + 'static,
164 P: PolicyNetwork<T>,
165 > TrustRegionOptimizer<T, P>
166{
167 pub fn new(config: TrustRegionConfig<T>, policy: P) -> Self {
169 Self {
170 config,
171 policy,
172 score_samples: None,
173 cost_constraint: None,
174 natural_grad_state: NaturalGradientState {
175 prev_gradients: None,
176 momentum: T::from(0.9).unwrap_or_else(|| T::zero()),
177 adaptive_lr_state: AdaptiveLRState {
178 learning_rate: T::from(0.01).unwrap_or_else(|| T::zero()),
179 adapt_factor: T::from(1.5).unwrap_or_else(|| T::zero()),
180 success_count: 0,
181 failure_count: 0,
182 },
183 },
184 update_count: 0,
185 }
186 }
187
188 pub fn set_score_samples(&mut self, samples: Array2<T>) {
196 self.score_samples = Some(samples);
197 }
198
199 pub fn clear_score_samples(&mut self) {
202 self.score_samples = None;
203 }
204
205 pub fn set_cost_constraint(&mut self, cost_gradient: Array1<T>, cost_surplus: T) {
215 self.cost_constraint = Some((cost_gradient, cost_surplus));
216 }
217
218 pub fn clear_cost_constraint(&mut self) {
220 self.cost_constraint = None;
221 }
222
223 pub fn update(&mut self, gradients: &Array1<T>) -> Result<RLOptimizationMetrics<T>> {
229 match self.config.method {
230 TrustRegionMethod::TRPO => self.update_trpo(gradients),
231 TrustRegionMethod::CPO => self.update_cpo(gradients),
232 TrustRegionMethod::Projection => self.update_projection(gradients),
233 TrustRegionMethod::NaturalGradient => self.update_natural_gradient(gradients),
234 }
235 }
236
237 fn update_trpo(&mut self, gradients: &Array1<T>) -> Result<RLOptimizationMetrics<T>> {
246 let report = self.trpo_step(gradients, None::<&mut SurrogateFn<'_, P, T>>)?;
247 self.update_count += 1;
248 Ok(Self::metrics_from_report(&report))
249 }
250
251 pub fn update_trpo_with_surrogate<F>(
265 &mut self,
266 gradients: &Array1<T>,
267 mut surrogate: F,
268 ) -> Result<RLOptimizationMetrics<T>>
269 where
270 F: FnMut(&P) -> Result<T>,
271 {
272 let report = self.trpo_step(gradients, Some(&mut surrogate))?;
273 self.update_count += 1;
274 Ok(Self::metrics_from_report(&report))
275 }
276
277 fn metrics_from_report(report: &TrustRegionStepReport<T>) -> RLOptimizationMetrics<T> {
278 let mut metrics = RLOptimizationMetrics {
279 kl_divergence: Some(report.kl),
280 ..Default::default()
281 };
282 metrics.policy_loss = -report.surrogate_improvement;
283 metrics
284 .custom_metrics
285 .insert("step_scale".to_string(), report.step_scale);
286 metrics.custom_metrics.insert(
287 "line_search_accepted".to_string(),
288 if report.accepted { T::one() } else { T::zero() },
289 );
290 metrics
291 }
292
293 fn trpo_step(
296 &mut self,
297 gradients: &Array1<T>,
298 surrogate: Option<&mut SurrogateFn<'_, P, T>>,
299 ) -> Result<TrustRegionStepReport<T>> {
300 let natural_grad = self.compute_natural_gradient(gradients)?;
302 let fvp = self.fisher_vector_product(&natural_grad)?;
303 let shs = self.dot(&natural_grad, &fvp);
304
305 if !matches!(
308 shs.partial_cmp(&tiny::<T>()),
309 Some(std::cmp::Ordering::Greater)
310 ) {
311 return Ok(TrustRegionStepReport {
312 accepted: false,
313 step_scale: T::zero(),
314 kl: T::zero(),
315 surrogate_improvement: T::zero(),
316 backtracks: 0,
317 });
318 }
319
320 let two = T::from(2.0).unwrap_or_else(|| T::one() + T::one());
322 let beta = (two * self.config.max_kl / shs).sqrt();
323 if !beta.is_finite() {
324 return Err(OptimError::ComputationError(
325 "TRPO step size sqrt(2δ/sᵀFs) is not finite".to_string(),
326 ));
327 }
328 let full_step = &natural_grad * beta;
329
330 self.line_search(gradients, &full_step, surrogate, None)
331 }
332
333 fn update_cpo(&mut self, gradients: &Array1<T>) -> Result<RLOptimizationMetrics<T>> {
356 let (cost_gradient, cost_surplus) =
357 match self.cost_constraint.clone() {
358 Some(pair) => pair,
359 None => return Err(OptimError::UnsupportedOperation(
360 "CPO requires a safety constraint: call set_cost_constraint(cost_gradient, \
361 cost_surplus) before update(), or select TrustRegionMethod::TRPO for the \
362 unconstrained problem"
363 .to_string(),
364 )),
365 };
366
367 if cost_gradient.len() != gradients.len() {
368 return Err(OptimError::DimensionMismatch(format!(
369 "cost gradient length ({}) does not match objective gradient length ({})",
370 cost_gradient.len(),
371 gradients.len()
372 )));
373 }
374
375 let two = T::from(2.0).unwrap_or_else(|| T::one() + T::one());
376 let delta = self.config.max_kl;
377
378 let hinv_g = self.conjugate_gradient(gradients)?;
379 let hinv_b = self.conjugate_gradient(&cost_gradient)?;
380
381 let q = self.dot(gradients, &hinv_g);
382 let r = self.dot(gradients, &hinv_b);
383 let s = self.dot(&cost_gradient, &hinv_b);
384
385 if !matches!(
387 s.partial_cmp(&tiny::<T>()),
388 Some(std::cmp::Ordering::Greater)
389 ) {
390 let report = self.trpo_step(gradients, None::<&mut SurrogateFn<'_, P, T>>)?;
391 self.update_count += 1;
392 return Ok(Self::metrics_from_report(&report));
393 }
394
395 let c = cost_surplus;
396 let b_coeff = two * delta - c * c / s;
397
398 let step = if c > T::zero()
399 && !matches!(
400 b_coeff.partial_cmp(&T::zero()),
401 Some(std::cmp::Ordering::Greater)
402 ) {
403 let scale = (two * delta / s).sqrt();
405 &hinv_b * (-scale)
406 } else {
407 let a_coeff = q - r * r / s;
408 let mut lambda = if a_coeff > T::zero() && b_coeff > T::zero() {
409 (a_coeff / b_coeff).sqrt()
410 } else {
411 (q / (two * delta)).sqrt()
412 };
413 let mut nu = (r + lambda * c) / s;
414 if nu < T::zero() {
415 nu = T::zero();
417 lambda = (q / (two * delta)).sqrt();
418 }
419 if !matches!(
420 lambda.partial_cmp(&tiny::<T>()),
421 Some(std::cmp::Ordering::Greater)
422 ) || !lambda.is_finite()
423 {
424 return Ok(Self::metrics_from_report(&TrustRegionStepReport {
425 accepted: false,
426 step_scale: T::zero(),
427 kl: T::zero(),
428 surrogate_improvement: T::zero(),
429 backtracks: 0,
430 }));
431 }
432 (&hinv_g - &(&hinv_b * nu)) / lambda
433 };
434
435 let report = self.line_search(
436 gradients,
437 &step,
438 None::<&mut SurrogateFn<'_, P, T>>,
439 Some((&cost_gradient, c)),
440 )?;
441 self.update_count += 1;
442
443 let mut metrics = Self::metrics_from_report(&report);
444 metrics.custom_metrics.insert("cost_surplus".to_string(), c);
445 Ok(metrics)
446 }
447
448 fn update_projection(&mut self, gradients: &Array1<T>) -> Result<RLOptimizationMetrics<T>> {
450 let projected_grad = self.project_to_trust_region(gradients)?;
452 self.apply_parameter_update(&projected_grad)?;
453
454 Ok(RLOptimizationMetrics::default())
455 }
456
457 fn update_natural_gradient(
459 &mut self,
460 gradients: &Array1<T>,
461 ) -> Result<RLOptimizationMetrics<T>> {
462 let natural_grad = self.compute_natural_gradient(gradients)?;
463 let lr = self.natural_grad_state.adaptive_lr_state.learning_rate;
464 let update_step = &natural_grad * lr;
465
466 self.apply_parameter_update(&update_step)?;
467
468 Ok(RLOptimizationMetrics::default())
469 }
470
471 fn compute_natural_gradient(&mut self, gradients: &Array1<T>) -> Result<Array1<T>> {
473 self.conjugate_gradient(gradients)
475 }
476
477 fn conjugate_gradient(&self, b: &Array1<T>) -> Result<Array1<T>> {
486 let n = b.len();
487 let mut x = Array1::zeros(n);
488 let mut r = b.clone();
489 let mut p = r.clone();
490 let mut rsold = self.dot(&r, &r);
491
492 if !matches!(
494 rsold.partial_cmp(&tiny::<T>()),
495 Some(std::cmp::Ordering::Greater)
496 ) {
497 return Ok(x);
498 }
499
500 for _i in 0..self.config.cg_iters {
501 let ap = self.fisher_vector_product(&p)?;
502 let pap = self.dot(&p, &ap);
503
504 if !matches!(
506 pap.abs().partial_cmp(&tiny::<T>()),
507 Some(std::cmp::Ordering::Greater)
508 ) || !pap.is_finite()
509 {
510 break;
511 }
512
513 let alpha = rsold / pap;
514
515 x = &x + &(&p * alpha);
516 r = &r - &(&ap * alpha);
517
518 let rsnew = self.dot(&r, &r);
519
520 if rsnew.sqrt() < self.config.cg_tolerance {
521 break;
522 }
523 if !matches!(
524 rsnew.partial_cmp(&tiny::<T>()),
525 Some(std::cmp::Ordering::Greater)
526 ) {
527 break;
528 }
529
530 let beta = rsnew / rsold;
531 p = &r + &(&p * beta);
532 rsold = rsnew;
533 }
534
535 Ok(x)
536 }
537
538 fn fisher_vector_product(&self, v: &Array1<T>) -> Result<Array1<T>> {
558 let damping = self.config.cg_damping;
560
561 match &self.score_samples {
562 Some(samples) if samples.nrows() > 0 => {
563 let n_samples = samples.nrows();
564 let dim = samples.ncols();
565
566 if dim != v.len() {
567 return Err(OptimError::DimensionMismatch(format!(
568 "Score sample dimension ({}) does not match vector dimension ({})",
569 dim,
570 v.len()
571 )));
572 }
573
574 let mut accum: Array1<T> = Array1::zeros(dim);
576 for row in samples.rows() {
577 let proj: T = row.iter().zip(v.iter()).map(|(&g, &x)| g * x).sum();
579 for (acc, &g) in accum.iter_mut().zip(row.iter()) {
581 *acc = *acc + g * proj;
582 }
583 }
584
585 let inv_n = T::one()
586 / T::from(n_samples).ok_or_else(|| {
587 OptimError::ComputationError(
588 "Failed to convert sample count to scalar type".to_string(),
589 )
590 })?;
591 accum.mapv_inplace(|x| x * inv_n);
592
593 let ridge = self.config.fisher_reg + damping;
595 Ok(&accum + &(v * ridge))
596 }
597 _ => Ok(v + &(v * damping)),
599 }
600 }
601
602 fn line_search(
616 &mut self,
617 gradients: &Array1<T>,
618 full_step: &Array1<T>,
619 mut surrogate: Option<&mut SurrogateFn<'_, P, T>>,
620 cost_constraint: Option<(&Array1<T>, T)>,
621 ) -> Result<TrustRegionStepReport<T>> {
622 let base = match surrogate {
624 Some(ref mut f) => f(&self.policy)?,
625 None => T::zero(),
626 };
627
628 let mut scale = T::one();
629 for attempt in 0..self.config.max_backtracks.max(1) {
630 let step = full_step * scale;
631
632 let kl = self.estimate_kl_divergence(&step, T::one())?;
634 let expected = self.dot(gradients, &step);
635
636 let cost_ok = match cost_constraint {
638 Some((cost_gradient, surplus)) => {
639 surplus + self.dot(cost_gradient, &step) <= T::zero()
640 }
641 None => true,
642 };
643
644 self.apply_parameter_update(&step)?;
645
646 let value = match surrogate {
647 Some(ref mut f) => f(&self.policy)?,
648 None => expected - kl,
650 };
651 let actual = value - base;
652
653 let ratio = if expected > tiny::<T>() {
654 actual / expected
655 } else {
656 T::neg_infinity()
657 };
658
659 let accept = kl <= self.config.max_kl
660 && cost_ok
661 && actual > T::zero()
662 && ratio > self.config.accept_ratio;
663
664 if accept {
665 self.natural_grad_state.adaptive_lr_state.success_count += 1;
666 return Ok(TrustRegionStepReport {
667 accepted: true,
668 step_scale: scale,
669 kl,
670 surrogate_improvement: actual,
671 backtracks: attempt,
672 });
673 }
674
675 self.apply_parameter_update(&(&step * -T::one()))?;
677 scale = scale * self.config.backtrack_coeff;
678 }
679
680 self.natural_grad_state.adaptive_lr_state.failure_count += 1;
682 Ok(TrustRegionStepReport {
683 accepted: false,
684 step_scale: T::zero(),
685 kl: T::zero(),
686 surrogate_improvement: T::zero(),
687 backtracks: self.config.max_backtracks.max(1),
688 })
689 }
690
691 fn estimate_kl_divergence(&self, direction: &Array1<T>, stepsize: T) -> Result<T> {
693 let fvp = self.fisher_vector_product(direction)?;
695 let kl_estimate = T::from(0.5).unwrap_or_else(|| T::zero())
696 * self.dot(direction, &fvp)
697 * stepsize
698 * stepsize;
699 Ok(kl_estimate)
700 }
701
702 fn project_to_trust_region(&self, gradients: &Array1<T>) -> Result<Array1<T>> {
704 let grad_norm = self.norm(gradients);
705 let max_norm = (T::from(2.0).unwrap_or_else(|| T::zero()) * self.config.max_kl).sqrt();
706
707 if grad_norm <= max_norm {
708 Ok(gradients.clone())
709 } else {
710 Ok(gradients * (max_norm / grad_norm))
711 }
712 }
713
714 fn apply_parameter_update(&mut self, update: &Array1<T>) -> Result<()> {
724 let params = self.policy.get_parameters();
725 let deltas = unflatten_named(¶ms, update)?;
726 self.policy.update_parameters(&deltas)
727 }
728
729 fn dot(&self, a: &Array1<T>, b: &Array1<T>) -> T {
731 a.iter().zip(b.iter()).map(|(&x, &y)| x * y).sum()
732 }
733
734 fn norm(&self, v: &Array1<T>) -> T {
736 self.dot(v, v).sqrt()
737 }
738}
739
740#[cfg(test)]
741mod tests {
742 use super::super::{ActionDistribution, DistributionType, PolicyEvaluation};
743 use super::*;
744 use approx::assert_abs_diff_eq;
745 use scirs2_core::ndarray::{arr1, arr2};
746 use std::cell::RefCell;
747 use std::collections::HashMap;
748
749 struct MockPolicy {
756 params: HashMap<String, Array1<f64>>,
757 last_gradients: RefCell<Option<HashMap<String, Array1<f64>>>>,
758 }
759
760 impl MockPolicy {
761 fn new() -> Self {
762 let mut params = HashMap::new();
763 params.insert("w".to_string(), arr1(&[0.0, 0.0, 0.0]));
764 Self {
765 params,
766 last_gradients: RefCell::new(None),
767 }
768 }
769 }
770
771 impl PolicyNetwork<f64> for MockPolicy {
772 fn evaluate_actions(
773 &self,
774 _observations: &Array2<f64>,
775 _actions: &Array2<f64>,
776 ) -> Result<PolicyEvaluation<f64>> {
777 Ok(PolicyEvaluation {
778 log_probs: arr1(&[0.0]),
779 entropy: arr1(&[0.0]),
780 metrics: HashMap::new(),
781 })
782 }
783
784 fn get_action_distribution(
785 &self,
786 _observations: &Array2<f64>,
787 ) -> Result<ActionDistribution<f64>> {
788 Ok(ActionDistribution {
789 mean: None,
790 std: None,
791 logits: None,
792 distribution_type: DistributionType::Gaussian,
793 })
794 }
795
796 fn update_parameters(&mut self, gradients: &HashMap<String, Array1<f64>>) -> Result<()> {
797 for (key, grad) in gradients {
799 if let Some(p) = self.params.get_mut(key) {
800 *p = &*p + grad;
801 }
802 }
803 *self.last_gradients.borrow_mut() = Some(gradients.clone());
804 Ok(())
805 }
806
807 fn get_parameters(&self) -> HashMap<String, Array1<f64>> {
808 self.params.clone()
809 }
810 }
811
812 fn make_optimizer(cg_damping: f64) -> TrustRegionOptimizer<f64, MockPolicy> {
813 let config = TrustRegionConfig::<f64> {
815 cg_damping,
816 fisher_reg: 0.0,
817 ..Default::default()
818 };
819 TrustRegionOptimizer::new(config, MockPolicy::new())
820 }
821
822 fn reference_fvp(samples: &Array2<f64>, v: &Array1<f64>, damping: f64) -> Array1<f64> {
824 let n = samples.nrows();
825 let dim = samples.ncols();
826 let mut out = Array1::<f64>::zeros(dim);
827 for row in samples.rows() {
828 let proj: f64 = row.iter().zip(v.iter()).map(|(&g, &x)| g * x).sum();
829 for (o, &g) in out.iter_mut().zip(row.iter()) {
830 *o += g * proj;
831 }
832 }
833 out.mapv_inplace(|x| x / n as f64);
834 &out + &(v * damping)
835 }
836
837 #[test]
838 fn test_fisher_vector_product_matches_empirical_formula() {
839 let damping = 0.1;
840 let mut opt = make_optimizer(damping);
841
842 let samples = arr2(&[[1.0, 2.0, 3.0], [0.5, -1.0, 2.0]]);
844 opt.set_score_samples(samples.clone());
845
846 let v = arr1(&[0.3, -0.7, 1.1]);
847 let got = opt
848 .fisher_vector_product(&v)
849 .expect("fisher-vector product");
850 let expected = reference_fvp(&samples, &v, damping);
851
852 assert_eq!(got.len(), expected.len());
853 for (g, e) in got.iter().zip(expected.iter()) {
854 assert_abs_diff_eq!(*g, *e, epsilon = 1e-10);
855 }
856 }
857
858 #[test]
859 fn test_fisher_vector_product_identity_fallback() {
860 let damping = 0.1;
861 let opt = make_optimizer(damping);
862 let v = arr1(&[1.0, -2.0, 4.0]);
864 let got = opt
865 .fisher_vector_product(&v)
866 .expect("fisher-vector product");
867 let expected = &v + &(&v * damping);
868 for (g, e) in got.iter().zip(expected.iter()) {
869 assert_abs_diff_eq!(*g, *e, epsilon = 1e-12);
870 }
871
872 let mut opt2 = make_optimizer(damping);
874 opt2.set_score_samples(Array2::<f64>::zeros((0, 3)));
875 let got2 = opt2
876 .fisher_vector_product(&v)
877 .expect("fisher-vector product");
878 for (g, e) in got2.iter().zip(expected.iter()) {
879 assert_abs_diff_eq!(*g, *e, epsilon = 1e-12);
880 }
881 }
882
883 #[test]
884 fn test_conjugate_gradient_solves_damped_system() {
885 let damping = 0.5;
886 let mut opt = make_optimizer(damping);
887 let samples = arr2(&[[1.0, 0.5, -0.3], [0.2, 1.5, 0.7], [-0.5, 0.1, 1.2]]);
888 opt.set_score_samples(samples.clone());
889
890 let b = arr1(&[1.0, -2.0, 0.5]);
891 let x = opt.conjugate_gradient(&b).expect("conjugate gradient");
892
893 let ax = opt
896 .fisher_vector_product(&x)
897 .expect("fisher-vector product");
898 let residual: f64 = ax
899 .iter()
900 .zip(b.iter())
901 .map(|(&a, &bv)| (a - bv) * (a - bv))
902 .sum::<f64>()
903 .sqrt();
904 assert!(
905 residual < 1e-6,
906 "CG residual too large: {residual} (x = {x:?})"
907 );
908 }
909
910 #[test]
911 fn test_apply_parameter_update_forwards_split_gradient() {
912 let damping = 0.1;
913 let mut opt = make_optimizer(damping);
914
915 let update = arr1(&[0.1, 0.2, 0.3]);
917 opt.apply_parameter_update(&update).expect("apply update");
918
919 let recorded = opt.policy.last_gradients.borrow();
921 let map = recorded.as_ref().expect("update_parameters was not called");
922 let w_grad = map.get("w").expect("missing 'w' gradient");
923 assert_eq!(w_grad.len(), 3);
924 assert_abs_diff_eq!(w_grad[0], 0.1, epsilon = 1e-12);
925 assert_abs_diff_eq!(w_grad[1], 0.2, epsilon = 1e-12);
926 assert_abs_diff_eq!(w_grad[2], 0.3, epsilon = 1e-12);
927
928 let params = opt.policy.get_parameters();
930 let w = params.get("w").expect("w parameter");
931 assert_abs_diff_eq!(w[0], 0.1, epsilon = 1e-12);
932 assert_abs_diff_eq!(w[1], 0.2, epsilon = 1e-12);
933 assert_abs_diff_eq!(w[2], 0.3, epsilon = 1e-12);
934 }
935
936 #[test]
937 fn test_apply_parameter_update_length_mismatch_errors() {
938 let mut opt = make_optimizer(0.1);
939 let bad = arr1(&[0.1, 0.2, 0.3, 0.4]);
941 assert!(opt.apply_parameter_update(&bad).is_err());
942 }
943
944 #[test]
945 fn test_kl_estimate_uses_real_damped_fisher() {
946 let damping = 0.2;
947 let mut opt = make_optimizer(damping);
948 let samples = arr2(&[[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]);
949 opt.set_score_samples(samples.clone());
950
951 let direction = arr1(&[1.0, 1.0, 1.0]);
952 let step = 0.5_f64;
953
954 let fvp = reference_fvp(&samples, &direction, damping);
956 let quad: f64 = direction.iter().zip(fvp.iter()).map(|(&d, &f)| d * f).sum();
957 let expected = 0.5 * quad * step * step;
958
959 let got = opt
960 .estimate_kl_divergence(&direction, step)
961 .expect("kl estimate");
962 assert_abs_diff_eq!(got, expected, epsilon = 1e-10);
963 }
964}