import os
import time
import math
import glob
import wandb
import random
import torch
import torchaudio
import itertools
import numpy as np
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
from torch.utils.data import DataLoader
from tqdm import tqdm
from typing import List


def mdct(x):
    N = x.shape[-1]
    n = torch.arange(N, device=x.device)
    k = torch.arange(N // 2, device=x.device)

    arg = (
        (math.pi / (2 * N))
        * ((2 * n + 1 + N // 2).view(-1, 1))
        * ((2 * k + 1).view(1, -1))
    )
    mdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5

    return torch.matmul(x, mdct_matrix)


def imdct(X):
    half_N = X.shape[-1]
    N = half_N * 2
    n = torch.arange(N, device=X.device)
    k = torch.arange(half_N, device=X.device)

    arg = (
        (math.pi / (2 * N))
        * ((2 * n + 1 + N // 2).view(-1, 1))
        * ((2 * k + 1).view(1, -1))
    )
    imdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5

    return torch.matmul(X, imdct_matrix.T) * 2.0


def audio_to_mdct_frames(audio, frame_size=1920, midside=False):
    batch, channels, samples = audio.shape
    hop_size = frame_size // 2

    if midside:
        mid = (audio[:, 0, :] + audio[:, 1, :]) / 2.0
        side = (audio[:, 0, :] - audio[:, 1, :]) / 2.0
        audio = torch.stack([mid, side], dim=1)

    n_frames = (samples - frame_size) // hop_size + 1

    frames = []
    for i in range(n_frames):
        start = i * hop_size
        frame = audio[:, :, start : start + frame_size]
        if frame.shape[-1] == frame_size:
            frames.append(frame)

    frames = torch.stack(frames, dim=2)

    window = torch.sin(
        torch.pi / frame_size * (torch.arange(frame_size, device=audio.device) + 0.5)
    )
    frames = frames * window.view(1, 1, 1, -1)

    shape = frames.shape
    frames_reshaped = frames.reshape(-1, frame_size)
    mdct_coeffs = mdct(frames_reshaped)

    return mdct_coeffs.reshape(shape[0], shape[1], shape[2], -1)


def mdct_frames_to_audio_slow(mdct_coeffs, frame_size=1920, midside=False):
    batch, channels, n_frames, half_frame_size = mdct_coeffs.shape
    hop_size = frame_size // 2

    shape = mdct_coeffs.shape
    coeffs_reshaped = mdct_coeffs.reshape(-1, half_frame_size)
    frames = imdct(coeffs_reshaped)
    frames = frames.reshape(shape[0], shape[1], shape[2], -1)

    window = torch.sin(
        torch.pi
        / frame_size
        * (torch.arange(frame_size, device=mdct_coeffs.device) + 0.5)
    )
    frames = frames * window.view(1, 1, 1, -1)

    total_samples = (n_frames - 1) * hop_size + frame_size
    output = torch.zeros(batch, channels, total_samples, device=mdct_coeffs.device)

    for i in range(n_frames):
        start = i * hop_size
        output[:, :, start : start + frame_size] += frames[:, :, i]

    if midside:
        mid = output[:, 0, :]
        side = output[:, 1, :]
        left = mid + side
        right = mid - side
        output = torch.stack([left, right], dim=1)

    return output


def mdct_frames_to_audio(mdct_coeffs, frame_size=1920, midside=False):
    batch, channels, n_frames, half_frame_size = mdct_coeffs.shape
    hop_size = frame_size // 2

    # Reshape and apply IMDCT
    shape = mdct_coeffs.shape
    coeffs_reshaped = mdct_coeffs.reshape(-1, half_frame_size)
    frames = imdct(coeffs_reshaped)
    frames = frames.reshape(shape[0], shape[1], shape[2], -1)

    # Apply window
    window = torch.sin(
        torch.pi
        / frame_size
        * (torch.arange(frame_size, device=mdct_coeffs.device) + 0.5)
    )
    frames = frames * window.view(1, 1, 1, -1)

    # Calculate total samples and create output shape
    total_samples = (n_frames - 1) * hop_size + frame_size

    # Reshape frames to prepare for folding
    frames = frames.permute(0, 1, 3, 2)  # [batch, channels, frame_size, n_frames]
    frames = frames.reshape(batch * channels, frame_size, n_frames)

    # Use fold operation to overlap-add frames
    output = torch.nn.functional.fold(
        frames,
        output_size=(1, total_samples),
        kernel_size=(1, frame_size),
        stride=(1, hop_size),
    )

    # Reshape output to expected dimensions
    output = output.view(batch, channels, total_samples)

    if midside:
        mid = output[:, 0, :]
        side = output[:, 1, :]
        left = mid + side
        right = mid - side
        output = torch.stack([left, right], dim=1)

    return output


class WaveformMDCTVAE(nn.Module):
    def __init__(
        self,
        frame_size: int = 1920,
        latent_dim: int = 256,
        hidden_dims: list = None,
        dropout: float = 0.1,
        midside: bool = False,
    ):
        super().__init__()

        self.frame_size = frame_size
        self.n_coeffs = frame_size // 2
        self.latent_dim = latent_dim
        self.midside = midside

        if hidden_dims is None:
            hidden_dims = [512, 256]

        # Encoder layers
        modules = []
        input_dim = 2 * self.n_coeffs

        for h_dim in hidden_dims:
            modules.append(
                nn.Sequential(
                    nn.Linear(input_dim, h_dim),
                    # nn.LayerNorm(h_dim),
                    nn.LeakyReLU(),
                    nn.Dropout(dropout),
                    nn.Linear(h_dim, h_dim),
                )
            )
            input_dim = h_dim

        self.encoder = nn.Sequential(*modules)
        self.fc_mu = nn.Linear(hidden_dims[-1], latent_dim)
        self.fc_var = nn.Linear(hidden_dims[-1], latent_dim)

        # Decoder layers
        modules = []
        hidden_dims.reverse()

        self.decoder_input = nn.Sequential(
            nn.Linear(latent_dim, hidden_dims[0]),
            nn.LayerNorm(hidden_dims[0]),
            nn.LeakyReLU(),
            nn.Dropout(dropout),
        )

        for i in range(len(hidden_dims) - 1):
            modules.append(
                nn.Sequential(
                    nn.Linear(hidden_dims[i], hidden_dims[i + 1]),
                    # nn.LayerNorm(hidden_dims[i + 1]),
                    nn.LeakyReLU(),
                    nn.Dropout(dropout),
                    nn.Linear(hidden_dims[i + 1], hidden_dims[i + 1]),
                )
            )

        self.decoder = nn.Sequential(*modules)
        self.final_layer = nn.Linear(hidden_dims[-1], 2 * self.n_coeffs)

    def _encode(self, mdct_frames: torch.Tensor) -> list[torch.Tensor]:
        batch_size, _, n_frames, _ = mdct_frames.shape
        ch1 = mdct_frames[:, 0, :, :]
        ch2 = mdct_frames[:, 1, :, :]

        x = torch.cat((ch1, ch2), dim=-1)
        result = self.encoder(x)
        mu = self.fc_mu(result)
        log_var = self.fc_var(result)

        return [mu, log_var]

    def _decode(self, z: torch.Tensor) -> torch.Tensor:
        batch_size, n_frames, _ = z.shape

        result = self.decoder_input(z)
        result = self.decoder(result)
        result = self.final_layer(result)
        ch1 = result[..., : self.n_coeffs]
        ch2 = result[..., self.n_coeffs :]
        result = torch.stack((ch1, ch2), dim=1)

        return result

    def reparameterize(self, mu: torch.Tensor, log_var: torch.Tensor) -> torch.Tensor:
        if self.training:
            std = torch.exp(0.5 * log_var)
            eps = torch.randn_like(std)
            return eps * std + mu
        else:
            return mu

    def forward(
        self, waveform: torch.Tensor
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Forward pass handling both waveform conversion and VAE operations.

        Args:
            waveform: Input audio tensor of shape (batch_size, channels, samples)

        Returns:
            tuple containing:
            - reconstructed waveform
            - reconstructed MDCT coefficients
            - original MDCT coefficients (for loss computation)
            - mu
            - log_var
        """
        # Convert input waveform to MDCT frames
        mdct_frames = audio_to_mdct_frames(
            waveform, frame_size=self.frame_size, midside=self.midside
        )  # (batch, channels, n_frames, n_coeffs)

        # Encode and decode
        mu, log_var = self._encode(mdct_frames)
        z = self.reparameterize(mu, log_var)
        mdct_recon = self._decode(z)

        # Convert back to waveform
        waveform_recon = mdct_frames_to_audio(
            mdct_recon, frame_size=self.frame_size, midside=self.midside
        )

        return waveform_recon, mdct_recon, mdct_frames, mu, log_var


def vae_loss(
    mdct_recon: torch.Tensor,
    mdct_original: torch.Tensor,
    waveform_recon: torch.Tensor,
    waveform_original: torch.Tensor,
    mu: torch.Tensor,
    log_var: torch.Tensor,
    kld_weight: float = 0.0,
) -> dict:
    """
    Separate loss function that can be called after forward pass.
    """
    # Reconstruction loss in MDCT domain
    # print(mdct_recon[0:2, 0, 0, :10], mdct_original[0:2, 0, 0, :10])
    recons_loss = F.mse_loss(mdct_recon, mdct_original)

    # recons_loss = F.mse_loss(waveform_recon, waveform_original)

    # KL divergence loss
    kld_loss = torch.mean(-0.5 * torch.sum(1 + log_var - mu**2 - log_var.exp(), dim=2))
    kld_loss = torch.mean(kld_loss)

    # Total loss
    loss = recons_loss  # + kld_weight * kld_loss

    return {"loss": loss, "reconstruction_loss": recons_loss, "kld_loss": kld_loss}


class BufferedAudioDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        filepaths: List[str],
        sample_rate: int,
        num_workers: int = 1,
        chunk_size_s: float = 10.0,
        buffer_size: int = 50_000,
    ):
        self.filepaths = filepaths
        self.sample_rate = sample_rate
        self.chunk_size_s = chunk_size_s
        self.buffer_size = buffer_size
        self.chunk_size_samples = int(chunk_size_s * sample_rate)
        self.num_workers = num_workers
        self.items_since_last_reload = buffer_size  # force a reload
        self.buffer = []

    def __len__(self):
        return self.buffer_size * self.num_workers

    def _reload_buffer(self):
        self.buffer = []
        rand_idxs = torch.randperm(len(self.filepaths))

        print("Reloading buffer...")
        # max rand_idxs repeat endlessly
        rand_idxs = itertools.cycle(rand_idxs)
        pbar = tqdm(rand_idxs, total=len(self.filepaths), desc="Loading audio buffer")
        for idx in pbar:
            if len(self.buffer) >= self.buffer_size:
                break

            try:
                filepath = self.filepaths[idx]
                audio, sr = torchaudio.load(filepath)

                if sr != self.sample_rate:
                    audio = torchaudio.functional.resample(audio, sr, self.sample_rate)

                # Pad if needed to ensure consistent chunk size
                if audio.shape[-1] < self.chunk_size_samples:
                    continue

                # Split into chunks
                chunks = audio.unfold(
                    -1, self.chunk_size_samples, self.chunk_size_samples
                )
                chunks = chunks.chunk(chunks.shape[1], dim=1)

                # Filter chunks by minimum length
                valid_chunks = [
                    chunk.squeeze(1)
                    for chunk in chunks
                    if chunk.shape[-1] >= self.chunk_size_samples
                ]

                # filter out chunks of silence
                valid_chunks = [
                    chunk for chunk in valid_chunks if (chunk.abs() ** 2).mean() > 0.001
                ]

                self.buffer.extend(valid_chunks)

                pbar.set_postfix({"buffer_size": len(self.buffer)})

            except Exception as e:
                print(f"Error loading {filepath}: {e}")
                continue
        self.items_since_last_reload = 0

    def __getitem__(self, _):
        if self.items_since_last_reload >= len(self.buffer):
            self._reload_buffer()

        # get a random preset and apply it to the audio
        buffer_idx = np.random.randint(0, len(self.buffer))
        audio = self.buffer[buffer_idx]

        # ensure nothing is out of range
        if audio.abs().max() > 1.0:
            audio = audio / audio.abs().max()

        # apply random gain reduction
        # if np.random.uniform() < 0.5:
        #    gain_reduction_db = np.random.uniform(-10, 0)
        #    audio *= 10 ** (gain_reduction_db / 20.0)

        # self.items_since_last_reload += 1

        return audio


def save_checkpoint(
    model,
    optimizer,
    run_config,
    global_step,
    checkpoint_dir,
):
    checkpoint = {
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "run_config": run_config,
        "global_step": global_step,
    }
    checkpoint_path = os.path.join(checkpoint_dir, f"last_ckpt.pt")
    print(f"Saving checkpoint to {checkpoint_path}")
    torch.save(checkpoint, checkpoint_path)


def validate(model, val_loader, kld_weight: float = 0.0):
    print("Validating...")
    pbar = tqdm(val_loader, total=len(val_loader))
    # accumulate loss, reconstruction loss, and kld loss
    total_loss = 0.0
    total_recons_loss = 0.0
    total_kld_loss = 0.0
    total_samples = 0
    for audio_batch in pbar:
        waveform_original = audio_batch.cuda()

        with torch.no_grad():
            waveform_recon, mdct_recon, mdct_original, mu, log_var = model(
                waveform_original
            )

            loss_dict = vae_loss(
                mdct_recon,
                mdct_original,
                waveform_recon,
                waveform_original,
                mu,
                log_var,
                kld_weight=kld_weight,
            )

            loss = loss_dict["loss"]
            loss = loss.mean()

            total_loss += loss.item()
            total_recons_loss += loss_dict["reconstruction_loss"].item()
            total_kld_loss += loss_dict["kld_loss"].item()
            total_samples += len(waveform_original)

    return (
        total_loss / total_samples,
        total_recons_loss / total_samples,
        total_kld_loss / total_samples,
    )


if __name__ == "__main__":

    run_start_time = time.strftime("%Y-%m-%d_%H-%M-%S")
    checkpoint_dir = f"/app/suno/christian/checkpoints/mdct-codec/{run_start_time}_s{random.randint(0, 9999)}"
    os.makedirs(checkpoint_dir, exist_ok=False)

    torch.set_float32_matmul_precision("medium")

    # Initialize distributed process group
    local_rank = int(os.environ.get("LOCAL_RANK", 0))
    print(f"Local rank: {local_rank}")
    # dist.init_process_group(backend="nccl")
    torch.cuda.set_device(local_rank)

    # set the seed differently for each process
    torch.manual_seed(local_rank)

    run_config = {
        "training": {
            "max_steps": 1_000_000,
            "run_name": "test",
            "project_name": "mdct-codec",
            "lr": 1e-4,
            "grad_clip_norm": 1.0,
            "kld_weight": 0.0001,
        },
        "model": {
            "frame_size": 1920,
            "latent_dim": 1920,
            "hidden_dims": [2048, 2048, 2048, 2048],
            "dropout": 0.0,
            "midside": False,
        },
        "dataset": {
            "train_audio_dir": "/app/suno/data/audio_2ch_48khz_lg/train/genius_hq",
            "val_audio_dir": "/app/suno/data/audio_2ch_48khz_lg/val/genius_hq",
            "batch_size": 32,
            "num_workers": 4,
            "chunk_size_s": 10.0,
            "buffer_size": 5_000,
            "sample_rate": 48_000,
        },
    }

    # Initialize wandb (only one process should do this)
    if local_rank == 0:
        wandb.init(
            project=run_config["training"]["project_name"],
            name=run_config["training"]["run_name"],
        )
        wandb.config.update(
            {"checkpoint_dir": checkpoint_dir, "run_config": run_config}
        )

    # find all filepaths in the train and val dirs
    train_filepaths = glob.glob(
        os.path.join(run_config["dataset"]["train_audio_dir"], "*.wav"), recursive=True
    )
    val_filepaths = glob.glob(
        os.path.join(run_config["dataset"]["val_audio_dir"], "*.wav"), recursive=True
    )

    # create train dataset
    train_dataset = BufferedAudioDataset(
        filepaths=train_filepaths,
        sample_rate=run_config["dataset"]["sample_rate"],
        num_workers=run_config["dataset"]["num_workers"],
        chunk_size_s=run_config["dataset"]["chunk_size_s"],
        buffer_size=run_config["dataset"]["buffer_size"],
    )

    train_loader = DataLoader(
        train_dataset,
        batch_size=run_config["dataset"]["batch_size"],
        num_workers=run_config["dataset"]["num_workers"],
        shuffle=True,
        pin_memory=True,
        persistent_workers=True,
        drop_last=True,
    )

    # create val dataset
    val_dataset = BufferedAudioDataset(
        filepaths=val_filepaths,
        sample_rate=run_config["dataset"]["sample_rate"],
        num_workers=run_config["dataset"]["num_workers"],
        chunk_size_s=run_config["dataset"]["chunk_size_s"],
        buffer_size=run_config["dataset"]["buffer_size"],
    )

    val_loader = DataLoader(
        val_dataset,
        batch_size=run_config["dataset"]["batch_size"],
        num_workers=run_config["dataset"]["num_workers"],
        persistent_workers=True,
        pin_memory=True,
        shuffle=False,
    )

    # create model and wrap in DDP
    model = WaveformMDCTVAE(
        frame_size=run_config["model"]["frame_size"],
        latent_dim=run_config["model"]["latent_dim"],
        hidden_dims=run_config["model"]["hidden_dims"],
        dropout=run_config["model"]["dropout"],
        midside=run_config["model"]["midside"],
    )
    num_params = sum(p.numel() for p in model.parameters())
    print(f"Number of parameters: {num_params/1e6:0.1f}M")
    model.cuda()
    # model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])
    # model = torch.compile(model)  # Add dynamo compilation

    optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"])

    warmup_scheduler = torch.optim.lr_scheduler.LinearLR(
        optimizer, start_factor=0.001, end_factor=1.0, total_iters=1000
    )
    cosine_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, run_config["training"]["max_steps"] - 1000
    )
    scheduler = torch.optim.lr_scheduler.ChainedScheduler(
        [warmup_scheduler, cosine_scheduler]
    )

    global_step = 0
    while global_step < run_config["training"]["max_steps"]:
        pbar = tqdm(train_loader, total=len(train_loader))
        for audio_batch in pbar:
            optimizer.zero_grad()

            # move to gpu
            audio_batch = audio_batch.cuda()

            # run the compression model
            waveform_recon, mdct_recon, mdct_original, mu, log_var = model(audio_batch)

            # Compute loss
            loss_dict = vae_loss(
                mdct_recon,
                mdct_original,
                waveform_recon,
                audio_batch,
                mu,
                log_var,
                kld_weight=run_config["training"]["kld_weight"],
            )
            loss = loss_dict["loss"]
            loss.backward()

            torch.nn.utils.clip_grad_norm_(
                model.parameters(), run_config["training"]["grad_clip_norm"]
            )

            optimizer.step()
            scheduler.step()

            loss = loss.mean()
            grad_norm = torch.norm(
                torch.stack(
                    [
                        torch.norm(p.grad)
                        for p in model.parameters()
                        if p.grad is not None
                    ]
                )
            )

            if local_rank == 0:
                pbar.set_postfix({"loss": loss.item()})
                wandb.log(
                    {
                        "train/loss": loss.item(),
                        "train/grad_norm": grad_norm.item(),
                        "train/reconstruction_loss": loss_dict[
                            "reconstruction_loss"
                        ].item(),
                        "train/kld_loss": loss_dict["kld_loss"].item(),
                        "trainer/lr": optimizer.param_groups[0]["lr"],
                        "trainer/global_step": global_step,
                    }
                )
                global_step += 1

        if local_rank == 0:
            print(f"Step {global_step} loss: {loss.item():.4f}")

            val_loss, val_recons_loss, val_kld_loss = validate(
                model, val_loader, kld_weight=run_config["training"]["kld_weight"]
            )

            wandb.log(
                {
                    "val/loss": val_loss,
                    "val/reconstruction_loss": val_recons_loss,
                    "val/kld_loss": val_kld_loss,
                    "trainer/global_step": global_step,
                }
            )
            save_checkpoint(model, optimizer, run_config, global_step, checkpoint_dir)

    print("Done!")
