import os
from typing import Optional
import torch
import torchtitan_npu
from torchtitan.config.manager import ConfigManager
from torchtitan.tools.logging import init_logger, logger
from torchtitan.train import Trainer
if __name__ == "__main__":
init_logger()
config_manager = ConfigManager()
config = config_manager.parse_args()
trainer: Optional[Trainer] = None
try:
trainer = Trainer(config)
if config.checkpoint.create_seed_checkpoint:
assert (
int(os.environ["WORLD_SIZE"]) == 1
), "Must create seed checkpoint using a single device, to disable sharding."
assert (
config.checkpoint.enable
), "Must enable checkpointing when creating a seed checkpoint."
trainer.checkpointer.save(curr_step=0, last_step=True)
logger.info("Created seed checkpoint")
else:
trainer.train()
except Exception:
if trainer:
trainer.close()
raise
else:
trainer.close()
torch.distributed.destroy_process_group()
logger.info("Process group destroyed")