From 92375271fcb54ec25a36944d101d74e4f02094d1 Mon Sep 17 00:00:00 2001 From: Sadi Kneipp Date: Tue, 22 Sep 2026 10:54:55 -0700 Subject: [PATCH] Remove redundant CloudPathwaysArrayHandler from pathwaysutils in favor of orbax. PiperOrigin-RevId: 986101467 --- pathwaysutils/_initialize.py | 14 +- pathwaysutils/persistence/orbax_handler.py | 352 --------------------- pathwaysutils/test/initialize_test.py | 15 + 3 files changed, 24 insertions(+), 357 deletions(-) delete mode 100644 pathwaysutils/persistence/orbax_handler.py diff --git a/pathwaysutils/_initialize.py b/pathwaysutils/_initialize.py index 27476c9..de7bb46 100644 --- a/pathwaysutils/_initialize.py +++ b/pathwaysutils/_initialize.py @@ -18,10 +18,11 @@ import os import jax +from orbax.checkpoint import type_handlers from orbax.checkpoint._src.metadata import array_metadata_store as array_metadata_store_lib +from orbax.checkpoint._src.serialization import cloud_pathways_array_handler from pathwaysutils import profiling from pathwaysutils import proxy_backend -from pathwaysutils.persistence import orbax_handler _logger = logging.getLogger(__name__) @@ -91,11 +92,14 @@ def initialize() -> None: _logger.debug("Detected Pathways-on-Cloud backend. Applying changes.") proxy_backend.register_backend_factory() profiling.monkey_patch_jax() - # TODO: b/365549911 - Remove when OCDBT-compatible if _is_persistence_enabled(): - orbax_handler.register_pathways_handlers( - timeout=datetime.timedelta(hours=1), - array_metadata_store=array_metadata_store_lib.Store(), + type_handlers.register_type_handler( + jax.Array, + cloud_pathways_array_handler.CloudPathwaysArrayHandler( + timeout=datetime.timedelta(hours=1), + array_metadata_store=array_metadata_store_lib.Store(), + ), + override=True, ) # Turn off JAX compilation cache because Pathways handles its own diff --git a/pathwaysutils/persistence/orbax_handler.py b/pathwaysutils/persistence/orbax_handler.py deleted file mode 100644 index 7d4deb4..0000000 --- a/pathwaysutils/persistence/orbax_handler.py +++ /dev/null @@ -1,352 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# https://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""TypeHandlers supporting Pathways backend.""" - -import collections -from collections.abc import Coroutine, Sequence -import concurrent.futures -import datetime -import logging -from typing import Any, cast - -import jax -from orbax.checkpoint import future -from orbax.checkpoint import type_handlers -from orbax.checkpoint._src.metadata import array_metadata as array_metadata_lib -from orbax.checkpoint._src.metadata import array_metadata_store as array_metadata_store_lib -from pathwaysutils.persistence import helper - - -_logger = logging.getLogger(__name__) - -ParamInfo = type_handlers.ParamInfo -SaveArgs = type_handlers.SaveArgs -RestoreArgs = type_handlers.RestoreArgs -ArrayRestoreArgs = type_handlers.ArrayRestoreArgs -ArrayMetadata = array_metadata_lib.ArrayMetadata - - -def extract_parent_dir_and_name( - infos: Sequence[ParamInfo], -) -> tuple[Sequence[str], Sequence[str]]: - """Extracts names and locations from ParamInfos.""" - parent_dirs = [str(info.parent_dir) for info in infos] - names = [str(info.name) for info in infos] - return parent_dirs, names - - -class CloudPathwaysArrayHandler(type_handlers.ArrayHandler): - """A TypeHandler for array types when using Pathways.""" - - def __init__( - self, - timeout: datetime.timedelta | None = None, - use_ocdbt: bool = False, - array_metadata_store: array_metadata_store_lib.Store | None = None, - ): - """Orbax array handler for Pathways on Cloud with Persistence API. - - Args: - timeout: Duration indicating the timeout for reading and writing arrays. - Default is 1 hour. - use_ocdbt: allows using Tensorstore OCDBT driver. - array_metadata_store: An optional store for writing and reading array - metadata. Only required for saving new-style jax random keys. - """ - if timeout is None: - timeout = datetime.timedelta(hours=1) - self.timeout = timeout - - if use_ocdbt: - raise ValueError("OCDBT not supported for Pathways.") - super().__init__(array_metadata_store=array_metadata_store) - - async def _background_serialize( - self, - futures_results: Sequence[concurrent.futures.Future[None]], - metadata_coroutine: Coroutine[Any, Any, None] | None = None, - ) -> None: - if metadata_coroutine: - await metadata_coroutine - - for future_result in futures_results: - future_result.result() - - def _wait_for_directory_creation_signals(self): - async def _no_op(): - pass - - # Wait for directory creation signals to be set. - future.CommitFutureAwaitingContractedSignals(_no_op()).result() - - async def serialize( - self, - values: Sequence[jax.Array], - infos: Sequence[ParamInfo], - args: Sequence[SaveArgs] | None = None, - ) -> Sequence[future.Future]: - """Uses Pathways Persistence API to serialize a jax array.""" - type_handlers.check_input_arguments(values, infos, args) - - if any([arg.dtype is not None for arg in args]): # pyrefly: ignore[not-iterable] - raise ValueError("Casting during save not supported for Pathways.") - - array_metadatas = [] - any_random_key = False - arrays = [] - for v, info, arg in zip(values, infos, args): # pyrefly: ignore[bad-argument-type] - ext_metadata = None - if jax.dtypes.issubdtype(v.dtype, jax.dtypes.prng_key): - any_random_key = True - impl = str(jax.random.key_impl(v)) - v = jax.random.key_data(v) - ext_metadata = {array_metadata_lib.RANDOM_KEY_IMPL: impl} - - array_metadatas.append( - ArrayMetadata( - param_name=info.name, - shape=v.shape, - dtype=(arg.dtype if arg is not None else v.dtype), # pyrefly: ignore[bad-argument-type] - write_shape=getattr(v, "local_shape", v.shape), - chunk_shape=getattr(v, "local_shape", v.shape), - use_ocdbt=False, - use_zarr3=False, - ext_metadata=ext_metadata, - ) - ) - arrays.append(v) - - if any_random_key and self._array_metadata_store is None: - raise ValueError( - "Array metadata store is not set with a checkpoint that requires" - f" it. Array metadata: {array_metadatas}" - ) - - metadata_coroutine = None - if self._array_metadata_store is not None: - metadata_coroutine = self._array_metadata_store.write( - checkpoint_dir=infos[0].parent_dir, - array_metadatas=array_metadatas, - process_index=0, - ) - - self._wait_for_directory_creation_signals() - locations, names = extract_parent_dir_and_name(infos) - # Group arrays by parent directory and device assignment so each batch - # satisfies SideChannelLoadedExecutable bulk persistence constraints. - grouped_writes: dict[ - tuple[str, tuple[Any, ...]], tuple[list[str], list[jax.Array]] - ] = collections.defaultdict(lambda: ([], [])) - for loc, name, arr in zip(locations, names, arrays): - # pylint:disable=protected-access - key = (loc, tuple(arr.sharding._device_assignment)) - # pylint:enable=protected-access - grouped_writes[key][0].append(name) - grouped_writes[key][1].append(arr) - - futures_results = [ - helper.write_arrays( - loc, group_names, group_arrays, timeout=self.timeout - ) - for (loc, _), (group_names, group_arrays) in grouped_writes.items() - ] - - return [ - future.CommitFutureAwaitingContractedSignals( - self._background_serialize(futures_results, metadata_coroutine), - name="cloud_pathways_array_handler", - ) - ] - - async def deserialize( - self, - infos: Sequence[ParamInfo], - args: Sequence[RestoreArgs] | None = None, - ) -> Sequence[jax.Array]: - """Uses Pathways Persistence API to deserialize a jax array.""" - if args is None: - raise ValueError("Must provide ArrayRestoreArgs to restore as jax.Array.") - type_handlers.check_input_arguments(infos, args) - - global_meshes = [] - mesh_axes = [] - global_shapes = [] - dtypes = [] - shardings = [] - - should_open_metadata = False - for arg in args: - if not isinstance(arg, ArrayRestoreArgs): - raise ValueError( - "To restore jax.Array, provide ArrayRestoreArgs; found" - f" {type(arg).__name__}" - ) - arg = cast(ArrayRestoreArgs, arg) - if arg.sharding is None and (arg.mesh is None or arg.mesh_axes is None): - raise ValueError( - "Sharding of jax.Array cannot be None. Provide `mesh`" - " and `mesh_axes` OR `sharding`." - ) - if arg.sharding is None: - global_meshes.append(arg.mesh) - mesh_axes.append(arg.mesh_axes) - shardings.append( - jax.sharding.NamedSharding(mesh=arg.mesh, spec=arg.mesh_axes) # pyrefly: ignore[bad-argument-type] - ) - else: - if not isinstance(arg.sharding, jax.sharding.NamedSharding): - raise ValueError("Pathways only supports jax.sharding.NamedSharding.") - sharding = cast(jax.sharding.NamedSharding, arg.sharding) - global_meshes.append(sharding.mesh) - mesh_axes.append(sharding.spec) - shardings.append(sharding) - if arg.global_shape is None or arg.dtype is None: - _logger.warning( - "Shape or dtype not provided for restoration. Provide these" - " properties for improved performance." - ) - should_open_metadata = True - global_shapes.append(arg.global_shape) - dtypes.append(arg.dtype) - - if should_open_metadata: - metadatas = await self.metadata(infos) - global_shapes = [ - m.shape if s is None else s for m, s in zip(metadatas, global_shapes) - ] - dtypes = [m.dtype if d is None else d for m, d in zip(metadatas, dtypes)] - - array_metadatas_cache = {} - if self._array_metadata_store is not None: - if array_metadatas := await self._array_metadata_store.read( - checkpoint_dir=infos[0].parent_dir, - process_index=0, - ): - if not isinstance(array_metadatas, list): - raise ValueError( - "Array metadata store returned unexpected result:" - f" {array_metadatas}" - ) - - array_metadatas_cache = { - array_metadata.param_name: array_metadata - for array_metadata in array_metadatas - } - - # Group inputs by parent_dir and global_mesh so that we can perform batched - # Array construction for each group. - inputs_by_location_and_mesh = collections.defaultdict(list) - for i, (info, global_mesh) in enumerate(zip(infos, global_meshes)): - inputs_by_location_and_mesh[(str(info.parent_dir), global_mesh)].append(i) - - results = cast(list[jax.Array], [None] * len(infos)) - - for (location, global_mesh), idxs in inputs_by_location_and_mesh.items(): - grouped_infos = [infos[idx] for idx in idxs] - grouped_global_shapes = [global_shapes[idx] for idx in idxs] - grouped_dtypes = [dtypes[idx] for idx in idxs] - grouped_shardings = [shardings[idx] for idx in idxs] - _, names = extract_parent_dir_and_name(grouped_infos) - - grouped_read_dtypes = [] - grouped_read_shapes = [] - grouped_read_shardings = [] - - # Typed PRNG keys (e.g. key) are written physically as raw key_data - # (uint32). Transform logical shapes, dtypes, and shardings into physical - # equivalents for persistence read requests. - for shape, dtype, sharding in zip( - grouped_global_shapes, grouped_dtypes, grouped_shardings - ): - if jax.dtypes.issubdtype(dtype, jax.dtypes.prng_key): - phys_struct = jax.eval_shape( - jax.random.key_data, jax.ShapeDtypeStruct(shape, dtype) - ) - read_dtype = phys_struct.dtype - read_shape = phys_struct.shape - # Append unpartitioned (None) specs for physical trailing dimensions - # added by key_data. - trailing_ndim = len(read_shape) - len(shape) - if isinstance(sharding, jax.sharding.NamedSharding): - read_sharding = jax.sharding.NamedSharding( - sharding.mesh, - jax.sharding.PartitionSpec( - *sharding.spec, *([None] * trailing_ndim) - ), - ) - else: - read_sharding = sharding - else: - read_dtype = dtype - read_shape = shape - read_sharding = sharding - - grouped_read_dtypes.append(read_dtype) - grouped_read_shapes.append(read_shape) - grouped_read_shardings.append(read_sharding) - - grouped_arrays, read_future = helper.read_arrays( - location, - names, - grouped_read_dtypes, - grouped_read_shapes, - grouped_read_shardings, - global_mesh.devices, - timeout=self.timeout, - ) - # each persistence call is awaited serially. - read_future.result() - for idx, info, arr in zip(idxs, grouped_infos, grouped_arrays): - orig_dtype = dtypes[idx] - # Re-wrap physical key_data uint32 arrays into target PRNG key. - if jax.dtypes.issubdtype(orig_dtype, jax.dtypes.prng_key): - arr = jax.random.wrap_key_data( - arr, dtype=orig_dtype # pyrefly: ignore[bad-argument-type] - ) - elif meta := array_metadatas_cache.get(info.name): - - assert isinstance( - meta, array_metadata_lib.SerializedArrayMetadata - ), f"Expecting SerializedArrayMetadata but got {type(meta)}." - if meta.ext_metadata: - assert isinstance(meta.ext_metadata, dict), ( - "Expecting ext_metadata to be a dict but got" - f" {type(meta.ext_metadata)}." - ) - - if impl := meta.ext_metadata.get( - array_metadata_lib.RANDOM_KEY_IMPL - ): - arr = jax.random.wrap_key_data(arr, impl=impl) - results[idx] = arr - - return results - - -def register_pathways_handlers( - timeout: datetime.timedelta | None = None, - array_metadata_store: array_metadata_store_lib.Store | None = None, -): - """Function that must be called before saving or restoring with Pathways.""" - _logger.debug( - "Registering CloudPathwaysArrayHandler (Pathways Persistence API)." - ) - type_handlers.register_type_handler( - jax.Array, - CloudPathwaysArrayHandler( - timeout=timeout, - array_metadata_store=array_metadata_store, - ), - override=True, - ) diff --git a/pathwaysutils/test/initialize_test.py b/pathwaysutils/test/initialize_test.py index 31eaa0f..9228a49 100644 --- a/pathwaysutils/test/initialize_test.py +++ b/pathwaysutils/test/initialize_test.py @@ -17,6 +17,8 @@ from absl.testing import absltest from absl.testing import parameterized import jax +from orbax.checkpoint import type_handlers +from orbax.checkpoint._src.serialization import cloud_pathways_array_handler from pathwaysutils import _initialize @@ -99,6 +101,19 @@ def test_persistence_enabled(self): del os.environ["ENABLE_PATHWAYS_PERSISTENCE"] self.assertFalse(_initialize._is_persistence_enabled()) + def test_initialize_registers_orbax_pathways_handler_when_persistence_enabled( + self, + ): + jax.config.update("jax_platforms", "proxy") + os.environ["ENABLE_PATHWAYS_PERSISTENCE"] = "1" + _initialize._initialization_count = 0 + + _initialize.initialize() + handler = type_handlers.get_type_handler(jax.Array) + self.assertIsInstance( + handler, cloud_pathways_array_handler.CloudPathwaysArrayHandler + ) + if __name__ == "__main__": absltest.main()