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#[derive(Clone, Debug, PartialEq, Eq)]
11pub enum GradientTransformError {
12 InvalidScalar,
14 InvalidDType,
16 ParameterMismatch(ParamId),
18 QuantizedGradient(ParamId),
20 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>(¶m.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 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 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 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 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}