Skip to main content

r2l_sampler/direct/
mod.rs

1// R2l sampler where each worker writes directly to the output buffer. This is preferred, when the
2// raw observations and rewards are to be stored.
3
4pub mod worker;
5
6use std::sync::Arc;
7
8use bimodal_array::ArrayHandle;
9use bimodal_array::bimodal_array;
10use r2l_core::buffers::buffer::TrajectoryBuffer;
11use r2l_core::buffers::buffer::TrajectoryView;
12use r2l_core::env::Env;
13use r2l_core::env::EnvBuilder;
14use r2l_core::env::EnvBuilderType;
15use r2l_core::error::Result;
16use r2l_core::models::Actor;
17use r2l_core::on_policy::algorithm::Sampler;
18use r2l_core::rng::{sample_u64, set_seed};
19
20use crate::RolloutMode;
21use crate::SamplerExecutionMode;
22use crate::direct::worker::ThreadHandle;
23use crate::direct::worker::ThreadWorker;
24use crate::direct::worker::ThreadWorkers;
25use crate::direct::worker::Worker;
26use crate::direct::worker::WorkerPool;
27
28/// Instruction returned by a [`DirectSamplerHook`] during rollout collection.
29pub enum SamplerHookResult {
30    /// Finish the current rollout.
31    Stop,
32    /// Collect data up to the supplied bound, then invoke the hook again.
33    Bound(RolloutMode),
34}
35
36/// Hook that controls the sequence of collection bounds for a raw sampler.
37pub trait DirectSamplerHook {
38    /// Environment type sampled by the hook's sampler.
39    type E: Env;
40
41    /// Returns the next collection instruction.
42    fn hook(&mut self, core: &mut DirectSamplerCore<Self::E>) -> SamplerHookResult;
43
44    /// Resets hook state before a new training or evaluation run.
45    fn reset(&mut self) {}
46}
47
48/// Mutable direct-sampler state exposed to [`DirectSamplerHook`] implementations.
49pub struct DirectSamplerCore<E: Env> {
50    buffers: ArrayHandle<TrajectoryBuffer<E::Tensor>>,
51    worker_pool: WorkerPool<E>,
52}
53
54impl<E: Env> DirectSamplerCore<E> {
55    /// Returns the per-environment output buffers.
56    pub fn buffers_mut(&mut self) -> &mut ArrayHandle<TrajectoryBuffer<E::Tensor>> {
57        &mut self.buffers
58    }
59
60    /// Resets every worker environment and clears its active episode state.
61    ///
62    /// # Errors
63    ///
64    /// Returns an error if an environment cannot be reset or a worker is interrupted.
65    pub fn reset_all_envs(&mut self) -> Result<()> {
66        self.worker_pool.reset_all_envs()
67    }
68
69    /// Builds sampler state from an environment collection and execution mode.
70    ///
71    /// # Panics
72    ///
73    /// Panics if an environment cannot be built.
74    #[must_use]
75    pub fn build<EB: EnvBuilder<Env = E>>(
76        env_builder: EnvBuilderType<EB>,
77        execution_mode: SamplerExecutionMode,
78    ) -> Self {
79        let num_envs = env_builder.num_envs();
80        let buffers: Vec<TrajectoryBuffer<E::Tensor>> = vec![TrajectoryBuffer::default(); num_envs];
81        let (buffers, buffer_handlers) = bimodal_array(buffers);
82        let worker_pool = match execution_mode {
83            SamplerExecutionMode::SingleThreaded => {
84                let workers: Vec<_> = buffer_handlers
85                    .into_iter()
86                    .enumerate()
87                    .map(|(idx, element_handle)| {
88                        let env = env_builder.build_idx(idx).unwrap();
89                        Worker::new(env, element_handle)
90                    })
91                    .collect();
92                WorkerPool::Vec(workers)
93            }
94            SamplerExecutionMode::MultiThreaded => {
95                let env_builder = Arc::new(env_builder);
96                let workers: Vec<_> = buffer_handlers
97                    .into_iter()
98                    .enumerate()
99                    .map(|(idx, element_handle)| {
100                        let (command_tx, command_rx) = crossbeam::channel::unbounded();
101                        let (res_tx, res_rx) = crossbeam::channel::unbounded();
102                        let env_builder = env_builder.clone();
103                        let worker_seed = sample_u64();
104                        let handle = std::thread::spawn(move || {
105                            set_seed(worker_seed);
106                            let env = env_builder.build_idx(idx).unwrap();
107                            let worker = Worker::new(env, element_handle);
108                            let mut thread_worker = ThreadWorker::new(worker, command_rx, res_tx);
109                            thread_worker.work();
110                        });
111                        ThreadHandle::new(handle, command_tx, res_rx)
112                    })
113                    .collect();
114                WorkerPool::Thread(ThreadWorkers::new(workers))
115            }
116        };
117        Self {
118            buffers,
119            worker_pool,
120        }
121    }
122}
123
124/// Rollout sampler whose workers write directly to output buffers.
125pub struct DirectSampler<E: Env, H: DirectSamplerHook<E = E>> {
126    core: DirectSamplerCore<E>,
127    hook: H,
128}
129
130impl<E: Env, H: DirectSamplerHook<E = E>> DirectSampler<E, H> {
131    pub fn new(core: DirectSamplerCore<E>, hook: H) -> Self {
132        Self { core, hook }
133    }
134
135    /// Builds a raw sampler and its environment workers.
136    pub fn build<EB: EnvBuilder<Env = E>>(
137        env_builder: EnvBuilderType<EB>,
138        hook: H,
139        execution_mode: SamplerExecutionMode,
140    ) -> Self {
141        Self {
142            core: DirectSamplerCore::build(env_builder, execution_mode),
143            hook,
144        }
145    }
146
147    /// Builds a homogeneous sampler from a shared environment builder.
148    ///
149    /// # Errors
150    ///
151    /// Returns an error if `num_envs` is zero.
152    pub fn build_from_env_builder(
153        env_builder: Arc<dyn EnvBuilder<Env = E>>,
154        num_envs: usize,
155        hook: H,
156        execution_mode: SamplerExecutionMode,
157    ) -> Result<Self>
158    where
159        E: 'static,
160    {
161        let env_builder = move || env_builder.build_env();
162        Ok(Self::build(
163            EnvBuilderType::homogeneous(env_builder, num_envs)?,
164            hook,
165            execution_mode,
166        ))
167    }
168}
169
170impl<E: Env, H: DirectSamplerHook<E = E>> Sampler for DirectSampler<E, H> {
171    type Tensor = E::Tensor;
172
173    fn reset_all_envs(&mut self) -> Result<()> {
174        self.core.reset_all_envs()?;
175        self.hook.reset();
176        Ok(())
177    }
178
179    fn collect_rollouts<A: Actor<Tensor = Self::Tensor> + Clone>(
180        &mut self,
181        actor: A,
182    ) -> Result<()> {
183        self.core.worker_pool.clear_buffers();
184        self.core.worker_pool.set_actor(&actor)?;
185        loop {
186            let result = self.hook.hook(&mut self.core);
187            match result {
188                SamplerHookResult::Bound(bound) => self.core.worker_pool.collect(bound)?,
189                SamplerHookResult::Stop => break,
190            }
191        }
192        Ok(())
193    }
194
195    fn trajectory_views(&mut self) -> impl AsRef<[TrajectoryView<'_, Self::Tensor>]> {
196        self.core
197            .buffers
198            .lock_map(|buffer| buffer.to_trajectory_view())
199            .unwrap()
200    }
201}