import dataclasses import json from math import inf, ceil import random import time import numpy as np from tqdm import tqdm from api import ( ComposerAPI, SearchCtx, ComposerAPIBasePerfRecorder, GenerationRequest, ) from data import Vocab, midi2wavtool, VocabError from data_gen import ( stream_examples, augment_example, ) from train_target import train_target_wants_text, WITH_TEXT from data_aug import Augmenter, TextAugmenter, AugmentError import sys import torch from embedder import Embedder from loader import load_config from logitbias_json import json2logitbias from model import ModelArgs, MusicalPositionEmbedTransformer from dump_training_examples import clip_props, separate_drums from model_accelerated import MusicalPositionEmbedTransformerAccelerated import itertools def generate_eval_data(config, vocab): seq_len = int(config["seq_len"]) seq_len_min = int(config["seq_len_min"]) eval_split = int(config["eval_split_idx"]) auger = Augmenter.from_config(vocab, config["aug"]) text_auger = TextAugmenter(is_eval=False) for ( id, descs, stripe, example, accomps, example_tags, accomps_tagss, ) in stream_examples( split=eval_split, minlen=seq_len_min, maxlen=seq_len, order="random", aug_limit=10, train_target=WITH_TEXT, ): try: main, combined_accomp, _, text = augment_example( stripe, vocab, auger, text_auger, descs, example, None, accomps, seq_len_min, seq_len, example_tags, accomps_tagss, ) except (VocabError, AugmentError): continue yield main, combined_accomp, text def gen_critic_data(model_prefix): config = load_config(f"{model_prefix}/config.json") vocab = Vocab.from_config(config) encoder = Embedder() if train_target_wants_text(config["train_target"]) else None model = MusicalPositionEmbedTransformer( vocab, dataclasses.replace( ModelArgs.from_config(vocab, config), max_batch_size=2, max_seq_len=config["inference"]["max_decoder_seq_len"], cache=False, enable_flash=True, ), encoder, ) model.load_state_dict(torch.load(f"{model_prefix}/model.pt", map_location="cpu")) model.eval() model = model.cuda() model_accelerated = MusicalPositionEmbedTransformerAccelerated( model.params, f"{model_prefix}/model.trt", f"{model_prefix}/model_one_step.trt", model.encoder, ) api = ComposerAPI(vocab, model_accelerated) class_embeds = [] nonclass_embeds = [] class_examples = [] nonclass_examples = [] limit = 1000 for main, combined_accomp, text in tqdm( itertools.islice(generate_eval_data(config, vocab), limit), total=limit ): encoder_input_ids, encoder_attention_mask = encoder.tokenize([text]) if len(main) == 0: continue is_drums = main[0]["note"] >= 1000 main_len = min(max(n["offBeat"] for n in main), 39) main = [n for n in main if n["offBeat"] <= main_len] prefix_len = random.random() * main_len * 0.9 prefix = [n for n in main if n["onBeat"] < prefix_len] try: main_vec, _ = vocab.midi_to_tensor( main, accompany=combined_accomp, strict=True, ) prefix_vec, _ = vocab.midi_to_tensor( prefix, accompany=combined_accomp, strict=True, rest_end=True, ) except VocabError: continue main_vec = main_vec[:-1] # remove end symbol fake_max_len = min(int(main_vec.shape[0] * 1.5), api.max_seq_len) fake_ys = np.zeros((api.inference_batch_size, fake_max_len, 9), dtype=np.int32) fake_ys[0, : prefix_vec.shape[0]] = prefix_vec dists = np.zeros( (api.inference_batch_size, fake_max_len, vocab.N), dtype=np.float32 ) logit_masks = np.zeros((api.inference_batch_size, vocab.N), dtype=np.float32) req = GenerationRequest(0.3, f"error || beat >= {main_len}") bias = json2logitbias( {("pitch < 1000" if is_drums else "pitch >= 1000"): -1000.0} ) search_ctx = SearchCtx( request_ctxts=[ api.setup_request_context( 0, req, 0.0, 0.0, bias, fake_ys, prefix_vec.shape[0] ) ], encoder_input_ids=encoder_input_ids, encoder_attention_mask=encoder_attention_mask, start=prefix_vec.numpy(), logit_masks=logit_masks, ys=fake_ys, dists=dists, pos=prefix_vec.shape[0], recorder=ComposerAPIBasePerfRecorder(), start_time=time.time(), deadline=inf, ) api.generate_ctx(search_ctx) # convert to midi to cut in the same way as the real example fake_full_clip = vocab.tensor_to_midi_using_embeds( fake_ys[0][: search_ctx.request_ctxts[0].stop_pos] )[-1] fake_bounded_clip = [n for n in fake_full_clip if n["offBeat"] <= main_len] try: fake_vec, _ = vocab.midi_to_tensor( fake_bounded_clip, accompany=combined_accomp, strict=True, ) fake_vec = fake_vec[:-1] # remove end symbol except VocabError: continue inp = torch.zeros( (2, max(main_vec.shape[0], fake_vec.shape[0]), 9), dtype=torch.long ) if inp.shape[1] > model.params.max_seq_len: continue inp[:, :, 0] = vocab.pad.index inp[0, : main_vec.shape[0]] = main_vec inp[1, : fake_vec.shape[0]] = fake_vec emb_real, emb_fake = model.forward( inp, return_embeddings=True, encoder_input_ids=encoder_input_ids, encoder_attention_mask=encoder_attention_mask, ).unbind(0) if torch.allclose(emb_real, emb_fake): continue class_embeds.append([emb_real]) nonclass_embeds.append([emb_fake]) class_examples.append(main_vec) nonclass_examples.append(fake_vec) # print("real:") # print(vocab.repr_tensor(main_vec)) # print("fake:") # print(vocab.repr_tensor(fake_vec)) # print("---") torch.save( { "class_embeds": class_embeds, "nonclass_embeds": nonclass_embeds, }, "data.pt", ) dump_examples(vocab, class_examples, nonclass_examples) def dump_examples(vocab, class_examples, nonclass_examples): clipss = [[] for i in range(4)] beat_accum = 0 for i, c_nc in enumerate(zip(class_examples, nonclass_examples)): clips = [vocab.tensor_to_midi_using_embeds(c)[-1] for c in c_nc] example_beats = ceil(max(n["offBeat"] for clip in clips for n in clip)) separated_clips = [] for clip in clips: separated_clips.extend(list(separate_drums(clip))) wavtool_clips = [ dict( midi2wavtool(c), **clip_props, loopEnd=example_beats, timelineStart=beat_accum, timelineEnd=beat_accum + example_beats, name=f"{n}", ) for c, n in zip( separated_clips, [ "TH |", "TD |", "FH |", "FD |", ], ) ] if all(len(clip["notes"]) == 0 for clip in wavtool_clips): continue beat_accum += example_beats for i, clip in enumerate(wavtool_clips): if len(clip["notes"]) > 0: clipss[i].append(clip) obj = { "content": [{"clips": cs, "automationPoints": []} for cs in clipss], "length": beat_accum, } from pprint import pprint with open("out.json", "w") as f: json.dump(obj, f) if __name__ == "__main__": from util import configure_logging configure_logging() with torch.no_grad(): gen_critic_data(sys.argv[1])