import bz2
import itertools
import pickle


import requests
import torch
from tqdm.autonotebook import tqdm
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    GPT2LMHeadModel,
    GPT2Tokenizer,
)

convai1_data = requests.get("http://convai.io/2017/data/train_full.json").json()
print(len(convai1_data))
convai2_data = requests.get(
    "http://convai.io/data/summer_wild_evaluation_dialogs.json"
).json()
print(len(convai2_data))

for dial in tqdm(convai1_data):
    quality = sum(
        [participant_score["quality"] for participant_score in dial["evaluation"]]
    ) / len(dial["evaluation"])
    dial["quality"] = quality
    utterances = [thread_line["text"] for thread_line in dial["thread"]]
    dial["utterances"] = utterances
    dial["predictions"] = dict()
    dial["id"] = str(dial["dialogId"])

convai1_data = [dial for dial in convai1_data if len(dial["utterances"]) > 2]
print(len(convai1_data))

for dial in tqdm(convai2_data):
    dial["quality"] = dial["eval_score"]
    utterances = [thread_line["text"] for thread_line in dial["dialog"]]
    dial["utterances"] = utterances
    dial["predictions"] = dict()
    dial["id"] = str(dial["dialog_id"])

convai2_data = [dial for dial in convai2_data if len(dial["utterances"]) > 2]
print(len(convai2_data))

convai_data_len = len(convai1_data) + len(convai2_data)


######################################################
### GPT2 scoring function
######################################################


def gpt2_sent_probability(PADDING_TEXT, text, model_tokenizer, model):

    tokenize_text = model_tokenizer.encode(text, add_special_tokens=False)[:512]
    tokenize_input = (
        [model_tokenizer.bos_token_id]
        + model_tokenizer.encode(PADDING_TEXT, add_special_tokens=False)[:510]
        + tokenize_text
        + [model_tokenizer.eos_token_id]
    )
    tokenize_text = tokenize_text + [model_tokenizer.eos_token_id]
    tokenize_text_len = len(tokenize_text)

    tokenize_input = torch.LongTensor(tokenize_input)

    if torch.cuda.is_available():
        tokenize_input = tokenize_input.cuda()

    with torch.no_grad():
        predicted_probs = model(tokenize_input)[0]
        predicted_probs = predicted_probs[-tokenize_text_len:-1]
        predicted_probs = torch.softmax(predicted_probs, dim=-1)

    predicted_probs = predicted_probs.detach().cpu().numpy().tolist()

    sentence_word_probs = list()

    for predicted_prob, token_id in zip(predicted_probs, tokenize_text):
        sentence_word_probs.append(predicted_prob[token_id])

    return sentence_word_probs


######################################################
### Compute GPT2 scores
######################################################


for GPT2_MODEL in tqdm(["gpt2", "gpt2-medium", "gpt2-large"]):

    model_tokenizer = GPT2Tokenizer.from_pretrained(GPT2_MODEL)
    model = GPT2LMHeadModel.from_pretrained(GPT2_MODEL)

    if torch.cuda.is_available():
        model = model.cuda()

    model = model.eval()

    for dial in tqdm(
        itertools.chain(convai1_data, convai2_data),
        total=convai_data_len,
        desc=GPT2_MODEL,
    ):
        utterances = dial["utterances"]

        # utterance pair

        sentences_word_probs = list()

        for u1, u2 in zip(utterances[:-1], utterances[1:]):
            try:
                sentence_word_probs = gpt2_sent_probability(
                    u1, u2, model_tokenizer, model
                )

                sentences_word_probs.append(sentence_word_probs)
            except Exception as ex:
                print(ex)

        dial["predictions"][f"{GPT2_MODEL}_pair_word_probs"] = sentences_word_probs

        # complete context

        sentences_word_probs = list()

        for u_idx in range(1, len(utterances)):
            try:
                context = " ".join(utterances[u_idx]).strip()
                response = utterances[u_idx]
                sentence_word_probs = gpt2_sent_probability(
                    context, response, model_tokenizer, model
                )

                sentences_word_probs.append(sentence_word_probs)
            except Exception as ex:
                print(ex)

        dial["predictions"][f"{GPT2_MODEL}_context_word_probs"] = sentences_word_probs


######################################################
### DialoGPT scoring function
######################################################


def dialogpt_sent_probability(PADDING_TEXT_list, text, model_tokenizer, model):

    tokenize_prefix = list()

    for utterance in PADDING_TEXT_list:
        tokenize_prefix.extend(
            model_tokenizer.encode(utterance, add_special_tokens=False)
        )

    tokenize_prefix = tokenize_prefix[:510]

    tokenize_text = model_tokenizer.encode(text, add_special_tokens=False)[:512]
    tokenize_input = (
        [model_tokenizer.bos_token_id]
        + tokenize_prefix
        + tokenize_text
        + [model_tokenizer.eos_token_id]
    )
    tokenize_text = tokenize_text + [model_tokenizer.eos_token_id]
    tokenize_text_len = len(tokenize_text)

    tokenize_input = torch.LongTensor(tokenize_input)

    if torch.cuda.is_available():
        tokenize_input = tokenize_input.cuda()

    with torch.no_grad():
        predicted_probs = model(tokenize_input)[0]
        predicted_probs = predicted_probs[-tokenize_text_len:-1]
        predicted_probs = torch.softmax(predicted_probs, dim=-1)

    predicted_probs = predicted_probs.detach().cpu().numpy().tolist()

    sentence_word_probs = list()

    for predicted_prob, token_id in zip(predicted_probs, tokenize_text):
        sentence_word_probs.append(predicted_prob[token_id])

    return sentence_word_probs


######################################################
### Compute DialogGPT scores
######################################################


for DIALOGPT_MODEL in tqdm(
    [
        "microsoft/DialoGPT-small",
        "microsoft/DialoGPT-medium",
        "microsoft/DialoGPT-large",
    ]
):
    model_tokenizer = AutoTokenizer.from_pretrained(DIALOGPT_MODEL)
    model = AutoModelForCausalLM.from_pretrained(DIALOGPT_MODEL)

    if torch.cuda.is_available():
        model = model.cuda()

    model = model.eval()

    for dial in tqdm(
        itertools.chain(convai1_data, convai2_data),
        total=convai_data_len,
        desc=DIALOGPT_MODEL,
    ):
        utterances = dial["utterances"]

        # utterance pair

        sentences_word_probs = list()

        for u1, u2 in zip(utterances[:-1], utterances[1:]):
            try:
                sentence_word_probs = dialogpt_sent_probability(
                    [u1], u2, model_tokenizer, model
                )

                sentences_word_probs.append(sentence_word_probs)
            except Exception as ex:
                print(ex)

        dial["predictions"][F"{DIALOGPT_MODEL}_word_probs"] = sentences_word_probs

        # complete context

        sentences_word_probs = list()

        for u_idx in range(1, len(utterances)):
            try:
                context = utterances[:u_idx]
                response = utterances[u_idx]
                sentence_word_probs = gpt2_sent_probability(
                    context, response, model_tokenizer, model
                )

                sentences_word_probs.append(sentence_word_probs)
            except Exception as ex:
                print(ex)

        dial["predictions"][f"{DIALOGPT_MODEL}_context_word_probs"] = sentences_word_probs


with bz2.open("./convai1_results.pickle.bz2", mode="wb") as fout:
    pickle.dump(convai1_data, fout)

with bz2.open("./convai2_results.pickle.bz2", mode="wb") as fout:
    pickle.dump(convai2_data, fout)
