Make StaticDataclass hashes deterministic across processes. - #2471
Open
copybara-service[bot] wants to merge 1 commit into
Open
copybara-service[bot] wants to merge 1 commit into
copybara-service[bot] wants to merge 1 commit into
Conversation
`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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Make StaticDataclass hashes deterministic across processes.
StaticDataclass.__hash__returnshash(self._full_tuple()), and_full_tuple()contains the class id string(
f"{cls.__module__}.{cls.__qualname__}") along with any string, bytes, or enumfields. Python randomizes
hash()forstrandbytesper process unlessPYTHONHASHSEEDis pinned, so the hash of aSolverorTimeStepCalculatordiffers between runs of the same binary.
Since these objects are passed as static arguments to jitted functions (they end
up in the
aux_dataofSimulationStepFn's pytree), a process-dependent hashbecomes 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:
strandbytesare hashed with SHA-256 and truncated to 64 bits,enummembers areexpanded into their qualified class name, member name and value, tuples are
converted elementwise, and
Noneis mapped to a fixed sentinel. Numbers andbooleans 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_dirgives warm-cache hitsacross processes without also having to set
PYTHONHASHSEED=0. On a SPARC fullpulse 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.