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