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
16pub trait IntoModuleValueAndGrad<'a, M, Args, Val, Err>
18where
19 M: ModuleParameters + 'a,
20 Args: Clone,
21{
22 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
135pub 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 #[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 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 let err = result.unwrap_err();
258 assert!(!err.what().contains("non-zero value"))
259 }
260}