Team Ai
Apppublic

jone/Music_Source_Separation

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
train.py300 linesDownload Raw Back to bytesep
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