Team Ai
Apppublic

unity2009/MusicGenAI

sourceHugging Facecc-by-nc-4.0updated 3y agoView on Hugging Face
0likes
test_multibanddiffusion.py54 linesDownload Raw Back to models
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