from torch import nn


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):
        if self.model_name == "xlm-roberta":
            from transformers import AutoTokenizer, XLMRobertaModel
            self.tokenizer = AutoTokenizer.from_pretrained("xlm-roberta-base")
            self.text_encoder = XLMRobertaModel.from_pretrained("xlm-roberta-base")
        else:
            raise ValueError("%s is not supported yet." % self.model_name)

    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