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 pub fn from_tensor(value: Tensor<B, D>) -> Self {
70 Param::initialized(ParamId::new(), value.require_grad())
73 }
74
75 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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) }
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) }
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 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 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); let param = param.train::<TestAutodiffBackend>();
564 assert!(param.is_require_grad());
565 assert!(param.require_grad); let param = param.no_grad();
568 assert!(!param.is_require_grad());
569 assert!(!param.require_grad); let param = param.valid();
572 assert!(!param.is_require_grad()); assert!(!param.require_grad); let param = param.train::<TestAutodiffBackend>();
576 assert!(!param.is_require_grad());
577 assert!(!param.require_grad); }
579}