"""Synchronize all ranks and tear down the distributed process group."""
if not dist.is_available() or not dist.is_initialized():
return
try:
empty_cache()
dist.barrier()
synchronize()
finally:
destroy_process_group()
def train(self) -> None:
"""Run the configured training loop."""
config: TrainerConfig = self.config
self.data_iterator = None
try:
self.on_train_begin()
self.data_iterator = HyperIter(
self.train_dataloader, use_background_prefetcher=config.dataloader.use_background_prefetcher
)
logger.info(
"Rank%s Start training. Global step: %s. Train iters: %s. Start epoch: %s. Train epochs: %s.",
self.local_rank, self.state.global_step, self.train_iters, self.state.epoch, self.train_epochs,
)
start_epoch = self.state.epoch
for epoch in range(start_epoch, self.train_epochs):
if epoch != start_epoch:
self.train_dataloader.set_epoch(epoch)
self.data_iterator = HyperIter(
self.train_dataloader, use_background_prefetcher=config.dataloader.use_background_prefetcher
)
self.state.epoch = epoch
self.on_epoch_begin()
start_step = self.state.global_step - epoch * self.train_steps
train_steps = min(self.train_steps, self.train_iters - epoch * self.train_steps)
for _ in range(start_step, train_steps):
try:
self.train_step(self.data_iterator)
except StopIteration:
logger.info(
"epoch:%s Dataloader finished with drop_last %s",
epoch,
config.dataloader.drop_last,
)
break
self.on_epoch_end()
self.state.epoch = epoch + 1
print_device_mem_info(f"VRAM usage after epoch {epoch + 1}")
self.data_iterator.stop()
self.data_iterator = None
finally:
# ExitStack runs every callback even when an earlier cleanup raises.
with ExitStack() as cleanup:
cleanup.callback(self.destroy_distributed)
cleanup.callback(synchronize)
cleanup.callback(self.on_train_end)
if self.data_iterator is not None:
cleanup.callback(self.data_iterator.stop)