dade2d8c3e
Previously mjx dataclasses would return np.ndarray as bytes for jax tracing to hash since they are not inherently hashable. This meant that any tracing that was cached would hold on to a copy of the numpy arrays in the dataclass. Children such as mjx.Model that store large numpy arrays would end up duplicating that data 3+ times in some cases. This eliminates O(N * array_size) memory duplication across N cached pytree traces, saving a lot of memory on models with heavy mesh/texture data. PiperOrigin-RevId: 880985898 Change-Id: I58d4e91fda4112e4c633818f92907d20138e962a