Skip to content
Open
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 @@ -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,
)

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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,
)

Expand Down Expand Up @@ -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,
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
4 changes: 2 additions & 2 deletions checkpoint/orbax/checkpoint/_src/logging/event_tracking.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
2 changes: 1 addition & 1 deletion checkpoint/orbax/checkpoint/_src/path/atomicity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
28 changes: 8 additions & 20 deletions checkpoint/orbax/checkpoint/checkpoint_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading