use crate::optim::optimizer::Optimizer;
use crate::tensor::Tensor;
pub struct SGD {
parameters: Vec<Tensor>,
lr: f32,
}
impl SGD {
pub fn new(parameters: Vec<Tensor>, lr: f32) -> Self {
Self { parameters, lr }
}
}
impl Optimizer for SGD {
fn step(&mut self) {
for p in &self.parameters {
if let Some(grad) = &p.grad {
let grad_borrow = grad.borrow();
let update = &*grad_borrow * self.lr;
*p.data.borrow_mut() -= &update;
}
}
}
fn zero_grad(&self) {
for p in &self.parameters {
if let Some(grad) = &p.grad {
grad.borrow_mut().fill(0.0);
}
}
}
}