From fd12c46168bfbea88689d360929a4ab9f9168d99 Mon Sep 17 00:00:00 2001 From: Adam Cogdell Date: Fri, 14 Aug 2026 13:13:36 -0700 Subject: [PATCH] No public description PiperOrigin-RevId: 964853113 --- .../_src/checkpointers/async_checkpointer.py | 21 ++++++++-- .../checkpointers/async_checkpointer_test.py | 38 +++++++++++++++++++ .../_src/checkpointers/checkpointer.py | 7 +++- .../_src/checkpointers/checkpointer_test.py | 25 ++++++++++++ .../base_pytree_checkpoint_handler.py | 8 ++-- .../handlers/composite_checkpoint_handler.py | 2 +- .../composite_checkpoint_handler_test.py | 20 ++++++++++ .../checkpoint/_src/logging/event_tracking.py | 4 +- .../orbax/checkpoint/_src/path/atomicity.py | 2 +- .../orbax/checkpoint/checkpoint_manager.py | 28 ++++---------- .../experimental/v1/_src/path/async_utils.py | 2 +- 11 files changed, 124 insertions(+), 33 deletions(-) diff --git a/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer.py b/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer.py index 7e49b11e98..58f053837f 100644 --- a/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer.py +++ b/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer.py @@ -105,7 +105,7 @@ def _background_wait_for_commit_futures( commit_duration_secs, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/tensorstore_duration_secs', + '/jax/orbax/write/background_ts_duration_secs', commit_duration_secs, ) @@ -161,6 +161,10 @@ def _background_wait_for_commit_futures( '/jax/checkpoint/write/async/thread_duration_sec', thread_duration_secs, ) + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/background_duration_secs', + thread_duration_secs, + ) logging.info( '[process=%s][thread=%s] Background save thread done. Time taken: %fs.', current_process, @@ -484,7 +488,7 @@ def _callback() -> None: checkpoint_start_time, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/finalize_duration_secs', + '/jax/orbax/write/background_finalize_secs', time.time() - finalize_start_time, ) operation_recorder = event_tracking.OperationRecorder( @@ -605,7 +609,13 @@ def save( ) operation_recorder.record_start(start_time=checkpoint_start_time) tmpdir = self.get_temporary_path(directory) + wait_prev_start_time = time.perf_counter() self.wait_until_finished() + wait_prev_duration_secs = time.perf_counter() - wait_prev_start_time + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/blocking_wait_prev_duration_secs', + wait_prev_duration_secs, + ) self.synchronize_next_awaitable_signal_operation_id() on_commit_callback = self._make_on_commit_callback( tmpdir, custom_metadata, checkpoint_start_time @@ -619,8 +629,13 @@ def save( ) ) blocking_end_time = time.time() + blocking_duration_secs = blocking_end_time - checkpoint_start_time + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/blocking_duration_secs', + blocking_duration_secs, + ) operation_recorder.record_blocking_completion( - blocking_end_time - checkpoint_start_time, + blocking_duration_secs, end_time=blocking_end_time, ) self._async_manager.start_async_commit( diff --git a/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer_test.py b/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer_test.py index d4a15d4527..66f0ecf891 100644 --- a/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer_test.py +++ b/checkpoint/orbax/checkpoint/_src/checkpointers/async_checkpointer_test.py @@ -512,6 +512,44 @@ def update(self, file_path: typing.PathLike, **kwargs: Any): # TODO(nikhilbansall): Open source this test. + def test_save_metrics(self): + handler = PyTreeCheckpointHandler() + checkpointer = AsyncCheckpointer(handler) + with mock.patch( + 'jax.monitoring.record_event_duration_secs' + ) as mock_record_duration: + checkpointer.save( + self.directory, args=PyTreeSaveArgs(item=self.pytree) + ) + checkpointer.wait_until_finished() + + recorded_metrics = [ + call[0][0] for call in mock_record_duration.call_args_list + ] + self.assertIn( + '/jax/orbax/write/blocking_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/blocking_wait_prev_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/blocking_tree_map_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/blocking_d2h_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/background_ts_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/background_duration_secs', recorded_metrics + ) + if multihost.is_primary_host(checkpointer._primary_host): + self.assertIn( + '/jax/orbax/write/background_finalize_secs', recorded_metrics + ) + checkpointer.close() + if __name__ == '__main__': multiprocess_test.main() diff --git a/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer.py b/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer.py index 307020e1ec..115ab3117e 100644 --- a/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer.py +++ b/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer.py @@ -270,8 +270,13 @@ def save( processes=self._active_processes, ) blocking_end_time = time.time() + blocking_duration_secs = blocking_end_time - checkpoint_start_time + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/blocking_duration_secs', + blocking_duration_secs, + ) operation_recorder.record_blocking_completion( - blocking_end_time - checkpoint_start_time, + blocking_duration_secs, blocking_end_time, ) diff --git a/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer_test.py b/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer_test.py index c8736dd269..b65043b1ab 100644 --- a/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer_test.py +++ b/checkpoint/orbax/checkpoint/_src/checkpointers/checkpointer_test.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +from unittest import mock from absl import flags from orbax.checkpoint._src.checkpointers import checkpointer as checkpointer_lib from orbax.checkpoint._src.checkpointers import checkpointer_test_utils @@ -29,6 +30,30 @@ class CheckpointerTest( def checkpointer(self, handler, **kwargs): return checkpointer_lib.Checkpointer(handler, **kwargs) + def test_save_metrics(self): + handler = checkpointer_test_utils.PyTreeCheckpointHandler() + checkpointer = self.checkpointer(handler) + with mock.patch( + 'jax.monitoring.record_event_duration_secs' + ) as mock_record_duration: + checkpointer.save( + self.directory, + args=checkpointer_test_utils.args.PyTreeSave(self.pytree), + ) + recorded_metrics = [ + call[0][0] for call in mock_record_duration.call_args_list + ] + self.assertIn( + '/jax/orbax/write/blocking_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/blocking_tree_map_duration_secs', recorded_metrics + ) + self.assertIn( + '/jax/orbax/write/blocking_d2h_duration_secs', recorded_metrics + ) + checkpointer.close() + if __name__ == '__main__': multiprocess_test.main() diff --git a/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py b/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py index 75a89d4f97..396c231992 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/base_pytree_checkpoint_handler.py @@ -746,11 +746,11 @@ async def async_save( ) ] jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/tree_mapping_duration_secs', + '/jax/orbax/write/blocking_tree_map_duration_secs', batch_requests_ready_time - start_time, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/d2h_duration_secs', + '/jax/orbax/write/blocking_d2h_duration_secs', total_serialization_initiated_time - batch_requests_ready_time, ) async_save_end_time = time.time() @@ -1210,7 +1210,7 @@ async def _write_metadata_file( time.time() - metadata_write_start_time, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/metadata_write_duration_secs', + '/jax/orbax/write/background_meta_write_secs', time.time() - metadata_write_start_time, ) @@ -1382,7 +1382,7 @@ async def merge_ocdbt_per_process_files(): time.time() - merge_start_time, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/ocdbt_merge_duration_secs', + '/jax/orbax/write/background_ocdbt_merge_secs', time.time() - merge_start_time, ) diff --git a/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler.py b/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler.py index 7a750ad3cd..59133b692e 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler.py @@ -1050,7 +1050,7 @@ def finalize(self, directory: epath.Path): handler.finalize(tmp_dir.get()) asyncio_utils.run_sync(tmp_dir.finalize()) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/item_finalize_duration_secs', + '/jax/orbax/write/background_finalize_secs', time.time() - item_finalize_start_time, ) diff --git a/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler_test.py b/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler_test.py index 92855b8efe..491d0c1a8c 100644 --- a/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler_test.py +++ b/checkpoint/orbax/checkpoint/_src/handlers/composite_checkpoint_handler_test.py @@ -1172,6 +1172,26 @@ def test_metadata_with_missing_metadata_file(self): self.assertIn('datasets', step_metadata.item_metadata) self.assertIsNone(step_metadata.item_metadata['datasets']) + def test_finalize_metric(self): + handler = CompositeCheckpointHandler(state=StandardCheckpointHandler()) + state = {'a': 1, 'b': 2} + with mock.patch( + 'jax.monitoring.record_event_duration_secs' + ) as mock_record_duration: + self.save( + handler, + self.directory, + CompositeArgs( + state=args_lib.StandardSave(state), + ), + ) + recorded_metrics = [ + call[0][0] for call in mock_record_duration.call_args_list + ] + self.assertIn( + '/jax/orbax/write/background_finalize_secs', recorded_metrics + ) + if __name__ == '__main__': absltest.main() diff --git a/checkpoint/orbax/checkpoint/_src/logging/event_tracking.py b/checkpoint/orbax/checkpoint/_src/logging/event_tracking.py index 1e0530a4ab..addf9e86e5 100644 --- a/checkpoint/orbax/checkpoint/_src/logging/event_tracking.py +++ b/checkpoint/orbax/checkpoint/_src/logging/event_tracking.py @@ -134,12 +134,12 @@ def record_blocking_completion(self, duration_secs: float, end_time: float): duration_secs, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/blocking_duration_secs', + '/jax/orbax/write/blocking_duration_secs', duration_secs, storage_type=self._storage_type, ) jax.monitoring.record_scalar( - '/jax/orbax/write/async/blocking_end_time', + '/jax/orbax/write/blocking_end_time', _seconds_to_milliseconds(end_time), storage_type=self._storage_type, ) diff --git a/checkpoint/orbax/checkpoint/_src/path/atomicity.py b/checkpoint/orbax/checkpoint/_src/path/atomicity.py index 10693a10a3..5c61bd1f8c 100644 --- a/checkpoint/orbax/checkpoint/_src/path/atomicity.py +++ b/checkpoint/orbax/checkpoint/_src/path/atomicity.py @@ -813,7 +813,7 @@ async def _create_paths( # time for savings. This can eventually be removed once we completely disable # sync directory creation. jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async_directory_creation_secs', + '/jax/orbax/write/background_dir_init_secs', directory_creation_secs, ) logging.vlog( diff --git a/checkpoint/orbax/checkpoint/checkpoint_manager.py b/checkpoint/orbax/checkpoint/checkpoint_manager.py index 22137ab279..6fcc371962 100644 --- a/checkpoint/orbax/checkpoint/checkpoint_manager.py +++ b/checkpoint/orbax/checkpoint/checkpoint_manager.py @@ -1527,16 +1527,10 @@ def save( '[process=%s] Saving checkpoint at step %d', process_index, step ) validation_duration = time.time() - validation_start_time - if is_async_checkpointer(self._checkpointer): - jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/validation_duration_secs', - validation_duration, - ) - else: - jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/validation_duration_secs', - validation_duration, - ) + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/blocking_val_duration_secs', + validation_duration, + ) step_stats.checkpointer_blocking_start_time = time.time() self._checkpointer.save( save_directory, args=args, custom_metadata=custom_metadata, force=True @@ -2096,16 +2090,10 @@ def wait_until_finished(self): '/jax/checkpoint/write/wait_for_prev_duration_secs', duration, ) - if self._finalize_thread.get() is not None: - jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async/wait_for_prev_duration_secs', - duration, - ) - else: - jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/wait_for_prev_duration_secs', - duration, - ) + jax.monitoring.record_event_duration_secs( + '/jax/orbax/write/blocking_wait_prev_duration_secs', + duration, + ) self._wait_for_prev_save_duration += duration def is_saving_in_progress(self) -> bool: diff --git a/checkpoint/orbax/checkpoint/experimental/v1/_src/path/async_utils.py b/checkpoint/orbax/checkpoint/experimental/v1/_src/path/async_utils.py index 523cf66330..b9d995e303 100644 --- a/checkpoint/orbax/checkpoint/experimental/v1/_src/path/async_utils.py +++ b/checkpoint/orbax/checkpoint/experimental/v1/_src/path/async_utils.py @@ -64,7 +64,7 @@ async def _create_paths( directory_creation_secs, ) jax.monitoring.record_event_duration_secs( - '/jax/orbax/write/async_directory_creation_secs', + '/jax/orbax/write/background_dir_init_secs', directory_creation_secs, ) logging.vlog(