Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion optvl/om_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,7 +194,10 @@ def om_set_avl_inputs(sys, inputs):
# add the parameters to the run
for ref in sys.ovl.ref_var_to_fort_var:
if ref in inputs:
val = inputs[ref][0]
if ref == "XYZref":
val = inputs[ref][0:3]
else:
val = inputs[ref][0]
sys.ovl.set_reference_data({ref: val})


Expand Down
3 changes: 2 additions & 1 deletion optvl/optvl_class.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,7 @@ class OVLSolver(object):
"Sref": ["CASE_R", "SREF"],
"Cref": ["CASE_R", "CREF"],
"Bref": ["CASE_R", "BREF"],
"XYZref": ["CASE_R", "XYZREF"],
}

case_derivs_to_fort_var = {
Expand Down Expand Up @@ -3201,7 +3202,7 @@ def _execute_jac_vec_prod_fwd(
mode: str = "AD",
step: float = 1e-7,
) -> Tuple[Dict[str, float], np.ndarray, Dict[str, float], Dict[str, float], np.ndarray, np.ndarray]:
"""Get partial derivatives in forward mode. This routine is usefulinternally and when creating wrappers for things like OpenMDAO
"""Get partial derivatives in forward mode. This routine is useful internally and when creating wrappers for things like OpenMDAO

Args:
con_seeds: Case constraint AD seeds
Expand Down
19 changes: 12 additions & 7 deletions tests/test_partial_derivs.py
Original file line number Diff line number Diff line change
Expand Up @@ -284,9 +284,10 @@ def test_rev_param(self):

def test_fwd_ref(self):
for ref_key in self.ovl_solver.ref_var_to_fort_var:
func_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[0]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0

func_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0}, mode="FD", step=1e-7)[0]
func_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[0]
func_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed}, mode="FD", step=1e-7)[0]

for func_key in func_seeds:
# print(f"{func_key} wrt {ref_key}", func_seeds[func_key], func_seeds_FD[func_key])
Expand All @@ -309,9 +310,9 @@ def test_fwd_ref(self):

def test_rev_ref(self):
for ref_key in self.ovl_solver.ref_var_to_fort_var:
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
self.ovl_solver.clear_ad_seeds_fast()

func_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[0]
func_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[0]
self.ovl_solver.clear_ad_seeds_fast()

for func_key in self.ovl_solver.case_var_to_fort_var:
Expand Down Expand Up @@ -490,8 +491,11 @@ def test_rev_param(self):

def test_fwd_ref(self):
for ref_key in self.ovl_solver.ref_var_to_fort_var:
res_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[1]
res_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0}, mode="FD", step=1e-7)[1]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
res_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[1]
res_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(
ref_seeds={ref_key: 1.0}, mode="FD", step=1e-7
)[1]

# print(f"res wrt {ref_key}", np.linalg.norm(res_seeds), np.linalg.norm(res_seeds_FD))
np.testing.assert_allclose(res_seeds, res_seeds_FD, atol=1e-6, err_msg=f"d(res) w.r.t.{ref_key}")
Expand All @@ -504,7 +508,8 @@ def test_rev_ref(self):
self.ovl_solver.clear_ad_seeds_fast()

for ref_key in self.ovl_solver.ref_var_to_fort_var:
res_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[1]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
res_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[1]
# do dot product
res_sum = np.sum(res_seeds_rev * res_seeds_fwd)
ref_sum = np.sum(ref_seeds_rev[ref_key])
Expand Down
13 changes: 8 additions & 5 deletions tests/test_stab_derivs_partial_derivs.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,9 +180,10 @@ def test_rev_gamma_u(self):

def test_fwd_ref(self):
for ref_key in self.ovl_solver.ref_var_to_fort_var:
res_u_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[5]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
res_u_seeds = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[5]

res_u_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0}, mode="FD", step=1e-5)[
res_u_seeds_FD = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed}, mode="FD", step=1e-5)[
5
]

Expand Down Expand Up @@ -381,9 +382,10 @@ def test_rev_gamma_u(self):

def test_fwd_ref(self):
for ref_key in self.ovl_solver.ref_var_to_fort_var:
sd_d = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[3]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
sd_d = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[3]

sd_d_fd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0}, mode="FD", step=1e-6)[3]
sd_d_fd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed}, mode="FD", step=1e-6)[3]

for deriv_func in sd_d:
sens_label = f"{deriv_func} wrt {ref_key}"
Expand Down Expand Up @@ -416,7 +418,8 @@ def test_rev_ref(self):
self.ovl_solver.clear_ad_seeds_fast()

for ref_key in self.ovl_solver.ref_var_to_fort_var:
stab_deriv_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: 1.0})[3]
ref_seed = np.ones(3) if ref_key == "XYZref" else 1.0
stab_deriv_seeds_fwd = self.ovl_solver._execute_jac_vec_prod_fwd(ref_seeds={ref_key: ref_seed})[3]

stab_deriv_sum = 0.0
for deriv_func in stab_deriv_seeds_fwd:
Expand Down
4 changes: 2 additions & 2 deletions tests/test_total_derivs.py
Original file line number Diff line number Diff line change
Expand Up @@ -355,7 +355,7 @@ def test_ref(self):
# print(f"{func_key:5} wrt {ref_key:5} | AD:{ad_dot: 5e} FD:{fd_dot: 5e} rel err:{rel_err:.2e}")

tol = 1e-13
if np.abs(ad_dot) < tol or np.abs(fd_dot) < tol:
if np.abs(np.linalg.norm(ad_dot)) < tol or np.abs(fd_dot) < tol:
# If either value is basically zero, use an absolute tolerance
np.testing.assert_allclose(
ad_dot,
Expand All @@ -381,7 +381,7 @@ def test_ref(self):
# f"{func_key} wrt {var_key:5} wrt {ref_key} | AD:{ad_dot: 5e} FD:{func_dot: 5e} rel err:{rel_err:.2e}"
# )
tol = 1e-8
if np.abs(ad_dot) < tol or np.abs(func_dot) < tol:
if np.abs(np.linalg.norm(ad_dot)) < tol or np.abs(func_dot) < tol:
# If either value is basically zero, use an absolute tolerance
np.testing.assert_allclose(
ad_dot,
Expand Down