Fix Dont combine option for SEGS Detector for AnimateDiff (#656)

Co-authored-by: maratz <maratz@ngrow.ai>
This commit is contained in:
Marat Zhanabekov
2024-06-28 22:51:30 +09:00
committed by GitHub
co-authored by maratz
parent f7df6e4445
commit cbc8384a9e
+5 -22
View File
@@ -412,11 +412,12 @@ class SimpleDetectorForAnimateDiff:
merged_mask = get_whole_merged_mask()
return segs_nodes.MaskToSEGS.doit(merged_mask, False, crop_factor, False, drop_size, contour_fill=True)[0]
def get_merged_neighboring_segs():
def get_segs(merged_neighboring=False):
pivot_segs = get_pivot_segs()
masks_by_frame = get_masked_frames()
masks_by_frame = get_merged_neighboring_mask(masks_by_frame)
if merged_neighboring:
masks_by_frame = get_merged_neighboring_mask(masks_by_frame)
new_segs = []
for seg in pivot_segs[1]:
@@ -435,33 +436,15 @@ class SimpleDetectorForAnimateDiff:
return pivot_segs[0], new_segs
def get_separated_segs():
pivot_segs = get_pivot_segs()
masks_by_frame = get_masked_frames()
new_segs = []
for seg in pivot_segs[1]:
cropped_mask = torch.zeros(seg.cropped_mask.shape, dtype=torch.float32, device="cpu").unsqueeze(0)
x1, y1, x2, y2 = seg.crop_region
for mask in masks_by_frame:
cropped_mask_at_frame = mask[y1:y2, x1:x2]
cropped_mask = torch.cat((cropped_mask, cropped_mask_at_frame), dim=0)
new_seg = SEG(seg.cropped_image, cropped_mask, seg.confidence, seg.crop_region, seg.bbox, seg.label, seg.control_net_wrapper)
new_segs.append(new_seg)
return pivot_segs[0], new_segs
# create result mask
if masking_mode == "Pivot SEGS":
return (get_pivot_segs(), )
elif masking_mode == "Combine neighboring frames":
return (get_merged_neighboring_segs(), )
return (get_segs(merged_neighboring=True), )
else: # elif masking_mode == "Don't combine":
return (get_separated_segs(), )
return (get_segs(merged_neighboring=False), )
def doit(self, bbox_detector, image_frames, bbox_threshold, bbox_dilation, crop_factor, drop_size,
sub_threshold, sub_dilation, sub_bbox_expansion, sam_mask_hint_threshold,