import os
import sys
import warnings
from dataclasses import dataclass
from pathlib import Path

import argbind
import auraloss
import torch
from audiotools import AudioSignal
from audiotools import ml
from audiotools.core import util
from audiotools.data import transforms
from audiotools.data.datasets import AudioDataset
from audiotools.data.datasets import AudioLoader
from audiotools.data.datasets import ConcatDataset
from audiotools.ml.decorators import timer
from audiotools.ml.decorators import Tracker
from audiotools.ml.decorators import when
from torch.utils.tensorboard import SummaryWriter

from dac.model.dac2 import DAC as DAC_import
from dac.model.discriminator4 import Discriminator as Discriminator_import
from dac.nn import loss as loss_import
from dac.utils.accelerator import Accelerator

USE_AURALOSS = False

warnings.filterwarnings("ignore", category=UserWarning)

# Enable cudnn autotuner to speed up training
# (can be altered by the funcs.seed function)
torch.backends.cudnn.benchmark = bool(int(os.getenv("CUDNN_BENCHMARK", 1)))
# Uncomment to trade memory for speed.

# Optimizers
AdamW = argbind.bind(torch.optim.AdamW, "generator", "discriminator")
Accelerator = argbind.bind(Accelerator, without_prefix=True)


@argbind.bind("generator", "discriminator")
def ExponentialLR(optimizer, gamma: float = 1.0):
    return torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma)


# Models
DAC = argbind.bind(DAC_import)
Discriminator = argbind.bind(Discriminator_import)

# Data
AudioDataset = argbind.bind(AudioDataset, "train", "val")
AudioLoader = argbind.bind(AudioLoader, "train", "val")

# Transforms
filter_fn = lambda fn: hasattr(fn, "transform") and fn.__qualname__ not in [
    "BaseTransform",
    "Compose",
    "Choose",
]
tfm = argbind.bind_module(transforms, "train", "val", filter_fn=filter_fn)

# Loss
filter_fn = lambda fn: hasattr(fn, "forward") and "Loss" in fn.__name__
losses = argbind.bind_module(loss_import, filter_fn=filter_fn)


def get_infinite_loader(dataloader):
    while True:
        for batch in dataloader:
            yield batch


@argbind.bind("train", "val")
def build_transform(
    augment_prob: float = 1.0,
    preprocess: list = ["Identity"],
    augment: list = ["Identity"],
    postprocess: list = ["Identity"],
):
    to_tfm = lambda l: [getattr(tfm, x)() for x in l]
    preprocess = transforms.Compose(*to_tfm(preprocess), name="preprocess")
    augment = transforms.Compose(*to_tfm(augment), name="augment", prob=augment_prob)
    postprocess = transforms.Compose(*to_tfm(postprocess), name="postprocess")
    transform = transforms.Compose(preprocess, augment, postprocess)
    return transform


@argbind.bind("train", "val", "test")
def build_dataset(
    sample_rate: int,
    folders: dict = None,
):
    # Give one loader per key/value of dictionary, where
    # value is a list of folders. Create a dataset for each one.
    # Concatenate the datasets with ConcatDataset, which
    # cycles through them.
    datasets = []
    for _, v in folders.items():
        loader = AudioLoader(sources=v)
        transform = build_transform()
        dataset = AudioDataset(loader, sample_rate, transform=transform)
        datasets.append(dataset)

    dataset = ConcatDataset(datasets)
    dataset.transform = transform
    return dataset


@dataclass
class State:
    generator: DAC
    optimizer_g: AdamW
    scheduler_g: ExponentialLR

    discriminator: Discriminator
    optimizer_d: AdamW
    scheduler_d: ExponentialLR

    mel_loss: (
        auraloss.freq.SumAndDifferenceSTFTLoss
        if USE_AURALOSS
        else losses.MelSpectrogramLoss
    )
    gan_loss: losses.GANLoss

    train_data: AudioDataset
    val_data: AudioDataset

    tracker: Tracker
    sample_rate: int


