diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..bf4d75c --- /dev/null +++ b/__init__.py @@ -0,0 +1,13 @@ +from .nodes.MakeFrame import(BreakFrames, GetKeyFrames, MakeGrid) + +NODE_CLASS_MAPPINGS = { + "BreakFrames": BreakFrames, + "GetKeyFrames": GetKeyFrames, + "MakeGrid": MakeGrid, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "BreakFrames": "BreakFrames", + "GetKeyFrames": "GetKeyFrames", + "MakeGrid": "MakeGrid", +} \ No newline at end of file diff --git a/makeframeutils.py b/makeframeutils.py new file mode 100644 index 0000000..bacb7b7 --- /dev/null +++ b/makeframeutils.py @@ -0,0 +1,194 @@ +import torch +import numpy as np +from pathlib import Path +import os +from torchvision.transforms import ToTensor, ToPILImage +from PIL import Image, ImageDraw, ImageFont + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") +#Get num closest to 8 +def cl8(num): + rem = num % 8 + if rem <= 4: + return round(num - rem) + else: + return round(num + (8 - rem)) + +def closest_lcm(n, div1, div2): + # Find the LCM + lcm = np.lcm(int(div1), int(div2)) + + # Find the multiple of LCM closest to n + lower_multiple = (n // lcm) * lcm # Largest multiple of lcm less than or equal to n + upper_multiple = lower_multiple + lcm # Smallest multiple of lcm greater than n + + # Find which multiple is closer to n + if n - lower_multiple > upper_multiple - n: + return upper_multiple + else: + return lower_multiple + +def normalize_size(images): + refimage = images[0] + refimage = refimage.resize((cl8(refimage.width), cl8(refimage.height)), Image.Resampling.LANCZOS) + return_images = [] + for i in range(len(images)): + if images[i].size != refimage.size: + images[i] = images[i].resize(refimage.size, Image.Resampling.LANCZOS) + return_images.append(images[i]) + np.lcm(6, 8) + return return_images + +def constrain_image(image, max_width, max_height): + width, height = image.size + aspect_ratio = width / float(height) + + if width > max_width or height > max_height: + if width / float(max_width) > height / float(max_height): + new_width = max_width + new_height = int(new_width / aspect_ratio) + else: + new_height = max_height + new_width = int(new_height * aspect_ratio) + image = image.resize((cl8(new_width), cl8(new_height)), Image.Resampling.LANCZOS) + + return image + +def padlist(lst, targetsize): + if targetsize <= len(lst): + return lst[:targetsize] + + last_elem = lst[-1] + num_repeats = targetsize - len(lst) + + return lst + [last_elem] * num_repeats + +def MakeGrid(images, rows, cols): + widths, heights = zip(*(i.size for i in images)) + + grid_width = max(widths) * cols + grid_height = max(heights) * rows + cell_width = grid_width // cols + cell_height = grid_height // rows + final_image = Image.new('RGB', (grid_width, grid_height)) + x_offset = 0 + y_offset = 0 + for i in range(len(images)): + final_image.paste(images[i], (x_offset, y_offset)) + x_offset += cell_width + if x_offset == grid_width: + x_offset = 0 + y_offset += cell_height + + # Save the final image + return final_image + +def BreakGrid(grid, rows, cols): + width = grid.width // cols + height = grid.height // rows + outimages = [] + for row in range(rows): + for col in range(cols): + left = col * width + top = row * height + right = left + width + bottom = top + height + current_img = grid.crop((left, top, right, bottom)) + outimages.append(current_img) + return outimages + +def ImgLabeler(img, text, size=72, color=(255,255,255)): + font = ImageFont.truetype("arial.ttf", size) + draw = ImageDraw.Draw(img) + + # Get text size + text_size = draw.textsize(text, font=font) + + # Calculate x, y coordinates of the text + x = (img.width - text_size[0]) / 2 + y = (img.height - text_size[1]) / 2 + + # Position for the text, centered + text_position = (x, y) + + # Draw the text onto the image + draw.text(text_position, text, font=font, fill=color) + + return img + +def load_and_preprocess(image_path): + with Image.open(image_path) as img: + return torch.tensor(np.array(img.convert('L')), device=device, dtype=torch.float32) + +def compute_histogram(tensor): + # Min and max values for grayscale images + min_val, max_val = 0, 255 + # Compute the histogram by counting the number of occurrences within each bin + hist = torch.histc(tensor, bins=255, min=min_val, max=max_val) + return hist / tensor.numel() # Normalize by the number of elements + +def get_iterated_path(directory, base_filename, extension='.png'): + # Construct the full file path + count = 0 + while True: + # Append a count to the filename if it's not the first file + if count == 0: + unique_filename = f"{base_filename}{extension}" + else: + unique_filename = f"{base_filename}_{count}{extension}" + full_file_path = os.path.join(directory, unique_filename) + + # Check if a file with this name already exists + if not os.path.exists(full_file_path): + break # Exit the loop once the image is saved + count += 1 + + return full_file_path + +def CheckMakeDir(dir): + try: + if not Path.exists(Path(dir)): + Path.mkdir(Path(dir), parents=True, exist_ok=True) + return True, "Success" + except: + st_out = f"Error: could not locate nor create directory {dir}." + print(st_out) + return False, st_out + +def apply_conditional_ema_pytorch(images: list[Image.Image], alpha: float, threshold: float) -> list[Image.Image]: + + processed_images = [] + ema_tensor = None + to_tensor = ToTensor() + to_pil = ToPILImage() + + for image in images: + frame_tensor = to_tensor(image).to(device) + + if ema_tensor is None: + # For the first frame, EMA is just the frame itself + ema_tensor = frame_tensor.clone() + else: + # Calculate EMA for each subsequent frame + deviation = torch.abs(frame_tensor - ema_tensor) + update_mask = deviation > threshold # Mask to identify where to update + + # EMA update + ema_tensor = alpha * frame_tensor + (1 - alpha) * ema_tensor + ema_tensor[update_mask] = frame_tensor[update_mask] # Apply conditional update + + processed_images.append(to_pil(ema_tensor)) + + return processed_images + +def cat_to_pils(tensor): + to_pil = ToPILImage() + tensor = tensor.permute(0, 3, 1, 2) + pils = [to_pil(tensor[i]) for i in range(tensor.shape[0])] + return pils + +def pil_to_cat(pimg): + to_tensor = ToTensor() + tensor = to_tensor(pimg).unsqueeze(0).permute(0, 2, 3, 1) + print(tensor.shape) + return tensor \ No newline at end of file diff --git a/nodes/MakeFrame.py b/nodes/MakeFrame.py new file mode 100644 index 0000000..9d5a99a --- /dev/null +++ b/nodes/MakeFrame.py @@ -0,0 +1,149 @@ +import cv2 +import torch +import torchvision.transforms as transforms +import numpy as np +import os +from PIL import Image +from .. import makeframeutils as mfu + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + +class BreakFrames: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "file_input": ("STRING", { + "multiline": False, + "default": "C:/Videos/video.mp4" + }), + }, + } + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("Frames",) + + FUNCTION = "breakframes" + CATEGORY = "Frames" + + def breakframes(self, file_input): + if not os.path.exists(file_input): + raise FileNotFoundError(f"File '{file_input} cannot be found.'") + # Open the video file + try: + video_capture = cv2.VideoCapture(file_input) + # Check if the anim was loaded + if not video_capture.isOpened(): + print(f"Error: Could not open video file {file_input}.") + return None + except: + print(f"Error: Could not open video file {file_input}.") + return None + # Read each frame from anim + frame_count = 0 + transformer = transforms.ToTensor() + tensors = [] + while True: + ret, frame = video_capture.read() + if ret: + frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + tensors.append(transformer(frame).to(device).unsqueeze(0)) + frame_count += 1 + else: + break + video_capture.release() + cat_frame_tensors = torch.cat(tensors, dim = 0).permute(0, 2, 3, 1) + + return (cat_frame_tensors,) + +class GetKeyFrames: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "frames": ("IMAGE", ), + "num_keyframes": ("INT", { + "default": 12, + "min": 2, + "max": 4096, + "step": 1 + }), + }, + } + RETURN_TYPES = ("IMAGE", "INT",) + RETURN_NAMES = ("Keyframes", "Keyframe indices",) + + FUNCTION = "getkeyframes" + CATEGORY = "Frames" + + def getkeyframes(self, frames, num_keyframes): + N = np.clip(num_keyframes, 2, len(frames)-1) # Save a spot for first frame + frames = frames + differences = [torch.norm(frames[i+1] - frames[i], p=2) for i in range(len(frames)-1)] + _, top_indices = torch.topk(torch.tensor(differences), k=N, largest=True) + keyframe_indices = sorted([index.item() + 1 for index in top_indices]) + keyframe_indices.insert(0, 0) + cat_keyframe_tensors = [frames[i].to(device).unsqueeze(0) for i in keyframe_indices] + cat_keyframe_tensors = torch.cat(cat_keyframe_tensors, dim = 0) + return (cat_keyframe_tensors, keyframe_indices) + +class MakeGrid: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "frames": ("IMAGE", ), + "grid_rows": ("INT", { + "default": 4, + "min": 2, + "max": 24, + "step": 1 + }), + "grid_cols": ("INT", { + "default": 4, + "min": 2, + "max": 24, + "step": 1 + }), + "max_width": ("INT", { + "default": 1024, + "min": 64, + "max": 4096, + "step": 8 + }), + "max_height": ("INT", { + "default": 1024, + "min": 64, + "max": 4096, + "step": 8 + }), + }, + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("Grid",) + + FUNCTION = "makegrid" + CATEGORY = "Frames" + + def makegrid(self, frames, grid_rows, grid_cols, max_width, max_height): + # Pad the list with extras; black space frustrates generation + + # Size needs to be divisible by 8 AND the grid dimension + if not grid_rows == 8: max_height = mfu.closest_lcm(max_height, 8, grid_rows) + if not grid_cols == 8: max_width = mfu.closest_lcm(max_width, 8, grid_cols) + # Build base grid + pils = mfu.cat_to_pils(frames) + pils = mfu.normalize_size(pils) #normalize sizes to /8 + grid = mfu.constrain_image(mfu.MakeGrid(pils, grid_rows, grid_cols), max_width, max_height) + grid = mfu.pil_to_cat(grid) + + return (grid,) \ No newline at end of file