Skip to content

Make StaticDataclass hashes deterministic across processes. - #2471

Open
copybara-service[bot] wants to merge 1 commit into
mainfrom
test_983683799
Open

copybara-service[bot] wants to merge 1 commit into
mainfrom
test_983683799

Conversation

@copybara-service

Copy link
Copy Markdown

Make StaticDataclass hashes deterministic across processes.

StaticDataclass.__hash__ returns hash(self._full_tuple()), and
_full_tuple() contains the class id string
(f"{cls.__module__}.{cls.__qualname__}") along with any string, bytes, or enum
fields. Python randomizes hash() for str and bytes per process unless
PYTHONHASHSEED is pinned, so the hash of a Solver or TimeStepCalculator
differs between runs of the same binary.

Since these objects are passed as static arguments to jitted functions (they end
up in the aux_data of SimulationStepFn's pytree), a process-dependent hash
becomes part of the JAX persistent compilation cache key. The cache therefore
missed on every fresh process and every simulation recompiled its step function
from scratch.

Map values with randomized hashes onto stable keys before hashing: str and
bytes are hashed with SHA-256 and truncated to 64 bits, enum members are
expanded into their qualified class name, member name and value, tuples are
converted elementwise, and None is mapped to a fixed sentinel. Numbers and
booleans already hash deterministically and are passed through unchanged.
Equality is untouched, as __eq__ still compares _full_tuple() directly.

With this change, enabling jax_compilation_cache_dir gives warm-cache hits
across processes without also having to set PYTHONHASHSEED=0. On a SPARC full
pulse simulation this takes the compilation phase from 39.4s to 17.5s on a warm
cache, and the end to end run from 74.6s to 51.5s.

`StaticDataclass.__hash__` returns `hash(self._full_tuple())`, and
`_full_tuple()` contains the class id string
(`f"{cls.__module__}.{cls.__qualname__}"`) along with any string, bytes, or enum
fields. Python randomizes `hash()` for `str` and `bytes` per process unless
`PYTHONHASHSEED` is pinned, so the hash of a `Solver` or `TimeStepCalculator`
differs between runs of the same binary.

Since these objects are passed as static arguments to jitted functions (they end
up in the `aux_data` of `SimulationStepFn`'s pytree), a process-dependent hash
becomes part of the JAX persistent compilation cache key. The cache therefore
missed on every fresh process and every simulation recompiled its step function
from scratch.

Map values with randomized hashes onto stable keys before hashing: `str` and
`bytes` are hashed with SHA-256 and truncated to 64 bits, `enum` members are
expanded into their qualified class name, member name and value, tuples are
converted elementwise, and `None` is mapped to a fixed sentinel. Numbers and
booleans already hash deterministically and are passed through unchanged.
Equality is untouched, as `__eq__` still compares `_full_tuple()` directly.

With this change, enabling `jax_compilation_cache_dir` gives warm-cache hits
across processes without also having to set `PYTHONHASHSEED=0`. On a SPARC full
pulse simulation this takes the compilation phase from 39.4s to 17.5s on a warm
cache, and the end to end run from 74.6s to 51.5s.

PiperOrigin-RevId: 983683799
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant