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()
