ruda_model/module/param/
running.rs1use super::ParamId;
2use crate::module::{
3 AutodiffModule, Content, Module, ModuleDisplay, ModuleDisplayDefault, ModuleMapper,
4 ModuleVisitor, Param,
5};
6
7use alloc::string::ToString;
8use alloc::vec::Vec;
9
10#[cfg(target_has_atomic = "ptr")]
11use alloc::sync::Arc;
12
13#[cfg(not(target_has_atomic = "ptr"))]
14use portable_atomic_util::Arc;
15
16use ruda_core::stub::Mutex;
17use ruda_tensor::api::{
18 Tensor,
19 backend::{AutodiffBackend, Backend},
20 ops::Device,
21};
22
23#[cfg(feature = "std")]
24mod threading {
25 pub(super) use std::collections::HashMap;
26 pub(super) use std::thread::ThreadId;
27
28 #[inline(always)]
29 pub(super) fn get_thread_current_id() -> ThreadId {
30 std::thread::current().id()
31 }
32}
33
34#[cfg(not(feature = "std"))]
35mod threading {
36 pub(super) use ruda_core::stub::ThreadId;
37 pub(super) use hashbrown::HashMap;
38
39 #[inline(always)]
40 pub(super) fn get_thread_current_id() -> ThreadId {
41 panic!("Current thread id is not available")
42 }
43}
44
45use threading::*;
47
48#[derive(Clone, Debug)]
54pub struct RunningState<V> {
55 id: ParamId,
56 values: Arc<Mutex<HashMap<ThreadId, V>>>,
57 value: Arc<Mutex<V>>,
58}
59
60impl<V> core::fmt::Display for RunningState<V> {
63 fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
64 write!(f, "RunningState(id={})", self.id)
65 }
66}
67
68impl<V> ModuleDisplayDefault for RunningState<V> {
69 fn content(&self, content: Content) -> Option<Content> {
70 content
71 .add_formatted(&"RunningState".to_string())
72 .optional()
73 }
74}
75
76impl<V> ModuleDisplay for RunningState<V> {}
77
78impl<const D: usize, B: Backend> Module<B> for RunningState<Tensor<B, D>> {
79 type Record = Param<Tensor<B, D>>;
80
81 fn visit<V: ModuleVisitor<B>>(&self, visitor: &mut V) {
82 let tensor = self.value.lock().unwrap();
83 let param = Param::initialized(self.id, tensor.clone());
84 visitor.visit_float(¶m)
85 }
86
87 fn map<M: ModuleMapper<B>>(self, mapper: &mut M) -> Self {
88 let mut tensor = self.value.lock().unwrap();
89 let param = Param::initialized(self.id, tensor.clone());
90 let param_out = mapper.map_float(param);
91 let (_, tensor_out, _) = param_out.consume();
92
93 *tensor = tensor_out;
94 core::mem::drop(tensor);
95
96 self
97 }
98
99 fn into_record(self) -> Self::Record {
100 Param::initialized(self.id, self.synchronized_value())
101 }
102
103 fn load_record(mut self, record: Self::Record) -> Self {
104 let mut pending = self.values.lock().unwrap();
105 let mut tensor = self.value.lock().unwrap();
106 *tensor = record.val().to_device(&tensor.device());
107 pending.clear();
108 self.id = record.id;
109
110 core::mem::drop(tensor);
111 core::mem::drop(pending);
112
113 self
114 }
115
116 fn to_device(self, device: &Device<B>) -> Self {
117 let mut pending = self.values.lock().unwrap();
118 let mut tensor = self.value.lock().unwrap();
119 let tensor_out = tensor.clone().to_device(device);
120 for update in pending.values_mut() {
121 *update = update.clone().to_device(device);
122 }
123
124 *tensor = tensor_out;
125 core::mem::drop(tensor);
126 core::mem::drop(pending);
127
128 self
129 }
130
131 fn fork(self, device: &Device<B>) -> Self {
132 self.to_device(device) }
134
135 fn collect_devices(&self, mut devices: Vec<Device<B>>) -> Vec<Device<B>> {
136 let device = self.value.lock().unwrap().device();
137
138 if !devices.contains(&device) {
139 devices.push(device)
140 }
141
142 devices
143 }
144}
145
146impl<const D: usize, B: Backend> RunningState<Tensor<B, D>> {
147 pub fn new(value: Tensor<B, D>) -> Self {
149 Self {
150 id: ParamId::new(),
151 values: Arc::new(Mutex::new(HashMap::new())),
152 value: Arc::new(Mutex::new(value)),
153 }
154 }
155
156 pub fn with_id(id: ParamId, value: Tensor<B, D>) -> Self {
158 Self {
159 id,
160 values: Arc::new(Mutex::new(HashMap::new())),
161 value: Arc::new(Mutex::new(value)),
162 }
163 }
164
165 pub fn from_record(record: Param<Tensor<B, D>>) -> Self {
167 let tensor = record.val();
168 Self {
169 id: record.id,
170 values: Arc::new(Mutex::new(HashMap::new())),
171 value: Arc::new(Mutex::new(tensor)),
172 }
173 }
174
175 pub fn update(&self, value: Tensor<B, D>) {
177 let thread_id = get_thread_current_id();
178 let mut map = self.values.lock().unwrap();
179
180 if map.contains_key(&thread_id) {
181 self.update_value(&mut map);
182 }
183
184 map.insert(thread_id, value);
185 }
186
187 pub fn value(&self) -> Tensor<B, D> {
193 let value = self.value.lock().unwrap();
194 value.clone()
195 }
196
197 pub fn value_sync(&self) -> Tensor<B, D> {
204 let thread_id = get_thread_current_id();
205 let mut map = self.values.lock().unwrap();
206
207 if map.contains_key(&thread_id) {
208 self.update_value(&mut map);
209 }
210
211 let value = self.value.lock().unwrap();
212 value.clone()
213 }
214
215 fn synchronized_value(&self) -> Tensor<B, D> {
216 let mut map = self.values.lock().unwrap();
217
218 if !map.is_empty() {
219 self.update_value(&mut map);
220 }
221 self.value.lock().unwrap().clone()
222 }
223
224 fn update_value(&self, map: &mut HashMap<ThreadId, Tensor<B, D>>) {
225 let mut value_updated: Option<Tensor<B, D>> = None;
226 let mut counter = 0;
227
228 for (_key, tensor) in map.drain() {
229 counter += 1;
230
231 value_updated = match value_updated {
232 Some(current) => {
233 let device = current.device();
234 Some(tensor.to_device(&device).add(current))
235 }
236 None => Some(tensor),
237 };
238 }
239
240 if let Some(value) = value_updated {
241 let value = value.div_scalar(counter);
242 let mut value_old = self.value.lock().unwrap();
243 *value_old = value;
244 }
245 }
246}
247
248impl<const D: usize, B: AutodiffBackend> AutodiffModule<B> for RunningState<Tensor<B, D>> {
249 type InnerModule = RunningState<Tensor<B::InnerBackend, D>>;
250
251 fn valid(&self) -> Self::InnerModule {
252 let value = self.synchronized_value();
253
254 RunningState::with_id(self.id, value.inner())
255 }
256
257 fn from_inner(module: Self::InnerModule) -> Self {
258 let value = module.synchronized_value();
259
260 RunningState::with_id(module.id, Tensor::from_inner(value))
261 }
262}