r2l_sampler/direct/
mod.rs1pub 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
28pub enum SamplerHookResult {
30 Stop,
32 Bound(RolloutMode),
34}
35
36pub trait DirectSamplerHook {
38 type E: Env;
40
41 fn hook(&mut self, core: &mut DirectSamplerCore<Self::E>) -> SamplerHookResult;
43
44 fn reset(&mut self) {}
46}
47
48pub struct DirectSamplerCore<E: Env> {
50 buffers: ArrayHandle<TrajectoryBuffer<E::Tensor>>,
51 worker_pool: WorkerPool<E>,
52}
53
54impl<E: Env> DirectSamplerCore<E> {
55 pub fn buffers_mut(&mut self) -> &mut ArrayHandle<TrajectoryBuffer<E::Tensor>> {
57 &mut self.buffers
58 }
59
60 pub fn reset_all_envs(&mut self) -> Result<()> {
66 self.worker_pool.reset_all_envs()
67 }
68
69 #[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
124pub 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 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 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}