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
28 changes: 28 additions & 0 deletions ccflow/tests/utils/test_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -778,6 +778,34 @@ def test_numpy_structured_array(self):
arr = np.array([(1, 2.0), (3, 4.0)], dtype=dt)
assert normalize_token(arr)[0] == "ndarray"

def test_numpy_structured_array_ignores_padding(self):
dt = np.dtype([("a", np.uint8), ("b", np.uint64)], align=True)
a = np.frombuffer(bytearray(dt.itemsize), dtype=dt)
b = np.frombuffer(bytearray(b"\xff" * dt.itemsize), dtype=dt)
for arr in (a, b):
arr["a"] = 7
arr["b"] = 12345
assert a.tobytes() != b.tobytes()
assert normalize_token(a) == normalize_token(b)

def test_numpy_structured_array_different_values(self):
dt = np.dtype([("a", np.uint8), ("b", np.uint64)], align=True)
a = np.zeros(1, dtype=dt)
b = np.zeros(1, dtype=dt)
a["b"] = 1
b["b"] = 2
assert normalize_token(a) != normalize_token(b)

def test_numpy_structured_array_nested(self):
dt = np.dtype([("p", [("x", np.float32), ("y", np.float32)]), ("m", np.int16, (2, 2))], align=True)
a = np.zeros(2, dtype=dt)
b = np.zeros(2, dtype=dt)
a["p"]["x"] = 1.5
b["p"]["x"] = 1.5
assert normalize_token(a) == normalize_token(b)
b["m"][0, 0, 0] = 4
assert normalize_token(a) != normalize_token(b)

def test_numpy_scalar(self):
s = np.int64(42)
assert normalize_token(s) == ("np_scalar", "int64", 42)
Expand Down
4 changes: 4 additions & 0 deletions ccflow/utils/tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -362,6 +362,10 @@ def _normalize_ndarray(obj):
# element-wise instead. Cycle detection is keyed on the array because tolist() returns a fresh list.
if obj.dtype.hasobject:
return _with_cycle_check(obj, lambda: ("ndarray", str(obj.dtype), obj.shape, normalize_token(obj.tolist())))
# Structured dtypes may carry padding between fields. Those bytes are never written, so hashing the
# raw buffer gives arrays with identical field values different tokens. Recurse per field to skip them.
if obj.dtype.names is not None:
return ("ndarray", str(obj.dtype), obj.shape, tuple((name, normalize_token(obj[name])) for name in obj.dtype.names))
return ("ndarray", str(obj.dtype), obj.shape, hashlib.sha256(np.ascontiguousarray(obj).tobytes()).hexdigest())

@normalize_token.register(np.ma.MaskedArray)
Expand Down
Loading