MCPcopy Create free account
hub / github.com/apple/axlearn / _run_serializer

Function _run_serializer

axlearn/common/array_serialization.py:450–495  ·  view source on GitHub ↗

Asynchronously serializes a list of tensors with _async_serialize.

(
    arrays: list[Tensor],
    tensorstore_specs: list[dict[str, Any]],
    d2h_futures: list[futures.Future],
    *,
    max_concurrent_bytes: Optional[int] = None,
    tensorstore_spec_modifier: Optional[TensorstoreSpecModifier] = None,
    max_data_shard_degree: int,
    shard_threshold_bytes: int,
)

Source from the content-addressed store, hash-verified

448 open=True,
449 assume_metadata=True,
450 context=serialization.TS_CONTEXT,
451 )
452
453 # Avoid additional copy of input array into the TensorStore chunk cache. If `arr_inp` is a
454 # jax.Array, the result of converting it to a NumPy array, as is done internally by TensorStore,
455 # is guaranteed to be immutable and therefore it is safe to retain reference indefinitely.
456 is_jax_array = isinstance(arr_inp, jax.Array)
457 await asyncio.gather(
458 *(
459 t[info.index].write(info.data, can_reference_source_data_indefinitely=is_jax_array)
460 for info in shard_infos
461 )
462 )
463 if limiter is not None:
464 await limiter.release_bytes(nbytes)
465
466
467async def _run_serializer(
468 arrays: list[Tensor],
469 tensorstore_specs: list[dict[str, Any]],
470 d2h_futures: list[futures.Future],
471 *,
472 max_concurrent_bytes: Optional[int] = None,
473 tensorstore_spec_modifier: Optional[TensorstoreSpecModifier] = None,
474 max_data_shard_degree: int,
475 shard_threshold_bytes: int,
476):
477 """Asynchronously serializes a list of tensors with _async_serialize."""
478 # We add 1 because LimitInFlightBytes expects a limit strictly greater than any request.
479 # pylint: disable=protected-access
480 limiter = (
481 serialization._LimitInFlightBytes(max_concurrent_bytes + 1)
482 if max_concurrent_bytes
483 else None
484 )
485 # pylint: enable=protected-access
486 future_writer = jax.tree.map(
487 functools.partial(
488 _async_serialize,
489 limiter=limiter,
490 max_data_shard_degree=max_data_shard_degree,
491 shard_threshold_bytes=shard_threshold_bytes,
492 tensorstore_spec_modifier=tensorstore_spec_modifier,
493 ),
494 arrays,
495 tensorstore_specs,
496 d2h_futures,
497 )
498 try:

Callers 4

serializeMethod · 0.85
serializeMethod · 0.85

Calls 1

mapMethod · 0.80