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

Method __init__

axlearn/common/checkpointer.py:455–480  ·  view source on GitHub ↗
(self, cfg: Config)

Source from the content-addressed store, hash-verified

453 shard_threshold_bytes: Optional[int] = None
454
455 def __init__(self, cfg: Config):
456 super().__init__(cfg)
457 cfg = self.config
458 # TODO(markblee): Consider making BoundedDataShardedAsyncCheckpointManager
459 # the default once stable.
460 if cfg.max_concurrent_gb is not None or cfg.max_data_shard_degree:
461 self._manager = BoundedDataShardedAsyncCheckpointManager(
462 max_concurrent_gb=cfg.max_concurrent_gb,
463 timeout_secs=cfg.timeout_secs,
464 max_data_shard_degree=cfg.max_data_shard_degree,
465 shard_threshold_bytes=cfg.shard_threshold_bytes,
466 )
467 else:
468 if cfg.shard_threshold_bytes is not None:
469 raise ValueError(
470 f"shard_threshold_bytes is set to {cfg.shard_threshold_bytes}, but "
471 "max_data_shard_degree is not set. It will not take any effect."
472 )
473 self._manager = GlobalAsyncCheckpointManager(timeout_secs=cfg.timeout_secs)
474 if cfg.max_concurrent_restore_gb is not None and cfg.max_concurrent_restore_gb <= 0:
475 raise ValueError(
476 f"max_concurrent_restore_gb must be strictly positive. "
477 f"Got {cfg.max_concurrent_restore_gb}"
478 )
479 self._max_concurrent_restore_gb = cfg.max_concurrent_restore_gb or 32
480 self._executor = futures.ThreadPoolExecutor()
481
482 @dataclasses.dataclass
483 class CheckpointSpec: # pylint: disable=too-many-instance-attributes

Callers

nothing calls this directly

Tested by

no test coverage detected