1use rand::Rng;
13use rand::SeedableRng;
14use rand_chacha::ChaCha8Rng;
15
16use crate::error::RloxError;
17
18#[derive(Debug, Clone)]
20pub struct DatasetStats {
21 pub n_transitions: usize,
22 pub n_episodes: usize,
23 pub obs_dim: usize,
24 pub act_dim: usize,
25 pub mean_return: f32,
26 pub std_return: f32,
27 pub min_return: f32,
28 pub max_return: f32,
29 pub mean_episode_length: f32,
30}
31
32#[derive(Debug, Clone)]
34pub struct OfflineBatch {
35 pub obs: Vec<f32>, pub next_obs: Vec<f32>, pub actions: Vec<f32>, pub rewards: Vec<f32>, pub terminated: Vec<u8>, pub obs_dim: usize,
41 pub act_dim: usize,
42}
43
44#[derive(Debug, Clone)]
46pub struct TrajectoryBatch {
47 pub obs: Vec<f32>, pub actions: Vec<f32>, pub rewards: Vec<f32>, pub returns_to_go: Vec<f32>, pub timesteps: Vec<u32>, pub mask: Vec<u8>, pub seq_len: usize,
54 pub obs_dim: usize,
55 pub act_dim: usize,
56}
57
58pub struct OfflineDatasetBuffer {
60 obs: Vec<f32>,
61 next_obs: Vec<f32>,
62 actions: Vec<f32>,
63 rewards: Vec<f32>,
64 terminated: Vec<u8>,
65 #[allow(dead_code)]
66 truncated: Vec<u8>,
67
68 episode_starts: Vec<usize>,
70 episode_lengths: Vec<usize>,
71 episode_returns: Vec<f32>,
72
73 obs_dim: usize,
74 act_dim: usize,
75 len: usize,
76
77 obs_mean: Option<Vec<f32>>,
79 obs_std: Option<Vec<f32>>,
80 reward_mean: Option<f32>,
81 reward_std: Option<f32>,
82}
83
84impl OfflineDatasetBuffer {
85 #[allow(clippy::too_many_arguments)]
90 pub fn from_arrays(
91 obs: Vec<f32>,
92 next_obs: Vec<f32>,
93 actions: Vec<f32>,
94 rewards: Vec<f32>,
95 terminated: Vec<u8>,
96 truncated: Vec<u8>,
97 obs_dim: usize,
98 act_dim: usize,
99 ) -> Result<Self, RloxError> {
100 let n = rewards.len();
101
102 if obs.len() != n * obs_dim {
103 return Err(RloxError::ShapeMismatch {
104 expected: format!("obs length = {} * {} = {}", n, obs_dim, n * obs_dim),
105 got: format!("{}", obs.len()),
106 });
107 }
108 if next_obs.len() != n * obs_dim {
109 return Err(RloxError::ShapeMismatch {
110 expected: format!("next_obs length = {}", n * obs_dim),
111 got: format!("{}", next_obs.len()),
112 });
113 }
114 if actions.len() != n * act_dim {
115 return Err(RloxError::ShapeMismatch {
116 expected: format!("actions length = {} * {} = {}", n, act_dim, n * act_dim),
117 got: format!("{}", actions.len()),
118 });
119 }
120 if terminated.len() != n || truncated.len() != n {
121 return Err(RloxError::ShapeMismatch {
122 expected: format!("terminated/truncated length = {}", n),
123 got: format!(
124 "terminated={}, truncated={}",
125 terminated.len(),
126 truncated.len()
127 ),
128 });
129 }
130
131 let mut episode_starts = vec![0usize];
133 let mut episode_returns = Vec::new();
134 let mut ep_return = 0.0f32;
135
136 for i in 0..n {
137 ep_return += rewards[i];
138 let done = terminated[i] != 0 || truncated[i] != 0;
139 if done || i == n - 1 {
140 episode_returns.push(ep_return);
141 if i + 1 < n {
142 episode_starts.push(i + 1);
143 }
144 ep_return = 0.0;
145 }
146 }
147
148 let episode_lengths: Vec<usize> = episode_starts
149 .windows(2)
150 .map(|w| w[1] - w[0])
151 .chain(std::iter::once(n - episode_starts.last().unwrap_or(&0)))
152 .collect();
153
154 Ok(Self {
155 obs,
156 next_obs,
157 actions,
158 rewards,
159 terminated,
160 truncated,
161 episode_starts,
162 episode_lengths,
163 episode_returns,
164 obs_dim,
165 act_dim,
166 len: n,
167 obs_mean: None,
168 obs_std: None,
169 reward_mean: None,
170 reward_std: None,
171 })
172 }
173
174 pub fn len(&self) -> usize {
176 self.len
177 }
178
179 pub fn is_empty(&self) -> bool {
180 self.len == 0
181 }
182
183 pub fn n_episodes(&self) -> usize {
185 self.episode_starts.len()
186 }
187
188 pub fn obs_dim(&self) -> usize {
189 self.obs_dim
190 }
191
192 pub fn act_dim(&self) -> usize {
193 self.act_dim
194 }
195
196 #[allow(clippy::needless_range_loop)]
198 pub fn compute_normalization(&mut self) {
199 let n = self.len;
200 let d = self.obs_dim;
201
202 let mut mean = vec![0.0f64; d];
204 for i in 0..n {
205 for j in 0..d {
206 mean[j] += self.obs[i * d + j] as f64;
207 }
208 }
209 for m in &mut mean {
210 *m /= n as f64;
211 }
212
213 let mut var = vec![0.0f64; d];
214 for i in 0..n {
215 for j in 0..d {
216 let diff = self.obs[i * d + j] as f64 - mean[j];
217 var[j] += diff * diff;
218 }
219 }
220 for v in &mut var {
221 *v = (*v / n as f64).sqrt().max(1e-8);
222 }
223
224 self.obs_mean = Some(mean.iter().map(|&x| x as f32).collect());
225 self.obs_std = Some(var.iter().map(|&x| x as f32).collect());
226
227 let r_mean = self.rewards.iter().map(|&r| r as f64).sum::<f64>() / n as f64;
229 let r_var = self
230 .rewards
231 .iter()
232 .map(|&r| {
233 let d = r as f64 - r_mean;
234 d * d
235 })
236 .sum::<f64>()
237 / n as f64;
238 self.reward_mean = Some(r_mean as f32);
239 self.reward_std = Some((r_var.sqrt().max(1e-8)) as f32);
240 }
241
242 pub fn sample(&self, batch_size: usize, seed: u64) -> OfflineBatch {
244 let mut rng = ChaCha8Rng::seed_from_u64(seed);
245 let d = self.obs_dim;
246 let a = self.act_dim;
247
248 let mut obs = Vec::with_capacity(batch_size * d);
249 let mut next_obs = Vec::with_capacity(batch_size * d);
250 let mut actions = Vec::with_capacity(batch_size * a);
251 let mut rewards = Vec::with_capacity(batch_size);
252 let mut terminated = Vec::with_capacity(batch_size);
253
254 for _ in 0..batch_size {
255 let idx = rng.random_range(0..self.len);
256
257 obs.extend_from_slice(&self.obs[idx * d..(idx + 1) * d]);
258 next_obs.extend_from_slice(&self.next_obs[idx * d..(idx + 1) * d]);
259 actions.extend_from_slice(&self.actions[idx * a..(idx + 1) * a]);
260 rewards.push(self.rewards[idx]);
261 terminated.push(self.terminated[idx]);
262 }
263
264 if let (Some(mean), Some(std)) = (&self.obs_mean, &self.obs_std) {
266 for i in 0..batch_size {
267 for j in 0..d {
268 obs[i * d + j] = (obs[i * d + j] - mean[j]) / std[j];
269 next_obs[i * d + j] = (next_obs[i * d + j] - mean[j]) / std[j];
270 }
271 }
272 }
273
274 OfflineBatch {
275 obs,
276 next_obs,
277 actions,
278 rewards,
279 terminated,
280 obs_dim: d,
281 act_dim: a,
282 }
283 }
284
285 pub fn sample_trajectories(
291 &self,
292 batch_size: usize,
293 seq_len: usize,
294 seed: u64,
295 ) -> TrajectoryBatch {
296 let mut rng = ChaCha8Rng::seed_from_u64(seed);
297 let d = self.obs_dim;
298 let a = self.act_dim;
299 let n_eps = self.n_episodes();
300
301 let total = batch_size * seq_len;
302 let mut obs = vec![0.0f32; total * d];
303 let mut actions = vec![0.0f32; total * a];
304 let mut rewards = vec![0.0f32; total];
305 let mut returns_to_go = vec![0.0f32; total];
306 let mut timesteps = vec![0u32; total];
307 let mut mask = vec![0u8; total];
308
309 for b in 0..batch_size {
310 let ep_idx = rng.random_range(0..n_eps);
311 let ep_start = self.episode_starts[ep_idx];
312 let ep_len = self.episode_lengths[ep_idx];
313
314 let max_start = ep_len.saturating_sub(seq_len);
316 let start_offset = rng.random_range(0..=max_start);
317 let actual_len = seq_len.min(ep_len - start_offset);
318
319 let mut rtg = vec![0.0f32; actual_len];
321 if actual_len > 0 {
322 rtg[actual_len - 1] = self.rewards[ep_start + start_offset + actual_len - 1];
323 for t in (0..actual_len - 1).rev() {
324 rtg[t] = self.rewards[ep_start + start_offset + t] + rtg[t + 1];
325 }
326 }
327
328 for (t, rtg_val) in rtg.iter().enumerate() {
329 let src_idx = ep_start + start_offset + t;
330 let dst_idx = b * seq_len + t;
331
332 obs[dst_idx * d..(dst_idx + 1) * d]
333 .copy_from_slice(&self.obs[src_idx * d..(src_idx + 1) * d]);
334 actions[dst_idx * a..(dst_idx + 1) * a]
335 .copy_from_slice(&self.actions[src_idx * a..(src_idx + 1) * a]);
336 rewards[dst_idx] = self.rewards[src_idx];
337 returns_to_go[dst_idx] = *rtg_val;
338 timesteps[dst_idx] = (start_offset + t) as u32;
339 mask[dst_idx] = 1;
340 }
341 }
342
343 TrajectoryBatch {
344 obs,
345 actions,
346 rewards,
347 returns_to_go,
348 timesteps,
349 mask,
350 seq_len,
351 obs_dim: d,
352 act_dim: a,
353 }
354 }
355
356 pub fn stats(&self) -> DatasetStats {
358 let returns = &self.episode_returns;
359 let n_eps = returns.len();
360
361 let mean_return = if n_eps > 0 {
362 returns.iter().sum::<f32>() / n_eps as f32
363 } else {
364 0.0
365 };
366
367 let std_return = if n_eps > 1 {
368 let var: f32 = returns
369 .iter()
370 .map(|&r| (r - mean_return).powi(2))
371 .sum::<f32>()
372 / (n_eps - 1) as f32;
373 var.sqrt()
374 } else {
375 0.0
376 };
377
378 let min_return = returns.iter().cloned().reduce(f32::min).unwrap_or(0.0);
379 let max_return = returns.iter().cloned().reduce(f32::max).unwrap_or(0.0);
380
381 let mean_ep_len = if n_eps > 0 {
382 self.episode_lengths.iter().sum::<usize>() as f32 / n_eps as f32
383 } else {
384 0.0
385 };
386
387 DatasetStats {
388 n_transitions: self.len,
389 n_episodes: n_eps,
390 obs_dim: self.obs_dim,
391 act_dim: self.act_dim,
392 mean_return,
393 std_return,
394 min_return,
395 max_return,
396 mean_episode_length: mean_ep_len,
397 }
398 }
399}
400
401#[cfg(test)]
402mod tests {
403 use super::*;
404
405 fn make_test_dataset(
406 n: usize,
407 obs_dim: usize,
408 act_dim: usize,
409 ep_len: usize,
410 ) -> OfflineDatasetBuffer {
411 let rewards = vec![1.0f32; n];
412 let mut terminated = vec![0u8; n];
413 let truncated = vec![0u8; n];
414
415 for (i, t) in terminated.iter_mut().enumerate().take(n) {
417 if (i + 1).is_multiple_of(ep_len) {
418 *t = 1;
419 }
420 }
421
422 OfflineDatasetBuffer::from_arrays(
423 vec![0.1f32; n * obs_dim],
424 vec![0.2f32; n * obs_dim],
425 vec![0.0f32; n * act_dim],
426 rewards,
427 terminated,
428 truncated,
429 obs_dim,
430 act_dim,
431 )
432 .unwrap()
433 }
434
435 #[test]
436 fn test_load_from_arrays() {
437 let buf = make_test_dataset(100, 4, 1, 10);
438 assert_eq!(buf.len(), 100);
439 assert_eq!(buf.obs_dim(), 4);
440 assert_eq!(buf.act_dim(), 1);
441 }
442
443 #[test]
444 fn test_episode_boundary_detection() {
445 let buf = make_test_dataset(100, 4, 1, 10);
446 assert_eq!(buf.n_episodes(), 10);
447 assert_eq!(buf.episode_lengths, vec![10; 10]);
448 }
449
450 #[test]
451 fn test_episode_returns() {
452 let buf = make_test_dataset(100, 4, 1, 10);
453 for &ret in &buf.episode_returns {
455 assert!((ret - 10.0).abs() < 1e-5);
456 }
457 }
458
459 #[test]
460 fn test_sample_uniform_shapes() {
461 let buf = make_test_dataset(1000, 4, 2, 100);
462 let batch = buf.sample(32, 42);
463 assert_eq!(batch.obs.len(), 32 * 4);
464 assert_eq!(batch.next_obs.len(), 32 * 4);
465 assert_eq!(batch.actions.len(), 32 * 2);
466 assert_eq!(batch.rewards.len(), 32);
467 assert_eq!(batch.terminated.len(), 32);
468 }
469
470 #[test]
471 fn test_sample_deterministic() {
472 let buf = make_test_dataset(1000, 4, 1, 100);
473 let b1 = buf.sample(32, 42);
474 let b2 = buf.sample(32, 42);
475 assert_eq!(b1.obs, b2.obs);
476 assert_eq!(b1.rewards, b2.rewards);
477 }
478
479 #[test]
480 fn test_sample_different_seeds() {
481 let n = 1000;
483 let obs_dim = 4;
484 let obs: Vec<f32> = (0..n * obs_dim).map(|i| i as f32 * 0.001).collect();
485 let mut terminated = vec![0u8; n];
486 for i in (99..n).step_by(100) {
487 terminated[i] = 1;
488 }
489 let buf = OfflineDatasetBuffer::from_arrays(
490 obs.clone(),
491 obs,
492 vec![0.0; n],
493 vec![1.0; n],
494 terminated,
495 vec![0; n],
496 obs_dim,
497 1,
498 )
499 .unwrap();
500
501 let b1 = buf.sample(32, 42);
502 let b2 = buf.sample(32, 99);
503 assert_ne!(
504 b1.obs, b2.obs,
505 "Different seeds should produce different samples"
506 );
507 }
508
509 #[test]
510 fn test_normalization() {
511 let mut buf = make_test_dataset(1000, 4, 1, 100);
512 buf.compute_normalization();
513 assert!(buf.obs_mean.is_some());
514 assert!(buf.obs_std.is_some());
515
516 let batch = buf.sample(32, 42);
517 let mean: f32 = batch.obs.iter().sum::<f32>() / batch.obs.len() as f32;
519 assert!(
520 mean.abs() < 1.0,
521 "Normalized mean should be near 0, got {mean}"
522 );
523 }
524
525 #[test]
526 fn test_sample_trajectories_shapes() {
527 let buf = make_test_dataset(1000, 4, 2, 100);
528 let batch = buf.sample_trajectories(8, 20, 42);
529 assert_eq!(batch.obs.len(), 8 * 20 * 4);
530 assert_eq!(batch.actions.len(), 8 * 20 * 2);
531 assert_eq!(batch.rewards.len(), 8 * 20);
532 assert_eq!(batch.returns_to_go.len(), 8 * 20);
533 assert_eq!(batch.timesteps.len(), 8 * 20);
534 assert_eq!(batch.mask.len(), 8 * 20);
535 }
536
537 #[test]
538 fn test_sample_trajectories_mask() {
539 let buf = make_test_dataset(50, 4, 1, 5); let batch = buf.sample_trajectories(4, 10, 42); for b in 0..4 {
545 let valid: usize = (0..10).map(|t| batch.mask[b * 10 + t] as usize).sum();
546 assert!(
547 valid <= 5,
548 "Valid mask count should be <= ep_len=5, got {valid}"
549 );
550 assert!(valid > 0, "Should have at least 1 valid step");
551 }
552 }
553
554 #[test]
555 fn test_sample_trajectories_returns_to_go() {
556 let buf = make_test_dataset(100, 4, 1, 10);
557 let batch = buf.sample_trajectories(1, 10, 42);
558
559 let mut prev_rtg = f32::MAX;
561 for t in 0..10 {
562 if batch.mask[t] == 1 {
563 assert!(
564 batch.returns_to_go[t] <= prev_rtg + 1e-5,
565 "RTG should be non-increasing, got {} after {}",
566 batch.returns_to_go[t],
567 prev_rtg
568 );
569 prev_rtg = batch.returns_to_go[t];
570 }
571 }
572 }
573
574 #[test]
575 fn test_stats() {
576 let buf = make_test_dataset(100, 4, 1, 10);
577 let stats = buf.stats();
578 assert_eq!(stats.n_transitions, 100);
579 assert_eq!(stats.n_episodes, 10);
580 assert_eq!(stats.obs_dim, 4);
581 assert_eq!(stats.act_dim, 1);
582 assert!((stats.mean_return - 10.0).abs() < 1e-5);
583 assert!((stats.mean_episode_length - 10.0).abs() < 1e-5);
584 }
585
586 #[test]
587 fn test_empty_dataset_error() {
588 let result =
589 OfflineDatasetBuffer::from_arrays(vec![], vec![], vec![], vec![], vec![], vec![], 4, 1);
590 assert!(result.is_ok());
592 assert_eq!(result.unwrap().len(), 0);
593 }
594
595 #[test]
596 fn test_mismatched_lengths_error() {
597 let result = OfflineDatasetBuffer::from_arrays(
598 vec![0.0; 40], vec![0.0; 40],
600 vec![0.0; 10], vec![0.0; 10],
602 vec![0; 5], vec![0; 10],
604 4,
605 1,
606 );
607 assert!(result.is_err());
608 }
609
610 #[test]
611 fn test_variable_episode_lengths() {
612 let n = 25; let obs_dim = 2;
615 let act_dim = 1;
616 let mut terminated = vec![0u8; n];
617 terminated[4] = 1; terminated[12] = 1; terminated[24] = 1; let buf = OfflineDatasetBuffer::from_arrays(
622 vec![0.0; n * obs_dim],
623 vec![0.0; n * obs_dim],
624 vec![0.0; n * act_dim],
625 vec![1.0; n],
626 terminated,
627 vec![0; n],
628 obs_dim,
629 act_dim,
630 )
631 .unwrap();
632
633 assert_eq!(buf.n_episodes(), 3);
634 assert_eq!(buf.episode_lengths, vec![5, 8, 12]);
635 }
636}