use std::io::{Read, Write};
use crate::tensor::Result;
mod sgd;
mod adam;
mod rmsprop;
mod adagrad;
mod radam;
mod nadam;
pub use sgd::{SGD, SGDBuilder};
pub use adam::{Adam, AdamBuilder, AdamW, AdamWBuilder};
pub use rmsprop::{RMSprop, RMSpropBuilder};
pub use adagrad::{Adagrad, AdagradBuilder};
pub use radam::RAdam;
pub use nadam::NAdam;
pub trait Optimizer {
fn step(&mut self) -> Result<()>;
fn zero_grad(&self);
fn lr(&self) -> f64;
fn set_lr(&mut self, lr: f64);
fn set_group_lr(&mut self, _group: usize, lr: f64) {
self.set_lr(lr);
}
fn scale_lr(&mut self, factor: f64) {
self.set_lr(self.lr() * factor);
}
fn reset_state(&mut self) {}
fn save_state_to(&self, _path: &str) -> Result<()> {
Err(crate::tensor::TensorError::new(
"Optimizer::save_state_to: this optimizer does not yet \
implement Stateful; optimizer state cannot be persisted. \
Open a follow-up to add Stateful for this optimizer.",
))
}
}
struct GroupMeta {
lr: f64,
range: std::ops::Range<usize>,
}
fn write_groups<W: Write>(w: &mut W, groups: &[GroupMeta]) -> Result<()> {
use crate::nn::checkpoint::{write_f64_le, write_i64_le, write_u32_le};
write_u32_le(w, groups.len() as u32)?;
for g in groups {
write_f64_le(w, g.lr)?;
write_i64_le(w, g.range.start as i64)?;
write_i64_le(w, g.range.end as i64)?;
}
Ok(())
}
fn read_groups<R: Read>(r: &mut R, n_params: usize, what: &str) -> Result<Vec<GroupMeta>> {
use crate::nn::checkpoint::{read_f64_le, read_i64_le, read_u32_le};
let ng = read_u32_le(r)? as usize;
let mut groups = Vec::with_capacity(ng.min(1024));
let mut expected_start = 0i64;
for i in 0..ng {
let lr = read_f64_le(r)?;
let start = read_i64_le(r)?;
let end = read_i64_le(r)?;
if start != expected_start || end < start || end > n_params as i64 {
return Err(crate::tensor::TensorError::new(&format!(
"{what}: corrupt optimizer state: group {i} range {start}..{end} \
(expected a contiguous partition of 0..{n_params})"
)));
}
expected_start = end;
groups.push(GroupMeta {
lr,
range: start as usize..end as usize,
});
}
if ng > 0 && expected_start != n_params as i64 {
return Err(crate::tensor::TensorError::new(&format!(
"{what}: corrupt optimizer state: groups cover 0..{expected_start} \
of {n_params} params"
)));
}
Ok(groups)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StateKind {
Sgd,
Adam,
AdamW,
RMSprop,
Adagrad,
RAdam,
NAdam,
GradScaler,
}
impl StateKind {
fn tag(self) -> u32 {
match self {
StateKind::Sgd => 1,
StateKind::Adam => 2,
StateKind::AdamW => 3,
StateKind::RMSprop => 4,
StateKind::Adagrad => 5,
StateKind::RAdam => 6,
StateKind::NAdam => 7,
StateKind::GradScaler => 8,
}
}
fn from_tag(tag: u32) -> Option<StateKind> {
Some(match tag {
1 => StateKind::Sgd,
2 => StateKind::Adam,
3 => StateKind::AdamW,
4 => StateKind::RMSprop,
5 => StateKind::Adagrad,
6 => StateKind::RAdam,
7 => StateKind::NAdam,
8 => StateKind::GradScaler,
_ => return None,
})
}
fn name(self) -> &'static str {
match self {
StateKind::Sgd => "SGD",
StateKind::Adam => "Adam",
StateKind::AdamW => "AdamW",
StateKind::RMSprop => "RMSprop",
StateKind::Adagrad => "Adagrad",
StateKind::RAdam => "RAdam",
StateKind::NAdam => "NAdam",
StateKind::GradScaler => "GradScaler",
}
}
}
pub(crate) const STATE_MAGIC: [u8; 4] = *b"FDLO";
pub(crate) const STATE_VERSION: u32 = 1;
fn write_state_header<W: Write>(w: &mut W, kind: StateKind) -> Result<()> {
use crate::nn::checkpoint::write_u32_le;
w.write_all(&STATE_MAGIC).map_err(|e| {
crate::tensor::TensorError::new(&format!("io: {}", e))
})?;
write_u32_le(w, STATE_VERSION)?;
write_u32_le(w, kind.tag())?;
Ok(())
}
fn read_state_header<R: Read>(r: &mut R, expected: StateKind, path: &str) -> Result<()> {
use crate::nn::checkpoint::read_u32_le;
let mut magic = [0u8; 4];
r.read_exact(&mut magic).map_err(|e| {
crate::tensor::TensorError::new(&format!("{path}: io: {}", e))
})?;
if magic != STATE_MAGIC {
return Err(crate::tensor::TensorError::new(&format!(
"{path}: not a current flodl optimizer state file (missing FDLO \
header). If this file was written by an earlier flodl, convert \
it once with flodl::nn::migrate_optim_state_file(src, dst, \
StateKind::{:?}) — the old format carries no type tag, so the \
kind must be supplied.",
expected
)));
}
let version = read_u32_le(r)?;
if version > STATE_VERSION {
return Err(crate::tensor::TensorError::new(&format!(
"{path}: state file version {version} is newer than this flodl \
supports (max {STATE_VERSION}) — upgrade flodl to load it"
)));
}
let tag = read_u32_le(r)?;
let found = StateKind::from_tag(tag).ok_or_else(|| {
crate::tensor::TensorError::new(&format!(
"{path}: unknown state kind tag {tag} (corrupt file, or written \
by a newer flodl)"
))
})?;
if found != expected {
return Err(crate::tensor::TensorError::new(&format!(
"{path}: state file was written by {} but is being loaded into {}",
found.name(), expected.name()
)));
}
Ok(())
}
pub trait Stateful {
fn state_kind(&self) -> StateKind;
fn save_state<W: Write>(&self, w: &mut W) -> Result<()>;
fn load_state<R: Read>(&mut self, r: &mut R) -> Result<()>;
fn save_state_file(&self, path: &str) -> Result<()> {
let kind = self.state_kind();
crate::nn::checkpoint::write_file_atomic(path, |mut w| {
write_state_header(&mut w, kind)?;
self.save_state(&mut w)
})
}
fn load_state_file(&mut self, path: &str) -> Result<()> {
let f = std::fs::File::open(path).map_err(|e| {
crate::tensor::TensorError::new(&format!("io: {}", e))
})?;
let expected = self.state_kind();
if path.ends_with(".gz") {
let mut r = flate2::read::GzDecoder::new(f);
read_state_header(&mut r, expected, path)?;
self.load_state(&mut r)
} else {
let mut r = std::io::BufReader::new(f);
read_state_header(&mut r, expected, path)?;
self.load_state(&mut r)
}
}
}
pub fn migrate_optim_state_file(src: &str, dst: &str, kind: StateKind) -> Result<()> {
use crate::nn::checkpoint::{read_u32_le, write_f64_le};
if matches!(kind, StateKind::Adagrad | StateKind::RAdam | StateKind::NAdam) {
return Err(crate::tensor::TensorError::new(&format!(
"migrate_optim_state_file: {} had no serialized state format \
before the FDLO header — nothing to migrate",
kind.name()
)));
}
let f = std::fs::File::open(src).map_err(|e| {
crate::tensor::TensorError::new(&format!("{src}: io: {}", e))
})?;
let mut r: Box<dyn Read> = if src.ends_with(".gz") {
Box::new(flate2::read::GzDecoder::new(f))
} else {
Box::new(std::io::BufReader::new(f))
};
let mut first = [0u8; 4];
r.read_exact(&mut first).map_err(|e| {
crate::tensor::TensorError::new(&format!("{src}: io: {}", e))
})?;
if first == STATE_MAGIC {
return Err(crate::tensor::TensorError::new(&format!(
"migrate_optim_state_file: {src} already has the current FDLO \
header — nothing to migrate"
)));
}
let io_err = |e: std::io::Error| {
crate::tensor::TensorError::new(&format!("io: {}", e))
};
fn migrate_adam_payload<R: Read, W: Write>(
r: &mut R, w: &mut W, count: u32,
) -> Result<()> {
use crate::nn::checkpoint::{
read_f64_le, read_i64_le, read_tensor_state,
write_f64_le, write_i64_le, write_u32_le, write_tensor_state,
};
write_u32_le(w, count)?;
let lr = read_f64_le(r)?;
write_f64_le(w, lr)?;
let t = read_i64_le(r)?;
for _ in 0..count {
let m = read_tensor_state(r, crate::tensor::Device::CPU)?;
let v = read_tensor_state(r, crate::tensor::Device::CPU)?;
write_tensor_state(w, m.as_ref())?;
write_tensor_state(w, v.as_ref())?;
write_i64_le(w, t)?;
}
std::io::copy(r, w).map_err(|e| {
crate::tensor::TensorError::new(&format!("io: {}", e))
})?;
Ok(())
}
crate::nn::checkpoint::write_file_atomic(dst, |mut w| {
write_state_header(&mut w, kind)?;
match kind {
StateKind::Adam => {
migrate_adam_payload(&mut r, &mut w, u32::from_le_bytes(first))?;
}
StateKind::AdamW => {
let mut rest = [0u8; 4];
r.read_exact(&mut rest).map_err(io_err)?;
let mut wd = [0u8; 8];
wd[..4].copy_from_slice(&first);
wd[4..].copy_from_slice(&rest);
write_f64_le(&mut w, f64::from_le_bytes(wd))?;
let count = read_u32_le(&mut r)?;
migrate_adam_payload(&mut r, &mut w, count)?;
}
_ => {
w.write_all(&first).map_err(io_err)?;
std::io::copy(&mut r, &mut w).map_err(io_err)?;
}
}
Ok(())
})
}
#[cfg(test)]
mod test_helpers {
use crate::nn::parameter::Parameter;
use crate::tensor::{Tensor, TensorOptions};
pub(super) fn make_param(name: &str, shape: &[i64]) -> Parameter {
let t = Tensor::randn(shape, TensorOptions {
dtype: crate::tensor::DType::Float32,
device: crate::tensor::test_device(),
}).unwrap();
Parameter::new(t, name)
}
pub(super) fn state_tmp(name: &str) -> String {
std::env::temp_dir()
.join(format!("flodl_optim_state_{}_{}", std::process::id(), name))
.to_string_lossy()
.into_owned()
}
}
#[cfg(test)]
mod tests {
use super::*;
use super::test_helpers::make_param;
use crate::nn::parameter::Parameter;
#[test]
fn test_empty_params_optimizers_no_panic() {
let empty: &[Parameter] = &[];
let mut adam = Adam::new(empty, 0.001);
adam.step().unwrap();
adam.zero_grad();
let mut sgd = SGD::new(empty, 0.01, 0.9);
sgd.step().unwrap();
sgd.zero_grad();
let mut adamw = AdamW::new(empty, 0.001, 0.01);
adamw.step().unwrap();
adamw.zero_grad();
let mut rmsprop = RMSprop::new(empty, 0.01);
rmsprop.step().unwrap();
rmsprop.zero_grad();
let mut adagrad = Adagrad::new(empty, 0.01);
adagrad.step().unwrap();
adagrad.zero_grad();
let mut radam = RAdam::new(empty, 0.01);
radam.step().unwrap();
radam.zero_grad();
let mut nadam = NAdam::new(empty, 0.01);
nadam.step().unwrap();
nadam.zero_grad();
}
#[test]
fn test_step_after_zero_grad_on_fresh_optimizer() {
let p = make_param("w", &[3, 2]);
let mut adam = Adam::new(std::slice::from_ref(&p), 0.001);
let mut sgd = SGD::new(std::slice::from_ref(&p), 0.01, 0.9);
adam.zero_grad();
adam.step().unwrap();
sgd.zero_grad();
sgd.step().unwrap();
let vals = p.variable.data().to_f32_vec().unwrap();
for (i, &v) in vals.iter().enumerate() {
assert!(v.is_finite(), "param[{}] should be finite after step-without-backward: {}", i, v);
}
}
#[test]
fn test_set_lr_all_optimizers() {
let p = make_param("w", &[2]);
let mut adam = Adam::new(std::slice::from_ref(&p), 0.001);
adam.set_lr(0.42);
assert!((adam.lr() - 0.42).abs() < 1e-12, "Adam set_lr failed");
let mut sgd = SGD::new(std::slice::from_ref(&p), 0.01, 0.0);
sgd.set_lr(0.42);
assert!((sgd.lr() - 0.42).abs() < 1e-12, "SGD set_lr failed");
let mut adamw = AdamW::new(std::slice::from_ref(&p), 0.001, 0.01);
adamw.set_lr(0.42);
assert!((adamw.lr() - 0.42).abs() < 1e-12, "AdamW set_lr failed");
let mut rmsprop = RMSprop::new(std::slice::from_ref(&p), 0.01);
rmsprop.set_lr(0.42);
assert!((rmsprop.lr() - 0.42).abs() < 1e-12, "RMSprop set_lr failed");
let mut nadam = NAdam::new(std::slice::from_ref(&p), 0.01);
nadam.set_lr(0.42);
assert!((nadam.lr() - 0.42).abs() < 1e-12, "NAdam set_lr failed");
let mut radam = RAdam::new(std::slice::from_ref(&p), 0.01);
radam.set_lr(0.42);
assert!((radam.lr() - 0.42).abs() < 1e-12, "RAdam set_lr failed");
let mut adagrad = Adagrad::new(std::slice::from_ref(&p), 0.01);
adagrad.set_lr(0.42);
assert!((adagrad.lr() - 0.42).abs() < 1e-12, "Adagrad set_lr failed");
}
}