Skip to main content

rlox_core/training/
vtrace.rs

1use crate::error::RloxError;
2
3/// Compute V-trace targets and policy gradient advantages (Espeholt et al. 2018).
4///
5/// Processes backwards from t=n-1 to t=0:
6///   rho_t = min(rho_bar, exp(log_rhos[t]))
7///   c_t   = min(c_bar,   exp(log_rhos[t]))
8///   non_terminal = 1.0 - dones[t]
9///   delta_t = rho_t * (rewards[t] + gamma * non_terminal * values[t+1] - values[t])
10///   vs[t]   = values[t] + delta_t + gamma * non_terminal * c_t * (vs[t+1] - values[t+1])
11///   pg_advantages[t] = rho_t * (rewards[t] + gamma * non_terminal * vs[t+1] - values[t])
12///
13/// Uses `bootstrap_value` for values[n] and vs[n], zeroed when the last step
14/// is terminal (`dones[n-1] == 1.0`).
15///
16/// Returns `(vs, pg_advantages)`.
17// V-trace requires 4 parallel slices + 4 scalar hyperparams; each is an independent math arg.
18#[allow(clippy::too_many_arguments)]
19pub fn compute_vtrace(
20    log_rhos: &[f32],
21    rewards: &[f32],
22    values: &[f32],
23    dones: &[f32],
24    bootstrap_value: f32,
25    gamma: f32,
26    rho_bar: f32,
27    c_bar: f32,
28) -> Result<(Vec<f32>, Vec<f32>), RloxError> {
29    let n = log_rhos.len();
30
31    if rewards.len() != n || values.len() != n || dones.len() != n {
32        return Err(RloxError::ShapeMismatch {
33            expected: format!("all slices length {n}"),
34            got: format!(
35                "log_rhos={}, rewards={}, values={}, dones={}",
36                n,
37                rewards.len(),
38                values.len(),
39                dones.len()
40            ),
41        });
42    }
43
44    if n == 0 {
45        return Ok((Vec::new(), Vec::new()));
46    }
47
48    let mut vs = vec![0.0f32; n];
49    let mut pg_advantages = vec![0.0f32; n];
50
51    // Handle last step (t = n-1) outside the loop
52    let last = n - 1;
53    {
54        let ratio = log_rhos[last].exp();
55        let rho_t = rho_bar.min(ratio);
56        let non_terminal = 1.0 - dones[last];
57        let next_value = bootstrap_value * non_terminal;
58
59        let delta_t = rho_t * (rewards[last] + gamma * next_value - values[last]);
60        // vs_next for the last step is bootstrap_value (zeroed if terminal)
61        let vs_next_val = bootstrap_value * non_terminal;
62        vs[last] = values[last]
63            + delta_t
64            + gamma * non_terminal * rho_bar.min(ratio).min(c_bar) * (vs_next_val - next_value);
65        pg_advantages[last] = rho_t * (rewards[last] + gamma * vs_next_val - values[last]);
66    }
67
68    // Iterate backwards for remaining steps
69    let mut vs_next = vs[last];
70
71    for t in (0..last).rev() {
72        let ratio = log_rhos[t].exp();
73        let rho_t = rho_bar.min(ratio);
74        let c_t = c_bar.min(ratio);
75        let non_terminal = 1.0 - dones[t];
76
77        let next_value = values[t + 1];
78
79        let delta_t = rho_t * (rewards[t] + gamma * non_terminal * next_value - values[t]);
80        vs[t] = values[t] + delta_t + gamma * non_terminal * c_t * (vs_next - next_value);
81        pg_advantages[t] = rho_t * (rewards[t] + gamma * non_terminal * vs_next - values[t]);
82
83        vs_next = vs[t];
84    }
85
86    Ok((vs, pg_advantages))
87}
88
89#[cfg(test)]
90mod tests {
91    use super::*;
92
93    #[test]
94    fn vtrace_empty_input() {
95        let (vs, adv) = compute_vtrace(&[], &[], &[], &[], 0.0, 0.99, 1.0, 1.0).unwrap();
96        assert!(vs.is_empty());
97        assert!(adv.is_empty());
98    }
99
100    #[test]
101    fn vtrace_mismatched_lengths() {
102        let result = compute_vtrace(&[0.0], &[1.0, 2.0], &[0.5], &[0.0], 0.0, 0.99, 1.0, 1.0);
103        assert!(result.is_err());
104    }
105
106    #[test]
107    fn vtrace_on_policy_matches_gae_like() {
108        // When log_rhos = 0 (on-policy), rho=1, c=1 => V-trace reduces to
109        // something close to GAE(lambda=1)
110        let log_rhos = vec![0.0; 3];
111        let rewards = vec![1.0, 1.0, 1.0];
112        let values = vec![0.0, 0.0, 0.0];
113        let bootstrap = 0.0;
114        let gamma = 0.99;
115
116        let dones = vec![0.0; 3];
117        let (vs, _adv) = compute_vtrace(
118            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, 1.0, 1.0,
119        )
120        .unwrap();
121
122        // On-policy with rho=c=1:
123        // t=2: delta = 1*(1 + 0.99*0 - 0) = 1, vs[2] = 0 + 1 + 0.99*1*(0-0) = 1
124        // t=1: delta = 1*(1 + 0.99*0 - 0) = 1, vs[1] = 0 + 1 + 0.99*1*(1-0) = 1.99
125        // t=0: delta = 1*(1 + 0.99*0 - 0) = 1, vs[0] = 0 + 1 + 0.99*1*(1.99-0) = 2.9701
126        assert!((vs[2] - 1.0).abs() < 1e-5);
127        assert!((vs[1] - 1.99).abs() < 1e-5);
128        assert!((vs[0] - 2.9701).abs() < 1e-4);
129    }
130
131    #[test]
132    fn vtrace_single_step() {
133        // Use rho_bar large enough so clipping doesn't engage
134        let log_rho = 0.5_f32;
135        let log_rhos = vec![log_rho];
136        let rewards = vec![1.0];
137        let values = vec![0.5];
138        let bootstrap = 0.0;
139        let gamma = 0.99;
140        let rho_bar = 10.0; // no clipping
141        let c_bar = 10.0;
142
143        let dones = vec![0.0];
144        let (vs, adv) = compute_vtrace(
145            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, rho_bar, c_bar,
146        )
147        .unwrap();
148
149        let rho = log_rho.exp(); // ~1.6487
150        let _c = c_bar.min(rho);
151        // t=0 (only step, n-1): next_value = bootstrap = 0, vs_next = bootstrap = 0
152        // delta = rho * (1.0 + 0.99*0 - 0.5) = rho * 0.5
153        // vs[0] = 0.5 + rho*0.5 + 0.99*c*(bootstrap - bootstrap) = 0.5 + rho*0.5
154        // pg_adv[0] = rho * (1.0 + 0.99*bootstrap - 0.5) = rho * 0.5
155        let expected_vs = 0.5 + rho * 0.5;
156        let expected_adv = rho * 0.5;
157
158        assert!(
159            (vs[0] - expected_vs).abs() < 1e-5,
160            "vs[0]={}, expected={}",
161            vs[0],
162            expected_vs
163        );
164        assert!(
165            (adv[0] - expected_adv).abs() < 1e-5,
166            "adv[0]={}, expected={}",
167            adv[0],
168            expected_adv
169        );
170    }
171
172    #[test]
173    fn vtrace_clipping_reduces_correction() {
174        // With very large importance ratio, clipping should limit the correction
175        let log_rhos = vec![5.0]; // exp(5) ~ 148, way above rho_bar=1
176        let rewards = vec![1.0];
177        let values = vec![0.0];
178        let bootstrap = 0.0;
179        let gamma = 0.99;
180
181        let dones = vec![0.0];
182        let (vs_clipped, _) = compute_vtrace(
183            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, 1.0, 1.0,
184        )
185        .unwrap();
186        let (vs_unclipped, _) = compute_vtrace(
187            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, 200.0, 200.0,
188        )
189        .unwrap();
190
191        // With rho_bar=1.0, rho is clamped to 1.0 => delta = 1*(1-0) = 1 => vs = 1.0
192        assert!((vs_clipped[0] - 1.0).abs() < 1e-5);
193        // Without clipping, rho ~ 148 => delta = 148*(1-0) = 148 => vs = 148
194        assert!(vs_unclipped[0] > 100.0);
195    }
196
197    #[test]
198    fn vtrace_output_lengths_match_input() {
199        let n = 10;
200        let log_rhos = vec![0.0; n];
201        let rewards = vec![1.0; n];
202        let values = vec![0.5; n];
203        let dones = vec![0.0; n];
204        let (vs, adv) =
205            compute_vtrace(&log_rhos, &rewards, &values, &dones, 0.0, 0.99, 1.0, 1.0).unwrap();
206        assert_eq!(vs.len(), n);
207        assert_eq!(adv.len(), n);
208    }
209
210    #[test]
211    fn vtrace_reference_implementation() {
212        // Reference: manually compute for a 3-step trajectory
213        let gamma = 0.9_f32;
214        let rho_bar = 1.5_f32;
215        let c_bar = 1.2_f32;
216
217        let log_rhos = vec![0.2, -0.3, 0.8];
218        let rewards = vec![1.0, 2.0, 3.0];
219        let values = vec![0.5, 1.0, 1.5];
220        let bootstrap = 2.0;
221
222        // Manually compute backwards:
223        // t=2: rho = min(1.5, exp(0.8)) = min(1.5, 2.2255) = 1.5
224        //       c  = min(1.2, 2.2255) = 1.2
225        //       next_val = bootstrap = 2.0
226        //       delta = 1.5 * (3.0 + 0.9*2.0 - 1.5) = 1.5 * 3.3 = 4.95
227        //       vs[2] = 1.5 + 4.95 + 0.9*1.2*(2.0 - 2.0) = 6.45
228        //       pg_adv[2] = 1.5 * (3.0 + 0.9*2.0 - 1.5) = 4.95
229        //       vs_next = 6.45
230        let rho_2 = 1.5_f32;
231        let c_2 = 1.2_f32;
232        let delta_2 = rho_2 * (3.0 + 0.9 * 2.0 - 1.5);
233        let vs_2 = 1.5 + delta_2 + 0.9 * c_2 * (2.0 - 2.0);
234        let pg_2 = rho_2 * (3.0 + 0.9 * 2.0 - 1.5);
235
236        // t=1: rho = min(1.5, exp(-0.3)) = min(1.5, 0.7408) = 0.7408
237        //       c  = min(1.2, 0.7408) = 0.7408
238        //       next_val = values[2] = 1.5
239        //       delta = 0.7408 * (2.0 + 0.9*1.5 - 1.0) = 0.7408 * 2.35 = 1.74088
240        //       vs[1] = 1.0 + 1.74088 + 0.9*0.7408*(6.45 - 1.5) = 1.0 + 1.74088 + 0.9*0.7408*4.95
241        //       pg_adv[1] = 0.7408 * (2.0 + 0.9*6.45 - 1.0) = 0.7408 * (2.0 + 5.805 - 1.0) = 0.7408 * 6.805
242        let rho_1 = (-0.3_f32).exp();
243        let c_1 = c_bar.min(rho_1);
244        let delta_1 = rho_1 * (2.0 + 0.9 * 1.5 - 1.0);
245        let vs_1 = 1.0 + delta_1 + 0.9 * c_1 * (vs_2 - 1.5);
246        let pg_1 = rho_1 * (2.0 + 0.9 * vs_2 - 1.0);
247
248        // t=0: rho = min(1.5, exp(0.2)) = min(1.5, 1.2214) = 1.2214
249        //       c  = min(1.2, 1.2214) = 1.2
250        //       next_val = values[1] = 1.0
251        //       delta = 1.2214 * (1.0 + 0.9*1.0 - 0.5) = 1.2214 * 1.4
252        //       vs[0] = 0.5 + delta + 0.9*1.2*(vs_1 - 1.0)
253        //       pg_adv[0] = 1.2214 * (1.0 + 0.9*vs_1 - 0.5)
254        let rho_0 = (0.2_f32).exp();
255        let c_0 = c_bar.min(rho_0);
256        let delta_0 = rho_0 * (1.0 + 0.9 * 1.0 - 0.5);
257        let vs_0 = 0.5 + delta_0 + 0.9 * c_0 * (vs_1 - 1.0);
258        let pg_0 = rho_0 * (1.0 + 0.9 * vs_1 - 0.5);
259
260        let dones = vec![0.0; 3];
261        let (vs, adv) = compute_vtrace(
262            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, rho_bar, c_bar,
263        )
264        .unwrap();
265
266        assert!(
267            (vs[0] - vs_0).abs() < 1e-4,
268            "vs[0]: got {}, expected {}",
269            vs[0],
270            vs_0
271        );
272        assert!(
273            (vs[1] - vs_1).abs() < 1e-4,
274            "vs[1]: got {}, expected {}",
275            vs[1],
276            vs_1
277        );
278        assert!(
279            (vs[2] - vs_2).abs() < 1e-4,
280            "vs[2]: got {}, expected {}",
281            vs[2],
282            vs_2
283        );
284        assert!(
285            (adv[0] - pg_0).abs() < 1e-4,
286            "adv[0]: got {}, expected {}",
287            adv[0],
288            pg_0
289        );
290        assert!(
291            (adv[1] - pg_1).abs() < 1e-4,
292            "adv[1]: got {}, expected {}",
293            adv[1],
294            pg_1
295        );
296        assert!(
297            (adv[2] - pg_2).abs() < 1e-4,
298            "adv[2]: got {}, expected {}",
299            adv[2],
300            pg_2
301        );
302    }
303
304    #[test]
305    fn vtrace_with_dones_resets_at_boundary() {
306        // 4-step trajectory with a done at t=1. Episode boundary should
307        // prevent rewards from leaking across episodes.
308        let gamma = 0.99_f32;
309        let log_rhos = vec![0.0; 4]; // on-policy
310        let rewards = vec![1.0, 1.0, 1.0, 1.0];
311        let values = vec![0.0; 4];
312        let dones = vec![0.0, 1.0, 0.0, 0.0]; // done at t=1
313        let bootstrap = 0.0;
314
315        let (vs_with_dones, _) = compute_vtrace(
316            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, 1.0, 1.0,
317        )
318        .unwrap();
319
320        // Without dones, rewards leak across episodes
321        let no_dones = vec![0.0; 4];
322        let (vs_no_dones, _) = compute_vtrace(
323            &log_rhos, &rewards, &values, &no_dones, bootstrap, gamma, 1.0, 1.0,
324        )
325        .unwrap();
326
327        // After the boundary (t=0), the done-aware version should produce
328        // a LOWER vs because future rewards beyond the boundary are zeroed.
329        assert!(
330            vs_with_dones[0] < vs_no_dones[0],
331            "vs_with_dones[0]={} should be < vs_no_dones[0]={}",
332            vs_with_dones[0],
333            vs_no_dones[0]
334        );
335
336        // Steps after the boundary (t=2, t=3) should be unaffected
337        assert!(
338            (vs_with_dones[3] - vs_no_dones[3]).abs() < 1e-5,
339            "t=3 should be identical"
340        );
341    }
342
343    #[test]
344    fn vtrace_without_dones_matches_old_behavior() {
345        // Passing all-zeros dones should reproduce the original behavior
346        let gamma = 0.9_f32;
347        let rho_bar = 1.5_f32;
348        let c_bar = 1.2_f32;
349        let log_rhos = vec![0.2, -0.3, 0.8];
350        let rewards = vec![1.0, 2.0, 3.0];
351        let values = vec![0.5, 1.0, 1.5];
352        let bootstrap = 2.0;
353        let dones = vec![0.0; 3];
354
355        let (vs, adv) = compute_vtrace(
356            &log_rhos, &rewards, &values, &dones, bootstrap, gamma, rho_bar, c_bar,
357        )
358        .unwrap();
359
360        // Manually computed reference (same as vtrace_reference_implementation)
361        let rho_2 = 1.5_f32;
362        let c_2 = 1.2_f32;
363        let delta_2 = rho_2 * (3.0 + 0.9 * 2.0 - 1.5);
364        let vs_2 = 1.5 + delta_2 + 0.9 * c_2 * (2.0 - 2.0);
365        let pg_2 = rho_2 * (3.0 + 0.9 * 2.0 - 1.5);
366
367        let rho_1 = (-0.3_f32).exp();
368        let c_1 = c_bar.min(rho_1);
369        let delta_1 = rho_1 * (2.0 + 0.9 * 1.5 - 1.0);
370        let vs_1 = 1.0 + delta_1 + 0.9 * c_1 * (vs_2 - 1.5);
371        let pg_1 = rho_1 * (2.0 + 0.9 * vs_2 - 1.0);
372
373        let rho_0 = (0.2_f32).exp();
374        let c_0 = c_bar.min(rho_0);
375        let delta_0 = rho_0 * (1.0 + 0.9 * 1.0 - 0.5);
376        let vs_0 = 0.5 + delta_0 + 0.9 * c_0 * (vs_1 - 1.0);
377        let pg_0 = rho_0 * (1.0 + 0.9 * vs_1 - 0.5);
378
379        assert!(
380            (vs[0] - vs_0).abs() < 1e-4,
381            "vs[0]: got {}, expected {}",
382            vs[0],
383            vs_0
384        );
385        assert!(
386            (vs[1] - vs_1).abs() < 1e-4,
387            "vs[1]: got {}, expected {}",
388            vs[1],
389            vs_1
390        );
391        assert!(
392            (vs[2] - vs_2).abs() < 1e-4,
393            "vs[2]: got {}, expected {}",
394            vs[2],
395            vs_2
396        );
397        assert!(
398            (adv[0] - pg_0).abs() < 1e-4,
399            "adv[0]: got {}, expected {}",
400            adv[0],
401            pg_0
402        );
403        assert!(
404            (adv[1] - pg_1).abs() < 1e-4,
405            "adv[1]: got {}, expected {}",
406            adv[1],
407            pg_1
408        );
409        assert!(
410            (adv[2] - pg_2).abs() < 1e-4,
411            "adv[2]: got {}, expected {}",
412            adv[2],
413            pg_2
414        );
415
416        // Suppress unused-variable warnings
417        let _ = (c_0, c_1, c_2, pg_0, pg_1, pg_2, delta_0, delta_1, delta_2);
418    }
419
420    #[test]
421    fn vtrace_dones_at_last_step_zeros_bootstrap() {
422        // When the last step is terminal, bootstrap should be zeroed
423        let gamma = 0.99_f32;
424        let log_rhos = vec![0.0]; // on-policy, single step
425        let rewards = vec![1.0];
426        let values = vec![0.5];
427        let bootstrap = 10.0; // large bootstrap to make the effect visible
428
429        // With done at last step
430        let dones_terminal = vec![1.0];
431        let (vs_term, adv_term) = compute_vtrace(
432            &log_rhos,
433            &rewards,
434            &values,
435            &dones_terminal,
436            bootstrap,
437            gamma,
438            1.0,
439            1.0,
440        )
441        .unwrap();
442
443        // Without done
444        let dones_none = vec![0.0];
445        let (vs_cont, adv_cont) = compute_vtrace(
446            &log_rhos,
447            &rewards,
448            &values,
449            &dones_none,
450            bootstrap,
451            gamma,
452            1.0,
453            1.0,
454        )
455        .unwrap();
456
457        // Terminal: delta = 1*(1.0 + 0.99*0*10 - 0.5) = 0.5, vs = 0.5 + 0.5 = 1.0
458        // Non-terminal: delta = 1*(1.0 + 0.99*10 - 0.5) = 10.4, vs = 0.5 + 10.4 = 10.9
459        assert!(
460            (vs_term[0] - 1.0).abs() < 1e-5,
461            "terminal vs[0]={}, expected 1.0",
462            vs_term[0]
463        );
464        assert!(
465            vs_cont[0] > vs_term[0],
466            "non-terminal vs should be larger due to bootstrap"
467        );
468
469        // Terminal advantage: rho*(r + gamma*0*vs_next - v) = 1*(1+0-0.5) = 0.5
470        assert!(
471            (adv_term[0] - 0.5).abs() < 1e-5,
472            "terminal adv[0]={}, expected 0.5",
473            adv_term[0]
474        );
475        assert!(
476            adv_cont[0] > adv_term[0],
477            "non-terminal adv should be larger"
478        );
479    }
480
481    mod proptests {
482        use super::*;
483        use proptest::prelude::*;
484
485        proptest! {
486            #[test]
487            fn vtrace_output_length_matches_input(n in 0..200usize) {
488                let log_rhos = vec![0.0; n];
489                let rewards = vec![1.0; n];
490                let values = vec![0.5; n];
491                let dones = vec![0.0; n];
492                let (vs, adv) = compute_vtrace(&log_rhos, &rewards, &values, &dones, 0.0, 0.99, 1.0, 1.0).unwrap();
493                prop_assert_eq!(vs.len(), n);
494                prop_assert_eq!(adv.len(), n);
495            }
496
497            #[test]
498            fn vtrace_on_policy_vs_are_finite(n in 1..100usize) {
499                let log_rhos = vec![0.0; n];
500                let rewards: Vec<f32> = (0..n).map(|i| (i as f32) * 0.1).collect();
501                let values: Vec<f32> = (0..n).map(|i| (i as f32) * 0.05).collect();
502                let dones = vec![0.0; n];
503                let (vs, adv) = compute_vtrace(&log_rhos, &rewards, &values, &dones, 0.0, 0.99, 1.0, 1.0).unwrap();
504                for i in 0..n {
505                    prop_assert!(vs[i].is_finite(), "vs[{}] is not finite: {}", i, vs[i]);
506                    prop_assert!(adv[i].is_finite(), "adv[{}] is not finite: {}", i, adv[i]);
507                }
508            }
509        }
510    }
511}