Skip to main content

ruda_optim/optim/simple/record/
base.rs

1
2use super::{AdaptorRecordItemV1, AdaptorRecordV1};
3use crate::optim::SimpleOptimizer;
4use ruda_model::record::{PrecisionSettings, Record};
5use ruda_model::tensor::backend::AutodiffBackend;
6use serde::{Deserialize, Serialize};
7
8/// [Optimizer adaptor](crate::optim::simple::adaptor::OptimizerAdaptor) record.
9///
10/// Records are versioned for backward compatibility, so old records can be loaded.
11pub enum AdaptorRecord<O, B>
12where
13    O: SimpleOptimizer<B::InnerBackend>,
14    B: AutodiffBackend,
15{
16    /// Version 1.
17    V1(AdaptorRecordV1<O, B::InnerBackend>),
18}
19
20/// [Optimizer adaptor](crate::optim::simple::adaptor::OptimizerAdaptor) record item.
21#[derive(Serialize, Deserialize, Clone)]
22#[serde(bound = "")]
23pub enum AdaptorRecordItem<
24    O: SimpleOptimizer<B::InnerBackend>,
25    B: AutodiffBackend,
26    S: PrecisionSettings,
27> {
28    /// Version 1.
29    V1(AdaptorRecordItemV1<O, B::InnerBackend, S>),
30}
31
32impl<O, B> Record<B> for AdaptorRecord<O, B>
33where
34    O: SimpleOptimizer<B::InnerBackend>,
35    B: AutodiffBackend,
36{
37    type Item<S: PrecisionSettings> = AdaptorRecordItem<O, B, S>;
38
39    fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
40        match self {
41            AdaptorRecord::V1(record) => AdaptorRecordItem::V1(record.into_item()),
42        }
43    }
44
45    fn from_item<S: PrecisionSettings>(item: Self::Item<S>, device: &B::Device) -> Self {
46        match item {
47            AdaptorRecordItem::V1(item) => Self::V1(AdaptorRecordV1::from_item(item, device)),
48        }
49    }
50}
51
52impl<O, B> Clone for AdaptorRecord<O, B>
53where
54    O: SimpleOptimizer<B::InnerBackend>,
55    B: AutodiffBackend,
56{
57    fn clone(&self) -> Self {
58        match self {
59            AdaptorRecord::V1(record) => Self::V1(record.clone()),
60        }
61    }
62}
63
64impl<O, B> AdaptorRecord<O, B>
65where
66    O: SimpleOptimizer<B::InnerBackend>,
67    B: AutodiffBackend,
68{
69    /// Converts the record into the optimizer state.
70    ///
71    /// # Returns
72    ///
73    /// The optimizer state.
74    pub fn into_state<const D: usize>(self) -> O::State<D> {
75        match self {
76            AdaptorRecord::V1(record) => record.into_state(),
77        }
78    }
79
80    /// Converts the optimizer state into the record.
81    ///
82    /// # Arguments
83    ///
84    /// * `state`: The optimizer state.
85    ///
86    /// # Returns
87    ///
88    /// The record.
89    pub fn from_state<const D: usize>(state: O::State<D>) -> Self {
90        Self::V1(AdaptorRecordV1::from_state(state))
91    }
92}