File size: 12,267 Bytes
9572863
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
from __gin__ import dynamic_registration
import __main__ as train_script
from music_spectrogram_diffusion import audio_codecs
from music_spectrogram_diffusion.models.diffusion import diffusion_utils
from music_spectrogram_diffusion.models.diffusion import models
from music_spectrogram_diffusion.models.diffusion import network
from music_spectrogram_diffusion import preprocessors
from music_spectrogram_diffusion import tasks
from music_spectrogram_diffusion import vocabularies
import seqio
from t5x import adafactor
from t5x import gin_utils
from t5x import partitioning
from t5x import trainer
from t5x import utils

# Macros:
# ==============================================================================
AUDIO_CODEC = @audio_codecs.MelGAN()
BATCH_SIZE = 1024
DATASET_NAME = 'mega'
EVAL_STEPS = 20
EVALUATOR_NUM_EXAMPLES = None
EVALUATOR_USE_MEMORY_CACHE = True
INFER_EVAL_TASK_NAME = @infer_eval/tasks.construct_task_name()
INFER_TASK_NAME = @infer/tasks.construct_task_name()
INPUT_VOCABULARY = @vocabularies.vocabulary_from_codec()
JSON_WRITE_N_RESULTS = 0
LABEL_SMOOTHING = 0.0
LOSS_NORMALIZING_FACTOR = None
MODEL = @models.ContextDiffusionModel()
MODEL_DIR = ''
NUM_MICROBATCHES = None
NUM_VELOCITY_BINS = 1
ONSETS_ONLY = False
OPTIMIZER = @adafactor.Adafactor()
PROGRAM_GRANULARITY = 'full'
TASK_FEATURE_LENGTHS = {'inputs': 2048, 'targets': 256, 'targets_context': 256}
TASK_PREFIX = 'synthesis_with_context'
TEST_TASK_NAME = %INFER_TASK_NAME
TRAIN_EVAL_TASK_NAME = %TRAIN_TASK_NAME
TRAIN_STEPS = 500000
TRAIN_TASK_NAME = @train/tasks.construct_task_name()
USE_CACHED_TASKS = True
USE_TIES = True
VOCAB_CONFIG = @vocabularies.VocabularyConfig()
Z_LOSS = 0.0001

# Parameters for adafactor.Adafactor:
# ==============================================================================
adafactor.Adafactor.decay_rate = 0.8
adafactor.Adafactor.logical_factor_rules = \
    @adafactor.standard_logical_factor_rules()
adafactor.Adafactor.step_offset = 0

# Parameters for vocabularies.build_codec:
# ==============================================================================
vocabularies.build_codec.vocab_config = %VOCAB_CONFIG

# Parameters for utils.CheckpointConfig:
# ==============================================================================
utils.CheckpointConfig.restore = None
utils.CheckpointConfig.save = @utils.SaveCheckpointConfig()

# Parameters for infer/tasks.construct_task_name:
# ==============================================================================
infer/tasks.construct_task_name.audio_codec = %AUDIO_CODEC
infer/tasks.construct_task_name.dataset_name = %DATASET_NAME
infer/tasks.construct_task_name.note_representation_config = \
    @tasks.NoteRepresentationConfig()
infer/tasks.construct_task_name.task_prefix = %TASK_PREFIX
infer/tasks.construct_task_name.task_suffix = 'test'
infer/tasks.construct_task_name.vocab_config = %VOCAB_CONFIG

# Parameters for infer_eval/tasks.construct_task_name:
# ==============================================================================
infer_eval/tasks.construct_task_name.audio_codec = %AUDIO_CODEC
infer_eval/tasks.construct_task_name.dataset_name = %DATASET_NAME
infer_eval/tasks.construct_task_name.note_representation_config = \
    @tasks.NoteRepresentationConfig()
infer_eval/tasks.construct_task_name.task_prefix = %TASK_PREFIX
infer_eval/tasks.construct_task_name.task_suffix = 'eval'
infer_eval/tasks.construct_task_name.vocab_config = %VOCAB_CONFIG

# Parameters for train/tasks.construct_task_name:
# ==============================================================================
train/tasks.construct_task_name.audio_codec = %AUDIO_CODEC
train/tasks.construct_task_name.dataset_name = %DATASET_NAME
train/tasks.construct_task_name.note_representation_config = \
    @tasks.NoteRepresentationConfig()
