Skip to main content

mlx_rs/nn/
value_and_grad.rs

1use crate::module::{update_parameters, ModuleParameters};
2use crate::transforms::keyed_value_and_grad;
3use crate::{error::Exception, Array};
4
5use crate::module::FlattenedModuleParam;
6
7fn trainable_params(model: &impl ModuleParameters) -> FlattenedModuleParam {
8    model
9        .trainable_parameters()
10        .flatten()
11        .into_iter()
12        .map(|(k, v)| (k, v.clone()))
13        .collect()
14}
15
16/// Helper trait for [`value_and_grad`]
17pub trait IntoModuleValueAndGrad<'a, M, Args, Val, Err>
18where
19    M: ModuleParameters + 'a,
20    Args: Clone,
21{
22    /// Computes the valud and gradient of the passed function `f(model, args)` with regard to the
23    /// model's trainable parameters.
24    fn into_module_value_and_grad(
25        self,
26    ) -> impl FnMut(&mut M, Args) -> Result<(Val, FlattenedModuleParam), Exception> + 'a;
27}
28
29impl<'a, F, M, Args> IntoModuleValueAndGrad<'a, M, Args, Vec<Array>, ()> for F
30where
31    M: ModuleParameters + 'a,
32    F: FnMut(&mut M, Args) -> Vec<Array> + 'a,
33    Args: Clone,
34{
35    fn into_module_value_and_grad(
36        mut self,
37    ) -> impl FnMut(&mut M, Args) -> Result<(Vec<Array>, FlattenedModuleParam), Exception> + 'a
38    {
39        move |model, arrays| {
40            let trainable_parameters = trainable_params(model);
41            let inner = |parameters: FlattenedModuleParam, arrays: Args| -> Vec<Array> {
42                let flattened_parameters = parameters.into_iter();
43                update_parameters(model, flattened_parameters);
44
45                self(model, arrays)
46            };
47            let mut vg = keyed_value_and_grad(inner);
48
49            let (v, g) = vg(trainable_parameters, arrays)?;
50            Ok((v, g))
51        }
52    }
53}
54
55impl<'a, F, M, Args> IntoModuleValueAndGrad<'a, M, Args, Vec<Array>, Exception> for F
56where
57    M: ModuleParameters + 'a,
58    F: FnMut(&mut M, Args) -> Result<Vec<Array>, Exception> + 'a,
59    Args: Clone,
60{
61    fn into_module_value_and_grad(
62        mut self,
63    ) -> impl FnMut(&mut M, Args) -> Result<(Vec<Array>, FlattenedModuleParam), Exception> + 'a
64    {
65        move |model, arrays| {
66            let trainable_parameters = trainable_params(model);
67            let inner =
68                |parameters: FlattenedModuleParam, arrays: Args| -> Result<Vec<Array>, Exception> {
69                    let flattened_parameters = parameters.into_iter().map(|(k, v)| (k, v.clone()));
70                    update_parameters(model, flattened_parameters);
71
72                    self(model, arrays)
73                };
74            let mut vg = keyed_value_and_grad(inner);
75
76            let (v, g) = vg(trainable_parameters, arrays)?;
77            Ok((v, g))
78        }
79    }
80}
81
82impl<'a, F, M, Args> IntoModuleValueAndGrad<'a, M, Args, Array, ()> for F
83where
84    M: ModuleParameters + 'a,
85    F: FnMut(&mut M, Args) -> Array + 'a,
86    Args: Clone,
87{
88    fn into_module_value_and_grad(
89        mut self,
90    ) -> impl FnMut(&mut M, Args) -> Result<(Array, FlattenedModuleParam), Exception> + 'a {
91        move |model, arrays| {
92            let trainable_parameters = trainable_params(model);
93            let inner = |parameters: FlattenedModuleParam, arrays: Args| -> Vec<Array> {
94                let flattened_parameters = parameters.into_iter().map(|(k, v)| (k, v.clone()));
95                update_parameters(model, flattened_parameters);
96
97                vec![self(model, arrays)]
98            };
99            let mut vg = keyed_value_and_grad(inner);
100
101            let (v, g) = vg(trainable_parameters, arrays)?;
102            let v = v.into_iter().next().expect("Expected a single value");
103            Ok((v, g))
104        }
105    }
106}
107
108impl<'a, F, M, Args> IntoModuleValueAndGrad<'a, M, Args, Array, Exception> for F
109where
110    M: ModuleParameters + 'a,
111    F: FnMut(&mut M, Args) -> Result<Array, Exception> + 'a,
112    Args: Clone,
113{
114    fn into_module_value_and_grad(
115        mut self,
116    ) -> impl FnMut(&mut M, Args) -> Result<(Array, FlattenedModuleParam), Exception> + 'a {
117        move |model, arrays| {
118            let trainable_parameters = trainable_params(model);
119            let inner =
120                |parameters: FlattenedModuleParam, arrays: Args| -> Result<Vec<Array>, Exception> {
121                    let flattened_parameters = parameters.into_iter().map(|(k, v)| (k, v.clone()));
122                    update_parameters(model, flattened_parameters);
123
124                    self(model, arrays).map(|v| vec![v])
125                };
126            let mut vg = keyed_value_and_grad(inner);
127
128            let (v, g) = vg(trainable_parameters, arrays)?;
129            let v = v.into_iter().next().expect("Expected a single value");
130            Ok((v, g))
131        }
132    }
133}
134
135/// Transform the passed function `f(model, args)` to a function that computes the gradients of `f`
136/// with regard to the model's trainable parameters and also its value.
137pub fn value_and_grad<'a, F, M, Args, Val, Err>(
138    f: F,
139) -> impl FnMut(&mut M, Args) -> Result<(Val, FlattenedModuleParam), Exception> + 'a
140where
141    M: ModuleParameters + 'a,
142    F: IntoModuleValueAndGrad<'a, M, Args, Val, Err>,
143    Args: Clone,
144{
145    f.into_module_value_and_grad()
146}
147
148#[cfg(test)]
149mod tests {
150    use crate::module::Module;
151    use crate::{error::Exception, Array, Dtype};
152
153    use crate::nn::{self, Linear};
154
155    fn assert_finite_nonzero(value: impl AsRef<Array>) {
156        let value = value.as_ref();
157        assert_eq!(value.dtype(), Dtype::Float32);
158        assert!(value.shape().is_empty());
159        let value = value.item_exact::<f32>();
160        assert!(value.is_finite());
161        assert_ne!(value, 0.0);
162    }
163
164    // The unit test below is adapted from `test_compiled_optimizer` in
165    // `mlx/python/tests/test_optimizers.py``
166    #[test]
167    fn test_value_and_grad() {
168        let mut model = Linear::new(2, 2).unwrap();
169        let x = crate::random::uniform::<_, f32>(1.0, 2.0, &[2, 2], None).unwrap();
170
171        let loss = |model: &mut Linear, x: &Array| -> Vec<Array> {
172            vec![model.forward(x).unwrap().sum(None).unwrap()]
173        };
174
175        let mut vg = nn::value_and_grad(loss);
176        let (v, g) = vg(&mut model, &x).unwrap();
177
178        assert_finite_nonzero(v[0].sum(None).unwrap());
179        assert_finite_nonzero(g["weight"].sum(None).unwrap());
180        assert_finite_nonzero(g["bias"].sum(None).unwrap());
181    }
182
183    #[test]
184    fn test_value_and_grad_with_unary_output() {
185        let mut model = Linear::new(2, 2).unwrap();
186        let x = crate::random::uniform::<_, f32>(1.0, 2.0, &[2, 2], None).unwrap();
187
188        let loss = |model: &mut Linear, x: &Array| -> Array {
189            model.forward(x).unwrap().sum(None).unwrap()
190        };
191
192        let mut vg = nn::value_and_grad(loss);
193        let (v, g) = vg(&mut model, &x).unwrap();
194
195        assert_finite_nonzero(v.sum(None).unwrap());
196        assert_finite_nonzero(g["weight"].sum(None).unwrap());
197        assert_finite_nonzero(g["bias"].sum(None).unwrap());
198    }
199
200    #[test]
201    fn test_fallible_module_value_and_grad() {
202        let mut model = Linear::new(2, 2).unwrap();
203        let x = crate::random::uniform::<_, f32>(1.0, 2.0, &[2, 2], None).unwrap();
204
205        let loss = |model: &mut Linear, x: &Array| -> Result<Vec<Array>, Exception> {
206            Ok(vec![model.forward(x)?.sum(None)?])
207        };
208
209        let mut vg = nn::value_and_grad(loss);
210        let (v, g) = vg(&mut model, &x).unwrap();
211
212        assert_finite_nonzero(v[0].sum(None).unwrap());
213        assert_finite_nonzero(g["weight"].sum(None).unwrap());
214        assert_finite_nonzero(g["bias"].sum(None).unwrap());
215    }
216
217    #[test]
218    fn test_value_and_grad_with_two_args() {
219        let mut model = Linear::new(2, 2).unwrap();
220        let x = crate::random::uniform::<_, f32>(1.0, 2.0, &[2, 2], None).unwrap();
221        let y = crate::ops::ones::<f32>(x.shape()).unwrap();
222
223        let loss =
224            |model: &mut Linear, (x, y): (&Array, &Array)| -> Result<Vec<Array>, Exception> {
225                model
226                    .forward(x)?
227                    .subtract(y)?
228                    .square()?
229                    .sum(None)
230                    .map(|v| vec![v])
231            };
232
233        let mut vg = nn::value_and_grad(loss);
234        let (v, g) = vg(&mut model, (&x, &y)).unwrap();
235
236        assert_finite_nonzero(v[0].sum(None).unwrap());
237        assert_finite_nonzero(g["weight"].sum(None).unwrap());
238        assert_finite_nonzero(g["bias"].sum(None).unwrap());
239    }
240
241    #[test]
242    fn test_value_and_grad_with_error() {
243        let mut model = Linear::new(2, 2).unwrap();
244        // Use a shape that is not compatible with the model
245        let x = crate::random::uniform::<_, f32>(1.0, 2.0, &[3, 3], None).unwrap();
246
247        let loss = |model: &mut Linear, x: &Array| -> Result<Vec<Array>, Exception> {
248            Ok(vec![model.forward(x)?.sum(None)?])
249        };
250
251        let mut vg = nn::value_and_grad(loss);
252        let result = vg(&mut model, &x);
253
254        assert!(result.is_err());
255
256        // Check that the error message is not just "mlx_closure returned a non-zero value"
257        let err = result.unwrap_err();
258        assert!(!err.what().contains("non-zero value"))
259    }
260}