CoolFace
Apppublic

mosibi/RVC_HFv2

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
parser.py245 linesDownload Raw Back to demucs
1# Copyright (c) Facebook, Inc. and its affiliates.2# All rights reserved.3#4# This source code is licensed under the license found in the5# LICENSE file in the root directory of this source tree.6 7import argparse8import os9from pathlib import Path10 11 12def get_parser():13    parser = argparse.ArgumentParser("demucs", description="Train and evaluate Demucs.")14    default_raw = None15    default_musdb = None16    if 'DEMUCS_RAW' in os.environ:17        default_raw = Path(os.environ['DEMUCS_RAW'])18    if 'DEMUCS_MUSDB' in os.environ:19        default_musdb = Path(os.environ['DEMUCS_MUSDB'])20    parser.add_argument(21        "--raw",22        type=Path,23        default=default_raw,24        help="Path to raw audio, can be faster, see python3 -m demucs.raw to extract.")25    parser.add_argument("--no_raw", action="store_const", const=None, dest="raw")26    parser.add_argument("-m",27                        "--musdb",28                        type=Path,29                        default=default_musdb,30                        help="Path to musdb root")31    parser.add_argument("--is_wav", action="store_true",32                        help="Indicate that the MusDB dataset is in wav format (i.e. MusDB-HQ).")33    parser.add_argument("--metadata", type=Path, default=Path("metadata/"),34                        help="Folder where metadata information is stored.")35    parser.add_argument("--wav", type=Path,36                        help="Path to a wav dataset. This should contain a 'train' and a 'valid' "37                             "subfolder.")38    parser.add_argument("--samplerate", type=int, default=44100)39    parser.add_argument("--audio_channels", type=int, default=2)40    parser.add_argument("--samples",41                        default=44100 * 10,42                        type=int,43                        help="number of samples to feed in")44    parser.add_argument("--data_stride",45                        default=44100,46                        type=int,47                        help="Stride for chunks, shorter = longer epochs")48    parser.add_argument("-w", "--workers", default=10, type=int, help="Loader workers")49    parser.add_argument("--eval_workers", default=2, type=int, help="Final evaluation workers")50    parser.add_argument("-d",51                        "--device",52                        help="Device to train on, default is cuda if available else cpu")53    parser.add_argument("--eval_cpu", action="store_true", help="Eval on test will be run on cpu.")54    parser.add_argument("--dummy", help="Dummy parameter, useful to create a new checkpoint file")55    parser.add_argument("--test", help="Just run the test pipeline + one validation. "56                                       "This should be a filename relative to the models/ folder.")57    parser.add_argument("--test_pretrained", help="Just run the test pipeline + one validation, "58                                                  "on a pretrained model. ")59 60    parser.add_argument("--rank", default=0, type=int)61    parser.add_argument("--world_size", default=1, type=int)62    parser.add_argument("--master")63 64    parser.add_argument("--checkpoints",65                        type=Path,66                        default=Path("checkpoints"),67                        help="Folder where to store checkpoints etc")68    parser.add_argument("--evals",69                        type=Path,70                        default=Path("evals"),71                        help="Folder where to store evals and waveforms")72    parser.add_argument("--save",73                        action="store_true",74                        help="Save estimated for the test set waveforms")75    parser.add_argument("--logs",76                        type=Path,77                        default=Path("logs"),78                        help="Folder where to store logs")79    parser.add_argument("--models",80                        type=Path,81                        default=Path("models"),82                        help="Folder where to store trained models")83    parser.add_argument("-R",84                        "--restart",85                        action='store_true',86                        help='Restart training, ignoring previous run')87 88    parser.add_argument("--seed", type=int, default=42)89    parser.add_argument("-e", "--epochs", type=int, default=180, help="Number of epochs")90    parser.add_argument("-r",91                        "--repeat",92                        type=int,93                        default=2,94                        help="Repeat the train set, longer epochs")95    parser.add_argument("-b", "--batch_size", type=int, default=64)96    parser.add_argument("--lr", type=float, default=3e-4)97    parser.add_argument("--mse", action="store_true", help="Use MSE instead of L1")98    parser.add_argument("--init", help="Initialize from a pre-trained model.")99 100    # Augmentation options101    parser.add_argument("--no_augment",102                        action="store_false",103                        dest="augment",104                        default=True,105                        help="No basic data augmentation.")106    parser.add_argument("--repitch", type=float, default=0.2,107                        help="Probability to do tempo/pitch change")108    parser.add_argument("--max_tempo", type=float, default=12,109                        help="Maximum relative tempo change in %% when using repitch.")110 111    parser.add_argument("--remix_group_size",112                        type=int,113                        default=4,114                        help="Shuffle sources using group of this size. Useful to somewhat "115                        "replicate multi-gpu training "116                        "on less GPUs.")117    parser.add_argument("--shifts",118                        type=int,119                        default=10,120                        help="Number of random shifts used for the shift trick.")121    parser.add_argument("--overlap",122                        type=float,123                        default=0.25,124                        help="Overlap when --split_valid is passed.")125 126    # See model.py for doc127    parser.add_argument("--growth",128                        type=float,129                        default=2.,130                        help="Number of channels between two layers will increase by this factor")131    parser.add_argument("--depth",132                        type=int,133                        default=6,134                        help="Number of layers for the encoder and decoder")135    parser.add_argument("--lstm_layers", type=int, default=2, help="Number of layers for the LSTM")136    parser.add_argument("--channels",137                        type=int,138                        default=64,139                        help="Number of channels for the first encoder layer")140    parser.add_argument("--kernel_size",141                        type=int,142                        default=8,143                        help="Kernel size for the (transposed) convolutions")144    parser.add_argument("--conv_stride",145                        type=int,146                        default=4,147                        help="Stride for the (transposed) convolutions")148    parser.add_argument("--context",149                        type=int,150                        default=3,151                        help="Context size for the decoder convolutions "152                        "before the transposed convolutions")153    parser.add_argument("--rescale",154                        type=float,155                        default=0.1,156                        help="Initial weight rescale reference")157    parser.add_argument("--no_resample", action="store_false",158                        default=True, dest="resample",159                        help="No Resampling of the input/output x2")160    parser.add_argument("--no_glu",161                        action="store_false",162                        default=True,163                        dest="glu",164                        help="Replace all GLUs by ReLUs")165    parser.add_argument("--no_rewrite",166                        action="store_false",167                        default=True,168                        dest="rewrite",169                        help="No 1x1 rewrite convolutions")170    parser.add_argument("--normalize", action="store_true")171    parser.add_argument("--no_norm_wav", action="store_false", dest='norm_wav', default=True)172 173    # Tasnet options174    parser.add_argument("--tasnet", action="store_true")175    parser.add_argument("--split_valid",176                        action="store_true",177                        help="Predict chunks by chunks for valid and test. Required for tasnet")178    parser.add_argument("--X", type=int, default=8)179 180    # Other options181    parser.add_argument("--show",182                        action="store_true",183                        help="Show model architecture, size and exit")184    parser.add_argument("--save_model", action="store_true",185                        help="Skip traning, just save final model "186                             "for the current checkpoint value.")187    parser.add_argument("--save_state",188                        help="Skip training, just save state "189                             "for the current checkpoint value. You should "190                             "provide a model name as argument.")191 192    # Quantization options193    parser.add_argument("--q-min-size", type=float, default=1,194                        help="Only quantize layers over this size (in MB)")195    parser.add_argument(196        "--qat", type=int, help="If provided, use QAT training with that many bits.")197 198    parser.add_argument("--diffq", type=float, default=0)199    parser.add_argument(200        "--ms-target", type=float, default=162,201        help="Model size target in MB, when using DiffQ. Best model will be kept "202             "only if it is smaller than this target.")203 204    return parser205 206 207def get_name(parser, args):208    """209    Return the name of an experiment given the args. Some parameters are ignored,210    for instance --workers, as they do not impact the final result.211    """212    ignore_args = set([213        "checkpoints",214        "deterministic",215        "eval",216        "evals",217        "eval_cpu",218        "eval_workers",219        "logs",220        "master",221        "rank",222        "restart",223        "save",224        "save_model",225        "save_state",226        "show",227        "workers",228        "world_size",229    ])230    parts = []231    name_args = dict(args.__dict__)232    for name, value in name_args.items():233        if name in ignore_args:234            continue235        if value != parser.get_default(name):236            if isinstance(value, Path):237                parts.append(f"{name}={value.name}")238            else:239                parts.append(f"{name}={value}")240    if parts:241        name = " ".join(parts)242    else:243        name = "default"244    return name245