extract and merge frames

This commit is contained in:
spacepxl
2024-12-12 21:23:05 -05:00
parent b46bca5672
commit d0b38f0b63
+85
View File
@@ -2158,6 +2158,87 @@ class CustomNoise:
return (Noise_CustomNoise(noise_latent),)
class ExtractNFrames:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"frames": ("INT", {"default": 16, "min": 2}),
},
"optional": {
"images": ("IMAGE",),
"masks": ("MASK",),
},
}
RETURN_TYPES = ("LIST", "IMAGE", "MASK")
RETURN_NAMES = ("index_list", "images", "masks")
FUNCTION = "extract"
CATEGORY = "image/filters/frames"
def extract(self, frames, images=None, masks=None):
original_length = 2
if images is not None: original_length = max(original_length, len(images))
if masks is not None: original_length = max(original_length, len(masks))
n = min(original_length, frames)
step = step = (original_length - 1) / (n - 1)
ids = [round(i * step) for i in range(n)]
while len(ids) < frames:
ids.append(ids[-1])
new_images = []
new_masks = []
for i in ids:
if images is not None:
new_images.append(images[min(i, len(images) - 1)].detach().clone())
else:
new_images.append(torch.zeros(512, 512, 3))
if masks is not None:
new_masks.append(masks[min(i, len(masks) - 1)].detach().clone())
else:
new_masks.append(torch.zeros(512, 512))
return (ids, torch.stack(new_images, dim=0), torch.stack(new_masks, dim=0))
class MergeFramesByIndex:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"index_list": ("LIST",),
"orig_images": ("IMAGE",),
"images": ("IMAGE",),
},
"optional": {
"orig_masks": ("MASK",),
"masks": ("MASK",),
},
}
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("images", "masks")
FUNCTION = "merge"
CATEGORY = "image/filters/frames"
def merge(self, index_list, orig_images, images, orig_masks=None, masks=None):
new_images = orig_images.detach().clone()
new_masks = torch.ones_like(new_images[..., 0]) # BHW
if orig_masks is not None:
for i in range(len(new_masks)):
new_masks[i] = orig_masks[min(i, len(orig_masks) - 1)].detach().clone()
for i, frame in enumerate(set(index_list)):
frame_mask = masks[i] if masks is not None else torch.ones_like(new_masks[i])
new_images[frame] *= (1 - frame_mask[..., None])
new_images[frame] += images[i].detach().clone() * frame_mask[..., None]
new_masks[frame] *= 0
return (new_images, new_masks)
NODE_CLASS_MAPPINGS = {
"AdainFilterLatent": AdainFilterLatent,
"AdainImage": AdainImage,
@@ -2184,6 +2265,7 @@ NODE_CLASS_MAPPINGS = {
"DilateErodeMask": DilateErodeMask,
"EnhanceDetail": EnhanceDetail,
"ExposureAdjust": ExposureAdjust,
"ExtractNFrames": ExtractNFrames,
"FrequencyCombine": FrequencyCombine,
"FrequencySeparate": FrequencySeparate,
"GameOfLife": GameOfLife,
@@ -2199,6 +2281,7 @@ NODE_CLASS_MAPPINGS = {
"LatentNormalizeShuffle": LatentNormalizeShuffle,
"LatentStats": LatentStats,
"MedianFilterImage": MedianFilterImage,
"MergeFramesByIndex": MergeFramesByIndex,
"ModelTest": ModelTest,
"NormalMapSimple": NormalMapSimple,
"OffsetLatentImage": OffsetLatentImage,
@@ -2240,6 +2323,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DilateErodeMask": "Dilate/Erode Mask",
"EnhanceDetail": "Enhance Detail",
"ExposureAdjust": "Exposure Adjust",
"ExtractNFrames": "Extract N Frames",
"FrequencyCombine": "Frequency Combine",
"FrequencySeparate": "Frequency Separate",
"GameOfLife": "Game Of Life",
@@ -2255,6 +2339,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LatentNormalizeShuffle": "LatentNormalizeShuffle",
"LatentStats": "Latent Stats",
"MedianFilterImage": "Median Filter Image",
"MergeFramesByIndex": "Merge Frames By Index",
"ModelTest": "Model Test",
"NormalMapSimple": "Normal Map (Simple)",
"OffsetLatentImage": "Offset Latent Image",