unity2009/MusicGenAI
0
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the license found in the5# LICENSE file in the root directory of this source tree.6 7import random8 9import numpy as np10import torch11from audiocraft.models.multibanddiffusion import MultiBandDiffusion, DiffusionProcess12from audiocraft.models import EncodecModel, DiffusionUnet13from audiocraft.modules import SEANetEncoder, SEANetDecoder14from audiocraft.modules.diffusion_schedule import NoiseSchedule15from audiocraft.quantization import DummyQuantizer16 17 18class TestMBD:19 20 def _create_mbd(self,21 sample_rate: int,22 channels: int,23 n_filters: int = 3,24 n_residual_layers: int = 1,25 ratios: list = [5, 4, 3, 2],26 num_steps: int = 1000,27 codec_dim: int = 128,28 **kwargs):29 frame_rate = np.prod(ratios)30 encoder = SEANetEncoder(channels=channels, dimension=codec_dim, n_filters=n_filters,31 n_residual_layers=n_residual_layers, ratios=ratios)32 decoder = SEANetDecoder(channels=channels, dimension=codec_dim, n_filters=n_filters,33 n_residual_layers=n_residual_layers, ratios=ratios)34 quantizer = DummyQuantizer()35 compression_model = EncodecModel(encoder, decoder, quantizer, frame_rate=frame_rate,36 sample_rate=sample_rate, channels=channels, **kwargs)37 diffusion_model = DiffusionUnet(chin=channels, num_steps=num_steps, codec_dim=codec_dim)38 schedule = NoiseSchedule(device='cpu', num_steps=num_steps)39 DP = DiffusionProcess(model=diffusion_model, noise_schedule=schedule)40 mbd = MultiBandDiffusion(DPs=[DP], codec_model=compression_model)41 return mbd42 43 def test_model(self):44 random.seed(1234)45 sample_rate = 24_00046 channels = 147 codec_dim = 12848 mbd = self._create_mbd(sample_rate=sample_rate, channels=channels, codec_dim=codec_dim)49 for _ in range(10):50 length = random.randrange(1, 10_000)51 x = torch.randn(2, channels, length)52 res = mbd.regenerate(x, sample_rate)53 assert res.shape == x.shape54 