View a markdown version of this page

Pos pemeriksaan berjenjang - Amazon SageMaker AI

Terjemahan disediakan oleh mesin penerjemah. Jika konten terjemahan yang diberikan bertentangan dengan versi bahasa Inggris aslinya, utamakan versi bahasa Inggris.

Pos pemeriksaan berjenjang

HyperPod pos pemeriksaan berjenjang terkelola menulis pos pemeriksaan ke memori CPU cluster Anda terlebih dahulu dan mereplikasi mereka di seluruh node. Ini secara berkala mempertahankan mereka ke penyimpanan yang tahan lama seperti Amazon S3. Karena tingkat cepat adalah memori daripada penyimpanan objek, Anda dapat memeriksa lebih sering dan kehilangan lebih sedikit kemajuan ketika node gagal. Untuk cara kerjanya dan cara mengkonfigurasinya, lihatHyperPod pos pemeriksaan berjenjang terkelola.

Pengaturan

Siapkan pos pemeriksaan berjenjang terkelola di cluster terlebih dahulu. Ini adalah kemampuan tingkat cluster, dan pengaturannya sama dengan kerangka kerja apa pun yang Anda latih. Untuk langkah-langkahnya, lihatSiapkan pos pemeriksaan berjenjang terkelola.

Kemudian instal paket checkpointing pada host yang menjalankan beban kerja Ray Anda.

pip install amzn-sagemaker-checkpointing

Gunakan dengan Ray Train

Simpan pos pemeriksaan SageMakerTieredStorageWriter alih-alih membiarkan Ray mengunggahnya. Bangun SageMakerCheckpointConfig bagian dalam fungsi pelatihan Anda dan berikan penulis keasync_save. Laporkan pos pemeriksaan secara asinkron sehingga pelatihan berlanjut saat pos pemeriksaan diunggah ke Amazon S3.

Contoh ini menggunakan pengunggahan pos pemeriksaan asinkron dalam dokumentasi Ray, yang memungkinkan Ray Train memulai utas latar belakang untuk menunggu unggahan selesai sementara pelatihan berlanjut pada langkah berikutnya. save_checkpointFungsi ini membuat pos pemeriksaan ke penyimpanan berjenjang dan melaporkannya ke Ray Train dengan. CheckpointUploadMode.ASYNC ray.train.reportmerekam pos pemeriksaan dengan Ray Train sehingga dapat melacak pos pemeriksaan terbaik, menegakkannum_to_keep, dan memulihkan dari pos pemeriksaan terbaru jika gagal. Untuk informasi selengkapnya tentang pos pemeriksaan Ray Train, lihat Menyimpan dan memuat pos pemeriksaan di dokumentasi 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), ), )

Poin kunci:

  • Simpan asinkron dengan penulisan serial. Hanya satu yang async_save bisa terbang dalam satu waktu. Utas latar belakang melakukan operasi kolektif yang mengharuskan semua peringkat untuk berpartisipasi. Memanggil async_save lagi sebelum yang sebelumnya selesai menyebabkan kebuntuan. _prev_checkpoint_futurePola memastikan setiap penyimpanan selesai sebelum yang berikutnya dimulai, sementara masih tumpang tindih unggahan dengan langkah pelatihan berikutnya.

  • DELETE_LOCAL_CHECKPOINT_AFTER_UPLOAD=Salah. Setel ini untuk mencegah Ray menghapus pos pemeriksaan yang dilaporkan. ray.train.report Karena pos pemeriksaan yang dilaporkan menunjuk ke jalur S3 yang dikelola oleh penulis penyimpanan berjenjang, menghapusnya akan menghapus pos pemeriksaan dari S3.

  • Lanjutkan dari pos pemeriksaan. load_checkpointmembaca dari penyimpanan berjenjang (memori cluster terlebih dahulu, lalu Amazon S3). Saat node diganti dan pekerjaan dimulai ulangFailureConfig, pemulihan akan dibaca dari tingkat memori cepat bila tersedia, menghindari unduhan Amazon S3 penuh.