use candle_core::test_utils::{to_vec0_round, to_vec2_round};
use anyhow::Result;
use candle_core::{Device, Tensor, Var};
use candle_nn::{Linear, Module, Optimizer};
use candle_optimisers::{
adam::{Adam, ParamsAdam},
Decay,
};
#[test]
fn adam_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam::default();
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.9000, 0.6967]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.7996);
Ok(())
}
#[test]
fn adam_weight_decay_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam {
weight_decay: Some(Decay::WeightDecay(0.6)),
..Default::default()
};
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.8997, 0.6964]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.7975);
Ok(())
}
#[test]
fn adamw_weight_decay_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam {
weight_decay: Some(Decay::DecoupledWeightDecay(0.6)),
..Default::default()
};
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.6901, 0.5677]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.6287);
Ok(())
}
#[test]
fn adam_amsgrad_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam {
amsgrad: true,
..Default::default()
};
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.9001, 0.6904]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.7978);
Ok(())
}
#[test]
fn adam_amsgrad_decay_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam {
amsgrad: true,
weight_decay: Some(Decay::WeightDecay(0.6)),
..Default::default()
};
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.8998, 0.6901]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.7955);
Ok(())
}
#[test]
fn adamw_amsgrad_decay_test() -> Result<()> {
let w_gen = Tensor::new(&[[3f32, 1.]], &Device::Cpu)?;
let b_gen = Tensor::new(-2f32, &Device::Cpu)?;
let gen = Linear::new(w_gen, Some(b_gen));
let sample_xs = Tensor::new(&[[2f32, 1.], [7., 4.], [-4., 12.], [5., 8.]], &Device::Cpu)?;
let sample_ys = gen.forward(&sample_xs)?;
let params = ParamsAdam {
weight_decay: Some(Decay::DecoupledWeightDecay(0.6)),
amsgrad: true,
..Default::default()
};
let w = Var::new(&[[0f32, 0.]], &Device::Cpu)?;
let b = Var::new(0f32, &Device::Cpu)?;
let mut n_sgd = Adam::new(vec![w.clone(), b.clone()], params)?;
let lin = Linear::new(w.as_tensor().clone(), Some(b.as_tensor().clone()));
for _step in 0..1000 {
let ys = lin.forward(&sample_xs)?;
let loss = ys.sub(&sample_ys)?.sqr()?.sum_all()?;
n_sgd.backward_step(&loss)?;
}
assert_eq!(to_vec2_round(&w, 4)?, &[[0.6901, 0.5648]]);
assert_eq!(to_vec0_round(&b, 4)?, 0.6287);
Ok(())
}