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

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.ml.decorators import timer
from audiotools.ml.decorators import Tracker
from audiotools.ml.decorators import when
import numpy as np
from torch.utils.tensorboard import SummaryWriter
import torchaudio.functional as aF

from dac.model.dac2 import DAC as DAC_import
from dac.model.discriminator2 import Discriminator as Discriminator_import
from dac.nn import loss as loss_import

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

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


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

    discriminator: Discriminator
    optimizer_d: AdamW
    scheduler_d: ExponentialLR

    mel_loss: auraloss.freq.SumAndDifferenceSTFTLoss
    gan_loss: losses.GANLoss

    tracker: Tracker
    sample_rate: int
    n_samples_val: int


@argbind.bind(without_prefix=True)
def load(
    args,
    accel: ml.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
    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"])

    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,
    )
    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"],
        n_samples_val=args["val/n_examples"],
    )


@timer()
@torch.no_grad()
def val_loop(batch, state, accel):
    state.generator.eval()
    batch = batch.to(accel.device)
    recons = state.generator(batch, state.sample_rate)["audio"]
    return {
        "loss": state.mel_loss(recons, batch),
        "mel/loss": state.mel_loss(recons, batch),
    }


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

    batch = batch.to(accel.device)

    signal = AudioSignal(batch, state.sample_rate)
    global_batch_size = signal.batch_size * accel.world_size

    with accel.autocast():
        out = state.generator(batch, state.sample_rate)
        recons = AudioSignal(out["audio"], state.sample_rate)
        commitment_loss = out["vq/commitment_loss"]
        codebook_loss = out["vq/codebook_loss"]

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

    b, _, t = signal.audio_data.shape

    with accel.autocast():
        output["mel/loss"] = state.mel_loss(recons.audio_data, signal.audio_data)
        (
            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["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"] = global_batch_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 save_samples(state, val_dataloader, val_idx, writer):
    state.tracker.print("Saving audio samples to TensorBoard")
    state.generator.eval()

    samples = [val_dataloader[idx] for idx in val_idx]
    batch = torch.stack(samples).to(accel.device)

    signal = AudioSignal(batch, state.sample_rate)

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

    audio_dict = {"recons": recons}
    if state.tracker.step == 0:
        audio_dict["signal"] = signal

    for k, v in audio_dict.items():
        for nb in range(v.batch_size):
            v[nb].cpu().write_audio_to_tb(
                f"{k}/sample_{nb}.wav", writer, state.tracker.step
            )

@torch.no_grad()
def save_golden_samples(state, golden_dir, output_dir):
    state.tracker.print("Saving golden audio samples")
    state.generator.eval()

    golden_dir = Path(golden_dir)
    output_dir = Path(output_dir)

    golden_files = list(golden_dir.glob("*.wav"))
    golden_files.sort()

    for golden_file in golden_files:
        signal = AudioSignal(golden_file, state.sample_rate, device=accel.device)
        recons = state.generator(signal.audio_data, signal.sample_rate)["audio"]
        recons = AudioSignal(recons, signal.sample_rate)

        recons.cpu().write(output_dir / f"{golden_file.stem}_{state.tracker.step}.wav")



def validate(state, val_dataloader, accel):
    for n, batch in enumerate(val_dataloader):
        output = val_loop(batch, state, accel)
        if n == state.n_samples_val - 1:
            break

    # 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


COMMON_SAMPLE_RATES = [8000, 16000, 24000, 32000, 44100, 48000]


def _cycle_sample_rate(waveform, from_sample_rate=48_000, to_sample_rate=8_000):
    assert(isinstance(waveform, torch.Tensor))
    assert(len(waveform.shape) == 2)
    resampled_waveform = aF.resample(waveform, from_sample_rate, to_sample_rate)
    cycled_waveform = aF.resample(resampled_waveform, to_sample_rate, from_sample_rate)
    assert(waveform.shape == cycled_waveform.shape)
    return cycled_waveform


class CustomDataLoader:
    def __init__(self, mem_map_path, duration_s=5, batch_size=8, sample_rate=48000, is_val=False):
        self.data = np.memmap(mem_map_path, dtype=np.int16, mode="r")
        self.duration_s = duration_s
        self.batch_size = batch_size
        self.sample_rate = sample_rate
        self.n_samples = int(round(self.duration_s*self.sample_rate))*2
        self.is_val = is_val

    # def __len__(self):
    #     return self.length

    def __getitem__(self, idx):
        arr = np.array(self.data[idx:idx+self.n_samples].reshape(-1, 2).T)
        arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max)
        return arr

    def __iter__(self):
        return self

    def __next__(self):
        out = []
        for _ in range(self.batch_size):
            idx = int(random.randint(0, len(self.data) - self.n_samples - 1) / 2) * 2
            arr = np.array(self.data[idx:idx+self.n_samples].reshape(-1, 2).T)
            arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max)
            # if not self.is_val and random.random() >= 0.9:
            #     arr = _cycle_sample_rate(
            #         arr,
            #         from_sample_rate=self.sample_rate,
            #         to_sample_rate=random.choice(COMMON_SAMPLE_RATES),
            #     )
            out.append(arr)
        return torch.stack(out)

@argbind.bind(without_prefix=True)
def train(
    args,
    accel: ml.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,
    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 = "test",
    log_interval: int = 100,
    golden_dir: str = None,
    output_dir: str = None,
):
    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)

    # 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_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", args["val/n_examples"])(val_loop)
    validate = tracker.log("val", "mean")(validate)

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

    ddp = int(os.environ.get("RANK", -1)) != -1
    if ddp:
        seed_offset = torch.distributed.get_rank()
    else:
        seed_offset = 0
    random.seed(6006 + seed_offset)

    train_dataloader = CustomDataLoader(
        args["train/memmap_path"],
        duration_s=args["train/duration"],
        batch_size=args["batch_size"],
        sample_rate=args["DAC.sample_rate"],
        is_val=False,
    )
    val_dataloader = CustomDataLoader(
        args["val/memmap_path"],
        duration_s=args["val/duration"],
        batch_size=args["batch_size"],
        sample_rate=args["DAC.sample_rate"],
        is_val=True,
    )

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

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

    with tracker.live:
        for tracker.step, batch in enumerate(train_dataloader, start=tracker.step):

            # batch = {
            #     'idx': tensor([0, 0]),
            #     'signal': <audiotools.core.audio_signal.AudioSignal at 0x7f32639ff340>,
            #     'source_idx': tensor([0, 0]),
            #     'item_idx': tensor([2, 2]),
            #     'source': ['../gpt/samples/mini', '../gpt/samples/mini'],
            #     'path': ['../gpt/samples/mini/obama_short.mp3', '../gpt/samples/mini/obama_short.mp3'],
            # }

            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_samples(state, val_dataloader, args["val/samples_offsets"], writer)
                if golden_dir is not None and output_dir is not None:
                    save_golden_samples(state, golden_dir, output_dir)

            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)
