Skip to main content

r2l_sampler/staged/
mod.rs

1mod worker;
2
3use std::sync::Arc;
4
5use bimodal_array::{ArrayHandle, bimodal_array, bimodal_array_with_factory};
6use itertools::Itertools;
7use r2l_core::{
8    buffers::buffer::{TrajectoryBuffer, TrajectoryView},
9    env::{
10        Env, EnvBuilder, EnvBuilderType,
11        normalizer::{ClippedNormalizer, NormalizerMode},
12    },
13    error::Result,
14    models::Actor,
15    on_policy::algorithm::Sampler,
16    rng::sample_u64,
17    tensor::R2lTensor,
18};
19
20use crate::{
21    RolloutMode, SamplerExecutionMode, SamplerHookResult,
22    staged::{
23        worker::ThreadHandle,
24        worker::{ThreadWorkerFactory, ThreadWorkers, VecWorkers, WorkerPool},
25    },
26};
27
28/// Hook that controls the sequence of collection bounds for a staged sampler.
29pub trait StagedSamplerHook {
30    /// Environment type sampled by the hook's sampler.
31    type E: Env<Tensor: R2lTensor>;
32
33    /// Returns the next collection instruction.
34    fn hook(&mut self, core: &mut StagedSamplerCore<Self::E>) -> SamplerHookResult;
35
36    /// Resets hook state before a new training or evaluation run.
37    fn reset(&mut self) {}
38}
39
40/// Mutable staged-sampler state exposed to hook implementations.
41pub struct StagedSamplerCore<E: Env> {
42    pool: WorkerPool<E>,
43    obs_normalizer: Option<ClippedNormalizer<E::Tensor>>,
44    last_states: ArrayHandle<E::Tensor>,
45    buffers: Vec<TrajectoryBuffer<E::Tensor>>,
46}
47
48impl<E: Env> StagedSamplerCore<E> {
49    /// Returns the per-environment output trajectory buffers mutably.
50    pub fn buffers_mut(&mut self) -> &mut Vec<TrajectoryBuffer<E::Tensor>> {
51        &mut self.buffers
52    }
53
54    /// Builds staged sampler state and its environment workers.
55    ///
56    /// # Panics
57    ///
58    /// Panics if an environment cannot be built or reset.
59    ///
60    /// # Errors
61    ///
62    /// Returns an error if the initial observations cannot be normalized.
63    pub fn build<EB: EnvBuilder<Env = E>>(
64        env_builder: &EnvBuilderType<EB>,
65        execution_mode: SamplerExecutionMode,
66        obs_normalizer: Option<ClippedNormalizer<E::Tensor>>,
67    ) -> Result<Self> {
68        let num_envs = env_builder.num_envs();
69        let buffers = vec![TrajectoryBuffer::default(); num_envs];
70        let (mut last_states, pool) = match execution_mode {
71            SamplerExecutionMode::SingleThreaded => Self::build_vec_workers(env_builder, num_envs),
72            SamplerExecutionMode::MultiThreaded => {
73                Self::build_thread_workers(env_builder, num_envs)
74            }
75        };
76        if let Some(obs_normalizer) = &obs_normalizer {
77            let mut last_states = last_states.lock().unwrap();
78            obs_normalizer.apply_slice_in_place(&mut last_states)?;
79        }
80        Ok(Self {
81            pool,
82            obs_normalizer,
83            last_states,
84            buffers,
85        })
86    }
87
88    fn build_vec_workers<EB: EnvBuilder<Env = E>>(
89        env_builder: &EnvBuilderType<EB>,
90        num_envs: usize,
91    ) -> (ArrayHandle<E::Tensor>, WorkerPool<E>) {
92        let mut envs = Vec::with_capacity(num_envs);
93        let mut initial_states = Vec::with_capacity(num_envs);
94        for env_idx in 0..num_envs {
95            let mut env = env_builder.build_idx(env_idx).unwrap();
96            let state = env.reset(sample_u64()).unwrap();
97            initial_states.push(state.clone());
98            envs.push(env);
99        }
100        let (last_states, last_state_handles) = bimodal_array(initial_states);
101        let workers = envs.into_iter().zip(last_state_handles).collect();
102        (last_states, WorkerPool::Vec(VecWorkers::new(workers)))
103    }
104
105    fn build_thread_workers<EB: EnvBuilder<Env = E>>(
106        env_builder: &EnvBuilderType<EB>,
107        num_envs: usize,
108    ) -> (ArrayHandle<E::Tensor>, WorkerPool<E>) {
109        let mut worker_handles = Vec::with_capacity(num_envs);
110        let factories = (0..num_envs)
111            .map(|idx| {
112                let (command_tx, command_rx) = crossbeam::channel::unbounded();
113                let (result_tx, result_rx) = crossbeam::channel::unbounded();
114                worker_handles.push(ThreadHandle::new(command_tx, result_rx));
115                let env_builder = env_builder.clone();
116                let env_builder = move || env_builder.build_idx(idx);
117                ThreadWorkerFactory::new(command_rx, result_tx, env_builder.clone(), sample_u64())
118            })
119            .collect();
120        let last_states = bimodal_array_with_factory(factories);
121        let workers = ThreadWorkers::new(worker_handles);
122        (last_states, WorkerPool::Thread(workers))
123    }
124
125    /// Collects a bounded rollout from the worker pool.
126    ///
127    /// # Errors
128    ///
129    /// Returns an error if an environment step, normalization, or worker operation fails.
130    pub fn collect(&mut self, bound: RolloutMode) -> Result<()> {
131        match bound {
132            RolloutMode::StepBound { n_steps } => {
133                for _ in 0..n_steps {
134                    self.step()?;
135                }
136            }
137            RolloutMode::EpisodeBound { n_episodes } => {
138                let mut episode_counts = vec![0; self.buffers.len()];
139                loop {
140                    let worker_idxs = episode_counts
141                        .iter()
142                        .positions(|count| *count < n_episodes)
143                        .collect::<Vec<_>>();
144                    if worker_idxs.is_empty() {
145                        break;
146                    }
147                    let terminations = self.step_indexed(&worker_idxs)?;
148                    for (idx, terminated) in worker_idxs.into_iter().zip(terminations) {
149                        if terminated {
150                            episode_counts[idx] += 1;
151                        }
152                    }
153                }
154            }
155        }
156        Ok(())
157    }
158
159    fn step_indexed(&mut self, indices: &[usize]) -> Result<Vec<bool>> {
160        let mut multi_memory = self.pool.step_indexed(indices)?;
161        if let Some(obs_normalizer) = &self.obs_normalizer {
162            let mut last_states = self.last_states.lock().unwrap();
163            let mut next_states = indices
164                .iter()
165                .map(|idx| last_states[*idx].clone())
166                .collect::<Vec<_>>();
167            obs_normalizer.apply_slice_in_place(&mut next_states)?;
168            for (idx, next_state) in indices.iter().zip(next_states) {
169                last_states[*idx] = next_state;
170            }
171            obs_normalizer
172                .with_mode(NormalizerMode::ReadOnly)
173                .apply_slice_in_place(multi_memory.next_states_mut())?;
174        }
175        let memories = multi_memory.into_stored_memories();
176        let terminations = memories
177            .iter()
178            .map(r2l_core::buffers::Memory::is_done)
179            .collect();
180        for (idx, memory) in indices.iter().zip(memories) {
181            self.buffers[*idx].push(memory);
182        }
183        Ok(terminations)
184    }
185
186    fn step(&mut self) -> Result<()> {
187        let mut multi_memory = self.pool.step()?;
188        if let Some(obs_normalizer) = &self.obs_normalizer {
189            let mut last_states = self.last_states.lock().unwrap();
190            obs_normalizer.apply_slice_in_place(&mut last_states)?;
191            obs_normalizer
192                .with_mode(NormalizerMode::ReadOnly)
193                .apply_slice_in_place(multi_memory.next_states_mut())?;
194        }
195        let memories = multi_memory.into_stored_memories();
196        for (idx, memory) in memories.into_iter().enumerate() {
197            self.buffers[idx].push(memory);
198        }
199        Ok(())
200    }
201
202    /// Clears all output trajectory buffers.
203    pub fn clear_buffers(&mut self) {
204        self.buffers
205            .iter_mut()
206            .for_each(r2l_core::buffers::buffer::TrajectoryBuffer::clear);
207    }
208
209    /// Installs a clone of `policy` on every worker.
210    ///
211    /// # Errors
212    ///
213    /// Returns an error if a worker has stopped or cannot acknowledge the update.
214    pub fn set_policy<A: Actor<Tensor = E::Tensor> + Clone>(&mut self, policy: &A) -> Result<()> {
215        self.pool.set_policy(policy)
216    }
217
218    /// Borrows all collected trajectories in worker order.
219    pub fn trajectory_views(&mut self) -> impl AsRef<[TrajectoryView<'_, E::Tensor>]> {
220        self.buffers
221            .iter()
222            .map(|buffer| buffer.to_trajectory_view())
223            .collect::<Vec<_>>()
224    }
225}
226
227/// Observation-normalizing rollout sampler controlled by a hook.
228pub struct StagedSampler<E: Env<Tensor: R2lTensor>, H: StagedSamplerHook<E = E>> {
229    core: StagedSamplerCore<E>,
230    hook: H,
231}
232
233impl<E: Env<Tensor: R2lTensor>, H: StagedSamplerHook<E = E>> StagedSampler<E, H> {
234    /// Creates a staged sampler from its core state and hook.
235    pub fn new(core: StagedSamplerCore<E>, hook: H) -> Self {
236        Self { core, hook }
237    }
238
239    /// Returns the shared observation normalizer, when configured.
240    pub fn obs_normalizer(&self) -> Option<&ClippedNormalizer<E::Tensor>> {
241        self.core.obs_normalizer.as_ref()
242    }
243
244    /// Builds a sampler with an existing shared observation normalizer.
245    ///
246    /// # Errors
247    ///
248    /// Returns an error if the staged sampler core cannot be initialized.
249    pub fn build_with_obs_normalizer<EB: EnvBuilder<Env = E>>(
250        env_builder: &EnvBuilderType<EB>,
251        hook: H,
252        execution_mode: SamplerExecutionMode,
253        obs_normalizer: Option<ClippedNormalizer<E::Tensor>>,
254    ) -> Result<Self> {
255        Ok(Self {
256            core: StagedSamplerCore::build(env_builder, execution_mode, obs_normalizer)?,
257            hook,
258        })
259    }
260
261    /// Builds a homogeneous sampler from a shared environment builder.
262    ///
263    /// # Errors
264    ///
265    /// Returns an error if `num_envs` is zero.
266    pub fn build_from_env_builder(
267        env_builder: Arc<dyn EnvBuilder<Env = E>>,
268        num_envs: usize,
269        hook: H,
270        execution_mode: SamplerExecutionMode,
271        obs_normalizer: Option<ClippedNormalizer<E::Tensor>>,
272    ) -> Result<Self>
273    where
274        E: 'static,
275    {
276        let env_builder = move || env_builder.build_env();
277        Self::build_with_obs_normalizer(
278            &EnvBuilderType::homogeneous(env_builder, num_envs)?,
279            hook,
280            execution_mode,
281            obs_normalizer,
282        )
283    }
284}
285
286impl<E: Env<Tensor: R2lTensor>, H: StagedSamplerHook<E = E>> Sampler for StagedSampler<E, H> {
287    type Tensor = E::Tensor;
288
289    fn reset_all_envs(&mut self) -> Result<()> {
290        self.core.pool.reset_all()?;
291        if let Some(obs_normalizer) = &self.core.obs_normalizer {
292            let mut last_states = self.core.last_states.lock().unwrap();
293            obs_normalizer.apply_slice_in_place(&mut last_states)?;
294        }
295        self.core.clear_buffers();
296        self.hook.reset();
297        Ok(())
298    }
299
300    fn collect_rollouts<A: Actor<Tensor = Self::Tensor> + Clone>(
301        &mut self,
302        actor: A,
303    ) -> Result<()> {
304        self.core.clear_buffers();
305        self.core.set_policy(&actor)?;
306        loop {
307            let result = self.hook.hook(&mut self.core);
308            match result {
309                SamplerHookResult::Bound(bound) => self.core.collect(bound)?,
310                SamplerHookResult::Stop => break,
311            }
312        }
313        Ok(())
314    }
315
316    fn trajectory_views(&mut self) -> impl AsRef<[TrajectoryView<'_, Self::Tensor>]> {
317        self.core.trajectory_views()
318    }
319}