Skip to main content

ruda_optim/optim/
gradient_transform.rs

1use alloc::collections::BTreeSet;
2use core::fmt;
3use ruda_model::{
4    module::{AutodiffModule, ModuleVisitor, Param, ParamId},
5    tensor::{DType, FloatDType, Tensor, TensorMetadata, backend::AutodiffBackend},
6};
7use super::GradientsParams;
8
9/// Invalid gradient metadata or an explicit scaling argument.
10#[derive(Clone, Debug, PartialEq, Eq)]
11pub enum GradientTransformError {
12    /// The multiplier/divisor is not finite or cannot be represented in the work dtype.
13    InvalidScalar,
14    /// The requested arithmetic dtype is not F32 or F64.
15    InvalidDType,
16    /// A gradient does not have its parameter's geometry or device.
17    ParameterMismatch(ParamId),
18    /// A gradient is quantized; arithmetic must not silently dequantize it.
19    QuantizedGradient(ParamId),
20    /// At least one registered gradient is not owned by the supplied module.
21    UnknownParameters,
22}
23
24impl fmt::Display for GradientTransformError {
25    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
26        match self {
27            Self::InvalidScalar => f.write_str("gradient scale must be representable and finite; a divisor must be positive"),
28            Self::InvalidDType => f.write_str("gradient work dtype must be F32 or F64"),
29            Self::ParameterMismatch(id) => write!(f,"gradient geometry/device does not match parameter {id}"),
30            Self::QuantizedGradient(id) => write!(f,"quantized gradient for parameter {id} cannot be scaled implicitly"),
31            Self::UnknownParameters => f.write_str("gradient container includes parameters outside the supplied module"),
32        }
33    }
34}
35
36#[cfg(feature = "std")]
37impl std::error::Error for GradientTransformError {}
38
39pub(super) fn validate_work_dtype(dtype: FloatDType) -> Result<(), GradientTransformError> {
40    if matches!(dtype,FloatDType::F32 | FloatDType::F64) { Ok(()) }
41    else { Err(GradientTransformError::InvalidDType) }
42}
43
44pub(super) fn representable(value: f64, dtype: FloatDType) -> bool {
45    value.is_finite() && (dtype == FloatDType::F64 || (value as f32).is_finite())
46}
47
48struct Metadata<'a> {
49    gradients: &'a GradientsParams,
50    seen: BTreeSet<ParamId>,
51    matched: usize,
52    error: Option<GradientTransformError>,
53}
54
55impl<B: AutodiffBackend> ModuleVisitor<B> for Metadata<'_> {
56    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B,D>>) {
57        if self.error.is_some() { return; }
58        let Some(primitive) = self.gradients.container.get::<B::InnerBackend>(&param.id) else { return; };
59        if self.seen.insert(param.id) { self.matched += 1; }
60        if matches!(primitive.dtype(),DType::QFloat(_)) {
61            self.error = Some(GradientTransformError::QuantizedGradient(param.id));
62            return;
63        }
64        if primitive.rank() != D {
65            self.error = Some(GradientTransformError::ParameterMismatch(param.id));
66            return;
67        }
68        let gradient = Tensor::<B::InnerBackend,D>::from_primitive(primitive);
69        let parameter = param.val();
70        if gradient.dims() != parameter.dims() || gradient.device() != parameter.device() {
71            self.error = Some(GradientTransformError::ParameterMismatch(param.id));
72        }
73    }
74}
75
76impl GradientsParams {
77    /// Check module membership, actual dimensions and device without reading values.
78    /// Tied parameter IDs are counted once; absent gradients stay absent.
79    pub fn validate_for<B: AutodiffBackend,M: AutodiffModule<B>>(
80        &self, module: &M,
81    ) -> Result<(),GradientTransformError> {
82        let mut visitor = Metadata { gradients:self,seen:BTreeSet::new(),matched:0,error:None };
83        module.visit(&mut visitor);
84        if let Some(error) = visitor.error { return Err(error); }
85        if visitor.matched != self.len() { return Err(GradientTransformError::UnknownParameters); }
86        Ok(())
87    }
88
89    /// Copy handles and optionally cast all present module gradients to F32/F64.
90    /// This does not modify parameter storage, IDs or the source container.
91    pub fn cast_for<B: AutodiffBackend,M: AutodiffModule<B>>(
92        &self, module: &M, dtype: FloatDType,
93    ) -> Result<Self,GradientTransformError> {
94        self.transformed::<B,M>(module,dtype,1.,false)
95    }
96
97    /// Multiply each present gradient in an explicit work dtype, once per tied ID.
98    /// Negative and zero multipliers are intentional caller-selected transforms.
99    pub fn scaled_for<B: AutodiffBackend,M: AutodiffModule<B>>(
100        &self, module: &M, multiplier: f64, dtype: FloatDType,
101    ) -> Result<Self,GradientTransformError> {
102        self.transformed::<B,M>(module,dtype,multiplier,false)
103    }
104
105    /// Divide by a positive finite scale in F32/F64 without clearing the source.
106    /// No clipping, nonfinite-step policy, optimizer update or dynamic scaling.
107    pub fn unscaled_for<B: AutodiffBackend,M: AutodiffModule<B>>(
108        &self, module: &M, divisor: f64, dtype: FloatDType,
109    ) -> Result<Self,GradientTransformError> {
110        self.transformed::<B,M>(module,dtype,divisor,true)
111    }
112
113    fn transformed<B: AutodiffBackend,M: AutodiffModule<B>>(
114        &self, module: &M, dtype: FloatDType, scalar: f64, divide: bool,
115    ) -> Result<Self,GradientTransformError> {
116        validate_work_dtype(dtype)?;
117        if !representable(scalar,dtype) || (divide && (scalar <= 0. ||
118            (dtype == FloatDType::F32 && scalar as f32 == 0.))) {
119            return Err(GradientTransformError::InvalidScalar);
120        }
121        self.validate_for::<B,M>(module)?;
122        let mut visitor = Transform {source:self,result:Self::new(),seen:BTreeSet::new(),dtype,scalar,divide};
123        module.visit(&mut visitor);
124        Ok(visitor.result)
125    }
126}
127
128struct Transform<'a> {
129    source: &'a GradientsParams,
130    result: GradientsParams,
131    seen: BTreeSet<ParamId>,
132    dtype: FloatDType,
133    scalar: f64,
134    divide: bool,
135}
136
137impl<B: AutodiffBackend> ModuleVisitor<B> for Transform<'_> {
138    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B,D>>) {
139        if !self.seen.insert(param.id) { return; }
140        let Some(gradient) = self.source.get::<B::InnerBackend,D>(param.id) else { return; };
141        let gradient = gradient.cast(self.dtype);
142        let value = if self.scalar == 1. { gradient }
143            else if self.divide { gradient.div_scalar(self.scalar) }
144            else { gradient.mul_scalar(self.scalar) };
145        self.result.register::<B::InnerBackend,D>(param.id,value);
146    }
147}