use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroupState};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use torsh_core::error::{Result, TorshError};
use torsh_tensor::Tensor;
pub struct ElasticAveragingSGD {
params: Vec<Arc<RwLock<Tensor>>>,
center_params: Vec<Tensor>,
lr: f32,
momentum: Option<f32>,
weight_decay: Option<f32>,
rho: f32,
communication_freq: usize,
step_counter: usize,
momentum_buffers: HashMap<String, Tensor>,
worker_rank: usize,
num_workers: usize,
}
impl ElasticAveragingSGD {
#[allow(clippy::too_many_arguments)]
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
momentum: Option<f32>,
weight_decay: Option<f32>,
rho: f32,
communication_freq: usize,
worker_rank: usize,
num_workers: usize,
) -> Result<Self> {
if rho <= 0.0 || rho > 1.0 {
return Err(TorshError::InvalidArgument(
"Elastic parameter rho must be in (0, 1]".to_string(),
));
}
if communication_freq == 0 {
return Err(TorshError::InvalidArgument(
"Communication frequency must be greater than 0".to_string(),
));
}
let center_params = params.iter().map(|p| p.read().clone()).collect();
Ok(Self {
params,
center_params,
lr,
momentum,
weight_decay,
rho,
communication_freq,
step_counter: 0,
momentum_buffers: HashMap::new(),
worker_rank,
num_workers,
})
}
pub fn new_default(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
worker_rank: usize,
num_workers: usize,
) -> Result<Self> {
Self::new(
params,
lr,
Some(0.9), Some(1e-4), 0.1, 10, worker_rank,
num_workers,
)
}
fn local_update(&mut self) -> Result<()> {
for (i, param_arc) in self.params.iter().enumerate() {
let mut param = param_arc.write();
let grad = param
.grad()
.ok_or_else(|| TorshError::AutogradError("No gradient available".to_string()))?;
let mut effective_grad = grad.clone();
if let Some(decay) = self.weight_decay {
effective_grad = effective_grad.add(¶m.mul_scalar(decay)?)?;
}
if let Some(momentum) = self.momentum {
let param_key = format!("param_{}", i);
if let Some(buf) = self.momentum_buffers.get(¶m_key) {
let new_buf = buf.mul_scalar(momentum)?.add(&effective_grad)?;
effective_grad = new_buf.clone();
self.momentum_buffers.insert(param_key, new_buf);
} else {
self.momentum_buffers
.insert(param_key, effective_grad.clone());
}
}
let elastic_force = param.sub(&self.center_params[i])?.mul_scalar(self.rho)?;
effective_grad = effective_grad.add(&elastic_force)?;
let update = effective_grad.mul_scalar(self.lr)?;
*param = param.sub(&update)?;
}
Ok(())
}
pub fn communicate(&mut self, all_worker_params: &[Vec<Tensor>]) -> Result<()> {
if all_worker_params.is_empty() {
return Ok(());
}
for i in 0..self.center_params.len() {
if i >= all_worker_params[0].len() {
break; }
let mut sum = all_worker_params[0][i].clone();
for worker_params in all_worker_params.iter().skip(1) {
if i < worker_params.len() {
sum = sum.add(&worker_params[i])?;
}
}
self.center_params[i] = sum.div_scalar(self.num_workers as f32)?;
}
Ok(())
}
pub fn get_worker_params(&self) -> Vec<Tensor> {
self.params.iter().map(|p| p.read().clone()).collect()
}
pub fn get_center_params(&self) -> &[Tensor] {
&self.center_params
}
pub fn worker_rank(&self) -> usize {
self.worker_rank
}
pub fn rho(&self) -> f32 {
self.rho
}
pub fn set_rho(&mut self, rho: f32) -> Result<()> {
if rho <= 0.0 || rho > 1.0 {
return Err(TorshError::InvalidArgument(
"Elastic parameter rho must be in (0, 1]".to_string(),
));
}
self.rho = rho;
Ok(())
}
pub fn communication_freq(&self) -> usize {
self.communication_freq
}
pub fn set_communication_freq(&mut self, freq: usize) -> Result<()> {
if freq == 0 {
return Err(TorshError::InvalidArgument(
"Communication frequency must be greater than 0".to_string(),
));
}
self.communication_freq = freq;
Ok(())
}
pub fn should_communicate(&self) -> bool {
self.step_counter % self.communication_freq == 0
}
pub fn step_counter(&self) -> usize {
self.step_counter
}
pub fn reset_step_counter(&mut self) {
self.step_counter = 0;
}
}
impl Optimizer for ElasticAveragingSGD {
fn step(&mut self) -> OptimizerResult<()> {
self.local_update()?;
self.step_counter += 1;
Ok(())
}
fn zero_grad(&mut self) {
for param in &self.params {
param.write().zero_grad();
}
}
fn get_lr(&self) -> Vec<f32> {
vec![self.lr]
}
fn set_lr(&mut self, lr: f32) {
self.lr = lr;
}
fn add_param_group(
&mut self,
params: Vec<Arc<RwLock<Tensor>>>,
_options: HashMap<String, f32>,
) {
for param in ¶ms {
self.center_params.push(param.read().clone());
}
self.params.extend(params);
}
fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
self.params.clone()
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let param_group = ParamGroupState {
lr: self.lr,
options: [
("momentum".to_string(), self.momentum.unwrap_or(0.0)),
("weight_decay".to_string(), self.weight_decay.unwrap_or(0.0)),
("rho".to_string(), self.rho),
(
"communication_freq".to_string(),
self.communication_freq as f32,
),
("worker_rank".to_string(), self.worker_rank as f32),
("num_workers".to_string(), self.num_workers as f32),
]
.iter()
.cloned()
.collect(),
param_count: self.params.len(),
};
let mut state = HashMap::new();
for (param_id, momentum_buffer) in &self.momentum_buffers {
let mut param_state = HashMap::new();
param_state.insert("momentum_buffer".to_string(), momentum_buffer.clone());
state.insert(param_id.clone(), param_state);
}
let mut global_state = HashMap::new();
global_state.insert("step_counter".to_string(), self.step_counter as f32);
Ok(OptimizerState {
optimizer_type: "EASGD".to_string(),
version: "0.1.0".to_string(),
param_groups: vec![param_group],
state,
global_state,
})
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
if state.optimizer_type != "EASGD" {
return Err(OptimizerError::InvalidParameter(format!(
"Expected EASGD optimizer state, got {}",
state.optimizer_type
)));
}
if let Some(param_group) = state.param_groups.first() {
self.lr = param_group.lr;
if let Some(&momentum) = param_group.options.get("momentum") {
self.momentum = if momentum > 0.0 { Some(momentum) } else { None };
}
if let Some(&weight_decay) = param_group.options.get("weight_decay") {
self.weight_decay = if weight_decay > 0.0 {
Some(weight_decay)
} else {
None
};
}
if let Some(&rho) = param_group.options.get("rho") {
self.rho = rho;
}
if let Some(&communication_freq) = param_group.options.get("communication_freq") {
self.communication_freq = communication_freq as usize;
}
if let Some(&worker_rank) = param_group.options.get("worker_rank") {
self.worker_rank = worker_rank as usize;
}
if let Some(&num_workers) = param_group.options.get("num_workers") {
self.num_workers = num_workers as usize;
}
if param_group.param_count != self.params.len() {
return Err(OptimizerError::InvalidParameter(format!(
"Parameter count mismatch: expected {}, got {}",
self.params.len(),
param_group.param_count
)));
}
} else {
return Err(OptimizerError::InvalidParameter(
"No parameter groups found in state".to_string(),
));
}
self.momentum_buffers.clear();
for (param_id, param_state) in state.state {
if let Some(momentum_buffer) = param_state.get("momentum_buffer") {
self.momentum_buffers
.insert(param_id, momentum_buffer.clone());
}
}
if let Some(&step_counter) = state.global_state.get("step_counter") {
self.step_counter = step_counter as usize;
}
self.center_params.clear();
for param in &self.params {
self.center_params.push(param.read().clone());
}
Ok(())
}
}
pub mod utils {
use super::*;
pub fn create_easgd_for_cluster_size(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
worker_rank: usize,
num_workers: usize,
) -> Result<ElasticAveragingSGD> {
let (rho, comm_freq) = match num_workers {
1..=2 => (0.2, 5), 3..=8 => (0.1, 10), 9..=16 => (0.05, 20), _ => (0.02, 50), };
ElasticAveragingSGD::new(
params,
lr,
Some(0.9), Some(1e-4), rho,
comm_freq,
worker_rank,
num_workers,
)
}
pub fn simulate_easgd_round(
workers: &mut [ElasticAveragingSGD],
steps_per_round: usize,
) -> Result<()> {
for worker in workers.iter_mut() {
for _ in 0..steps_per_round {
worker.step()?;
}
}
let all_worker_params: Vec<Vec<Tensor>> = workers
.iter()
.map(|worker| worker.get_worker_params())
.collect();
for worker in workers.iter_mut() {
worker.communicate(&all_worker_params)?;
}
Ok(())
}
pub fn calculate_easgd_metrics(workers: &[ElasticAveragingSGD]) -> HashMap<String, f32> {
let mut metrics = HashMap::new();
if workers.is_empty() {
return metrics;
}
let mut total_divergence = 0.0;
let mut param_count = 0;
if !workers.is_empty() && !workers[0].params.is_empty() {
let num_params = workers[0].params.len();
for param_idx in 0..num_params {
let mut param_values = Vec::new();
for worker in workers {
if param_idx < worker.params.len() {
let param = worker.params[param_idx].read();
if let Ok(value) = param.get(&[0]) {
param_values.push(value);
}
}
}
if param_values.len() > 1 {
let mean = param_values.iter().sum::<f32>() / param_values.len() as f32;
let variance = param_values.iter().map(|v| (v - mean).powi(2)).sum::<f32>()
/ param_values.len() as f32;
total_divergence += variance.sqrt();
param_count += 1;
}
}
}
if param_count > 0 {
metrics.insert(
"average_parameter_divergence".to_string(),
total_divergence / param_count as f32,
);
}
metrics.insert("num_workers".to_string(), workers.len() as f32);
if let Some(first_worker) = workers.first() {
metrics.insert("rho".to_string(), first_worker.rho());
metrics.insert(
"communication_freq".to_string(),
first_worker.communication_freq() as f32,
);
metrics.insert(
"step_counter".to_string(),
first_worker.step_counter() as f32,
);
}
metrics
}
}