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 @@ -585,12 +585,18 @@ class MemoryOptions(_ActiveContextGuard):
not prioritized. Note that any "prioritized" keys are assumed to be
lightweight, and `transfer_concurrent_bytes` will be ignored for
them.
serialization_status_callback: A callback object that is called at various
points during the save process per keypath, allowing for monitoring or
control over the save process.
"""

write_concurrent_bytes: int | None = None
read_concurrent_bytes: int | None = None
transfer_concurrent_bytes: int | None = None
is_prioritized_key_fn: serialization_types.IsPrioritizedKeyFn | None = None
serialization_status_callback: (
serialization_types.SerializationStatusCallback | None
) = None


@dataclasses.dataclass(kw_only=True)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
from orbax.checkpoint.experimental.v1._src.context import context as context_lib
from orbax.checkpoint.experimental.v1._src.context import options as ocp_options
from orbax.checkpoint.experimental.v1._src.saving import saving
from orbax.checkpoint.experimental.v1._src.serialization import registration
from orbax.checkpoint.experimental.v1._src.serialization import types as serialization_types



Expand Down Expand Up @@ -95,6 +97,44 @@ def is_prioritized_key_fn(path):
f'Expected call not found in {mock_handler_class.call_args_list}',
)

def test_memory_options_callback_propagation(self):
class DummyCallback(serialization_types.SerializationStatusCallback):

def key_priority(
self,
keypath: serialization_types.tree_types.PyTreeKeyPath,
) -> serialization_types.TransferPriority:
del keypath
return serialization_types.TransferPriority.ASYNCHRONOUS_DEPRIORITIZED

def on_transfer_start(
self, keypath: serialization_types.tree_types.PyTreeKeyPath
) -> None:
pass

def on_transfer_end(
self, keypath: serialization_types.tree_types.PyTreeKeyPath
) -> None:
pass

def on_write_start(
self, keypath: serialization_types.tree_types.PyTreeKeyPath
) -> None:
pass

def on_write_end(
self, keypath: serialization_types.tree_types.PyTreeKeyPath
) -> None:
pass

callback = DummyCallback()
ctx = context_lib.Context()
ctx.memory.serialization_status_callback = callback

# Assert get_array_handler propagates it.
handler = registration.get_array_handler(ctx)
self.assertEqual(handler._callback, callback)


if __name__ == '__main__':
absltest.main()
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,7 @@ def _create_v0_saving_paraminfo(
ts_context=serialization_context.ts_context,
value_typestr=None, # TODO(dnlng): Add value typestr.
enable_pinned_host_transfer=saving_options.enable_pinned_host_transfer, # pyrefly: ignore[bad-argument-type]
keypath=param.keypath,
)


Expand Down Expand Up @@ -156,6 +157,7 @@ def _create_v0_restore_paraminfo(
raise_array_data_missing_error=loading_options.raise_array_data_missing_error,
use_zarr3=deserialization_context.zarr3_checkpoint,
write_shape=write_shape,
keypath=param.keypath,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@ def _create_v0_saving_paraminfo(
ocdbt_target_data_file_size=saving_options.ocdbt_target_data_file_size,
ts_context=serialization_context.ts_context,
value_typestr='np.ndarray',
keypath=param.keypath,
)


Expand Down Expand Up @@ -129,6 +130,7 @@ def _create_v0_restore_paraminfo(
ts_context=deserialization_context.ts_context,
raise_array_data_missing_error=loading_options.raise_array_data_missing_error,
use_zarr3=deserialization_context.zarr3_checkpoint,
keypath=param.keypath,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@ def get_array_handler(
enable_replica_parallel_separate_folder=saving_options.enable_replica_parallel_separate_folder,
enable_write_sharding_file=saving_options.enable_write_sharding_file,
array_metadata_store=saving_options.array_metadata_store,
callback=context.memory_options.serialization_status_callback,
)
if loading_options.use_load_and_broadcast:
load_and_broadcast_kwargs = dict(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ def _create_v0_saving_paraminfo(
ocdbt_target_data_file_size=saving_options.ocdbt_target_data_file_size,
ts_context=serialization_context.ts_context,
value_typestr="scalar",
keypath=param.keypath,
)


Expand Down Expand Up @@ -102,6 +103,7 @@ def _create_v0_restore_paraminfo(
ts_context=deserialization_context.ts_context,
raise_array_data_missing_error=loading_options.raise_array_data_missing_error,
use_zarr3=deserialization_context.zarr3_checkpoint,
keypath=param.keypath,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,8 @@
PLACEHOLDER = ...

IsPrioritizedKeyFn = serialization_types.IsPrioritizedKeyFn
SerializationStatusCallback = serialization_types.SerializationStatusCallback
TransferPriority = serialization_types.TransferPriority

### STANDARD PYTREE LEAF TYPES

Expand Down
Loading