From 07b0ec71566ab67df41207a5b3d5599275c5e211 Mon Sep 17 00:00:00 2001 From: Laurent Erignoux Date: Thu, 31 Jul 2025 09:34:55 +0800 Subject: [PATCH] Moving models download to a dedicated node. --- __init__.py | 8 ++++-- stable_3d.py | 80 +++++++++++++++++++++++++++++++++++++++------------- 2 files changed, 66 insertions(+), 22 deletions(-) diff --git a/__init__.py b/__init__.py index 6e1aee4..b028d0a 100644 --- a/__init__.py +++ b/__init__.py @@ -1,13 +1,15 @@ -from .stable_3d import Stable3DGenerate3D +from .stable_3d import Stable3DGenerate3D, Stable3DLoadModels NODE_CLASS_MAPPINGS = { - "Stable3DGenerate3D": Stable3DGenerate3D + "Stable3DGenerate3D": Stable3DGenerate3D, + "Stable3DLoadModels": Stable3DLoadModels } # A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { - "Stable3DGenerate3D": "Stable-3D Generate 3D" + "Stable3DGenerate3D": "Stable-3D Generate 3D", + "Stable3DLoadModels": "Stable-3D Load Models" } __all__ = [NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS] diff --git a/stable_3d.py b/stable_3d.py index 5b62e55..cb19c03 100644 --- a/stable_3d.py +++ b/stable_3d.py @@ -7,21 +7,19 @@ import numpy import sys import torch import trimesh +from huggingface_hub import snapshot_download from PIL import Image from PIL.PngImagePlugin import PngInfo import folder_paths sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'Stable3DGen')) - from hi3dgen.pipelines import Hi3DGenPipeline log = logging.getLogger(__name__) MAX_SEED = numpy.iinfo(numpy.int32).max -WEIGHTS_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'weights') -os.makedirs(WEIGHTS_DIR, exist_ok=True) # Initialize normal predictor """ @@ -47,43 +45,81 @@ except Exception as e: local_cache_dir='./weights' ) """ -# Loads model to ~/.cache/torch/hub/ -normal_predictor = torch.hub.load("Stable-X/StableNormal", "StableNormal_turbo", trust_repo=True) + +class Stable3DLoadModels: + """ + A node to load the models necessary for Stable3D + Node will download the models from huggingface or torch if missing. + """ + def __init__(self): + self.weights_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'weights') + os.makedirs(self.weights_dir, exist_ok=True) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "trellis_model": ( + "STRING", + { + "tooltip": "The trellis model to use", + "default": "Stable-X/trellis-normal-v0-1" + } + ), + "normal_model": ( + "STRING", + { + "tooltip": "The normal generation model", + "default": "Stable-X/yoso-normal-v1-8-1" + } + ), + "birefnet_model": ( + "STRING", + { + "default": "ZhengPeng7/BiRefNet", + "tooltip": "the Birefnet model." + } + ) + }, + } + + CATEGORY = "stable_3d_gen" + DESCRIPTION = "Load the models necessary for Stable3D Gen" + FUNCTION = "load_models" + INPUT_IS_LIST = False + OUTPUT_NODE = False + RETURN_NAMES = ("Trellis model", "Normal predictor", "Birefnet model") + RETURN_TYPES = ("TRELLIS_MODEL", "STABLE3D_NORMAL", "STABLE3D_BIREFNET") -def cache_weights(weights_dir: str) -> dict: +def load_models(self, trellis_model, normal_model, birefnet_model): """ Load weights locally if missing. Needs to be adapted to match ComfyUI Models storage """ - import os - from huggingface_hub import snapshot_download - os.makedirs(weights_dir, exist_ok=True) - model_ids = [ - "Stable-X/trellis-normal-v0-1", - "Stable-X/yoso-normal-v1-8-1", - "ZhengPeng7/BiRefNet", - ] + model_ids = [trellis_model, normal_model, birefnet_model] cached_paths = {} + loaded_models = [] for model_id in model_ids: log.info(f"Caching weights for: {model_id}") - local_path = os.path.join(weights_dir, model_id.split("/")[-1]) + local_path = os.path.join(self.weights_dir, model_id.split("/")[-1]) if os.path.exists(local_path): log.info(f"Already cached at: {local_path}") cached_paths[model_id] = local_path + loaded_models.append(local_path) continue log.info(f"Downloading and caching model: {model_id}") local_path = snapshot_download(repo_id=model_id, local_dir=os.path.join(weights_dir, model_id.split("/")[-1]), force_download=False) cached_paths[model_id] = local_path log.info(f"Cached at: {local_path}") + # Loads model to ~/.cache/torch/hub/ + normal_predictor = torch.hub.load("Stable-X/StableNormal", "StableNormal_turbo", trust_repo=True) + # torch.hub.load('facebookresearch/dinov2', name, pretrained=True) - return cached_paths - - -cache_weights(WEIGHTS_DIR) + return (loaded_models[0], normal_predictor, loaded_models[2]) class Stable3DGenerate3D: """ @@ -99,6 +135,9 @@ class Stable3DGenerate3D: def INPUT_TYPES(s): return { "required": { + "trellis_model": ("TRELLIS_MODEL", ), + "normal_predictor": ("STABLE3D_NORMAL", ), + "birefnet_model": ("STABLE3D_BIREFNET", ), "image": ("IMAGE",), "seed": ( "INT", @@ -164,6 +203,9 @@ class Stable3DGenerate3D: def generate_3d( self, + trellis_model, + normal_predictor, + birefnet_model, image, seed=-1, ss_guidance_strength=3,