Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -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(())
}
}