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

Method serialize

axlearn/common/array_serialization.py:1023–1055  ·  view source on GitHub ↗
(
        self,
        arrays: list[Tensor],
        tensorstore_specs: list[dict],
        *,
        on_commit_callback: Callable[[], None],
        additional_futures: Optional[list[futures.Future]] = None,
    )

Source from the content-addressed store, hash-verified

1021 def __init__(self, *args, **kwargs):
1022 super().__init__(*args, **kwargs)
1023 self._loop = asyncio.new_event_loop()
1024 self._loop_thread = threading.Thread(target=self._loop.run_forever, daemon=True)
1025 self._loop_thread.start()
1026 self._single_thread_pool = ThreadPoolExecutor(max_workers=1)
1027 # Use 80% CPU cores for parallel ckpt loading
1028 self._multi_thread_pool = ThreadPoolExecutor(max_workers=int(os.cpu_count() * 0.8))
1029
1030 def stop(self):
1031 """Cleans up any internal threads."""
1032 self._loop.call_soon_threadsafe(self._loop.stop)
1033 self._loop_thread.join()
1034 self._single_thread_pool.shutdown()
1035
1036 def __del__(self):
1037 self.stop()
1038 return super().__del__()
1039
1040 def serialize(
1041 self,
1042 arrays: list[Tensor],
1043 tensorstore_specs: list[dict],
1044 *,
1045 on_commit_callback: Callable[[], None],
1046 additional_futures: Optional[list[futures.Future]] = None,
1047 ):
1048 # Inject S3 endpoint once at the public entry point; downstream `ts.open` calls
1049 # inherit the rewritten specs. See `maybe_inject_s3_endpoint`.
1050 for s in tensorstore_specs:
1051 maybe_inject_s3_endpoint(s)
1052
1053 logging.info("Waiting for previous serialization to finish.")
1054 self.wait_until_finished()
1055
1056 commit_futures = [[] for _ in range(len(tensorstore_specs))]
1057
1058 # pylint: disable-next=redefined-outer-name

Callers

nothing calls this directly

Calls 4

_run_serializerFunction · 0.85
wait_until_finishedMethod · 0.45
resultMethod · 0.45
tree_flattenMethod · 0.45

Tested by

no test coverage detected