CoolFace
Apppublic

softwareweaver/MusicGen

sourceHugging Facecc-by-nc-4.0updated 11mo agoView on Hugging Face
0likes
test_codebooks_patterns.py247 linesDownload Raw Back to modules
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 pytest8import torch9 10from audiocraft.modules.codebooks_patterns import (11    DelayedPatternProvider,12    ParallelPatternProvider,13    Pattern,14    UnrolledPatternProvider,15)16 17 18class TestParallelPatternProvider:19 20    @pytest.mark.parametrize("n_q", [1, 4, 32])21    @pytest.mark.parametrize("timesteps", [0, 1, 16, 100])22    def test_get_pattern(self, n_q: int, timesteps: int):23        provider = ParallelPatternProvider(n_q)24        pattern = provider.get_pattern(timesteps)25        # + 1 to account for 1st step26        assert len(pattern.layout) == timesteps + 127 28    @pytest.mark.parametrize("n_q", [1, 4, 32])29    @pytest.mark.parametrize("timesteps", [8, 16, 100])30    def test_pattern_content(self, n_q: int, timesteps: int):31        provider = ParallelPatternProvider(n_q)32        pattern = provider.get_pattern(timesteps)33        for s, v in enumerate(pattern.layout):34            for i, code in enumerate(v):35                assert i == code.q36                assert code.t == s - 1  # account for the 1st empty step37 38    @pytest.mark.parametrize("n_q", [1, 4, 32])39    @pytest.mark.parametrize("timesteps", [8, 16, 100])40    def test_pattern_max_delay(self, n_q: int, timesteps: int):41        provider = ParallelPatternProvider(n_q)42        pattern = provider.get_pattern(timesteps)43        assert pattern.max_delay == 044        assert len(pattern.valid_layout) == len(pattern.layout) - pattern.max_delay45 46 47class TestDelayedPatternProvider:48 49    @pytest.mark.parametrize("n_q", [1, 4, 32])50    @pytest.mark.parametrize("timesteps", [0, 1, 16, 100])51    def test_get_pattern(self, n_q: int, timesteps: int):52        delays = [53            list(range(n_q)),54            [0] + [1] * (n_q - 1),55            [0] + [4] * (n_q - 1),56        ]57        for delay in delays:58            provider = DelayedPatternProvider(n_q, delay)59            pattern = provider.get_pattern(timesteps)60            # + 1 to account for 1st step61            assert len(pattern.layout) == timesteps + max(delay) + 162 63    @pytest.mark.parametrize("n_q", [1, 4, 32])64    @pytest.mark.parametrize("timesteps", [8, 16, 100])65    def test_pattern_content(self, n_q: int, timesteps: int):66        provider = DelayedPatternProvider(n_q)67        pattern = provider.get_pattern(timesteps)68        for s, v in enumerate(pattern.layout):69            for i, code in enumerate(v):70                assert i == code.q71                assert code.t == max(0, s - code.q - 1)72 73    @pytest.mark.parametrize("timesteps", [8, 16, 100])74    @pytest.mark.parametrize("delay", [[0, 1, 2, 3], [0, 1, 1, 1], [0, 3, 3, 3], [0, 3]])75    def test_pattern_max_delay(self, timesteps: int, delay: list):76        provider = DelayedPatternProvider(len(delay), delay)77        pattern = provider.get_pattern(timesteps)78        assert pattern.max_delay == max(delay)79        assert len(pattern.valid_layout) == len(pattern.layout) - pattern.max_delay80 81 82class TestUnrolledPatternProvider:83 84    @pytest.mark.parametrize("timesteps", [0, 1, 16])85    @pytest.mark.parametrize("flattening", [[0, 1, 2], [0, 1, 1]])86    @pytest.mark.parametrize("delays", [[0, 0, 0], [0, 5, 5]])87    def test_get_pattern(self, timesteps: int, flattening: list, delays: list):88        n_q = len(flattening)89        max_delay = max(delays)90        provider = UnrolledPatternProvider(n_q, flattening, delays)91        pattern = provider.get_pattern(timesteps)92        assert len(pattern.layout) == provider.num_virtual_steps(timesteps) + max_delay93 94    @pytest.mark.parametrize("timesteps", [0, 1, 16])95    @pytest.mark.parametrize("flattening", [[0, 1, 2], [0, 1, 1]])96    @pytest.mark.parametrize("delays", [[0, 0, 0], [0, 5, 5]])97    def test_pattern_max_delay(self, timesteps: int, flattening: list, delays: list):98        n_q = len(flattening)99        max_delay = max(delays)100        provider = UnrolledPatternProvider(n_q, flattening, delays)101        pattern = provider.get_pattern(timesteps)102        assert pattern.max_delay == max_delay103 104 105class TestPattern:106 107    def ref_build_pattern_sequence(self, z: torch.Tensor, pattern: Pattern, special_token: int):108        """Reference method to build the sequence from the pattern without using fancy scatter."""109        bs, n_q, T = z.shape110        z = z.cpu().numpy()111        assert n_q == pattern.n_q112        assert T <= pattern.timesteps113        inp = torch.full((bs, n_q, len(pattern.layout)), special_token, dtype=torch.long).numpy()114        inp[:] = special_token115        for s, v in enumerate(pattern.layout):116            for (t, q) in v:117                if t < T:118                    inp[:, q, s] = z[:, q, t]119        return torch.from_numpy(inp)120 121    def ref_revert_pattern_sequence(self, z: torch.Tensor, pattern: Pattern, special_token: int):122        """Reference method to revert the sequence from the pattern without using fancy scatter."""123        z = z.cpu().numpy()124        bs, n_q, S = z.shape125        assert pattern.n_q == n_q126        inp = torch.full((bs, pattern.n_q, pattern.timesteps), special_token, dtype=torch.long).numpy()127        inp[:] = special_token128        for s, v in enumerate(pattern.layout):129            for (t, q) in v:130                if t < pattern.timesteps:131                    inp[:, q, t] = z[:, q, s]132        return torch.from_numpy(inp)133 134    def ref_revert_pattern_logits(self, z: torch.Tensor, pattern: Pattern, special_token: float):135        """Reference method to revert the logits from the pattern without using fancy scatter."""136        z = z.cpu().numpy()137        bs, card, n_q, S = z.shape138        assert pattern.n_q == n_q139        ref_layout = pattern.layout140        inp = torch.full((bs, card, pattern.n_q, pattern.timesteps), special_token, dtype=torch.float).numpy()141        inp[:] = special_token142        for s, v in enumerate(ref_layout[1:]):143            if s < S:144                for (t, q) in v:145                    if t < pattern.timesteps:146                        inp[:, :, q, t] = z[:, :, q, s]147        return torch.from_numpy(inp)148 149    def _get_pattern_providers(self, n_q: int):150        pattern_provider_1 = ParallelPatternProvider(n_q)151        pattern_provider_2 = DelayedPatternProvider(n_q, list(range(n_q)))152        pattern_provider_3 = DelayedPatternProvider(n_q, [0] + [1] * (n_q - 1))153        pattern_provider_4 = UnrolledPatternProvider(154            n_q, flattening=list(range(n_q)), delays=[0] * n_q155        )156        pattern_provider_5 = UnrolledPatternProvider(157            n_q, flattening=[0] + [1] * (n_q - 1), delays=[0] * n_q158        )159        pattern_provider_6 = UnrolledPatternProvider(160            n_q, flattening=[0] + [1] * (n_q - 1), delays=[0] + [5] * (n_q - 1)161        )162        return [163            pattern_provider_1,164            pattern_provider_2,165            pattern_provider_3,166            pattern_provider_4,167            pattern_provider_5,168            pattern_provider_6,169        ]170 171    @pytest.mark.parametrize("n_q", [1, 4, 32])172    @pytest.mark.parametrize("timesteps", [16, 72])173    def test_build_pattern_sequence(self, n_q: int, timesteps: int):174        bs = 2175        card = 256176        special_token = card177 178        pattern_providers = self._get_pattern_providers(n_q)179        for pattern_provider in pattern_providers:180            pattern = pattern_provider.get_pattern(timesteps)181            # we can correctly build the sequence from the pattern182            z = torch.randint(0, card, (bs, n_q, timesteps))183            ref_res = self.ref_build_pattern_sequence(z, pattern, special_token)184            res, indexes, mask = pattern.build_pattern_sequence(z, special_token)185            assert (res == ref_res).float().mean() == 1.0186 187            # expected assertion fails on the number of timesteps188            invalid_timesteps = [timesteps + 1]189            if pattern.num_sequence_steps != pattern.timesteps:190                invalid_timesteps.append(pattern.num_sequence_steps)191            for i_timesteps in invalid_timesteps:192                z2 = torch.randint(0, card, (bs, n_q, i_timesteps))193                with pytest.raises(AssertionError):194                    pattern.build_pattern_sequence(z2, special_token)195 196            # expected assertion fails on the number of codebooks197            invalid_qs = [0, n_q - 1, n_q + 1]198            for i_q in invalid_qs:199                z3 = torch.randint(0, card, (bs, i_q, timesteps))200                with pytest.raises(AssertionError):201                    pattern.build_pattern_sequence(z3, special_token)202 203    @pytest.mark.parametrize("n_q", [1, 4, 32])204    @pytest.mark.parametrize("timesteps", [16, 72])205    def test_revert_pattern_sequence(self, n_q: int, timesteps: int):206        bs = 2207        card = 256208        special_token = card209 210        pattern_providers = self._get_pattern_providers(n_q)211        for pattern_provider in pattern_providers:212            pattern = pattern_provider.get_pattern(timesteps)213            # this works assuming previous tests are successful214            z = torch.randint(0, card, (bs, n_q, timesteps))215            s = self.ref_build_pattern_sequence(z, pattern, special_token)216            ref_out = self.ref_revert_pattern_sequence(s, pattern, special_token)217            # ensure our reference script retrieve the original sequence218            assert z.shape == ref_out.shape219            assert (z == ref_out).float().mean() == 1.0220            # now we can test the scatter version221            out, indexes, mask = pattern.revert_pattern_sequence(s, special_token)222            assert out.shape == ref_out.shape223            assert (out == ref_out).float().mean() == 1.0224 225    @pytest.mark.parametrize("n_q", [1, 4, 32])226    @pytest.mark.parametrize("timesteps", [16, 72])227    @pytest.mark.parametrize("card", [1, 2, 256, 1024])228    def test_revert_pattern_logits(self, n_q: int, timesteps: int, card: int):229        bs = 2230        special_token = card231        logits_special_token = float('nan')232 233        pattern_providers = self._get_pattern_providers(n_q)234        for pattern_provider in pattern_providers:235            pattern = pattern_provider.get_pattern(timesteps)236            # this works assuming previous tests are successful237            z = torch.randint(0, card, (bs, n_q, timesteps))238            s = self.ref_build_pattern_sequence(z, pattern, special_token)239            logits = torch.randn((bs, card, n_q, s.shape[-1]))240            ref_out = self.ref_revert_pattern_logits(logits, pattern, logits_special_token)241            # ensure our reference script retrieve the original sequence242            assert ref_out.shape == torch.Size([bs, card, n_q, timesteps])243            # now we can test the scatter version244            out, indexes, mask = pattern.revert_pattern_logits(logits, logits_special_token)245            assert out.shape == ref_out.shape246            assert (out == ref_out).float().mean() == 1.0247