Uses the T5 transformer encoder for text
| 56 | |
| 57 | |
| 58 | class FrozenT5Embedder(AbstractEncoder): |
| 59 | """Uses the T5 transformer encoder for text""" |
| 60 | def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77, freeze=True): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl |
| 61 | super().__init__() |
| 62 | self.tokenizer = T5Tokenizer.from_pretrained(version) |
| 63 | self.transformer = T5EncoderModel.from_pretrained(version) |
| 64 | self.device = device |
| 65 | self.max_length = max_length # TODO: typical value? |
| 66 | if freeze: |
| 67 | self.freeze() |
| 68 | |
| 69 | def freeze(self): |
| 70 | self.transformer = self.transformer.eval() |
| 71 | #self.train = disabled_train |
| 72 | for param in self.parameters(): |
| 73 | param.requires_grad = False |
| 74 | |
| 75 | def forward(self, text): |
| 76 | batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True, |
| 77 | return_overflowing_tokens=False, padding="max_length", return_tensors="pt") |
| 78 | tokens = batch_encoding["input_ids"].to(self.device) |
| 79 | outputs = self.transformer(input_ids=tokens) |
| 80 | |
| 81 | z = outputs.last_hidden_state |
| 82 | return z |
| 83 | |
| 84 | def encode(self, text): |
| 85 | return self(text) |
| 86 | |
| 87 | |
| 88 | class FrozenCLIPEmbedder(AbstractEncoder): |