import math
import torch
import torch.nn.functional as F
import time
from tqdm import tqdm
import torch.distributed as dist
from pathlib import Path
from queue import Queue, Empty
from threading import Thread, Event

try:
    import wandb
except ImportError:
    wandb = None


class DummyOptimizer(torch.optim.Optimizer):
    def __init__(self):
        self.param_groups = [{"lr": 0}]
        None

    def step(self):
        None

    def zero_grad(self, set_to_none=False):
        None


class DummyScheduler:
    def step(self):
        None


class TrainState:
    """Track number of steps, examples, and tokens processed"""

    step: int = 0  # Steps in the current epoch
    accum_step: int = 0  # Number of gradient accumulation steps
    samples: int = 0  # total # of examples used
    tokens: int = 0  # total # of tokens processed
    best_eval_acc: float = 0  # best eval accuracy so far
    skipped_updates: int = 0  # number of updates skipped due to NaN gradients


def make_perplexity_loss(pad_idx):
    def perplexity(out, target):
        x = out.detach().contiguous().view(-1, out.size(-1))
        y = target.detach().contiguous().view(-1)
        return F.cross_entropy(x, y, ignore_index=pad_idx, reduction="sum")

    return perplexity


def run_epoch(
    data_iter,
    make_eval_data_iter,
    model,
    loss_compute,
    optimizer,
    scheduler,
    batch_size,
    mode="train",
    num_batches=None,
    num_eval_batches=None,
    accum_iter=1,
    eval_iter=1,
    enable_checkpoint=False,
    epoch=0,
    train_state=TrainState(),
    quiet=False,
    scaler=None,
    max_examples_per_rank=None,
    checkpoint_path="./checkpoints",
    world_size=1,
):
    """Train a single epoch"""

    last_eval_acc = None

    num_kept_checkpoints = 5

    def checkpoint(step, acc):
        # only keep the last num_kept_checkpoints checkpoints
        files_by_mtime = list(
            sorted((p.stat().st_mtime, p) for p in Path(checkpoint_path).glob("*.pt"))
        )
        if len(files_by_mtime) > num_kept_checkpoints:
            for _, p in files_by_mtime[:-num_kept_checkpoints]:
                p.unlink()
        torch.save(
            {
                "model_state_dict": model.state_dict(),
                "optimizer_state_dict": optimizer.state_dict(),
                "scheduler_state_dict": scheduler.state_dict(),
                "train_state": train_state,
            },
            f"{checkpoint_path}/epoch_{epoch}_step_{step}_accuracy_{acc}.pt",
        )

    def verify_ranks_are_equal():
        dist.barrier()
        for p in model.parameters():
            if dist.get_rank() == 0:
                lst = [torch.zeros_like(p.data) for _ in range(world_size)]
                dist.gather(p.data, gather_list=lst)
                for i in range(1, world_size):
                    if not torch.allclose(p.data, lst[i]):
                        raise RuntimeError("ranks are not equal")
            else:
                dist.gather(p.data)

    def eval():
        nonlocal last_eval_acc
        model.eval()
        with torch.no_grad():
            verify_ranks_are_equal()
            eval_loss, eval_tokens, eval_accurate_count, _ = run_epoch(
                make_eval_data_iter(),
                None,
                model,
                make_perplexity_loss(model.module.params.padding_idx),
                DummyOptimizer(),
                DummyScheduler(),
                batch_size,
                num_batches=num_eval_batches,
                mode="eval",
                quiet=quiet,
                scaler=scaler,
                world_size=world_size,
                checkpoint_path=checkpoint_path,
            )
            dist.all_reduce(eval_loss)
            dist.all_reduce(eval_tokens)
            dist.all_reduce(eval_accurate_count)
            dist.barrier()

            perplexity = eval_loss.item() / eval_tokens.item()
            accuracy = eval_accurate_count.item() / eval_tokens.item()

            if not quiet:
                tqdm.write(f"eval_perplexity: {perplexity}, eval_accuracy: {accuracy}")
            if wandb.run is not None:
                wandb.log(
                    {"eval_perplexity": perplexity, "eval_accuracy": accuracy},
                    commit=False,
                )
            last_eval_acc = accuracy
        model.train()

    start = time.time()
    total_tokens = torch.tensor([0], device="cuda", dtype=torch.long)
    total_loss = torch.tensor([0], device="cuda", dtype=torch.float)
    total_accurate_count = torch.tensor([0], device="cuda", dtype=torch.long)
    accum_tokens = torch.tensor([0], device="cuda", dtype=torch.long)
    display_tokens = torch.tensor([0], device="cuda", dtype=torch.long)
    display_tokens_postdesc = torch.tensor([0], device="cuda", dtype=torch.long)
    display_loss = torch.tensor([0], device="cuda", dtype=torch.float)
    display_accurate_count = torch.tensor([0], device="cuda", dtype=torch.long)
    display_accurate_count_postdesc = torch.tensor([0], device="cuda", dtype=torch.long)
    display_overflow_examples = torch.tensor([0], device="cuda", dtype=torch.long)
    display_batch_underruns = torch.tensor([0], device="cuda", dtype=torch.long)
    display_batch_fill = torch.tensor([0], device="cuda", dtype=torch.float)
    n_accum = 0
    tqdm_total = num_batches
    if max_examples_per_rank is not None:
        tqdm_total = min(
            tqdm_total if tqdm_total is not None else +math.inf,
            max_examples_per_rank // batch_size,
        )

    iterator = (i for i in tqdm(data_iter, total=tqdm_total, disable=quiet))

    batch_q = Queue(20)
    batch_evt = Event()

    def batch_thd_entry():
        while True:
            try:
                nextbatch = next(iterator)
            except StopIteration:
                nextbatch = None
            except Exception as e:
                nextbatch = e
            batch_q.put(nextbatch)
            if batch_evt.is_set():
                iterator.close()
                break
            if nextbatch is None:
                break

    batch_thd = Thread(target=batch_thd_entry)
    batch_thd.start()

    try:
        i = 0
        verify_ranks_are_equal()
        while True:
            try:
                batch = batch_q.get(block=False)
            except Empty:
                display_batch_underruns[0] += 1
                batch = batch_q.get()
            if isinstance(batch, Exception):
                raise batch
            if batch is None:
                break
            out = model(
                batch.tgt,
                encoder_input_ids=batch.encoder_input_ids,
                encoder_attention_mask=batch.encoder_attention_mask,
            )
            loss_node = loss_compute(out, batch.tgt_y)
            accum_tokens += batch.ntokens
            if mode == "train" or mode == "train+log":
                if scaler is not None:
                    scaler.scale(loss_node).backward()
                else:
                    loss_node.backward()
                train_state.step += 1
                train_state.samples += batch.tgt.shape[0]
                train_state.tokens += batch.ntokens
                if i % accum_iter == 0:
                    if scaler is not None:
                        scaler.unscale_(optimizer)
                    # scale gradients by total number of non-pad tokens accumulated
                    dist.all_reduce(accum_tokens)
                    # compensate for DDP dividing by number of workers (no way to disable)
                    if accum_tokens > 0:
                        factor = world_size / accum_tokens
                        for p in model.parameters():
                            if p.grad is not None:
                                p.grad *= factor
                    total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
                    if torch.logical_or(total_norm.isnan(), total_norm.isinf()):
                        train_state.skipped_updates += 1
                    if scaler is not None:
                        scaler.step(optimizer)
                        scaler.update()
                    else:
                        optimizer.step()
                    if eval_iter is not None and i // accum_iter % eval_iter == 0:
                        eval()
                        if (
                            enable_checkpoint
                            and last_eval_acc > train_state.best_eval_acc
                        ):
                            checkpoint(i, last_eval_acc)
                            train_state.best_eval_acc = last_eval_acc
                    if eval_iter is None and enable_checkpoint and i % 1000 == 0:
                        checkpoint(i, 0)
                    optimizer.zero_grad(set_to_none=True)
                    accum_tokens[0] = 0
                    n_accum += 1
                    train_state.accum_step += 1
                    scheduler.step()

            accurate_count = torch.count_nonzero(
                torch.logical_and(
                    torch.argmax(out, dim=-1) == batch.tgt_y,
                    batch.tgt[:, :, 0] != model.module.params.padding_idx,
                )
            )

            def map_zero_to_high(x):
                x[x == 0] = 1e6
                return x

            description_anchor_argmax = map_zero_to_high(
                torch.argmax(
                    (
                        batch.tgt[:, :, 0]
                        == model.module.vocab.description_anchor.index
                    ).to(torch.long),
                    dim=-1,
                    keepdim=True,
                )
            )
            postdesc_mask = torch.logical_and(
                torch.arange(batch.tgt.shape[1], dtype=torch.long, device="cuda")
                .unsqueeze(0)
                .expand(batch.tgt.shape[0], -1)
                > description_anchor_argmax,
                batch.tgt[:, :, 0] != model.module.params.padding_idx,
            )
            tokens_postdesc = torch.count_nonzero(postdesc_mask)
            accurate_count_postdesc = torch.count_nonzero(
                torch.logical_and(
                    torch.argmax(out, dim=-1) == batch.tgt_y, postdesc_mask
                )
            )
            del postdesc_mask

            batch_size = batch.tgt.shape[0]
            total_loss += loss_node.detach()
            total_tokens += batch.ntokens
            total_accurate_count += accurate_count
            display_tokens += batch.ntokens
            display_tokens_postdesc += tokens_postdesc
            display_loss += loss_node.detach()
            display_accurate_count += accurate_count
            display_accurate_count_postdesc += accurate_count_postdesc
            display_overflow_examples += torch.count_nonzero(
                description_anchor_argmax == 1e6
            )
            del description_anchor_argmax
            display_batch_fill += batch.fill
            if i % 40 == 1 and (mode == "train" or mode == "train+log"):
                lr = optimizer.param_groups[0]["lr"]
                dist.all_reduce(display_tokens)
                dist.all_reduce(display_tokens_postdesc)
                dist.all_reduce(display_loss)
                dist.all_reduce(display_accurate_count)
                dist.all_reduce(display_accurate_count_postdesc)
                dist.all_reduce(display_overflow_examples)
                dist.all_reduce(display_batch_underruns)
                dist.all_reduce(display_batch_fill)
                elapsed = time.time() - start
                tokens_per_second = display_tokens.item() / elapsed
                loss_per_token = display_loss.item() / display_tokens.item()
                accuracy = display_accurate_count.item() / display_tokens.item()
                accuracy_postdesc = display_accurate_count_postdesc.item() / (
                    display_tokens_postdesc.item() + 1e-6
                )
                batch_underruns = display_batch_underruns.item()
                batch_fill = display_batch_fill.item() / (40 * world_size)
                overflow_pd_ratio = display_overflow_examples.item() / (
                    40 * world_size * batch_size
                )
                if not quiet:
                    tqdm.write(
                        (
                            "Epo Step: %6d / %d | Acc Step: %3d | Loss: %6.2f "
                            + "| Acc: %6.2f | AccPD: %6.2f | Tok / Sec: %7.1f | LR: %6.1e "
                            + "| Fill: %6.2f | OverflowPD: %6.2f | Underruns: %d"
                        )
                        % (
                            i,
                            num_batches,
                            n_accum,
                            loss_per_token,
                            accuracy,
                            accuracy_postdesc,
                            tokens_per_second,
                            lr,
                            batch_fill,
                            overflow_pd_ratio,
                            batch_underruns,
                        )
                    )
                start = time.time()
                display_tokens[0] = 0
                display_tokens_postdesc[0] = 0
                display_loss[0] = 0
                display_accurate_count[0] = 0
                display_accurate_count_postdesc[0] = 0
                display_overflow_examples[0] = 0
                display_batch_underruns[0] = 0
                display_batch_fill[0] = 0
                if wandb.run is not None:
                    log_obj = {
                        "loss": loss_per_token,
                        "accuracy": accuracy,
                        "accuracy_postdesc": accuracy_postdesc,
                        "lr": lr,
                        "tokens_per_second": tokens_per_second,
                        "epoch": 1.0 + epoch + i / num_batches,
                        "batch_fill": batch_fill,
                        "overflow_examples": overflow_pd_ratio,
                        "seq_len": batch.length,
                        "batch_underruns": batch_underruns,
                    }
                    if scaler is not None:
                        log_obj["scale"] = scaler.get_scale()
                        log_obj["skipped_updates"] = train_state.skipped_updates

                    wandb.log(log_obj)
                train_state.skipped_updates = 0

            del loss_node
            if (
                max_examples_per_rank is not None
                and train_state.step * batch_size > max_examples_per_rank
            ):
                break
            i += 1
        return total_loss, total_tokens, total_accurate_count, train_state
    finally:
        batch_evt.set()
        try:
            batch_q.get(timeout=2)
        except Empty:
            pass


def rate(step, model_size, factor, min_factor, steps_in_epoch):
    """
    we have to default the step to 1 for LambdaLR function
    to avoid zero raising to negative power.
    """
    max_lr = factor * model_size ** (-0.5)
    if step > steps_in_epoch:
        return min_factor * max_lr
    return max_lr * (
        min_factor
        + (1 + math.cos(step * math.pi / steps_in_epoch)) / 2 * (1 - min_factor)
    )
