diff --git a/nodes.py b/nodes.py index 843b74c..ec857b8 100644 --- a/nodes.py +++ b/nodes.py @@ -83,8 +83,13 @@ class MarigoldDepthEstimation: break if checkpoint_path is None: - raise FileNotFoundError("No checkpoint directory found.") - + try: + from huggingface_hub import snapshot_download + checkpoint_path = os.path.join(script_directory, "../../models/diffusers/Marigold") + snapshot_download(repo_id="Bingxin/Marigold", ignore_patterns=["*.bin"], local_dir=checkpoint_path, local_dir_use_symlinks=False) + + except: + raise FileNotFoundError("No checkpoint directory found.") 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() if use_fp16 else self.marigold_pipeline.to(device) self.marigold_pipeline.unet.eval() # Set the model to evaluation mode diff --git a/requirements.txt b/requirements.txt index 10b21e0..6c8ce53 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,5 @@ diffusers>=0.20.1 matplotlib scipy torch>=2.0.1 -transformers>=4.32.1 \ No newline at end of file +transformers>=4.32.1 +huggingface-hub \ No newline at end of file