Spaces:
Running on Zero
Running on Zero
| """ | |
| Train STAGE for DRUMS generation, with either a mixture or a metronome track as context | |
| """ | |
| import torch | |
| import lightning as L | |
| import hyperparameters as hp | |
| from conditioning.condition_type import ConditionType | |
| from conditioning.conditioning_method import ConditioningMethod | |
| from conditioning.prompt_processor import InterleavedContextPromptProcessor | |
| from conditioning.t5embedder import T5EmbedderGPU | |
| from data.stem import Stem | |
| from training.train import train | |
| import config as cfg | |
| RUN_NAME = "stage-drums" | |
| # distributed strategy - set to DDP to train on multi-gpu machines | |
| STRATEGY = None | |
| def launch_train(): | |
| lm_params = hp.PretrainedSmallLmParams(sep_token=2049) | |
| conditioning_params = hp.ConditioningParams( | |
| embedder_types={ | |
| ConditionType.DESCRIPTION: T5EmbedderGPU, | |
| }, | |
| conditioning_methods={ | |
| ConditionType.DESCRIPTION: ConditioningMethod.CROSS_ATTENTION, | |
| }, | |
| conditioning_dropout=0.5) | |
| prompt_params = hp.PromptProcessorParams( | |
| keep_only_valid_steps=True, | |
| model_class=InterleavedContextPromptProcessor, | |
| context_dropout=0.1) | |
| encodec_params = hp.pretrained_encodec_meta_32khz_params | |
| model_params = hp.MusicgenParams(encodec_params=encodec_params, | |
| prompt_processor_params=prompt_params, | |
| conditioning_params=conditioning_params, | |
| lm_params=lm_params) | |
| batch_size_train = 2 | |
| max_steps = 100_000 | |
| max_time = "00:24:00:00" | |
| n_samples_per_epoch = 10_000 | |
| accumulate_grad_batches = 4 | |
| dataset_params = hp.StemmedDatasetParams( | |
| clip_length_in_seconds=10, | |
| sample_rate=32_000, | |
| root_dir=cfg.moises_path(), | |
| single_stem=True, | |
| target_stem=Stem.DRUMS, | |
| min_context_seconds=5, | |
| use_style_conditioning=True, | |
| use_beat_conditioning=True, | |
| type_of_context="stems or beats", | |
| add_click=False, | |
| sync_chunks=False, | |
| bpm_in_caption=False, | |
| batch_size_train=batch_size_train, | |
| batch_size_test=12, | |
| num_workers=11, | |
| speed_transform_p=0.5, | |
| pitch_transform_p=0.5, | |
| n_samples_per_epoch=n_samples_per_epoch) | |
| train( | |
| model_params=model_params, | |
| dataset_params=dataset_params, | |
| max_time=max_time, | |
| max_steps=max_steps, | |
| accumulate_grad_batches=accumulate_grad_batches, | |
| run_name=RUN_NAME, | |
| distributed_strategy=STRATEGY, | |
| log=True, | |
| kill_on_end=False, | |
| ) | |
| if __name__ == "__main__": | |
| launch_train() | |