mosibi/RVC_HFv2
0
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 