Added nodes "Continue 5" and "Continue 6" + "Cuda Reset" / Removed not working nodes about Masks
This commit is contained in:
@@ -2708,7 +2708,52 @@ class Trellis2Continue4:
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, input_1, input_2, input_3, input_4):
|
||||
return (input_1, input_2, input_3, input_4)
|
||||
return (input_1, input_2, input_3, input_4)
|
||||
|
||||
class Trellis2Continue5:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_1": (any,),
|
||||
"input_2": (any,),
|
||||
"input_3": (any,),
|
||||
"input_4": (any,),
|
||||
"input_5": (any,),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any, any, any, any, any,)
|
||||
RETURN_NAMES = ("output_1", "output_2", "output_3", "output_4", "output_5")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, input_1, input_2, input_3, input_4, input_5):
|
||||
return (input_1, input_2, input_3, input_4, input_5)
|
||||
|
||||
class Trellis2Continue6:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_1": (any,),
|
||||
"input_2": (any,),
|
||||
"input_3": (any,),
|
||||
"input_4": (any,),
|
||||
"input_5": (any,),
|
||||
"input_6": (any,),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any, any, any, any, any, any)
|
||||
RETURN_NAMES = ("output_1", "output_2", "output_3", "output_4", "output_5", "output_6")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, input_1, input_2, input_3, input_4, input_5, input_6):
|
||||
return (input_1, input_2, input_3, input_4, input_5, input_6)
|
||||
|
||||
class Trellis2MeshWithVoxelToMeshlibMesh:
|
||||
@classmethod
|
||||
@@ -4074,219 +4119,25 @@ class Trellis2MultiViewTexturing:
|
||||
except ValueError:
|
||||
print(f"[MultiView] Warning: Could not parse angles: {angle_string}")
|
||||
return []
|
||||
|
||||
class Trellis2ApplyMaskToImage:
|
||||
|
||||
class Trellis2CudaReset:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"color_image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
# If True, pastes the white mask behind the subject like in your 3rd image
|
||||
"show_target_mask_behind": ("BOOLEAN", {"default": True}),
|
||||
"input_1": (any,),
|
||||
},
|
||||
"optional": {
|
||||
# Optional: if auto-background detection fails, you can feed a mask of the source image here
|
||||
"optional_source_mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "adapt_to_mask"
|
||||
RETURN_TYPES = (any, )
|
||||
RETURN_NAMES = ("output_1", )
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def adapt_to_mask(self, color_image, mask, show_target_mask_behind, optional_source_mask=None):
|
||||
B, H, W, C = color_image.shape
|
||||
_, mH, mW = mask.shape
|
||||
|
||||
out_images = []
|
||||
|
||||
for i in range(B):
|
||||
img = color_image[i] # (H, W, C)
|
||||
tgt_mask = mask[i] if i < mask.shape[0] else mask[-1] # (mH, mW)
|
||||
|
||||
# 1. Find the Bounding Box of the target mask
|
||||
tgt_coords = torch.nonzero(tgt_mask > 0.5)
|
||||
if tgt_coords.size(0) == 0:
|
||||
out_images.append(img.unsqueeze(0)) # If mask is blank, return original
|
||||
continue
|
||||
|
||||
tgt_ymin, tgt_xmin = torch.min(tgt_coords, dim=0)[0]
|
||||
tgt_ymax, tgt_xmax = torch.max(tgt_coords, dim=0)[0]
|
||||
tgt_ymax += 1
|
||||
tgt_xmax += 1
|
||||
|
||||
tgt_h = (tgt_ymax - tgt_ymin).item()
|
||||
tgt_w = (tgt_xmax - tgt_xmin).item()
|
||||
|
||||
# 2. Find the Bounding Box of the subject in the source image
|
||||
if optional_source_mask is not None:
|
||||
src_m = optional_source_mask[i] if i < optional_source_mask.shape[0] else optional_source_mask[-1]
|
||||
src_coords = torch.nonzero(src_m > 0.5)
|
||||
else:
|
||||
if C == 4:
|
||||
# If RGBA, use the alpha channel to find the subject
|
||||
src_coords = torch.nonzero(img[:, :, 3] > 0.1)
|
||||
else:
|
||||
# Auto-detect background using the top-left pixel
|
||||
bg_color = img[0, 0, :]
|
||||
diff = torch.abs(img - bg_color).sum(dim=-1)
|
||||
src_coords = torch.nonzero(diff > 0.1)
|
||||
|
||||
if src_coords.size(0) == 0:
|
||||
out_images.append(img.unsqueeze(0))
|
||||
continue
|
||||
|
||||
src_ymin, src_xmin = torch.min(src_coords, dim=0)[0]
|
||||
src_ymax, src_xmax = torch.max(src_coords, dim=0)[0]
|
||||
src_ymax += 1
|
||||
src_xmax += 1
|
||||
|
||||
# 3. Crop the source subject
|
||||
cropped_img = img[src_ymin:src_ymax, src_xmin:src_xmax, :] # (src_H, src_W, C)
|
||||
|
||||
# 4. Resize (Stretch/Shrink) the crop to match the target mask's bounding box
|
||||
cropped_img_permuted = cropped_img.permute(2, 0, 1).unsqueeze(0) # (1, C, H, W) for interpolate
|
||||
resized_img = F.interpolate(
|
||||
cropped_img_permuted,
|
||||
size=(tgt_h, tgt_w),
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
resized_img = resized_img.squeeze(0).permute(1, 2, 0) # Back to (tgt_h, tgt_w, C)
|
||||
|
||||
# 5. Create Output Canvas
|
||||
canvas = torch.zeros((mH, mW, C), dtype=img.dtype, device=img.device)
|
||||
|
||||
# Fill with the original background color (using top-left pixel)
|
||||
if C >= 3:
|
||||
bg_color = img[0, 0, :3]
|
||||
canvas[:, :, :3] = bg_color
|
||||
if C == 4:
|
||||
canvas[:, :, 3] = 1.0 # Solid alpha for background
|
||||
|
||||
# Replicate the white mask background seen in your 3rd image
|
||||
if show_target_mask_behind:
|
||||
mask_expanded = tgt_mask.unsqueeze(-1).repeat(1, 1, C)
|
||||
# Where mask is 1, make canvas white
|
||||
if C == 3:
|
||||
canvas = torch.where(mask_expanded > 0.5, torch.ones_like(canvas), canvas)
|
||||
elif C == 4:
|
||||
white_rgba = torch.ones_like(canvas)
|
||||
canvas = torch.where(mask_expanded > 0.5, white_rgba, canvas)
|
||||
|
||||
# 6. Paste the stretched subject into the target coordinates
|
||||
canvas[tgt_ymin:tgt_ymax, tgt_xmin:tgt_xmax, :] = resized_img
|
||||
|
||||
out_images.append(canvas.unsqueeze(0))
|
||||
|
||||
return (torch.cat(out_images, dim=0),)
|
||||
|
||||
class Trellis2TargetFitMask:
|
||||
"""
|
||||
A powerful ComfyUI node that non-uniformly stretches or shrinks
|
||||
an image to perfectly fill the bounding box of a target mask.
|
||||
The output background is determined by the input image's top-left pixel.
|
||||
"""
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"color_image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
# If set to True, pastes the white mask shape behind the subject
|
||||
# exactly as seen in your provided 'result' image example.
|
||||
"paste_mask_shape_behind": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "apply_transform"
|
||||
CATEGORY = "image/transform"
|
||||
|
||||
def apply_transform(self, color_image, mask, paste_mask_shape_behind):
|
||||
# Image shapes: (Batch, Height, Width, Channels)
|
||||
# Mask shape: (Batch, Height, Width)
|
||||
B, H, W, C = color_image.shape
|
||||
mBatch, mHeight, mWidth = mask.shape
|
||||
|
||||
# Determine background color using the source image's top-left pixel
|
||||
# (Assuming the background is solid)
|
||||
source_bg_color = color_image[0, 0, 0, :3].unsqueeze(0).unsqueeze(0)
|
||||
|
||||
results = []
|
||||
|
||||
for i in range(B):
|
||||
img = color_image[i] # (H, W, C)
|
||||
tgt_mask = mask[i] if i < mBatch else mask[-1] # (mH, mW)
|
||||
|
||||
# --- 1. Get Target Bounding Box from the white mask region ---
|
||||
tgt_coords = torch.nonzero(tgt_mask > 0.5)
|
||||
|
||||
if tgt_coords.size(0) == 0:
|
||||
# Target mask is empty, append original or solid color
|
||||
results.append(img.unsqueeze(0))
|
||||
continue
|
||||
|
||||
tgt_ymin, tgt_xmin = torch.min(tgt_coords, dim=0)[0]
|
||||
tgt_ymax, tgt_xmax = torch.max(tgt_coords, dim=0)[0]
|
||||
|
||||
# Calculate required height and width
|
||||
tgt_h_box = (tgt_ymax - tgt_ymin).item() + 1
|
||||
tgt_w_box = (tgt_xmax - tgt_xmin).item() + 1
|
||||
|
||||
# --- 2. Auto-detect subject in source image ---
|
||||
# Using basic thresholding against the detected bg color
|
||||
diff = torch.abs(img[..., :3] - source_bg_color).sum(dim=-1)
|
||||
src_coords = torch.nonzero(diff > 0.1) # Threshold to isolate subject
|
||||
|
||||
if src_coords.size(0) == 0:
|
||||
results.append(img.unsqueeze(0))
|
||||
continue
|
||||
|
||||
src_ymin, src_xmin = torch.min(src_coords, dim=0)[0]
|
||||
src_ymax, src_xmax = torch.max(src_coords, dim=0)[0]
|
||||
|
||||
# Crop the subject
|
||||
# ymax/xmax +1 for correct slicing
|
||||
subject_crop = img[src_ymin:src_ymax+1, src_xmin:src_xmax+1, :]
|
||||
|
||||
# --- 3. Resize (Stretch/Shrink) to Target Box size ---
|
||||
# Interpolate requires (B, C, H, W)
|
||||
subject_crop_p = subject_crop.permute(2, 0, 1).unsqueeze(0)
|
||||
|
||||
warped_subject_p = F.interpolate(
|
||||
subject_crop_p,
|
||||
size=(tgt_h_box, tgt_w_box),
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# Squeeze back to (H, W, C)
|
||||
warped_subject = warped_subject_p.squeeze(0).permute(1, 2, 0)
|
||||
|
||||
# --- 4. Composite the result ---
|
||||
# Start with a canvas matching mask dimensions, filled with bg color
|
||||
canvas = source_bg_color.repeat(mHeight, mWidth, 1)
|
||||
|
||||
if paste_mask_shape_behind:
|
||||
# Map mask to canvas dimensions for the underlay
|
||||
mask_map = tgt_mask.unsqueeze(-1).repeat(1, 1, 3)
|
||||
# Where mask is white, fill canvas with pure white
|
||||
canvas = torch.where(mask_map > 0.5, torch.ones_like(canvas), canvas)
|
||||
|
||||
# Paste the non-uniformly warped subject into the target bbox
|
||||
canvas[tgt_ymin:tgt_ymin+tgt_h_box, tgt_xmin:tgt_xmin+tgt_w_box, :] = warped_subject
|
||||
|
||||
results.append(canvas.unsqueeze(0))
|
||||
|
||||
# Join batch
|
||||
final_result = torch.cat(results, dim=0)
|
||||
|
||||
return (final_result,)
|
||||
def process(self, input_1):
|
||||
reset_cuda()
|
||||
return (input_1,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Trellis2LoadModel": Trellis2LoadModel,
|
||||
@@ -4336,8 +4187,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"Trellis2MultiViewTexturing": Trellis2MultiViewTexturing,
|
||||
"Trellis2Continue3": Trellis2Continue3,
|
||||
"Trellis2Continue4": Trellis2Continue4,
|
||||
"Trellis2ApplyMaskToImage": Trellis2ApplyMaskToImage,
|
||||
"Trellis2TargetFitMask": Trellis2TargetFitMask,
|
||||
"Trellis2Continue5": Trellis2Continue5,
|
||||
"Trellis2Continue6": Trellis2Continue6,
|
||||
"Trellis2CudaReset": Trellis2CudaReset,
|
||||
}
|
||||
|
||||
|
||||
@@ -4389,6 +4241,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Trellis2MultiViewTexturing": "Trellis2 - Projection MultiView Texturing",
|
||||
"Trellis2Continue3": "Trellis2 - Continue 3",
|
||||
"Trellis2Continue4": "Trellis2 - Continue 4",
|
||||
"Trellis2ApplyMaskToImage": "Trellis2 - Apply Mask To Image",
|
||||
"Trellis2TargetFitMask": "Trellis2 - Target Fit Mask",
|
||||
"Trellis2Continue5": "Trellis2 - Continue 5",
|
||||
"Trellis2Continue6": "Trellis2 - Continue 6",
|
||||
"Trellis2CudaReset": "Trellis2 - Cuda Reset",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user