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:
Binary file not shown.
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user