Skip to main content

atom/
tests.rs

1// -------------------------------------------------------------------------
2// SPDX-License-Identifier: EPL-2.0 OR GPL-2.0-or-later
3// SPDX-FileCopyrightText: Bradley M. Bell <bradbell@seanet.com>
4// SPDX-FileContributor: 2025 Bradley M. Bell
5// -------------------------------------------------------------------------
6//
7use rustad::{
8    AD,
9    ADfn,
10    ad_from_vector,
11    start_recording,
12    stop_recording,
13    call_atom,
14    IndexT,
15};
16//
17//
18// V
19use super::V;
20//
21// value_callback_f
22// f(x) = x[0] * x[0] + x[1] * x[1] + ...
23fn value_callback_f(
24    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
25) -> ADfn<V> {
26    //
27    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}
33//
34// callback_forward_zero_value
35pub fn callback_forward_zero_value(
36    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
37) {
38    //
39    // f
40    let f = value_callback_f(sumsq_atom_id, call_info, trace);
41    //
42    // x, y
43    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}
48//
49// callback_forward_zero_ad
50pub fn callback_forward_zero_ad(
51    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
52) {
53    //
54    let f = value_callback_f(sumsq_atom_id, call_info, trace);
55    //
56    // g(x) = f(x)
57    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    //
63    // x, y
64    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    //
69    assert_eq!( y[0], x[0]*x[0] + x[1]*x[1] );
70}
71//
72// callback_forward_one_value
73pub fn callback_forward_one_value(
74    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
75) {
76    //
77    // f
78    let f = value_callback_f(sumsq_atom_id, call_info, trace);
79    //
80    // x, dx, dy
81    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}
88//
89// callback_forward_one_ad
90pub fn callback_forward_one_ad(
91    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
92) {
93    //
94    // f
95    let f = value_callback_f(sumsq_atom_id, call_info, trace);
96    //
97    // av, dx1, adx1
98    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    //
105    // g
106    // callback to sumsq_forward_one_ad
107    // g(x) = f'(x) * dx1 = 2 * ( x[0] * dx1[0] + x[2] * dx1[2] + ... )
108    let ady              = f.forward_one_ad(&av, adx1, trace);
109    let g                = stop_recording(ady);
110    //
111    // x, v, y
112    // check forward_zero_value
113    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    //
118    // check forward_one_value
119    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    //
123    // check reverse_one_value
124    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}
129//
130// callback_reverse_one_value
131pub fn callback_reverse_one_value(
132    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
133) {
134    //
135    // f
136    let f = value_callback_f(sumsq_atom_id, call_info, trace);
137    //
138    // x, dy, dx
139    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}
147//
148// callback_reverse_one_ad
149pub fn callback_reverse_one_ad(
150    sumsq_atom_id : IndexT , call_info : IndexT, trace : bool
151) {
152    //
153    // f
154    let f = value_callback_f(sumsq_atom_id, call_info, trace);
155    //
156    // av, dy1, ady1
157    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    //
164    // g
165    // callback to sumsq_reverse_one_ad
166    // g(x) = dy1 * f'(x) = 2 * ( dy1[0] * x[0], dy1[0] * x[1], ... )
167    let adx              = f.reverse_one_ad(&av, ady1, trace);
168    let g                = stop_recording(adx);
169    //
170    // x, v
171    // check forward_zero_value
172    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    //
178    // check forward_one_value
179    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    //
184    // check reverse_one_value
185    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}