View a markdown version of this page

階層型チェックポイント - Amazon SageMaker AI

翻訳は機械翻訳により提供されています。提供された翻訳内容と英語版の間で齟齬、不一致または矛盾がある場合、英語版が優先します。

階層型チェックポイント

HyperPod マネージド階層型チェックポイントは、最初にチェックポイントをクラスターの CPU メモリに書き込み、ノード間でレプリケートします。定期的に Amazon S3 などの耐久性の高いストレージに保持されます。高速階層はオブジェクトストレージではなくメモリであるため、ノードに障害が発生したときにチェックポイントをより頻繁に実行し、進行状況を減らすことができます。仕組みと設定方法については、「」を参照してくださいHyperPod マネージド階層型チェックポイント。

セットアップ

最初にクラスターでマネージド階層型チェックポイントを設定します。これはクラスターレベルの機能であり、セットアップはトレーニングするフレームワークと同じです。手順についてはマネージド階層型チェックポイントの設定を参照してください。

次に、Ray ワークロードを実行するホストにチェックポイントパッケージをインストールします。

pip install amzn-sagemaker-checkpointing

Ray Train で使用する

Ray がチェックポイントをアップロードするSageMakerTieredStorageWriterのではなく、 を介してチェックポイントを保存します。トレーニング関数SageMakerCheckpointConfig内に を構築し、ライターを に渡しますasync_save。チェックポイントを非同期に報告して、チェックポイントが Amazon S3 にアップロードされている間にトレーニングが続行されるようにします。

この例では、Ray ドキュメントで非同期チェックポイントアップロードを使用します。これにより、Ray Train はバックグラウンドスレッドを開始してアップロードが完了するまで待機し、トレーニングは次のステップに進みます。このsave_checkpoint関数はチェックポイントを階層型ストレージにステージングし、 を使用して Ray Train に報告しますCheckpointUploadMode.ASYNC。 は Ray Train を使用してチェックポイントray.train.reportを記録します。これにより、最適なチェックポイントを追跡し、 を適用しnum_to_keep、障害時に最新のチェックポイントから復元できます。Ray Train チェックポイントの詳細については、Ray ドキュメントの「チェックポイントの保存とロード」を参照してください。

import os import torch.distributed as dist from torch.distributed.checkpoint import async_save, load import ray import ray.train from ray.train import ( Checkpoint, CheckpointConfig, CheckpointUploadMode, RunConfig, ScalingConfig, FailureConfig, ) from ray.train.torch import TorchTrainer from amzn_sagemaker_checkpointing.config.sagemaker_checkpoint_config import SageMakerCheckpointConfig from amzn_sagemaker_checkpointing.checkpointing.filesystem.filesystem import ( SageMakerTieredStorageWriter, SageMakerTieredStorageReader, ) S3_PATH = "s3://my-bucket/checkpoints" EXPERIMENT_NAME = "my-experiment" def create_checkpoint_config(): """Create a checkpoint config scoped to this training job.""" return SageMakerCheckpointConfig( namespace=EXPERIMENT_NAME, world_size=dist.get_world_size(), s3_tier_base_path=S3_PATH, ) # Track the previous checkpoint future to avoid async_save deadlock. # Only one async_save can be in flight at a time because background # threads perform collectives that require all ranks to participate. _prev_checkpoint_future = None def save_checkpoint(model, optimizer, config, step, metrics): """Save a checkpoint asynchronously and report it to Ray Train.""" global _prev_checkpoint_future # Wait for the previous checkpoint to finish before starting a new one. if _prev_checkpoint_future is not None: _prev_checkpoint_future.result() config.save_to_s3 = True writer = SageMakerTieredStorageWriter(checkpoint_config=config, step=step) state_dict = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "step": step, } future = async_save(state_dict=state_dict, storage_writer=writer) _prev_checkpoint_future = future def wait_for_upload(checkpoint, name): future.result() return checkpoint ray.train.report( metrics=metrics, checkpoint=Checkpoint(writer.s3_checkpoint_dir), checkpoint_upload_mode=CheckpointUploadMode.ASYNC, checkpoint_upload_fn=wait_for_upload, delete_local_checkpoint_after_upload=False, ) def load_checkpoint(model, optimizer, config): """Load the latest checkpoint if one exists. Returns the next step number.""" reader = SageMakerTieredStorageReader(checkpoint_config=config) state_dict = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "step": 0, } try: load(state_dict, storage_reader=reader) except FileNotFoundError: # No checkpoint found, start from scratch. return 0 model.load_state_dict(state_dict["model"]) optimizer.load_state_dict(state_dict["optimizer"]) return state_dict["step"] + 1 def train_func(config): device = ray.train.torch.get_device() model = build_model().to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3) ckpt_config = create_checkpoint_config() start_step = load_checkpoint(model, optimizer, ckpt_config) for step in range(start_step, config["max_steps"]): loss = train_step(model, optimizer) if step % config["checkpoint_freq"] == 0: save_checkpoint( model, optimizer, ckpt_config, step, metrics={"loss": loss, "step": step}, ) trainer = TorchTrainer( train_func, train_loop_config={"max_steps": 1000, "checkpoint_freq": 10}, scaling_config=ScalingConfig(num_workers=4, use_gpu=True), run_config=RunConfig( name=EXPERIMENT_NAME, storage_path=S3_PATH, failure_config=FailureConfig(max_failures=3), checkpoint_config=CheckpointConfig(num_to_keep=3), ), )

キーポイント:

  • シリアル化された書き込みによる非同期保存。一度に 1 つのみ飛行async_saveできます。バックグラウンドスレッドは、すべてのランクが参加する必要がある集合演算を実行します。前の手順が完了する前に async_save を再度呼び出すと、デッドロックが発生します。この_prev_checkpoint_futureパターンにより、各保存が次のトレーニングステップと重複したまま、次の保存が開始される前に完了します。

  • delete_local_checkpoint_after_upload=False。これを設定して、Ray が を通じて報告されたチェックポイントを削除しないようにしますray.train.report。報告されたチェックポイントは階層型ストレージライターによって管理される S3 パスを指すため、削除するとチェックポイントが S3 から削除されます。

  • チェックポイントから再開します。 は階層型ストレージ (最初にクラスターメモリ、次に Amazon S3) からload_checkpoint読み取ります。ノードが置き換えられ、 を介してジョブが再起動するとFailureConfig、利用可能な場合、リカバリは高速メモリ層から読み取り、Amazon S3 の完全なダウンロードを回避します。