Skip to main content

ruda_model/module/param/
running.rs

1use 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
45// Re-export items from the disabled/enabled blocks
46use threading::*;
47
48/// A state that can be updated during the forward pass while being thread safe.
49///
50/// # Note
51///
52/// The state value is the average of all updates on all threads.
53#[derive(Clone, Debug)]
54pub struct RunningState<V> {
55    id: ParamId,
56    values: Arc<Mutex<HashMap<ThreadId, V>>>,
57    value: Arc<Mutex<V>>,
58}
59
60// Implement display for the module
61
62impl<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(&param)
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) // Same thing here since no grad.
133    }
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    /// Create a new running state.
148    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    /// Create a new running state.
157    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    /// Create a new running state from a record.
166    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    /// Update the value on the current thread.
176    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    /// Get the current value,
188    ///
189    /// # Note
190    ///
191    /// The current value might be outdated by one update.
192    pub fn value(&self) -> Tensor<B, D> {
193        let value = self.value.lock().unwrap();
194        value.clone()
195    }
196
197    /// Get the current value and make sure it is sync.
198    ///
199    /// # Note
200    ///
201    /// Don't use this function after an update on the same thread where other threads might have to
202    /// register their update before the actual synchronization needs to happen.
203    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}