1pub(crate) mod adagrad;
8pub(crate) mod sgd;
9
10pub use adagrad::{AdaGrad, AdaGradConfig};
11pub use sgd::{Sgd, SgdConfig};
12
13#[derive(Debug, Clone)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[non_exhaustive]
17pub enum Optimizer {
18 Sgd(Sgd),
20 AdaGrad(AdaGrad),
22}
23
24impl Optimizer {
25 pub fn sgd(feature_count: usize, config: SgdConfig) -> Result<Self, RillError> {
27 Ok(Optimizer::Sgd(Sgd::new(feature_count, config)?))
28 }
29
30 pub fn adagrad(feature_count: usize, config: AdaGradConfig) -> Result<Self, RillError> {
32 Ok(Optimizer::AdaGrad(AdaGrad::new(feature_count, config)?))
33 }
34
35 pub fn param_count(&self) -> usize {
37 match self {
38 Optimizer::Sgd(o) => o.param_count(),
39 Optimizer::AdaGrad(o) => o.param_count(),
40 }
41 }
42
43 pub fn samples_seen(&self) -> u64 {
45 match self {
46 Optimizer::Sgd(o) => o.samples_seen(),
47 Optimizer::AdaGrad(o) => o.samples_seen(),
48 }
49 }
50
51 pub fn step(
56 &mut self,
57 weights: &mut [f64],
58 intercept: &mut f64,
59 grad_weights: &[f64],
60 grad_intercept: f64,
61 ) -> Result<(), RillError> {
62 match self {
63 Optimizer::Sgd(o) => o.step(weights, intercept, grad_weights, grad_intercept),
64 Optimizer::AdaGrad(o) => o.step(weights, intercept, grad_weights, grad_intercept),
65 }
66 }
67
68 pub fn reset(&mut self) {
70 match self {
71 Optimizer::Sgd(o) => o.reset(),
72 Optimizer::AdaGrad(o) => o.reset(),
73 }
74 }
75}
76
77#[cfg(feature = "serde")]
78impl crate::persistence::ValidateState for Optimizer {
79 fn validate_state(&self) -> Result<(), RillError> {
80 match self {
81 Optimizer::Sgd(o) => o.validate_state(),
82 Optimizer::AdaGrad(o) => o.validate_state(),
83 }
84 }
85}
86
87use crate::error::RillError;