Moving image pre-processing to a dedicated custom node.

This commit is contained in:
Laurent Erignoux
2025-08-01 14:23:13 +08:00
committed by Erignoux Laurent
parent 9c5a3b8850
commit 3cdcef758f
2 changed files with 67 additions and 20 deletions
+5 -3
View File
@@ -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]
+62 -17
View File
@@ -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,