Moving image pre-processing to a dedicated custom node.
This commit is contained in:
committed by
Erignoux Laurent
parent
9c5a3b8850
commit
3cdcef758f
+5
-3
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user