Skip to main content

ruda_model/module/
base.rs

1use super::{Param, ParamId, Quantizer};
2use crate::{
3    record::Record,
4    tensor::backend::{AutodiffBackend, Backend},
5};
6use alloc::{string::String, vec::Vec};
7pub use ruda_model_macros::Module;
8use ruda_tensor::api::{Bool, Int, Tensor, ops::Device};
9
10/// Type alias to `Vec<B::Device>` which supports `no_std` environments, but automatically using
11/// the `alloc` crate.
12pub type Devices<B> = Vec<Device<B>>;
13
14// At the moment, our plan is to continue experimenting with the macro internally and monitor its development.
15// We may consider making it public in the future.
16macro_rules! module {
17    (map=$module:ident, ops=$item:expr) => {{
18        struct Mapper;
19        impl<B: Backend> ModuleMapper<B> for Mapper {
20            fn map_float<const D: usize>(
21                &mut self,
22                param: Param<Tensor<B, D>>,
23            ) -> Param<Tensor<B, D>> {
24                let func = $item;
25                func(param)
26            }
27
28            fn map_int<const D: usize>(
29                &mut self,
30                param: Param<Tensor<B, D, Int>>,
31            ) -> Param<Tensor<B, D, Int>> {
32                param
33            }
34
35            fn map_bool<const D: usize>(
36                &mut self,
37                param: Param<Tensor<B, D, Bool>>,
38            ) -> Param<Tensor<B, D, Bool>> {
39                param
40            }
41        }
42        let mut mapper = Mapper;
43        $module.map(&mut mapper)
44    }};
45    (visit_float=$module:ident, ops=$item:expr, state=$state_ty:ty, init=$init:expr) => {{
46        struct Visitor<'a, B: Backend> {
47            state: &'a mut $state_ty,
48            backend: core::marker::PhantomData<B>,
49        }
50        impl<'a, B: Backend> ModuleVisitor<B> for Visitor<'a, B> {
51            fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
52                let func = $item;
53                func(&param.val(), &mut self.state)
54            }
55        }
56        #[allow(clippy::redundant_closure_call)]
57        let mut state = $init();
58        let mut visitor = Visitor {
59            state: &mut state,
60            backend: core::marker::PhantomData,
61        };
62        $module.visit(&mut visitor);
63        state
64    }};
65}
66
67/// Trait for all neural network modules.
68///
69/// Modules should be created using the [derive](ruda_model_macros::Module) attribute.
70/// This will make your module trainable, savable and loadable via
71/// `state` and `load`.
72///
73/// # Example
74///
75/// A module should have a [backend](crate::tensor::backend::Backend) defined as a generic
76/// parameter B. This will be used by the [derive](ruda_model_macros::Module) attribute to generate the code
77/// necessary to optimize and train the module on any backend.
78///
79/// ```rust, ignore
80/// use ruda_model::module::Module;
81/// use ruda_nn::Linear;
82/// use ruda_tensor::api::{Tensor, backend::Backend};
83///
84/// #[derive(Module, Debug)]
85/// struct MyModule<B: Backend> {
86///   my_param: Linear<B>,
87///   my_other_field: usize,
88/// }
89/// ```
90pub trait Module<B: Backend>: Clone + Send + core::fmt::Debug {
91    /// Type to save and load the module.
92    type Record: Record<B>;
93
94    /// Return all the devices found in the underneath module tree added to the given vector
95    /// without duplicates.
96    fn collect_devices(&self, devices: Devices<B>) -> Devices<B>;
97
98    /// Return all the devices found in the underneath module tree without duplicates.
99    fn devices(&self) -> Devices<B> {
100        self.collect_devices(Devices::<B>::new())
101    }
102
103    /// Fork the module and all of its sub-modules to the given device.
104    ///
105    /// # Notes
106    ///
107    /// This is similar to [to_device](Module::to_device), but it ensures the output module on the
108    /// new device will have its own autodiff graph.
109    fn fork(self, device: &B::Device) -> Self;
110
111    /// Move the module and all of its sub-modules to the given device.
112    ///
113    /// # Warnings
114    ///
115    /// The operation supports autodiff and it will be registered when activated. However, this may
116    /// not be what you want. The output model will be an intermediary model, meaning that you
117    /// can't optimize it with gradient descent. If you want to optimize the output network on the
118    /// target device, use [fork](Module::fork) instead.
119    fn to_device(self, device: &B::Device) -> Self;
120
121    /// Convert floating parameter storage while preserving IDs, shared parameters
122    /// and frozen/trainable settings. Converted trainable parameters are new leaves;
123    /// gradients through the conversion itself are not retained.
124    fn to_dtype(self, dtype: ruda_tensor::FloatDType) -> Self {
125        self.map(&mut super::precision::DtypeMapper::new(dtype))
126    }
127
128    /// Convert only explicitly selected floating parameter IDs, retaining original
129    /// unselected tensors and integer/bool storage. Shared selected roles reuse the
130    /// same converted node per ID/trainability, preserving Param mappers and flags.
131    /// IDs not present in the module are not mapped; an empty selection changes nothing.
132    /// Converted trainable values are new leaves, as with [`Module::to_dtype`].
133    fn to_dtype_selected(self, dtype: ruda_tensor::FloatDType, parameter_ids: &[ParamId]) -> Self {
134        self.map(&mut super::precision::DtypeMapper::new_selected(dtype, parameter_ids))
135    }
136
137    /// Each tensor in the module tree will not require grad.
138    ///
139    /// # Warnings
140    ///
141    /// This should not be used for inference, use [valid](AutodiffModule::valid) when using
142    /// AD modules. This is mostly useful when performing partial finetuning, which is updating only
143    /// a small fraction of the parameters instead of finetuning all of them.
144    fn no_grad(self) -> Self {
145        module!(
146            map = self,
147            ops = |param: Param<Tensor<B, D>>| param.set_require_grad(false)
148        )
149    }
150
151    /// Move the module and all of its sub-modules to the autodiff backend.
152    ///
153    /// # Notes
154    ///
155    /// * Only plain modules (not already on an autodiff backend) can be moved.
156    /// * Calling `train()` on a module that is already on an autodiff backend
157    ///   will result in a type error, because the module's inner backend does not match.
158    fn train<AB>(self) -> <Self as HasAutodiffModule<AB>>::TrainModule
159    where
160        AB: AutodiffBackend<InnerBackend = B>,
161        Self: HasAutodiffModule<AB>,
162    {
163        <Self as HasAutodiffModule<AB>>::TrainModule::from_inner(self)
164    }
165
166    /// Get the number of parameters the module has, including all of its sub-modules.
167    fn num_params(&self) -> usize {
168        module!(
169            visit_float = self,
170            ops = |tensor: &Tensor<B, D>, state: &mut usize| {
171                *state += tensor.shape().num_elements();
172            },
173            state = usize,
174            init = || 0
175        )
176    }
177    /// Visit each tensor parameter in the module with a [visitor](ModuleVisitor).
178    fn visit<Visitor: ModuleVisitor<B>>(&self, visitor: &mut Visitor);
179
180    /// Map each tensor parameter in the module with a [mapper](ModuleMapper).
181    fn map<Mapper: ModuleMapper<B>>(self, mapper: &mut Mapper) -> Self;
182
183    /// Load the module state from a record.
184    fn load_record(self, record: Self::Record) -> Self;
185
186    /// Convert the module into a record containing the state.
187    fn into_record(self) -> Self::Record;
188
189    #[cfg(feature = "std")]
190    /// Save the module to a file using the provided [file recorder](crate::record::FileRecorder).
191    ///
192    /// List of supported file recorders:
193    ///
194    /// * [default](crate::record::DefaultFileRecorder)
195    /// * [bincode](crate::record::BinFileRecorder)
196    /// * [bincode compressed with gzip](crate::record::BinGzFileRecorder)
197    /// * [json pretty](crate::record::PrettyJsonFileRecorder)
198    /// * [json compressed with gzip](crate::record::JsonGzFileRecorder)
199    /// * [named mpk](crate::record::NamedMpkFileRecorder)
200    /// * [named mpk compressed with gzip](crate::record::NamedMpkGzFileRecorder)
201    ///
202    /// ## Notes
203    ///
204    /// The file extension is automatically added depending on the file recorder provided, you
205    /// don't have to specify it.
206    fn save_file<FR, PB>(
207        self,
208        file_path: PB,
209        recorder: &FR,
210    ) -> Result<(), crate::record::RecorderError>
211    where
212        FR: crate::record::FileRecorder<B>,
213        PB: Into<std::path::PathBuf>,
214    {
215        let record = Self::into_record(self);
216        recorder.record(record, file_path.into())
217    }
218
219    #[cfg(feature = "std")]
220    /// Load the module from a file using the provided [file recorder](crate::record::FileRecorder).
221    ///
222    /// The recorder should be the same as the one used to save the module, see
223    /// [save_file](Self::save_file).
224    ///
225    /// ## Notes
226    ///
227    /// The file extension is automatically added depending on the file recorder provided, you
228    /// don't have to specify it.
229    fn load_file<FR, PB>(
230        self,
231        file_path: PB,
232        recorder: &FR,
233        device: &B::Device,
234    ) -> Result<Self, crate::record::RecorderError>
235    where
236        FR: crate::record::FileRecorder<B>,
237        PB: Into<std::path::PathBuf>,
238    {
239        let record = recorder.load(file_path.into(), device)?;
240
241        Ok(self.load_record(record))
242    }
243
244    /// Quantize the weights of the module.
245    fn quantize_weights(self, quantizer: &mut Quantizer) -> Self {
246        self.map(quantizer)
247    }
248}
249
250/// Module visitor trait for traversing and inspecting module parameters.
251pub trait ModuleVisitor<B: Backend> {
252    /// Visit a float parameter in the module.
253    ///
254    /// # Parameters
255    /// - `param`: The float parameter to visit
256    #[allow(unused_variables)]
257    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {}
258
259    /// Visit an int parameter in the module.
260    ///
261    /// # Parameters
262    /// - `param`: The integer parameter to visit
263    #[allow(unused_variables)]
264    fn visit_int<const D: usize>(&mut self, param: &Param<Tensor<B, D, Int>>) {}
265
266    /// Visit a bool parameter in the module.
267    ///
268    /// # Parameters
269    /// - `param`: The boolean parameter to visit
270    #[allow(unused_variables)]
271    fn visit_bool<const D: usize>(&mut self, param: &Param<Tensor<B, D, Bool>>) {}
272
273    /// Called when entering a submodule.
274    ///
275    /// # Parameters
276    /// - `name`: The name of the submodule being entered
277    /// - `container_type`: The type of the container with format:
278    ///   - For user-defined structs: "Struct:TypeName" (e.g., "Struct:Linear")
279    ///   - For user-defined enums: "Enum:TypeName" (e.g., "Enum:MyEnum")
280    ///   - For Vec containers: "Vec" (name is the index)
281    ///   - For Tuple containers: "Tuple" (name is the index)
282    ///   - For Array containers: "Array" (name is the index)
283    ///
284    /// Note: Option containers do not call enter_module/exit_module to preserve
285    /// the field name in the path (e.g., "bias" instead of "bias.Some")
286    #[allow(unused_variables)]
287    fn enter_module(&mut self, name: &str, container_type: &str) {}
288
289    /// Called when exiting a submodule.
290    ///
291    /// # Parameters
292    /// - `name`: The name of the submodule being exited
293    /// - `container_type`: The type of the container with format:
294    ///   - For user-defined structs: "Struct:TypeName" (e.g., "Struct:Linear")
295    ///   - For user-defined enums: "Enum:TypeName" (e.g., "Enum:MyEnum")
296    ///   - For Vec containers: "Vec" (name is the index)
297    ///   - For Tuple containers: "Tuple" (name is the index)
298    ///   - For Array containers: "Array" (name is the index)
299    ///
300    /// Note: Option containers do not call enter_module/exit_module to preserve
301    /// the field name in the path (e.g., "bias" instead of "bias.Some")
302    #[allow(unused_variables)]
303    fn exit_module(&mut self, name: &str, container_type: &str) {}
304
305    /// Visit a float tensor with its full module path.
306    ///
307    /// # Parameters
308    /// - `path`: The path components to the tensor as a slice (e.g., &["encoder", "layer1", "weight"]).
309    ///   Each element represents a module name in the hierarchy, with the final element
310    ///   being the parameter name. This allows efficient reuse of the path stack.
311    /// - `id`: The unique identifier of the parameter
312    /// - `tensor`: The float tensor to visit
313    #[allow(unused_variables)]
314    fn visit_float_with_path<const D: usize>(
315        &mut self,
316        path: &[String],
317        id: ParamId,
318        tensor: &Tensor<B, D>,
319    ) {
320    }
321
322    /// Visit an int tensor with its full module path.
323    ///
324    /// # Parameters
325    /// - `path`: The path components to the tensor as a slice (e.g., &["encoder", "layer1", "weight"]).
326    ///   Each element represents a module name in the hierarchy, with the final element
327    ///   being the parameter name. This allows efficient reuse of the path stack.
328    /// - `id`: The unique identifier of the parameter
329    /// - `tensor`: The integer tensor to visit
330    #[allow(unused_variables)]
331    fn visit_int_with_path<const D: usize>(
332        &mut self,
333        path: &[String],
334        id: ParamId,
335        tensor: &Tensor<B, D, Int>,
336    ) {
337    }
338
339    /// Visit a bool tensor with its full module path.
340    ///
341    /// # Parameters
342    /// - `path`: The path components to the tensor as a slice (e.g., &["encoder", "layer1", "weight"]).
343    ///   Each element represents a module name in the hierarchy, with the final element
344    ///   being the parameter name. This allows efficient reuse of the path stack.
345    /// - `id`: The unique identifier of the parameter
346    /// - `tensor`: The boolean tensor to visit
347    #[allow(unused_variables)]
348    fn visit_bool_with_path<const D: usize>(
349        &mut self,
350        path: &[String],
351        id: ParamId,
352        tensor: &Tensor<B, D, Bool>,
353    ) {
354    }
355}
356
357/// Module mapper trait for transforming module parameters.
358pub trait ModuleMapper<B: Backend> {
359    /// Called when entering a submodule.
360    ///
361    /// # Parameters
362    /// - `name`: The name of the submodule being entered
363    /// - `container_type`: The type of the container with format:
364    ///   - For user-defined structs: "Struct:TypeName" (e.g., "Struct:Linear")
365    ///   - For user-defined enums: "Enum:TypeName" (e.g., "Enum:MyEnum")
366    ///   - For Vec containers: "Vec" (name is the index)
367    ///   - For Tuple containers: "Tuple" (name is the index)
368    ///   - For Array containers: "Array" (name is the index)
369    ///
370    /// Note: Option containers do not call enter_module/exit_module to preserve
371    /// the field name in the path (e.g., "bias" instead of "bias.Some")
372    #[allow(unused_variables)]
373    fn enter_module(&mut self, name: &str, container_type: &str) {}
374
375    /// Called when exiting a submodule.
376    ///
377    /// # Parameters
378    /// - `name`: The name of the submodule being exited
379    /// - `container_type`: The type of the container with format:
380    ///   - For user-defined structs: "Struct:TypeName" (e.g., "Struct:Linear")
381    ///   - For user-defined enums: "Enum:TypeName" (e.g., "Enum:MyEnum")
382    ///   - For Vec containers: "Vec" (name is the index)
383    ///   - For Tuple containers: "Tuple" (name is the index)
384    ///   - For Array containers: "Array" (name is the index)
385    ///
386    /// Note: Option containers do not call enter_module/exit_module to preserve
387    /// the field name in the path (e.g., "bias" instead of "bias.Some")
388    #[allow(unused_variables)]
389    fn exit_module(&mut self, name: &str, container_type: &str) {}
390
391    /// Map a float parameter in the module.
392    ///
393    /// # Parameters
394    /// - `param`: The float parameter to transform
395    ///
396    /// # Returns
397    /// The transformed parameter
398    #[allow(unused_variables)]
399    fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
400        let (id, tensor, mapper) = param.consume();
401        Param::from_mapped_value(id, tensor, mapper)
402    }
403
404    /// Map an int parameter in the module.
405    ///
406    /// # Parameters
407    /// - `param`: The integer parameter to transform
408    ///
409    /// # Returns
410    /// The transformed parameter
411    #[allow(unused_variables)]
412    fn map_int<const D: usize>(
413        &mut self,
414        param: Param<Tensor<B, D, Int>>,
415    ) -> Param<Tensor<B, D, Int>> {
416        let (id, tensor, mapper) = param.consume();
417        Param::from_mapped_value(id, tensor, mapper)
418    }
419
420    /// Map a bool parameter in the module.
421    ///
422    /// # Parameters
423    /// - `param`: The boolean parameter to transform
424    ///
425    /// # Returns
426    /// The transformed parameter
427    #[allow(unused_variables)]
428    fn map_bool<const D: usize>(
429        &mut self,
430        param: Param<Tensor<B, D, Bool>>,
431    ) -> Param<Tensor<B, D, Bool>> {
432        let (id, tensor, mapper) = param.consume();
433        Param::from_mapped_value(id, tensor, mapper)
434    }
435}
436
437/// Module with auto-differentiation backend.
438pub trait AutodiffModule<B: AutodiffBackend>: Module<B> + Send + core::fmt::Debug {
439    /// Inner module without auto-differentiation.
440    type InnerModule: Module<B::InnerBackend>;
441
442    /// Returns the same module, but on the inner backend without auto-differentiation.
443    fn valid(&self) -> Self::InnerModule;
444
445    /// Wraps an inner module back into an auto-diff module.
446    fn from_inner(module: Self::InnerModule) -> Self;
447}
448
449/// Helper trait to associate a module with its autodiff version.
450pub trait HasAutodiffModule<B: AutodiffBackend> {
451    /// The module with auto-differentiation.
452    type TrainModule: AutodiffModule<B, InnerModule = Self>;
453}
454
455#[cfg(test)]
456mod tests {
457    use super::*;
458
459    use crate::TestAutodiffBackend;
460    use crate::test_utils::SimpleLinear;
461
462    #[test]
463    fn test_module_val_train_stateful() {
464        let device = Default::default();
465        let module = SimpleLinear::<TestAutodiffBackend>::new(4, 4, &device);
466
467        assert!(module.weight.is_require_grad());
468        assert!(module.weight.require_grad);
469
470        let module = module.valid();
471        assert!(!module.weight.is_require_grad());
472        assert!(module.weight.require_grad); // stateful
473
474        // Without `HasAutodiffModule`, we would need to specify the module type as well, which would be annoying
475        // let module: SimpleLinear<TestAutodiffBackend> = module.train();
476        let module = module.train::<TestAutodiffBackend>();
477        assert!(module.weight.is_require_grad());
478        assert!(module.weight.require_grad); // stateful
479
480        let module = module.no_grad();
481        assert!(!module.weight.is_require_grad());
482        assert!(!module.weight.require_grad); // stateful
483
484        let module = module.valid();
485        assert!(!module.weight.is_require_grad()); // always
486        assert!(!module.weight.require_grad); // stateful
487
488        let module = module.train::<TestAutodiffBackend>();
489        assert!(!module.weight.is_require_grad());
490        assert!(!module.weight.require_grad); // stateful
491    }
492}