softwareweaver/MusicGen
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 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 