Skip to main content

rlox_rl_ops/
kl.rs

1//! Token-level KL divergence operations (exact and Schulman 2020 estimator).
2//!
3//! The `impl_kl_ops!` macro generates both `f64_ops` and `f32_ops` sub-modules,
4//! each exposing the same function set for their respective float type.  The
5//! f64 variants are re-exported at module level for ergonomic use as the default.
6
7macro_rules! impl_kl_ops {
8    ($mod_name:ident, $float:ty) => {
9        pub mod $mod_name {
10            use crate::error::RlOpsError;
11
12            /// GRPO group advantage: `(reward - mean) / std`.
13            /// Returns zeros if std < 1e-8.
14            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            /// Token-level KL divergence: `sum(exp(log_p) * (log_p - log_q))`.
37            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            /// Batched GRPO group advantages: process all groups in a single call.
56            ///
57            /// `rewards` is a flat slice of length `n_prompts * group_size`.
58            /// Returns a Vec of the same length with per-group z-score normalisation.
59            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            /// Token-level KL divergence using the Schulman (2020) estimator:
94            /// `sum(exp(log_p - log_q) - (log_p - log_q) - 1)`.
95            ///
96            /// This is the estimator used by TRL (HuggingFace). It is unbiased and
97            /// numerically more stable than the exact `exp(log_p) * (log_p - log_q)`.
98            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            /// Batched token-level KL divergence: process all sequences in a single call.
120            ///
121            /// `log_probs_policy` and `log_probs_ref` are flat slices of length `batch * seq_len`.
122            /// Returns a Vec of length `batch` with per-sequence KL values.
123            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            /// Batched token-level KL divergence using the Schulman (2020) estimator.
170            ///
171            /// `log_probs_policy` and `log_probs_ref` are flat slices of length `batch * seq_len`.
172            /// Returns a Vec of length `batch` with per-sequence KL values.
173            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
228// Re-export f64 versions at module level for ergonomic use as the default.
229pub 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    // --- Batched KL tests ---
367
368    #[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    // --- f32 variant tests ---
404
405    #[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}