Le traduzioni sono generate tramite traduzione automatica. In caso di conflitto tra il contenuto di una traduzione e la versione originale in Inglese, quest'ultima prevarrà.
Checkpoint a più livelli
HyperPod il checkpoint gestito su più livelli scrive prima i checkpoint nella memoria della CPU del cluster e li replica tra i nodi. Periodicamente li mantiene su uno storage durevole come Amazon S3. Poiché il livello più rapido è costituito dalla memoria e non dallo storage a oggetti, è possibile effettuare controlli più spesso e perdere meno progressi in caso di guasto di un nodo. Per sapere come funziona e come configurarlo, consultaHyperPod checkpoint gestito su più livelli.
Configurazione
Configura prima il checkpoint gestito su più livelli sul cluster. È una funzionalità a livello di cluster e la configurazione è la stessa indipendentemente dal framework con cui ti alleni. Per la procedura, consultare Configura il checkpoint gestito su più livelli.
Quindi installa il pacchetto checkpointing sugli host che eseguono il tuo carico di lavoro Ray.
pip install amzn-sagemaker-checkpointing
Usalo con Ray Train
Salva i checkpoint SageMakerTieredStorageWriter invece di lasciare che Ray li carichi. Crea una funzione di formazione SageMakerCheckpointConfig interna e passa a chi scrive. async_save Segnala il checkpoint in modo asincrono in modo che la formazione continui mentre il checkpoint viene caricato su Amazon S3.
Questo esempio utilizza il caricamento asincrono dei checkpoint nella documentazione di Ray, che consente save_checkpoint funzione indirizza il checkpoint allo storage su più livelli e lo segnala a Ray Train con. CheckpointUploadMode.ASYNC ray.train.reportregistra il checkpoint con Ray Train in modo che possa tracciare i checkpoint migliorinum_to_keep, applicarli e ripristinarli dall'ultimo checkpoint in caso di guasto. Per ulteriori informazioni sul checkpoint di Ray Train, consulta Salvare e caricare i checkpoint nella documentazione di 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), ), )
Punti chiave:
-
Salvataggio asincrono con scritture serializzate.
async_savePuò essere in volo solo uno alla volta. I thread in background eseguono operazioni collettive che richiedono la partecipazione di tutti i ranghi. Unaasync_savenuova chiamata prima del completamento di quella precedente causa una situazione di stallo. Lo_prev_checkpoint_futureschema assicura che ogni salvataggio termini prima dell'inizio di quello successivo, pur continuando a sovrapporre il caricamento alla fase di addestramento successiva. -
delete_local_checkpoint_after_upload=false. Impostalo per impedire a Ray di eliminare il checkpoint segnalato.
ray.train.reportPoiché il checkpoint segnalato punta a un percorso S3 gestito dal tiered storage writer, la sua eliminazione rimuoverebbe il checkpoint da S3. -
Riprendi dal checkpoint.
load_checkpointlegge dallo storage su più livelli (prima la memoria del cluster, poi Amazon S3). Quando un nodo viene sostituito e il processo viene riavviatoFailureConfig, il ripristino viene letto dal livello di memoria veloce, quando disponibile, evitando il download completo di Amazon S3.