import os
import tqdm
import torch
import numpy as np
from einops import rearrange
import torch.distributed as dist
from torch.utils.data import DataLoader, DistributedSampler, Dataset
from torch.nn.parallel import DistributedDataParallel as DDP
from suno_utils.models.ditto.ditto import Ditto
from suno_utils.audio import Audio
from suno_utils.utils.text import read_jsonl

import sys

sys.path.append("/home/minz/glockenspiel/ditto-training")
from ditto.data_loaders.cleaned_ytmsd import CleanedYTMSDDataset


class CleanedYTMSDDataset(Dataset):
    def __init__(
        self,
        data_path="/app/suno/data/audio_mono_24khz/cleaned_ytmsd",
        split="train",
        input_length_s=30.0,
        sample_rate=24000,
        num_samples=-1,
    ):
        assert split in ["train", "valid"]
        self.data_path = data_path
        self.split = split
        self.input_length_s = input_length_s
        self.sample_rate = sample_rate
        self.num_samples = num_samples
        self.max_chunks = 8  # 30s * 8 chunks = 240 s

        # load files
        self.metadata = read_jsonl(os.path.join(data_path, "%s.jsonl" % split))
        print("%d files are available for %s set" % (len(self.metadata), split))

    def __getitem__(self, index):
        # read data
        metadata = self.metadata[index]

        # load audio
        audio_path = metadata["filepath"]
        audio = Audio.from_file(audio_path, sample_rate=self.sample_rate)

        # max 4 min
        wav = audio.get_segment(from_s=0.0, to_s=240).array_float
        num_chunks = len(wav) // self.input_length_s // self.sample_rate
        wav[int(num_chunks * self.input_length_s * self.sample_rate) :] = 0.0

        # append short
        if len(wav) < int(self.sample_rate * self.max_chunks * self.input_length_s):
            pad = int(self.sample_rate * self.max_chunks * self.input_length_s) - len(
                wav
            )
            wav = np.pad(wav, (0, pad), mode="constant", constant_values=0)

        if num_chunks == 0:
            audio_path = "invalid"

        # make stacks
        wav = wav.reshape(self.max_chunks, int(self.input_length_s * self.sample_rate))

        return wav, num_chunks, audio_path

    def __len__(self):
        if self.num_samples > 0:
            return self.num_samples
        else:
            return len(self.metadata)


# data parallel
dist.init_process_group(backend="nccl")
gpu_id = torch.distributed.get_rank()
torch.cuda.set_device(gpu_id)

# load model
ditto = Ditto(
    music_encoder_name="musicfm_concat",
    latent_dim=128,
    model_path="/home/minz/logs/ditto_cleaned_ytmsd_128/step_370k.pt",
    is_flash=False,
)
ditto = ditto.eval().bfloat16().cuda(gpu_id)
ditto = DDP(ditto, device_ids=[gpu_id])

# train data
train_ds = CleanedYTMSDDataset(split="train")
train_sampler = DistributedSampler(train_ds, num_replicas=4, rank=gpu_id, shuffle=False)
train_dl = DataLoader(
    dataset=train_ds,
    batch_size=4,
    shuffle=False,
    sampler=train_sampler,
    drop_last=False,
    num_workers=8,
)
train_memmap_fn = "/app/suno/minz/ytmsd_train.dat"
train_memmap_file_fn = "/app/suno/minz/ytmsd_train_files.dat"
train_embeddings = np.memmap(
    train_memmap_fn, dtype=np.float32, mode="w+", shape=(len(train_ds), 128)
)
train_filenames = np.memmap(
    train_memmap_file_fn, dtype=f"S{100}", mode="w+", shape=(len(train_ds),)
)

# for multi gpu
total_samples = len(train_ds)
samples_per_gpu = total_samples // dist.get_world_size()
offset = samples_per_gpu * dist.get_rank()

for batch_idx, (wav, counter, fn) in enumerate(tqdm.tqdm(train_dl)):
    b = len(wav)
    wav = rearrange(wav, "b c t -> (b c) t")
    wav = wav.cuda().bfloat16()
    emb = ditto.module.music_to_latent(wav).cpu().detach().numpy()
    emb = rearrange(emb, "(b c) d -> b c d", b=b, c=8)
    emb = emb.sum(axis=1) / counter.numpy()[:, np.newaxis]

    start_idx = offset + batch_idx * train_dl.batch_size
    end_idx = start_idx + b
    train_embeddings[start_idx:end_idx] = emb
    train_filenames[start_idx:end_idx] = list(fn)

    if batch_idx % 10 == 0:
        train_embeddings.flush()
        train_filenames.flush()
train_embeddings.flush()
train_filenames.flush()


# valid data
valid_ds = CleanedYTMSDDataset(split="valid")
valid_sampler = DistributedSampler(valid_ds, num_replicas=4, rank=gpu_id, shuffle=False)
valid_dl = DataLoader(
    dataset=valid_ds,
    batch_size=4,
    shuffle=False,
    sampler=valid_sampler,
    drop_last=False,
    num_workers=8,
)
valid_memmap_fn = "/app/suno/minz/ytmsd_valid.dat"
valid_memmap_file_fn = "/app/suno/minz/ytmsd_valid_files.dat"
valid_embeddings = np.memmap(
    valid_memmap_fn, dtype=np.float32, mode="w+", shape=(len(valid_ds), 128)
)
valid_filenames = np.memmap(
    valid_memmap_file_fn, dtype=f"S{100}", mode="w+", shape=(len(valid_ds),)
)

# for multi gpu
total_samples = len(valid_ds)
samples_per_gpu = total_samples // dist.get_world_size()
offset = samples_per_gpu * dist.get_rank()

for batch_idx, (wav, counter, fn) in enumerate(tqdm.tqdm(valid_dl)):
    b = len(wav)
    wav = rearrange(wav, "b c t -> (b c) t")
    wav = wav.cuda().bfloat16()
    emb = ditto.module.music_to_latent(wav).cpu().detach().numpy()
    emb = rearrange(emb, "(b c) d -> b c d", b=b, c=8)
    emb = emb.sum(axis=1) / counter.numpy()[:, np.newaxis]

    start_idx = offset + batch_idx * valid_dl.batch_size
    end_idx = start_idx + b
    valid_embeddings[start_idx:end_idx] = emb
    valid_filenames[start_idx:end_idx] = list(fn)

    if batch_idx % 10 == 0:
        valid_embeddings.flush()
        valid_filenames.flush()
valid_embeddings.flush()
valid_filenames.flush()
