Method
__init__
(self, version="google/t5-v1_1-large", device="cuda", max_length=77, freeze=True)
Source from the content-addressed store, hash-verified
| 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() |
Callers
nothing calls this directly
Tested by
no test coverage detected