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