import torch
import torch.nn as nn
from torch.cuda.amp import autocast
from torch.nn import CrossEntropyLoss
from transformers import AutoModelWithLMHead


class GPT2OdinModel(nn.Module):
    def __init__(self, use_amp: bool = False):
        super(GPT2OdinModel, self).__init__()
        self.transformer = AutoModelWithLMHead.from_pretrained("distilgpt2").base_model
        config = self.transformer.config
        self.linear_g_component = nn.Linear(
            in_features=config.n_embd, out_features=1, bias=True
        )
        self.linear_h_component = nn.Linear(
            in_features=config.n_embd, out_features=config.vocab_size, bias=True
        )
        self.use_amp = use_amp

    def forward(
        self,
        input_ids,
        attention_mask,
        labels=None,
    ):
        with autocast(self.use_amp):
            hidden_states = self.transformer(
                input_ids=input_ids, attention_mask=attention_mask
            ).last_hidden_state

            g_logits = torch.sigmoid(self.linear_g_component(hidden_states))
            h_prod = self.linear_h_component(hidden_states)
            lm_logits = h_prod / g_logits

            loss = None
            if labels is not None:
                # Shift so that tokens < n predict n
                shift_logits = lm_logits[..., :-1, :].contiguous()
                shift_labels = labels[..., 1:].contiguous()
                # Flatten the tokens
                loss_fct = CrossEntropyLoss()
                loss = loss_fct(
                    shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)
                )

            if loss:
                output = (
                    loss,
                    lm_logits,
                )
            else:
                output = (lm_logits, h_prod, g_logits)

            return output
