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
2 changes: 2 additions & 0 deletions checkpoint/orbax/checkpoint/_src/metadata/array_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,8 @@ class ArrayMetadata:
ext_metadata: ExtMetadata | None = (
None # to contain any extension metadata for an array.
)
compression_algorithm: str = 'none'
compression_level: int | None = None


@dataclasses.dataclass(frozen=True, kw_only=True)
Expand Down
132 changes: 75 additions & 57 deletions checkpoint/orbax/checkpoint/_src/serialization/jax_array_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,34 +280,71 @@ def _record_logical_metrics(
)


def _record_compression_metrics(
direction: types.IoDirection,
logical_bytes: int,
raw_bytes: int,
storage_type: str,
custom_prefix: str = '',
metadatas: Sequence[ts_utils.ArrayMetadata] | None = None,
) -> None:
"""Logs and records compression ratio and metrics."""
if logical_bytes <= 0:
return
ratio = float(raw_bytes) / logical_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',
multihost.process_index(),
direction.value.capitalize(),
ratio,
humanize.naturalsize(raw_bytes, binary=True),
humanize.naturalsize(logical_bytes, binary=True),
algo_str,
level_str,
)
jax.monitoring.record_scalar(
f'/jax/orbax/{direction.value}/worker/io/compression_ratio',
ratio,
storage_type=storage_type,
custom_prefix=custom_prefix,
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(
direction: types.IoDirection,
logical_bytes: int,
duration: float,
storage_type: str,
initial_ts_metrics: Sequence[dict[str, Any]] | None = None,
custom_prefix: str = '',
metadatas: Sequence[ts_utils.ArrayMetadata] | None = None,
):
"""Records raw metrics collected from TensorStore."""
if initial_ts_metrics is None:
return

try:
final_ts_metrics = ts.experimental_collect_matching_metrics('/tensorstore/')
except Exception: # pylint: disable=broad-except
final_ts_metrics = None

final_ts_metrics = ts_utils.collect_tensorstore_metrics()
if final_ts_metrics is None:
return

initial_bytes = ts_utils.get_total_bytes_from_tensorstore(
initial_ts_metrics, direction
raw_bytes = ts_utils.get_tensorstore_raw_bytes_delta(
initial_ts_metrics, final_ts_metrics, direction
)
final_bytes = ts_utils.get_total_bytes_from_tensorstore(
final_ts_metrics, direction
)
raw_bytes = final_bytes - initial_bytes

if raw_bytes <= 0:
return

Expand Down Expand Up @@ -335,29 +372,14 @@ def _record_raw_metrics(
custom_prefix=custom_prefix,
)

if logical_bytes > 0:
ratio = float(raw_bytes) / logical_bytes
logging.info(
'[process=%d] %s ratio (raw/logical): %.3f (%s / %s)',
multihost.process_index(),
direction.value.capitalize(),
ratio,
humanize.naturalsize(raw_bytes, binary=True),
humanize.naturalsize(logical_bytes, binary=True),
)
jax.monitoring.record_scalar(
f'/jax/orbax/{direction.value}/worker/io/compression_ratio',
ratio,
storage_type=storage_type,
custom_prefix=custom_prefix,
)
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,
)
_record_compression_metrics(
direction,
logical_bytes,
raw_bytes,
storage_type,
custom_prefix=custom_prefix,
metadatas=metadatas,
)


def _log_io_metrics(
Expand All @@ -367,6 +389,7 @@ def _log_io_metrics(
parent_dir: epath.Path,
initial_ts_metrics: Sequence[dict[str, Any]] | None = None,
custom_prefix: str = '',
metadatas: Sequence[ts_utils.ArrayMetadata] | None = None,
):
"""Logs and records IO telemetry metrics for array serialization/deserialization."""
duration = time.time() - start_time
Expand All @@ -386,6 +409,7 @@ def _log_io_metrics(
storage_type,
initial_ts_metrics=initial_ts_metrics,
custom_prefix=custom_prefix,
metadatas=metadatas,
)


Expand All @@ -404,12 +428,7 @@ def _worker_serialize_arrays(
ext_metadata: Dict[str, Any],
):
"""Worker function to serialize arrays."""
try:
initial_ts_metrics = ts.experimental_collect_matching_metrics(
'/tensorstore/'
)
except Exception: # pylint: disable=broad-except
initial_ts_metrics = None
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
total_start_time = time.time()
rslices_per_array = _get_replica_slices(
arrays,
Expand All @@ -419,7 +438,7 @@ def _worker_serialize_arrays(
max_replicas_for_replica_parallel,
)

asyncio_utils.run_sync(
array_metadatas = asyncio_utils.run_sync(
_async_serialize_replica_slices(
rslices_per_array,
infos,
Expand All @@ -440,6 +459,7 @@ def _worker_serialize_arrays(
start_time=total_start_time,
parent_dir=infos[0].parent_dir,
initial_ts_metrics=initial_ts_metrics,
metadatas=array_metadatas,
)


Expand Down Expand Up @@ -551,22 +571,20 @@ def _serialize_arrays_batches_without_dispatcher(
async def _serialize_without_dispatcher():
if not prioritized and not deprioritized:
return
try:
initial_ts_metrics = ts.experimental_collect_matching_metrics(
'/tensorstore/'
)
except Exception: # pylint: disable=broad-except
initial_ts_metrics = None
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
total_start_time = time.time()
logical_bytes = 0
all_array_metadatas: list[ts_utils.ArrayMetadata] = []

if prioritized_values_on_host:
logical_bytes += sum(v.nbytes for v in prioritized_values_on_host)
await async_serialize_replica_slices_batch(
b_metadatas = await async_serialize_replica_slices_batch(
prioritized_values_on_host,
prioritized_infos,
prioritized_args,
)
if b_metadatas:
all_array_metadatas.extend(b_metadatas)
_on_batch_callback(prioritized_infos, callback.on_write_end)
if deprioritized:
assert device_host_max_bytes is not None
Expand All @@ -586,11 +604,13 @@ async def _serialize_without_dispatcher():
b_arrays_on_host = replica_slices_transfer_arrays_to_host(b_arrays)
_on_batch_callback(b_infos, callback.on_transfer_end)
logical_bytes += sum(v.nbytes for v in b_arrays_on_host)
await async_serialize_replica_slices_batch(
b_metadatas = await async_serialize_replica_slices_batch(
b_arrays_on_host,
b_infos,
b_args,
)
if b_metadatas:
all_array_metadatas.extend(b_metadatas)
_on_batch_callback(b_infos, callback.on_write_end)

info_sample = prioritized[0][1] if prioritized else deprioritized[0][1]
Expand All @@ -600,6 +620,7 @@ async def _serialize_without_dispatcher():
start_time=total_start_time,
parent_dir=info_sample.parent_dir,
initial_ts_metrics=initial_ts_metrics,
metadatas=all_array_metadatas,
)

return future.CommitFutureAwaitingContractedSignals(
Expand Down Expand Up @@ -783,7 +804,7 @@ async def _async_serialize_replica_slices(
enable_replica_parallel_separate_folder: bool,
use_replica_parallel: bool,
ext_metadata: Dict[str, Any],
) -> None:
) -> Sequence[ts_utils.ArrayMetadata]:
"""This function contains the logic from ArrayHandler._background_serialize."""
write_coros = []
array_metadatas = []
Expand Down Expand Up @@ -884,6 +905,8 @@ async def _async_serialize_replica_slices(
ocdbt_transaction.abort()
raise

return array_metadatas


def _wrap_random_key_data(
array_metadatas: Any,
Expand Down Expand Up @@ -1014,12 +1037,7 @@ async def _deserialize_arrays(
array_metadata_store: array_metadata_store_lib.Store | None,
) -> Sequence[jax.Array]:
"""Deserializes arrays and applies array_metadata if available."""
try:
initial_ts_metrics = ts.experimental_collect_matching_metrics(
'/tensorstore/'
)
except Exception: # pylint: disable=broad-except
initial_ts_metrics = None
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
total_start_time = time.time()

async def _async_deserialize(
Expand Down
Loading
Loading