use candle_core::{DType, Device, Result, Tensor, Var};
pub const ADAGRAD_EPS: f64 = 1e-10;
pub struct RowAdagrad {
acc: Tensor,
lr: f64,
eps: f64,
full_count: Option<Tensor>,
}
impl RowAdagrad {
pub fn new(n_rows: usize, lr: f64, dev: &Device) -> Result<Self> {
Ok(Self {
acc: Tensor::zeros(n_rows, DType::F32, dev)?,
lr,
eps: ADAGRAD_EPS,
full_count: None,
})
}
fn advance(&mut self, g2: &Tensor) -> Result<Tensor> {
self.acc = (&self.acc + g2)?.detach();
self.acc
.sqrt()?
.affine(1.0 / self.lr, self.eps / self.lr)?
.recip()
}
pub fn step(&mut self, var: &Var, grad: &Tensor) -> Result<()> {
let row_sq = grad.sqr()?.sum(1)?;
self.step_with_row_sq(var, grad, &row_sq)
}
pub fn step_with_row_sq(&mut self, var: &Var, grad: &Tensor, row_sq: &Tensor) -> Result<()> {
let grad = grad.detach();
let h = var.dims()[1] as f64;
let step = self.advance(&row_sq.detach().affine(1.0 / h, 0.0)?)?;
var.set(
&var.as_tensor()
.sub(&grad.broadcast_mul(&step.unsqueeze(1)?)?)?,
)
}
pub fn step_with_bias(
&mut self,
row: &Var,
bias: &Var,
g_row: &Tensor,
g_bias: &Tensor,
row_mask: Option<&Tensor>,
decay: f64,
) -> Result<()> {
let h = row.dims()[1] as f64;
let g_row = match row_mask {
None => g_row.detach(),
Some(m) => g_row.detach().broadcast_mul(m)?,
};
let g_bias = g_bias.detach();
let row_sq = g_row.sqr()?.sum(1)?; let count = match row_mask {
Some(m) => m.squeeze(1)?.affine(h, 1.0)?,
None => match self.full_count.as_ref() {
Some(c) => c.clone(),
None => {
let c = Tensor::full((h + 1.0) as f32, row_sq.dims()[0], row_sq.device())?;
self.full_count = Some(c.clone());
c
}
},
};
let g2 = (&row_sq + g_bias.sqr()?)?.div(&count)?;
let step = self.advance(&g2)?;
let mut new_row = row.as_tensor().clone();
if decay != 1.0 {
let touched = row_sq.gt(0f32)?.to_dtype(DType::F32)?;
let factor = touched.affine(decay - 1.0, 1.0)?; new_row = new_row.broadcast_mul(&factor.unsqueeze(1)?)?;
}
let new_row = new_row.sub(&g_row.broadcast_mul(&step.unsqueeze(1)?)?)?;
row.set(&new_row)?;
bias.set(&bias.as_tensor().sub(&(g_bias * step)?)?)
}
pub fn step_bias(&mut self, bias: &Var, g_bias: &Tensor) -> Result<()> {
let g = g_bias.detach();
let step = self.advance(&g.sqr()?)?;
bias.set(&bias.as_tensor().sub(&(g * step)?)?)
}
#[must_use]
pub fn accumulator(&self) -> &Tensor {
&self.acc
}
}
#[cfg(test)]
#[path = "optim_tests.rs"]
mod optim_tests;