File size: 2,625 Bytes
2e1dc7f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
"""
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()