MAVT / train.py
Anbinh93's picture
Initial upload: code + configs + Stage 3 live progress (rgat-demo branch)
251713e verified
Raw
History Blame Contribute Delete
1.73 kB
#!/usr/bin/env python3
"""MAVT training entry point using Lightning CLI.
Usage examples:
# Stage 1: image-only with synthetic data (smoke test)
python train.py fit \
--config configs/train/stage1_image.yaml \
--model.training_stage 1 \
--trainer.max_steps 500
# Stage 1: real data, single GPU
python train.py fit --config configs/train/stage1_image.yaml \
--data.image_root /path/to/open-images
# Stage 1: DDP on 8 GPUs
python train.py fit --config configs/train/stage1_image.yaml \
--trainer.devices 8 --trainer.strategy ddp \
--data.image_root /path/to/open-images
# Stage 2: resume from stage 1 checkpoint
python train.py fit --config configs/train/stage2_video.yaml \
--ckpt_path checkpoints/stage1/mavt-stage1-best.ckpt
# WandB logging
python train.py fit --config configs/train/stage1_image.yaml \
--trainer.logger.class_path lightning.pytorch.loggers.WandbLogger \
--trainer.logger.init_args.project mavt \
--trainer.logger.init_args.name stage1-image
"""
import torch.multiprocessing as mp
from lightning.pytorch.cli import LightningCLI
from mavt.training.lightning_module import MAVTLightningModule
from mavt.data.datamodule import MAVTDataModule
# Use file_system sharing strategy: file_descriptor (default) leaks FDs across
# DataLoader workers and exhausts the system-wide ENFILE limit on busy nodes
# (saw "OSError: Too many open files in system" mid-training with 8 workers).
mp.set_sharing_strategy('file_system')
def main():
cli = LightningCLI(
model_class=MAVTLightningModule,
datamodule_class=MAVTDataModule,
save_config_callback=None, # avoid overwrite on multi-run
)
if __name__ == '__main__':
main()