Skip to main content

rlox_core/buffer/
offline.rs

1//! Read-only offline dataset buffer for offline RL algorithms.
2//!
3//! Unlike [`ReplayBuffer`], this buffer is loaded once from a static dataset
4//! and never modified. It supports:
5//! - Uniform i.i.d. transition sampling (for TD3+BC, IQL, CQL, BC)
6//! - Trajectory subsequence sampling (for Decision Transformer)
7//! - Return-conditioned sampling (for return-conditioned methods)
8//! - Dataset normalization statistics
9//!
10//! Designed for D4RL/Minari-scale datasets (1M+ transitions).
11
12use rand::Rng;
13use rand::SeedableRng;
14use rand_chacha::ChaCha8Rng;
15
16use crate::error::RloxError;
17
18/// Statistics about the loaded dataset.
19#[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/// A batch of i.i.d. sampled transitions.
33#[derive(Debug, Clone)]
34pub struct OfflineBatch {
35    pub obs: Vec<f32>,       // [batch_size * obs_dim]
36    pub next_obs: Vec<f32>,  // [batch_size * obs_dim]
37    pub actions: Vec<f32>,   // [batch_size * act_dim]
38    pub rewards: Vec<f32>,   // [batch_size]
39    pub terminated: Vec<u8>, // [batch_size]
40    pub obs_dim: usize,
41    pub act_dim: usize,
42}
43
44/// A batch of contiguous trajectory subsequences.
45#[derive(Debug, Clone)]
46pub struct TrajectoryBatch {
47    pub obs: Vec<f32>,           // [batch_size * seq_len * obs_dim]
48    pub actions: Vec<f32>,       // [batch_size * seq_len * act_dim]
49    pub rewards: Vec<f32>,       // [batch_size * seq_len]
50    pub returns_to_go: Vec<f32>, // [batch_size * seq_len]
51    pub timesteps: Vec<u32>,     // [batch_size * seq_len]
52    pub mask: Vec<u8>,           // [batch_size * seq_len] (1 = valid, 0 = padding)
53    pub seq_len: usize,
54    pub obs_dim: usize,
55    pub act_dim: usize,
56}
57
58/// Read-only offline dataset buffer.
59pub 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 boundary tracking
69    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    // Normalization (computed lazily)
78    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    /// Create from flat arrays.
86    ///
87    /// Arrays must be row-major: obs has length `n * obs_dim`, etc.
88    // Each parameter is a distinct flat array or dimension; grouping into a struct would complicate the Python FFI caller.
89    #[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        // Detect episode boundaries
132        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    /// Number of transitions in the dataset.
175    pub fn len(&self) -> usize {
176        self.len
177    }
178
179    pub fn is_empty(&self) -> bool {
180        self.len == 0
181    }
182
183    /// Number of episodes in the dataset.
184    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    /// Compute and cache normalization statistics.
197    #[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        // Obs mean and std
203        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        // Reward mean and std
228        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    /// Sample i.i.d. transitions uniformly.
243    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        // Apply normalization if available
265        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    /// Sample contiguous trajectory subsequences.
286    ///
287    /// Each sample is a contiguous window of `seq_len` transitions from a
288    /// single episode. If the episode is shorter than `seq_len`, the sequence
289    /// is right-padded with zeros and the mask indicates valid positions.
290    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            // Random start within episode
315            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            // Compute returns-to-go for this episode segment
320            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    /// Get dataset statistics.
357    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        // Mark episode boundaries
416        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        // Each episode has 10 steps with reward 1.0 → return = 10.0
454        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        // Use varying obs so different indices produce different data
482        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        // Normalized obs should have roughly zero mean
518        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        // Short episodes → padding
540        let buf = make_test_dataset(50, 4, 1, 5); // 10 episodes of length 5
541        let batch = buf.sample_trajectories(4, 10, 42); // request seq_len=10
542
543        // Each trajectory comes from ep_len=5 episode, so at most 5 valid
544        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        // Returns-to-go should be decreasing within valid region
560        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        // Empty is technically valid (0 transitions)
591        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], // 10 * 4
599            vec![0.0; 40],
600            vec![0.0; 10], // 10 * 1
601            vec![0.0; 10],
602            vec![0; 5], // WRONG: should be 10
603            vec![0; 10],
604            4,
605            1,
606        );
607        assert!(result.is_err());
608    }
609
610    #[test]
611    fn test_variable_episode_lengths() {
612        // Create dataset with variable-length episodes
613        let n = 25; // episodes: 5 + 8 + 12 = 25
614        let obs_dim = 2;
615        let act_dim = 1;
616        let mut terminated = vec![0u8; n];
617        terminated[4] = 1; // episode 1: steps 0-4
618        terminated[12] = 1; // episode 2: steps 5-12
619        terminated[24] = 1; // episode 3: steps 13-24
620
621        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}