File size: 1,733 Bytes
251713e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | #!/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()
|