1use rustad::{
8 AD,
9 ADfn,
10 ad_from_vector,
11 start_recording,
12 stop_recording,
13 call_atom,
14 IndexT,
15};
16use super::V;
20fn value_callback_f(
24 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
25) -> ADfn<V> {
26 let x : Vec<V> = vec![ 1.0 , 2.0 ];
28 let ax = start_recording(x);
29 let ay = call_atom(ax, sumsq_atom_id, call_info, trace);
30 let f = stop_recording(ay);
31 f
32}
33pub fn callback_forward_zero_value(
36 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
37) {
38 let f = value_callback_f(sumsq_atom_id, call_info, trace);
41 let x : Vec<V> = vec![ 3.0 , 4.0 ];
44 let mut v : Vec<V> = Vec::new();
45 let y = f.forward_zero_value(&mut v , x.clone(), trace);
46 assert_eq!( y[0], x[0]*x[0] + x[1]*x[1] );
47}
48pub fn callback_forward_zero_ad(
51 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
52) {
53 let f = value_callback_f(sumsq_atom_id, call_info, trace);
55 let x : Vec<V> = vec![ 3.0 , 4.0 ];
58 let ax = start_recording(x);
59 let mut av : Vec< AD<V> > = Vec::new();
60 let ay = f.forward_zero_ad(&mut av , ax.clone(), trace);
61 let g = stop_recording(ay);
62 let x : Vec<V> = vec![ 3.0 , 4.0 ];
65 let mut v : Vec<V> = Vec::new();
66 let y = g.forward_zero_value(&mut v , x.clone(), trace);
67 assert_eq!( y[0], x[0]*x[0] + x[1]*x[1] );
68 assert_eq!( y[0], x[0]*x[0] + x[1]*x[1] );
70}
71pub fn callback_forward_one_value(
74 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
75) {
76 let f = value_callback_f(sumsq_atom_id, call_info, trace);
79 let x : Vec<V> = vec![ 3.0 , 4.0 ];
82 let mut v : Vec<V> = Vec::new();
83 f.forward_zero_value(&mut v , x.clone(), trace);
84 let dx : Vec<V> = vec![ 5.0, 6.0 ];
85 let dy = f.forward_one_value(&v , dx.clone(), trace);
86 assert_eq!( dy[0], 2.0 * x[0]*dx[0] + 2.0 * x[1]*dx[1] );
87}
88pub fn callback_forward_one_ad(
91 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
92) {
93 let f = value_callback_f(sumsq_atom_id, call_info, trace);
96 let x : Vec<V> = vec![ 3.0 , 4.0 ];
99 let ax = start_recording(x);
100 let mut av : Vec< AD<V> > = Vec::new();
101 f.forward_zero_ad(&mut av , ax, trace);
102 let dx1 : Vec<V> = vec![ 5.0, 6.0 ];
103 let adx1 = ad_from_vector(dx1.clone());
104 let ady = f.forward_one_ad(&av, adx1, trace);
109 let g = stop_recording(ady);
110 let x : Vec<V> = vec![ 3.0 , 4.0 ];
114 let mut v : Vec<V> = Vec::new();
115 let y = g.forward_zero_value(&mut v , x.clone(), trace);
116 assert_eq!( y[0], 2.0 * ( x[0] * dx1[0] + x[1] * dx1[1] ) );
117 let dx2 : Vec<V> = vec![ 7.0, 8.0 ];
120 let dy = g.forward_one_value(&v , dx2.clone(), trace);
121 assert_eq!( dy[0], 2.0 * ( dx2[0] * dx1[0] + dx2[1] * dx1[1] ) );
122 let dy : Vec<V> = vec![ 9.0 ];
125 let dx = g.reverse_one_value(&v , dy.clone(), trace);
126 assert_eq!( dx[0], 2.0 * dy[0] * dx1[0] );
127 assert_eq!( dx[1], 2.0 * dy[0] * dx1[1] );
128}
129pub fn callback_reverse_one_value(
132 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
133) {
134 let f = value_callback_f(sumsq_atom_id, call_info, trace);
137 let x : Vec<V> = vec![ 3.0 , 4.0 ];
140 let mut v : Vec<V> = Vec::new();
141 f.forward_zero_value(&mut v , x.clone(), trace);
142 let dy : Vec<V> = vec![ 5.0 ];
143 let dx = f.reverse_one_value(&v , dy.clone(), trace);
144 assert_eq!( dx[0], 2.0 * x[0]*dy[0] );
145 assert_eq!( dx[1], 2.0 * x[1]*dy[0] );
146}
147pub fn callback_reverse_one_ad(
150 sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
151) {
152 let f = value_callback_f(sumsq_atom_id, call_info, trace);
155 let x : Vec<V> = vec![ 3.0 , 4.0 ];
158 let ax = start_recording(x);
159 let mut av : Vec< AD<V> > = Vec::new();
160 f.forward_zero_ad(&mut av , ax, trace);
161 let dy1 : Vec<V> = vec![ 5.0 ];
162 let ady1 = ad_from_vector(dy1.clone());
163 let adx = f.reverse_one_ad(&av, ady1, trace);
168 let g = stop_recording(adx);
169 let x : Vec<V> = vec![ 3.0 , 4.0 ];
173 let mut v : Vec<V> = Vec::new();
174 let y = g.forward_zero_value(&mut v , x.clone(), trace);
175 assert_eq!( y[0], 2.0 * dy1[0] * x[0] );
176 assert_eq!( y[1], 2.0 * dy1[0] * x[1] );
177 let dx : Vec<V> = vec![ 6.0, 7.0 ];
180 let dy = g.forward_one_value(&v, dx.clone(), trace);
181 assert_eq!( dy[0], 2.0 * dy1[0] * dx[0] );
182 assert_eq!( dy[1], 2.0 * dy1[0] * dx[1] );
183 let dy2 : Vec<V> = vec![ 8.0, 9.0 ];
186 let dx = g.reverse_one_value(&v, dy2.clone(), trace);
187 assert_eq!( dx[0], 2.0 * dy1[0] * dy2[0] );
188 assert_eq!( dx[1], 2.0 * dy1[0] * dy2[1] );
189}