1use 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#[derive(Debug, Clone, Copy)]
18pub enum HERStrategy {
19 Final,
21 Future {
23 k: usize,
25 },
26 Episode,
28}
29
30impl Default for HERStrategy {
31 fn default() -> Self {
32 HERStrategy::Future { k: 4 }
33 }
34}
35
36#[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 #[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 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 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 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 for _ in 0..n_relabeled {
168 let ep_idx = complete[rng.random_range(0..complete.len())];
170 let ep = &episodes[ep_idx];
171
172 let trans_offset = rng.random_range(0..ep.length);
174 let trans_idx = (ep.start + trans_offset) % self.capacity;
175
176 let (obs, next_obs, action, _reward, terminated, truncated) =
178 self.buffer.get(trans_idx);
179
180 let relabel_offset = match self.strategy {
182 HERStrategy::Final => ep.length - 1,
183 HERStrategy::Future { .. } => {
184 if trans_offset >= ep.length - 1 {
185 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 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 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 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 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 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 pub fn len(&self) -> usize {
257 self.buffer.len()
258 }
259
260 pub fn is_empty(&self) -> bool {
262 self.buffer.is_empty()
263 }
264
265 pub fn num_complete_episodes(&self) -> usize {
267 self.tracker.num_complete_episodes()
268 }
269}
270
271#[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 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; HERBuffer::new(
308 capacity,
309 obs_dim,
310 1, goal_dim,
312 core_dim, core_dim + goal_dim, HERStrategy::default(), 0.05, )
317 }
318
319 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; 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); }
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 for _ in 0..10 {
515 push_goal_episode(&mut buf, 10, goal_dim);
516 }
517
518 let batch = buf.sample_with_relabeling(32, 0.0, 42).unwrap();
520 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 for _ in 0..10 {
535 push_goal_episode(&mut buf, 10, goal_dim);
536 }
537 assert_eq!(buf.len(), 50);
538 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); 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}