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
28pub trait StagedSamplerHook {
30 type E: Env<Tensor: R2lTensor>;
32
33 fn hook(&mut self, core: &mut StagedSamplerCore<Self::E>) -> SamplerHookResult;
35
36 fn reset(&mut self) {}
38}
39
40pub 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 pub fn buffers_mut(&mut self) -> &mut Vec<TrajectoryBuffer<E::Tensor>> {
51 &mut self.buffers
52 }
53
54 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 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 pub fn clear_buffers(&mut self) {
204 self.buffers
205 .iter_mut()
206 .for_each(r2l_core::buffers::buffer::TrajectoryBuffer::clear);
207 }
208
209 pub fn set_policy<A: Actor<Tensor = E::Tensor> + Clone>(&mut self, policy: &A) -> Result<()> {
215 self.pool.set_policy(policy)
216 }
217
218 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
227pub 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 pub fn new(core: StagedSamplerCore<E>, hook: H) -> Self {
236 Self { core, hook }
237 }
238
239 pub fn obs_normalizer(&self) -> Option<&ClippedNormalizer<E::Tensor>> {
241 self.core.obs_normalizer.as_ref()
242 }
243
244 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 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}