use crate::{lookahead::Lookahead, radam::RAdam};
use parking_lot::RwLock;
use std::sync::Arc;
use torsh_tensor::Tensor;
pub type Ranger = Lookahead<RAdam>;
impl Ranger {
pub fn new_ranger(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
let radam = RAdam::new(params, Some(lr), None, None, None, None);
Lookahead::with_defaults(radam)
}
#[allow(clippy::too_many_arguments)]
pub fn with_radam_params(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
beta1: f32,
beta2: f32,
eps: f32,
weight_decay: f32,
) -> Self {
let radam = RAdam::new(
params,
Some(lr),
Some(beta1),
Some(beta2),
Some(eps),
Some(weight_decay),
);
Lookahead::with_defaults(radam)
}
pub fn with_lookahead_params(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
alpha: f32,
k: usize,
) -> Self {
let radam = RAdam::new(params, Some(lr), None, None, None, None);
Lookahead::new(radam, alpha, k)
}
#[allow(clippy::too_many_arguments)]
pub fn with_all_params(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
beta1: f32,
beta2: f32,
eps: f32,
weight_decay: f32,
alpha: f32,
k: usize,
) -> Self {
let radam = RAdam::new(
params,
Some(lr),
Some(beta1),
Some(beta2),
Some(eps),
Some(weight_decay),
);
Lookahead::new(radam, alpha, k)
}
pub fn radam(&self) -> &RAdam {
self.base_optimizer()
}
pub fn radam_mut(&mut self) -> &mut RAdam {
self.base_optimizer_mut()
}
pub fn beta1(&self) -> f32 {
self.radam().beta1()
}
pub fn beta2(&self) -> f32 {
self.radam().beta2()
}
pub fn eps(&self) -> f32 {
self.radam().eps()
}
pub fn weight_decay(&self) -> f32 {
self.radam().weight_decay()
}
pub fn set_beta1(&mut self, beta1: f32) {
self.radam_mut().set_beta1(beta1);
}
pub fn set_beta2(&mut self, beta2: f32) {
self.radam_mut().set_beta2(beta2);
}
pub fn set_eps(&mut self, eps: f32) {
self.radam_mut().set_eps(eps);
}
pub fn set_weight_decay(&mut self, weight_decay: f32) {
self.radam_mut().set_weight_decay(weight_decay);
}
pub fn ranger21(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
let radam = RAdam::new(
params,
Some(lr),
Some(0.95), Some(0.999), Some(1e-7), Some(0.0), );
Lookahead::new(radam, 0.8, 6) }
pub fn with_warm_restarts(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
let radam = RAdam::new(
params,
Some(lr),
Some(0.9), Some(0.999), Some(1e-6), Some(0.0), );
Lookahead::new(radam, 0.6, 4) }
pub fn set_rectification(&mut self, enabled: bool) {
self.radam_mut().set_rectification(enabled);
}
pub fn is_rectification_enabled(&self) -> bool {
self.radam().is_rectification_enabled()
}
pub fn rectification_coefficient(&self) -> Option<f32> {
self.radam().rectification_coefficient()
}
}
pub struct RangerBuilder {
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
beta1: f32,
beta2: f32,
eps: f32,
weight_decay: f32,
alpha: f32,
k: usize,
}
impl RangerBuilder {
pub fn new(params: Vec<Arc<RwLock<Tensor>>>) -> Self {
Self {
params,
lr: 1e-3,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
weight_decay: 0.0,
alpha: 0.5,
k: 5,
}
}
pub fn lr(mut self, lr: f32) -> Self {
self.lr = lr;
self
}
pub fn beta1(mut self, beta1: f32) -> Self {
self.beta1 = beta1;
self
}
pub fn beta2(mut self, beta2: f32) -> Self {
self.beta2 = beta2;
self
}
pub fn eps(mut self, eps: f32) -> Self {
self.eps = eps;
self
}
pub fn weight_decay(mut self, weight_decay: f32) -> Self {
self.weight_decay = weight_decay;
self
}
pub fn alpha(mut self, alpha: f32) -> Self {
self.alpha = alpha;
self
}
pub fn k(mut self, k: usize) -> Self {
self.k = k;
self
}
pub fn build(self) -> Ranger {
Ranger::with_all_params(
self.params,
self.lr,
self.beta1,
self.beta2,
self.eps,
self.weight_decay,
self.alpha,
self.k,
)
}
}
impl Ranger {
pub fn builder(params: Vec<Arc<RwLock<Tensor>>>) -> RangerBuilder {
RangerBuilder::new(params)
}
pub fn for_vision(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
Self::with_all_params(
params, lr, 0.9, 0.999, 1e-8, 1e-4, 0.5, 5, )
}
pub fn for_nlp(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
Self::with_all_params(
params, lr, 0.9, 0.999, 1e-8, 0.01, 0.6, 4, )
}
pub fn for_rl(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> Self {
Self::with_all_params(
params, lr, 0.95, 0.999, 1e-6, 0.0, 0.8, 3, )
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{Optimizer, OptimizerResult};
use parking_lot::RwLock;
use std::sync::Arc;
use torsh_tensor::creation::ones;
#[test]
fn test_ranger_creation() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::new_ranger(vec![param], 0.001);
assert_eq!(optimizer.get_lr()[0], 0.001);
assert_eq!(optimizer.alpha(), 0.5);
assert_eq!(optimizer.k(), 5);
assert_eq!(optimizer.beta1(), 0.9);
assert_eq!(optimizer.beta2(), 0.999);
}
#[test]
fn test_ranger_with_radam_params() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::with_radam_params(vec![param], 0.002, 0.95, 0.9999, 1e-7, 0.01);
assert_eq!(optimizer.get_lr()[0], 0.002);
assert_eq!(optimizer.beta1(), 0.95);
assert_eq!(optimizer.beta2(), 0.9999);
assert_eq!(optimizer.eps(), 1e-7);
assert_eq!(optimizer.weight_decay(), 0.01);
assert_eq!(optimizer.alpha(), 0.5);
assert_eq!(optimizer.k(), 5);
}
#[test]
fn test_ranger_with_lookahead_params() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::with_lookahead_params(vec![param], 0.003, 0.8, 6);
assert_eq!(optimizer.get_lr()[0], 0.003);
assert_eq!(optimizer.alpha(), 0.8);
assert_eq!(optimizer.k(), 6);
assert_eq!(optimizer.beta1(), 0.9);
assert_eq!(optimizer.beta2(), 0.999);
}
#[test]
fn test_ranger_with_all_params() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer =
Ranger::with_all_params(vec![param], 0.005, 0.95, 0.9999, 1e-7, 0.02, 0.7, 8);
assert_eq!(optimizer.get_lr()[0], 0.005);
assert_eq!(optimizer.beta1(), 0.95);
assert_eq!(optimizer.beta2(), 0.9999);
assert_eq!(optimizer.eps(), 1e-7);
assert_eq!(optimizer.weight_decay(), 0.02);
assert_eq!(optimizer.alpha(), 0.7);
assert_eq!(optimizer.k(), 8);
}
#[test]
fn test_ranger21() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::ranger21(vec![param], 0.001);
assert_eq!(optimizer.get_lr()[0], 0.001);
assert_eq!(optimizer.beta1(), 0.95);
assert_eq!(optimizer.beta2(), 0.999);
assert_eq!(optimizer.eps(), 1e-7);
assert_eq!(optimizer.alpha(), 0.8);
assert_eq!(optimizer.k(), 6);
}
#[test]
fn test_ranger_builder() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::builder(vec![param])
.lr(0.002)
.beta1(0.95)
.beta2(0.9999)
.eps(1e-7)
.weight_decay(0.01)
.alpha(0.6)
.k(4)
.build();
assert_eq!(optimizer.get_lr()[0], 0.002);
assert_eq!(optimizer.beta1(), 0.95);
assert_eq!(optimizer.beta2(), 0.9999);
assert_eq!(optimizer.eps(), 1e-7);
assert_eq!(optimizer.weight_decay(), 0.01);
assert_eq!(optimizer.alpha(), 0.6);
assert_eq!(optimizer.k(), 4);
}
#[test]
fn test_ranger_setters() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let mut optimizer = Ranger::new_ranger(vec![param], 0.001);
optimizer.set_beta1(0.95);
optimizer.set_beta2(0.9999);
optimizer.set_eps(1e-7);
optimizer.set_weight_decay(0.02);
optimizer.set_alpha(0.7);
optimizer.set_k(7);
assert_eq!(optimizer.beta1(), 0.95);
assert_eq!(optimizer.beta2(), 0.9999);
assert_eq!(optimizer.eps(), 1e-7);
assert_eq!(optimizer.weight_decay(), 0.02);
assert_eq!(optimizer.alpha(), 0.7);
assert_eq!(optimizer.k(), 7);
}
#[test]
fn test_ranger_domain_specific() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let vision_opt = Ranger::for_vision(vec![param.clone()], 0.001);
assert_eq!(vision_opt.weight_decay(), 1e-4);
let nlp_opt = Ranger::for_nlp(vec![param.clone()], 0.001);
assert_eq!(nlp_opt.weight_decay(), 0.01);
assert_eq!(nlp_opt.alpha(), 0.6);
assert_eq!(nlp_opt.k(), 4);
let rl_opt = Ranger::for_rl(vec![param.clone()], 0.001);
assert_eq!(rl_opt.beta1(), 0.95);
assert_eq!(rl_opt.alpha(), 0.8);
assert_eq!(rl_opt.k(), 3);
}
#[test]
fn test_ranger_step() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
{
let mut p = param.write();
let grad = ones(&[2, 2]).unwrap().mul_scalar(0.1)?;
p.set_grad(Some(grad));
}
let mut optimizer = Ranger::new_ranger(vec![param.clone()], 0.01);
optimizer.step()?;
assert_eq!(optimizer.step_count(), 1);
Ok(())
}
#[test]
fn test_ranger_rectification() {
let param = Arc::new(RwLock::new(ones(&[2, 2]).unwrap()));
let mut optimizer = Ranger::new_ranger(vec![param], 0.01);
assert!(optimizer.is_rectification_enabled());
optimizer.set_rectification(false);
assert!(optimizer.is_rectification_enabled());
optimizer.set_rectification(true);
assert!(optimizer.is_rectification_enabled());
}
#[test]
fn test_ranger_warm_restarts() {
let param = Arc::new(RwLock::new(ones(&[2, 3]).unwrap()));
let optimizer = Ranger::with_warm_restarts(vec![param], 0.001);
assert_eq!(optimizer.eps(), 1e-6);
assert_eq!(optimizer.alpha(), 0.6);
assert_eq!(optimizer.k(), 4);
}
}