import json import torch from tqdm import tqdm from data import Vocab, midi2wavtool from data_gen import DataGenerator, Dataset, Augmenter from embedder import Embedder import data_aug import dataset_classes import sys clip_props = { "type": "MIDI", "color": "#f65943", "fadeIn": 0, "fadeOut": 0, "readStart": 0, "loopStart": 0, "lifted": False, "ccs": {}, } # clip -> harmonic clip, drum clip def separate_drums(clip): return ( [n for n in clip if n["note"] < 1000], [dict(n, note=n["note"] - 1012) for n in clip if n["note"] >= 1000], ) class Dumper: def __init__(self, config): self.clipss = [[] for i in range(4)] self.config = config self.vocab = Vocab.from_config(config) self.batch_size = int(config["batch_size"]) self.seq_len = int(config["seq_len"]) self.seq_len_min = int(config["seq_len_min"]) self.seq_len_max = int(config["seq_len_max"]) self.train_target = config["train_target"] self.example_beats = ( self.vocab.embed_length_max // self.vocab.quantize_divisions ) self.num_batches = 0 self.embedder = Embedder(load_pretrained_weights=False) self.dataset = Dataset.dynamic_from_config(config, "dataset_class") self.auger = Augmenter.dynamic_from_config( self.vocab, config, "augmenter_class" ) print(self.dataset) print(self.auger) train_split_idx = int(config["train_split_idx"]) self.data_gen = DataGenerator( self.dataset, self.vocab, split=train_split_idx, ranksize=(0, 1), batch_size=self.batch_size, seq_len=self.seq_len, seq_len_min=self.seq_len_min, seq_len_max=self.seq_len_max, train_target=self.train_target, parallelism=1, to_device="cpu", text_tokenize=self.embedder.tokenize, example_continuous=False, pack_batch=False, augmenter=self.auger, ) def dump(self, batch): for i, id, t in zip(range(batch.tgt.shape[0]), batch.ids, batch.texts): tgt_cpu = ( batch.tgt[i].cpu() if isinstance(batch.tgt, torch.Tensor) else batch.tgt[i] ) symlen = torch.count_nonzero(tgt_cpu[:, 0] != self.vocab.pad.index) examples_out = self.num_batches * self.batch_size + i print("---") print(self.vocab.repr_tensor(tgt_cpu[:symlen])) clips = self.vocab.tensor_to_midi_using_embeds(tgt_cpu) separated_clips = [] for clip in reversed(clips): separated_clips.extend(list(separate_drums(clip))) if batch.encoder_input_ids is not None: tn = self.embedder.detokenize(batch.encoder_input_ids[i].unsqueeze(0))[ 0 ] else: tn = "no input ids" print(f'detokenized prompt: "{tn}"') wavtool_clips = [] for c, n in zip( separated_clips, [ f"MH | {symlen}", "MD |", "AH |", "AD |", ], ): clipsz = midi2wavtool(c, split_on_desc_anchor=True) clip_start = self.example_beats * examples_out clip_end = clip_start + ( self.example_beats if len(clipsz) == 1 else clipsz[1]["notes"][0]["start"] ) clip_internal_offset = 0 subclips = [] for clip in clipsz: subclips.append( dict( clip_props, notes=[ dict( n, start=n["start"] - clip_internal_offset, end=n["end"] - clip_internal_offset, ) for n in clip["notes"] ], loopEnd=clip_end - clip_start, timelineStart=clip_start, timelineEnd=clip_end, name=f"{n} {tn}", embed_len=self.embedder.tokenize([t])[0].shape[1], ) ) clip_internal_offset += clip_end - clip_start clip_start = clip_end clip_end = self.example_beats * (examples_out + 1) wavtool_clips.append(subclips) if all( len(clip["notes"]) == 0 for subclips in wavtool_clips for clip in subclips ): continue for i, subclips in enumerate(wavtool_clips): for clip in subclips: if len(clip["notes"]) > 0: self.clipss[i].append(clip) self.num_batches += 1 def write(self, path): obj = { "content": [{"clips": cs, "automationPoints": []} for cs in self.clipss], "length": self.example_beats * self.num_batches * self.batch_size, } with open(path, "w") as f: json.dump(obj, f) def dump_training_examples(self): for batch in self.data_gen.generate(force_batches=10, order="random"): self.dump(batch) self.write("out.json") if __name__ == "__main__": from util import configure_logging configure_logging() with open(sys.argv[1], "r") as f: config = json.load(f) d = Dumper(config) d.dump_training_examples()