(
self,
arrays: list[Tensor],
tensorstore_specs: list[dict],
*,
on_commit_callback: Callable[[], None],
additional_futures: Optional[list[futures.Future]] = None,
)
| 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 |
nothing calls this directly
no test coverage detected