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()