From 505cfd4f828a0d9188e42259674798c1051c6fb8 Mon Sep 17 00:00:00 2001 From: Daniel Ng Date: Fri, 4 Sep 2026 05:22:57 -0700 Subject: [PATCH] Internal PiperOrigin-RevId: 976265655 --- .../_src/serialization/jax_array_handlers.py | 35 +-- .../_src/serialization/tensorstore_utils.py | 53 ++-- .../serialization/tensorstore_utils_test.py | 242 ++++++------------ 3 files changed, 113 insertions(+), 217 deletions(-) diff --git a/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py b/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py index 20e71a057..97245c148 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py @@ -330,21 +330,16 @@ def _record_raw_metrics( logical_bytes: int, duration: float, storage_type: str, - initial_ts_metrics: Sequence[dict[str, Any]] | None = None, + initial_raw_bytes: int | None = None, custom_prefix: str = '', metadatas: Sequence[ts_utils.ArrayMetadata] | None = None, ): """Records raw metrics collected from TensorStore.""" - if initial_ts_metrics is None: + if initial_raw_bytes is None: return - final_ts_metrics = ts_utils.collect_tensorstore_metrics() - if final_ts_metrics is None: - return - - raw_bytes = ts_utils.get_tensorstore_raw_bytes_delta( - initial_ts_metrics, final_ts_metrics, direction - ) + final_raw_bytes = ts_utils.get_tensorstore_raw_bytes(direction) + raw_bytes = max(0, final_raw_bytes - initial_raw_bytes) if raw_bytes <= 0: return @@ -387,7 +382,7 @@ def _log_io_metrics( logical_bytes: int, start_time: float, parent_dir: epath.Path, - initial_ts_metrics: Sequence[dict[str, Any]] | None = None, + initial_raw_bytes: int | None = None, custom_prefix: str = '', metadatas: Sequence[ts_utils.ArrayMetadata] | None = None, ): @@ -407,7 +402,7 @@ def _log_io_metrics( logical_bytes, duration, storage_type, - initial_ts_metrics=initial_ts_metrics, + initial_raw_bytes=initial_raw_bytes, custom_prefix=custom_prefix, metadatas=metadatas, ) @@ -428,7 +423,9 @@ def _worker_serialize_arrays( ext_metadata: Dict[str, Any], ): """Worker function to serialize arrays.""" - initial_ts_metrics = ts_utils.collect_tensorstore_metrics() + initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes( + types.IoDirection.WRITE + ) total_start_time = time.time() rslices_per_array = _get_replica_slices( arrays, @@ -458,7 +455,7 @@ def _worker_serialize_arrays( logical_bytes=total_io_bytes, start_time=total_start_time, parent_dir=infos[0].parent_dir, - initial_ts_metrics=initial_ts_metrics, + initial_raw_bytes=initial_raw_bytes, metadatas=array_metadatas, ) @@ -571,7 +568,9 @@ def _serialize_arrays_batches_without_dispatcher( async def _serialize_without_dispatcher(): if not prioritized and not deprioritized: return - initial_ts_metrics = ts_utils.collect_tensorstore_metrics() + initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes( + types.IoDirection.WRITE + ) total_start_time = time.time() logical_bytes = 0 all_array_metadatas: list[ts_utils.ArrayMetadata] = [] @@ -619,7 +618,7 @@ async def _serialize_without_dispatcher(): logical_bytes=logical_bytes, start_time=total_start_time, parent_dir=info_sample.parent_dir, - initial_ts_metrics=initial_ts_metrics, + initial_raw_bytes=initial_raw_bytes, metadatas=all_array_metadatas, ) @@ -1037,7 +1036,9 @@ async def _deserialize_arrays( array_metadata_store: array_metadata_store_lib.Store | None, ) -> Sequence[jax.Array]: """Deserializes arrays and applies array_metadata if available.""" - initial_ts_metrics = ts_utils.collect_tensorstore_metrics() + initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes( + types.IoDirection.READ + ) total_start_time = time.time() async def _async_deserialize( @@ -1141,7 +1142,7 @@ async def _async_deserialize( logical_bytes=logical_bytes, start_time=total_start_time, parent_dir=infos[0].parent_dir, - initial_ts_metrics=initial_ts_metrics, + initial_raw_bytes=initial_raw_bytes, ) return ret diff --git a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py index 551dd4b94..d67c611f5 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py @@ -1155,54 +1155,35 @@ def array_metadata_from_tensorstore( ) -def get_total_bytes_from_tensorstore( - metrics: Sequence[dict[str, Any]], direction: types.IoDirection +def get_tensorstore_raw_bytes( + direction: types.IoDirection = types.IoDirection.WRITE, ) -> int: - """Sums bytes_read or bytes_written from all kvstore drivers in metrics.""" - total = 0 - if direction == types.IoDirection.WRITE: - suffix = '/bytes_written' - elif direction == types.IoDirection.READ: - suffix = '/bytes_read' - else: - raise ValueError(f'Invalid direction: {direction}') + """Collects and returns total raw bytes read or written by TensorStore.""" + suffix = ( + '/bytes_written' + if direction == types.IoDirection.WRITE + else '/bytes_read' + ) + # Querying `/tensorstore/kvstore/` returns only a few kvstore driver + # metrics (e.g. file, gcs, s3), making metric collection and + # extraction fast without scanning all TensorStore metrics. + metrics: Sequence[dict[str, Any]] = ts.experimental_collect_matching_metrics( + '/tensorstore/kvstore/' + ) + # Sum the metric values for the given suffix. + total = 0 for m in metrics: if not isinstance(m, dict): continue name = m.get('name', '') - if name.startswith('/tensorstore/kvstore/') and name.endswith(suffix): + if name.endswith(suffix): for val in m.get('values', []): if isinstance(val, dict): total += val.get('value', 0) return total -def get_tensorstore_raw_bytes_delta( - initial_metrics: Sequence[dict[str, Any]] | None, - final_metrics: Sequence[dict[str, Any]] | None, - direction: types.IoDirection = types.IoDirection.WRITE, -) -> int: - """Computes transferred raw bytes delta between two metric snapshots.""" - if initial_metrics is None or final_metrics is None: - return 0 - try: - initial_bytes = get_total_bytes_from_tensorstore(initial_metrics, direction) - final_bytes = get_total_bytes_from_tensorstore(final_metrics, direction) - return max(0, final_bytes - initial_bytes) - except Exception: # pylint: disable=broad-except - logging.exception('Failed to compute TensorStore raw bytes delta.') - return 0 - - -def collect_tensorstore_metrics() -> Sequence[dict[str, Any]] | None: - """Safely collects TensorStore driver metrics.""" - try: - return ts.experimental_collect_matching_metrics('/tensorstore') - except Exception: # pylint: disable=broad-except - return None - - def resolve_compression_settings( metadatas: Sequence[ArrayMetadata], ) -> tuple[str, str]: diff --git a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py index c807e4100..caab8eda2 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py @@ -16,6 +16,7 @@ import math import os import tempfile +from typing import Any import unittest from absl.testing import absltest @@ -1185,138 +1186,6 @@ def test_get_ts_context( self.assertDictEqual(expected_spec, context.spec.to_json()) -class GetTotalBytesFromTensorstoreTest(parameterized.TestCase): - - def test_get_total_bytes_written(self): - metrics = [ - { - 'name': '/tensorstore/kvstore/gcs/bytes_written', - 'values': [{'value': 100}, {'value': 200}], - }, - { - 'name': '/tensorstore/kvstore/gfile/bytes_written', - 'values': [{'value': 50}], - }, - { - 'name': '/tensorstore/kvstore/gcs/bytes_read', - 'values': [{'value': 500}], - }, - { - 'name': '/other/metric/bytes_written', - 'values': [{'value': 1000}], - }, - ] - self.assertEqual( - ts_utils.get_total_bytes_from_tensorstore( - metrics, serialization_types.IoDirection.WRITE - ), - 350, - ) - - def test_get_total_bytes_read(self): - metrics = [ - { - 'name': '/tensorstore/kvstore/gcs/bytes_written', - 'values': [{'value': 100}], - }, - { - 'name': '/tensorstore/kvstore/gcs/bytes_read', - 'values': [{'value': 500}, {'value': 250}], - }, - { - 'name': '/tensorstore/kvstore/gfile/bytes_read', - 'values': [{'value': 50}], - }, - ] - self.assertEqual( - ts_utils.get_total_bytes_from_tensorstore( - metrics, serialization_types.IoDirection.READ - ), - 800, - ) - - @parameterized.named_parameters( - ('with_compression', True), - ('without_compression', False), - ) - def test_get_total_bytes_with_real_ops(self, use_compression): - initial_metrics_write = ts.experimental_collect_matching_metrics( - '/tensorstore/' - ) - - tempdir = self.create_tempdir().full_path - info = serialization_types.ParamInfo( - name='arr', - parent_dir=epath.Path(tempdir), - use_compression=use_compression, - ) - write_spec = ts_utils.build_array_write_spec( - info, - global_shape=(100000,), - local_shape=(100000,), - dtype=np.dtype(np.int32), - use_ocdbt=True, - ) - write_context = ts_utils.get_ts_context(use_ocdbt=True) - store = ts.open( - write_spec.json, - context=write_context, - create=True, - delete_existing=True, - dtype=np.int32, - shape=(100000,), - ).result() - store.write(np.arange(100000, dtype=np.int32)).result() - - final_metrics_write = ts.experimental_collect_matching_metrics( - '/tensorstore/' - ) - - bytes_written = ts_utils.get_total_bytes_from_tensorstore( - final_metrics_write, serialization_types.IoDirection.WRITE - ) - ts_utils.get_total_bytes_from_tensorstore( - initial_metrics_write, serialization_types.IoDirection.WRITE - ) - - self.assertGreater(bytes_written, 0) - - initial_metrics_read = ts.experimental_collect_matching_metrics( - '/tensorstore/' - ) - - read_spec = ts_utils.build_array_read_spec(info, use_ocdbt=True) - read_context = ts_utils.get_ts_context(use_ocdbt=True) - read_store = ts.open( - read_spec.json, - open=True, - context=read_context, - dtype=np.int32, - shape=(100000,), - ).result() - read_store.read().result() - - final_metrics_read = ts.experimental_collect_matching_metrics( - '/tensorstore/' - ) - - bytes_read = ts_utils.get_total_bytes_from_tensorstore( - final_metrics_read, serialization_types.IoDirection.READ - ) - ts_utils.get_total_bytes_from_tensorstore( - initial_metrics_read, serialization_types.IoDirection.READ - ) - - self.assertGreater(bytes_read, 0) - - if not use_compression: - # Logical size is 100000 * 4 bytes = 400000 bytes. - self.assertLess(abs(bytes_written - 400000) / 400000, 0.01) - self.assertLess(abs(bytes_read - 400000) / 400000, 0.01) - else: - # Verify that compression actually reduced the bytes written/read. - self.assertLess(bytes_written, 300000) - self.assertLess(bytes_read, 300000) - - def _is_using_file_driver() -> bool: return ts_utils.DEFAULT_DRIVER == 'file' @@ -1512,47 +1381,92 @@ def test_commit_temporary_metadata_mode( self._verify_kvstack_spec(kvstore_tspec['base'], expected_base_path) -class GetTensorStoreRawBytesDeltaTest(parameterized.TestCase): +class GetTensorStoreRawBytesTest(parameterized.TestCase): - def test_none_metrics(self): - self.assertEqual(ts_utils.get_tensorstore_raw_bytes_delta(None, None), 0) - self.assertEqual(ts_utils.get_tensorstore_raw_bytes_delta([], None), 0) - self.assertEqual(ts_utils.get_tensorstore_raw_bytes_delta(None, []), 0) + def test_get_tensorstore_raw_bytes(self): + bytes_written = ts_utils.get_tensorstore_raw_bytes( + serialization_types.IoDirection.WRITE + ) + self.assertIsInstance(bytes_written, int) + self.assertGreaterEqual(bytes_written, 0) - def test_delta_calculation(self): - initial = [{ - 'name': '/tensorstore/kvstore/ocdbt/bytes_written', - 'values': [{'value': 100}], - }] - final = [{ - 'name': '/tensorstore/kvstore/ocdbt/bytes_written', - 'values': [{'value': 350}], - }] - delta = ts_utils.get_tensorstore_raw_bytes_delta( - initial, final, serialization_types.IoDirection.WRITE + bytes_read = ts_utils.get_tensorstore_raw_bytes( + serialization_types.IoDirection.READ ) - self.assertEqual(delta, 250) + self.assertIsInstance(bytes_read, int) + self.assertGreaterEqual(bytes_read, 0) - def test_negative_delta_returns_zero(self): - initial = [{ - 'name': '/tensorstore/kvstore/ocdbt/bytes_written', - 'values': [{'value': 500}], + def test_mock_metrics(self): + fake_metrics = [{ + 'name': '/tensorstore/kvstore/file/bytes_written', + 'values': [{'value': 1234}], }] - final = [{ - 'name': '/tensorstore/kvstore/ocdbt/bytes_written', - 'values': [{'value': 300}], - }] - delta = ts_utils.get_tensorstore_raw_bytes_delta( - initial, final, serialization_types.IoDirection.WRITE + with unittest.mock.patch.object( + ts, 'experimental_collect_matching_metrics', return_value=fake_metrics + ): + self.assertEqual( + ts_utils.get_tensorstore_raw_bytes( + serialization_types.IoDirection.WRITE + ), + 1234, + ) + + @parameterized.product( + use_ocdbt=(False, True), + use_zarr3=(False, True), + ) + def test_live_tensorstore_configurations(self, use_ocdbt, use_zarr3): + tmpdir = tempfile.mkdtemp() + driver = 'zarr3' if use_zarr3 else 'zarr' + kvstore_spec = {'driver': 'file', 'path': tmpdir} + if use_ocdbt: + kvstore_spec = {'driver': 'ocdbt', 'base': kvstore_spec} + + metadata: dict[str, Any] = {'shape': [100, 100]} + if use_zarr3: + metadata['data_type'] = 'float32' + metadata['chunk_grid'] = { + 'name': 'regular', + 'configuration': {'chunk_shape': [50, 50]}, + } + else: + metadata['dtype'] = '