325 lines
12 KiB
Python
325 lines
12 KiB
Python
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
|
|
|
|
import folder_paths
|
|
|
|
sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'Stable3DGen'))
|
|
from hi3dgen.pipelines import Hi3DGenPipeline
|
|
|
|
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)
|
|
|
|
trellis_folder = folder_paths.get_folder_paths("trellis")[0]
|
|
hi3dgen_pipeline = Hi3DGenPipeline.from_pretrained(os.path.join(trellis_folder, 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."}
|
|
),
|
|
"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_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
|
|
|
|
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)
|
|
image = Image.fromarray(numpy_image_scaled)
|
|
|
|
# FIXME We should properly handle batch here.
|
|
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 (self.pil_to_tensor(normal_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.")
|
|
for (batch_number, image) in enumerate(normal_images):
|
|
numpy_image = 255. * image.cpu().numpy()
|
|
pil_image = Image.fromarray(numpy.clip(numpy_image, 0, 255).astype(numpy.uint8))
|
|
|
|
outputs = hi3dgen_pipeline.run(
|
|
pil_image,
|
|
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)
|
|
|
|
# FiXME we should return all the files in case of batch.
|
|
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")
|