Download config.gin from harp-dev/music-spectrogram-diffusion-checkpoints: direct link, hf CLI and curl.
- Browser
- Download file 12.3 kB
-
https://huggingface.co/harp-dev/music-spectrogram-diffusion-checkpoints/resolve/main/config.gin
- Command line
-
hf download hf://harp-dev/music-spectrogram-diffusion-checkpoints/config.gin
-
curl -L -o config.gin https://huggingface.co/harp-dev/music-spectrogram-diffusion-checkpoints/resolve/main/config.gin
12.3 kB
| 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 | |