from data_gen import DataGenerator, Dataset, RawExample, TrivialAugmenter from data import wavtool2midi, Vocab import json import torch from data_tests import vocab_test_inputs import pytest from embedder import Embedder class MockDataset(Dataset): def __init__(self, split_elems): super().__init__() self.split_elems = split_elems self.shuffle_called = False self.requested_ranksize = None self.requested_order = None def num_examples(self): return {split: len(elems) for split, elems in self.split_elems.items()} def stream_examples_impl(self, split, ranksize, order): self.requested_ranksize = ranksize self.requested_order = order elems = self.split_elems[split] for elem in elems: yield elem def shuffle(self): self.shuffle_called = True def make_mock_dataset(): def make_split(name): return [ RawExample( id=i, descs=[f"split {name} example {i}"], example=wavtool2midi( json.loads(v[0][0] if isinstance(v[0], list) else v[0]) ), accomps=([] if v[1] is None else [wavtool2midi(json.loads(v[1]))]), ) for i, v in enumerate(vocab_test_inputs) ] return MockDataset( { 0: make_split("train"), 1: make_split("eval"), } ) @pytest.mark.parametrize("batch_size", [1, 4, 64]) @pytest.mark.parametrize("forcebatches", [None, 2, 20]) @pytest.mark.parametrize("parallelism", [1, 4]) def test_data_generator_batches(batch_size, forcebatches, parallelism): really_test_data_generator( batch_size, 0, forcebatches, parallelism, "random", (0, 1) ) def test_data_generator_eval_split(): really_test_data_generator(1, 1, None, 1, "random", (0, 1)) def test_data_generator_id_order(): really_test_data_generator(1, 0, None, 1, "id", (0, 1)) def test_data_generator_ranksize(): really_test_data_generator(1, 0, None, 1, "random", (3, 4)) @pytest.mark.parametrize("seq_len", [3, 7, 64]) @pytest.mark.parametrize("pack_batch", [False, True]) def test_data_generator_example_continuous(seq_len, pack_batch): dataset = make_mock_dataset() vocab = Vocab(24, 96, 24, 48, 20, 127, 2, 1, 24 * 4, 4, 4 * 4 * 24, 4, 24) data_gen = DataGenerator( dataset, vocab, 0, (0, 1), 1, seq_len, 1, 256, "tgt", 1, "cpu", None, True, pack_batch, TrivialAugmenter(), ) refs = [] for raw_example in dataset.split_elems[0]: refs.append( vocab.midi_to_tensor( raw_example.example, accompany=(raw_example.accomps[0] if raw_example.accomps else None), strict=True, ) ) frags = [] for batch in data_gen.generate(order="random"): s = batch.tgt[0] frags.append(s[s[:, 0] != vocab.pad.index]) reconstructed = torch.cat(frags, dim=0) begins = torch.nonzero(reconstructed[:, 0] == vocab.begin.index).view(-1).tolist() begins.append(reconstructed.shape[0]) for begin, end in zip(begins, begins[1:]): seq = reconstructed[begin:end] for i, ref in enumerate(refs): if torch.equal(seq, ref) or ( ref[-1, 0] == vocab.end.index and torch.equal(seq, ref[:-1]) ): del refs[i] break else: assert False, f"unexpected example in batch: {seq}" assert len(refs) == 0, "missing examples" def really_test_data_generator( batch_size, split, forcebatches, parallelism, order, ranksize ): dataset = make_mock_dataset() vocab = Vocab(24, 96, 24, 48, 20, 127, 2, 1, 24 * 4, 4, 4 * 4 * 24, 4, 24) encoder = Embedder(load_pretrained_weights=False) data_gen = DataGenerator( dataset, vocab, split, ranksize, batch_size, 128, 1, None, "tgt", parallelism, "cpu", encoder.tokenize, False, False, TrivialAugmenter(), ) assert not dataset.shuffle_called assert dataset.requested_ranksize is None for i in range(2): data_gen.shuffle() assert dataset.shuffle_called refs = [] for raw_example in dataset.split_elems[split]: refs.append( ( vocab.midi_to_tensor( raw_example.example, accompany=( raw_example.accomps[0] if raw_example.accomps else None ), strict=True, ), raw_example.descs, ) ) batch_count = 0 for batch in data_gen.generate(force_batches=forcebatches, order=order): batch_count += 1 is_empty_batch = True for i in range(batch.tgt.shape[0]): tgt_valid = batch.tgt[i][batch.tgt[i][:, 0] != vocab.pad.index] if tgt_valid.shape[0] == 0: continue is_empty_batch = False for j, (ref_tensor, ref_descs) in enumerate(refs): if ( torch.equal(tgt_valid, ref_tensor) and ref_descs[0] == batch.texts[i] ): # verify encoding enc_toks = batch.encoder_input_ids[ i, batch.encoder_attention_mask[i] > 0 ] enc_str = encoder.detokenize(enc_toks.view(1, -1))[0] assert enc_str == batch.texts[i] del refs[j] break else: assert False, f"unexpected example in batch: {tgt_valid}" assert forcebatches is not None or not is_empty_batch, "unexpected empty batch" if forcebatches is None or forcebatches * batch_size >= len( dataset.split_elems[split] ): assert len(refs) == 0, f"missing examples from batches: {refs}" if forcebatches is not None: assert batch_count == forcebatches assert dataset.requested_order == order assert dataset.requested_ranksize == ranksize