use crate::{
Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
};
use torsh_core::error::{Result, TorshError};
use torsh_tensor::Tensor;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Clone)]
pub struct BaseOptimizer {
pub(crate) param_groups: Vec<ParamGroup>,
pub(crate) state: HashMap<String, HashMap<String, Tensor>>,
#[allow(dead_code)]
pub(crate) optimizer_type: String,
pub(crate) defaults: HashMap<String, f32>,
}
impl BaseOptimizer {
#[allow(dead_code)]
pub(crate) fn apply_weight_decay(
&self,
param: &mut Tensor,
weight_decay: f32,
) -> OptimizerResult<()> {
if weight_decay != 0.0 {
let decay = param
.mul_scalar(weight_decay)
.map_err(OptimizerError::TensorError)?;
*param = param.sub(&decay).map_err(OptimizerError::TensorError)?;
}
Ok(())
}
#[allow(dead_code)]
pub(crate) fn param_id(param: &Arc<RwLock<Tensor>>) -> String {
format!("{:p}", param.as_ref())
}
#[allow(dead_code)]
pub(crate) fn init_state(&mut self, param_id: String) {
self.state.entry(param_id).or_default();
}
#[allow(dead_code)]
pub(crate) fn get_or_create_state(
&mut self,
param_id: &str,
state_name: &str,
init_fn: impl FnOnce() -> Tensor,
) -> Tensor {
self.state
.get_mut(param_id)
.expect("state should exist for param_id")
.entry(state_name.to_string())
.or_insert_with(init_fn)
.clone()
}
#[allow(dead_code)]
pub(crate) fn update_state(&mut self, param_id: &str, state_name: &str, value: Tensor) {
self.state
.get_mut(param_id)
.expect("state should exist for param_id")
.insert(state_name.to_string(), value);
}
#[allow(dead_code)]
pub(crate) fn init_state_with_zeros(
&mut self,
param_id: String,
param: &Tensor,
state_names: &[&str],
) -> OptimizerResult<()> {
let state = self.state.entry(param_id).or_default();
for &name in state_names {
if !state.contains_key(name) {
let zeros = torsh_tensor::creation::zeros_like(param)
.map_err(OptimizerError::TensorError)?;
state.insert(name.to_string(), zeros);
}
}
Ok(())
}
#[allow(dead_code)]
pub(crate) fn init_adam_state(
&mut self,
param_id: String,
param: &Tensor,
amsgrad: bool,
) -> OptimizerResult<()> {
let state_names = if amsgrad {
vec!["step", "exp_avg", "exp_avg_sq", "max_exp_avg_sq"]
} else {
vec!["step", "exp_avg", "exp_avg_sq"]
};
self.init_state_with_zeros(param_id, param, &state_names)
}
#[allow(dead_code)]
pub(crate) fn init_sgd_state(
&mut self,
param_id: String,
param: &Tensor,
momentum: bool,
) -> OptimizerResult<()> {
let state_names = if momentum {
vec!["momentum_buffer"]
} else {
vec![]
};
if !state_names.is_empty() {
self.init_state_with_zeros(param_id, param, &state_names)
} else {
self.init_state(param_id);
Ok(())
}
}
#[allow(dead_code)]
pub(crate) fn apply_weight_decay_to_grad(
&self,
grad: &mut Tensor,
param: &Tensor,
weight_decay: f32,
) -> OptimizerResult<()> {
if weight_decay != 0.0 {
let weight_decay_term = param
.mul_scalar(weight_decay)
.map_err(OptimizerError::TensorError)?;
*grad = grad
.add_op(&weight_decay_term)
.map_err(OptimizerError::TensorError)?;
}
Ok(())
}
#[allow(dead_code)]
pub(crate) fn get_step_count(
&mut self,
param_id: &str,
increment: bool,
) -> OptimizerResult<i32> {
let state = self
.state
.get_mut(param_id)
.expect("state should exist for param_id");
let step_tensor = state.get_mut("step").expect("step state should exist");
if increment {
step_tensor
.add_scalar_(1.0)
.map_err(OptimizerError::TensorError)?;
}
let step = step_tensor.to_vec().map_err(OptimizerError::TensorError)?[0] as i32;
Ok(step)
}
#[allow(dead_code)]
pub(crate) fn compute_bias_correction(&self, betas: (f32, f32), step: i32) -> (f32, f32) {
let bias_correction1 = 1.0 - betas.0.powi(step);
let bias_correction2 = 1.0 - betas.1.powi(step);
(bias_correction1, bias_correction2)
}
#[allow(dead_code)]
pub(crate) fn update_exp_avg(
&self,
exp_avg: &mut Tensor,
grad: &Tensor,
beta: f32,
) -> OptimizerResult<()> {
exp_avg
.mul_scalar_(beta)
.map_err(OptimizerError::TensorError)?;
let grad_term = grad
.mul_scalar(1.0 - beta)
.map_err(OptimizerError::TensorError)?;
*exp_avg = exp_avg
.add(&grad_term)
.map_err(OptimizerError::TensorError)?;
Ok(())
}
#[allow(dead_code)]
pub(crate) fn update_exp_avg_sq(
&self,
exp_avg_sq: &mut Tensor,
grad: &Tensor,
beta: f32,
) -> OptimizerResult<()> {
exp_avg_sq
.mul_scalar_(beta)
.map_err(OptimizerError::TensorError)?;
let grad_squared = grad.mul_op(grad).map_err(OptimizerError::TensorError)?;
let grad_sq_term = grad_squared
.mul_scalar(1.0 - beta)
.map_err(OptimizerError::TensorError)?;
*exp_avg_sq = exp_avg_sq
.add(&grad_sq_term)
.map_err(OptimizerError::TensorError)?;
Ok(())
}
#[allow(dead_code)]
pub(crate) fn clip_gradient(&self, grad: &mut Tensor, max_norm: f32) -> OptimizerResult<f32> {
let norm = grad.norm().map_err(OptimizerError::TensorError)?;
let norm_value = norm.to_vec().map_err(OptimizerError::TensorError)?[0];
if norm_value > max_norm {
let scale = max_norm / norm_value;
*grad = grad
.mul_scalar(scale)
.map_err(OptimizerError::TensorError)?;
}
Ok(norm_value)
}
#[allow(dead_code)]
pub(crate) fn validate_gradients(&self) -> bool {
self.param_groups
.iter()
.all(|group| group.params.iter().all(|param| param.read().has_grad()))
}
pub(crate) fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
collect_parameters(&self.param_groups)
}
}
pub(crate) fn collect_parameters(param_groups: &[ParamGroup]) -> Vec<Arc<RwLock<Tensor>>> {
param_groups
.iter()
.flat_map(|group| group.params.iter().cloned())
.collect()
}
impl Optimizer for BaseOptimizer {
fn step(&mut self) -> OptimizerResult<()> {
Err(OptimizerError::TensorError(TorshError::Other(
"Optimizer step not yet implemented - scirs2 integration pending".to_string(),
)))
}
fn zero_grad(&mut self) {
for group in &self.param_groups {
for param in &group.params {
param.write().zero_grad();
}
}
}
fn get_lr(&self) -> Vec<f32> {
self.param_groups.iter().map(|g| g.lr).collect()
}
fn set_lr(&mut self, lr: f32) {
for group in &mut self.param_groups {
group.lr = lr;
}
}
fn add_param_group(
&mut self,
params: Vec<Arc<RwLock<Tensor>>>,
mut options: HashMap<String, f32>,
) {
let lr = options
.remove("lr")
.unwrap_or_else(|| self.defaults.get("lr").copied().unwrap_or(1e-3));
let mut group = ParamGroup::new(params, lr);
group.options = options;
self.param_groups.push(group);
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
self.create_state_dict(None)
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
state
.validate()
.map_err(|e| OptimizerError::StateError(e.to_string()))?;
if state.param_groups.len() != self.param_groups.len() {
return Err(OptimizerError::StateError(
"Loaded state dict has different number of parameter groups".to_string(),
));
}
for (i, (group, state_group)) in self
.param_groups
.iter()
.zip(state.param_groups.iter())
.enumerate()
{
if group.params.len() != state_group.param_count {
return Err(OptimizerError::StateError(format!(
"Parameter count mismatch in group {}: expected {}, got {}",
i,
group.params.len(),
state_group.param_count
)));
}
}
for (group, state_group) in self.param_groups.iter_mut().zip(state.param_groups.iter()) {
group.lr = state_group.lr;
group.options = state_group.options.clone();
}
self.state = state.state;
for (key, value) in state.global_state {
self.defaults.insert(key, value);
}
Ok(())
}
}
impl BaseOptimizer {
#[allow(dead_code)]
pub(crate) fn create_state_dict(
&self,
additional_global_state: Option<HashMap<String, f32>>,
) -> OptimizerResult<OptimizerState> {
let param_groups = self
.param_groups
.iter()
.map(|g| ParamGroupState::from_param_group(g))
.collect();
let mut optimizer_state = OptimizerState::new(self.optimizer_type.clone());
optimizer_state.param_groups = param_groups;
optimizer_state.state = self.state.clone();
for (key, value) in &self.defaults {
optimizer_state.global_state.insert(key.clone(), *value);
}
if let Some(additional) = additional_global_state {
for (key, value) in additional {
optimizer_state.global_state.insert(key, value);
}
}
Ok(optimizer_state)
}
}
pub mod functional {
use super::*;
pub fn clip_grad_before_step<O: Optimizer>(
_optimizer: &O,
max_norm: Option<f32>,
_norm_type: f32,
) -> f32 {
if let Some(_max_norm) = max_norm {
0.0
} else {
0.0
}
}
}