Spaces:
Running on Zero
Running on Zero
File size: 4,803 Bytes
0122a25 | 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 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """CLI interface using PyTorch Lightning."""
from __future__ import annotations
import logging
import os.path as osp
import torch
from absl import app # pylint: disable=no-name-in-module
from torch.utils.collect_env import get_pretty_env_info
from mapdet3d.common.logging import dump_config, rank_zero_info, setup_logger
from mapdet3d.common.typing import ArgsType
from mapdet3d.common.util import set_tf32
from mapdet3d.config import instantiate_classes
from mapdet3d.config.typing import ExperimentConfig
from mapdet3d.engine.callbacks import (
Callback,
LRSchedulerCallback,
VisualizerCallback,
)
from mapdet3d.engine.data_module import DataModule
from mapdet3d.engine.flag import (
_CKPT,
_CONFIG,
_GPUS,
_NODES,
_RESUME,
_SHOW_CONFIG,
_VIS,
_WANDB,
)
from mapdet3d.engine.parser import pprints_config
from mapdet3d.engine.trainer import PLTrainer
from mapdet3d.engine.training_module import TrainingModule
def main(argv: ArgsType) -> None:
"""Main entry point for the CLI.
Example to run this script:
>>> python -m mapdet3d.pl.run fit --config configs/faster_rcnn/faster_rcnn_coco.py
"""
# Get config
mode = argv[1]
assert mode in {"fit", "test"}, f"Invalid mode: {mode}"
config: ExperimentConfig = _CONFIG.value
num_gpus = _GPUS.value
num_nodes = _NODES.value
# Setup logging
logger = logging.getLogger("mapdet3d")
logger_pl = logging.getLogger("pytorch_lightning")
log_file = osp.join(config.output_dir, f"log_{config.timestamp}.txt")
setup_logger(logger, log_file)
setup_logger(logger_pl, log_file)
# Dump config
config_file = osp.join(
config.output_dir, f"config_{config.timestamp}.yaml"
)
dump_config(config, config_file)
rank_zero_info("Environment info: %s", get_pretty_env_info())
# PyTorch Setting
set_tf32(config.use_tf32, config.tf32_matmul_precision)
torch.hub.set_dir(f"{config.work_dir}/.cache/torch/hub")
# Setup device
if num_gpus > 0:
config.pl_trainer.accelerator = "gpu"
config.pl_trainer.devices = num_gpus
else:
config.pl_trainer.accelerator = "cpu"
config.pl_trainer.devices = 1
if num_nodes > 1:
config.pl_trainer.num_nodes = num_nodes
# Wandb
config.pl_trainer.wandb = _WANDB.value
trainer_args = instantiate_classes(config.pl_trainer).to_dict()
if _SHOW_CONFIG.value:
rank_zero_info(pprints_config(config))
# Instantiate classes
if mode == "fit":
train_data_connector = instantiate_classes(config.train_data_connector)
loss = instantiate_classes(config.loss)
else:
train_data_connector = None
loss = None
if config.test_data_connector is not None:
test_data_connector = instantiate_classes(config.test_data_connector)
else:
test_data_connector = None
# Callbacks
vis = _VIS.value
callbacks: list[Callback] = []
for cb in config.callbacks:
callback = instantiate_classes(cb)
assert isinstance(callback, Callback), (
"Callback must be a subclass of Callback. "
f"Provided callback: {cb} is not!"
)
if not vis and isinstance(callback, VisualizerCallback):
rank_zero_info(
f"{callback.visualizer} is not used. "
"Please set --vis=True to use it."
)
continue
callbacks.append(callback)
# Add needed callbacks
callbacks.append(LRSchedulerCallback())
# Checkpoint path
ckpt_path = _CKPT.value
# Resume training
resume = _RESUME.value
if resume:
if ckpt_path is None:
resume_ckpt_path = osp.join(
config.output_dir, "checkpoints/last.ckpt"
)
else:
resume_ckpt_path = ckpt_path
else:
resume_ckpt_path = None
trainer = PLTrainer(callbacks=callbacks, **trainer_args)
hyper_params = trainer_args
if config.get("params", None) is not None:
hyper_params.update(config.params.to_dict())
training_module = TrainingModule(
config.model,
config.optimizers,
loss,
train_data_connector,
test_data_connector,
hyper_params,
config.seed,
ckpt_path if not resume else None,
config.check_unused_parameters,
)
data_module = DataModule(config.data)
if mode == "fit":
trainer.fit(
training_module, datamodule=data_module, ckpt_path=resume_ckpt_path
)
elif mode == "test":
trainer.test(training_module, datamodule=data_module, verbose=False)
def entrypoint() -> None:
"""Entry point for the CLI."""
app.run(main)
if __name__ == "__main__":
entrypoint()
|