View a markdown version of this page

Mehrstufiges Checkpointing - Amazon SageMaker KI

Die vorliegende Übersetzung wurde maschinell erstellt. Im Falle eines Konflikts oder eines Widerspruchs zwischen dieser übersetzten Fassung und der englischen Fassung (einschließlich infolge von Verzögerungen bei der Übersetzung) ist die englische Fassung maßgeblich.

Mehrstufiges Checkpointing

HyperPod verwaltetes mehrstufiges Checkpointing schreibt Checkpoints zuerst in den CPU-Speicher Ihres Clusters und repliziert sie knotenübergreifend. Sie werden in regelmäßigen Abständen auf einem dauerhaften Speicher wie Amazon S3 gespeichert. Da es sich bei der schnellen Ebene um Arbeitsspeicher und nicht um Objektspeicher handelt, können Sie häufiger Checkpoints durchführen und verlieren weniger Fortschritt, wenn ein Knoten ausfällt. Informationen zur Funktionsweise und Konfiguration finden Sie unterHyperPod verwaltetes mehrstufiges Checkpointing.

Einrichtung

Richten Sie zuerst verwaltete mehrstufige Checkpoints auf dem Cluster ein. Es handelt sich um eine Funktion auf Cluster-Ebene, und das Setup ist unabhängig vom Framework, mit dem Sie trainieren, identisch. Die Schritte finden Sie in Richten Sie verwaltetes mehrstufiges Checkpointing ein.

Installieren Sie dann das Checkpointing-Paket auf den Hosts, auf denen Ihr Ray-Workload ausgeführt wird.

pip install amzn-sagemaker-checkpointing

Verwenden Sie es mit Ray Train

Speichere Checkpoints durch, SageMakerTieredStorageWriter anstatt Ray sie hochladen zu lassen. Erstellen Sie eine SageMakerCheckpointConfig interne Trainingsfunktion und übergeben Sie den Writer anasync_save. Melden Sie den Checkpoint asynchron, damit das Training fortgesetzt wird, während der Checkpoint auf Amazon S3 hochgeladen wird.

In diesem Beispiel wird das asynchrone Checkpoint-Upload in der Ray-Dokumentation verwendet, sodass Ray Train einen Hintergrund-Thread starten kann, um auf den Abschluss des Uploads zu warten, während das Training mit dem nächsten Schritt fortgesetzt wird. Die save_checkpoint Funktion verschiebt den Checkpoint in einen mehrstufigen Speicher und meldet ihn mit an Ray Train. CheckpointUploadMode.ASYNC ray.train.reportzeichnet den Checkpoint mit Ray Train auf, sodass das System die besten Checkpoints verfolgennum_to_keep, durchsetzen und bei einem Ausfall vom letzten Checkpoint aus wiederherstellen kann. Weitere Informationen zum Ray Train-Checkpointing finden Sie in der Ray-Dokumentation unter Speichern und Laden von Checkpoints.

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), ), )

Die wichtigsten Punkte:

  • Asynchrones Speichern mit serialisierten Schreibvorgängen. Es async_save kann jeweils nur einer im Flug sein. Die Hintergrund-Threads führen kollektive Operationen aus, an denen alle Ränge teilnehmen müssen. Ein async_save erneuter Aufruf, bevor der vorherige abgeschlossen ist, führt zu einem Deadlock. Das _prev_checkpoint_future Muster stellt sicher, dass jeder Speichervorgang abgeschlossen ist, bevor der nächste beginnt, während der Upload immer noch mit dem nächsten Trainingsschritt überlappt wird.

  • delete_local_checkpoint_after_upload=Falsch. Stellen Sie diese Option ein, um zu verhindern, dass Ray den Checkpoint löscht, über ray.train.report den berichtet wurde. Da der gemeldete Checkpoint auf einen S3-Pfad verweist, der vom Tiered Storage Writer verwaltet wird, würde das Löschen des Checkpoints den Checkpoint aus S3 entfernen.

  • Fahren Sie am Checkpoint fort. load_checkpointliest aus mehrstufigem Speicher (zuerst Cluster-Speicher, dann Amazon S3). Wenn ein Knoten ausgetauscht wird und der Job über neu gestartet wirdFailureConfig, liest die Wiederherstellung Daten aus der schnellen Speicherebene, sofern verfügbar, wodurch ein vollständiger Amazon S3-Download vermieden wird.