diff --git a/Stable3DGen/hi3dgen/pipelines/hi3dgen.py b/Stable3DGen/hi3dgen/pipelines/hi3dgen.py index 5415241..2e9f06d 100755 --- a/Stable3DGen/hi3dgen/pipelines/hi3dgen.py +++ b/Stable3DGen/hi3dgen/pipelines/hi3dgen.py @@ -31,6 +31,7 @@ # This modified file is released under the same license. from typing import * from contextlib import contextmanager +import os import torch import torch.nn as nn import torch.nn.functional as F @@ -199,7 +200,7 @@ class Hi3DGenPipeline(Pipeline): """Lazy loading of the BiRefNet model""" from transformers import AutoImageProcessor, Mask2FormerForUniversalSegmentation, AutoModelForImageSegmentation self.birefnet_model = AutoModelForImageSegmentation.from_pretrained( - 'weights/BiRefNet', + os.path.join(os.path.dirname(os.path.abspath(__file__)), '../../..', 'weights', 'BiRefNet'), trust_remote_code=True ).to(self.device) self.birefnet_model.eval() diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..500956e --- /dev/null +++ b/__init__.py @@ -0,0 +1,15 @@ +from .stable_3d import Stable3DGenerate3D, Stable3DPreprocessMesh + + +NODE_CLASS_MAPPINGS = { + "Stable3DGenerate3D": Stable3DGenerate3D, + "Stable3DPreprocessMesh": Stable3DPreprocessMesh +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "Stable3DGenerate3D": "Stable-3D Generate 3D", + "Stable3DPreprocessMesh": "Stable-3D Preprocess Mesh" +} + +__all__ = [NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS] diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..3e70a0c --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,25 @@ +[project] +dependencies = [ + "diffusers~=0.28", + "accelerate~=1.9", + "triton~=3.2", + "kornia~=0.8", + "timm~=0.6", + "transformers~=4.46", + "trimesh~=4.7", + "scikit-image~=0.25", +] + + +name = "ComfyUI-Stable3DGen" +description = "A ComfyUI custom node to generate 3D assets using Stable3D" +version = "1.0.0" +license = { file = "LICENSE" } + +[project.urls] +Repository = "https://github.com/lerignoux/ComfyUI-Stable3DGen.git" + +[tool.comfy] +PublisherId = "lerignoux" +DisplayName = "ComfyUI Stable3DGen" +Icon = "https://avatars.githubusercontent.com/u/171443259?s=48&v=4" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..c1f8061 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,20 @@ +--extra-index-url https://download.pytorch.org/whl/cu124 +# For StableNormal +diffusers~=0.28 +accelerate~=1.9 +triton + +# For BirefNet +kornia~=0.8 +timm~=0.6 +transformers~=4.46 + +trimesh~=4.7 +scikit-image~=0.25 +xformers~=0.0 + +torch==2.5.1 +torchaudio==2.5.1 +torchvision==0.20.1 +torchsde +spconv diff --git a/stable_3d.py b/stable_3d.py new file mode 100644 index 0000000..0d701b6 --- /dev/null +++ b/stable_3d.py @@ -0,0 +1,263 @@ +import datetime +import io +import json +import logging +import os +import numpy +import sys +import torch +import trimesh +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 +TMP_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'tmp') +WEIGHTS_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'weights') +os.makedirs(TMP_DIR, exist_ok=True) +os.makedirs(WEIGHTS_DIR, exist_ok=True) + +# Initialize normal predictor +""" +predictor_model = os.path.join(torch.hub.get_dir(), 'hugoycj_StableNormal_main') +log.info(f"Loading torch predictor model: {predictor_model}") +try: + normal_predictor = torch.hub.load( + predictor_model, + "StableNormal_turbo", + yoso_version='yoso-normal-v1-8-1', + source='local', + local_cache_dir='./weights', + pretrained=True + ) +except Exception as e: + new_model = "hugoycj/StableNormal" + log.info(f"Failed loading local torch {predictor_model} downloading {new_model}, {e}") + normal_predictor = torch.hub.load( + "hugoycj/StableNormal", + "StableNormal_turbo", + trust_repo=True, + yoso_version='yoso-normal-v1-8-1', + local_cache_dir='./weights' + ) +""" +# Loads model to ~/.cache/torch/hub/ +normal_predictor = torch.hub.load("Stable-X/StableNormal", "StableNormal_turbo", trust_repo=True) + + +def cache_weights(weights_dir: str) -> dict: + """ + 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", + ] + cached_paths = {} + 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]) + if os.path.exists(local_path): + log.info(f"Already cached at: {local_path}") + cached_paths[model_id] = 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}") + + # torch.hub.load('facebookresearch/dinov2', name, pretrained=True) + + return cached_paths + + +cache_weights(WEIGHTS_DIR) + +class Stable3DGenerate3D: + """ + A node to generate a Stable3D asset + """ + def __init__(self): + self.output_dir = folder_paths.get_output_directory() + self.temp_dir = folder_paths.get_temp_directory() + self.compress_level = 4 + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "seed": ( + "INT", + { + "tooltip": "The generation seed" + } + ), + "ss_guidance_strength": ( + "INT", + { + "default": 3, + "tooltip": "the titles for each slide." + } + ), + "ss_sampling_steps": ( + "INT", + { + "default": 50, + "step": 1, + "tooltip": "" + } + ), + "slat_guidance_strength": ( + "INT", + { + "default": 3, + "tooltip": "" + } + ), + "slat_sampling_steps": ( + "INT", + { + "default": 6, + "tooltip": "" + } + ) + }, + } + + CATEGORY = "stable_3d_gen" + DESCRIPTION = "Generates a Stable 3D Gen mesh_prompt from an imput image" + FUNCTION = "generate_3d" + INPUT_IS_LIST = False + OUTPUT_NODE = True + RETURN_NAMES = ("mesh_file_path",) + RETURN_TYPES = ("STRING",) + + def save_3d_asset(self, generated_mesh, filename=None): + """ + Save the 3d asset file to user output directory + """ + output_id = datetime.datetime.now().strftime("%Y-%m-%d-%H%M%S") + if filename is None: + filename = f"{output_id}_mesh.glb" + if '.glb' not in filename: + filename = f"{filename}.glb" + mesh_path = os.path.join(self.output_directory, filename) + + trimesh_mesh = generated_mesh.to_trimesh(transform_pose=True) + trimesh_mesh.export(mesh_path) + + return mesh_path + + def generate_3d( + self, + image, + seed=-1, + ss_guidance_strength=3, + ss_sampling_steps=50, + slat_guidance_strength=3, + slat_sampling_steps=6 + ): + if image is None: + return None, None, None + + if seed == -1: + seed = numpy.random.randint(0, MAX_SEED) + + hi3dgen_pipeline = Hi3DGenPipeline.from_pretrained("custom_nodes/ComfyUI-Stable3DGen/weights/trellis-normal-v0-1") + hi3dgen_pipeline.cuda() + + image = torch.rand(1, 512, 512, 3) # Example tensor + numpy_image = image.squeeze(0).cpu().numpy() + numpy_image_scaled = numpy.clip(numpy_image * 255, 0, 255).astype(numpy.uint8) + pil_image = Image.fromarray(numpy_image_scaled) + + # FIXME We should properly handle batch here. + image = hi3dgen_pipeline.preprocess_image(pil_image, resolution=512) + normal_image = normal_predictor(pil_image, resolution=512, match_input_resolution=True, data_type='object') + + outputs = hi3dgen_pipeline.run( + normal_image, + seed=seed, + formats=["mesh",], + preprocess_image=False, + sparse_structure_sampler_params={ + "steps": ss_sampling_steps, + "cfg_strength": ss_guidance_strength, + }, + slat_sampler_params={ + "steps": slat_sampling_steps, + "cfg_strength": slat_guidance_strength, + }, + ) + generated_mesh = outputs['mesh'][0] + + # Save outputs + import datetime + output_id = datetime.datetime.now().strftime("%Y%m%d%H%M%S") + os.makedirs(os.path.join(TMP_DIR, output_id), exist_ok=True) + mesh_path = f"{TMP_DIR}/{output_id}/mesh.glb" + + # Export mesh + trimesh_mesh = generated_mesh.to_trimesh(transform_pose=True) + + trimesh_mesh.export(mesh_path) + + return normal_image, mesh_path, mesh_path + + @classmethod + def IS_CHANGED(s, images, seed, ss_guidance_strength, ss_sampling_steps, slat_guidance_strength, slat_sampling_steps): + return float("NaN") + + +class Stable3DPreprocessMesh: + """ + A node to generate a glb 3d Object from the Stable3D asset + """ + def __init__(self): + self.output_dir = folder_paths.get_output_directory() + self.temp_dir = folder_paths.get_temp_directory() + self.compress_level = 4 + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mesh_prompt": ("MESH_PROMPT", { + "multiline": True, + "default": "Mesh prompt", + "tooltip": "The prompt to generate the mesh" + }) + }, + } + + CATEGORY = "stable_3d_gen" + DESCRIPTION = "Generates a glb 3d Object from the Stable 3D mesh prompt" + FUNCTION = "preprocess_mesh" + INPUT_IS_LIST = True + OUTPUT_NODE = True + RETURN_NAMES = ("filename",) + RETURN_TYPES = ("STRING",) + + def preprocess_mesh(self, mesh_prompt): + print("Processing mesh") + mesh_file = f"{mesh_prompt}.glb" + trimesh_mesh = trimesh.load_mesh(mesh_prompt) + trimesh_mesh.export(mesh_file) + return mesh_file