Skip to content

Commit f84b0b1

Browse files
committed
Adding comparison script
1 parent 75138b8 commit f84b0b1

1 file changed

Lines changed: 42 additions & 0 deletions

File tree

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,42 @@
1+
import numpy as np
2+
import sys
3+
4+
sys.path.append('./jaxnp_hash/')
5+
6+
from jan_example import h_max_gamma_over_KY_jax as jan_hfun
7+
from ibcdfo.manifold_sampling import h_max_gamma_over_KY as old_hfun
8+
9+
10+
a = np.load("jans_msp_output_0.npz", allow_pickle=True)
11+
b = np.load("old_msp_output_0.npz", allow_pickle=True)
12+
13+
for key in ["X", "F", "h_msp"]:
14+
print(f"\n{key}:")
15+
print(" same shape:", a[key].shape == b[key].shape)
16+
print(" exactly equal:", np.array_equal(a[key], b[key]))
17+
print(" allclose:", np.allclose(a[key], b[key], rtol=1e-12, atol=1e-12, equal_nan=True))
18+
19+
if a[key].shape == b[key].shape:
20+
diff = a[key] - b[key]
21+
print(" max abs diff:", np.nanmax(np.abs(diff)))
22+
23+
24+
diff_h = a["h_msp"] - b["h_msp"]
25+
rows = np.unique(np.where(diff_h)[0])
26+
27+
print("\nRows where h_msp differs:", rows)
28+
29+
for row in rows:
30+
Frow = a["F"][row]
31+
32+
print(f"\n=== row {row} ===")
33+
print("Frow:", Frow)
34+
print("saved jan h_msp:", a["h_msp"][row])
35+
print("saved old h_msp:", b["h_msp"][row])
36+
print("saved diff:", a["h_msp"][row] - b["h_msp"][row])
37+
38+
jan_val = jan_hfun(Frow)
39+
old_val = old_hfun(Frow)
40+
41+
print("jan_hfun(Frow):", jan_val)
42+
print("old_hfun(Frow):", old_val)

0 commit comments

Comments
 (0)