@argbind.bind(without_prefix=True)
def load(
    args,
    accel: Accelerator,
    tracker: Tracker,
    save_path: str,
    resume: bool = False,
    tag: str = "latest",
    load_weights: bool = False,
):
    generator, g_extra = None, {}
    discriminator, d_extra = None, {}

    if resume:
        kwargs = {
            "folder": f"{save_path}/{tag}",
            "map_location": "cpu",
            "package": not load_weights,
        }
        tracker.print(f"Resuming from {str(Path('.').absolute())}/{kwargs['folder']}")
        if (Path(kwargs["folder"]) / "dac").exists():
            generator, g_extra = DAC.load_from_folder(**kwargs)
        if (Path(kwargs["folder"]) / "discriminator").exists():
            discriminator, d_extra = Discriminator.load_from_folder(**kwargs)

    generator = DAC() if generator is None else generator
    print(generator)

    discriminator = Discriminator() if discriminator is None else discriminator

    # tracker.print(generator)
    # tracker.print(discriminator)

    generator = accel.prepare_model(generator)
    discriminator = accel.prepare_model(discriminator)

    with argbind.scope(args, "generator"):
        optimizer_g = AdamW(generator.parameters(), use_zero=accel.use_ddp)
        scheduler_g = ExponentialLR(optimizer_g)
    with argbind.scope(args, "discriminator"):
        optimizer_d = AdamW(discriminator.parameters(), use_zero=accel.use_ddp)
        scheduler_d = ExponentialLR(optimizer_d)

    if "optimizer.pth" in g_extra:
        optimizer_g.load_state_dict(g_extra["optimizer.pth"])
    if "scheduler.pth" in g_extra:
        scheduler_g.load_state_dict(g_extra["scheduler.pth"])
    if "tracker.pth" in g_extra:
        tracker.load_state_dict(g_extra["tracker.pth"])

    if "optimizer.pth" in d_extra:
        optimizer_d.load_state_dict(d_extra["optimizer.pth"])
    if "scheduler.pth" in d_extra:
        scheduler_d.load_state_dict(d_extra["scheduler.pth"])

    sample_rate = accel.unwrap(generator).sample_rate
    with argbind.scope(args, "train"):
        train_data = build_dataset(sample_rate)
    with argbind.scope(args, "val"):
        val_data = build_dataset(sample_rate)

    if USE_AURALOSS:
        mel_loss = auraloss.freq.SumAndDifferenceSTFTLoss(
            sample_rate=args["DAC.sample_rate"],
            fft_sizes=[2048, 1024, 512, 256, 128, 64, 32],
            hop_sizes=[512, 256, 128, 64, 32, 16, 8],
            win_lengths=[2048, 1024, 512, 256, 128, 64, 32],
            perceptual_weighting=True,
        )
    else:
        # mel_loss = losses.BandSplitSpectrogramLoss()
        mel_loss = losses.MelSpectrogramLoss()
    gan_loss = losses.GANLoss(discriminator)

    return State(
        generator=generator,
        optimizer_g=optimizer_g,
        scheduler_g=scheduler_g,
        discriminator=discriminator,
        optimizer_d=optimizer_d,
        scheduler_d=scheduler_d,
        mel_loss=mel_loss,
        gan_loss=gan_loss,
        tracker=tracker,
        sample_rate=args["DAC.sample_rate"],
        train_data=train_data,
        val_data=val_data,
    )