train/tasks.construct_task_name.task_prefix = %TASK_PREFIX
train/tasks.construct_task_name.task_suffix = 'train'
train/tasks.construct_task_name.vocab_config = %VOCAB_CONFIG

# Parameters for models.ContextDiffusionModel:
# ==============================================================================
models.ContextDiffusionModel.audio_codec = %AUDIO_CODEC
models.ContextDiffusionModel.diffusion_config = @diffusion_utils.DiffusionConfig()
models.ContextDiffusionModel.input_vocabulary = %INPUT_VOCABULARY
models.ContextDiffusionModel.module = @network.ContinuousContextTransformer()
models.ContextDiffusionModel.optimizer_def = %OPTIMIZER
models.ContextDiffusionModel.output_vocabulary = \
    @seqio.vocabularies.PassThroughVocabulary()

# Parameters for network.ContinuousContextTransformer:
# ==============================================================================
network.ContinuousContextTransformer.config = @network.T5Config()

# Parameters for utils.create_learning_rate_scheduler:
# ==============================================================================
utils.create_learning_rate_scheduler.base_learning_rate = 0.001
utils.create_learning_rate_scheduler.factors = 'constant'
utils.create_learning_rate_scheduler.warmup_steps = 1000

# Parameters for infer_eval/utils.DatasetConfig:
# ==============================================================================
infer_eval/utils.DatasetConfig.batch_size = %BATCH_SIZE
infer_eval/utils.DatasetConfig.mixture_or_task_name = %INFER_EVAL_TASK_NAME
infer_eval/utils.DatasetConfig.pack = False
infer_eval/utils.DatasetConfig.seed = 42
infer_eval/utils.DatasetConfig.shuffle = False
infer_eval/utils.DatasetConfig.split = 'eval'
infer_eval/utils.DatasetConfig.task_feature_lengths = %TASK_FEATURE_LENGTHS
infer_eval/utils.DatasetConfig.use_cached = %USE_CACHED_TASKS

# Parameters for train/utils.DatasetConfig:
# ==============================================================================
train/utils.DatasetConfig.batch_size = %BATCH_SIZE
train/utils.DatasetConfig.mixture_or_task_name = %TRAIN_TASK_NAME
train/utils.DatasetConfig.pack = False
train/utils.DatasetConfig.seed = None
train/utils.DatasetConfig.shuffle = True
train/utils.DatasetConfig.split = 'train'
train/utils.DatasetConfig.task_feature_lengths = %TASK_FEATURE_LENGTHS
train/utils.DatasetConfig.use_cached = %USE_CACHED_TASKS

# Parameters for train_eval/utils.DatasetConfig:
# ==============================================================================
train_eval/utils.DatasetConfig.batch_size = %BATCH_SIZE
train_eval/utils.DatasetConfig.mixture_or_task_name = %TRAIN_EVAL_TASK_NAME
train_eval/utils.DatasetConfig.pack = False
train_eval/utils.DatasetConfig.seed = 42
train_eval/utils.DatasetConfig.shuffle = False
train_eval/utils.DatasetConfig.split = 'eval'
train_eval/utils.DatasetConfig.task_feature_lengths = %TASK_FEATURE_LENGTHS
train_eval/utils.DatasetConfig.use_cached = %USE_CACHED_TASKS

# Parameters for diffusion_utils.DiffusionConfig:
# ==============================================================================
diffusion_utils.DiffusionConfig.classifier_free_guidance = \
    @diffusion_utils.ClassifierFreeGuidanceConfig()
diffusion_utils.DiffusionConfig.sampler = @diffusion_utils.SamplerConfig()
diffusion_utils.DiffusionConfig.train_schedule = \
    @train/diffusion_utils.DiffusionSchedule()

# Parameters for models.DiffusionModel.loss_fn:
# ==============================================================================
models.DiffusionModel.loss_fn.label_smoothing = %LABEL_SMOOTHING
models.DiffusionModel.loss_fn.loss_normalizing_factor = %LOSS_NORMALIZING_FACTOR
models.DiffusionModel.loss_fn.z_loss = %Z_LOSS

# Parameters for sampler/diffusion_utils.DiffusionSchedule:
# ==============================================================================
sampler/diffusion_utils.DiffusionSchedule.name = 'cosine'
sampler/diffusion_utils.DiffusionSchedule.num_steps = 1000

