1macro_rules! impl_kl_ops {
8 ($mod_name:ident, $float:ty) => {
9 pub mod $mod_name {
10 use crate::error::RlOpsError;
11
12 pub fn compute_group_advantages(rewards: &[$float]) -> Vec<$float> {
15 if rewards.is_empty() {
16 return Vec::new();
17 }
18
19 let n = rewards.len() as $float;
20 let mean = rewards.iter().sum::<$float>() / n;
21 let variance = rewards
22 .iter()
23 .map(|&r| (r - mean) * (r - mean))
24 .sum::<$float>()
25 / n;
26 let std = variance.sqrt();
27
28 if std < 1e-8 as $float {
29 return vec![0.0 as $float; rewards.len()];
30 }
31
32 let inv_std = 1.0 as $float / std;
33 rewards.iter().map(|&r| (r - mean) * inv_std).collect()
34 }
35
36 pub fn compute_token_kl(
38 log_probs_policy: &[$float],
39 log_probs_ref: &[$float],
40 ) -> Result<$float, RlOpsError> {
41 if log_probs_policy.len() != log_probs_ref.len() {
42 return Err(RlOpsError::ShapeMismatch {
43 expected: format!("len={}", log_probs_policy.len()),
44 got: format!("len={}", log_probs_ref.len()),
45 });
46 }
47
48 Ok(log_probs_policy
49 .iter()
50 .zip(log_probs_ref.iter())
51 .map(|(&log_p, &log_q)| log_p.exp() * (log_p - log_q))
52 .sum())
53 }
54
55 pub fn compute_batch_group_advantages(
60 rewards: &[$float],
61 group_size: usize,
62 ) -> Result<Vec<$float>, RlOpsError> {
63 if group_size == 0 {
64 return Err(RlOpsError::ShapeMismatch {
65 expected: "group_size > 0".to_string(),
66 got: "0".to_string(),
67 });
68 }
69 if rewards.len() % group_size != 0 {
70 return Err(RlOpsError::ShapeMismatch {
71 expected: format!("len divisible by {group_size}"),
72 got: format!("len={}", rewards.len()),
73 });
74 }
75
76 const PAR_ELEMENT_THRESHOLD: usize = 4096;
77 if rewards.len() >= PAR_ELEMENT_THRESHOLD {
78 use rayon::prelude::*;
79 let out: Vec<$float> = rewards
80 .par_chunks_exact(group_size)
81 .flat_map_iter(|group| compute_group_advantages(group))
82 .collect();
83 Ok(out)
84 } else {
85 let mut out = Vec::with_capacity(rewards.len());
86 for group in rewards.chunks_exact(group_size) {
87 out.extend_from_slice(&compute_group_advantages(group));
88 }
89 Ok(out)
90 }
91 }
92
93 pub fn compute_token_kl_schulman(
99 log_probs_policy: &[$float],
100 log_probs_ref: &[$float],
101 ) -> Result<$float, RlOpsError> {
102 if log_probs_policy.len() != log_probs_ref.len() {
103 return Err(RlOpsError::ShapeMismatch {
104 expected: format!("len={}", log_probs_policy.len()),
105 got: format!("len={}", log_probs_ref.len()),
106 });
107 }
108
109 Ok(log_probs_policy
110 .iter()
111 .zip(log_probs_ref.iter())
112 .map(|(&log_p, &log_q)| {
113 let r = log_p - log_q;
114 r.exp() - r - 1.0 as $float
115 })
116 .sum())
117 }
118
119 pub fn compute_batch_token_kl(
124 log_probs_policy: &[$float],
125 log_probs_ref: &[$float],
126 seq_len: usize,
127 ) -> Result<Vec<$float>, RlOpsError> {
128 if log_probs_policy.len() != log_probs_ref.len() {
129 return Err(RlOpsError::ShapeMismatch {
130 expected: format!("len={}", log_probs_policy.len()),
131 got: format!("len={}", log_probs_ref.len()),
132 });
133 }
134 if seq_len == 0 {
135 return Err(RlOpsError::ShapeMismatch {
136 expected: "seq_len > 0".to_string(),
137 got: "0".to_string(),
138 });
139 }
140 if log_probs_policy.len() % seq_len != 0 {
141 return Err(RlOpsError::ShapeMismatch {
142 expected: format!("len divisible by {seq_len}"),
143 got: format!("len={}", log_probs_policy.len()),
144 });
145 }
146
147 const PAR_ELEMENT_THRESHOLD: usize = 4096;
148 let batch_size = log_probs_policy.len() / seq_len;
149
150 let kl_for_seq = |i: usize| -> $float {
151 let off = i * seq_len;
152 let ps = &log_probs_policy[off..off + seq_len];
153 let qs = &log_probs_ref[off..off + seq_len];
154 ps.iter()
155 .zip(qs.iter())
156 .map(|(&log_p, &log_q)| log_p.exp() * (log_p - log_q))
157 .sum()
158 };
159
160 let out = if log_probs_policy.len() >= PAR_ELEMENT_THRESHOLD {
161 use rayon::prelude::*;
162 (0..batch_size).into_par_iter().map(kl_for_seq).collect()
163 } else {
164 (0..batch_size).map(kl_for_seq).collect()
165 };
166 Ok(out)
167 }
168
169 pub fn compute_batch_token_kl_schulman(
174 log_probs_policy: &[$float],
175 log_probs_ref: &[$float],
176 seq_len: usize,
177 ) -> Result<Vec<$float>, RlOpsError> {
178 if log_probs_policy.len() != log_probs_ref.len() {
179 return Err(RlOpsError::ShapeMismatch {
180 expected: format!("len={}", log_probs_policy.len()),
181 got: format!("len={}", log_probs_ref.len()),
182 });
183 }
184 if seq_len == 0 {
185 return Err(RlOpsError::ShapeMismatch {
186 expected: "seq_len > 0".to_string(),
187 got: "0".to_string(),
188 });
189 }
190 if log_probs_policy.len() % seq_len != 0 {
191 return Err(RlOpsError::ShapeMismatch {
192 expected: format!("len divisible by {seq_len}"),
193 got: format!("len={}", log_probs_policy.len()),
194 });
195 }
196
197 const PAR_ELEMENT_THRESHOLD: usize = 4096;
198 let batch_size = log_probs_policy.len() / seq_len;
199
200 let kl_for_seq = |i: usize| -> $float {
201 let off = i * seq_len;
202 let ps = &log_probs_policy[off..off + seq_len];
203 let qs = &log_probs_ref[off..off + seq_len];
204 ps.iter()
205 .zip(qs.iter())
206 .map(|(&log_p, &log_q)| {
207 let r = log_p - log_q;
208 r.exp() - r - 1.0 as $float
209 })
210 .sum()
211 };
212
213 let out = if log_probs_policy.len() >= PAR_ELEMENT_THRESHOLD {
214 use rayon::prelude::*;
215 (0..batch_size).into_par_iter().map(kl_for_seq).collect()
216 } else {
217 (0..batch_size).map(kl_for_seq).collect()
218 };
219 Ok(out)
220 }
221 }
222 };
223}
224
225impl_kl_ops!(f64_ops, f64);
226impl_kl_ops!(f32_ops, f32);
227
228pub use f64_ops::*;
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 #[test]
236 fn test_group_advantages_basic() {
237 let rewards = [1.0, 0.5, 0.8];
238 let adv = compute_group_advantages(&rewards);
239 assert_eq!(adv.len(), 3);
240 let mean: f64 = adv.iter().sum::<f64>() / adv.len() as f64;
241 assert!(mean.abs() < 1e-10);
242 }
243
244 #[test]
245 fn test_group_advantages_constant_rewards() {
246 let rewards = [5.0, 5.0, 5.0];
247 let adv = compute_group_advantages(&rewards);
248 assert!(adv.iter().all(|&v| v == 0.0));
249 }
250
251 #[test]
252 fn test_group_advantages_empty() {
253 let adv = compute_group_advantages(&[]);
254 assert!(adv.is_empty());
255 }
256
257 #[test]
258 fn test_token_kl_identical() {
259 let log_p = [-1.0, -2.0, -0.5];
260 let kl = compute_token_kl(&log_p, &log_p).unwrap();
261 assert!(kl.abs() < 1e-15);
262 }
263
264 #[test]
265 fn test_token_kl_known_value() {
266 let log_p = [-1.0];
267 let log_q = [-2.0];
268 let kl = compute_token_kl(&log_p, &log_q).unwrap();
269 assert!((kl - (-1.0_f64).exp()).abs() < 1e-10);
270 }
271
272 #[test]
273 fn test_token_kl_mismatched_lengths_returns_err() {
274 let result = compute_token_kl(&[1.0, 2.0], &[1.0]);
275 assert!(result.is_err());
276 }
277
278 #[test]
279 fn token_kl_mismatched_lengths_returns_err_not_panic() {
280 let log_p = vec![-1.0f64, -2.0];
281 let log_q = vec![-1.0f64];
282 let result = compute_token_kl(&log_p, &log_q);
283 assert!(result.is_err(), "mismatched lengths must return Err");
284 }
285
286 #[test]
287 fn token_kl_matching_lengths_returns_ok() {
288 let log_p = vec![-1.0f64, -2.0, -0.5];
289 let log_q = vec![-1.0f64, -2.0, -0.5];
290 let result = compute_token_kl(&log_p, &log_q);
291 assert!(result.is_ok());
292 assert!(result.unwrap().abs() < 1e-15);
293 }
294
295 #[test]
296 fn token_kl_empty_slices_returns_zero() {
297 let result = compute_token_kl(&[], &[]);
298 assert!(result.is_ok());
299 assert_eq!(result.unwrap(), 0.0);
300 }
301
302 #[test]
303 fn token_kl_nan_input_propagates_to_output() {
304 let log_p = vec![f64::NAN];
305 let log_q = vec![-1.0f64];
306 let result = compute_token_kl(&log_p, &log_q);
307 if let Ok(v) = result {
308 assert!(v.is_nan(), "NaN input should produce NaN output");
309 }
310 }
311
312 #[test]
313 fn token_kl_inf_input_does_not_panic() {
314 let log_p = vec![f64::INFINITY];
315 let log_q = vec![-1.0f64];
316 let _result = compute_token_kl(&log_p, &log_q);
317 }
318
319 #[test]
320 fn token_kl_known_value_still_correct_after_refactor() {
321 let log_p = vec![-1.0f64];
322 let log_q = vec![-2.0f64];
323 let kl = compute_token_kl(&log_p, &log_q).unwrap();
324 assert!((kl - (-1.0_f64).exp()).abs() < 1e-10);
325 }
326
327 #[test]
328 fn test_batch_group_advantages() {
329 let rewards = [1.0, 2.0, 3.0, 10.0, 10.0, 10.0];
330 let adv = compute_batch_group_advantages(&rewards, 3).unwrap();
331 assert_eq!(adv.len(), 6);
332 let g1_mean: f64 = adv[..3].iter().sum::<f64>() / 3.0;
333 assert!(g1_mean.abs() < 1e-10);
334 assert!(adv[3..6].iter().all(|&v| v == 0.0));
335 }
336
337 #[test]
338 fn test_batch_group_advantages_bad_size() {
339 assert!(compute_batch_group_advantages(&[1.0, 2.0, 3.0], 2).is_err());
340 assert!(compute_batch_group_advantages(&[1.0], 0).is_err());
341 }
342
343 #[test]
344 fn test_token_kl_schulman_identical() {
345 let log_p = [-1.0, -2.0, -0.5];
346 let kl = compute_token_kl_schulman(&log_p, &log_p).unwrap();
347 assert!(kl.abs() < 1e-15);
348 }
349
350 #[test]
351 fn test_token_kl_schulman_known_value() {
352 let log_p = [-1.0];
353 let log_q = [-2.0];
354 let kl = compute_token_kl_schulman(&log_p, &log_q).unwrap();
355 assert!((kl - (1.0_f64.exp() - 2.0)).abs() < 1e-10);
356 }
357
358 #[test]
359 fn test_token_kl_schulman_non_negative() {
360 let log_p = [-0.5, -1.0, -3.0, 0.0];
361 let log_q = [-1.0, -0.5, -0.1, -2.0];
362 let kl = compute_token_kl_schulman(&log_p, &log_q).unwrap();
363 assert!(kl >= 0.0, "Schulman KL should be non-negative, got {kl}");
364 }
365
366 #[test]
369 fn test_batch_token_kl_matches_unbatched() {
370 let log_p = vec![-1.0, -2.0, -0.5, -1.5, -0.3, -2.5];
371 let log_q = vec![-1.1, -1.9, -0.6, -1.4, -0.4, -2.4];
372 let batched = compute_batch_token_kl(&log_p, &log_q, 3).unwrap();
373 let kl0 = compute_token_kl(&log_p[..3], &log_q[..3]).unwrap();
374 let kl1 = compute_token_kl(&log_p[3..], &log_q[3..]).unwrap();
375 assert_eq!(batched.len(), 2);
376 assert!((batched[0] - kl0).abs() < 1e-12);
377 assert!((batched[1] - kl1).abs() < 1e-12);
378 }
379
380 #[test]
381 fn test_batch_token_kl_schulman_matches_unbatched() {
382 let log_p = vec![-1.0, -2.0, -0.5, -1.5, -0.3, -2.5];
383 let log_q = vec![-1.1, -1.9, -0.6, -1.4, -0.4, -2.4];
384 let batched = compute_batch_token_kl_schulman(&log_p, &log_q, 3).unwrap();
385 let kl0 = compute_token_kl_schulman(&log_p[..3], &log_q[..3]).unwrap();
386 let kl1 = compute_token_kl_schulman(&log_p[3..], &log_q[3..]).unwrap();
387 assert_eq!(batched.len(), 2);
388 assert!((batched[0] - kl0).abs() < 1e-12);
389 assert!((batched[1] - kl1).abs() < 1e-12);
390 }
391
392 #[test]
393 fn test_batch_token_kl_bad_seq_len() {
394 assert!(compute_batch_token_kl(&[1.0, 2.0, 3.0], &[1.0, 2.0, 3.0], 0).is_err());
395 assert!(compute_batch_token_kl(&[1.0, 2.0, 3.0], &[1.0, 2.0, 3.0], 2).is_err());
396 }
397
398 #[test]
399 fn test_batch_token_kl_mismatched_lengths() {
400 assert!(compute_batch_token_kl(&[1.0, 2.0], &[1.0], 1).is_err());
401 }
402
403 #[test]
406 fn test_f32_token_kl_identical() {
407 let log_p: Vec<f32> = vec![-1.0, -2.0, -0.5];
408 let kl = f32_ops::compute_token_kl(&log_p, &log_p).unwrap();
409 assert!(kl.abs() < 1e-6);
410 }
411
412 #[test]
413 fn test_f32_batch_token_kl_schulman() {
414 let log_p: Vec<f32> = vec![-1.0, -2.0, -0.5, -1.5];
415 let log_q: Vec<f32> = vec![-1.1, -1.9, -0.6, -1.4];
416 let batched = f32_ops::compute_batch_token_kl_schulman(&log_p, &log_q, 2).unwrap();
417 assert_eq!(batched.len(), 2);
418 let kl0 = f32_ops::compute_token_kl_schulman(&log_p[..2], &log_q[..2]).unwrap();
419 assert!((batched[0] - kl0).abs() < 1e-6);
420 }
421
422 #[test]
423 fn test_f32_group_advantages() {
424 let rewards: Vec<f32> = vec![1.0, 2.0, 3.0];
425 let adv = f32_ops::compute_group_advantages(&rewards);
426 assert_eq!(adv.len(), 3);
427 let mean: f32 = adv.iter().sum::<f32>() / 3.0;
428 assert!(mean.abs() < 1e-5);
429 }
430}