| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| adafactor.Adafactor.decay_rate = 0.8 |
| adafactor.Adafactor.logical_factor_rules = \ |
| @adafactor.standard_logical_factor_rules() |
| adafactor.Adafactor.step_offset = 0 |
|
|
| |
| |
| vocabularies.build_codec.vocab_config = %VOCAB_CONFIG |
|
|
| |
| |
| utils.CheckpointConfig.restore = None |
| utils.CheckpointConfig.save = @utils.SaveCheckpointConfig() |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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() |
|
|
| |
| |
| network.ContinuousContextTransformer.config = @network.T5Config() |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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 |
|
|
| |
| |
| 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() |
|
|
| |
| |
| 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 |
|
|
| |
| |
| sampler/diffusion_utils.DiffusionSchedule.name = 'cosine' |
| sampler/diffusion_utils.DiffusionSchedule.num_steps = 1000 |
|
|
| |
| |
| train/diffusion_utils.DiffusionSchedule.name = 'cosine' |
|
|
| |
| |
| 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 |
|
|
| |
| |
| seqio.JSONLogger.write_n_results = %JSON_WRITE_N_RESULTS |
|
|
| |
| |
| preprocessors.map_midi_programs.granularity_type = %PROGRAM_GRANULARITY |
|
|
| |
| |
| tasks.NoteRepresentationConfig.include_ties = True |
| tasks.NoteRepresentationConfig.onsets_only = False |
|
|
| |
| |
| vocabularies.num_embeddings.vocabulary = %INPUT_VOCABULARY |
|
|
| |
| |
| seqio.vocabularies.PassThroughVocabulary.size = 0 |
|
|
| |
| |
| partitioning.PjitPartitioner.model_parallel_submesh = None |
| partitioning.PjitPartitioner.num_partitions = 1 |
|
|
| |
| |
| diffusion_utils.SamplerConfig.schedule = \ |
| @sampler/diffusion_utils.DiffusionSchedule() |
|
|
| |
| |
| utils.SaveCheckpointConfig.dtype = 'float32' |
| utils.SaveCheckpointConfig.keep = None |
| utils.SaveCheckpointConfig.period = 10000 |
| utils.SaveCheckpointConfig.save_dataset = False |
|
|
| |
| |
| 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() |
|
|
| |
| |
| 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 |
|
|
| |
| |
| trainer.Trainer.learning_rate_fn = @utils.create_learning_rate_scheduler() |
| trainer.Trainer.num_microbatches = %NUM_MICROBATCHES |
|
|
| |
| |
| vocabularies.vocabulary_from_codec.codec = @vocabularies.build_codec() |
|
|
| |
| |
| vocabularies.VocabularyConfig.num_velocity_bins = %NUM_VELOCITY_BINS |
|
|