1use crate::error::RloxError;
2
3#[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 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 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 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 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 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 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; 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(); let _c = c_bar.min(rho);
151 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 let log_rhos = vec![5.0]; 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 assert!((vs_clipped[0] - 1.0).abs() < 1e-5);
193 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 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 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 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 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 let gamma = 0.99_f32;
309 let log_rhos = vec![0.0; 4]; 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]; 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 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 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 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 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 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 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 let gamma = 0.99_f32;
424 let log_rhos = vec![0.0]; let rewards = vec![1.0];
426 let values = vec![0.5];
427 let bootstrap = 10.0; 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 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 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 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}