From d7f330771d829725a92d7c47a5212a17dd384826 Mon Sep 17 00:00:00 2001 From: Daniel Ng Date: Wed, 9 Sep 2026 09:38:54 -0700 Subject: [PATCH] Internal PiperOrigin-RevId: 978599342 --- .../_src/serialization/jax_array_handlers.py | 65 +++++---- .../_src/serialization/type_handlers_test.py | 125 ++++++++++++++++++ 2 files changed, 161 insertions(+), 29 deletions(-) diff --git a/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py b/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py index 97245c148..0cea6f73b 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py @@ -246,14 +246,16 @@ def _record_logical_metrics( custom_prefix: str = '', ): """Records logical bytes, throughput, and duration to JAX monitoring.""" - logical_throughput = logical_bytes / duration if duration > 0 else 0 + logical_throughput = logical_bytes / duration if duration > 0 else None logging.info( '[process=%d] %s throughput: %s/s (total gbytes: %s) (time elapsed: %s s)' ' (per-host)%s', multihost.process_index(), f'/jax/orbax/{direction.value}/worker/io/requested', - humanize.naturalsize(logical_throughput, binary=True, format='%.3f'), + humanize.naturalsize(logical_throughput, binary=True, format='%.3f') + if logical_throughput is not None + else 'N/A', humanize.naturalsize(logical_bytes, binary=True), duration, f' (prefix: {custom_prefix})' if custom_prefix else '', @@ -272,12 +274,14 @@ def _record_logical_metrics( storage_type=storage_type, custom_prefix=custom_prefix, ) - jax.monitoring.record_scalar( - f'/jax/orbax/{direction.value}/worker/io/requested/throughput/gbytes_per_sec', - logical_throughput / (1024**3), - storage_type=storage_type, - custom_prefix=custom_prefix, - ) + + if logical_throughput is not None: + jax.monitoring.record_scalar( + f'/jax/orbax/{direction.value}/worker/io/requested/throughput/gbytes_per_sec', + logical_throughput / (1024**3), + storage_type=storage_type, + custom_prefix=custom_prefix, + ) def _record_compression_metrics( @@ -289,20 +293,20 @@ def _record_compression_metrics( metadatas: Sequence[ts_utils.ArrayMetadata] | None = None, ) -> None: """Logs and records compression ratio and metrics.""" - if logical_bytes <= 0: + if logical_bytes <= 0 or raw_bytes <= 0: return - ratio = float(raw_bytes) / logical_bytes + ratio = float(logical_bytes) / raw_bytes algo_str = 'none' level_str = 'None' if metadatas is not None: algo_str, level_str = ts_utils.resolve_compression_settings(metadatas) logging.info( - '[process=%d] %s ratio (raw/logical): %.3f (%s / %s), algo=%s, level=%s', + '[process=%d] %s ratio (logical/raw): %.3f (%s / %s), algo=%s, level=%s', multihost.process_index(), direction.value.capitalize(), ratio, - humanize.naturalsize(raw_bytes, binary=True), humanize.naturalsize(logical_bytes, binary=True), + humanize.naturalsize(raw_bytes, binary=True), algo_str, level_str, ) @@ -314,15 +318,6 @@ def _record_compression_metrics( compression_algorithm=algo_str, compression_level=level_str, ) - if direction == types.IoDirection.WRITE: - jax.monitoring.record_scalar( - '/jax/orbax/write/worker/io/compressed_gbytes', - raw_bytes / (1024**3), - storage_type=storage_type, - custom_prefix=custom_prefix, - compression_algorithm=algo_str, - compression_level=level_str, - ) def _record_raw_metrics( @@ -343,29 +338,41 @@ def _record_raw_metrics( if raw_bytes <= 0: return - raw_throughput = raw_bytes / duration if duration > 0 else 0 + raw_throughput = raw_bytes / duration if duration > 0 else None logging.info( '[process=%d] Raw %s throughput: %s/s (total gbytes: %s) (time elapsed:' ' %s s) (per-host)%s', multihost.process_index(), f'/jax/orbax/{direction.value}/worker/io/raw', - humanize.naturalsize(raw_throughput, binary=True, format='%.3f'), + humanize.naturalsize(raw_throughput, binary=True, format='%.3f') + if raw_throughput is not None + else 'N/A', humanize.naturalsize(raw_bytes, binary=True), duration, f' (prefix: {custom_prefix})' if custom_prefix else '', ) + algo_str = 'none' + level_str = 'None' + if metadatas is not None: + algo_str, level_str = ts_utils.resolve_compression_settings(metadatas) + # raw/gbytes records the actual physical raw bytes transferred to/from the + # storage backend, including compression (e.g. zstd) when enabled or + # uncompressed chunk bytes when disabled. jax.monitoring.record_scalar( f'/jax/orbax/{direction.value}/worker/io/raw/gbytes', raw_bytes / (1024**3), storage_type=storage_type, custom_prefix=custom_prefix, + compression_algorithm=algo_str, + compression_level=level_str, ) - jax.monitoring.record_scalar( - f'/jax/orbax/{direction.value}/worker/io/raw/throughput/gbytes_per_sec', - raw_throughput / (1024**3), - storage_type=storage_type, - custom_prefix=custom_prefix, - ) + if raw_throughput is not None: + jax.monitoring.record_scalar( + f'/jax/orbax/{direction.value}/worker/io/raw/throughput/gbytes_per_sec', + raw_throughput / (1024**3), + storage_type=storage_type, + custom_prefix=custom_prefix, + ) _record_compression_metrics( direction, diff --git a/checkpoint/orbax/checkpoint/_src/serialization/type_handlers_test.py b/checkpoint/orbax/checkpoint/_src/serialization/type_handlers_test.py index e6dec1f53..a54dd86d1 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/type_handlers_test.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/type_handlers_test.py @@ -13,6 +13,7 @@ # limitations under the License. import asyncio +import collections import dataclasses import threading import tracemalloc @@ -1418,5 +1419,129 @@ def test_replace_updates_multiple_fields(self): self.assertEqual(info.parent_dir, path) +class RecordCompressionMetricsTest(parameterized.TestCase): + + def test_record_compression_metrics_calculation(self): + recorded_scalars = {} + + def _mock_record_scalar(name: str, value: Any, *_args, **_kwargs): + recorded_scalars[name] = value + + with mock.patch( + 'jax.monitoring.record_scalar', side_effect=_mock_record_scalar + ): + # Canonical ratio: logical / raw = 200 / 100 = 2.0 + jax_array_handlers._record_compression_metrics( + direction=types.IoDirection.WRITE, + logical_bytes=200, + raw_bytes=100, + storage_type='gcs', + ) + self.assertEqual( + recorded_scalars['/jax/orbax/write/worker/io/compression_ratio'], 2.0 + ) + + # Read direction canonical ratio: 300 / 150 = 2.0 + recorded_scalars.clear() + jax_array_handlers._record_compression_metrics( + direction=types.IoDirection.READ, + logical_bytes=300, + raw_bytes=150, + storage_type='gcs', + ) + self.assertEqual( + recorded_scalars['/jax/orbax/read/worker/io/compression_ratio'], 2.0 + ) + + # Zero or negative inputs should not record + recorded_scalars.clear() + jax_array_handlers._record_compression_metrics( + direction=types.IoDirection.WRITE, + logical_bytes=0, + raw_bytes=100, + storage_type='gcs', + ) + self.assertNotIn( + '/jax/orbax/write/worker/io/compression_ratio', recorded_scalars + ) + + recorded_scalars.clear() + jax_array_handlers._record_compression_metrics( + direction=types.IoDirection.WRITE, + logical_bytes=100, + raw_bytes=0, + storage_type='gcs', + ) + self.assertNotIn( + '/jax/orbax/write/worker/io/compression_ratio', recorded_scalars + ) + + # Ensure write/worker/io/compressed_gbytes is never recorded + recorded_scalars.clear() + jax_array_handlers._record_compression_metrics( + direction=types.IoDirection.WRITE, + logical_bytes=200, + raw_bytes=100, + storage_type='gcs', + ) + self.assertNotIn( + '/jax/orbax/write/worker/io/compressed_gbytes', recorded_scalars + ) + + def test_record_raw_metrics_records_compression_tags(self): + recorded_scalars = {} + recorded_kwargs = collections.defaultdict(list) + + def _mock_record_scalar(name: str, value: Any, *_args, **kwargs): + recorded_scalars[name] = value + recorded_kwargs[name].append(kwargs) + + with mock.patch( + 'jax.monitoring.record_scalar', side_effect=_mock_record_scalar + ), mock.patch.object( + jax_array_handlers.ts_utils, + 'get_tensorstore_raw_bytes', + return_value=1024**3, + ): + # Write direction + jax_array_handlers._record_raw_metrics( + direction=types.IoDirection.WRITE, + logical_bytes=2 * (1024**3), + duration=1.0, + storage_type='gcs', + initial_raw_bytes=0, + ) + self.assertEqual( + recorded_scalars['/jax/orbax/write/worker/io/raw/gbytes'], 1.0 + ) + write_raw_kwargs = recorded_kwargs[ + '/jax/orbax/write/worker/io/raw/gbytes' + ][0] + self.assertEqual(write_raw_kwargs.get('compression_algorithm'), 'none') + self.assertEqual(write_raw_kwargs.get('compression_level'), 'None') + self.assertNotIn( + '/jax/orbax/write/worker/io/compressed_gbytes', recorded_scalars + ) + + # Read direction + recorded_scalars.clear() + recorded_kwargs.clear() + jax_array_handlers._record_raw_metrics( + direction=types.IoDirection.READ, + logical_bytes=2 * (1024**3), + duration=1.0, + storage_type='gcs', + initial_raw_bytes=0, + ) + self.assertEqual( + recorded_scalars['/jax/orbax/read/worker/io/raw/gbytes'], 1.0 + ) + read_raw_kwargs = recorded_kwargs[ + '/jax/orbax/read/worker/io/raw/gbytes' + ][0] + self.assertEqual(read_raw_kwargs.get('compression_algorithm'), 'none') + self.assertEqual(read_raw_kwargs.get('compression_level'), 'None') + + if __name__ == '__main__': multiprocess_test.main()