import datetime import logging import os import numpy import sys import torch from huggingface_hub import snapshot_download from PIL import Image from transformers import AutoModelForImageSegmentation from comfy.utils import common_upscale import folder_paths sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'Stable3DGen')) from hi3dgen.pipelines import Hi3DGenPipeline # noqa E402 log = logging.getLogger(__name__) MAX_SEED = numpy.iinfo(numpy.int32).max def set_spconv_algo(func): def wrapper(*args, **kwargs): old_spconv_algo = os.environ.get('SPCONV_ALGO') os.environ['SPCONV_ALGO'] = 'native' try: return func(*args, **kwargs) finally: if old_spconv_algo is None: del os.environ['SPCONV_ALGO'] else: os.environ['SPCONV_ALGO'] = old_spconv_algo return wrapper 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.models_path = os.path.join(folder_paths.models_dir, "trellis") folder_paths.add_model_folder_path("trellis", self.models_path) os.makedirs(self.models_path, 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 Background removal model." } ) }, } CATEGORY = "stable_3d_gen" DESCRIPTION = "Download if necessary and load the models necessary for Stable3D Gen preprocessing and generation." FUNCTION = "load_models" INPUT_IS_LIST = False OUTPUT_NODE = False RETURN_NAMES = ("hi3dgen pipeline", "Normal predictor") RETURN_TYPES = ("HI3DGEN_PIPELINE", "STABLE3D_NORMAL") def load_birefnet_model(self, hi3dgen_pipeline, birefnet_model_name): """ Custom birefnet model loader on the hi3dGen pipeline to customize model location """ hi3dgen_pipeline.birefnet_model = AutoModelForImageSegmentation.from_pretrained( os.path.join(self.models_path, birefnet_model_name), trust_remote_code=True ).to(hi3dgen_pipeline.device) hi3dgen_pipeline.birefnet_model.eval() return def download_model(self, model_id): log.info(f"Caching weights for: {model_id}") local_path = os.path.join(self.models_path, model_id) if os.path.exists(local_path): log.info(f"Already cached at: {local_path}") return local_path log.info(f"Downloading and caching model: {model_id} to trellis model folder") local_path = snapshot_download( repo_id=model_id, local_dir=os.path.join(self.models_path, model_id), force_download=False ) return local_path @set_spconv_algo def load_models(self, trellis_model, normal_model, birefnet_model): """ Load weights locally if missing. Models are downloaded in ComfyUI/models/trellis folder torch libraries are downloaded in ~/.cache/torch/hub/ """ model_ids = [trellis_model, normal_model, birefnet_model] loaded_models = [] for model_id in model_ids: loaded_models.append(self.download_model(model_id)) # Loads yoso pedictor model to ~/models/trellis and StableNormal_turbo predictor library to ~/.cache/torch/hub/ yozo_model_folder = os.path.join(self.models_path, "Stable-X") normal_predictor = torch.hub.load( "hugoycj/StableNormal", "StableNormal_turbo", trust_repo=True, yoso_version='yoso-normal-v1-8-1', local_cache_dir=yozo_model_folder ) # download dinov2 feature detection model and library to ~/.cache/torch/hub/ torch.hub.load('facebookresearch/dinov2', 'dinov2_vitl14_reg', pretrained=True) hi3dgen_pipeline = Hi3DGenPipeline.from_pretrained(os.path.join(self.models_path, trellis_model)) hi3dgen_pipeline.cuda() self.load_birefnet_model(hi3dgen_pipeline, birefnet_model) return (hi3dgen_pipeline, normal_predictor) class Stable3DPreprocessImage: """ A node to Preprocess an input image into a normal representation for 3d generation """ def __init__(self): self.temp_directory = folder_paths.get_temp_directory() @classmethod def INPUT_TYPES(s): return { "required": { "hi3dgen_pipeline": ( "HI3DGEN_PIPELINE", {"tooltip": "The hi3dgen pipeline containing the remove background setup (BiRefNet)."} ), "normal_predictor": ( "STABLE3D_NORMAL", {"tooltip": "The normal predictor model to generate the image normal."} ), "images": ("IMAGE",) }, } CATEGORY = "stable_3d_gen" DESCRIPTION = "Preprocess an input image into normal suitable for Stable 3D Generation." FUNCTION = "preprocess_image" INPUT_IS_LIST = False OUTPUT_NODE = True RETURN_NAMES = ("normal_images",) RETURN_TYPES = ("IMAGE",) def pil_to_tensor(self, image): image = numpy.array(image).astype(numpy.float32) / 255.0 image = torch.from_numpy(image)[None,] return image def save_normal_image(self, normal_image): output_id = datetime.datetime.now().strftime("%Y-%m-%d-%H%M%S") filename = f"{output_id}_normal.png" path = os.path.join(self.temp_directory, filename) normal_image.save(path) log.debug(f"normal_image saved as {path}") return path @staticmethod def uniformize_images(images): for index in range(1, len(images)): if images[index].shape[1:] != images[0].shape[1:]: images[index] = common_upscale( images[index].movedim(-1, 1), images[0].shape[2], images[0].shape[1], "bilinear", "center" ).movedim(1, -1) return images def preprocess_image(self, hi3dgen_pipeline, normal_predictor, images): normal_images = [] for (_, img) in enumerate(images): numpy_image = 255. * img.cpu().numpy() numpy_image_scaled = Image.fromarray(numpy.clip(numpy_image, 0, 255).astype(numpy.uint8)) image = hi3dgen_pipeline.preprocess_image(numpy_image_scaled, resolution=1024) normal_image = normal_predictor(image, resolution=768, match_input_resolution=True, data_type='object') self.save_normal_image(normal_image) normal_images.append(self.pil_to_tensor(normal_image)) if len(normal_images) > 1: output_image = torch.cat(self.uniformize_images(normal_images), dim=0) else: output_image = normal_images[0] return (output_image,) class Stable3DGenerate3D: """ A node to generate a Stable3D asset """ def __init__(self): self.output_directory = folder_paths.get_output_directory() self.temp_directory = folder_paths.get_temp_directory() self.compress_level = 4 pass @classmethod def INPUT_TYPES(s): return { "required": { "hi3dgen_pipeline": ("HI3DGEN_PIPELINE", ), "normal_images": ("IMAGE",), "seed": ( "INT", { "default": -1, "min": -1, "max": MAX_SEED, "tooltip": "The generation seed, use -1 for a random seed." } ), "ss_guidance_strength": ( "FLOAT", { "default": 3.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "The sparse structure guidance strength" } ), "ss_sampling_steps": ( "INT", { "default": 50, "min": 1, "max": 50, "step": 1, "tooltip": "The sparse structure sampling steps. increasing the steps increase generation time" } ), "slat_guidance_strength": ( "FLOAT", { "default": 3.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "" } ), "slat_sampling_steps": ( "INT", { "default": 6, "min": 1, "max": 50, "step": 1, "tooltip": "The slat sampling steps. increasing the steps increase generation time" } ) }, } CATEGORY = "stable_3d_gen" DESCRIPTION = "Generates a Stable 3D Gen mesh 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 """ if filename is None: output_id = datetime.datetime.now().strftime("%Y-%m-%d-%H%M%S") filename = f"{output_id}_mesh.glb" if '.glb' not in filename: filename = f"{filename}.glb" mesh_path = os.path.join(self.output_directory, filename) log.info(f"Saving 3d mesh to {mesh_path}.") trimesh_mesh = generated_mesh.to_trimesh(transform_pose=True) trimesh_mesh.export(mesh_path) return mesh_path @set_spconv_algo def generate_3d( self, hi3dgen_pipeline, normal_images, seed=-1, ss_guidance_strength=3.0, ss_sampling_steps=50, slat_guidance_strength=3.0, slat_sampling_steps=6 ): if seed == -1: seed = numpy.random.randint(0, MAX_SEED) log.info("Starting 3d mesh generation.") pil_input = [] for (_, image) in enumerate(normal_images): numpy_image = 255. * image.cpu().numpy() pil_input.append(Image.fromarray(numpy.clip(numpy_image, 0, 255).astype(numpy.uint8))) method = 'run' if len(pil_input) == 1: pil_input = pil_input[0] elif len(pil_input) > 1: method = 'run_multi_image' else: raise ValueError("No image provided to run 3d pipeline.") outputs = getattr(hi3dgen_pipeline, method)( pil_input, 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] saved_path = self.save_3d_asset(generated_mesh) filename = saved_path.split("output/")[1] return (filename,) @classmethod def IS_CHANGED( s, images, seed, ss_guidance_strength, ss_sampling_steps, slat_guidance_strength, slat_sampling_steps ): # FIXME We should properly handle re-generation depending on the input parameters. return float("NaN")