diff --git a/nodes/deep_bump.py b/nodes/deep_bump.py index ff88337..7fae9a4 100644 --- a/nodes/deep_bump.py +++ b/nodes/deep_bump.py @@ -2,13 +2,16 @@ import tempfile from pathlib import Path import numpy as np + +# torch must be imported prior to onnx for the CUDAProvider. +import torch # isort:skip import onnxruntime as ort -import torch from PIL import Image from ..errors import ModelNotFound from ..log import mklog from ..utils import ( + download_model, get_model_path, tensor2pil, tiles_infer, @@ -23,7 +26,12 @@ log = mklog(__name__) # - COLOR to NORMALS def color_to_normals( - color_img, overlap, progress_callback, *, save_temp=False + color_img, + overlap, + progress_callback, + *, + save_temp=False, + auto_download=False, ): """Compute a normal map from the given color map. @@ -67,7 +75,13 @@ def color_to_normals( log.debug("DeepBump Color → Normals : loading model") model = get_model_path("deepbump", "deepbump256.onnx") if not model or not model.exists(): - raise ModelNotFound(f"deepbump ({model})") + if not auto_download: + raise ModelNotFound(f"deepbump ({model})") + log.debug("Downloading models...") + download_model( + "https://github.com/HugoTini/DeepBump/raw/master/deepbump256.onnx", + "deepbump", + ) providers = [ "TensorrtExecutionProvider", @@ -351,6 +365,9 @@ class MTB_DeepBump: ), "normals_to_height_seamless": ("BOOLEAN", {"default": True}), }, + "optional": { + "auto_download": ("BOOLEAN", {"default": True}), + }, } RETURN_TYPES = ("IMAGE",) @@ -366,6 +383,7 @@ class MTB_DeepBump: color_to_normals_overlap="SMALL", normals_to_curvature_blur_radius="SMALL", normals_to_height_seamless=True, + auto_download=False, ): images = tensor2pil(image) out_images = [] @@ -380,7 +398,10 @@ class MTB_DeepBump: # Apply processing if mode == "Color to Normals": out_img = color_to_normals( - in_img, color_to_normals_overlap, None + in_img, + color_to_normals_overlap, + None, + auto_download=auto_download, ) if mode == "Normals to Curvature": out_img = normals_to_curvature( diff --git a/nodes/graph_utils.py b/nodes/graph_utils.py index b1a444c..95e5b82 100644 --- a/nodes/graph_utils.py +++ b/nodes/graph_utils.py @@ -758,7 +758,7 @@ class MTB_TensorOps: RETURN_TYPES = ("IMAGE",) FUNCTION = "apply" - CATEGORY = "tensor_ops" + CATEGORY = "mtb/tensor_ops" def apply( self, diff --git a/utils.py b/utils.py index 6083219..b041391 100644 --- a/utils.py +++ b/utils.py @@ -15,7 +15,9 @@ from enum import Enum from functools import reduce from pathlib import Path from typing import TypeVar +from urllib.parse import urlparse +import comfy.utils import folder_paths import numpy as np import numpy.typing as npt @@ -858,6 +860,62 @@ def tiles_split(img, tile_size, stride_size): # region MODEL Utilities + + +def download_model(model_url: str, destination: str): + if isinstance(model_url, list): + for url in model_url: + download_model(url, destination) + return + + filename = Path(urlparse(model_url).path).name + + if "drive.google.com" in model_url: + try: + import gdown + except ImportError: + log.info("Installing gdown") + subprocess.check_call( + [ + sys.executable, + "-m", + "pip", + "install", + "gdown", + ] + ) + import gdown + + if "/folders/" in model_url: + # download folder + try: + gdown.download_folder( + model_url, output=destination, resume=True + ) + except TypeError: + gdown.download_folder(model_url, output=destination) + + return + # download from google drive + gdown.download(model_url, destination, quiet=False, resume=True) + return True + response = requests.get(model_url, stream=True) + total_size = int(response.headers.get("content-length", 0)) + + destination_path = get_model_path(destination, filename) + destination_path.parent.mkdir(exist_ok=True) + + pbar = comfy.utils.ProgressBar(total_size) + with open(destination_path, "wb") as file: + for data in response.iter_content(chunk_size=4096): + file.write(data) + pbar.update(len(data)) + + log.info( + f"Downloaded model from {model_url} to {destination_path}", + ) + + def download_antelopev2(): antelopev2_url = ( "https://drive.google.com/uc?id=18wEUfMNohBJ4K3Ly5wpTejPfDzp-8fI8"