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 @@ -330,21 +330,16 @@ def _record_raw_metrics(
logical_bytes: int,
duration: float,
storage_type: str,
initial_ts_metrics: Sequence[dict[str, Any]] | None = None,
initial_raw_bytes: int | None = None,
custom_prefix: str = '',
metadatas: Sequence[ts_utils.ArrayMetadata] | None = None,
):
"""Records raw metrics collected from TensorStore."""
if initial_ts_metrics is None:
if initial_raw_bytes is None:
return

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

raw_bytes = ts_utils.get_tensorstore_raw_bytes_delta(
initial_ts_metrics, final_ts_metrics, direction
)
final_raw_bytes = ts_utils.get_tensorstore_raw_bytes(direction)
raw_bytes = max(0, final_raw_bytes - initial_raw_bytes)
if raw_bytes <= 0:
return

Expand Down Expand Up @@ -387,7 +382,7 @@ def _log_io_metrics(
logical_bytes: int,
start_time: float,
parent_dir: epath.Path,
initial_ts_metrics: Sequence[dict[str, Any]] | None = None,
initial_raw_bytes: int | None = None,
custom_prefix: str = '',
metadatas: Sequence[ts_utils.ArrayMetadata] | None = None,
):
Expand All @@ -407,7 +402,7 @@ def _log_io_metrics(
logical_bytes,
duration,
storage_type,
initial_ts_metrics=initial_ts_metrics,
initial_raw_bytes=initial_raw_bytes,
custom_prefix=custom_prefix,
metadatas=metadatas,
)
Expand All @@ -428,7 +423,9 @@ def _worker_serialize_arrays(
ext_metadata: Dict[str, Any],
):
"""Worker function to serialize arrays."""
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes(
types.IoDirection.WRITE
)
total_start_time = time.time()
rslices_per_array = _get_replica_slices(
arrays,
Expand Down Expand Up @@ -458,7 +455,7 @@ def _worker_serialize_arrays(
logical_bytes=total_io_bytes,
start_time=total_start_time,
parent_dir=infos[0].parent_dir,
initial_ts_metrics=initial_ts_metrics,
initial_raw_bytes=initial_raw_bytes,
metadatas=array_metadatas,
)

Expand Down Expand Up @@ -571,7 +568,9 @@ def _serialize_arrays_batches_without_dispatcher(
async def _serialize_without_dispatcher():
if not prioritized and not deprioritized:
return
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes(
types.IoDirection.WRITE
)
total_start_time = time.time()
logical_bytes = 0
all_array_metadatas: list[ts_utils.ArrayMetadata] = []
Expand Down Expand Up @@ -619,7 +618,7 @@ async def _serialize_without_dispatcher():
logical_bytes=logical_bytes,
start_time=total_start_time,
parent_dir=info_sample.parent_dir,
initial_ts_metrics=initial_ts_metrics,
initial_raw_bytes=initial_raw_bytes,
metadatas=all_array_metadatas,
)

Expand Down Expand Up @@ -1037,7 +1036,9 @@ async def _deserialize_arrays(
array_metadata_store: array_metadata_store_lib.Store | None,
) -> Sequence[jax.Array]:
"""Deserializes arrays and applies array_metadata if available."""
initial_ts_metrics = ts_utils.collect_tensorstore_metrics()
initial_raw_bytes = ts_utils.get_tensorstore_raw_bytes(
types.IoDirection.READ
)
total_start_time = time.time()

async def _async_deserialize(
Expand Down Expand Up @@ -1141,7 +1142,7 @@ async def _async_deserialize(
logical_bytes=logical_bytes,
start_time=total_start_time,
parent_dir=infos[0].parent_dir,
initial_ts_metrics=initial_ts_metrics,
initial_raw_bytes=initial_raw_bytes,
)
return ret

Expand Down
53 changes: 17 additions & 36 deletions checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1155,54 +1155,35 @@ def array_metadata_from_tensorstore(
)


def get_total_bytes_from_tensorstore(
metrics: Sequence[dict[str, Any]], direction: types.IoDirection
def get_tensorstore_raw_bytes(
direction: types.IoDirection = types.IoDirection.WRITE,
) -> int:
"""Sums bytes_read or bytes_written from all kvstore drivers in metrics."""
total = 0
if direction == types.IoDirection.WRITE:
suffix = '/bytes_written'
elif direction == types.IoDirection.READ:
suffix = '/bytes_read'
else:
raise ValueError(f'Invalid direction: {direction}')
"""Collects and returns total raw bytes read or written by TensorStore."""
suffix = (
'/bytes_written'
if direction == types.IoDirection.WRITE
else '/bytes_read'
)
# Querying `/tensorstore/kvstore/` returns only a few kvstore driver
# metrics (e.g. file, gcs, s3), making metric collection and
# extraction fast without scanning all TensorStore metrics.
metrics: Sequence[dict[str, Any]] = ts.experimental_collect_matching_metrics(
'/tensorstore/kvstore/'
)

# Sum the metric values for the given suffix.
total = 0
for m in metrics:
if not isinstance(m, dict):
continue
name = m.get('name', '')
if name.startswith('/tensorstore/kvstore/') and name.endswith(suffix):
if name.endswith(suffix):
for val in m.get('values', []):
if isinstance(val, dict):
total += val.get('value', 0)
return total


def get_tensorstore_raw_bytes_delta(
initial_metrics: Sequence[dict[str, Any]] | None,
final_metrics: Sequence[dict[str, Any]] | None,
direction: types.IoDirection = types.IoDirection.WRITE,
) -> int:
"""Computes transferred raw bytes delta between two metric snapshots."""
if initial_metrics is None or final_metrics is None:
return 0
try:
initial_bytes = get_total_bytes_from_tensorstore(initial_metrics, direction)
final_bytes = get_total_bytes_from_tensorstore(final_metrics, direction)
return max(0, final_bytes - initial_bytes)
except Exception: # pylint: disable=broad-except
logging.exception('Failed to compute TensorStore raw bytes delta.')
return 0


def collect_tensorstore_metrics() -> Sequence[dict[str, Any]] | None:
"""Safely collects TensorStore driver metrics."""
try:
return ts.experimental_collect_matching_metrics('/tensorstore')
except Exception: # pylint: disable=broad-except
return None


def resolve_compression_settings(
metadatas: Sequence[ArrayMetadata],
) -> tuple[str, str]:
Expand Down
Loading
Loading