Skip to main content

rlox_core/buffer/
her.rs

1//! Hindsight Experience Replay (HER) buffer.
2//!
3//! Stores transitions with goal information and performs goal relabeling
4//! during sampling (Andrychowicz et al., 2017). Supports Final, Future(k),
5//! and Episode relabeling strategies.
6
7use rand::Rng;
8use rand::SeedableRng;
9use rand_chacha::ChaCha8Rng;
10
11use crate::error::RloxError;
12
13use super::episode::{EpisodeMeta, EpisodeTracker};
14use super::ringbuf::{ReplayBuffer, SampledBatch};
15
16/// HER goal relabeling strategy.
17#[derive(Debug, Clone, Copy)]
18pub enum HERStrategy {
19    /// Replace goal with the final state achieved in the episode.
20    Final,
21    /// Replace goal with a future state sampled uniformly from the remainder.
22    Future {
23        /// Number of relabeled goals per original transition. Default: 4.
24        k: usize,
25    },
26    /// Replace goal with a random state from the episode.
27    Episode,
28}
29
30impl Default for HERStrategy {
31    fn default() -> Self {
32        HERStrategy::Future { k: 4 }
33    }
34}
35
36/// Hindsight Experience Replay buffer.
37///
38/// Stores transitions with goal information and performs goal relabeling
39/// during sampling. The obs vector layout is:
40/// `[obs_core | achieved_goal | desired_goal | ...]`
41#[derive(Debug)]
42pub struct HERBuffer {
43    buffer: ReplayBuffer,
44    tracker: EpisodeTracker,
45    obs_dim: usize,
46    act_dim: usize,
47    goal_dim: usize,
48    achieved_goal_start: usize,
49    desired_goal_start: usize,
50    capacity: usize,
51    strategy: HERStrategy,
52    goal_tolerance: f32,
53}
54
55impl HERBuffer {
56    /// Create a new HER buffer.
57    ///
58    /// # Arguments
59    /// * `capacity` - maximum transitions
60    /// * `obs_dim` - full observation dimension (includes goal components)
61    /// * `act_dim` - action dimension
62    /// * `goal_dim` - goal vector dimension
63    /// * `achieved_goal_start` - index within obs where achieved goal starts
64    /// * `desired_goal_start` - index within obs where desired goal starts
65    /// * `strategy` - relabeling strategy
66    /// * `goal_tolerance` - tolerance for sparse reward computation
67    // All 8 params are distinct required dimensions/config for a HER buffer; no natural grouping without breaking the public API.
68    #[allow(clippy::too_many_arguments)]
69    pub fn new(
70        capacity: usize,
71        obs_dim: usize,
72        act_dim: usize,
73        goal_dim: usize,
74        achieved_goal_start: usize,
75        desired_goal_start: usize,
76        strategy: HERStrategy,
77        goal_tolerance: f32,
78    ) -> Self {
79        Self {
80            buffer: ReplayBuffer::new(capacity, obs_dim, act_dim),
81            tracker: EpisodeTracker::new(capacity),
82            obs_dim,
83            act_dim,
84            goal_dim,
85            achieved_goal_start,
86            desired_goal_start,
87            capacity,
88            strategy,
89            goal_tolerance,
90        }
91    }
92
93    /// Push a single transition, notifying the episode tracker.
94    pub fn push_slices(
95        &mut self,
96        obs: &[f32],
97        next_obs: &[f32],
98        action: &[f32],
99        reward: f32,
100        terminated: bool,
101        truncated: bool,
102    ) -> Result<(), RloxError> {
103        let write_pos = self.buffer.write_pos();
104        let was_full = self.buffer.len() == self.capacity;
105
106        if was_full {
107            self.tracker.invalidate_overwritten(write_pos, 1);
108        }
109
110        self.buffer
111            .push_slices(obs, next_obs, action, reward, terminated, truncated)?;
112
113        let done = terminated || truncated;
114        self.tracker.notify_push(write_pos, done);
115
116        Ok(())
117    }
118
119    /// Sample a batch with HER relabeling.
120    ///
121    /// `her_ratio` controls the fraction of samples that get relabeled goals.
122    /// The remaining samples use their original goals.
123    pub fn sample_with_relabeling(
124        &self,
125        batch_size: usize,
126        her_ratio: f32,
127        seed: u64,
128    ) -> Result<SampledBatch, RloxError> {
129        if self.buffer.is_empty() {
130            return Err(RloxError::BufferError("buffer is empty".into()));
131        }
132
133        let episodes = self.tracker.episodes();
134        let complete: Vec<usize> = episodes
135            .iter()
136            .enumerate()
137            .filter(|(_, ep)| ep.complete)
138            .map(|(i, _)| i)
139            .collect();
140
141        if complete.is_empty() {
142            return Err(RloxError::BufferError(
143                "no complete episodes for HER relabeling".into(),
144            ));
145        }
146
147        let mut rng = ChaCha8Rng::seed_from_u64(seed);
148        let n_relabeled = ((batch_size as f32) * her_ratio).ceil() as usize;
149        let n_original = batch_size - n_relabeled;
150
151        let mut batch = SampledBatch::with_capacity(batch_size, self.obs_dim, self.act_dim);
152
153        // Sample original (unrelabeled) transitions
154        if n_original > 0 {
155            let original = self.buffer.sample(n_original, rng.random())?;
156            batch.observations.extend_from_slice(&original.observations);
157            batch
158                .next_observations
159                .extend_from_slice(&original.next_observations);
160            batch.actions.extend_from_slice(&original.actions);
161            batch.rewards.extend_from_slice(&original.rewards);
162            batch.terminated.extend_from_slice(&original.terminated);
163            batch.truncated.extend_from_slice(&original.truncated);
164        }
165
166        // Sample relabeled transitions
167        for _ in 0..n_relabeled {
168            // Pick a random complete episode
169            let ep_idx = complete[rng.random_range(0..complete.len())];
170            let ep = &episodes[ep_idx];
171
172            // Pick a random transition within the episode
173            let trans_offset = rng.random_range(0..ep.length);
174            let trans_idx = (ep.start + trans_offset) % self.capacity;
175
176            // Get the original transition
177            let (obs, next_obs, action, _reward, terminated, truncated) =
178                self.buffer.get(trans_idx);
179
180            // Compute the relabel index based on strategy
181            let relabel_offset = match self.strategy {
182                HERStrategy::Final => ep.length - 1,
183                HERStrategy::Future { .. } => {
184                    if trans_offset >= ep.length - 1 {
185                        // Already at the end, use the same position
186                        trans_offset
187                    } else {
188                        rng.random_range((trans_offset + 1)..ep.length)
189                    }
190                }
191                HERStrategy::Episode => rng.random_range(0..ep.length),
192            };
193            let relabel_idx = (ep.start + relabel_offset) % self.capacity;
194
195            // Get the achieved goal from the relabel transition's next_obs
196            let (_, relabel_next_obs, _, _, _, _) = self.buffer.get(relabel_idx);
197            let new_goal = &relabel_next_obs
198                [self.achieved_goal_start..self.achieved_goal_start + self.goal_dim];
199
200            // Create modified observation with new desired goal
201            let mut new_obs = obs.to_vec();
202            new_obs[self.desired_goal_start..self.desired_goal_start + self.goal_dim]
203                .copy_from_slice(new_goal);
204
205            let mut new_next_obs = next_obs.to_vec();
206            new_next_obs[self.desired_goal_start..self.desired_goal_start + self.goal_dim]
207                .copy_from_slice(new_goal);
208
209            // Compute new reward based on achieved goal in next_obs vs new desired goal
210            let achieved_in_next =
211                &next_obs[self.achieved_goal_start..self.achieved_goal_start + self.goal_dim];
212            let new_reward = sparse_goal_reward(achieved_in_next, new_goal, self.goal_tolerance);
213
214            batch.observations.extend_from_slice(&new_obs);
215            batch.next_observations.extend_from_slice(&new_next_obs);
216            batch.actions.extend_from_slice(action);
217            batch.rewards.push(new_reward);
218            batch.terminated.push(terminated);
219            batch.truncated.push(truncated);
220        }
221
222        batch.batch_size = batch_size;
223
224        Ok(batch)
225    }
226
227    /// Compute relabeling indices for a given episode and transition.
228    ///
229    /// Returns indices (offsets within the episode) to use as substitute goals.
230    pub fn compute_relabel_indices(
231        &self,
232        episode: &EpisodeMeta,
233        transition_offset: usize,
234        seed: u64,
235    ) -> Vec<usize> {
236        let mut rng = ChaCha8Rng::seed_from_u64(seed);
237        match self.strategy {
238            HERStrategy::Final => vec![episode.length - 1],
239            HERStrategy::Future { k } => {
240                if transition_offset >= episode.length - 1 {
241                    // At the last step, can only relabel with itself
242                    vec![transition_offset; k]
243                } else {
244                    (0..k)
245                        .map(|_| rng.random_range((transition_offset + 1)..episode.length))
246                        .collect()
247                }
248            }
249            HERStrategy::Episode => {
250                vec![rng.random_range(0..episode.length)]
251            }
252        }
253    }
254
255    /// Number of valid transitions currently stored.
256    pub fn len(&self) -> usize {
257        self.buffer.len()
258    }
259
260    /// Whether the buffer is empty.
261    pub fn is_empty(&self) -> bool {
262        self.buffer.is_empty()
263    }
264
265    /// Number of complete episodes currently tracked.
266    pub fn num_complete_episodes(&self) -> usize {
267        self.tracker.num_complete_episodes()
268    }
269}
270
271/// Compute sparse goal-conditioned reward.
272///
273/// Returns `0.0` if `||achieved - desired||_2 < tolerance`, else `-1.0`.
274///
275/// Uses squared distance comparison to avoid a costly `sqrt`.
276#[inline]
277pub fn sparse_goal_reward(achieved: &[f32], desired: &[f32], tolerance: f32) -> f32 {
278    let dist_sq: f32 = achieved
279        .iter()
280        .zip(desired.iter())
281        .map(|(&a, &d)| (a - d) * (a - d))
282        .sum();
283    if dist_sq < tolerance * tolerance {
284        0.0
285    } else {
286        -1.0
287    }
288}
289
290#[cfg(test)]
291mod tests {
292    use super::*;
293
294    /// Helper: build an obs vector with embedded achieved and desired goals.
295    /// Layout: [core(2) | achieved_goal(goal_dim) | desired_goal(goal_dim)]
296    fn make_obs(core: &[f32], achieved: &[f32], desired: &[f32]) -> Vec<f32> {
297        let mut obs = Vec::with_capacity(core.len() + achieved.len() + desired.len());
298        obs.extend_from_slice(core);
299        obs.extend_from_slice(achieved);
300        obs.extend_from_slice(desired);
301        obs
302    }
303
304    fn make_her_buffer(capacity: usize, goal_dim: usize) -> HERBuffer {
305        let core_dim = 2;
306        let obs_dim = core_dim + goal_dim * 2; // core + achieved + desired
307        HERBuffer::new(
308            capacity,
309            obs_dim,
310            1, // act_dim
311            goal_dim,
312            core_dim,               // achieved_goal_start
313            core_dim + goal_dim,    // desired_goal_start
314            HERStrategy::default(), // Future { k: 4 }
315            0.05,                   // goal_tolerance
316        )
317    }
318
319    /// Push an episode where the agent moves from origin toward a goal.
320    fn push_goal_episode(buf: &mut HERBuffer, length: usize, goal_dim: usize) {
321        let desired_goal = vec![10.0; goal_dim];
322        for i in 0..length {
323            let progress = (i as f32 + 1.0) / length as f32;
324            let achieved = vec![10.0 * progress; goal_dim];
325            let next_achieved = vec![10.0 * (progress + 1.0 / length as f32).min(1.0); goal_dim];
326            let core = vec![progress, progress];
327
328            let obs = make_obs(&core, &achieved, &desired_goal);
329            let next_obs = make_obs(
330                &[progress + 0.1, progress + 0.1],
331                &next_achieved,
332                &desired_goal,
333            );
334            let action = vec![0.0];
335            let reward = -1.0; // sparse: not at goal yet
336            let done = i == length - 1;
337
338            buf.push_slices(&obs, &next_obs, &action, reward, done, false)
339                .unwrap();
340        }
341    }
342
343    #[test]
344    fn test_her_new_is_empty() {
345        let buf = make_her_buffer(100, 3);
346        assert_eq!(buf.len(), 0);
347        assert!(buf.is_empty());
348    }
349
350    #[test]
351    fn test_her_push_increments() {
352        let mut buf = make_her_buffer(100, 3);
353        push_goal_episode(&mut buf, 5, 3);
354        assert_eq!(buf.len(), 5);
355        assert_eq!(buf.num_complete_episodes(), 1);
356    }
357
358    #[test]
359    fn test_final_strategy_uses_last_state() {
360        let goal_dim = 2;
361        let core_dim = 2;
362        let obs_dim = core_dim + goal_dim * 2;
363        let mut buf = HERBuffer::new(
364            100,
365            obs_dim,
366            1,
367            goal_dim,
368            core_dim,
369            core_dim + goal_dim,
370            HERStrategy::Final,
371            0.05,
372        );
373        push_goal_episode(&mut buf, 5, goal_dim);
374
375        let ep = &buf.tracker.episodes()[0];
376        let indices = buf.compute_relabel_indices(ep, 2, 42);
377        assert_eq!(indices.len(), 1);
378        assert_eq!(indices[0], 4); // last step of episode (length 5)
379    }
380
381    #[test]
382    fn test_future_strategy_picks_future_state() {
383        let goal_dim = 2;
384        let core_dim = 2;
385        let obs_dim = core_dim + goal_dim * 2;
386        let mut buf = HERBuffer::new(
387            100,
388            obs_dim,
389            1,
390            goal_dim,
391            core_dim,
392            core_dim + goal_dim,
393            HERStrategy::Future { k: 4 },
394            0.05,
395        );
396        push_goal_episode(&mut buf, 10, goal_dim);
397
398        let ep = &buf.tracker.episodes()[0];
399        let indices = buf.compute_relabel_indices(ep, 3, 42);
400        assert_eq!(indices.len(), 4);
401        for &idx in &indices {
402            assert!(
403                idx > 3,
404                "future index {idx} should be > transition offset 3"
405            );
406            assert!(idx < 10, "future index {idx} should be < episode length 10");
407        }
408    }
409
410    #[test]
411    fn test_episode_strategy_picks_any_state() {
412        let goal_dim = 2;
413        let core_dim = 2;
414        let obs_dim = core_dim + goal_dim * 2;
415        let mut buf = HERBuffer::new(
416            100,
417            obs_dim,
418            1,
419            goal_dim,
420            core_dim,
421            core_dim + goal_dim,
422            HERStrategy::Episode,
423            0.05,
424        );
425        push_goal_episode(&mut buf, 10, goal_dim);
426
427        let ep = &buf.tracker.episodes()[0];
428        let indices = buf.compute_relabel_indices(ep, 5, 42);
429        assert_eq!(indices.len(), 1);
430        assert!(indices[0] < 10);
431    }
432
433    #[test]
434    fn test_sparse_goal_reward_achieved() {
435        let achieved = [1.0, 2.0, 3.0];
436        let desired = [1.0, 2.0, 3.0];
437        assert_eq!(sparse_goal_reward(&achieved, &desired, 0.05), 0.0);
438    }
439
440    #[test]
441    fn test_sparse_goal_reward_not_achieved() {
442        let achieved = [1.0, 2.0, 3.0];
443        let desired = [10.0, 20.0, 30.0];
444        assert_eq!(sparse_goal_reward(&achieved, &desired, 0.05), -1.0);
445    }
446
447    #[test]
448    fn test_relabel_indices_future_k4() {
449        let goal_dim = 2;
450        let core_dim = 2;
451        let obs_dim = core_dim + goal_dim * 2;
452        let mut buf = HERBuffer::new(
453            100,
454            obs_dim,
455            1,
456            goal_dim,
457            core_dim,
458            core_dim + goal_dim,
459            HERStrategy::Future { k: 4 },
460            0.05,
461        );
462        push_goal_episode(&mut buf, 10, goal_dim);
463
464        let ep = &buf.tracker.episodes()[0];
465        let indices = buf.compute_relabel_indices(ep, 3, 42);
466        assert_eq!(indices.len(), 4);
467        for &idx in &indices {
468            assert!(idx > 3 && idx < 10);
469        }
470    }
471
472    #[test]
473    fn test_relabel_indices_deterministic() {
474        let goal_dim = 2;
475        let core_dim = 2;
476        let obs_dim = core_dim + goal_dim * 2;
477        let mut buf = HERBuffer::new(
478            100,
479            obs_dim,
480            1,
481            goal_dim,
482            core_dim,
483            core_dim + goal_dim,
484            HERStrategy::Future { k: 4 },
485            0.05,
486        );
487        push_goal_episode(&mut buf, 10, goal_dim);
488
489        let ep = &buf.tracker.episodes()[0];
490        let i1 = buf.compute_relabel_indices(ep, 3, 42);
491        let i2 = buf.compute_relabel_indices(ep, 3, 42);
492        assert_eq!(i1, i2);
493    }
494
495    #[test]
496    fn test_her_sample_batch_shape() {
497        let goal_dim = 3;
498        let mut buf = make_her_buffer(100, goal_dim);
499        push_goal_episode(&mut buf, 10, goal_dim);
500
501        let batch = buf.sample_with_relabeling(8, 0.8, 42).unwrap();
502        let obs_dim = 2 + goal_dim * 2;
503        assert_eq!(batch.batch_size, 8);
504        assert_eq!(batch.observations.len(), 8 * obs_dim);
505        assert_eq!(batch.actions.len(), 8);
506        assert_eq!(batch.rewards.len(), 8);
507    }
508
509    #[test]
510    fn test_her_ratio_controls_relabeling() {
511        let goal_dim = 2;
512        let mut buf = make_her_buffer(200, goal_dim);
513        // Push multiple episodes so we have enough data
514        for _ in 0..10 {
515            push_goal_episode(&mut buf, 10, goal_dim);
516        }
517
518        // With ratio=0.0, no relabeling => all rewards should be -1.0 (original)
519        let batch = buf.sample_with_relabeling(32, 0.0, 42).unwrap();
520        // All original rewards are -1.0
521        for &r in &batch.rewards {
522            assert_eq!(
523                r, -1.0,
524                "with ratio=0, all rewards should be original (-1.0)"
525            );
526        }
527    }
528
529    #[test]
530    fn test_her_with_ring_wrap() {
531        let goal_dim = 2;
532        let mut buf = make_her_buffer(50, goal_dim);
533        // Push 100 transitions (wraps around)
534        for _ in 0..10 {
535            push_goal_episode(&mut buf, 10, goal_dim);
536        }
537        assert_eq!(buf.len(), 50);
538        // Should still be able to sample
539        let result = buf.sample_with_relabeling(4, 0.8, 42);
540        assert!(result.is_ok());
541    }
542
543    #[test]
544    fn test_her_empty_buffer_errors() {
545        let buf = make_her_buffer(100, 3);
546        let result = buf.sample_with_relabeling(4, 0.8, 42);
547        assert!(result.is_err());
548    }
549
550    mod proptests {
551        use super::*;
552        use proptest::prelude::*;
553
554        proptest! {
555            #[test]
556            fn prop_relabel_indices_in_range(
557                ep_len in 2usize..20,
558                trans_offset in 0usize..19,
559            ) {
560                let trans_offset = trans_offset.min(ep_len - 1);
561                let goal_dim = 2;
562                let core_dim = 2;
563                let obs_dim = core_dim + goal_dim * 2;
564                let buf = HERBuffer::new(
565                    100, obs_dim, 1, goal_dim, core_dim, core_dim + goal_dim,
566                    HERStrategy::Future { k: 4 }, 0.05,
567                );
568                let ep = EpisodeMeta { start: 0, length: ep_len, complete: true };
569                let indices = buf.compute_relabel_indices(&ep, trans_offset, 42);
570                for &idx in &indices {
571                    prop_assert!(idx < ep_len, "index {idx} >= episode length {ep_len}");
572                }
573            }
574
575            #[test]
576            fn prop_future_indices_strictly_future(
577                ep_len in 3usize..20,
578                trans_offset in 0usize..18,
579            ) {
580                let trans_offset = trans_offset.min(ep_len - 2); // ensure room for future
581                let goal_dim = 2;
582                let core_dim = 2;
583                let obs_dim = core_dim + goal_dim * 2;
584                let buf = HERBuffer::new(
585                    100, obs_dim, 1, goal_dim, core_dim, core_dim + goal_dim,
586                    HERStrategy::Future { k: 4 }, 0.05,
587                );
588                let ep = EpisodeMeta { start: 0, length: ep_len, complete: true };
589                let indices = buf.compute_relabel_indices(&ep, trans_offset, 42);
590                for &idx in &indices {
591                    prop_assert!(idx > trans_offset,
592                        "future index {idx} should be > offset {trans_offset}");
593                }
594            }
595
596            #[test]
597            fn prop_sparse_reward_binary(
598                a0 in -10.0f32..10.0,
599                a1 in -10.0f32..10.0,
600                d0 in -10.0f32..10.0,
601                d1 in -10.0f32..10.0,
602            ) {
603                let r = sparse_goal_reward(&[a0, a1], &[d0, d1], 0.05);
604                prop_assert!(r == 0.0 || r == -1.0, "reward should be 0.0 or -1.0, got {r}");
605            }
606        }
607    }
608}