View a markdown version of this page

Point de contrôle à plusieurs niveaux - Amazon SageMaker AI

Les traductions sont fournies par des outils de traduction automatique. En cas de conflit entre le contenu d'une traduction et celui de la version originale en anglais, la version anglaise prévaudra.

Point de contrôle à plusieurs niveaux

HyperPod le point de contrôle hiérarchisé géré écrit d'abord les points de contrôle dans la mémoire du processeur de votre cluster et les reproduit sur tous les nœuds. Il les conserve périodiquement sur un stockage durable tel qu'Amazon S3. Comme le niveau le plus rapide est la mémoire plutôt que le stockage d'objets, vous pouvez effectuer des points de contrôle plus souvent et perdre moins de progression en cas de défaillance d'un nœud. Pour savoir comment cela fonctionne et comment le configurer, consultezHyperPod points de contrôle hiérarchisés gérés.

Configuration

Configurez d'abord le point de contrôle hiérarchisé géré sur le cluster. Il s'agit d'une fonctionnalité au niveau du cluster, et la configuration est la même quel que soit le framework avec lequel vous vous entraînez. Pour les étapes, consultez Mettre en place des points de contrôle hiérarchisés gérés.

Installez ensuite le package de contrôle sur les hôtes qui exécutent votre charge de travail Ray.

pip install amzn-sagemaker-checkpointing

Utilisez-le avec Ray Train

Enregistrez les points de contrôle au SageMakerTieredStorageWriter lieu de laisser Ray les télécharger. Créez une fonction de formation SageMakerCheckpointConfig interne et passez le rédacteur àasync_save. Signalez le point de contrôle de manière asynchrone afin que la formation se poursuive pendant le chargement du point de contrôle sur Amazon S3.

Cet exemple utilise le téléchargement de points de contrôle asynchrones dans la documentation Ray, ce qui permet à Ray Train de lancer un fil de discussion en arrière-plan pour attendre la fin du téléchargement pendant que la formation se poursuit à l'étape suivante. La save_checkpoint fonction place le point de contrôle dans le stockage hiérarchisé et le signale à Ray Train avec. CheckpointUploadMode.ASYNC ray.train.reportenregistre le point de contrôle avec Ray Train afin qu'il puisse suivre les meilleurs points de contrôle, les appliquer num_to_keep et les rétablir à partir du dernier point de contrôle en cas de panne. Pour plus d'informations sur le point de contrôle Ray Train, consultez la section Enregistrement et chargement des points de contrôle dans la documentation 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), ), )

Points clés :

  • Sauvegarde asynchrone avec écritures sérialisées. Une seule personne async_save peut être en vol à la fois. Les fils d'arrière-plan effectuent des opérations collectives qui nécessitent la participation de tous les grades. Un async_save nouvel appel avant la fin du précédent entraîne un blocage. Le _prev_checkpoint_future modèle garantit que chaque sauvegarde se termine avant le début de la suivante, tout en faisant chevaucher le téléchargement avec l'étape d'entraînement suivante.

  • DELETE_LOCAL_CHECKPOINT_AFTER_UPLOAD=False. Réglez cette option pour empêcher Ray de supprimer le point de contrôle signalé. ray.train.report Étant donné que le point de contrôle signalé pointe vers un chemin S3 géré par l'enregistreur de stockage hiérarchisé, sa suppression supprimerait le point de contrôle de S3.

  • Reprenez le poste de contrôle. load_checkpointlit depuis le stockage hiérarchisé (d'abord la mémoire du cluster, puis Amazon S3). Lorsqu'un nœud est remplacé et que la tâche redémarre viaFailureConfig, la restauration se lit à partir du niveau de mémoire rapide lorsqu'il est disponible, évitant ainsi un téléchargement complet d'Amazon S3.