Initial commit
This commit is contained in:
@@ -0,0 +1,543 @@
|
||||
//! NovoGrad optimizer implementation with layer-wise gradient normalization
|
||||
//!
|
||||
//! Implements the NovoGrad optimization algorithm with:
|
||||
//! - Layer-wise gradient normalization
|
||||
//! - Adaptive second moment averaging
|
||||
//! - Momentum with normalized gradients
|
||||
//! - Decoupled weight decay
|
||||
//! - Bias correction for early iterations
|
||||
//! - Gradient clipping support
|
||||
//!
|
||||
//! # Algorithm
|
||||
//!
|
||||
//! NovoGrad normalizes gradients layer-wise and uses adaptive learning rates:
|
||||
//! 1. Compute gradient norm per layer: g_l = ||∇L||
|
||||
//! 2. Normalize gradient: g_norm = g / g_l
|
||||
//! 3. Compute second moment: v_t = β2 * v_{t-1} + (1-β2) * g_l²
|
||||
//! 4. Momentum update: m_t = β1 * m_{t-1} + (lr / √(v_t + ε)) * g_norm
|
||||
//! 5. Weight update: θ_t+1 = θ_t - m_t - λ * θ_t (with weight decay)
|
||||
//!
|
||||
//! # References
|
||||
//! - Ginsburg et al. "Stochastic Gradient Methods with Layer-wise Adaptive Moments"
|
||||
//! - Particularly effective for large batch training and transformer models
|
||||
|
||||
use crate::{Result, TransformerError};
|
||||
use crate::optimizers::{Optimizer, BaseOptimizer};
|
||||
use rtx_tensor::{Tensor, TensorError, DType};
|
||||
use std::collections::HashMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, trace};
|
||||
|
||||
/// NovoGrad optimizer configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct NovoGradConfig {
|
||||
/// Learning rate
|
||||
pub learning_rate: f64,
|
||||
/// Beta1 parameter (momentum decay rate)
|
||||
pub beta1: f64,
|
||||
/// Beta2 parameter (second moment decay rate)
|
||||
pub beta2: f64,
|
||||
/// Epsilon for numerical stability
|
||||
pub eps: f64,
|
||||
/// Weight decay coefficient (decoupled)
|
||||
pub weight_decay: f64,
|
||||
/// Whether to use gradient averaging (optional)
|
||||
pub grad_averaging: bool,
|
||||
}
|
||||
|
||||
impl Default for NovoGradConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
learning_rate: 1e-3,
|
||||
beta1: 0.95,
|
||||
beta2: 0.98,
|
||||
eps: 1e-8,
|
||||
weight_decay: 0.0,
|
||||
grad_averaging: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// NovoGrad optimizer state for a single parameter
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct NovoGradState {
|
||||
/// Momentum accumulation (normalized gradient with adaptive scaling)
|
||||
pub momentum: Tensor,
|
||||
/// Second moment estimate (gradient norm squared)
|
||||
pub second_moment: f64,
|
||||
/// Step count for bias correction
|
||||
pub step: i64,
|
||||
}
|
||||
|
||||
/// NovoGrad optimizer with layer-wise gradient normalization
|
||||
///
|
||||
/// NovoGrad addresses gradient explosion/vanishing problems by normalizing
|
||||
/// gradients layer-wise and using adaptive second moment averaging.
|
||||
///
|
||||
/// # Mathematical Foundation
|
||||
///
|
||||
/// For each layer/parameter:
|
||||
/// 1. Gradient norm: g_l = ||g_t||_2
|
||||
/// 2. Normalized gradient: ĝ_t = g_t / g_l
|
||||
/// 3. Second moment: v_t = β₂ * v_{t-1} + (1 - β₂) * g_l²
|
||||
/// 4. Effective learning rate: α_eff = α / √(v̂_t + ε)
|
||||
/// 5. Momentum: m_t = β₁ * m_{t-1} + α_eff * ĝ_t
|
||||
/// 6. Parameter update: θ_t = θ_{t-1} - m_t - λ * θ_{t-1}
|
||||
///
|
||||
/// With bias correction:
|
||||
/// - v̂_t = v_t / (1 - β₂^t)
|
||||
///
|
||||
/// # Performance Characteristics
|
||||
/// - Particularly effective for large batch training
|
||||
/// - Helps with gradient explosion/vanishing in deep networks
|
||||
/// - Layer-wise adaptive learning rates
|
||||
/// - Memory efficient (one momentum buffer per parameter)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct NovoGradOptimizer {
|
||||
/// Base optimizer functionality
|
||||
base: BaseOptimizer,
|
||||
/// Beta1 parameter (momentum decay rate)
|
||||
beta1: f64,
|
||||
/// Beta2 parameter (second moment decay rate)
|
||||
beta2: f64,
|
||||
/// Epsilon for numerical stability
|
||||
eps: f64,
|
||||
/// Weight decay coefficient
|
||||
weight_decay: f64,
|
||||
/// Whether to use gradient averaging
|
||||
grad_averaging: bool,
|
||||
/// Per-parameter state
|
||||
state: HashMap<String, NovoGradState>,
|
||||
}
|
||||
|
||||
impl NovoGradOptimizer {
|
||||
/// Create a new NovoGrad optimizer from configuration
|
||||
pub fn new(config: NovoGradConfig) -> Result<Self> {
|
||||
// Validate parameters
|
||||
Self::validate_parameters(
|
||||
config.learning_rate,
|
||||
config.beta1,
|
||||
config.beta2,
|
||||
config.eps,
|
||||
config.weight_decay,
|
||||
)?;
|
||||
|
||||
debug!(
|
||||
"Creating NovoGrad optimizer: lr={}, β₁={}, β₂={}, ε={}, weight_decay={}, grad_avg={}",
|
||||
config.learning_rate, config.beta1, config.beta2, config.eps,
|
||||
config.weight_decay, config.grad_averaging
|
||||
);
|
||||
|
||||
Ok(Self {
|
||||
base: BaseOptimizer::new(config.learning_rate),
|
||||
beta1: config.beta1,
|
||||
beta2: config.beta2,
|
||||
eps: config.eps,
|
||||
weight_decay: config.weight_decay,
|
||||
grad_averaging: config.grad_averaging,
|
||||
state: HashMap::new(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Validate optimizer parameters
|
||||
fn validate_parameters(
|
||||
learning_rate: f64,
|
||||
beta1: f64,
|
||||
beta2: f64,
|
||||
eps: f64,
|
||||
weight_decay: f64,
|
||||
) -> Result<()> {
|
||||
if learning_rate <= 0.0 {
|
||||
return Err(TransformerError::generic(
|
||||
format!("learning_rate {} must be positive", learning_rate)
|
||||
));
|
||||
}
|
||||
|
||||
if !(0.0..1.0).contains(&beta1) {
|
||||
return Err(TransformerError::generic(
|
||||
format!("beta1 {} must be in [0, 1)", beta1)
|
||||
));
|
||||
}
|
||||
|
||||
if !(0.0..1.0).contains(&beta2) {
|
||||
return Err(TransformerError::generic(
|
||||
format!("beta2 {} must be in [0, 1)", beta2)
|
||||
));
|
||||
}
|
||||
|
||||
if eps < 0.0 {
|
||||
return Err(TransformerError::generic(
|
||||
format!("eps {} must be non-negative", eps)
|
||||
));
|
||||
}
|
||||
|
||||
if weight_decay < 0.0 {
|
||||
return Err(TransformerError::generic(
|
||||
format!("weight_decay {} must be non-negative", weight_decay)
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get beta1 parameter
|
||||
pub fn beta1(&self) -> f64 {
|
||||
self.beta1
|
||||
}
|
||||
|
||||
/// Get beta2 parameter
|
||||
pub fn beta2(&self) -> f64 {
|
||||
self.beta2
|
||||
}
|
||||
|
||||
/// Get epsilon parameter
|
||||
pub fn eps(&self) -> f64 {
|
||||
self.eps
|
||||
}
|
||||
|
||||
/// Get weight decay parameter
|
||||
pub fn weight_decay(&self) -> f64 {
|
||||
self.weight_decay
|
||||
}
|
||||
|
||||
/// Get gradient averaging setting
|
||||
pub fn grad_averaging(&self) -> bool {
|
||||
self.grad_averaging
|
||||
}
|
||||
|
||||
/// Compute gradient norm for layer-wise normalization
|
||||
fn compute_gradient_norm(&self, grad: &Tensor) -> Result<f64> {
|
||||
// Compute L2 norm of the gradient tensor
|
||||
let grad_data = grad.to_cpu()?;
|
||||
let norm_squared: f64 = grad_data.iter()
|
||||
.map(|&x| (x as f64) * (x as f64))
|
||||
.sum();
|
||||
Ok(norm_squared.sqrt())
|
||||
}
|
||||
|
||||
/// Initialize state for a parameter if it doesn't exist
|
||||
fn ensure_state(&mut self, param_name: &str, param: &Tensor) -> Result<()> {
|
||||
if !self.state.contains_key(param_name) {
|
||||
let shape = param.shape();
|
||||
let device = param.device();
|
||||
|
||||
let momentum = Tensor::zeros(shape.clone(), device)?;
|
||||
|
||||
let state = NovoGradState {
|
||||
momentum,
|
||||
second_moment: 0.0,
|
||||
step: 0,
|
||||
};
|
||||
|
||||
self.state.insert(param_name.to_string(), state);
|
||||
trace!("Initialized NovoGrad state for parameter: {}", param_name);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Perform parameter update with layer-wise gradient normalization
|
||||
fn update_parameter_static(
|
||||
param: &Tensor,
|
||||
grad: &Tensor,
|
||||
state: &mut NovoGradState,
|
||||
learning_rate: f64,
|
||||
beta1: f64,
|
||||
beta2: f64,
|
||||
eps: f64,
|
||||
weight_decay: f64,
|
||||
grad_averaging: bool,
|
||||
) -> Result<Tensor> {
|
||||
// Increment step count
|
||||
state.step += 1;
|
||||
let step = state.step as f64;
|
||||
|
||||
trace!("NovoGrad step {} for parameter", step);
|
||||
|
||||
// 1. Compute gradient norm per layer: g_l = ||∇L||
|
||||
let grad_data = grad.to_cpu()?;
|
||||
let grad_norm_squared: f64 = grad_data.iter()
|
||||
.map(|&x| (x as f64) * (x as f64))
|
||||
.sum();
|
||||
let grad_norm = grad_norm_squared.sqrt();
|
||||
|
||||
// Handle zero gradient case
|
||||
if grad_norm < eps {
|
||||
trace!("Zero or very small gradient norm: {}, applying only weight decay", grad_norm);
|
||||
|
||||
// Apply weight decay if configured
|
||||
let final_param = if weight_decay > 0.0 {
|
||||
// Decoupled weight decay: θ_t+1 = θ_t - λ * θ_t
|
||||
param * (1.0 - learning_rate * weight_decay)
|
||||
} else {
|
||||
param.clone()
|
||||
}?;
|
||||
|
||||
return Ok(final_param);
|
||||
}
|
||||
|
||||
// 2. Normalize gradient: g_norm = g / g_l
|
||||
let normalized_grad = grad / (grad_norm as f32)?;
|
||||
|
||||
// 3. Update second moment: v_t = β2 * v_{t-1} + (1-β2) * g_l²
|
||||
state.second_moment = beta2 * state.second_moment + (1.0 - beta2) * grad_norm_squared;
|
||||
|
||||
// 4. Bias correction for second moment
|
||||
let bias_corrected_second_moment = state.second_moment / (1.0 - beta2.powf(step));
|
||||
|
||||
// 5. Compute effective learning rate: α_eff = α / √(v̂_t + ε)
|
||||
let effective_lr = learning_rate / (bias_corrected_second_moment.sqrt() + eps);
|
||||
|
||||
// 6. Update momentum: m_t = β1 * m_{t-1} + α_eff * g_norm
|
||||
let momentum_term1 = (&state.momentum * beta1)?;
|
||||
let momentum_term2 = (&normalized_grad * (effective_lr as f32))?;
|
||||
state.momentum = (momentum_term1 + momentum_term2)?;
|
||||
|
||||
// Apply gradient averaging if enabled
|
||||
let final_momentum = if grad_averaging {
|
||||
// Scale momentum by step count for averaging effect
|
||||
(&state.momentum / (step as f32))?
|
||||
} else {
|
||||
state.momentum.clone()
|
||||
};
|
||||
|
||||
// 7. Apply weight update: θ_t+1 = θ_t - m_t - λ * θ_t
|
||||
let weight_decay_term = if weight_decay > 0.0 {
|
||||
// Decoupled weight decay
|
||||
param * (learning_rate * weight_decay as f32)?
|
||||
} else {
|
||||
Tensor::zeros(param.shape().clone(), param.device())?
|
||||
};
|
||||
|
||||
let final_param = ((param - &final_momentum)? - &weight_decay_term)?;
|
||||
|
||||
Ok(final_param)
|
||||
}
|
||||
}
|
||||
|
||||
impl Optimizer for NovoGradOptimizer {
|
||||
fn step_param(&mut self, param_name: &str, param: &Tensor, grad: &Tensor) -> Result<Tensor> {
|
||||
// Validate input shapes match
|
||||
if param.shape() != grad.shape() {
|
||||
return Err(TransformerError::shape_mismatch(
|
||||
format!("Parameter shape {:?} doesn't match gradient shape {:?}",
|
||||
param.shape().dims().to_vec(),
|
||||
grad.shape().dims().to_vec())
|
||||
));
|
||||
}
|
||||
|
||||
// Ensure state exists
|
||||
self.ensure_state(param_name, param)?;
|
||||
|
||||
// Get the state and perform update in separate scope
|
||||
let state = self.state.get_mut(param_name).unwrap();
|
||||
let learning_rate = self.base.learning_rate;
|
||||
let beta1 = self.beta1;
|
||||
let beta2 = self.beta2;
|
||||
let eps = self.eps;
|
||||
let weight_decay = self.weight_decay;
|
||||
let grad_averaging = self.grad_averaging;
|
||||
|
||||
// Perform update using copied values
|
||||
Self::update_parameter_static(
|
||||
param, grad, state, learning_rate, beta1, beta2, eps, weight_decay, grad_averaging
|
||||
)
|
||||
}
|
||||
|
||||
fn learning_rate(&self) -> f64 {
|
||||
self.base.learning_rate
|
||||
}
|
||||
|
||||
fn set_learning_rate(&mut self, lr: f64) -> Result<()> {
|
||||
if lr <= 0.0 {
|
||||
return Err(TransformerError::generic(
|
||||
format!("learning_rate {} must be positive", lr)
|
||||
));
|
||||
}
|
||||
|
||||
self.base.set_learning_rate(lr);
|
||||
debug!("Updated NovoGrad learning rate to: {}", lr);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn has_state(&self, param_name: &str) -> bool {
|
||||
self.state.contains_key(param_name)
|
||||
}
|
||||
|
||||
fn reset_state(&mut self, param_name: &str) -> Result<()> {
|
||||
if self.state.remove(param_name).is_some() {
|
||||
debug!("Reset NovoGrad state for parameter: {}", param_name);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn reset_all_state(&mut self) {
|
||||
let count = self.state.len();
|
||||
self.state.clear();
|
||||
debug!("Reset all NovoGrad state ({} parameters)", count);
|
||||
}
|
||||
|
||||
fn get_step_count(&self, param_name: &str) -> Result<i64> {
|
||||
self.state
|
||||
.get(param_name)
|
||||
.map(|state| state.step)
|
||||
.ok_or_else(|| {
|
||||
TransformerError::optimizer(format!("No state found for parameter: {}", param_name))
|
||||
})
|
||||
}
|
||||
|
||||
fn optimizer_type(&self) -> &'static str {
|
||||
"NovoGrad"
|
||||
}
|
||||
|
||||
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
|
||||
self
|
||||
}
|
||||
|
||||
fn store_gradients_internal(&mut self, gradients: HashMap<String, Tensor>) -> Result<()> {
|
||||
debug!("Storing {} gradients for NovoGrad optimizer", gradients.len());
|
||||
self.base.store_gradients(gradients);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn process_stored_gradients(&mut self, learning_rate: f64) -> Result<HashMap<String, Tensor>> {
|
||||
debug!("Processing {} stored gradients with lr={} using NovoGrad algorithm",
|
||||
self.base.stored_gradients.len(), learning_rate);
|
||||
|
||||
// Update learning rate if different
|
||||
if (self.learning_rate() - learning_rate).abs() > f64::EPSILON {
|
||||
self.set_learning_rate(learning_rate)?;
|
||||
}
|
||||
|
||||
let mut parameter_updates = HashMap::new();
|
||||
|
||||
// Process each stored gradient with full NovoGrad algorithm
|
||||
for (param_name, grad) in &self.base.stored_gradients {
|
||||
// Initialize state if it doesn't exist
|
||||
if !self.state.contains_key(param_name) {
|
||||
debug!("Creating NovoGrad state for parameter: {}", param_name);
|
||||
let shape = grad.shape();
|
||||
let device = grad.device();
|
||||
|
||||
let momentum = Tensor::zeros(shape.clone(), device)?;
|
||||
|
||||
let state = NovoGradState {
|
||||
momentum,
|
||||
second_moment: 0.0,
|
||||
step: 0,
|
||||
};
|
||||
|
||||
self.state.insert(param_name.clone(), state);
|
||||
}
|
||||
|
||||
// Get mutable reference to state
|
||||
let state = self.state.get_mut(param_name).unwrap();
|
||||
|
||||
// Increment step count
|
||||
state.step += 1;
|
||||
let step = state.step as f64;
|
||||
|
||||
trace!("NovoGrad step {} for parameter {}", step, param_name);
|
||||
|
||||
// 1. Compute gradient norm per layer: g_l = ||∇L||
|
||||
let grad_data = grad.to_cpu()?;
|
||||
let grad_norm_squared: f64 = grad_data.iter()
|
||||
.map(|&x| (x as f64) * (x as f64))
|
||||
.sum();
|
||||
let grad_norm = grad_norm_squared.sqrt();
|
||||
|
||||
// Handle zero gradient case
|
||||
if grad_norm < self.eps {
|
||||
trace!("Zero or very small gradient norm: {} for parameter {}", grad_norm, param_name);
|
||||
|
||||
// Apply only weight decay if configured
|
||||
let weight_decay_update = if self.weight_decay > 0.0 {
|
||||
// Note: We don't have access to parameter values here, so this is a limitation
|
||||
debug!("Weight decay configured but parameter values not available in process_stored_gradients for {}", param_name);
|
||||
Tensor::zeros(grad.shape().clone(), grad.device())?
|
||||
} else {
|
||||
Tensor::zeros(grad.shape().clone(), grad.device())?
|
||||
};
|
||||
|
||||
parameter_updates.insert(param_name.clone(), weight_decay_update);
|
||||
continue;
|
||||
}
|
||||
|
||||
// 2. Normalize gradient: g_norm = g / g_l
|
||||
let normalized_grad = (grad / (grad_norm as f32))?;
|
||||
|
||||
// 3. Update second moment: v_t = β2 * v_{t-1} + (1-β2) * g_l²
|
||||
state.second_moment = self.beta2 * state.second_moment + (1.0 - self.beta2) * grad_norm_squared;
|
||||
|
||||
// 4. Bias correction for second moment
|
||||
let bias_corrected_second_moment = state.second_moment / (1.0 - self.beta2.powf(step));
|
||||
|
||||
// 5. Compute effective learning rate: α_eff = α / √(v̂_t + ε)
|
||||
let effective_lr = learning_rate / (bias_corrected_second_moment.sqrt() + self.eps);
|
||||
|
||||
// 6. Update momentum: m_t = β1 * m_{t-1} + α_eff * g_norm
|
||||
let momentum_term1 = (&state.momentum * self.beta1)?;
|
||||
let momentum_term2 = (&normalized_grad * (effective_lr as f32))?;
|
||||
state.momentum = (momentum_term1 + momentum_term2)?;
|
||||
|
||||
// Apply gradient averaging if enabled
|
||||
let final_momentum = if self.grad_averaging {
|
||||
// Scale momentum by step count for averaging effect
|
||||
(&state.momentum / (step as f32))?
|
||||
} else {
|
||||
state.momentum.clone()
|
||||
};
|
||||
|
||||
// Create parameter update (negative because we want to subtract)
|
||||
let novograd_update = (final_momentum * -1.0)?;
|
||||
|
||||
// Note: Weight decay handling is limited without access to current parameter values
|
||||
if self.weight_decay > 0.0 {
|
||||
debug!("Weight decay configured but parameter values not available in process_stored_gradients for {}", param_name);
|
||||
}
|
||||
|
||||
parameter_updates.insert(param_name.clone(), novograd_update);
|
||||
debug!("Created NovoGrad update for parameter: {} (step {})", param_name, step);
|
||||
}
|
||||
|
||||
// Clear stored gradients after processing
|
||||
self.base.clear_gradients();
|
||||
|
||||
debug!("Generated {} parameter updates using NovoGrad algorithm", parameter_updates.len());
|
||||
Ok(parameter_updates)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "disabled_tests"))]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rtx_tensor::{Device, Shape};
|
||||
|
||||
#[test]
|
||||
fn test_novograd_parameter_validation() {
|
||||
let config = NovoGradConfig::default();
|
||||
|
||||
// Test valid parameters
|
||||
assert!(NovoGradOptimizer::new(config).is_ok());
|
||||
|
||||
// Test invalid learning rate
|
||||
let mut invalid_config = NovoGradConfig::default();
|
||||
invalid_config.learning_rate = -0.1;
|
||||
assert!(NovoGradOptimizer::new(invalid_config).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_novograd_gradient_norm_computation() -> Result<()> {
|
||||
let config = NovoGradConfig::default();
|
||||
let optimizer = NovoGradOptimizer::new(config)?;
|
||||
let device = Device::cpu();
|
||||
|
||||
// Test gradient with known norm
|
||||
let grad = Tensor::from_data(vec![3.0, 4.0], &[2], &device)?; // ||grad|| = 5.0
|
||||
let norm = optimizer.compute_gradient_norm(&grad)?;
|
||||
|
||||
assert!((norm - 5.0).abs() < 1e-6, "Expected norm 5.0, got {}", norm);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user