Cache empty text embedding (#7)

* Cache empty text embedding to reduce a bit of time and avoid CLIP loading

* It is only 9MB so moving to global won't affect RAM much

* Complete
This commit is contained in:
Fannovel16
2023-12-14 17:31:13 +02:00
committed by GitHub
parent 3fb3d6fd69
commit 1a5017007a
2 changed files with 3 additions and 1 deletions
Binary file not shown.
+3 -1
View File
@@ -21,6 +21,8 @@ def colorizedepth(depth_map, colorize_method):
depth_colored_hwc = chw2hwc(depth_colored) depth_colored_hwc = chw2hwc(depth_colored)
return depth_colored_hwc return depth_colored_hwc
empty_text_embed = torch.load(os.path.join(__file__, '..', "empty_text_embed.pt"), map_location="cpu")
class MarigoldDepthEstimation: class MarigoldDepthEstimation:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -78,7 +80,7 @@ class MarigoldDepthEstimation:
if checkpoint_path is None: if checkpoint_path is None:
raise FileNotFoundError("No checkpoint directory found.") raise FileNotFoundError("No checkpoint directory found.")
self.marigold_pipeline = MarigoldPipeline.from_pretrained(checkpoint_path, enable_xformers=False) self.marigold_pipeline = MarigoldPipeline.from_pretrained(checkpoint_path, enable_xformers=False, empty_text_embed=empty_text_embed)
self.marigold_pipeline = self.marigold_pipeline.to(device).half() self.marigold_pipeline = self.marigold_pipeline.to(device).half()
self.marigold_pipeline.unet.eval() # Set the model to evaluation mode self.marigold_pipeline.unet.eval() # Set the model to evaluation mode