1
2use super::SimpleOptimizer;
3use super::adaptor::OptimizerAdaptor;
4use super::decay::{WeightDecay, WeightDecayConfig};
5use super::momentum::{Momentum, MomentumConfig, MomentumState};
6use crate::LearningRate;
7use crate::grad_clipping::GradientClippingConfig;
8use ruda_model::config::Config;
9use ruda_model::module::AutodiffModule;
10use ruda_model::record::Record;
11use ruda_model::tensor::Tensor;
12use ruda_model::tensor::backend::{AutodiffBackend, Backend};
13
14#[derive(Config, Debug)]
16pub struct SgdConfig {
17 weight_decay: Option<WeightDecayConfig>,
19 momentum: Option<MomentumConfig>,
21 gradient_clipping: Option<GradientClippingConfig>,
23}
24
25#[derive(Clone)]
29pub struct Sgd<B: Backend> {
30 momentum: Option<Momentum<B>>,
31 weight_decay: Option<WeightDecay>,
32}
33
34#[derive(Record, Clone, new)]
36pub struct SgdState<B: Backend, const D: usize> {
37 pub momentum: Option<MomentumState<B, D>>,
39}
40
41impl SgdConfig {
42 pub fn build<B: Backend>(&self) -> Sgd<B> {
44 Sgd {
45 momentum: self.momentum.as_ref().map(Momentum::new),
46 weight_decay: self.weight_decay.as_ref().map(WeightDecay::new),
47 }
48 }
49
50 pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
52 &self,
53 ) -> OptimizerAdaptor<Sgd<B::InnerBackend>, M, B> {
54 let mut optim = OptimizerAdaptor::from(self.build());
55 if let Some(config) = &self.gradient_clipping {
56 optim = optim.with_grad_clipping(config.init());
57 }
58 optim
59 }
60}
61
62impl<B: Backend> SimpleOptimizer<B> for Sgd<B> {
63 type State<const D: usize> = SgdState<B, D>;
64
65 fn step<const D: usize>(
66 &self,
67 lr: LearningRate,
68 tensor: Tensor<B, D>,
69 mut grad: Tensor<B, D>,
70 state: Option<Self::State<D>>,
71 ) -> (Tensor<B, D>, Option<Self::State<D>>) {
72 let mut state_momentum = None;
73
74 if let Some(state) = state {
75 state_momentum = state.momentum;
76 }
77
78 if let Some(weight_decay) = &self.weight_decay {
79 grad = weight_decay.transform(grad, tensor.clone());
80 }
81
82 if let Some(momentum) = &self.momentum {
83 let (grad_out, state) = momentum.transform(grad, state_momentum);
84 state_momentum = Some(state);
85 grad = grad_out;
86 }
87
88 let state = SgdState::new(state_momentum);
89 let delta = grad.mul_scalar(lr);
90
91 (tensor - delta, Some(state))
92 }
93
94 fn to_device<const D: usize>(mut state: Self::State<D>, device: &B::Device) -> Self::State<D> {
95 state.momentum = state.momentum.map(|state| state.to_device(device));
96 state
97 }
98}
99
100#[cfg(test)]
101mod tests {
102 use super::*;
103 use crate::{
104 TestAutodiffBackend, TestBackend,
105 grad_clipping::GradientClipping,
106 optim::{GradientsParams, Optimizer},
107 };
108 use ruda_model::tensor::{Distribution, Shape};
109 use ruda_nn::{Linear, LinearConfig};
110
111 const LEARNING_RATE: LearningRate = 0.02;
112
113 #[test]
114 fn with_updated_params_should_have_state() {
115 let device = Default::default();
116 let layer = layer::<TestAutodiffBackend>(&device);
117 let mut optim = sgd_with_all();
118 let loss = layer.forward(random_tensor::<TestAutodiffBackend>(&device));
119 let grads = loss.backward();
120 let grads = GradientsParams::from_grads(grads, &layer);
121 let _layer = optim.step(LEARNING_RATE, layer, grads);
122
123 let record = optim.to_record();
124
125 assert!(!record.is_empty());
126 }
127
128 #[test]
129 fn without_updated_params_should_not_have_state() {
130 let optim = sgd_with_all();
131 let record = optim.to_record();
132 assert!(record.is_empty());
133 }
134
135 #[test]
136 fn can_attach_gradient_clipping() {
137 let optim = sgd_with_all().with_grad_clipping(GradientClipping::Value(0.5));
138 assert!(optim.has_gradient_clipping());
139 }
140
141 #[test]
142 fn should_load_state() {
143 let device = Default::default();
144 let layer = layer::<TestAutodiffBackend>(&device);
145 let mut optim = sgd_with_all();
146 let loss = layer.forward(random_tensor(&device));
147 let grads = loss.backward();
148 let grads = GradientsParams::from_grads(grads, &layer);
149 let _layer = optim.step(LEARNING_RATE, layer, grads);
150
151 let record = optim.to_record();
152 let optim_new = sgd_with_all();
153 let record_new = optim_new.to_record();
154 let optim_new = optim_new.load_record(record.clone());
155 let state_restored = optim_new.to_record();
156
157 assert_ne!(record.len(), record_new.len());
158 assert_eq!(record.len(), state_restored.len());
159 }
160
161 fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
162 Tensor::<B, 2>::random(Shape::new([2, 20]), Distribution::Default, device)
163 }
164
165 fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
166 LinearConfig::new(20, 20).init(device)
167 }
168
169 fn sgd_with_all()
170 -> OptimizerAdaptor<Sgd<TestBackend>, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
171 SgdConfig {
172 weight_decay: Some(WeightDecayConfig { penalty: 0.05 }),
173 momentum: Some(MomentumConfig {
174 momentum: 0.9,
175 dampening: 0.1,
176 nesterov: true,
177 }),
178 gradient_clipping: None,
179 }
180 .init()
181 }
182}