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 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
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_savepeut ê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. Unasync_savenouvel appel avant la fin du précédent entraîne un blocage. Le_prev_checkpoint_futuremodè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.