Skip to main content

ruda_model/module/param/
base.rs

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/// Type alias for the init function stored in `Uninitialized`.
34/// On targets without atomics, `portable_atomic_util::Arc` needs `Box` indirection
35/// for unsized types, mirroring the `Mapper` pattern above.
36#[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
56/// Parameters are the fundamental building blocks of [modules](crate::module::Module) where they
57/// serve as containers for [tensors](crate::tensor::Tensor) that can be updated during
58/// training, and loaded during inference. If you don't want to save the tensors
59/// and/or don't want to update it during training, you don't need this type to wrap your tensor.
60///
61/// # Core Lazy Initialization Architecture
62///
63/// `Param<T>` has a dual-state design using `SyncOnceCell<T>`:
64///
65/// ## State Management
66///
67/// **Two possible states:**
68///
69/// 1. **Initialized**: `state: SyncOnceCell<T>` contains value, `initialization: None`
70/// 2. **Uninitialized (Lazy)**: `state` is empty, `initialization: Some(RwLock<Option<Uninitialized<T>>>)`
71pub struct Param<T: Parameter> {
72    /// The unique ID of this parameter. This is used by eg. optimizers to associate a gradient with a specific parameter.
73    pub id: ParamId,
74    /// The SyncOnceCell holding the initialized parameter value.
75    /// Empty for uninitialized parameters, populated after first access or explicit initialization.
76    pub(crate) state: SyncOnceCell<T>,
77    /// The deferred initialization state for lazy parameters.
78    ///
79    /// **State Transitions:**
80    /// - Initialized params: `None`
81    /// - Uninitialized params: `Some(RwLock<Some(Uninitialized<T>)>)`
82    /// - After lazy init triggers: `Some(RwLock<None>)` (inner Option is taken)
83    pub(crate) initialization: Option<RwLock<Option<Uninitialized<T>>>>,
84    pub(crate) param_mapper: ParamMapper<T>,
85    // For stateful `module.valid()` <> `module.train()`
86    pub(crate) require_grad: bool,
87}
88
89#[derive(Clone)]
90/// Applies transformations when loading and saving parameters.
91///
92/// # Mapper System
93///
94/// `ParamMapper<T>` allows applying transformations during serialization and deserialization:
95/// - `load: Option<Mapper<T>>` - transformation during deserialization (applied in `transform_for_load()`)
96/// - `save: Option<Mapper<T>>` - transformation during serialization (applied in `transform_for_save()`)
97///
98/// These are commonly used for:
99/// - Quantization/dequantization
100/// - Precision conversion (e.g., FP32 ↔ FP16)
101/// - Custom parameter transformations
102pub 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    /// Applies the transformation when loading the given parameter.
119    pub fn on_load(&self, param: T) -> T {
120        match &self.load {
121            Some(mapper) => mapper(param),
122            None => param,
123        }
124    }
125    /// Applies the transformation when saving the given parameter.
126    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
155/// Trait that defines what is necessary for a type to be a parameter.
156pub trait Parameter: Clone + core::fmt::Debug + Send {
157    /// The device type to be used.
158    type Device: Clone;
159
160    /// Fetch the device.
161    fn device(&self) -> Self::Device;
162
163    /// Fetch the gradient requirement.
164    fn is_require_grad(&self) -> bool;
165
166    /// Set the gradient requirement.
167    fn set_require_grad(self, require_grad: bool) -> Self;
168}
169
170/// The deferred initialization state for lazy parameters.
171#[allow(clippy::type_complexity)]
172pub(crate) struct Uninitialized<P: Parameter> {
173    /// The initialization function. Called with `(device, is_require_grad) -> Parameter`.
174    /// Wrapped in `Arc` so that cloning a `Param` preserves the lazy state without
175    /// triggering initialization. Each clone holds its own `Uninitialized` state and
176    /// will run the init function separately on first access (producing independent values).
177    init: InitFn<P>,
178    /// The target device on which the parameter should be initialized.
179    /// Used by `lazy_device()` to provide device information without triggering initialization.
180    pub(crate) device: P::Device,
181    /// The gradient requirement for the parameter.
182    /// Used by `lazy_is_require_grad()` to provide gradient settings without triggering initialization.
183    pub(crate) is_require_grad: bool,
184    /// The shape of the tensor parameter.
185    /// Used by `lazy_shape()` to provide shape information without triggering initialization.
186    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    /// Runs the initialization function.
202    ///
203    /// This is called by [Param::val] when accessing an uninitialized parameter for the first time.
204    /// The function is given the stored device and gradient requirement, and returns the initialized parameter.
205    ///
206    /// Although this takes `&self` (the `Arc<dyn Fn>` is callable multiple times), callers
207    /// are expected to invoke this only once per `Param` instance. The caller (`val()`) takes
208    /// the `Uninitialized` out of its `Option` via `take()` to enforce single-initialization.
209    fn initialize(&self) -> P {
210        (self.init)(&self.device, self.is_require_grad)
211    }
212}
213
214impl<T: Parameter> Param<T> {
215    /// Create a new parameter that is already initialized.
216    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    /// Create a new parameter that is not already initialized.
228    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    /// Gets the parameter value, initializing it lazily if needed.
253    ///
254    /// For initialized parameters, this returns a clone of the cached value.
255    /// For uninitialized parameters, this triggers initialization:
256    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    /// Check if the parameter has been initialized.
272    ///
273    /// Returns `true` if the parameter's value has been computed and cached,
274    /// `false` if it's still lazy and will be initialized on first access.
275    pub fn is_initialized(&self) -> bool {
276        self.state.get().is_some()
277    }
278
279    /// Read the initialized or declared lazy gradient requirement without running the initializer.
280    /// A backend/device without autodiff can still produce an untracked tensor when initialized.
281    pub fn planned_is_require_grad(&self) -> bool {
282        self.lazy_is_require_grad()
283    }
284
285    /// Gets the parameter's value while consuming the parameter.
286    pub fn into_value(self) -> T {
287        self.consume().1
288    }
289
290    /// Gets the parameter id and value while consuming the parameter.
291    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    /// Execute the given function on the inner value.
300    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    /// Create an initialized parameter with the given id, value, and param mapper.
315    ///
316    /// This is a helper method for creating parameters while preserving the param mapper,
317    /// typically used in ModuleMapper implementations.
318    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    /// Runs a transformation on the parameter when loading.
330    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    /// Runs a transformation on the parameter when saving.
337    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    /// Execute the given function on the inner value.
344    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    /// The device on which the parameter is or will be initialized, **without triggering initialization**.
374    ///
375    /// This is critical for the load optimization: when loading tensors into an uninitialized parameter,
376    /// we need to know the target device to move the loaded tensor appropriately, but we don't want to
377    /// trigger the initialization function (which would allocate an unnecessary tensor).
378    ///
379    /// Use this instead of [crate::tensor::Tensor::device] when you need the device but want to
380    /// preserve lazy initialization.
381    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    /// The gradient requirement on which the parameter is or will be initialized, **without triggering initialization**.
396    ///
397    /// Similar to [lazy_device](Self::lazy_device), this is critical for the load optimization.
398    /// When loading tensors into an uninitialized parameter, we need to apply the correct gradient
399    /// setting to the loaded tensor without triggering the initialization function.
400    ///
401    /// # Notes
402    ///
403    /// This is a crate-private function, since users are not expected to use `is_require_grad` of an
404    /// uninitialized module to then override its value. All low-level functions should be provided
405    /// by `ruda` and should handle those details.
406    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    /// Override the gradient requirement for the current parameter.
421    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 uninitialized, clone the lazy state without triggering initialization.
449        // This avoids allocating tensor memory for params that may never be used
450        // (e.g., when cloning a module just to load weights into it).
451        // The clone gets its own SyncOnceCell and RwLock, so initializing one
452        // does not affect the other.
453        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        // Already initialized (or init was already consumed): clone the value.
467        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    // Param<T> should be Sync so that models can be shared across threads
498    // (e.g. parallel inference with rayon).
499    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    /// Concurrent lazy initialization must not panic.
510    ///
511    /// Multiple threads call `val()` on an uninitialized `Param` simultaneously.
512    /// `SyncOnceCell::get_or_init` guarantees only one thread runs the initializer;
513    /// the others block and receive the same value.
514    #[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        // Share across threads via &param (requires Sync).
531        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            // All threads must get the same value.
537            let expected = results[0].to_data();
538            for result in &results[1..] {
539                assert_eq!(result.to_data(), expected);
540            }
541        });
542    }
543}