import os
import time
import wandb
import math
import torch
import random
import auraloss
import torchaudio
import numpy as np
import pyloudnorm as pyln
import scipy.signal as signal

from tqdm import tqdm
from suno_utils.utils.text import read_jsonl
from typing import Tuple

import torch.distributed as dist
import torch.multiprocessing as mp
from torch.nn.parallel import DistributedDataParallel as DDP


def biquad(
    gain_db: float,
    cutoff_freq: float,
    q_factor: float,
    sample_rate: float,
    filter_type: str,
) -> Tuple[np.ndarray, np.ndarray]:
    """Use design parameters to generate coefficients for a specific filter type."""
    A = 10 ** (gain_db / 40.0)
    w0 = 2.0 * np.pi * (cutoff_freq / sample_rate)
    alpha = np.sin(w0) / (2.0 * q_factor)
    cos_w0 = np.cos(w0)
    sqrt_A = np.sqrt(A)

    if filter_type == "high_shelf":
        b0 = A * ((A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha)
        b1 = -2 * A * ((A - 1) + (A + 1) * cos_w0)
        b2 = A * ((A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha)
        a0 = (A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha
        a1 = 2 * ((A - 1) - (A + 1) * cos_w0)
        a2 = (A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha
    elif filter_type == "low_shelf":
        b0 = A * ((A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha)
        b1 = 2 * A * ((A - 1) - (A + 1) * cos_w0)
        b2 = A * ((A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha)
        a0 = (A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha
        a1 = -2 * ((A - 1) + (A + 1) * cos_w0)
        a2 = (A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha
    elif filter_type == "peaking":
        b0 = 1 + alpha * A
        b1 = -2 * cos_w0
        b2 = 1 - alpha * A
        a0 = 1 + alpha / A
        a1 = -2 * cos_w0
        a2 = 1 - alpha / A

    b = np.array([b0, b1, b2]) / a0
    a = np.array([1.0, a1 / a0, a2 / a0])
    return b, a


def apply_stereo_to_mono(audio: torch.Tensor, sample_rate: float):
    return audio.mean(dim=0, keepdims=True).repeat(2, 1)


def apply_channel_imbalance(
    audio: torch.Tensor, sample_rate: float, imbalance: float = 0.0
):
    if not -1 <= imbalance <= 1 or audio.shape[-2] != 2:
        raise ValueError("Invalid input")
    out = audio.clone()
    l_gain, r_gain = (1.0 - imbalance, 1.0) if imbalance > 0 else (1.0, 1.0 + imbalance)
    out[0, :], out[1, :] = out[0, :] * l_gain, out[1, :] * r_gain
    return out


def apply_highpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):
    return torchaudio.functional.highpass_biquad(audio, sample_rate, cutoff_hz)


def apply_lowpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):
    return torchaudio.functional.lowpass_biquad(audio, sample_rate, cutoff_hz)


def apply_noise(
    audio: torch.Tensor,
    sample_rate: float,
    gain_db: float = 0.0,
    noise_type: str = "white",
):
    gain_lin = 10 ** (gain_db / 20.0)
    noise = torch.randn_like(audio)

    if noise_type == "white":
        return audio + gain_lin * noise
    elif noise_type == "pink":
        b = torch.tensor([0.049922035, -0.095993537, 0.050612699, -0.004408786])
        a = torch.tensor([1, -2.494956002, 2.017265875, -0.522189400])
        noise = torchaudio.functional.filtfilt(noise, a, b)
        noise /= noise.abs().max()
        return audio + gain_lin * noise
    else:
        raise ValueError(f"Invalid noise type: {noise_type}")


def apply_shelving_filter(
    audio: torch.Tensor,
    sample_rate: float,
    gain_db: float,
    cutoff_freq: float,
    q_factor: float,
    filter_type: str,
):
    # convert x to numpy
    audio = audio.numpy()
    b, a = biquad(
        gain_db,
        cutoff_freq,
        q_factor,
        sample_rate,
        filter_type,
    )
    x = signal.lfilter(b, a, audio).astype(np.float32)
    return torch.from_numpy(x)


# randomized corrputions


def apply_random_noise(audio: torch.Tensor, sample_rate: float):
    noise_type = random.choice(["white", "pink"])
    if noise_type == "white":
        noise_gain = random.uniform(-96, -48)
    else:
        noise_gain = random.uniform(-48, -12)
    return apply_noise(audio, sample_rate, noise_gain, noise_type)


def apply_random_stereo_to_mono(audio: torch.Tensor, sample_rate: float):
    return apply_stereo_to_mono(audio, sample_rate)


def apply_random_channel_imbalance(audio: torch.Tensor, sample_rate: float):
    imbalance = random.uniform(-1.0, 1.0)
    return apply_channel_imbalance(audio, sample_rate, imbalance)


def apply_random_filter(audio: torch.Tensor, sample_rate: float):
    filter_type = random.choice(["highpass", "lowpass", "high_shelf", "low_shelf"])
    if filter_type == "highpass":
        cutoff_freq = random.uniform(20, 4000)
        return apply_highpass(audio, sample_rate, cutoff_freq)
    elif filter_type == "lowpass":
        cutoff_freq = random.uniform(1000, 16000)
        return apply_lowpass(audio, sample_rate, cutoff_freq)
    else:
        gain_db = random.uniform(-12, 12)
        if filter_type == "high_shelf":
            cutoff_freq = random.uniform(6000, 20000)
        else:
            cutoff_freq = random.uniform(20, 2000)
        q_factor = random.uniform(0.1, 10.0)
        return apply_shelving_filter(
            audio, sample_rate, gain_db, cutoff_freq, q_factor, filter_type
        )


def corrupt(waveform_tensor, sample_rate):
    """
    Apply a random number of corruptions (at least one) to the input waveform.
    """
    # List of available random corruption functions
    corruption_fns = [
        apply_random_noise,
        apply_random_stereo_to_mono,
        apply_random_channel_imbalance,
        apply_random_filter,
    ]
    n_corr = random.randint(1, len(corruption_fns))  # at least one
    selected = random.sample(corruption_fns, n_corr)
    print(selected)
    out = waveform_tensor.clone()
    for fn in selected:
        out = fn(out, sample_rate)
    return torch.tanh(out)


# audio pair dataset
# takes in a metas files
# metas has the following structure:
# {
#     "id": "<id>",
#     "input_local_path": "<audio_path>",
#     "target_local_path": "<audio_path>",
# }
class AudioPairDataset(torch.utils.data.Dataset):
    def __init__(self, metas_filepath, chunk_size_samples=262144, buffer_size=1000):
        self.metas_filepath = metas_filepath
        self.chunk_size_samples = chunk_size_samples
        self.buffer_size = buffer_size

        metas = read_jsonl(metas_filepath)
        print(f"Loaded {len(metas)} metas from {metas_filepath}")
        self.metas = metas
        self.buffer = []
        self.items_since_last_reload = self.buffer_size

    def __len__(self):
        return len(self.metas)

    def _reload_buffer(self):
        self.buffer = []

        rand_indices = np.random.permutation(len(self.metas))
        finished = False

        if int(os.environ["LOCAL_RANK"]) == 0:
            print(f"Reloading buffer with {len(rand_indices)} items")

        for idx in rand_indices:
            meta = self.metas[idx]
            # input_audio, sr = torchaudio.load(meta["input_filepath"])
            target_audio, sr = torchaudio.load(meta["target_filepath"])

            # make inut audio lowpass filtered
            # input_audio = torchaudio.functional.lowpass_filter(input_audio, sr, 4000)
            # compute number of chunks in the input audio
            num_chunks = target_audio.shape[1] // self.chunk_size_samples

            for i in range(num_chunks):
                start_sample = i * self.chunk_size_samples
                end_sample = start_sample + self.chunk_size_samples
                target_chunk = target_audio[:, start_sample:end_sample]

                # apply random corruptions
                corrupted_chunk = corrupt(target_chunk, sr)

                self.buffer.append((corrupted_chunk, target_chunk))

                if len(self.buffer) >= self.buffer_size:
                    finished = True
                    break

            if finished:
                break

        pbar.set_postfix(buffer_size=len(self.buffer))

    def __getitem__(self, idx):
        if self.items_since_last_reload == self.buffer_size:
            self._reload_buffer()
            self.items_since_last_reload = 0

        input_audio_chunk, target_audio_chunk = self.buffer[
            self.items_since_last_reload
        ]
        self.items_since_last_reload += 1

        return input_audio_chunk, target_audio_chunk


class AudioPatcher:
    def __init__(self, patch_size):
        self.patch_size = patch_size

    def to_patches(self, audio):
        # audio: (bs, 2, seq_len)
        bs, channels, seq_len = audio.shape
        assert (
            seq_len % self.patch_size == 0
        ), "Sequence length must be divisible by patch size"
        num_patches = seq_len // self.patch_size
        patches = audio.view(
            bs, channels, num_patches, self.patch_size
        )  # (bs, 2, num_patches, patch_size)
        return patches

    def flatten_patches(self, patches):
        # patches: (bs, 2, num_patches, patch_size)
        bs, channels, num_patches, patch_size = patches.shape
        flattened = patches.permute(0, 2, 1, 3).reshape(
            bs, num_patches * channels, patch_size
        )
        # (bs, num_patches * 2, patch_size)
        return flattened

    def unflatten_patches(self, flattened, original_channels=2):
        # flattened: (bs, num_patches * channels, patch_size)
        bs, total_patches, patch_size = flattened.shape
        num_patches = total_patches // original_channels
        patches = flattened.view(
            bs, num_patches, original_channels, patch_size
        ).permute(0, 2, 1, 3)
        # (bs, 2, num_patches, patch_size)
        return patches

    def reconstruct_audio(self, patches):
        # patches: (bs, 2, num_patches, patch_size)
        bs, channels, num_patches, patch_size = patches.shape
        audio = patches.reshape(bs, channels, num_patches * patch_size)
        return audio


class SinusoidalPositionalEncoding(torch.nn.Module):
    def __init__(self, hidden_dim):
        super().__init__()
        position = torch.arange(10000).unsqueeze(1)
        div_term = torch.exp(
            torch.arange(0, hidden_dim, 2) * -(math.log(10000.0) / hidden_dim)
        )
        pe = torch.zeros(10000, hidden_dim)
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        self.register_buffer("pe", pe)

    def forward(self, x):
        # x: (batch, num_patches, hidden_dim)
        return x + self.pe[: x.size(1)]


class RefinerTransformer(torch.nn.Module):
    def __init__(
        self,
        hidden_dim=1024,
        num_heads=8,
        num_transformer_layers=12,
        dropout=0.1,
        patch_size=1024,
    ):
        super().__init__()
        self.patch_size = patch_size

        self.patcher = AudioPatcher(patch_size)
        self.pos_embed = SinusoidalPositionalEncoding(hidden_dim)
        self.input_layer = torch.nn.Linear(patch_size, hidden_dim)

        # Transformer encoder
        encoder_layer = torch.nn.TransformerEncoderLayer(
            d_model=hidden_dim,
            nhead=num_heads,
            dim_feedforward=hidden_dim * 4,
            dropout=dropout,
        )
        self.transformer = torch.nn.TransformerEncoder(
            encoder_layer, num_transformer_layers
        )

        self.output_layer = torch.nn.Linear(hidden_dim, patch_size)

    def forward(self, input_audio):
        # input_audio shape: (batch_size, 2, seq_len)
        # first split into patches, which becomes (batch_size, 2, num_patches, patch_size)
        patches = self.patcher.to_patches(input_audio)
        flattened = self.patcher.flatten_patches(patches)
        x = self.input_layer(flattened)  # (batch_size, num_patches, hidden_dim)
        x = self.pos_embed(x)  # (batch_size, num_patches, hidden_dim)
        x = self.transformer(x)  # (batch_size, num_patches, hidden_dim)
        x = self.output_layer(x)  # (batch_size, num_patches, 1)
        x = self.patcher.unflatten_patches(x)  # (batch_size, 2, seq_len)
        # fold the patches back together
        output_audio = self.patcher.reconstruct_audio(x)
        return input_audio + output_audio


class HDemucsRefiner(torch.nn.Module):
    def __init__(self, **kwargs):
        super().__init__()
        self.hdemucs = torchaudio.models.HDemucs(sources=["output"], **kwargs)

    def forward(self, input_audio):
        return self.hdemucs.forward(input_audio).sum(dim=1)


class DummyRefiner(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = torch.nn.Conv1d(2, 2, kernel_size=3, padding=1)

    def forward(self, x):
        return self.linear(x)


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, "last_ckpt.pt")
    torch.save(checkpoint, checkpoint_path)
    print(f"Saved checkpoint to {checkpoint_path}")


def set_seed(base_seed=42):
    rank = int(os.environ.get("LOCAL_RANK", 0))
    seed = base_seed + rank
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)


def setup():
    dist.init_process_group(backend="nccl")
    torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))


def reduce_tensor(tensor, average=True):
    """
    Reduces a tensor from all processes so rank 0 can log the average value.
    """
    rt = tensor.clone()
    dist.all_reduce(rt, op=dist.ReduceOp.SUM)
    if average:
        rt /= dist.get_world_size()
    return rt


def train_step(
    batch,
    model,
    optimizer,
    time_loss_fn,
    freq_loss_fn,
    run_config,
):
    input_audio, target_audio = batch
    input_audio = input_audio.cuda()
    target_audio = target_audio.cuda()

    output_audio = model(input_audio)
    time_loss = time_loss_fn(output_audio, target_audio)
    freq_loss = freq_loss_fn(output_audio, target_audio)
    loss = (
        run_config["training"]["time_loss_weight"] * time_loss
        + run_config["training"]["freq_loss_weight"] * freq_loss
    )

    return time_loss, freq_loss, loss


if __name__ == "__main__":

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

    torch.set_float32_matmul_precision("medium")

    setup()
    set_seed()

    local_rank = int(os.environ["LOCAL_RANK"])
    world_size = int(os.environ["WORLD_SIZE"])
    device = torch.device("cuda", local_rank)
    torch.cuda.set_device(local_rank)

    run_config = {
        "training": {
            "max_steps": 100_000,
            "run_name": "refiner-10s",
            "project_name": "refiner",
            "lr": 1e-4,
            "grad_clip_norm": 1.0,
            "preload_ckpt": None,
            "preload_optimizer": False,
            "warmup_steps": 500,
            "ckpt_every": 1000,
            "val_every": 1000,
            "time_loss_weight": 2.0,
            "freq_loss_weight": 1.0,
        },
        "model": {
            "hidden_dim": 2048,
            "num_heads": 16,
            "num_transformer_layers": 12,
            "dropout": 0.1,
        },
        "dataset": {
            "train_metas_filepath": "/mnt/localdisk/tmp_cjs/v2-infill-data-v1/metas_tr.jsonl",
            "val_metas_filepath": "/mnt/localdisk/tmp_cjs/v2-infill-data-v1/metas_val.jsonl",
            "batch_size": 16,
            "num_workers": 4,
            "chunk_size_samples": 524288,
            "buffer_size": 1000,
        },
    }

    # setup the data
    train_dataset = AudioPairDataset(
        run_config["dataset"]["train_metas_filepath"],
        run_config["dataset"]["chunk_size_samples"],
        run_config["dataset"]["buffer_size"],
    )
    val_dataset = AudioPairDataset(
        run_config["dataset"]["val_metas_filepath"],
        run_config["dataset"]["chunk_size_samples"],
        run_config["dataset"]["buffer_size"],
    )
    train_sampler = torch.utils.data.distributed.DistributedSampler(
        train_dataset, num_replicas=world_size, rank=local_rank, shuffle=True
    )
    train_loader = torch.utils.data.DataLoader(
        train_dataset,
        batch_size=run_config["dataset"]["batch_size"],
        sampler=train_sampler,
        num_workers=run_config["dataset"]["num_workers"],
        pin_memory=True,
    )
    val_loader = torch.utils.data.DataLoader(
        val_dataset,
        batch_size=run_config["dataset"]["batch_size"],
        num_workers=run_config["dataset"]["num_workers"],
        shuffle=False,
        pin_memory=True,
    )

    # setup the model
    # model = Refiner(**run_config["model"])
    model = HDemucsRefiner()
    model = model.cuda()
    print(f"Model has {sum(p.numel() for p in model.parameters()):,} parameters")

    # wrap the model in DDP
    model = DDP(model, device_ids=[local_rank], find_unused_parameters=True)
    # setup the optimizer
    optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"])

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

    # setup losses
    time_loss_fn = torch.nn.L1Loss()
    # time_loss_fn = auraloss.time.SI_SDR()
    freq_loss_fn = auraloss.freq.SumAndDifferenceSTFTLoss(
        fft_sizes=[1024, 2048, 4096, 8192],
        hop_sizes=[128, 256, 512, 1024],
        win_lengths=[1024, 2048, 4096, 8192],
    )

    global_step = 0

    # setup wandb
    if local_rank == 0:
        wandb.init(
            project="refiner",
            name=run_config["training"]["run_name"],
            config=run_config,
        )
        wandb.config.update(
            {"checkpoint_dir": checkpoint_dir, "run_config": run_config}
        )

    while global_step < run_config["training"]["max_steps"]:
        print(f"[{local_rank}] Starting training loop")
        torch.distributed.barrier()

        train_sampler.set_epoch(global_step)

        pbar = tqdm(train_loader, total=len(train_loader))
        # Track time for iterations per second calculation
        start_time = time.time()
        iter_times = []

        for batch in pbar:

            optimizer.zero_grad()

            loss = torch.tensor(0.0, device=device)

            time_loss, freq_loss, loss = train_step(
                batch,
                model,
                optimizer,
                time_loss_fn,
                freq_loss_fn,
                run_config,
            )

            loss.backward()

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

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

            optimizer.step()
            scheduler.step()

            global_step += 1

            if local_rank == 0:
                # Calculate iterations per second
                iter_time = time.time() - start_time
                iter_times.append(iter_time)
                if len(iter_times) > 100:  # Keep a moving window
                    iter_times.pop(0)
                ips = 1.0 / (sum(iter_times) / len(iter_times))
                start_time = time.time()

                pbar.set_postfix(
                    loss=loss.item(), grad_norm=grad_norm.item(), ips=f"{ips:.2f}"
                )

            if (
                global_step % run_config["training"]["ckpt_every"] == 0
                and local_rank == 0
            ):
                save_checkpoint(
                    model, optimizer, run_config, global_step, checkpoint_dir
                )

            if local_rank == 0:
                wandb.log(
                    {
                        "train/loss": loss.item(),
                        "train/grad_norm": grad_norm.item(),
                        "train/time_loss": time_loss.item(),
                        "train/freq_loss": freq_loss.item(),
                        "train/global_step": global_step,
                        "train/lr": optimizer.param_groups[0]["lr"],
                        "train/iterations_per_second": ips,
                    },
                )
