1use super::ParamId;
2use super::sync_once_cell::SyncOnceCell;
3use alloc::format;
4
5#[cfg(not(target_has_atomic = "ptr"))]
6use alloc::boxed::Box;
7use ruda_core::stub::RwLock;
8use ruda_tensor::api::Shape;
9use core::ops::Deref;
10
11#[cfg(target_has_atomic = "ptr")]
12use alloc::sync::Arc;
13
14#[cfg(not(target_has_atomic = "ptr"))]
15use portable_atomic_util::Arc;
16
17#[cfg(target_has_atomic = "ptr")]
18type Mapper<T> = Arc<dyn Fn(T) -> T + Send + Sync>;
19
20#[cfg(not(target_has_atomic = "ptr"))]
21type Mapper<T> = Arc<Box<dyn Fn(T) -> T + Send + Sync>>;
22
23#[cfg(target_has_atomic = "ptr")]
24fn new_mapper<T, F: Fn(T) -> T + Send + Sync + 'static>(func: F) -> Mapper<T> {
25 Arc::new(func)
26}
27
28#[cfg(not(target_has_atomic = "ptr"))]
29fn new_mapper<T, F: Fn(T) -> T + Send + Sync + 'static>(func: F) -> Mapper<T> {
30 Arc::new(Box::new(func))
31}
32
33#[cfg(target_has_atomic = "ptr")]
37type InitFn<P> = Arc<dyn Fn(&<P as Parameter>::Device, bool) -> P + Send + Sync>;
38
39#[cfg(not(target_has_atomic = "ptr"))]
40type InitFn<P> = Arc<Box<dyn Fn(&<P as Parameter>::Device, bool) -> P + Send + Sync>>;
41
42#[cfg(target_has_atomic = "ptr")]
43fn new_init_fn<P: Parameter, F: Fn(&P::Device, bool) -> P + Send + Sync + 'static>(
44 func: F,
45) -> InitFn<P> {
46 Arc::new(func)
47}
48
49#[cfg(not(target_has_atomic = "ptr"))]
50fn new_init_fn<P: Parameter, F: Fn(&P::Device, bool) -> P + Send + Sync + 'static>(
51 func: F,
52) -> InitFn<P> {
53 Arc::new(Box::new(func))
54}
55
56pub struct Param<T: Parameter> {
72 pub id: ParamId,
74 pub(crate) state: SyncOnceCell<T>,
77 pub(crate) initialization: Option<RwLock<Option<Uninitialized<T>>>>,
84 pub(crate) param_mapper: ParamMapper<T>,
85 pub(crate) require_grad: bool,
87}
88
89#[derive(Clone)]
90pub struct ParamMapper<T: Parameter> {
103 load: Option<Mapper<T>>,
104 save: Option<Mapper<T>>,
105}
106
107impl<T: Parameter> core::fmt::Debug for ParamMapper<T> {
108 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
109 f.write_fmt(format_args!(
110 "ParamMapper {{ load: {}, save: {} }}",
111 self.load.is_some(),
112 self.save.is_some()
113 ))
114 }
115}
116
117impl<T: Parameter> ParamMapper<T> {
118 pub fn on_load(&self, param: T) -> T {
120 match &self.load {
121 Some(mapper) => mapper(param),
122 None => param,
123 }
124 }
125 pub fn on_save(&self, param: T) -> T {
127 match &self.save {
128 Some(mapper) => mapper(param),
129 None => param,
130 }
131 }
132}
133
134impl<T: Parameter> Default for ParamMapper<T> {
135 fn default() -> Self {
136 Self {
137 load: None,
138 save: None,
139 }
140 }
141}
142
143impl<T: Parameter> core::fmt::Display for Param<T> {
144 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
145 f.write_str(format!("Param: {}", self.id).as_str())
146 }
147}
148
149impl<T: Parameter> core::fmt::Debug for Param<T> {
150 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
151 f.write_str(format!("Param: {} - {:?}", self.id, self.param_mapper).as_str())
152 }
153}
154
155pub trait Parameter: Clone + core::fmt::Debug + Send {
157 type Device: Clone;
159
160 fn device(&self) -> Self::Device;
162
163 fn is_require_grad(&self) -> bool;
165
166 fn set_require_grad(self, require_grad: bool) -> Self;
168}
169
170#[allow(clippy::type_complexity)]
172pub(crate) struct Uninitialized<P: Parameter> {
173 init: InitFn<P>,
178 pub(crate) device: P::Device,
181 pub(crate) is_require_grad: bool,
184 pub(crate) shape: Shape,
187}
188
189impl<P: Parameter> Clone for Uninitialized<P> {
190 fn clone(&self) -> Self {
191 Self {
192 init: self.init.clone(),
193 device: self.device.clone(),
194 is_require_grad: self.is_require_grad,
195 shape: self.shape.clone(),
196 }
197 }
198}
199
200impl<P: Parameter> Uninitialized<P> {
201 fn initialize(&self) -> P {
210 (self.init)(&self.device, self.is_require_grad)
211 }
212}
213
214impl<T: Parameter> Param<T> {
215 pub fn initialized(id: ParamId, value: T) -> Self {
217 let require_grad = value.is_require_grad();
218 Self {
219 id,
220 state: SyncOnceCell::initialized(value),
221 initialization: None,
222 param_mapper: Default::default(),
223 require_grad,
224 }
225 }
226
227 pub fn uninitialized<F>(
229 id: ParamId,
230 init: F,
231 device: T::Device,
232 is_require_grad: bool,
233 shape: Shape,
234 ) -> Self
235 where
236 F: Fn(&T::Device, bool) -> T + Send + Sync + 'static,
237 {
238 Self {
239 id,
240 state: SyncOnceCell::new(),
241 initialization: Some(RwLock::new(Some(Uninitialized {
242 init: new_init_fn(init),
243 device,
244 is_require_grad,
245 shape,
246 }))),
247 param_mapper: Default::default(),
248 require_grad: is_require_grad,
249 }
250 }
251
252 pub fn val(&self) -> T {
257 self.state
258 .get_or_init(|| {
259 let mut result = self
260 .initialization
261 .as_ref()
262 .expect("Should have an initialization when no state provided.")
263 .write()
264 .unwrap();
265 let state = result.take().expect("Should exist when not initialized");
266 state.initialize()
267 })
268 .clone()
269 }
270
271 pub fn is_initialized(&self) -> bool {
276 self.state.get().is_some()
277 }
278
279 pub fn planned_is_require_grad(&self) -> bool {
282 self.lazy_is_require_grad()
283 }
284
285 pub fn into_value(self) -> T {
287 self.consume().1
288 }
289
290 pub fn consume(self) -> (ParamId, T, ParamMapper<T>) {
292 let tensor = self.val();
293
294 core::mem::drop(self.state);
295
296 (self.id, tensor, self.param_mapper)
297 }
298
299 pub fn map<F: FnOnce(T) -> T>(self, func: F) -> Self {
301 let (id, tensor, param_mapper) = self.consume();
302 let tensor = func(tensor);
303 let require_grad = tensor.is_require_grad();
304
305 Self {
306 id,
307 state: SyncOnceCell::initialized(tensor),
308 initialization: None,
309 param_mapper,
310 require_grad,
311 }
312 }
313
314 pub fn from_mapped_value(id: ParamId, value: T, param_mapper: ParamMapper<T>) -> Self {
319 let require_grad = value.is_require_grad();
320 Self {
321 id,
322 state: SyncOnceCell::initialized(value),
323 initialization: None,
324 param_mapper,
325 require_grad,
326 }
327 }
328
329 pub fn load_mapper<F: Fn(T) -> T + Send + Sync + 'static>(mut self, func: F) -> Self {
331 self.param_mapper.load = Some(new_mapper(func));
332
333 self
334 }
335
336 pub fn save_mapper<F: Fn(T) -> T + Send + Sync + 'static>(mut self, func: F) -> Self {
338 self.param_mapper.save = Some(new_mapper(func));
339
340 self
341 }
342
343 pub fn init_mapper<F: Fn(T) -> T + Send + Sync + 'static>(self, func: F) -> Self
345 where
346 T: 'static,
347 {
348 let initialization = match &self.initialization {
349 Some(init) => init,
350 None => return self.map(func),
351 };
352
353 let mut init = initialization.write().unwrap();
354
355 match init.as_mut() {
356 Some(value) => {
357 let prev = value.init.clone();
358
359 value.init = new_init_fn(move |a, b| {
360 let tensor = prev(a, b);
361 func(tensor)
362 });
363 core::mem::drop(init);
364 self
365 }
366 None => {
367 core::mem::drop(init);
368 self.map(func)
369 }
370 }
371 }
372
373 pub fn lazy_device(&self) -> T::Device {
382 let initialization = match &self.initialization {
383 Some(init) => init,
384 None => return self.device(),
385 };
386
387 let init = initialization.read().unwrap();
388
389 match init.as_ref() {
390 Some(value) => value.device.clone(),
391 None => self.device(),
392 }
393 }
394
395 pub(crate) fn lazy_is_require_grad(&self) -> bool {
407 let initialization = match &self.initialization {
408 Some(init) => init,
409 None => return self.is_require_grad(),
410 };
411
412 let init = initialization.read().unwrap();
413
414 match init.as_ref() {
415 Some(value) => value.is_require_grad,
416 None => self.is_require_grad(),
417 }
418 }
419
420 pub fn set_require_grad(mut self, require_grad: bool) -> Self {
422 self.require_grad = require_grad;
423 let initialization = match &self.initialization {
424 Some(init) => init,
425 None => return self.map(|tensor| tensor.set_require_grad(require_grad)),
426 };
427
428 let mut init = initialization.write().unwrap();
429 let mut is_lazy = false;
430
431 if let Some(value) = init.as_mut() {
432 is_lazy = true;
433 value.is_require_grad = require_grad;
434 };
435
436 core::mem::drop(init);
437
438 if is_lazy {
439 return self;
440 }
441
442 self.map(|tensor| tensor.set_require_grad(require_grad))
443 }
444}
445
446impl<T: Parameter> Clone for Param<T> {
447 fn clone(&self) -> Self {
448 if let Some(init_lock) = &self.initialization {
454 let init_guard = init_lock.read().unwrap();
455 if let Some(uninit) = init_guard.as_ref() {
456 return Self {
457 id: self.id,
458 state: SyncOnceCell::new(),
459 initialization: Some(RwLock::new(Some(uninit.clone()))),
460 param_mapper: self.param_mapper.clone(),
461 require_grad: self.require_grad,
462 };
463 }
464 }
465
466 let mut param = Param::initialized(self.id, self.val());
468 param.param_mapper = self.param_mapper.clone();
469 param.require_grad = self.require_grad;
470 param
471 }
472}
473
474impl<T: Parameter> Deref for Param<T> {
475 type Target = T;
476
477 fn deref(&self) -> &Self::Target {
478 self.state.get_or_init(|| {
479 let mut result = self
480 .initialization
481 .as_ref()
482 .expect("Should have an initialization when no state provided.")
483 .write()
484 .unwrap();
485
486 let state = result.take().expect("Should exist when not initialized");
487 state.initialize()
488 })
489 }
490}
491
492#[cfg(test)]
493mod tests {
494 use super::*;
495 use ruda_tensor::api::{Tensor, backend::Backend};
496
497 fn _assert_sync<T: Sync>() {}
500
501 #[test]
502 fn param_is_sync() {
503 fn check<B: Backend>() {
504 _assert_sync::<Param<Tensor<B, 2>>>();
505 }
506 check::<ruda_tensor_host::Host>();
507 }
508
509 #[cfg(feature = "std")]
515 #[test]
516 fn param_concurrent_lazy_init() {
517 use alloc::vec::Vec;
518
519 type B = ruda_tensor_host::Host;
520 let device = Default::default();
521
522 let param: Param<Tensor<B, 2>> = Param::uninitialized(
523 ParamId::new(),
524 |device, _require_grad| Tensor::zeros([2, 3], device),
525 device,
526 false,
527 [2, 3].into(),
528 );
529
530 std::thread::scope(|s| {
532 let handles: Vec<_> = (0..4).map(|_| s.spawn(|| param.val())).collect();
533
534 let results: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect();
535
536 let expected = results[0].to_data();
538 for result in &results[1..] {
539 assert_eq!(result.to_data(), expected);
540 }
541 });
542 }
543}