use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroupState};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use torsh_core::error::Result;
use torsh_tensor::Tensor;
pub struct Lookahead<O: Optimizer> {
base_optimizer: O,
slow_weights: HashMap<String, Tensor>,
alpha: f32,
k: usize,
step_count: usize,
}
impl<O: Optimizer> Lookahead<O> {
pub fn new(base_optimizer: O, alpha: f32, k: usize) -> Self {
Self {
base_optimizer,
slow_weights: HashMap::new(),
alpha,
k,
step_count: 0,
}
}
pub fn with_defaults(base_optimizer: O) -> Self {
Self::new(base_optimizer, 0.5, 5)
}
fn get_param_key(param_ref: &Arc<RwLock<Tensor>>) -> String {
format!("param_{:p}", Arc::as_ptr(param_ref))
}
fn initialize_slow_weights(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
for param_ref in params {
let param_key = Self::get_param_key(param_ref);
if !self.slow_weights.contains_key(¶m_key) {
let param = param_ref.read();
self.slow_weights.insert(param_key, param.clone());
}
}
Ok(())
}
fn update_slow_weights(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
for param_ref in params {
let param_key = Self::get_param_key(param_ref);
if let Some(slow_weight) = self.slow_weights.get_mut(¶m_key) {
let param = param_ref.read();
let diff = param.sub(slow_weight)?;
let update = diff.mul_scalar(self.alpha)?;
*slow_weight = slow_weight.add(&update)?;
drop(param); let mut param_mut = param_ref.write();
*param_mut = slow_weight.clone();
}
}
Ok(())
}
pub fn base_optimizer(&self) -> &O {
&self.base_optimizer
}
pub fn base_optimizer_mut(&mut self) -> &mut O {
&mut self.base_optimizer
}
pub fn alpha(&self) -> f32 {
self.alpha
}
pub fn set_alpha(&mut self, alpha: f32) {
self.alpha = alpha;
}
pub fn k(&self) -> usize {
self.k
}
pub fn set_k(&mut self, k: usize) {
self.k = k;
}
pub fn step_count(&self) -> usize {
self.step_count
}
}
impl<O: Optimizer> Optimizer for Lookahead<O> {
fn step(&mut self) -> OptimizerResult<()> {
let params = self.base_optimizer.parameters();
if params.is_empty() {
return Err(OptimizerError::StateError(
"Lookahead requires access to the base optimizer's parameters, but \
`Optimizer::parameters()` returned an empty list. The wrapped optimizer \
must implement `parameters()` to expose its parameter tensors."
.to_string(),
));
}
if self.slow_weights.is_empty() {
self.initialize_slow_weights(¶ms)?;
}
self.base_optimizer.step()?;
self.step_count += 1;
if self.step_count % self.k == 0 {
self.update_slow_weights(¶ms)?;
}
Ok(())
}
fn zero_grad(&mut self) {
self.base_optimizer.zero_grad();
}
fn get_lr(&self) -> Vec<f32> {
self.base_optimizer.get_lr()
}
fn set_lr(&mut self, lr: f32) {
self.base_optimizer.set_lr(lr);
}
fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
self.base_optimizer.add_param_group(params, options);
}
fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
self.base_optimizer.parameters()
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut base_state = self.base_optimizer.state_dict()?;
for (key, tensor) in &self.slow_weights {
let mut param_state = HashMap::new();
param_state.insert("slow_weight".to_string(), tensor.clone());
base_state
.state
.insert(format!("lookahead_{}", key), param_state);
}
let mut lookahead_options = HashMap::new();
lookahead_options.insert("alpha".to_string(), self.alpha);
lookahead_options.insert("k".to_string(), self.k as f32);
lookahead_options.insert("step_count".to_string(), self.step_count as f32);
let lookahead_group = ParamGroupState {
lr: 0.0, options: lookahead_options,
param_count: 0,
};
base_state.param_groups.push(lookahead_group);
Ok(base_state)
}
fn load_state_dict(&mut self, mut state: OptimizerState) -> OptimizerResult<()> {
if let Some(lookahead_group) = state.param_groups.pop() {
if let Some(&alpha) = lookahead_group.options.get("alpha") {
self.alpha = alpha;
}
if let Some(&k) = lookahead_group.options.get("k") {
self.k = k as usize;
}
if let Some(&step_count) = lookahead_group.options.get("step_count") {
self.step_count = step_count as usize;
}
}
self.slow_weights.clear();
let mut base_state_entries = HashMap::new();
for (key, param_state) in state.state {
if key.starts_with("lookahead_") {
if let Some(tensor) = param_state.get("slow_weight") {
let param_key = key
.strip_prefix("lookahead_")
.expect("prefix should exist after starts_with check")
.to_string();
self.slow_weights.insert(param_key, tensor.clone());
}
} else {
base_state_entries.insert(key, param_state);
}
}
let base_state = OptimizerState {
param_groups: state.param_groups,
state: base_state_entries,
global_state: state.global_state,
optimizer_type: state.optimizer_type,
version: state.version,
};
self.base_optimizer.load_state_dict(base_state)
}
}
pub fn lookahead_adam(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Lookahead<crate::adam::Adam> {
let adam = crate::adam::Adam::new(params, Some(lr), None, None, None, false);
Lookahead::with_defaults(adam)
}
pub fn lookahead_sgd(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Lookahead<crate::sgd::SGD> {
let sgd = crate::sgd::SGD::new(params, lr, None, None, None, false);
Lookahead::with_defaults(sgd)
}
pub fn lookahead_radam(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
) -> Lookahead<crate::radam::RAdam> {
let radam = crate::radam::RAdam::new(params, Some(lr), None, None, None, None);
Lookahead::with_defaults(radam)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::adam::Adam;
use parking_lot::RwLock;
use std::sync::Arc;
use torsh_tensor::creation::ones;
#[test]
fn test_lookahead_creation() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let optimizer = Lookahead::new(adam, 0.5, 5);
assert_eq!(optimizer.alpha(), 0.5);
assert_eq!(optimizer.k(), 5);
assert_eq!(optimizer.step_count(), 0);
}
#[test]
fn test_lookahead_with_defaults() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let optimizer = Lookahead::with_defaults(adam);
assert_eq!(optimizer.alpha(), 0.5);
assert_eq!(optimizer.k(), 5);
assert_eq!(optimizer.step_count(), 0);
}
#[test]
fn test_lookahead_setters() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let mut optimizer = Lookahead::new(adam, 0.5, 5);
optimizer.set_alpha(0.8);
optimizer.set_k(10);
assert_eq!(optimizer.alpha(), 0.8);
assert_eq!(optimizer.k(), 10);
}
#[test]
fn test_lookahead_base_optimizer_access() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let mut optimizer = Lookahead::new(adam, 0.5, 5);
let lr = optimizer.base_optimizer().get_lr();
assert_eq!(lr[0], 0.001);
optimizer.base_optimizer_mut().set_lr(0.002);
let new_lr = optimizer.base_optimizer().get_lr();
assert_eq!(new_lr[0], 0.002);
}
#[test]
fn test_lookahead_lr_operations() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let mut optimizer = Lookahead::new(adam, 0.5, 5);
assert_eq!(optimizer.get_lr()[0], 0.001);
optimizer.set_lr(0.002);
assert_eq!(optimizer.get_lr()[0], 0.002);
}
#[test]
fn test_lookahead_zero_grad() {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap();
p.set_grad(Some(grad));
assert!(p.grad().is_some());
}
let adam = Adam::new(vec![param.clone()], Some(0.001), None, None, None, false);
let mut optimizer = Lookahead::new(adam, 0.5, 5);
optimizer.zero_grad();
let p = param.read();
assert!(p.grad().is_none());
}
#[test]
fn test_lookahead_step_counting() {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap().mul_scalar(0.1).unwrap();
p.set_grad(Some(grad));
}
let adam = Adam::new(vec![param.clone()], Some(0.001), None, None, None, false);
let mut optimizer = Lookahead::new(adam, 0.5, 3);
optimizer.step().unwrap();
assert_eq!(optimizer.step_count(), 1);
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap().mul_scalar(0.1).unwrap();
p.set_grad(Some(grad));
}
optimizer.step().unwrap();
assert_eq!(optimizer.step_count(), 2);
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap().mul_scalar(0.1).unwrap();
p.set_grad(Some(grad));
}
optimizer.step().unwrap();
assert_eq!(optimizer.step_count(), 3);
}
#[test]
fn test_lookahead_convenience_functions() {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
let _adam_lookahead = lookahead_adam(vec![param.clone()], 0.001);
let _sgd_lookahead = lookahead_sgd(vec![param.clone()], 0.01);
let _radam_lookahead = lookahead_radam(vec![param.clone()], 0.001);
}
#[test]
fn test_lookahead_exposes_base_parameters() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let adam = Adam::new(vec![param.clone()], Some(0.001), None, None, None, false);
let optimizer = Lookahead::new(adam, 0.5, 5);
let params = optimizer.parameters();
assert_eq!(params.len(), 1);
assert!(Arc::ptr_eq(¶ms[0], ¶m));
}
#[test]
fn test_lookahead_updates_slow_weights() {
use crate::sgd::SGD;
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap();
p.set_grad(Some(grad));
}
let sgd = SGD::new(vec![param.clone()], 0.1, None, None, None, false);
let mut optimizer = Lookahead::new(sgd, 0.5, 1);
optimizer.step().unwrap();
let result = param.read().to_vec().unwrap();
for v in &result {
assert!(
(v - 0.95).abs() < 1e-5,
"expected blended slow weight 0.95, got {v}"
);
}
}
#[test]
fn test_lookahead_state_dict() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
let adam = Adam::new(vec![param], Some(0.001), None, None, None, false);
let optimizer = Lookahead::new(adam, 0.6, 7);
let state = optimizer.state_dict()?;
assert!(!state.param_groups.is_empty());
let lookahead_group = state.param_groups.last().unwrap();
assert_eq!(lookahead_group.options.get("alpha"), Some(&0.6));
assert_eq!(lookahead_group.options.get("k"), Some(&7.0));
Ok(())
}
}