ruda-optim 0.21.57

Ruda optimizer updates, gradient state, clipping and learning-rate schedules.

use super::GradientsParams;
use ruda_model::module::{AutodiffModule, ModuleVisitor, Param, ParamId};
use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
use core::marker::PhantomData;
use hashbrown::HashSet;

#[cfg(not(feature = "std"))]
use alloc::vec::Vec;

pub struct GradientsParamsConverter<'a, M: AutodiffModule<B>, B: AutodiffBackend> {
    grads: &'a mut B::Gradients,
    grads_params: &'a mut GradientsParams,
    phatom: PhantomData<M>,
    filter: Option<HashSet<ParamId>>,
}

impl<'a, M: AutodiffModule<B>, B: AutodiffBackend> GradientsParamsConverter<'a, M, B> {
    /// Create a gradient visitor with an optional parameter-ID selection.
    pub fn new(
        grads: &'a mut B::Gradients,
        grads_params: &'a mut GradientsParams,
        filter: Option<Vec<ParamId>>,
    ) -> Self {
        Self {
            grads,
            grads_params,
            phatom: PhantomData,
            filter: filter.map(|ids| ids.into_iter().collect()),
        }
    }
}

#[derive(new)]
pub struct GradientsParamsChangeDevice<'a, M: AutodiffModule<B>, B: AutodiffBackend> {
    device: &'a B::Device,
    grads: &'a mut GradientsParams,
    phatom: PhantomData<M>,
}

impl<B, M> ModuleVisitor<B> for GradientsParamsConverter<'_, M, B>
where
    B: AutodiffBackend,
    M: AutodiffModule<B>,
{
    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
        if let Some(filter) = self.filter.as_ref()
            && !filter.contains(&param.id)
        {
            return;
        }

        let Some(grad) = param.val().grad_remove(self.grads) else {
            return;
        };

        let grad = match self.grads_params.remove::<B::InnerBackend, D>(param.id) {
            Some(previous) => previous.add(grad),
            None => grad,
        };
        self.grads_params
            .register::<B::InnerBackend, D>(param.id, grad);
    }
}

impl<B, M> ModuleVisitor<B> for GradientsParamsChangeDevice<'_, M, B>
where
    B: AutodiffBackend,
    M: AutodiffModule<B>,
{
    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
        let Some(grad) = self.grads.remove::<B::InnerBackend, D>(param.id) else {
            return;
        };

        self.grads
            .register::<B::InnerBackend, D>(param.id, grad.to_device(self.device));
    }
}