from collections import OrderedDict

import numpy as np
import torch
from datasets import load_from_disk
from torch.cuda.amp import GradScaler
from torch.utils.data import DataLoader
from tqdm import tqdm
from transformers import AdamW, AutoTokenizer, get_linear_schedule_with_warmup

import wandb
from models import GPT2LMHeadModel, GPT2LMHeadModelBackground, GPT2OdinModel

TOKENIZER = AutoTokenizer.from_pretrained("distilgpt2")


MODELS = {
    "gpt2": GPT2LMHeadModel,
    "gpt2background": GPT2LMHeadModelBackground,
    "gpt2odin": GPT2OdinModel,
}

DATA_SPLITS = OrderedDict(
    [
        ("train", "./data/encoded_training"),
        ("validation_date_in_domain_in", "./data/encoded_validation_date_in_domain_in"),
        (
            "validation_date_in_domain_out",
            "./data/encoded_validation_date_in_domain_out",
        ),
        (
            "validation_date_out_domain_in",
            "./data/encoded_validation_date_out_domain_in",
        ),
        (
            "validation_date_out_domain_out",
            "./data/encoded_validation_date_out_domain_out",
        ),
    ]
)

LOSS_LOG_TEMPLATE = {
    "train_loss": None,
    "validation_date_in_domain_in_loss": None,
    "validation_date_in_domain_out_loss": None,
    "validation_date_out_domain_in_loss": None,
    "validation_date_out_domain_out_loss": None,
    "epoch": None,
    "batch_idx": None,
}


def _collate_batch(examples):
    """Collate `examples` into a batch, using the information in `tokenizer` for padding if necessary."""
    # Tensorize if necessary.
    if isinstance(examples[0], (list, tuple)):
        examples = [torch.tensor(e, dtype=torch.long) for e in examples]

    # Check if padding is necessary.
    length_of_first = examples[0].size(0)
    are_tensors_same_length = all(x.size(0) == length_of_first for x in examples)
    if are_tensors_same_length:
        return torch.stack(examples, dim=0)

    # Creating the full tensor and filling it with our data.
    max_length = max(x.size(0) for x in examples)
    result = examples[0].new_full([len(examples), max_length], 0)
    for i, example in enumerate(examples):
        result[i, : example.shape[0]] = example
    return result


def _collate_fn(examples):
    result = dict()

    for example in examples:
        for key, val in example.items():
            if key not in result:
                result[key] = list()

            result[key].append(val)

    return {
        "attention_mask": _collate_batch(result["attention_mask"]),
        "input_ids": _collate_batch(result["input_ids"]),
    }


def _get_data_loaders(batch_size: int):
    data_loaders = list()

    for split_name, split_path in DATA_SPLITS.items():

        if "validation" in split_name:
            split_datset = load_from_disk(
                split_path,
                keep_in_memory=False,
            )
        else:
            split_datset = load_from_disk(split_path, keep_in_memory=False)

        split_dataloader = DataLoader(
            split_datset,
            batch_size=batch_size,
            shuffle="train" in split_name,
            pin_memory=True,
            collate_fn=_collate_fn,
            num_workers=4,
            prefetch_factor=2,
        )
        data_loaders.append(split_dataloader)

    return tuple(data_loaders)


def train_model(
    model_name: str,
    batch_size: int,
    cuda: bool,
    epochs: int,
    log_wandb: bool,
    gradient_accumulation: int,
    use_amp: bool,
    train_ratio_perc: float,
    eval_ratio_perc: float,
):
    torch.manual_seed(42)
    np.random.seed(0)

    if cuda:
        device = torch.device("cuda")
    else:
        device = torch.device("cpu")

    (
        train_dataloader,
        val_da_in_do_in_dataloader,
        val_da_in_do_out_dataloader,
        val_da_out_do_in_dataloader,
        val_da_out_do_out_dataloader,
    ) = _get_data_loaders(batch_size=batch_size)

    val_dataloaders = OrderedDict(
        [
            ("validation_date_in_domain_in", val_da_in_do_in_dataloader),
            ("validation_date_in_domain_out", val_da_in_do_out_dataloader),
            ("validation_date_out_domain_in", val_da_out_do_in_dataloader),
            ("validation_date_out_domain_out", val_da_out_do_out_dataloader),
        ]
    )

    model: torch.nn.Module = MODELS[model_name](use_amp=use_amp).to(device)

    num_train_steps = int(
        round(len(train_dataloader) * train_ratio_perc / 100)
    )

    num_train_update_steps = int(round(num_train_steps / gradient_accumulation))

    optimizer = AdamW(params=model.parameters())
    scheduler = get_linear_schedule_with_warmup(
        optimizer=optimizer,
        num_warmup_steps=int(round(0.2 * num_train_update_steps)),
        num_training_steps=num_train_update_steps,
    )
    if use_amp:
        scaler = GradScaler()

    if log_wandb:
        wandb.init(project="ood_dbdc")
        wandb.config.model_name = model_name
        wandb.config.cuda = cuda
        wandb.config.dataloader = train_dataloader.__dict__
        wandb.config.model_config = model.transformer.config.__dict__

    for epoch in range(epochs):
        model = model.train()
        optimizer.zero_grad()

        for batch_step, batch in enumerate(
            tqdm(train_dataloader, desc="Training"), start=1
        ):
            batch = {key: val.to(device) for key, val in batch.items()}
            batch["labels"] = batch["input_ids"]

            output = model(**batch)

            loss = output[0]
            loss = loss / gradient_accumulation
            if use_amp:
                scaler.scale(loss).backward()
            else:
                loss.backward()

            if batch_step % gradient_accumulation == 0:
                if use_amp:
                    scaler.unscale_(optimizer)

                    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)

                    scaler.step(optimizer)
                    scaler.update()
                else:
                    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
                    optimizer.step()

                optimizer.zero_grad()
                scheduler.step()

                if log_wandb:
                    loss = loss.detach().cpu().item()

                    log_line = LOSS_LOG_TEMPLATE.copy()
                    log_line["train_loss"] = loss
                    log_line["epoch"] = epoch
                    log_line["batch_idx"] = batch_step

                    wandb.log(log_line)

                if batch_step >= num_train_steps:
                    break

        optimizer.zero_grad()
        model = model.eval()

        torch.save(model, f"./{model_name}_{epoch}.pth")

        with torch.no_grad():
            for val_name, val_dataloader in val_dataloaders.items():
                num_val_steps = int(round(len(val_dataloader) * eval_ratio_perc / 100))
                for batch_step, batch in enumerate(
                    tqdm(val_dataloader, desc=f"{val_name}"), start=1
                ):
                    batch = {key: val.to(device) for key, val in batch.items()}
                    batch["labels"] = batch["input_ids"]

                    output = model(**batch)

                    if log_wandb:
                        loss = output[0].detach().cpu().item()

                        log_line = LOSS_LOG_TEMPLATE.copy()
                        log_line[f"{val_name}_loss"] = loss
                        log_line["epoch"] = epoch
                        log_line["batch_idx"] = batch_step

                        wandb.log(log_line)

                    if batch_step >= num_val_steps:
                        break
