diff --git a/ldm/modules/embedding_manager.py b/ldm/modules/embedding_manager.py index 677bc4ad3a4..128b575e9bb 100644 --- a/ldm/modules/embedding_manager.py +++ b/ldm/modules/embedding_manager.py @@ -215,11 +215,13 @@ def save(self, ckpt_path): ckpt_path, ) - def load(self, ckpt_path): + def load(self, ckpt_path, full=True): ckpt = torch.load(ckpt_path, map_location='cpu') - - self.string_to_token_dict = ckpt['string_to_token'] - self.string_to_param_dict = ckpt['string_to_param'] + self.string_to_token_dict = ckpt["string_to_token"] + self.string_to_param_dict = ckpt["string_to_param"] + if not full: + for key, value in self.string_to_param_dict.items(): + self.string_to_param_dict[key] = torch.nn.Parameter(value.half()) def get_embedding_norms_squared(self): all_params = torch.cat( diff --git a/ldm/simplet2i.py b/ldm/simplet2i.py index ffec5fda2b4..8c793a2b92a 100644 --- a/ldm/simplet2i.py +++ b/ldm/simplet2i.py @@ -488,7 +488,7 @@ def load_model(self): ) model = self._load_model_from_config(config, self.weights) if self.embedding_path is not None: - model.embedding_manager.load(self.embedding_path) + model.embedding_manager.load(self.embedding_path, self.full_precision) self.model = model.to(self.device) # model.to doesn't change the cond_stage_model.device used to move the tokenizer output, so set it here self.model.cond_stage_model.device = self.device