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"]
