Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,10 @@ def _background_wait_for_commit_futures(
'/jax/checkpoint/write/async/commit_duration_sec',
commit_duration_secs,
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/tensorstore_duration_secs',
commit_duration_secs,
)

if process_count > 1:
# All processes will wait at the barrier. When all processes are at the
Expand Down Expand Up @@ -392,6 +396,7 @@ def _make_on_commit_callback(
)

def _callback() -> None:
finalize_start_time = time.time()
if utils.is_primary_host(self._primary_host):
# Update StepMetadata after the handler save is complete.
# (blocking write)
Expand Down Expand Up @@ -430,6 +435,10 @@ def _callback() -> None:
tmpdir,
checkpoint_start_time,
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/finalize_duration_secs',
time.time() - finalize_start_time,
)
operation_recorder = event_tracking.OperationRecorder(
tmpdir.get_final(),
operation_type=event_tracking.OperationType.SAVE,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -273,6 +273,7 @@ def save(

# Ensure save operation atomicity and record time saved by checkpoint.
if multihost.is_primary_host(self._primary_host):
finalize_start_time = time.time()
# finalize does a final StepMetadata update.
self._handler.finalize(tmpdir.get())
asyncio_utils.run_sync(
Expand All @@ -281,6 +282,10 @@ def save(
checkpoint_start_time=checkpoint_start_time,
)
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/finalize_duration_secs',
time.time() - finalize_start_time,
)
multihost.sync_global_processes(
multihost.unique_barrier_key(
'Checkpointer:finalize',
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -746,6 +746,14 @@ async def async_save(
),
)
]
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/tree_mapping_duration_secs',
batch_requests_ready_time - start_time,
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/d2h_duration_secs',
total_serialization_initiated_time - batch_requests_ready_time,
)
async_save_end_time = time.time()
logging.info(
'[process=%s][thread=%s] Initiated Pytree async_save. Time taken:'
Expand Down Expand Up @@ -1202,6 +1210,10 @@ async def _write_metadata_file(
'/jax/checkpoint/write/async/metadata_write_duration_secs',
time.time() - metadata_write_start_time,
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/metadata_write_duration_secs',
time.time() - metadata_write_start_time,
)

async def _write_metadata_after_commits(
self,
Expand Down Expand Up @@ -1370,6 +1382,10 @@ async def merge_ocdbt_per_process_files():
'/jax/checkpoint/write/async/ocdbt_merge_duration_secs',
time.time() - merge_start_time,
)
jax.monitoring.record_event_duration_secs(
'/jax/orbax/write/async/ocdbt_merge_duration_secs',
time.time() - merge_start_time,
)

finalize_coros.append(merge_ocdbt_per_process_files())

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1042,8 +1042,13 @@ def finalize(self, directory: epath.Path):
if tmp_dir is None or handler is None:
# Not an error, as some items may not have been saved.
continue
item_finalize_start_time = time.time()
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',
time.time() - item_finalize_start_time,
)

# Remove the temporary path once it has been finalized.
self._current_temporary_paths.pop(item_name)
Expand Down
22 changes: 22 additions & 0 deletions checkpoint/orbax/checkpoint/checkpoint_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1390,6 +1390,7 @@ def save(
step_stats.step = step
step_stats.checkpoint_manager_blocking_start_time = time.time()
step_stats.directory = str(self.directory)
validation_start_time = time.time()

if items is None and args is None:
raise ValueError('Must provide `args` for `save`.')
Expand Down Expand Up @@ -1507,6 +1508,17 @@ def save(
logging.info(
'[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,
)
step_stats.checkpointer_blocking_start_time = time.time()
self._checkpointer.save(
save_directory, args=args, custom_metadata=custom_metadata, force=True
Expand Down Expand Up @@ -2066,6 +2078,16 @@ 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,
)
self._wait_for_prev_save_duration += duration

def is_saving_in_progress(self) -> bool:
Expand Down
Loading