MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / FrozenT5Embedder

Class FrozenT5Embedder

ldm/modules/encoders/modules.py:58–85  ·  view source on GitHub ↗

Uses the T5 transformer encoder for text

Source from the content-addressed store, hash-verified

56
57
58class 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
88class FrozenCLIPEmbedder(AbstractEncoder):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected