import torch
from torch.cuda.amp import autocast
from torch.distributions.bernoulli import Bernoulli
from torch.nn import Module
from transformers import AutoModelWithLMHead


class GPT2LMHeadModelBackground(Module):
    def __init__(self, mu=0.15, use_amp: bool = False):
        super(GPT2LMHeadModelBackground, self).__init__()
        self.transformer = AutoModelWithLMHead.from_pretrained("distilgpt2")
        self.bernoulli_dist = Bernoulli(torch.tensor([mu]))
        self.use_amp = use_amp

    def forward(
        self,
        input_ids,
        attention_mask,
        labels=None,
    ):
        with autocast(self.use_amp):
            if self.training:
                input_size = input_ids.size()
                pertrubation_mask = self.bernoulli_dist.sample(sample_shape=input_size)
                random_ints = torch.randint(
                    low=0, high=self.transformer.config.vocab_size, size=input_size
                )

                for batch_dim in range(input_size[0]):
                    for sample_dim in range(input_size[1]):
                        if pertrubation_mask[batch_dim, sample_dim] == 0:
                            continue

                        # Very small chance, 1 / self.config.vocab_size,
                        # that the pertrubed input gets the same id,
                        # i.e. it does not change
                        input_ids[batch_dim, sample_dim] = random_ints[
                            batch_dim, sample_dim
                        ]

            return self.transformer(
                input_ids=input_ids, attention_mask=attention_mask, labels=labels
            )
