As traduções são geradas por tradução automática. Em caso de conflito entre o conteúdo da tradução e da versão original em inglês, a versão em inglês prevalecerá.
Ponto de verificação hierárquico
HyperPod o checkpoint hierárquico gerenciado grava primeiro os pontos de verificação na memória da CPU do cluster e os replica entre os nós. Ele os mantém periodicamente em um armazenamento durável, como o Amazon S3. Como o nível mais rápido é a memória e não o armazenamento de objetos, você pode fazer o checkpoint com mais frequência e perder menos progresso quando um nó falha. Para saber como ele funciona e como configurá-lo, consulteHyperPod ponto de verificação hierárquico gerenciado.
Configuração
Configure primeiro o ponto de verificação hierárquico gerenciado no cluster. É um recurso em nível de cluster e a configuração é a mesma, seja qual for a estrutura com a qual você treina. Para obter as etapas, consulte Configurar pontos de verificação hierárquicos gerenciados.
Em seguida, instale o pacote de checkpoint nos hosts que executam sua carga de trabalho do Ray.
pip install amzn-sagemaker-checkpointing
Use-o com Ray Train
Salve os pontos de verificação SageMakerTieredStorageWriter em vez de deixar que Ray os envie. Crie uma função de treinamento SageMakerCheckpointConfig interna e passe o redator paraasync_save. Relate o ponto de verificação de forma assíncrona para que o treinamento continue enquanto o ponto de verificação é carregado para o Amazon S3.
Este exemplo usa o upload assíncrono de pontos de verificação save_checkpoint função transforma o ponto de verificação em armazenamento hierárquico e o reporta ao Ray Train com. CheckpointUploadMode.ASYNC ray.train.reportregistra o ponto de verificação com o Ray Train para que ele possa rastrear os melhores pontos de verificaçãonum_to_keep, aplicar e restaurar a partir do ponto de verificação mais recente em caso de falha. Para obter mais informações sobre os pontos de verificação do Ray Train, consulte Salvando e carregando pontos de verificação
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), ), )
Principais pontos:
-
Salvamento assíncrono com gravações serializadas. Somente um
async_savepode voar por vez. Os segmentos em segundo plano realizam operações coletivas que exigem a participação de todas as classificações. Ligarasync_savenovamente antes que a anterior seja concluída causa um impasse. O_prev_checkpoint_futurepadrão garante que cada salvamento termine antes do início do próximo, ao mesmo tempo em que sobrepõe o upload à próxima etapa de treinamento. -
delete_local_checkpoint_after_upload=Falso. Defina isso para evitar que Ray exclua o ponto de verificação informado.
ray.train.reportComo o ponto de verificação relatado aponta para um caminho do S3 gerenciado pelo gravador de armazenamento hierárquico, excluí-lo removeria o ponto de verificação do S3. -
Retomar a partir do ponto de verificação.
load_checkpointlê do armazenamento hierárquico (primeiro a memória do cluster e depois o Amazon S3). Quando um nó é substituído e a tarefa é reiniciadaFailureConfig, a recuperação é lida a partir do nível de memória rápida, quando disponível, evitando um download completo do Amazon S3.