jone/Music_Source_Separation
3
1import argparse2import logging3import os4import pathlib5from functools import partial6from typing import List, NoReturn7 8import pytorch_lightning as pl9from pytorch_lightning.plugins import DDPPlugin10 11from bytesep.callbacks import get_callbacks12from bytesep.data.augmentors import Augmentor13from bytesep.data.batch_data_preprocessors import (14 get_batch_data_preprocessor_class,15)16from bytesep.data.data_modules import DataModule, Dataset17from bytesep.data.samplers import SegmentSampler18from bytesep.losses import get_loss_function19from bytesep.models.lightning_modules import (20 LitSourceSeparation,21 get_model_class,22)23from bytesep.optimizers.lr_schedulers import get_lr_lambda24from bytesep.utils import (25 create_logging,26 get_pitch_shift_factor,27 read_yaml,28 check_configs_gramma,29)30 31 32def get_dirs(33 workspace: str, task_name: str, filename: str, config_yaml: str, gpus: int34) -> List[str]:35 r"""Get directories.36 37 Args:38 workspace: str39 task_name, str, e.g., 'musdb18'40 filenmae: str41 config_yaml: str42 gpus: int, e.g., 0 for cpu and 8 for training with 8 gpu cards43 44 Returns:45 checkpoints_dir: str46 logs_dir: str47 logger: pl.loggers.TensorBoardLogger48 statistics_path: str49 """50 51 # save checkpoints dir52 checkpoints_dir = os.path.join(53 workspace,54 "checkpoints",55 task_name,56 filename,57 "config={},gpus={}".format(pathlib.Path(config_yaml).stem, gpus),58 )59 os.makedirs(checkpoints_dir, exist_ok=True)60 61 # logs dir62 logs_dir = os.path.join(63 workspace,64 "logs",65 task_name,66 filename,67 "config={},gpus={}".format(pathlib.Path(config_yaml).stem, gpus),68 )69 os.makedirs(logs_dir, exist_ok=True)70 71 # loggings72 create_logging(logs_dir, filemode='w')73 logging.info(args)74 75 # tensorboard logs dir76 tb_logs_dir = os.path.join(workspace, "tensorboard_logs")77 os.makedirs(tb_logs_dir, exist_ok=True)78 79 experiment_name = os.path.join(task_name, filename, pathlib.Path(config_yaml).stem)80 logger = pl.loggers.TensorBoardLogger(save_dir=tb_logs_dir, name=experiment_name)81 82 # statistics path83 statistics_path = os.path.join(84 workspace,85 "statistics",86 task_name,87 filename,88 "config={},gpus={}".format(pathlib.Path(config_yaml).stem, gpus),89 "statistics.pkl",90 )91 os.makedirs(os.path.dirname(statistics_path), exist_ok=True)92 93 return checkpoints_dir, logs_dir, logger, statistics_path94 95 96def _get_data_module(97 workspace: str, config_yaml: str, num_workers: int, distributed: bool98) -> DataModule:99 r"""Create data_module. Mini-batch data can be obtained by:100 101 code-block:: python102 103 data_module.setup()104 for batch_data_dict in data_module.train_dataloader():105 print(batch_data_dict.keys())106 break107 108 Args:109 workspace: str110 config_yaml: str111 num_workers: int, e.g., 0 for non-parallel and 8 for using cpu cores112 for preparing data in parallel113 distributed: bool114 115 Returns:116 data_module: DataModule117 """118 119 configs = read_yaml(config_yaml)120 input_source_types = configs['train']['input_source_types']121 indexes_path = os.path.join(workspace, configs['train']['indexes_dict'])122 sample_rate = configs['train']['sample_rate']123 segment_seconds = configs['train']['segment_seconds']124 mixaudio_dict = configs['train']['augmentations']['mixaudio']125 augmentations = configs['train']['augmentations']126 max_pitch_shift = max(127 [128 augmentations['pitch_shift'][source_type]129 for source_type in input_source_types130 ]131 )132 batch_size = configs['train']['batch_size']133 steps_per_epoch = configs['train']['steps_per_epoch']134 135 segment_samples = int(segment_seconds * sample_rate)136 ex_segment_samples = int(segment_samples * get_pitch_shift_factor(max_pitch_shift))137 138 # sampler139 train_sampler = SegmentSampler(140 indexes_path=indexes_path,141 segment_samples=ex_segment_samples,142 mixaudio_dict=mixaudio_dict,143 batch_size=batch_size,144 steps_per_epoch=steps_per_epoch,145 )146 147 # augmentor148 augmentor = Augmentor(augmentations=augmentations)149 150 # dataset151 train_dataset = Dataset(augmentor, segment_samples)152 153 # data module154 data_module = DataModule(155 train_sampler=train_sampler,156 train_dataset=train_dataset,157 num_workers=num_workers,158 distributed=distributed,159 )160 161 return data_module162 163 164def train(args) -> NoReturn:165 r"""Train & evaluate and save checkpoints.166 167 Args:168 workspace: str, directory of workspace169 gpus: int170 config_yaml: str, path of config file for training171 """172 173 # arugments & parameters174 workspace = args.workspace175 gpus = args.gpus176 config_yaml = args.config_yaml177 filename = args.filename178 179 num_workers = 8180 distributed = True if gpus > 1 else False181 evaluate_device = "cuda" if gpus > 0 else "cpu"182 183 # Read config file.184 configs = read_yaml(config_yaml)185 check_configs_gramma(configs)186 task_name = configs['task_name']187 target_source_types = configs['train']['target_source_types']188 target_sources_num = len(target_source_types)189 channels = configs['train']['channels']190 batch_data_preprocessor_type = configs['train']['batch_data_preprocessor']191 model_type = configs['train']['model_type']192 loss_type = configs['train']['loss_type']193 optimizer_type = configs['train']['optimizer_type']194 learning_rate = float(configs['train']['learning_rate'])195 precision = configs['train']['precision']196 early_stop_steps = configs['train']['early_stop_steps']197 warm_up_steps = configs['train']['warm_up_steps']198 reduce_lr_steps = configs['train']['reduce_lr_steps']199 200 # paths201 checkpoints_dir, logs_dir, logger, statistics_path = get_dirs(202 workspace, task_name, filename, config_yaml, gpus203 )204 205 # training data module206 data_module = _get_data_module(207 workspace=workspace,208 config_yaml=config_yaml,209 num_workers=num_workers,210 distributed=distributed,211 )212 213 # batch data preprocessor214 BatchDataPreprocessor = get_batch_data_preprocessor_class(215 batch_data_preprocessor_type=batch_data_preprocessor_type216 )217 218 batch_data_preprocessor = BatchDataPreprocessor(219 target_source_types=target_source_types220 )221 222 # model223 Model = get_model_class(model_type=model_type)224 model = Model(input_channels=channels, target_sources_num=target_sources_num)225 226 # loss function227 loss_function = get_loss_function(loss_type=loss_type)228 229 # callbacks230 callbacks = get_callbacks(231 task_name=task_name,232 config_yaml=config_yaml,233 workspace=workspace,234 checkpoints_dir=checkpoints_dir,235 statistics_path=statistics_path,236 logger=logger,237 model=model,238 evaluate_device=evaluate_device,239 )240 # callbacks = []241 242 # learning rate reduce function243 lr_lambda = partial(244 get_lr_lambda, warm_up_steps=warm_up_steps, reduce_lr_steps=reduce_lr_steps245 )246 247 # pytorch-lightning model248 pl_model = LitSourceSeparation(249 batch_data_preprocessor=batch_data_preprocessor,250 model=model,251 optimizer_type=optimizer_type,252 loss_function=loss_function,253 learning_rate=learning_rate,254 lr_lambda=lr_lambda,255 )256 257 # trainer258 trainer = pl.Trainer(259 checkpoint_callback=False,260 gpus=gpus,261 callbacks=callbacks,262 max_steps=early_stop_steps,263 accelerator="ddp",264 sync_batchnorm=True,265 precision=precision,266 replace_sampler_ddp=False,267 plugins=[DDPPlugin(find_unused_parameters=True)],268 profiler='simple',269 )270 271 # Fit, evaluate, and save checkpoints.272 trainer.fit(pl_model, data_module)273 274 275if __name__ == "__main__":276 277 parser = argparse.ArgumentParser(description="")278 subparsers = parser.add_subparsers(dest="mode")279 280 parser_train = subparsers.add_parser("train")281 parser_train.add_argument(282 "--workspace", type=str, required=True, help="Directory of workspace."283 )284 parser_train.add_argument("--gpus", type=int, required=True)285 parser_train.add_argument(286 "--config_yaml",287 type=str,288 required=True,289 help="Path of config file for training.",290 )291 292 args = parser.parse_args()293 args.filename = pathlib.Path(__file__).stem294 295 if args.mode == "train":296 train(args)297 298 else:299 raise Exception("Error argument!")300 