from transformers import AutoTokenizer, T5EncoderModel, T5Config import torch import torch.nn as nn from typing import Optional class Embedder(nn.Module): def __init__(self, max_length=330, load_pretrained_weights=True, frozen=True): super().__init__() self.tokenizer = AutoTokenizer.from_pretrained("byt5-small", use_fast=False) self.max_length = max_length if load_pretrained_weights: self.encoder = T5EncoderModel.from_pretrained("google/byt5-small") else: self.encoder = T5EncoderModel(T5Config.from_pretrained("byt5-small")) self.encoder.eval() if frozen: for p in self.encoder.parameters(): p.requires_grad = False def tokenize(self, sentences): assert all(isinstance(s, str) for s in sentences) tokenized = self.tokenizer( [s.strip() for s in sentences], padding=True, truncation=True, max_length=self.max_length, return_tensors="pt", ) return tokenized["input_ids"], tokenized["attention_mask"] def detokenize(self, token_ids): return self.tokenizer.batch_decode(token_ids, skip_special_tokens=True) def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.FloatTensor] = None, ): output = self.encoder(input_ids=input_ids, attention_mask=attention_mask) return output["last_hidden_state"]