Skip to main content

ruda_model/module/param/
constant.rs

1use alloc::{format, string::ToString};
2use core::{fmt::Display, marker::PhantomData};
3
4use crate::{
5    module::{
6        AutodiffModule, Content, Devices, Module, ModuleDisplay, ModuleDisplayDefault,
7        ModuleMapper, ModuleVisitor,
8    },
9    record::{PrecisionSettings, Record},
10};
11use ruda_tensor::api::{
12    BasicAutodiffOps, BasicOps, Tensor,
13    backend::{AutodiffBackend, Backend},
14    ops::Device,
15};
16
17#[deprecated(
18    since = "0.21.0",
19    note = "ConstantRecord is misleading as it doesn't persist data. Use EmptyRecord instead."
20)]
21/// A record representing the absence of persistent module state.
22pub type ConstantRecord = EmptyRecord;
23
24/// A record representing the absence of persistent module state.
25///
26/// `EmptyRecord` is used for modules that do not store any data to be
27/// serialized or restored (e.g., modules marked with `#[module(skip)]`
28/// or modules without parameters).
29///
30/// This record contains no fields and serializes to `None`.
31#[derive(Debug, Clone, Copy, new, Default, PartialEq, Eq)]
32pub struct EmptyRecord;
33
34impl serde::Serialize for EmptyRecord {
35    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
36    where
37        S: serde::Serializer,
38    {
39        // nothing to serialize
40        S::serialize_none(serializer)
41    }
42}
43
44impl<'de> serde::Deserialize<'de> for EmptyRecord {
45    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
46    where
47        D: serde::Deserializer<'de>,
48    {
49        deserializer.deserialize_option(serde::de::IgnoredAny).ok();
50        Ok(EmptyRecord::new())
51    }
52}
53
54impl<B: Backend> Record<B> for EmptyRecord {
55    type Item<S: PrecisionSettings> = EmptyRecord;
56
57    fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
58        self
59    }
60
61    fn from_item<S: PrecisionSettings>(item: Self::Item<S>, _device: &B::Device) -> Self {
62        item
63    }
64}
65/// Constant macro.
66#[macro_export]
67macro_rules! empty {
68    (module) => {
69        type Record = $crate::module::EmptyRecord;
70
71        fn visit<V: $crate::module::ModuleVisitor<B>>(&self, _visitor: &mut V) {
72            // Nothing to do
73        }
74
75        fn map<M: $crate::module::ModuleMapper<B>>(self, _mapper: &mut M) -> Self {
76            self
77        }
78
79        fn load_record(self, _record: Self::Record) -> Self {
80            self
81        }
82
83        fn into_record(self) -> Self::Record {
84            $crate::module::EmptyRecord::new()
85        }
86
87        fn to_device(self, _: &B::Device) -> Self {
88            self
89        }
90
91        fn fork(self, _: &B::Device) -> Self {
92            self
93        }
94
95        fn collect_devices(&self, devices: $crate::module::Devices<B>) -> $crate::module::Devices<B> {
96            devices
97        }
98    };
99
100    (ad_module, $type:ty) => {
101        type InnerModule = $type;
102
103        fn valid(&self) -> Self::InnerModule {
104            self.clone()
105        }
106
107        fn from_inner(module: Self::InnerModule) -> Self {
108            module
109        }
110    };
111
112    ($type:ty) => {
113        impl<B: $crate::tensor::backend::Backend> $crate::module::Module<B> for $type {
114            $crate::empty!(module);
115        }
116
117        impl<B: $crate::tensor::backend::AutodiffBackend> $crate::module::AutodiffModule<B> for $type {
118            $crate::empty!(ad_module, $type);
119        }
120
121        impl $crate::module::ModuleDisplayDefault for $type {
122            fn content(&self, content: $crate::module::Content) -> Option<$crate::module::Content> {
123                let string = format!("{}", self);
124                content.add_formatted(&string).optional()
125            }
126        }
127
128        impl $crate::module::ModuleDisplay for $type {}
129    };
130}
131
132// TODO: breaking change for these constant types (currently empty record, non-persistent)?
133
134// General Types
135empty!(alloc::string::String);
136empty!(bool);
137
138// Float Types
139empty!(f64);
140empty!(f32);
141empty!(half::bf16);
142empty!(half::f16);
143
144// Unsigned Integer Types
145empty!(usize);
146empty!(u64);
147empty!(u32);
148empty!(u16);
149empty!(u8);
150
151// Signed Integer Types
152empty!(isize);
153empty!(i64);
154empty!(i32);
155empty!(i16);
156empty!(i8);
157
158impl ruda_model::module::ModuleDisplay for str {}
159impl ruda_model::module::ModuleDisplayDefault for str {
160    fn content(&self, content: ruda_model::module::Content) -> Option<ruda_model::module::Content> {
161        content.add_formatted(&self).optional()
162    }
163}
164
165// TODO: tensor record should persist
166impl<const D: usize, B: Backend, K: BasicOps<B>> Module<B> for Tensor<B, D, K> {
167    type Record = EmptyRecord;
168
169    fn visit<V: ModuleVisitor<B>>(&self, _visitor: &mut V) {}
170
171    fn map<M: ModuleMapper<B>>(self, _mapper: &mut M) -> Self {
172        self
173    }
174
175    fn into_record(self) -> Self::Record {
176        EmptyRecord
177    }
178
179    fn load_record(self, _record: Self::Record) -> Self {
180        self
181    }
182
183    fn to_device(self, device: &B::Device) -> Self {
184        self.to_device(device)
185    }
186
187    fn fork(self, device: &B::Device) -> Self {
188        self.to_device(device)
189    }
190
191    fn collect_devices(&self, mut devices: Devices<B>) -> Devices<B> {
192        let device = self.device();
193
194        if !devices.contains(&device) {
195            devices.push(device)
196        }
197
198        devices
199    }
200}
201
202impl<const D: usize, B: Backend, K: BasicOps<B>> ModuleDisplayDefault for Tensor<B, D, K> {
203    fn content(&self, content: Content) -> Option<Content> {
204        let string = format!("Tensor {{rank: {D}, shape: {:?}}}", self.shape().as_slice());
205        content.add_single(&string).optional()
206    }
207}
208
209impl<const D: usize, B: Backend, K: BasicOps<B>> ModuleDisplay for Tensor<B, D, K> {}
210
211impl<const D: usize, B: AutodiffBackend, K: BasicAutodiffOps<B>> AutodiffModule<B>
212    for Tensor<B, D, K>
213{
214    type InnerModule = Tensor<B::InnerBackend, D, K::InnerKind>;
215
216    fn valid(&self) -> Self::InnerModule {
217        self.clone().inner()
218    }
219
220    fn from_inner(tensor: Self::InnerModule) -> Self {
221        Tensor::from_inner(tensor)
222    }
223}
224
225impl<B: Backend> Module<B> for PhantomData<B> {
226    type Record = EmptyRecord;
227
228    fn visit<V: ModuleVisitor<B>>(&self, _visitor: &mut V) {
229        // Nothing to do
230    }
231
232    fn map<M: ModuleMapper<B>>(self, _mapper: &mut M) -> Self {
233        self
234    }
235
236    fn load_record(self, _record: Self::Record) -> Self {
237        self
238    }
239
240    fn into_record(self) -> Self::Record {
241        EmptyRecord::new()
242    }
243
244    fn to_device(self, _: &Device<B>) -> Self {
245        self
246    }
247
248    fn fork(self, _: &Device<B>) -> Self {
249        self
250    }
251
252    fn collect_devices(&self, devices: Devices<B>) -> Devices<B> {
253        devices
254    }
255}
256
257impl<B: Backend> ModuleDisplayDefault for PhantomData<B> {
258    fn content(&self, content: Content) -> Option<Content> {
259        content.add_single(&"PhantomData".to_string()).optional()
260    }
261}
262
263impl<B: Backend> ModuleDisplay for PhantomData<B> {}
264
265impl<B: AutodiffBackend> AutodiffModule<B> for PhantomData<B> {
266    type InnerModule = PhantomData<B::InnerBackend>;
267
268    fn valid(&self) -> Self::InnerModule {
269        PhantomData
270    }
271
272    fn from_inner(_module: Self::InnerModule) -> Self {
273        PhantomData
274    }
275}
276
277/// Container to satisfy the Module trait for types that are not modules.
278#[derive(Clone, Debug)]
279#[deprecated(
280    since = "0.21.0",
281    note = "Ignored<T> is deprecated. Use #[module(skip)] for non-persistent fields (same behavior)."
282)]
283pub struct Ignored<T>(pub T);
284
285#[allow(deprecated)]
286impl<B, T> Module<B> for Ignored<T>
287where
288    B: Backend,
289    T: Sync + Send + core::fmt::Debug + Clone,
290{
291    type Record = EmptyRecord;
292
293    fn visit<V: ModuleVisitor<B>>(&self, _visitor: &mut V) {
294        // Nothing to do
295    }
296
297    fn map<M: ModuleMapper<B>>(self, _mapper: &mut M) -> Self {
298        self
299    }
300
301    fn load_record(self, _record: Self::Record) -> Self {
302        self
303    }
304
305    fn into_record(self) -> Self::Record {
306        EmptyRecord::new()
307    }
308
309    fn to_device(self, _: &Device<B>) -> Self {
310        self
311    }
312
313    fn fork(self, _: &Device<B>) -> Self {
314        self
315    }
316
317    fn collect_devices(&self, devices: Devices<B>) -> Devices<B> {
318        devices
319    }
320}
321
322#[allow(deprecated)]
323impl<T> ModuleDisplayDefault for Ignored<T>
324where
325    T: Sync + Send + core::fmt::Debug + Clone,
326{
327    fn content(&self, content: Content) -> Option<Content> {
328        // For now, just print the debug representation of the ignored value
329        content.add_single(&format!("{:?}", self.0)).optional()
330    }
331}
332
333#[allow(deprecated)]
334impl<T> ModuleDisplay for Ignored<T> where T: Sync + Send + core::fmt::Debug + Clone {}
335
336#[allow(deprecated)]
337impl<T> Display for Ignored<T>
338where
339    T: Sync + Send + core::fmt::Debug + Clone,
340{
341    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
342        write!(f, "{:?}", self.0)
343    }
344}
345
346#[allow(deprecated)]
347impl<B: AutodiffBackend, T> AutodiffModule<B> for Ignored<T>
348where
349    B: AutodiffBackend,
350    T: Sync + Send + core::fmt::Debug + Clone,
351{
352    type InnerModule = Ignored<T>;
353
354    fn valid(&self) -> Self::InnerModule {
355        self.clone()
356    }
357
358    fn from_inner(module: Self::InnerModule) -> Self {
359        module
360    }
361}
362
363#[allow(deprecated)]
364// Implement deref for Ignored
365impl<T> core::ops::Deref for Ignored<T> {
366    type Target = T;
367
368    fn deref(&self) -> &Self::Target {
369        &self.0
370    }
371}
372
373#[cfg(all(test, feature = "std"))]
374mod tests {
375    use core::marker::PhantomData;
376
377    use ruda_tensor::api::backend::Backend;
378    use ruda_tensor::api::{Device, Tensor};
379
380    use crate::TestBackend;
381    use crate::{
382        TestAutodiffBackend,
383        record::{BinBytesRecorder, FullPrecisionSettings, Recorder},
384    };
385    use ruda_model::module::Module;
386
387
388    #[test]
389    fn tensor_load_record_setting() {
390        let device: &Device<TestAutodiffBackend> = &Default::default();
391        let tensor = Tensor::<TestAutodiffBackend, 2>::ones([3, 3], device);
392
393        let byte_recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
394        let bytes = Recorder::<TestAutodiffBackend>::record(
395            &byte_recorder,
396            tensor.clone().into_record(),
397            (),
398        )
399        .unwrap();
400
401        let no_grad_is_require_grad = tensor
402            .clone()
403            .no_grad()
404            .load_record(
405                Recorder::<TestAutodiffBackend>::load(&byte_recorder, bytes.clone(), device)
406                    .unwrap(),
407            )
408            .is_require_grad();
409
410        let with_default_is_require_grad = tensor
411            .load_record(
412                Recorder::<TestAutodiffBackend>::load(&byte_recorder, bytes.clone(), device)
413                    .unwrap(),
414            )
415            .is_require_grad();
416
417        assert!(!no_grad_is_require_grad);
418        assert!(!with_default_is_require_grad);
419    }
420
421    #[test]
422    fn empty_module_with_phantom() {
423        #[derive(Module, Debug, new)]
424        struct EmptyModule<B: Backend> {
425            _phantom: PhantomData<B>,
426        }
427
428        let _module = EmptyModule::<TestBackend>::new();
429
430        assert_eq!(core::mem::size_of::<EmptyModule<TestBackend>>(), 0);
431    }
432}