@timer()
@torch.no_grad()
def val_loop(batch, state, accel):
    state.generator.eval()
    batch = util.prepare_batch(batch, accel.device)
    signal = state.val_data.transform(
        batch["signal"].clone(), **batch["transform_args"]
    )

    out = state.generator(signal.audio_data, signal.sample_rate)
    recons = AudioSignal(out["audio"], signal.sample_rate)

    if USE_AURALOSS:
        mel_loss = state.mel_loss(recons.audio_data, signal.audio_data)
    else:
        signal_arr = signal.audio_data
        recons_arr = recons.audio_data
        # separate channels for stereo
        b, _, t = signal.audio_data.shape
        signal_flat = AudioSignal(signal_arr.reshape(b * 2, 1, t), state.sample_rate)
        recons_flat = AudioSignal(recons_arr.reshape(b * 2, 1, t), state.sample_rate)
        mel_loss = state.mel_loss(recons_flat, signal_flat)
        # mono signal
        signal_flat = AudioSignal(
            signal_arr.mean(dim=1, keepdim=True), state.sample_rate
        )
        recons_flat = AudioSignal(
            recons_arr.mean(dim=1, keepdim=True), state.sample_rate
        )
        mel_loss = mel_loss + state.mel_loss(recons_flat, signal_flat)
        mel_loss = mel_loss / 2

    return {
        "loss": mel_loss,
        "mel/loss": mel_loss,
    }


@timer()
def train_loop(state, batch, accel, lambdas):
    state.generator.train()
    state.discriminator.train()
    output = {}

    batch = util.prepare_batch(batch, accel.device)
    with torch.no_grad():
        signal = state.train_data.transform(
            batch["signal"].clone(), **batch["transform_args"]
        )

    with accel.autocast():
        out = state.generator(signal.audio_data, signal.sample_rate)
        recons = AudioSignal(out["audio"], signal.sample_rate)
        commitment_loss = (
            out["vq/commitment_loss"] if "vq/commitment_loss" in out else 0
        )
        codebook_loss = out["vq/codebook_loss"] if "vq/codebook_loss" in out else 0
        entropy_loss = out["vq/entropy_loss"] if "vq/entropy_loss" in out else 0
        codebook_entropy = (
            out["vq/codebook_entropy"] if "vq/codebook_entropy" in out else 0
        )
        orthogonal_loss = (
            out["vq/orthogonal_loss"] if "vq/orthogonal_loss" in out else 0
        )
        aux_loss = out["aux_loss"] if "aux_loss" in out else 0
        kl = out["kl"] if "kl" in out else 0

    with accel.autocast():
        output["adv/disc_loss"] = state.gan_loss.discriminator_loss(recons, signal)

    state.optimizer_d.zero_grad()
    accel.backward(output["adv/disc_loss"])
    accel.scaler.unscale_(state.optimizer_d)
    output["other/grad_norm_d"] = torch.nn.utils.clip_grad_norm_(
        state.discriminator.parameters(), 10.0
    )
    accel.step(state.optimizer_d)
    state.scheduler_d.step()

    with accel.autocast():
        if USE_AURALOSS:
            output["mel/loss"] = state.mel_loss(recons.audio_data, signal.audio_data)
        else:
            signal_arr = signal.audio_data
            recons_arr = recons.audio_data
            # separate channels for stereo
            b, _, t = signal.audio_data.shape
            signal_flat = AudioSignal(
                signal_arr.reshape(b * 2, 1, t), state.sample_rate
            )
            recons_flat = AudioSignal(
                recons_arr.reshape(b * 2, 1, t), state.sample_rate
            )
            mel_loss = state.mel_loss(recons_flat, signal_flat)
            # mono signal
            signal_flat = AudioSignal(
                signal_arr.mean(dim=1, keepdim=True), state.sample_rate
            )
            recons_flat = AudioSignal(
                recons_arr.mean(dim=1, keepdim=True), state.sample_rate
            )
            mel_loss = mel_loss + state.mel_loss(recons_flat, signal_flat)
            mel_loss = mel_loss / 2
            output["mel/loss"] = mel_loss
        (
            output["adv/gen_loss"],
            output["adv/feat_loss"],
        ) = state.gan_loss.generator_loss(recons, signal)
        output["vq/commitment_loss"] = commitment_loss
        output["vq/codebook_loss"] = codebook_loss
        output["vq/entropy_loss"] = entropy_loss
        output["vq/codebook_entropy"] = codebook_entropy
        output["vq/orthogonal_loss"] = orthogonal_loss
        output["aux_loss"] = aux_loss
        output["kl"] = kl
        output["loss"] = sum([v * output[k] for k, v in lambdas.items() if k in output])

    state.optimizer_g.zero_grad()
    accel.backward(output["loss"])
    accel.scaler.unscale_(state.optimizer_g)
    output["other/grad_norm"] = torch.nn.utils.clip_grad_norm_(
        state.generator.parameters(), 1e3
    )
    accel.step(state.optimizer_g)
    state.scheduler_g.step()
    accel.update()

    output["other/learning_rate"] = state.optimizer_g.param_groups[0]["lr"]
    output["other/batch_size"] = signal.batch_size * accel.world_size

    return {k: v for k, v in sorted(output.items())}


