use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use torsh_tensor::Tensor;
#[derive(Debug, Clone)]
pub struct EWCConfig {
pub importance: f32,
pub fisher_sample_size: usize,
pub diagonal_fisher: bool,
}
impl Default for EWCConfig {
fn default() -> Self {
Self {
importance: 1000.0,
fisher_sample_size: 200,
diagonal_fisher: true,
}
}
}
pub struct EWCOptimizer<O: Optimizer> {
base_optimizer: O,
config: EWCConfig,
fisher_information: HashMap<String, Tensor>,
optimal_params: HashMap<String, Tensor>,
current_task: usize,
param_groups: Vec<Arc<RwLock<Tensor>>>,
}
impl<O: Optimizer> EWCOptimizer<O> {
pub fn new(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
config: EWCConfig,
) -> OptimizerResult<Self> {
Ok(Self {
base_optimizer,
config,
fisher_information: HashMap::new(),
optimal_params: HashMap::new(),
current_task: 0,
param_groups: params,
})
}
pub fn with_defaults(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
) -> OptimizerResult<Self> {
Self::new(base_optimizer, params, EWCConfig::default())
}
pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
self.optimal_params
.insert(param_key.clone(), param_read.clone());
}
self.compute_fisher_diagonal()?;
self.current_task += 1;
Ok(())
}
fn compute_fisher_diagonal(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
if let Some(grad) = param_read.grad() {
let fisher = grad.mul(&grad)?;
if let Some(existing_fisher) = self.fisher_information.get(¶m_key) {
let accumulated = existing_fisher.add(&fisher)?;
self.fisher_information.insert(param_key, accumulated);
} else {
self.fisher_information.insert(param_key, fisher);
}
}
}
Ok(())
}
fn apply_ewc_penalty(&mut self) -> OptimizerResult<()> {
if self.optimal_params.is_empty() {
return Ok(());
}
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
if let (Some(fisher), Some(optimal)) = (
self.fisher_information.get(¶m_key),
self.optimal_params.get(¶m_key),
) {
let mut param_write = param.write();
let diff = param_write.sub(optimal)?;
let penalty = fisher.mul(&diff)?;
let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
if let Some(grad) = param_write.grad() {
let new_grad = grad.add(&scaled_penalty)?;
param_write.set_grad(Some(new_grad));
}
}
}
Ok(())
}
}
impl<O: Optimizer> Optimizer for EWCOptimizer<O> {
fn step(&mut self) -> OptimizerResult<()> {
self.apply_ewc_penalty()?;
self.base_optimizer.step()
}
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.param_groups.clone()
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = self.base_optimizer.state_dict()?;
state.optimizer_type = format!("EWC({})", state.optimizer_type);
state
.global_state
.insert("current_task".to_string(), self.current_task as f32);
state
.global_state
.insert("importance".to_string(), self.config.importance);
Ok(state)
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
self.base_optimizer.load_state_dict(state)
}
}
#[derive(Debug, Clone)]
pub struct SIConfig {
pub damping: f32,
pub importance: f32,
}
impl Default for SIConfig {
fn default() -> Self {
Self {
damping: 0.1,
importance: 1.0,
}
}
}
pub struct SIOptimizer<O: Optimizer> {
base_optimizer: O,
config: SIConfig,
path_integral: HashMap<String, Tensor>,
prev_params: HashMap<String, Tensor>,
importance: HashMap<String, Tensor>,
current_task: usize,
param_groups: Vec<Arc<RwLock<Tensor>>>,
}
impl<O: Optimizer> SIOptimizer<O> {
pub fn new(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
config: SIConfig,
) -> OptimizerResult<Self> {
let mut prev_params = HashMap::new();
let mut path_integral = HashMap::new();
for (i, param) in params.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
prev_params.insert(param_key.clone(), param_read.clone());
let shape_owned = param_read.shape().dims().to_vec();
drop(param_read);
let zeros = torsh_tensor::creation::zeros(&shape_owned)?;
path_integral.insert(param_key, zeros);
}
Ok(Self {
base_optimizer,
config,
path_integral,
prev_params,
importance: HashMap::new(),
current_task: 0,
param_groups: params,
})
}
pub fn with_defaults(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
) -> OptimizerResult<Self> {
Self::new(base_optimizer, params, SIConfig::default())
}
fn update_path_integral(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
if let Some(grad) = param_read.grad() {
if let Some(prev_param) = self.prev_params.get(¶m_key) {
let delta = param_read.sub(prev_param)?;
let contribution = grad.mul(&delta)?;
let neg_contribution = contribution.mul_scalar(-1.0)?;
if let Some(omega) = self.path_integral.get_mut(¶m_key) {
*omega = omega.add(&neg_contribution)?;
}
self.prev_params.insert(param_key, param_read.clone());
}
}
}
Ok(())
}
pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
if let Some(omega) = self.path_integral.get(¶m_key) {
if let Some(prev_param) = self.prev_params.get(¶m_key) {
let param_read = param.read();
let delta = param_read.sub(prev_param)?;
let delta_sq = delta.mul(&delta)?;
let denom = delta_sq.add_scalar(self.config.damping)?;
let task_importance = omega.div(&denom)?;
if let Some(existing) = self.importance.get(¶m_key) {
let accumulated = existing.add(&task_importance)?;
self.importance.insert(param_key.clone(), accumulated);
} else {
self.importance.insert(param_key.clone(), task_importance);
}
let shape_owned = param_read.shape().dims().to_vec();
let zeros = torsh_tensor::creation::zeros(&shape_owned)?;
self.path_integral.insert(param_key, zeros);
}
}
}
self.current_task += 1;
Ok(())
}
fn apply_si_penalty(&mut self) -> OptimizerResult<()> {
if self.importance.is_empty() {
return Ok(());
}
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
if let (Some(importance), Some(prev_param)) = (
self.importance.get(¶m_key),
self.prev_params.get(¶m_key),
) {
let mut param_write = param.write();
let diff = param_write.sub(prev_param)?;
let penalty = importance.mul(&diff)?;
let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
if let Some(grad) = param_write.grad() {
let new_grad = grad.add(&scaled_penalty)?;
param_write.set_grad(Some(new_grad));
}
}
}
Ok(())
}
}
impl<O: Optimizer> Optimizer for SIOptimizer<O> {
fn step(&mut self) -> OptimizerResult<()> {
self.update_path_integral()?;
self.apply_si_penalty()?;
self.base_optimizer.step()
}
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.param_groups.clone()
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = self.base_optimizer.state_dict()?;
state.optimizer_type = format!("SI({})", state.optimizer_type);
state
.global_state
.insert("current_task".to_string(), self.current_task as f32);
Ok(state)
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
self.base_optimizer.load_state_dict(state)
}
}
#[derive(Debug, Clone)]
pub struct MASConfig {
pub importance: f32,
pub n_samples: usize,
}
impl Default for MASConfig {
fn default() -> Self {
Self {
importance: 1.0,
n_samples: 100,
}
}
}
pub struct MASOptimizer<O: Optimizer> {
base_optimizer: O,
config: MASConfig,
importance: HashMap<String, Tensor>,
optimal_params: HashMap<String, Tensor>,
current_task: usize,
param_groups: Vec<Arc<RwLock<Tensor>>>,
}
impl<O: Optimizer> MASOptimizer<O> {
pub fn new(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
config: MASConfig,
) -> OptimizerResult<Self> {
Ok(Self {
base_optimizer,
config,
importance: HashMap::new(),
optimal_params: HashMap::new(),
current_task: 0,
param_groups: params,
})
}
pub fn with_defaults(
base_optimizer: O,
params: Vec<Arc<RwLock<Tensor>>>,
) -> OptimizerResult<Self> {
Self::new(base_optimizer, params, MASConfig::default())
}
pub fn compute_importance(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
if let Some(grad) = param_read.grad() {
let grad_abs = grad.abs()?;
if let Some(existing) = self.importance.get(¶m_key) {
let accumulated = existing.add(&grad_abs)?;
self.importance.insert(param_key, accumulated);
} else {
self.importance.insert(param_key, grad_abs);
}
}
}
Ok(())
}
pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
self.optimal_params.insert(param_key, param_read.clone());
}
self.current_task += 1;
Ok(())
}
fn apply_mas_penalty(&mut self) -> OptimizerResult<()> {
if self.importance.is_empty() {
return Ok(());
}
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
if let (Some(importance), Some(optimal)) = (
self.importance.get(¶m_key),
self.optimal_params.get(¶m_key),
) {
let mut param_write = param.write();
let diff = param_write.sub(optimal)?;
let penalty = importance.mul(&diff)?;
let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
if let Some(grad) = param_write.grad() {
let new_grad = grad.add(&scaled_penalty)?;
param_write.set_grad(Some(new_grad));
}
}
}
Ok(())
}
}
impl<O: Optimizer> Optimizer for MASOptimizer<O> {
fn step(&mut self) -> OptimizerResult<()> {
self.apply_mas_penalty()?;
self.base_optimizer.step()
}
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.param_groups.clone()
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = self.base_optimizer.state_dict()?;
state.optimizer_type = format!("MAS({})", state.optimizer_type);
state
.global_state
.insert("current_task".to_string(), self.current_task as f32);
Ok(state)
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
self.base_optimizer.load_state_dict(state)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sgd::SGD;
use torsh_tensor::creation::randn;
#[test]
fn test_ewc_config_default() {
let config = EWCConfig::default();
assert_eq!(config.importance, 1000.0);
assert_eq!(config.fisher_sample_size, 200);
assert!(config.diagonal_fisher);
}
#[test]
fn test_ewc_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
assert_eq!(optimizer.current_task, 0);
Ok(())
}
#[test]
fn test_ewc_consolidate_task() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[5, 5])?;
p.set_grad(Some(grad));
}
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let mut optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
optimizer.consolidate_task()?;
assert_eq!(optimizer.current_task, 1);
assert!(!optimizer.optimal_params.is_empty());
Ok(())
}
#[test]
fn test_si_config_default() {
let config = SIConfig::default();
assert_eq!(config.damping, 0.1);
assert_eq!(config.importance, 1.0);
}
#[test]
fn test_si_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let optimizer = SIOptimizer::with_defaults(base, vec![param])?;
assert_eq!(optimizer.current_task, 0);
Ok(())
}
#[test]
fn test_si_consolidate_task() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[3, 3])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[3, 3])?;
p.set_grad(Some(grad));
}
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let mut optimizer = SIOptimizer::with_defaults(base, vec![param])?;
optimizer.consolidate_task()?;
assert_eq!(optimizer.current_task, 1);
Ok(())
}
#[test]
fn test_mas_config_default() {
let config = MASConfig::default();
assert_eq!(config.importance, 1.0);
assert_eq!(config.n_samples, 100);
}
#[test]
fn test_mas_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let optimizer = MASOptimizer::with_defaults(base, vec![param])?;
assert_eq!(optimizer.current_task, 0);
Ok(())
}
#[test]
fn test_mas_compute_importance() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[5, 5])?;
p.set_grad(Some(grad));
}
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let mut optimizer = MASOptimizer::with_defaults(base, vec![param])?;
optimizer.compute_importance()?;
assert!(!optimizer.importance.is_empty());
Ok(())
}
#[test]
fn test_ewc_step() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[2, 2])?;
p.set_grad(Some(grad));
}
let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
let mut optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
optimizer.step()?;
Ok(())
}
}