from torch import nn
from transformers import AutoTokenizer, XLMRobertaModel


class TextEncoder(nn.Module):
    """
    Text encoder
    """

    def __init__(
        self,
        model_name="xlm-roberta",
    ):
        super(TextEncoder, self).__init__()

        self.model_name = model_name
        self.get_encoder()

    def get_encoder(self):
        self.tokenizer = AutoTokenizer.from_pretrained("xlm-roberta-base")
        self.text_encoder = XLMRobertaModel.from_pretrained("xlm-roberta-base")

    def get_embeddings(self, text):
        if self.model_name == "xlm-roberta":
            device = self.text_encoder.device
            inputs = self.tokenizer(text, padding=True, return_tensors="pt").to(device)
            outputs = self.text_encoder(**inputs)
            last_hidden_states = outputs.last_hidden_state
            return last_hidden_states

    def forward(self, text):
        emb = self.get_embeddings(text)
        # return emb.mean(dim=1)
        return emb[
            :, 0, :
        ]  # we take the first token to represent the sequence by appending [CLS] token
