|
| 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