@@ -6,6 +6,10 @@ pub mod utils;
66
77#[ cfg( test) ]
88pub mod test_utils {
9+ pub mod decompose;
10+ pub mod gradient_type;
11+ pub mod verify_gradient;
12+
913 use std:: f32:: consts:: FRAC_PI_2 ;
1014
1115 use geometry:: { arc:: Arc , circle:: Circle , direction:: Direction , line_segment:: LineSegment } ;
@@ -27,78 +31,4 @@ pub mod test_utils {
2731 PathSegment :: LineSegment ( LineSegment ( point![ 4.0 , 1.0 ] , point![ 4.0 , 4.0 ] ) ) ,
2832 ]
2933 }
30-
31- pub mod verify_gradient {
32- use std:: fmt:: Debug ;
33-
34- use approx:: { assert_relative_eq, AbsDiffEq , RelativeEq } ;
35- use num_traits:: { real:: Real , NumAssignOps } ;
36-
37- use crate :: traits:: {
38- decompose:: Decompose ,
39- gradient_type:: { Gradient , GradientType } ,
40- } ;
41-
42- pub fn verify_gradient <
43- A : Clone + Debug + Decompose < F > + GradientType ,
44- F : Real + NumAssignOps ,
45- > (
46- func : & impl Fn ( A ) -> F ,
47- grad : & impl Fn ( A ) -> Gradient < A > ,
48- epsilon : <Gradient < A > as AbsDiffEq >:: Epsilon ,
49- x : A ,
50- ) where
51- Gradient < A > : Debug + RelativeEq + Decompose < F > ,
52- <Gradient < A > as AbsDiffEq >:: Epsilon : From < f32 > ,
53- {
54- let real_gradient = grad ( x. clone ( ) ) ;
55- let numerical_gradient = numerical_grad ( func, x) ;
56-
57- assert_relative_eq ! ( real_gradient, numerical_gradient, epsilon = epsilon) ;
58- }
59-
60- fn numerical_grad < A : Clone + Decompose < F > + GradientType , F : Real + NumAssignOps > (
61- func : & impl Fn ( A ) -> F ,
62- x : A ,
63- ) -> Gradient < A >
64- where
65- Gradient < A > : Decompose < F > ,
66- {
67- let decomposed = ( 0 ..A :: N )
68- . map ( |i| numerical_nth_derivative ( func, i, x. clone ( ) ) )
69- . collect ( ) ;
70-
71- Gradient :: < A > :: compose ( decomposed)
72- }
73-
74- fn numerical_nth_derivative < A : Decompose < F > , F : Real + NumAssignOps > (
75- func : & impl Fn ( A ) -> F ,
76- n : usize ,
77- x : A ,
78- ) -> F {
79- let eps = F :: from ( 1e-4 ) . unwrap ( ) ;
80-
81- let middle = x. decompose ( ) ;
82- let above = {
83- let mut above = middle. clone ( ) ;
84- above[ n] += eps;
85-
86- A :: compose ( above)
87- } ;
88- let below = {
89- let mut below = middle. clone ( ) ;
90- below[ n] -= eps;
91-
92- A :: compose ( below)
93- } ;
94-
95- let sample_above = func ( above) ;
96- let sample_below = func ( below) ;
97-
98- let difference = sample_above - sample_below;
99- let sample_distance = F :: from ( 2.0 ) . unwrap ( ) * eps;
100-
101- difference / sample_distance
102- }
103- }
10434}
0 commit comments