New node: Extend Image for Outpainting

This commit is contained in:
Luis Quesada
2024-06-08 20:45:10 +02:00
parent 3ff35d3b14
commit dd523bda1b
4 changed files with 164 additions and 27 deletions
+4
View File
@@ -10,6 +10,8 @@ Check ComfyUI here: https://github.com/comfyanonymous/ComfyUI
"✂️ Inpaint Stitch" is a node that stitches the inpainted image back into the original image without altering unmasked areas.
"✂️ Extend Image for Outpainting" is a node that extends an image and masks in order to use the power of Inpaint Crop and Stich (rescaling, blur, blend, restitching) for outpainting.
The main advantages of inpainting only in a masked area with these nodes are:
- It's much faster than sampling the whole image.
- It enables setting the right amount of context from the image for the prompt to be more accurately represented in the generated picture.
@@ -60,6 +62,8 @@ If you want to inpaint with SDXL, use forced size = 1024.
# Changelog
## 2024-06-07
- Added the "Extend Image for Outpainting" node that allows leveraging the power of Inpaint Crop and Stitch (rescaling, blur, blend, restitching) for Outpainting.
## 2024-06-07
- Added a blending radius for seamless inpainting.
- Added a blur mask setting that grows and blurs the mask, providing better support
## 2024-06-01
+5 -2
View File
@@ -1,16 +1,19 @@
from .inpaint_cropandstitch import InpaintCrop
from .inpaint_cropandstitch import InpaintStitch
from .inpaint_cropandstitch import InpaintExtendOutpaint
WEB_DIRECTORY = "js"
NODE_CLASS_MAPPINGS = {
"InpaintCrop": InpaintCrop,
"InpaintStitch": InpaintStitch
"InpaintStitch": InpaintStitch,
"InpaintExtendOutpaint": InpaintExtendOutpaint,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"InpaintCrop": "✂️ Inpaint Crop",
"InpaintStitch": "✂️ Inpaint Stitch"
"InpaintStitch": "✂️ Inpaint Stitch",
"InpaintExtendOutpaint": "✂️ Extend Image for Outpainting",
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+110
View File
@@ -434,6 +434,7 @@ class InpaintCrop:
return (stitch, cropped_image, cropped_mask)
class InpaintStitch:
"""
ComfyUI-InpaintCropAndStitch
@@ -550,3 +551,112 @@ class InpaintStitch:
cropped_output = output[:, start_y:start_y + initial_height, start_x:start_x + initial_width, :]
output = cropped_output
return (output,)
class InpaintExtendOutpaint:
"""
ComfyUI-InpaintCropAndStitch
https://github.com/lquesada/ComfyUI-InpaintCropAndStitch
This node extends an image for inpainting with Inpaint Crop and Stitch.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
"mode": (["factors", "pixels"], {"default": "factors"}),
"expand_up_pixels": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
"expand_up_factor": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 100.0, "step": 0.01}),
"expand_down_pixels": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
"expand_down_factor": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 100.0, "step": 0.01}),
"expand_left_pixels": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
"expand_left_factor": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 100.0, "step": 0.01}),
"expand_right_pixels": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1}),
"expand_right_factor": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"optional_context_mask": ("MASK",),
}
}
CATEGORY = "inpaint"
RETURN_TYPES = ("IMAGE", "MASK", "MASK")
RETURN_NAMES = ("image", "mask", "context_mask")
FUNCTION = "inpaint_extend"
def inpaint_extend(self, image, mask, mode, expand_up_pixels, expand_up_factor, expand_down_pixels, expand_down_factor, expand_left_pixels, expand_left_factor, expand_right_pixels, expand_right_factor, optional_context_mask=None):
assert image.shape[0] == mask.shape[0], "Batch size of images and masks must be the same"
if optional_context_mask is not None:
assert optional_context_mask.shape[0] == image.shape[0], "Batch size of optional_context_masks must be the same as images or None"
results_image = []
results_mask = []
results_context_mask = []
batch_size = image.shape[0]
for b in range(batch_size):
one_image = image[b].unsqueeze(0) # Adding batch dimension
one_mask = mask[b].unsqueeze(0) # Adding batch dimension
one_context_mask = optional_context_mask[b].unsqueeze(0) if optional_context_mask is not None else None
if one_mask.shape[1] != one_image.shape[1] or one_mask.shape[2] != one_image.shape[2]:
assert False, "mask size must match image size"
if one_context_mask is not None and (one_context_mask.shape[1] != one_image.shape[1] or one_context_mask.shape[2] != one_image.shape[2]):
assert False, "context_mask size must match image size"
# Get original dimensions
orig_height, orig_width = one_image.shape[1], one_image.shape[2]
if mode == "factors":
# Calculate new dimensions based on factors
new_height = int(orig_height * (expand_up_factor + expand_down_factor - 1))
new_width = int(orig_width * (expand_left_factor + expand_right_factor - 1))
up_padding = int(orig_height * (expand_up_factor - 1))
down_padding = new_height - orig_height - up_padding
left_padding = int(orig_width * (expand_left_factor - 1))
right_padding = new_width - orig_width - left_padding
elif mode == "pixels":
# Calculate new dimensions based on pixel expansion
new_height = orig_height + expand_up_pixels + expand_down_pixels
new_width = orig_width + expand_left_pixels + expand_right_pixels
up_padding = expand_up_pixels
down_padding = expand_down_pixels
left_padding = expand_left_pixels
right_padding = expand_right_pixels
else:
raise ValueError("Mode must be either 'factors' or 'pixels'")
# Expand image
new_image = torch.zeros((one_image.shape[0], new_height, new_width, one_image.shape[3]), dtype=one_image.dtype)
new_image[:, up_padding:up_padding + orig_height, left_padding:left_padding + orig_width, :] = one_image.squeeze(0)
# Expand mask
new_mask = torch.ones((one_mask.shape[0], new_height, new_width), dtype=one_mask.dtype)
new_mask[:, up_padding:up_padding + orig_height, left_padding:left_padding + orig_width] = one_mask.squeeze(0)
# Expand context mask if present
if one_context_mask is not None:
new_context_mask = torch.ones((one_context_mask.shape[0], new_height, new_width), dtype=one_context_mask.dtype)
new_context_mask[:, up_padding:up_padding + orig_height, left_padding:left_padding + orig_width] = one_context_mask.squeeze(0)
# Append results
results_image.append(new_image.squeeze(0))
results_mask.append(new_mask.squeeze(0))
if one_context_mask is not None:
results_context_mask.append(new_context_mask.squeeze(0))
# Stack the results to form batches
output_image = torch.stack(results_image, dim=0)
output_mask = torch.stack(results_mask, dim=0)
output_context_mask = None
if optional_context_mask is not None:
output_context_mask = torch.stack(results_context_mask, dim=0)
return (output_image, output_mask, output_context_mask)
+45 -25
View File
@@ -3,31 +3,51 @@ import { app } from "../../scripts/app.js";
// Some fragments of this code are from https://github.com/LucianoCirino/efficiency-nodes-comfyui
function inpaintCropHandler(node) {
if (node.comfyClass != "InpaintCrop") {
return;
}
toggleWidget(node, findWidgetByName(node, "force_width"));
toggleWidget(node, findWidgetByName(node, "force_height"));
toggleWidget(node, findWidgetByName(node, "rescale_factor"));
toggleWidget(node, findWidgetByName(node, "min_width"));
toggleWidget(node, findWidgetByName(node, "min_height"));
toggleWidget(node, findWidgetByName(node, "max_width"));
toggleWidget(node, findWidgetByName(node, "max_height"));
toggleWidget(node, findWidgetByName(node, "padding"));
if (findWidgetByName(node, "mode").value == "free size") {
toggleWidget(node, findWidgetByName(node, "rescale_factor"), true);
toggleWidget(node, findWidgetByName(node, "padding"), true);
}
else if (findWidgetByName(node, "mode").value == "ranged size") {
toggleWidget(node, findWidgetByName(node, "min_width"), true);
toggleWidget(node, findWidgetByName(node, "min_height"), true);
toggleWidget(node, findWidgetByName(node, "max_width"), true);
toggleWidget(node, findWidgetByName(node, "max_height"), true);
toggleWidget(node, findWidgetByName(node, "padding"), true);
}
else if (findWidgetByName(node, "mode").value == "forced size") {
toggleWidget(node, findWidgetByName(node, "force_width"), true);
toggleWidget(node, findWidgetByName(node, "force_height"), true);
if (node.comfyClass == "InpaintCrop") {
toggleWidget(node, findWidgetByName(node, "force_width"));
toggleWidget(node, findWidgetByName(node, "force_height"));
toggleWidget(node, findWidgetByName(node, "rescale_factor"));
toggleWidget(node, findWidgetByName(node, "min_width"));
toggleWidget(node, findWidgetByName(node, "min_height"));
toggleWidget(node, findWidgetByName(node, "max_width"));
toggleWidget(node, findWidgetByName(node, "max_height"));
toggleWidget(node, findWidgetByName(node, "padding"));
if (findWidgetByName(node, "mode").value == "free size") {
toggleWidget(node, findWidgetByName(node, "rescale_factor"), true);
toggleWidget(node, findWidgetByName(node, "padding"), true);
}
else if (findWidgetByName(node, "mode").value == "ranged size") {
toggleWidget(node, findWidgetByName(node, "min_width"), true);
toggleWidget(node, findWidgetByName(node, "min_height"), true);
toggleWidget(node, findWidgetByName(node, "max_width"), true);
toggleWidget(node, findWidgetByName(node, "max_height"), true);
toggleWidget(node, findWidgetByName(node, "padding"), true);
}
else if (findWidgetByName(node, "mode").value == "forced size") {
toggleWidget(node, findWidgetByName(node, "force_width"), true);
toggleWidget(node, findWidgetByName(node, "force_height"), true);
}
} else if (node.comfyClass == "InpaintExtendOutpaint") {
toggleWidget(node, findWidgetByName(node, "expand_up_pixels"));
toggleWidget(node, findWidgetByName(node, "expand_up_factor"));
toggleWidget(node, findWidgetByName(node, "expand_down_pixels"));
toggleWidget(node, findWidgetByName(node, "expand_down_factor"));
toggleWidget(node, findWidgetByName(node, "expand_left_pixels"));
toggleWidget(node, findWidgetByName(node, "expand_left_factor"));
toggleWidget(node, findWidgetByName(node, "expand_right_pixels"));
toggleWidget(node, findWidgetByName(node, "expand_right_factor"));
if (findWidgetByName(node, "mode").value == "factors") {
toggleWidget(node, findWidgetByName(node, "expand_up_factor"), true);
toggleWidget(node, findWidgetByName(node, "expand_down_factor"), true);
toggleWidget(node, findWidgetByName(node, "expand_left_factor"), true);
toggleWidget(node, findWidgetByName(node, "expand_right_factor"), true);
}
if (findWidgetByName(node, "mode").value == "pixels") {
toggleWidget(node, findWidgetByName(node, "expand_up_pixels"), true);
toggleWidget(node, findWidgetByName(node, "expand_down_pixels"), true);
toggleWidget(node, findWidgetByName(node, "expand_left_pixels"), true);
toggleWidget(node, findWidgetByName(node, "expand_right_pixels"), true);
}
}
return;
}