def checkpoint(state, save_iters, save_path):
    metadata = {"logs": state.tracker.history}

    tags = ["latest"]
    state.tracker.print(f"Saving to {str(Path('.').absolute())}")
    if state.tracker.is_best("val", "mel/loss"):
        state.tracker.print(f"Best generator so far")
        tags.append("best")
    if state.tracker.step in save_iters:
        tags.append(f"{state.tracker.step // 1000}k")

    for tag in tags:
        generator_extra = {
            "optimizer.pth": state.optimizer_g.state_dict(),
            "scheduler.pth": state.scheduler_g.state_dict(),
            "tracker.pth": state.tracker.state_dict(),
            "metadata.pth": metadata,
        }
        accel.unwrap(state.generator).metadata = metadata
        accel.unwrap(state.generator).save_to_folder(
            f"{save_path}/{tag}", generator_extra, package=False
        )
        discriminator_extra = {
            "optimizer.pth": state.optimizer_d.state_dict(),
            "scheduler.pth": state.scheduler_d.state_dict(),
        }
        accel.unwrap(state.discriminator).save_to_folder(
            f"{save_path}/{tag}", discriminator_extra, package=False
        )


@torch.no_grad()
def cycle(state, audio_path):
    import numpy as np

    state.generator.eval()

    signal = AudioSignal(audio_path, device="cpu").resample(state.sample_rate)
    # print(signal.audio_data.shape)
    audio_data = signal.audio_data
    # if mono make stereo
    if audio_data.shape[1] == 1:
        audio_data = np.repeat(audio_data, 2, axis=1)
    recons = state.generator(audio_data, signal.sample_rate)["audio"]
    recons = AudioSignal(recons, signal.sample_rate)
    return recons


@torch.no_grad()
def save_golden_samples(state):
    from tempfile import NamedTemporaryFile
    import wandb

    state.tracker.print("Saving golden samples to wandb")

    golden_samples = os.listdir(
        "/home/minz/glockenspiel/descript-audio-codec/data/golden/"
    )
    for i, sample in enumerate(golden_samples):
        recons = cycle(
            state,
            f"/home/minz/glockenspiel/descript-audio-codec/data/golden/{sample}",
        )
        with NamedTemporaryFile(suffix=".mp3") as f:
            recons.cpu().write(f.name)
            wandb.log({f"{sample}": wandb.Audio(f.name)}, step=state.tracker.step)


def validate(state, val_dataloader, accel):
    for batch in val_dataloader:
        output = val_loop(batch, state, accel)
    # Consolidate state dicts if using ZeroRedundancyOptimizer
    if hasattr(state.optimizer_g, "consolidate_state_dict"):
        state.optimizer_g.consolidate_state_dict()
        state.optimizer_d.consolidate_state_dict()
    return output