# Parameters for train/diffusion_utils.DiffusionSchedule:
# ==============================================================================
train/diffusion_utils.DiffusionSchedule.name = 'cosine'

# Parameters for seqio.Evaluator:
# ==============================================================================
seqio.Evaluator.logger_cls = \
    [@seqio.PyLoggingLogger, @seqio.TensorBoardLogger, @seqio.JSONLogger]
seqio.Evaluator.num_examples = %EVALUATOR_NUM_EXAMPLES
seqio.Evaluator.use_memory_cache = %EVALUATOR_USE_MEMORY_CACHE

# Parameters for seqio.JSONLogger:
# ==============================================================================
seqio.JSONLogger.write_n_results = %JSON_WRITE_N_RESULTS

# Parameters for preprocessors.map_midi_programs:
# ==============================================================================
preprocessors.map_midi_programs.granularity_type = %PROGRAM_GRANULARITY

# Parameters for tasks.NoteRepresentationConfig:
# ==============================================================================
tasks.NoteRepresentationConfig.include_ties = True
tasks.NoteRepresentationConfig.onsets_only = False

# Parameters for vocabularies.num_embeddings:
# ==============================================================================
vocabularies.num_embeddings.vocabulary = %INPUT_VOCABULARY

# Parameters for seqio.vocabularies.PassThroughVocabulary:
# ==============================================================================
seqio.vocabularies.PassThroughVocabulary.size = 0

# Parameters for partitioning.PjitPartitioner:
# ==============================================================================
partitioning.PjitPartitioner.model_parallel_submesh = None
partitioning.PjitPartitioner.num_partitions = 1

# Parameters for diffusion_utils.SamplerConfig:
# ==============================================================================
diffusion_utils.SamplerConfig.schedule = \
    @sampler/diffusion_utils.DiffusionSchedule()

# Parameters for utils.SaveCheckpointConfig:
# ==============================================================================
utils.SaveCheckpointConfig.dtype = 'float32'
utils.SaveCheckpointConfig.keep = None
utils.SaveCheckpointConfig.period = 10000
utils.SaveCheckpointConfig.save_dataset = False

# Parameters for network.T5Config:
# ==============================================================================
network.T5Config.context_positions = 'terminal_relative'
network.T5Config.decoder_cross_attend_style = 'concat_encodings'
network.T5Config.dropout_rate = 0.1
network.T5Config.dtype = 'float32'
network.T5Config.emb_dim = 768
network.T5Config.head_dim = 64
network.T5Config.mlp_activations = ('gelu', 'linear')
network.T5Config.mlp_dim = 2048
network.T5Config.num_decoder_layers = 12
network.T5Config.num_encoder_layers = 12
network.T5Config.num_heads = 12
network.T5Config.position_encoding = 'fixed_permuted_offset'
network.T5Config.vocab_size = @vocabularies.num_embeddings()

# Parameters for train_script.train:
# ==============================================================================
train_script.train.checkpoint_cfg = @utils.CheckpointConfig()
train_script.train.eval_period = 10000
train_script.train.eval_steps = %EVAL_STEPS
train_script.train.infer_eval_dataset_cfg = @infer_eval/utils.DatasetConfig()
train_script.train.inference_evaluator_cls = @seqio.Evaluator
train_script.train.model = %MODEL
train_script.train.model_dir = ''
train_script.train.partitioner = @partitioning.PjitPartitioner()
train_script.train.random_seed = None
train_script.train.summarize_config_fn = @gin_utils.summarize_gin_config
train_script.train.total_steps = %TRAIN_STEPS
train_script.train.train_dataset_cfg = @train/utils.DatasetConfig()
train_script.train.train_eval_dataset_cfg = @train_eval/utils.DatasetConfig()
train_script.train.trainer_cls = @trainer.Trainer

# Parameters for trainer.Trainer:
# ==============================================================================
trainer.Trainer.learning_rate_fn = @utils.create_learning_rate_scheduler()
trainer.Trainer.num_microbatches = %NUM_MICROBATCHES

# Parameters for vocabularies.vocabulary_from_codec:
# ==============================================================================
vocabularies.vocabulary_from_codec.codec = @vocabularies.build_codec()

# Parameters for vocabularies.VocabularyConfig:
# ==============================================================================
vocabularies.VocabularyConfig.num_velocity_bins = %NUM_VELOCITY_BINS