Skip to main content

ruda_model/module/param/
tensor.rs

1use super::{Param, ParamId, Parameter};
2use crate::module::{
3    AutodiffModule, Content, HasAutodiffModule, Module, ModuleDisplay, ModuleDisplayDefault,
4    ModuleMapper, ModuleVisitor,
5};
6use crate::tensor::{
7    Tensor,
8    backend::{AutodiffBackend, Backend},
9};
10use alloc::{format, string::ToString, vec::Vec};
11use ruda_tensor::api::{Bool, Float, Int, TensorData, ops::Device};
12
13impl<B: Backend, const D: usize> Parameter for Tensor<B, D, Float> {
14    type Device = B::Device;
15
16    fn device(&self) -> Self::Device {
17        Tensor::device(self)
18    }
19
20    fn is_require_grad(&self) -> bool {
21        Tensor::is_require_grad(self)
22    }
23
24    fn set_require_grad(self, require_grad: bool) -> Self {
25        Tensor::set_require_grad(self, require_grad)
26    }
27}
28
29impl<B: Backend, const D: usize> Parameter for Tensor<B, D, Int> {
30    type Device = B::Device;
31
32    fn device(&self) -> Self::Device {
33        Tensor::device(self)
34    }
35
36    fn is_require_grad(&self) -> bool {
37        false
38    }
39
40    fn set_require_grad(self, _require_grad: bool) -> Self {
41        self
42    }
43}
44
45impl<B: Backend, const D: usize> Parameter for Tensor<B, D, Bool> {
46    type Device = B::Device;
47
48    fn device(&self) -> Self::Device {
49        Tensor::device(self)
50    }
51
52    fn is_require_grad(&self) -> bool {
53        false
54    }
55
56    fn set_require_grad(self, _require_grad: bool) -> Self {
57        self
58    }
59}
60
61impl<B: Backend, const D: usize> Param<Tensor<B, D>> {
62    /// Create a new parameter from a float tensor.
63    ///
64    /// # Warnings
65    ///
66    /// We strongly recommend using [Param::uninitialized] if you are using this method to
67    /// initialize parameters inside a module, since the tensor initialization will be lazy,
68    /// making the loading of weights more performant.
69    pub fn from_tensor(value: Tensor<B, D>) -> Self {
70        // When creating a parameter from a float tensor, we automatically mark it as requiring
71        // gradients, so that it can be updated by an optimizer.
72        Param::initialized(ParamId::new(), value.require_grad())
73    }
74
75    /// The shape of the parameter, **without triggering initialization**.
76    ///
77    /// This is critical for shape validation during loading: when applying tensors to an
78    /// uninitialized parameter, we need to validate the shape without triggering the
79    /// initialization function (which would allocate an unnecessary tensor).
80    ///
81    /// Use this instead of [crate::tensor::Tensor::shape] when you need the shape but want to
82    /// preserve lazy initialization.
83    pub fn lazy_shape(&self) -> ruda_tensor::api::Shape {
84        let initialization = match &self.initialization {
85            Some(init) => init,
86            None => return self.shape(),
87        };
88
89        let init = initialization.read().unwrap();
90
91        match init.as_ref() {
92            Some(value) => value.shape.clone(),
93            None => self.shape(),
94        }
95    }
96
97    /// Create a new parameter from data.
98    pub fn from_data<T>(data: T, device: &B::Device) -> Self
99    where
100        T: Into<TensorData>,
101    {
102        let data: TensorData = data.into();
103        // When creating a parameter from a float tensor, we automatically mark it as requiring
104        // gradients, so that it can be updated by an optimizer.
105        B::memory_persistent_allocations(device, data, |data| {
106            let value = Tensor::from_data(data, device);
107            Param::initialized(ParamId::new(), value.require_grad())
108        })
109    }
110
111    /// Transform a parameter for loading by applying load transformations.
112    ///
113    /// This method is used to restore a parameter from a tensor (typically during deserialization).
114    /// It ensures the tensor is moved to the expected device, applies the param mapper's
115    /// `on_load` transformation, and preserves the autodiff settings (require_grad).
116    pub fn transform_for_load(self, tensor: Tensor<B, D>, param_id: ParamId) -> Self {
117        let mut new_tensor = tensor;
118
119        let mapper = self.param_mapper.clone();
120
121        let expected_device = self.lazy_device();
122        let expected_require_grad = self.lazy_is_require_grad();
123
124        // Make sure we load the tensor into the same module device.
125        if new_tensor.device() != expected_device {
126            new_tensor = new_tensor.to_device(&expected_device).detach();
127        }
128
129        new_tensor = mapper.on_load(new_tensor);
130
131        // Make sure we load the tensor with the same autodiff setting.
132        new_tensor = new_tensor.set_require_grad(expected_require_grad);
133
134        let mut loaded = Self::initialized(param_id, new_tensor);
135        loaded.param_mapper = mapper;
136        loaded.require_grad = self.require_grad;
137        loaded
138    }
139
140    /// Transform a parameter for saving by applying save transformations.
141    ///
142    /// This method is used to prepare a parameter for saving (typically during serialization).
143    /// It applies the param mapper's `on_save` transformation, which can be used
144    /// to modify the tensor before serialization (e.g., quantization, precision conversion).
145    pub fn transform_for_save(&self) -> Self {
146        let mut tensor = self.val();
147        let mapper = self.param_mapper.clone();
148
149        tensor = mapper.on_save(tensor);
150
151        Self::initialized(self.id, tensor)
152    }
153}
154
155impl<B: Backend, const D: usize> Param<Tensor<B, D, Int>> {
156    /// The shape of the parameter, **without triggering initialization**.
157    ///
158    /// This is critical for shape validation during loading: when applying tensors to an
159    /// uninitialized parameter, we need to validate the shape without triggering the
160    /// initialization function (which would allocate an unnecessary tensor).
161    ///
162    /// Use this instead of [crate::tensor::Tensor::shape] when you need the shape but want to
163    /// preserve lazy initialization.
164    pub fn lazy_shape(&self) -> ruda_tensor::api::Shape {
165        let initialization = match &self.initialization {
166            Some(init) => init,
167            None => return self.shape(),
168        };
169
170        let init = initialization.read().unwrap();
171
172        match init.as_ref() {
173            Some(value) => value.shape.clone(),
174            None => self.shape(),
175        }
176    }
177
178    /// Transform a parameter for loading by applying load transformations.
179    ///
180    /// This method is used to restore a parameter from a tensor (typically during deserialization).
181    /// It ensures the tensor is moved to the expected device and applies the param mapper's
182    /// `on_load` transformation.
183    pub fn transform_for_load(self, tensor: Tensor<B, D, Int>, param_id: ParamId) -> Self {
184        let mut new_tensor = tensor;
185
186        let mapper = self.param_mapper.clone();
187
188        let expected_device = self.lazy_device();
189
190        // Make sure we load the tensor into the same module device.
191        if new_tensor.device() != expected_device {
192            new_tensor = new_tensor.to_device(&expected_device);
193        }
194
195        new_tensor = mapper.on_load(new_tensor);
196
197        let mut loaded = Self::initialized(param_id, new_tensor);
198        loaded.param_mapper = mapper;
199        loaded
200    }
201
202    /// Transform a parameter for saving by applying save transformations.
203    ///
204    /// This method is used to prepare a parameter for saving (typically during serialization).
205    /// It applies the param mapper's `on_save` transformation, which can be used
206    /// to modify the tensor before serialization (e.g., quantization, precision conversion).
207    pub fn transform_for_save(&self) -> Self {
208        let mut tensor = self.val();
209        let mapper = self.param_mapper.clone();
210
211        tensor = mapper.on_save(tensor);
212
213        Self::initialized(self.id, tensor)
214    }
215}
216
217impl<B: Backend, const D: usize> Param<Tensor<B, D, Bool>> {
218    /// The shape of the parameter, **without triggering initialization**.
219    ///
220    /// This is critical for shape validation during loading: when applying tensors to an
221    /// uninitialized parameter, we need to validate the shape without triggering the
222    /// initialization function (which would allocate an unnecessary tensor).
223    ///
224    /// **Returns:**
225    /// - For uninitialized params: the shape from the `Uninitialized` struct
226    /// - For initialized params: the actual shape from the tensor
227    ///
228    /// Use this instead of [crate::tensor::Tensor::shape] when you need the shape but want to
229    /// preserve lazy initialization.
230    pub fn lazy_shape(&self) -> ruda_tensor::api::Shape {
231        let initialization = match &self.initialization {
232            Some(init) => init,
233            None => return self.shape(),
234        };
235
236        let init = initialization.read().unwrap();
237
238        match init.as_ref() {
239            Some(value) => value.shape.clone(),
240            None => self.shape(),
241        }
242    }
243
244    /// Transform a parameter for loading by applying load transformations.
245    ///
246    /// This method is used to restore a parameter from a tensor (typically during deserialization).
247    /// It ensures the tensor is moved to the expected device and applies the param mapper's
248    /// `on_load` transformation.
249    pub fn transform_for_load(self, tensor: Tensor<B, D, Bool>, param_id: ParamId) -> Self {
250        let mut new_tensor = tensor;
251
252        let mapper = self.param_mapper.clone();
253
254        let expected_device = self.lazy_device();
255
256        // Make sure we load the tensor into the same module device.
257        if new_tensor.device() != expected_device {
258            new_tensor = new_tensor.to_device(&expected_device);
259        }
260
261        new_tensor = mapper.on_load(new_tensor);
262
263        let mut loaded = Self::initialized(param_id, new_tensor);
264        loaded.param_mapper = mapper;
265        loaded
266    }
267
268    /// Transform a parameter for saving by applying save transformations.
269    ///
270    /// This method is used to prepare a parameter for saving (typically during serialization).
271    /// It applies the param mapper's `on_save` transformation, which can be used
272    /// to modify the tensor before serialization (e.g., quantization, precision conversion).
273    pub fn transform_for_save(&self) -> Self {
274        let mut tensor = self.val();
275        let mapper = self.param_mapper.clone();
276
277        tensor = mapper.on_save(tensor);
278
279        Self::initialized(self.id, tensor)
280    }
281}
282
283impl<const D: usize, B: Backend> Module<B> for Param<Tensor<B, D>> {
284    type Record = Param<Tensor<B, D>>;
285
286    fn visit<V: ModuleVisitor<B>>(&self, visitor: &mut V) {
287        visitor.visit_float(self)
288    }
289
290    fn map<M: ModuleMapper<B>>(self, mapper: &mut M) -> Self {
291        mapper.map_float(self)
292    }
293
294    fn into_record(self) -> Self::Record {
295        self.transform_for_save()
296    }
297
298    fn load_record(self, record: Self::Record) -> Self {
299        let (record_param_id, record_tensor, _) = record.consume();
300        self.transform_for_load(record_tensor, record_param_id)
301    }
302
303    fn to_device(self, device: &Device<B>) -> Self {
304        let require_grad = self.require_grad;
305        let mut param = self.map(|tensor| tensor.to_device(device));
306        param.require_grad = require_grad;
307        param
308    }
309
310    fn fork(self, device: &Device<B>) -> Self {
311        let require_grad = self.require_grad;
312        let mut param = self.map(|tensor| {
313            let is_require_grad = tensor.is_require_grad();
314            let mut tensor = tensor.to_device(device).detach();
315
316            if is_require_grad {
317                tensor = tensor.require_grad();
318            }
319
320            tensor
321        });
322        param.require_grad = require_grad;
323        param
324    }
325
326    fn collect_devices(&self, mut devices: Vec<Device<B>>) -> Vec<Device<B>> {
327        let device = self.val().device();
328
329        if !devices.contains(&device) {
330            devices.push(device)
331        }
332
333        devices
334    }
335}
336
337impl<const D: usize, B: Backend> ModuleDisplayDefault for Param<Tensor<B, D>> {
338    fn content(&self, content: Content) -> Option<Content> {
339        let id = if content.display_settings.show_param_id() {
340            format!(", id: {}", self.id)
341        } else {
342            "".to_string()
343        };
344        let string = format!(
345            "ParamTensor {{rank: {D}, shape: {:?}, kind: float{id}}}",
346            self.shape().as_slice()
347        );
348        content.add_formatted(&string).optional()
349    }
350}
351impl<const D: usize, B: Backend> ModuleDisplay for Param<Tensor<B, D>> {}
352
353impl<const D: usize, B: Backend> Module<B> for Param<Tensor<B, D, Int>> {
354    type Record = Param<Tensor<B, D, Int>>;
355
356    fn visit<V: ModuleVisitor<B>>(&self, visitor: &mut V) {
357        visitor.visit_int(self)
358    }
359
360    fn map<M: ModuleMapper<B>>(self, mapper: &mut M) -> Self {
361        mapper.map_int(self)
362    }
363
364    fn into_record(self) -> Self::Record {
365        self.transform_for_save()
366    }
367
368    fn load_record(self, record: Self::Record) -> Self {
369        let (record_param_id, record_tensor, _) = record.consume();
370        self.transform_for_load(record_tensor, record_param_id)
371    }
372
373    fn to_device(self, device: &Device<B>) -> Self {
374        self.map(|tensor| tensor.to_device(device))
375    }
376
377    fn fork(self, device: &Device<B>) -> Self {
378        self.to_device(device) // Don't support autodiff.
379    }
380
381    fn collect_devices(&self, mut devices: Vec<Device<B>>) -> Vec<Device<B>> {
382        let device = self.val().device();
383
384        if !devices.contains(&device) {
385            devices.push(device)
386        }
387
388        devices
389    }
390}
391
392impl<const D: usize, B: Backend> ModuleDisplayDefault for Param<Tensor<B, D, Int>> {
393    fn content(&self, content: Content) -> Option<Content> {
394        let id = if content.display_settings.show_param_id() {
395            format!(", id: {}", self.id)
396        } else {
397            "".to_string()
398        };
399        let string = format!(
400            "ParamTensor {{rank: {D}, shape: {:?}, kind: int{id}}}",
401            self.shape().as_slice()
402        );
403        content.add_formatted(&string).optional()
404    }
405}
406impl<const D: usize, B: Backend> ModuleDisplay for Param<Tensor<B, D, Int>> {}
407
408impl<const D: usize, B: Backend> Module<B> for Param<Tensor<B, D, Bool>> {
409    type Record = Param<Tensor<B, D, Bool>>;
410
411    fn visit<V: ModuleVisitor<B>>(&self, visitor: &mut V) {
412        visitor.visit_bool(self)
413    }
414
415    fn map<M: ModuleMapper<B>>(self, mapper: &mut M) -> Self {
416        mapper.map_bool(self)
417    }
418
419    fn into_record(self) -> Self::Record {
420        self.transform_for_save()
421    }
422
423    fn load_record(self, record: Self::Record) -> Self {
424        let (record_param_id, record_tensor, _) = record.consume();
425        self.transform_for_load(record_tensor, record_param_id)
426    }
427
428    fn to_device(self, device: &Device<B>) -> Self {
429        self.map(|tensor| tensor.to_device(device))
430    }
431
432    fn fork(self, device: &Device<B>) -> Self {
433        self.to_device(device) // Don't support autodiff.
434    }
435
436    fn collect_devices(&self, mut devices: Vec<Device<B>>) -> Vec<Device<B>> {
437        let device = self.val().device();
438
439        if !devices.contains(&device) {
440            devices.push(device)
441        }
442
443        devices
444    }
445}
446
447impl<const D: usize, B: Backend> ModuleDisplayDefault for Param<Tensor<B, D, Bool>> {
448    fn content(&self, content: Content) -> Option<Content> {
449        let id = if content.display_settings.show_param_id() {
450            format!(", id: {}", self.id)
451        } else {
452            "".to_string()
453        };
454
455        let string = format!(
456            "ParamTensor {{rank: {D}, shape: {:?}, kind: bool{id}}}",
457            self.shape().as_slice()
458        );
459        content.add_formatted(&string).optional()
460    }
461}
462
463impl<const D: usize, B: Backend> ModuleDisplay for Param<Tensor<B, D, Bool>> {}
464
465impl<const D: usize, B: AutodiffBackend> AutodiffModule<B> for Param<Tensor<B, D>> {
466    type InnerModule = Param<Tensor<B::InnerBackend, D>>;
467
468    fn valid(&self) -> Self::InnerModule {
469        // Preserve initialized param `require_grad` state, but reset the inner value's
470        let require_grad = self.require_grad;
471        let mut param = Param::initialized(self.id, self.val().inner().set_require_grad(false));
472        param.require_grad = require_grad;
473        param
474    }
475
476    fn from_inner(module: Self::InnerModule) -> Self {
477        // Reinstate the param's `require_grad` state
478        let tensor = Tensor::from_inner(module.val()).set_require_grad(module.require_grad);
479        Param::initialized(module.id, tensor)
480    }
481}
482
483impl<const D: usize, B: AutodiffBackend> HasAutodiffModule<B>
484    for Param<Tensor<B::InnerBackend, D>>
485{
486    type TrainModule = Param<Tensor<B, D>>;
487}
488
489impl<const D: usize, B: AutodiffBackend> AutodiffModule<B> for Param<Tensor<B, D, Int>> {
490    type InnerModule = Param<Tensor<B::InnerBackend, D, Int>>;
491
492    fn valid(&self) -> Self::InnerModule {
493        Param::initialized(self.id, self.val().inner())
494    }
495
496    fn from_inner(module: Self::InnerModule) -> Self {
497        Param::initialized(module.id, Tensor::from_inner(module.val()))
498    }
499}
500
501impl<const D: usize, B: AutodiffBackend> AutodiffModule<B> for Param<Tensor<B, D, Bool>> {
502    type InnerModule = Param<Tensor<B::InnerBackend, D, Bool>>;
503
504    fn valid(&self) -> Self::InnerModule {
505        Param::initialized(self.id, self.val().inner())
506    }
507
508    fn from_inner(module: Self::InnerModule) -> Self {
509        Param::initialized(module.id, Tensor::from_inner(module.val()))
510    }
511}
512
513#[cfg(all(test, feature = "std"))]
514mod tests {
515    use super::*;
516    use crate::{
517        TestAutodiffBackend,
518        module::Module,
519        record::{BinBytesRecorder, FullPrecisionSettings, Recorder},
520    };
521
522    #[test]
523    fn test_load_record_setting() {
524        let device = Default::default();
525        let tensor = Tensor::<TestAutodiffBackend, 2>::ones([3, 3], &device).require_grad();
526
527        let byte_recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
528        let bytes = byte_recorder
529            .record(
530                Param::initialized(ParamId::new(), tensor.clone()).into_record(),
531                (),
532            )
533            .unwrap();
534
535        let no_grad_is_require_grad = Param::initialized(ParamId::new(), tensor.clone())
536            .no_grad()
537            .load_record(byte_recorder.load(bytes.clone(), &device).unwrap())
538            .is_require_grad();
539
540        let with_default_is_require_grad = Param::initialized(ParamId::new(), tensor)
541            .load_record(byte_recorder.load(bytes, &device).unwrap())
542            .is_require_grad();
543
544        assert!(!no_grad_is_require_grad);
545        assert!(with_default_is_require_grad);
546    }
547
548    #[test]
549    fn test_param_require_grad_stateful() {
550        let device = Default::default();
551        let tensor = Tensor::<TestAutodiffBackend, 2>::ones([3, 3], &device).require_grad();
552
553        let param = Param::initialized(ParamId::new(), tensor);
554        assert!(param.is_require_grad());
555        assert!(param.require_grad);
556
557        let param = param.valid();
558        assert!(!param.is_require_grad());
559        assert!(param.require_grad); // stateful
560
561        // Without `HasAutodiffModule`, we would need to specify the param type as well, which would be annoying:
562        // let param: Param<Tensor<TestAutodiffBackend, _>> = param.train();
563        let param = param.train::<TestAutodiffBackend>();
564        assert!(param.is_require_grad());
565        assert!(param.require_grad); // stateful
566
567        let param = param.no_grad();
568        assert!(!param.is_require_grad());
569        assert!(!param.require_grad); // stateful
570
571        let param = param.valid();
572        assert!(!param.is_require_grad()); // always
573        assert!(!param.require_grad); // stateful
574
575        let param = param.train::<TestAutodiffBackend>();
576        assert!(!param.is_require_grad());
577        assert!(!param.require_grad); // stateful
578    }
579}