Skip to content

Commit 7174be6

Browse files
c-dilksclaude
andauthored
perf: hipo2npz memory efficiency (#1373)
Co-authored-by: Claude <noreply@anthropic.com>
1 parent e9201d9 commit 7174be6

3 files changed

Lines changed: 360 additions & 204 deletions

File tree

‎.github/workflows/ci.yml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -357,7 +357,7 @@ jobs:
357357
- name: untar build
358358
run: tar xzvf coatjava.tar.gz
359359
- name: hipo2npz
360-
run: ./coatjava/bin/hipo2npz rec.hipo rec.npz RUN::config,REC::Event,REC::Particle
360+
run: ./coatjava/bin/hipo2npz rec.hipo rec.npz
361361
- name: hipo2npz-dump
362362
run: ./coatjava/bin/hipo2npz-dump rec.npz 1
363363

‎bin/hipo2npz-diff‎

Lines changed: 132 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,132 @@
1+
#!/usr/bin/env python3
2+
3+
######################################
4+
# author: generated by Claude Sonnet 5
5+
######################################
6+
7+
"""
8+
hipo2npz-diff - Compare two NumPy .npz files and report differences.
9+
10+
Usage:
11+
hipo2npz-diff file1.npz file2.npz
12+
hipo2npz-diff file1.npz file2.npz --rtol 1e-5 --atol 1e-8
13+
hipo2npz-diff file1.npz file2.npz --exact
14+
hipo2npz-diff file1.npz file2.npz --verbose
15+
16+
Exit codes:
17+
0 - files are equivalent (given tolerance)
18+
1 - differences found
19+
2 - error (bad file, etc.)
20+
"""
21+
22+
import argparse
23+
import sys
24+
25+
import numpy as np
26+
27+
28+
def diff_npz(path_a, path_b, rtol, atol, exact, verbose):
29+
try:
30+
a = np.load(path_a, allow_pickle=True)
31+
b = np.load(path_b, allow_pickle=True)
32+
except Exception as e:
33+
print(f"Error loading files: {e}", file=sys.stderr)
34+
sys.exit(2)
35+
36+
keys_a = set(a.files)
37+
keys_b = set(b.files)
38+
39+
only_in_a = sorted(keys_a - keys_b)
40+
only_in_b = sorted(keys_b - keys_a)
41+
common = sorted(keys_a & keys_b)
42+
43+
has_diff = False
44+
45+
if only_in_a:
46+
has_diff = True
47+
print(f"Keys only in {path_a}:")
48+
for k in only_in_a:
49+
print(f" - {k}")
50+
51+
if only_in_b:
52+
has_diff = True
53+
print(f"Keys only in {path_b}:")
54+
for k in only_in_b:
55+
print(f" - {k}")
56+
57+
for key in common:
58+
arr_a, arr_b = a[key], b[key]
59+
60+
if arr_a.shape != arr_b.shape:
61+
has_diff = True
62+
print(f"[{key}] shape mismatch: {arr_a.shape} vs {arr_b.shape}")
63+
continue
64+
65+
if arr_a.dtype != arr_b.dtype and verbose:
66+
print(f"[{key}] dtype differs: {arr_a.dtype} vs {arr_b.dtype}")
67+
68+
is_float = np.issubdtype(arr_a.dtype, np.floating)
69+
70+
try:
71+
if exact:
72+
# equal_nan=True: NaNs in the same position count as matching,
73+
# since NaN != NaN by IEEE rules but that's rarely what you want here
74+
equal = np.array_equal(arr_a, arr_b, equal_nan=is_float)
75+
else:
76+
equal = np.allclose(arr_a, arr_b, rtol=rtol, atol=atol, equal_nan=True)
77+
except TypeError:
78+
# Non-numeric / object arrays: fall back to plain equality (no equal_nan support)
79+
equal = np.array_equal(arr_a, arr_b)
80+
81+
if not equal:
82+
has_diff = True
83+
print(f"[{key}] values differ", end="")
84+
try:
85+
fa = arr_a.astype(np.float64)
86+
fb = arr_b.astype(np.float64)
87+
diff = np.abs(fa - fb)
88+
89+
# Positions where exactly one side is NaN (a "real" mismatch, not just
90+
# matching NaNs) vs. positions with a finite numeric difference
91+
nan_mismatch = np.isnan(fa) != np.isnan(fb)
92+
finite_mask = ~np.isnan(fa) & ~np.isnan(fb)
93+
94+
n_diff = np.count_nonzero(
95+
finite_mask & (diff > (atol + rtol * np.abs(fb)))
96+
)
97+
n_nan_mismatch = np.count_nonzero(nan_mismatch)
98+
99+
max_diff = np.max(diff[finite_mask]) if finite_mask.any() else 0.0
100+
print(
101+
f" (max abs diff = {max_diff:.6g}, "
102+
f"{n_diff}/{arr_a.size} elements differ, "
103+
f"{n_nan_mismatch} NaN-mismatch positions)"
104+
)
105+
except (TypeError, ValueError):
106+
print()
107+
elif verbose:
108+
print(f"[{key}] OK (identical within tolerance)")
109+
110+
if not has_diff:
111+
print(f"No differences found between {path_a} and {path_b}"
112+
+ ("" if exact else f" (rtol={rtol}, atol={atol})"))
113+
114+
return has_diff
115+
116+
117+
def main():
118+
parser = argparse.ArgumentParser(description="Diff two .npz files.")
119+
parser.add_argument("file_a", help="First .npz file")
120+
parser.add_argument("file_b", help="Second .npz file")
121+
parser.add_argument("--rtol", type=float, default=1e-5, help="Relative tolerance for float comparison (default: 1e-5)")
122+
parser.add_argument("--atol", type=float, default=1e-8, help="Absolute tolerance for float comparison (default: 1e-8)")
123+
parser.add_argument("--exact", action="store_true", help="Require exact equality instead of tolerance-based comparison")
124+
parser.add_argument("--verbose", action="store_true", help="Print status for matching keys too")
125+
args = parser.parse_args()
126+
127+
has_diff = diff_npz(args.file_a, args.file_b, args.rtol, args.atol, args.exact, args.verbose)
128+
sys.exit(1 if has_diff else 0)
129+
130+
131+
if __name__ == "__main__":
132+
main()

0 commit comments

Comments
 (0)