ruda_optim/optim/simple/record/
v1.rs1
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
11pub enum AdaptorRecordV1<O: SimpleOptimizer<B>, B: Backend> {
13 Rank0(O::State<0>),
15
16 Rank1(O::State<1>),
18
19 Rank2(O::State<2>),
21
22 Rank3(O::State<3>),
24
25 Rank4(O::State<4>),
27
28 Rank5(O::State<5>),
30
31 Rank6(O::State<6>),
33
34 Rank7(O::State<7>),
36
37 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#[derive(Serialize, Deserialize, Clone)]
59#[serde(bound = "")]
60pub enum AdaptorRecordItemV1<O: SimpleOptimizer<B>, B: Backend, S: PrecisionSettings> {
61 Rank0(<O::State<0> as Record<B>>::Item<S>),
63
64 Rank1(<O::State<1> as Record<B>>::Item<S>),
66
67 Rank2(<O::State<2> as Record<B>>::Item<S>),
69
70 Rank3(<O::State<3> as Record<B>>::Item<S>),
72
73 Rank4(<O::State<4> as Record<B>>::Item<S>),
75
76 Rank5(<O::State<5> as Record<B>>::Item<S>),
78
79 Rank6(<O::State<6> as Record<B>>::Item<S>),
81
82 Rank7(<O::State<7> as Record<B>>::Item<S>),
84
85 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 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 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}