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(¶m.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}