NguyenDinhHieu commited on
Commit
64033cd
·
verified ·
1 Parent(s): 6b10f1e

Update ldm/modules/encoders/modules.py

Browse files
Files changed (1) hide show
  1. ldm/modules/encoders/modules.py +5 -4
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
- from open_clip import get_tokenizer
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 = get_tokenizer(arch)
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 = open_clip.tokenize(text)
178
- z = self.encode_with_transformer(tokens.to(self.device))
 
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):