1use crate::error::RlOpsError;
4use crate::estimator::AdvantageEstimator;
5use crate::kl::f32_ops::{compute_batch_group_advantages, compute_group_advantages};
6
7pub struct GroupRelativeEstimator;
16
17impl AdvantageEstimator for GroupRelativeEstimator {
18 fn compute(&self, rewards: &[f32], group_size: usize) -> Result<Vec<f32>, RlOpsError> {
25 compute_batch_group_advantages(rewards, group_size)
26 }
27}
28
29pub fn compute_single_group_advantages(rewards: &[f32]) -> Vec<f32> {
34 compute_group_advantages(rewards)
35}
36
37#[cfg(test)]
38mod tests {
39 use super::*;
40
41 #[test]
51 fn advantage_estimator_trait_binary_rewards() {
52 let estimator = GroupRelativeEstimator;
53 let rewards = [1.0f32, 0.0, 1.0, 0.0];
54 let adv = estimator
55 .compute(&rewards, 4)
56 .expect("compute must succeed");
57 assert_eq!(adv.len(), 4);
58
59 let expected = [1.0f32, -1.0, 1.0, -1.0];
60 for (i, (&got, &exp)) in adv.iter().zip(expected.iter()).enumerate() {
61 assert!(
62 (got - exp).abs() < 1e-5,
63 "adv[{i}]: expected {exp}, got {got}"
64 );
65 }
66 }
67
68 #[test]
69 fn advantage_estimator_trait_constant_group_returns_zeros() {
70 let estimator = GroupRelativeEstimator;
71 let rewards = [5.0f32, 5.0, 5.0, 5.0];
72 let adv = estimator
73 .compute(&rewards, 4)
74 .expect("compute must succeed");
75 assert!(
76 adv.iter().all(|&v| v == 0.0),
77 "constant rewards must produce all-zero advantages"
78 );
79 }
80
81 #[test]
82 fn advantage_estimator_trait_multiple_groups() {
83 let estimator = GroupRelativeEstimator;
84 let rewards = [1.0f32, 0.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0];
86 let adv = estimator
87 .compute(&rewards, 4)
88 .expect("compute must succeed");
89 assert_eq!(adv.len(), 8);
90 assert!((adv[0] - 1.0).abs() < 1e-5);
92 assert!((adv[1] + 1.0).abs() < 1e-5);
93 assert!(adv[4..8].iter().all(|&v| v == 0.0));
95 }
96
97 #[test]
98 fn advantage_estimator_trait_bad_group_size_returns_err() {
99 let estimator = GroupRelativeEstimator;
100 let result = estimator.compute(&[1.0f32, 0.0, 1.0], 2);
102 assert!(result.is_err(), "non-divisible group_size must return Err");
103
104 let result2 = estimator.compute(&[1.0f32], 0);
106 assert!(result2.is_err(), "group_size=0 must return Err");
107 }
108
109 #[test]
110 fn single_group_convenience_function() {
111 let rewards = [2.0f32, 0.0];
112 let adv = compute_single_group_advantages(&rewards);
113 assert_eq!(adv.len(), 2);
114 assert!((adv[0] - 1.0).abs() < 1e-5);
115 assert!((adv[1] + 1.0).abs() < 1e-5);
116 }
117
118 #[test]
120 fn advantage_estimator_is_dyn_compatible() {
121 use std::sync::Arc;
122 let estimator: Arc<dyn AdvantageEstimator> = Arc::new(GroupRelativeEstimator);
123 let rewards = [1.0f32, 0.0, 1.0, 0.0];
124 let adv = estimator.compute(&rewards, 4).unwrap();
125 assert_eq!(adv.len(), 4);
126 }
127}