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)]
21pub type ConstantRecord = EmptyRecord;
23
24#[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 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#[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 }
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
132empty!(alloc::string::String);
136empty!(bool);
137
138empty!(f64);
140empty!(f32);
141empty!(half::bf16);
142empty!(half::f16);
143
144empty!(usize);
146empty!(u64);
147empty!(u32);
148empty!(u16);
149empty!(u8);
150
151empty!(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
165impl<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 }
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#[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 }
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 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)]
364impl<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}