Spaces:
Paused
Paused
Update ldm/modules/encoders/modules.py
Browse files
ldm/modules/encoders/modules.py
CHANGED
|
@@ -4,7 +4,7 @@ from torch.utils.checkpoint import checkpoint
|
|
| 4 |
|
| 5 |
from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
|
| 6 |
|
| 7 |
-
|
| 8 |
from ldm.util import default, count_params
|
| 9 |
from huggingface_hub import hf_hub_download
|
| 10 |
|
|
@@ -151,7 +151,7 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder):
|
|
| 151 |
)
|
| 152 |
|
| 153 |
self.model = open_clip.create_model(arch, pretrained=pretrained_path)
|
| 154 |
-
self.tokenizer =
|
| 155 |
|
| 156 |
if hasattr(self.model, "visual"):
|
| 157 |
delattr(self.model, "visual")
|
|
@@ -174,8 +174,9 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder):
|
|
| 174 |
param.requires_grad = False
|
| 175 |
|
| 176 |
def forward(self, text):
|
| 177 |
-
tokens =
|
| 178 |
-
|
|
|
|
| 179 |
return z
|
| 180 |
|
| 181 |
def encode_with_transformer(self, text):
|
|
|
|
| 4 |
|
| 5 |
from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
|
| 6 |
|
| 7 |
+
import open_clip
|
| 8 |
from ldm.util import default, count_params
|
| 9 |
from huggingface_hub import hf_hub_download
|
| 10 |
|
|
|
|
| 151 |
)
|
| 152 |
|
| 153 |
self.model = open_clip.create_model(arch, pretrained=pretrained_path)
|
| 154 |
+
self.tokenizer = open_clip.tokenizer.SimpleTokenizer()
|
| 155 |
|
| 156 |
if hasattr(self.model, "visual"):
|
| 157 |
delattr(self.model, "visual")
|
|
|
|
| 174 |
param.requires_grad = False
|
| 175 |
|
| 176 |
def forward(self, text):
|
| 177 |
+
tokens = [self.tokenizer.encode(t) for t in text]
|
| 178 |
+
tokens = torch.tensor(tokens).to(self.device)
|
| 179 |
+
z = self.encode_with_transformer(tokens)
|
| 180 |
return z
|
| 181 |
|
| 182 |
def encode_with_transformer(self, text):
|