From 0631af4b7cbc5f9055d4e596345e9f0c23de3532 Mon Sep 17 00:00:00 2001 From: Erignoux Laurent Date: Thu, 31 Jul 2025 14:52:29 +0800 Subject: [PATCH] Fixing image preprocessing --- stable_3d.py | 31 +++++++++++++++++-------------- 1 file changed, 17 insertions(+), 14 deletions(-) diff --git a/stable_3d.py b/stable_3d.py index 63bf068..a4eff6a 100644 --- a/stable_3d.py +++ b/stable_3d.py @@ -123,9 +123,9 @@ class Stable3DPreprocessImage: def INPUT_TYPES(s): return { "required": { - "background_removal_model": ( - "STABLE3D_BIREFNET", - {"tooltip": "The model to remove background (BiRefNet)."} + "hi3dgen_pipeline": ( + "HI3DGEN_PIPELINE", + {"tooltip": "The hi3dgen pipeline containing the remove background setup (BiRefNet)."} ), "normal_predictor": ( "STABLE3D_NORMAL", @@ -140,33 +140,36 @@ class Stable3DPreprocessImage: FUNCTION = "preprocess_image" INPUT_IS_LIST = False OUTPUT_NODE = True - RETURN_NAMES = ("normal_image",) + 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.info(f"normal_image saved as {path}") + return 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 + def preprocess_image(self, hi3dgen_pipeline, normal_predictor, image): + # FIXME We should support properly batch mode here. 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) + 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) + image = hi3dgen_pipeline.preprocess_image(image, resolution=1024) + normal_image = normal_predictor(image, resolution=768, match_input_resolution=True, data_type='object') self.save_normal_image(normal_image) - return normal_image + return (self.pil_to_tensor(normal_image),) class Stable3DGenerate3D: