| import os |
| import copy |
| import pytorch_lightning as pl |
| import datetime |
| import wandb |
|
|
| from pytorch_lightning.loggers import WandbLogger |
| from model.model_interface import ModelInterface |
| from scripts.dataset.data_interface import DataInterface |
| from pytorch_lightning.strategies import DDPStrategy |
|
|
|
|
| def load_wandb(config): |
| |
| wandb_config = config.setting.wandb_config |
| wandb_logger = WandbLogger(project=wandb_config.project, config=config, |
| name=wandb_config.name, |
| settings=wandb.Settings(start_method='fork')) |
| |
| return wandb_logger |
|
|
|
|
| def load_model(config): |
| |
| model_config = copy.deepcopy(config) |
| kwargs = model_config.pop('kwargs') |
| model_config.update(kwargs) |
| return ModelInterface.init_model(**model_config) |
|
|
|
|
| def load_dataset(config): |
| |
| dataset_config = copy.deepcopy(config) |
| kwargs = dataset_config.pop('kwargs') |
| dataset_config.update(kwargs) |
| return DataInterface.init_dataset(**dataset_config) |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| |
| def load_strategy(config): |
| config = copy.deepcopy(config) |
| if "timeout" in config.keys(): |
| timeout = int(config.pop('timeout')) |
| config["timeout"] = datetime.timedelta(seconds=timeout) |
|
|
| return DDPStrategy(**config) |
|
|
|
|
| |
| def load_trainer(config): |
| trainer_config = copy.deepcopy(config.Trainer) |
| |
| |
| if trainer_config.logger: |
| trainer_config.logger = load_wandb(config) |
| else: |
| trainer_config.logger = False |
|
|
| |
| |
| |
| |
| strategy = load_strategy(trainer_config.pop('strategy')) |
| return pl.Trainer(**trainer_config, strategy=strategy, callbacks=[]) |
|
|