@argbind.bind(without_prefix=True)
def train(
    args,
    accel: Accelerator,
    seed: int = 0,
    save_path: str = "ckpt",
    num_iters: int = 250000,
    save_iters: list = [10000, 50000, 100000, 200000],
    sample_freq: int = 10000,
    valid_freq: int = 1000,
    batch_size: int = 12,
    val_batch_size: int = 10,
    num_workers: int = 8,
    val_idx: list = [0, 1, 2, 3, 4, 5, 6, 7],
    lambdas: dict = {
        "mel/loss": 100.0,
        "adv/feat_loss": 2.0,
        "adv/gen_loss": 1.0,
        "vq/commitment_loss": 0.25,
        "vq/codebook_loss": 1.0,
    },
    wandb_log: bool = False,
    wandb_project: str = "dac_test",
    wandb_run_name: str = "mw scale notanh bst fine",
    log_interval: int = 100,
):
    util.seed(seed)
    Path(save_path).mkdir(exist_ok=True, parents=True)
    writer = (
        SummaryWriter(log_dir=f"{save_path}/logs") if accel.local_rank == 0 else None
    )
    tracker = Tracker(
        writer=writer, log_file=f"{save_path}/log.txt", rank=accel.local_rank
    )

    state = load(args, accel, tracker, save_path, resume=True, load_weights=True)
    train_dataloader = accel.prepare_dataloader(
        state.train_data,
        start_idx=state.tracker.step * batch_size,
        num_workers=num_workers,
        batch_size=batch_size,
        collate_fn=state.train_data.collate,
    )
    train_dataloader = get_infinite_loader(train_dataloader)
    val_dataloader = accel.prepare_dataloader(
        state.val_data,
        start_idx=0,
        num_workers=num_workers,
        batch_size=val_batch_size,
        collate_fn=state.val_data.collate,
        persistent_workers=True if num_workers > 0 else False,
    )

    master_process = accel.rank == 0
    if master_process and wandb_log:
        import wandb

        wandb.init(
            project=wandb_project,
            name=wandb_run_name,
            config=args,
        )

    # Wrap the functions so that they neatly track in TensorBoard + progress bars
    # and only run when specific conditions are met.
    global train_loop, val_loop, validate, save_golden_samples, checkpoint
    train_loop = tracker.log("train", "value", history=False)(
        tracker.track("train", num_iters, completed=state.tracker.step)(train_loop)
    )
    val_loop = tracker.track("val", len(val_dataloader))(val_loop)
    validate = tracker.log("val", "mean")(validate)

    # These functions run only on the 0-rank process
    save_golden_samples = when(lambda: accel.local_rank == 0)(save_golden_samples)
    checkpoint = when(lambda: accel.local_rank == 0)(checkpoint)

    with tracker.live:
        for tracker.step, batch in enumerate(train_dataloader, start=tracker.step):
            train_out = train_loop(state, batch, accel, lambdas)

            if master_process and wandb_log and tracker.step % log_interval == 0:
                wandb.log(train_out, step=tracker.step)

            last_iter = (
                tracker.step == num_iters - 1 if num_iters is not None else False
            )
            if tracker.step % sample_freq == 0 or last_iter:
                save_golden_samples(state)

            if tracker.step % valid_freq == 0 or last_iter:
                val_out = validate(state, val_dataloader, accel)
                if master_process and wandb_log:
                    # add val_ to keys to avoid confusion with train metrics
                    val_out = {f"val_{k}": v for k, v in val_out.items()}
                    wandb.log(val_out, step=tracker.step)
                checkpoint(state, save_iters, save_path)
                # Reset validation progress bar, print summary since last validation.
                tracker.done("val", f"Iteration {tracker.step}")

            if last_iter:
                break


if __name__ == "__main__":
    args = argbind.parse_args()
    args["args.debug"] = int(os.getenv("LOCAL_RANK", 0)) == 0
    with argbind.scope(args):
        with Accelerator() as accel:
            if accel.local_rank != 0:
                sys.tracebacklimit = 0
            train(args, accel)
