diff --git a/__init__.py b/__init__.py index b028d0a..81d4c89 100644 --- a/__init__.py +++ b/__init__.py @@ -1,15 +1,17 @@ -from .stable_3d import Stable3DGenerate3D, Stable3DLoadModels +from .stable_3d import Stable3DGenerate3D, Stable3DLoadModels, Stable3DPreprocessImage NODE_CLASS_MAPPINGS = { "Stable3DGenerate3D": Stable3DGenerate3D, - "Stable3DLoadModels": Stable3DLoadModels + "Stable3DLoadModels": Stable3DLoadModels, + "Stable3DPreprocessImage": Stable3DPreprocessImage } # A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { "Stable3DGenerate3D": "Stable-3D Generate 3D", - "Stable3DLoadModels": "Stable-3D Load Models" + "Stable3DLoadModels": "Stable-3D Load Models", + "Stable3DPreprocessImage": "Stable-3D Preprocess Image" } __all__ = [NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS] diff --git a/stable_3d.py b/stable_3d.py index b903bda..c8fadc4 100644 --- a/stable_3d.py +++ b/stable_3d.py @@ -77,7 +77,7 @@ class Stable3DLoadModels: "STRING", { "default": "ZhengPeng7/BiRefNet", - "tooltip": "the Birefnet model." + "tooltip": "the Background removal model." } ) }, @@ -121,6 +121,64 @@ def load_models(self, trellis_model, normal_model, birefnet_model): return (loaded_models[0], normal_predictor, loaded_models[2]) + +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": { + "background_removal_model": ( + "STABLE3D_BIREFNET", + {"tooltip": "The model to remove background (BiRefNet)."} + ), + "normal_predictor": ( + "STABLE3D_NORMAL", + {"tooltip": "The normal predictor model to generate the image normal."} + ), + "image": ("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_image",) + RETURN_TYPES = ("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) + + def preprocess_image(self, background_removal_model, normal_predictor, image): + + 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') + + self.save_normal_image(normal_image) + + return normal_image + + class Stable3DGenerate3D: """ A node to generate a Stable3D asset @@ -136,9 +194,7 @@ class Stable3DGenerate3D: return { "required": { "trellis_model": ("TRELLIS_MODEL", ), - "normal_predictor": ("STABLE3D_NORMAL", ), - "birefnet_model": ("STABLE3D_BIREFNET", ), - "image": ("IMAGE",), + "normal_image": ("IMAGE",), "seed": ( "INT", { @@ -207,16 +263,14 @@ class Stable3DGenerate3D: def generate_3d( self, trellis_model, - normal_predictor, - birefnet_model, - image, + normal_image, seed=-1, ss_guidance_strength=3, ss_sampling_steps=50, slat_guidance_strength=3, slat_sampling_steps=6 ): - if image is None: + if normal_image is None: return None, None, None if seed == -1: @@ -225,15 +279,6 @@ class Stable3DGenerate3D: 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,