Compare commits

...
Author SHA1 Message Date
Dr.Lt.Data c2824aa943 wip 2024-02-09 13:22:07 +09:00
Dr.Lt.Data 1b86633194 wip 2024-02-08 13:10:25 +09:00
5 changed files with 115 additions and 2 deletions
+1
View File
@@ -227,6 +227,7 @@ NODE_CLASS_MAPPINGS = {
"ImpactDilateMaskInSEGS": DilateMaskInSEGS,
"ImpactGaussianBlurMaskInSEGS": GaussianBlurMaskInSEGS,
"ImpactScaleBy_BBOX_SEG_ELT": SEG_ELT_BBOX_ScaleBy,
"ImpactSEGSLatentComposite": SEGSLatentComposite,
"BboxDetectorCombined_v2": BboxDetectorCombined,
"SegmDetectorCombined_v2": SegmDetectorCombined,
+10
View File
@@ -1739,6 +1739,16 @@ class BBoxDetectorBasedOnCLIPSeg:
self.aux = x
def get_seg_size(seg):
x1, y1, x2, y2 = seg.crop_region
return x2-x1, y2-y1
def get_bbox_size(seg):
x1, y1, x2, y2 = seg.bbox
return x2-x1, y2-y1
def update_node_status(node, text, progress=None):
if PromptServer.instance.client_id is None:
return
+92 -1
View File
@@ -1,7 +1,10 @@
import os
import sys
import torch
import impact.impact_server
import nodes
from nodes import MAX_RESOLUTION
from impact.utils import *
@@ -123,6 +126,94 @@ class SEGSDetailer:
return (segs, cnet_pil_list)
class SEGSLatentComposite:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"src_segs": ("SEGS", ),
"dest_segs": ("SEGS", ),
"target_latent": ("LATENT", ),
"vae": ("VAE", ),
},
"optional": {"ref_image_opt": ("IMAGE", ), }
}
RETURN_TYPES = ("LATENT", )
FUNCTION = "doit"
CATEGORY = "ImpactPack/__for_testing"
@staticmethod
def doit(src_segs, dest_segs, target_latent, vae, ref_image_opt=None):
apply_count = min(len(src_segs[1]), len(dest_segs[1]))
if apply_count == 0:
print(f"[Impact Pack] SEGSLatentComposite: src_segs or dest_segs is empty")
return (target_latent, )
if src_segs[1][0].cropped_image is None:
if ref_image_opt is None:
print(f"[Impact Pack] SEGSLatentComposite: there is no cropped_image nor ref_image_opt")
return (target_latent, )
src_segs = DefaultImageForSEGS().doit(src_segs, ref_image_opt, False)[0]
target_latent = target_latent.copy()
target_latent['samples'] = torch.clone(target_latent['samples'])
for i in range(0, apply_count):
seg1 = src_segs[1][i]
seg2 = dest_segs[1][i]
w1, h1 = core.get_bbox_size(seg1)
w2, h2 = core.get_bbox_size(seg2)
scale_factor = 1
if w1-w2 < h1-h2:
# fit to vertical
target_h = h2
scale_factor = target_h/h1
target_w = int(w1 * scale_factor)
dy = 0
dx = int((w2-target_w)/2//8)
elif w1-w2 > h1-h2:
# fit to horizontal
target_w = w2
scale_factor = target_w/w1
target_h = int(h1*scale_factor)
dx = 0
dy = int((h2-target_h)/2//8)
else:
# same ratio
target_h, target_w = w2, h2
scale_factor = target_w/w1
dx, dy = 0, 0
ax1, ay1, ax2, ay2 = seg1.bbox
bx1, by1, _, _ = seg1.crop_region
ax1, ax2, ay1, ay2 = (ax1-bx1), (ax2-bx1), (ay1-by1), (ay2-by1)
seg1_bbox_image = seg1.cropped_image[:, ay1:ay2, ax1:ax2, :]
seg1_image = utils.tensor_resize(seg1_bbox_image, target_w, target_h)
seg1_samples = nodes.VAEEncode().encode(vae, seg1_image)[0]['samples']
seg1_mask = utils.make_3d_mask(torch.from_numpy(seg1.cropped_mask))
seg1_mask = seg1_mask[ :, ay1:ay2, ax1:ax2]
seg1_mask = utils.resize_mask(seg1_mask, tuple(seg1_samples.shape[2:4]))
seg1_mask = seg1_mask.unsqueeze(0)
x1, y1, _, _ = seg2.bbox
x1 = int(x1//8 + dx)
y1 = int(y1//8 + dy)
x2 = x1 + seg1_samples.shape[3]
y2 = y1 + seg1_samples.shape[2]
target_samples = target_latent['samples'][:, :, y1:y2, x1:x2]
target_samples = seg1_samples*seg1_mask + target_samples*(1.0-seg1_mask)
target_latent['samples'][:, :, y1:y2, x1:x2] = target_samples
return (target_latent,)
class SEGSPaste:
@classmethod
def INPUT_TYPES(s):
@@ -1350,7 +1441,7 @@ class MakeTileSEGS:
FUNCTION = "doit"
CATEGORY = "ImpactPack/__for_testing"
CATEGORY = "ImpactPack/Util"
def doit(self, images, bbox_size, crop_factor, min_overlap, filter_segs_dilation, mask_irregularity=0, irregular_mask_mode="Reuse fast", filter_in_segs_opt=None, filter_out_segs_opt=None):
if bbox_size <= 2*min_overlap:
+1 -1
View File
@@ -251,7 +251,7 @@ class ConcatConditionings:
RETURN_TYPES = ("CONDITIONING", )
FUNCTION = "doit"
CATEGORY = "ImpactPack/__for_testing"
CATEGORY = "ImpactPack/Util"
def doit(self, **kwargs):
conditioning_to = list(kwargs.values())[0]
+11
View File
@@ -524,6 +524,17 @@ def make_3d_mask(mask):
return mask
def get_mask_size(mask):
if len(mask.shape) == 2:
return mask.shape[1], mask.shape[0]
elif len(mask.shape) == 3:
return mask.shape[2], mask.shape[1]
elif len(mask.shape) == 4:
return mask.shape[3], mask.shape[2]
raise Exception("unexpected mask dimension")
def is_same_device(a, b):
a_device = torch.device(a) if isinstance(a, str) else a
b_device = torch.device(b) if isinstance(b, str) else b