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,
)
| 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 | |
| 467 | async 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: |