Skip to main content

ruda_optim/optim/simple/record/
v1.rs

1
2use crate::optim::SimpleOptimizer;
3use ruda_model::record::{PrecisionSettings, Record};
4use ruda_model::tensor::backend::Backend;
5use core::any::Any;
6use serde::{Deserialize, Serialize};
7
8#[cfg(not(feature = "std"))]
9use alloc::boxed::Box;
10
11/// [Optimizer adaptor](crate::optim::simple::adaptor::OptimizerAdaptor) record item.
12pub enum AdaptorRecordV1<O: SimpleOptimizer<B>, B: Backend> {
13    /// Rank 0.
14    Rank0(O::State<0>),
15
16    /// Rank 1.
17    Rank1(O::State<1>),
18
19    /// Rank 2.
20    Rank2(O::State<2>),
21
22    /// Rank 3.
23    Rank3(O::State<3>),
24
25    /// Rank 4.
26    Rank4(O::State<4>),
27
28    /// Rank 5.
29    Rank5(O::State<5>),
30
31    /// Rank 6.
32    Rank6(O::State<6>),
33
34    /// Rank 7.
35    Rank7(O::State<7>),
36
37    /// Rank 8.
38    Rank8(O::State<8>),
39}
40
41impl<O: SimpleOptimizer<B>, B: Backend> Clone for AdaptorRecordV1<O, B> {
42    fn clone(&self) -> Self {
43        match self {
44            AdaptorRecordV1::Rank0(record) => AdaptorRecordV1::Rank0(record.clone()),
45            AdaptorRecordV1::Rank1(record) => AdaptorRecordV1::Rank1(record.clone()),
46            AdaptorRecordV1::Rank2(record) => AdaptorRecordV1::Rank2(record.clone()),
47            AdaptorRecordV1::Rank3(record) => AdaptorRecordV1::Rank3(record.clone()),
48            AdaptorRecordV1::Rank4(record) => AdaptorRecordV1::Rank4(record.clone()),
49            AdaptorRecordV1::Rank5(record) => AdaptorRecordV1::Rank5(record.clone()),
50            AdaptorRecordV1::Rank6(record) => AdaptorRecordV1::Rank6(record.clone()),
51            AdaptorRecordV1::Rank7(record) => AdaptorRecordV1::Rank7(record.clone()),
52            AdaptorRecordV1::Rank8(record) => AdaptorRecordV1::Rank8(record.clone()),
53        }
54    }
55}
56
57/// [Optimizer adaptor](crate::optim::simple::adaptor::OptimizerAdaptor) record item.
58#[derive(Serialize, Deserialize, Clone)]
59#[serde(bound = "")]
60pub enum AdaptorRecordItemV1<O: SimpleOptimizer<B>, B: Backend, S: PrecisionSettings> {
61    /// Rank 0.
62    Rank0(<O::State<0> as Record<B>>::Item<S>),
63
64    /// Rank 1.
65    Rank1(<O::State<1> as Record<B>>::Item<S>),
66
67    /// Rank 2.
68    Rank2(<O::State<2> as Record<B>>::Item<S>),
69
70    /// Rank 3.
71    Rank3(<O::State<3> as Record<B>>::Item<S>),
72
73    /// Rank 4.
74    Rank4(<O::State<4> as Record<B>>::Item<S>),
75
76    /// Rank 5.
77    Rank5(<O::State<5> as Record<B>>::Item<S>),
78
79    /// Rank 6.
80    Rank6(<O::State<6> as Record<B>>::Item<S>),
81
82    /// Rank 7.
83    Rank7(<O::State<7> as Record<B>>::Item<S>),
84
85    /// Rank 8.
86    Rank8(<O::State<8> as Record<B>>::Item<S>),
87}
88
89impl<O, B> AdaptorRecordV1<O, B>
90where
91    O: SimpleOptimizer<B>,
92    B: Backend,
93{
94    /// Convert the record into the state.
95    ///
96    /// # Returns
97    ///
98    /// The state.
99    ///
100    /// # Panics
101    ///
102    /// Panics if the state dimension is not supported.
103    pub fn into_state<const D: usize>(self) -> O::State<D> {
104        let boxed_state: Box<dyn Any> = match self {
105            AdaptorRecordV1::Rank0(s) => Box::new(s),
106            AdaptorRecordV1::Rank1(s) => Box::new(s),
107            AdaptorRecordV1::Rank2(s) => Box::new(s),
108            AdaptorRecordV1::Rank3(s) => Box::new(s),
109            AdaptorRecordV1::Rank4(s) => Box::new(s),
110            AdaptorRecordV1::Rank5(s) => Box::new(s),
111            AdaptorRecordV1::Rank6(s) => Box::new(s),
112            AdaptorRecordV1::Rank7(s) => Box::new(s),
113            AdaptorRecordV1::Rank8(s) => Box::new(s),
114        };
115        let state = boxed_state
116            .downcast::<O::State<D>>()
117            .expect("Unsupported state dimension, dimension up to 8 are supported.");
118        *state
119    }
120
121    /// Convert the state into the record.
122    ///
123    /// # Arguments
124    ///
125    /// * `state`: The state.
126    ///
127    /// # Returns
128    ///
129    /// The record.
130    pub fn from_state<const D: usize>(state: O::State<D>) -> Self {
131        let state: Box<dyn Any> = Box::new(state);
132
133        match D {
134            0 => AdaptorRecordV1::Rank0(*state.downcast().unwrap()),
135            1 => AdaptorRecordV1::Rank1(*state.downcast().unwrap()),
136            2 => AdaptorRecordV1::Rank2(*state.downcast().unwrap()),
137            3 => AdaptorRecordV1::Rank3(*state.downcast().unwrap()),
138            4 => AdaptorRecordV1::Rank4(*state.downcast().unwrap()),
139            5 => AdaptorRecordV1::Rank5(*state.downcast().unwrap()),
140            6 => AdaptorRecordV1::Rank6(*state.downcast().unwrap()),
141            7 => AdaptorRecordV1::Rank7(*state.downcast().unwrap()),
142            8 => AdaptorRecordV1::Rank8(*state.downcast().unwrap()),
143            _ => panic!("Unsupported state dimension, dimension up to 8 are supported."),
144        }
145    }
146}
147
148impl<O, B> Record<B> for AdaptorRecordV1<O, B>
149where
150    O: SimpleOptimizer<B>,
151    B: Backend,
152{
153    type Item<S: PrecisionSettings> = AdaptorRecordItemV1<O, B, S>;
154
155    fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
156        match self {
157            AdaptorRecordV1::Rank0(record) => AdaptorRecordItemV1::Rank0(record.into_item()),
158            AdaptorRecordV1::Rank1(record) => AdaptorRecordItemV1::Rank1(record.into_item()),
159            AdaptorRecordV1::Rank2(record) => AdaptorRecordItemV1::Rank2(record.into_item()),
160            AdaptorRecordV1::Rank3(record) => AdaptorRecordItemV1::Rank3(record.into_item()),
161            AdaptorRecordV1::Rank4(record) => AdaptorRecordItemV1::Rank4(record.into_item()),
162            AdaptorRecordV1::Rank5(record) => AdaptorRecordItemV1::Rank5(record.into_item()),
163            AdaptorRecordV1::Rank6(record) => AdaptorRecordItemV1::Rank6(record.into_item()),
164            AdaptorRecordV1::Rank7(record) => AdaptorRecordItemV1::Rank7(record.into_item()),
165            AdaptorRecordV1::Rank8(record) => AdaptorRecordItemV1::Rank8(record.into_item()),
166        }
167    }
168
169    fn from_item<S: PrecisionSettings>(item: Self::Item<S>, device: &B::Device) -> Self {
170        match item {
171            AdaptorRecordItemV1::Rank0(item) => {
172                AdaptorRecordV1::Rank0(<O::State<0> as Record<B>>::from_item(item, device))
173            }
174            AdaptorRecordItemV1::Rank1(item) => {
175                AdaptorRecordV1::Rank1(<O::State<1> as Record<B>>::from_item(item, device))
176            }
177            AdaptorRecordItemV1::Rank2(item) => {
178                AdaptorRecordV1::Rank2(<O::State<2> as Record<B>>::from_item(item, device))
179            }
180            AdaptorRecordItemV1::Rank3(item) => {
181                AdaptorRecordV1::Rank3(<O::State<3> as Record<B>>::from_item(item, device))
182            }
183            AdaptorRecordItemV1::Rank4(item) => {
184                AdaptorRecordV1::Rank4(<O::State<4> as Record<B>>::from_item(item, device))
185            }
186            AdaptorRecordItemV1::Rank5(item) => {
187                AdaptorRecordV1::Rank5(<O::State<5> as Record<B>>::from_item(item, device))
188            }
189            AdaptorRecordItemV1::Rank6(item) => {
190                AdaptorRecordV1::Rank6(<O::State<6> as Record<B>>::from_item(item, device))
191            }
192            AdaptorRecordItemV1::Rank7(item) => {
193                AdaptorRecordV1::Rank7(<O::State<7> as Record<B>>::from_item(item, device))
194            }
195            AdaptorRecordItemV1::Rank8(item) => {
196                AdaptorRecordV1::Rank8(<O::State<8> as Record<B>>::from_item(item, device))
197            }
198        }
199    }
200}