Fix Dont combine option for SEGS Detector for AnimateDiff (#656)
Co-authored-by: maratz <maratz@ngrow.ai>
This commit is contained in:
co-authored by
maratz
parent
f7df6e4445
commit
cbc8384a9e
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user