From 9a2425f12fd3a78eae8f3b2f54ef41463e56de9d Mon Sep 17 00:00:00 2001 From: AbyssYuan0 Date: Thu, 18 Jan 2024 16:11:55 +0800 Subject: [PATCH] =?UTF-8?q?=E5=91=BD=E5=90=8D=E8=A7=84=E6=A0=BC=E5=8C=96?= =?UTF-8?q?=EF=BC=8CGC=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- __init__.py | 31 +++++++++++- color_editor.py | 110 +++++++++++++++++++++++++++++++++++++++++ line_editor.py | 129 ++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 268 insertions(+), 2 deletions(-) create mode 100644 color_editor.py create mode 100644 line_editor.py diff --git a/__init__.py b/__init__.py index 55443ff..238d808 100644 --- a/__init__.py +++ b/__init__.py @@ -7,8 +7,9 @@ import torch import comfy.utils from .videoCut import getCutList, video_to_frames, cutToDir, frames_to_video from .seg import get_masks -from .thick_lines_from_canny import fill_white_segments, find_largest_white_component -from .remove_line import get_colors, find_similar_colors, most_common_fuzzy_color +from .line_editor import fill_white_segments, find_largest_white_component +from .color_editor import get_colors, find_similar_colors, most_common_fuzzy_color +import gc def getImageSize(IMAGE) -> tuple[int, int]: @@ -17,6 +18,10 @@ def getImageSize(IMAGE) -> tuple[int, int]: return size +def maskTensorToImgTensor(maskTensor): + return maskTensor.reshape((-1, 1, maskTensor.shape[-2], maskTensor.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + + def tensorToImg(imageTensor): imaget = imageTensor[0] i = 255. * imaget.cpu().numpy() @@ -757,6 +762,7 @@ class ApplyMaskToImage: def apply_mask_to_image(self, image, mask): image = tensorToImg(image) + mask = maskTensorToImgTensor(mask) mask = tensorToImg(mask) mask = mask.convert("L") @@ -1071,6 +1077,26 @@ class IdentifyLinesBasedOnBorderColor: return (msk_img, mask,) +class GarbageCollect: + def __init__(self) -> None: + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start": ("STRING", {"default": None}), + }, + } + + CATEGORY = "badger" + RETURN_TYPES = () + FUNCTION = "garbage_collect" + OUTPUT_NODE = True + def garbage_collect(self, start): + gc.collect() + + NODE_CLASS_MAPPINGS = { "ImageOverlap-badger": ImageOverlap, "FloatToInt-badger": FloatToInt, @@ -1097,6 +1123,7 @@ NODE_CLASS_MAPPINGS = { "GetUUID-badger": GetUUID, "GetDirName-badger": GetDirName, "IdentifyLinesBasedOnBorderColor-badger": IdentifyLinesBasedOnBorderColor, + "GarbageCollect-badger": GarbageCollect, } diff --git a/color_editor.py b/color_editor.py new file mode 100644 index 0000000..53f89a1 --- /dev/null +++ b/color_editor.py @@ -0,0 +1,110 @@ +from collections import defaultdict +import numpy as np +from PIL import Image + + +def rgb_to_hex(rgb_colr): + return '#{:02x}{:02x}{:02x}'.format(*rgb_colr) + + +def hex_to_rgb(hex_color): + hex_color = hex_color.lstrip('#') + return tuple(int(hex_color[i:i + 2], 16) for i in (0, 2, 4)) + + +def get_colors(PIL_img, n): + color_list = [] + img = PIL_img.convert('RGBA') # 确保图片是RGBA模式 + + # 获取图片尺寸 + width, height = img.size + + for y in range(height): + count = 0 + # 从左到右扫描 + for x in range(width): + r, g, b, a = img.getpixel((x, y)) + if a != 0: + count += 1 + if count <= n: + color = (r, g, b) + color_list.append(rgb_to_hex(color)) + else: + count = 0 + count = 0 + # 从右到左扫描 + for x in range(width - 1, -1, -1): + r, g, b, a = img.getpixel((x, y)) + if a != 0: + count += 1 + if count <= n: + color = (r, g, b) + color_list.append(rgb_to_hex(color)) + else: + count = 0 + + + return color_list + + +def color_distance(c1, c2): + (r1, g1, b1) = c1 + (r2, g2, b2) = c2 + return np.sqrt((r1 - r2) ** 2 + (g1 - g2) ** 2 + (b1 - b2) ** 2) + + +def average_color(colors): + r = int(np.mean([c[0] for c in colors])) + g = int(np.mean([c[1] for c in colors])) + b = int(np.mean([c[2] for c in colors])) + return f"{r:02x}{g:02x}{b:02x}" + + +def fuzzy_color_grouping(colors, threshold): + groups = defaultdict(list) + + for color in colors: + rgb = hex_to_rgb(color) + placed = False + + for group_color in groups: + if color_distance(rgb, hex_to_rgb(group_color)) < threshold: + groups[group_color].append(rgb) + placed = True + break + + if not placed: + groups[color].append(rgb) + + return groups + + +def most_common_fuzzy_color(colors, threshold): + groups = fuzzy_color_grouping(colors, threshold) + largest_group = max(groups, key=lambda k: len(groups[k])) + return average_color(groups[largest_group]) + + +def is_color_similar(color1, color2, threshold): + return all(abs(c1 - c2) <= threshold for c1, c2 in zip(color1, color2)) + + +def find_similar_colors(image, color_string, threshold): + # 转换颜色字符串为RGB元组 + target_color = tuple(int(color_string[i:i + 2], 16) for i in (0, 2, 4)) + + # 创建一个同样大小的黑色背景图像 + output_image = Image.new('RGB', image.size, (0, 0, 0)) + pixels = image.load() + output_pixels = output_image.load() + + # 遍历每个像素点,检查颜色是否接近目标颜色 + for x in range(image.width): + for y in range(image.height): + if is_color_similar(pixels[x, y], target_color, threshold): + # 将接近的颜色设置为白色 + output_pixels[x, y] = (255, 255, 255) + + return output_image + + diff --git a/line_editor.py b/line_editor.py new file mode 100644 index 0000000..ac6df21 --- /dev/null +++ b/line_editor.py @@ -0,0 +1,129 @@ +from PIL import Image +from collections import deque + +def draw_line(pixels, x0, y0, x1, y1): + """Draw a white line from (x0, y0) to (x1, y1) on the provided pixels map.""" + dx = abs(x1 - x0) + dy = abs(y1 - y0) + sx = 1 if x0 < x1 else -1 + sy = 1 if y0 < y1 else -1 + err = dx - dy + + while True: + pixels[x0, y0] = 255 + if x0 == x1 and y0 == y1: + break + e2 = 2 * err + if e2 > -dy: + err -= dy + x0 += sx + if e2 < dx: + err += dx + y0 += sy + +def fill_white_segments(original_image, low_threshold, high_threshold): + # Load the original image and convert it to grayscale + original_image = original_image.convert('L') + original_pixels = original_image.load() + width, height = original_image.size + + low_threshold = int(width*low_threshold) + high_threshold = int(width*high_threshold) + + # Create a new black image to draw the lines + new_image = Image.new('L', (width, height), 0) + new_pixels = new_image.load() + + # Scan horizontally + for y in range(height): + point_a = None + for x in range(width): + if original_pixels[x, y] == 255: + if point_a is None: + point_a = (x, y) + else: + if x - point_a[0] < high_threshold and x - point_a[0] > low_threshold : + draw_line(new_pixels, point_a[0], point_a[1], x, y) + point_a = (x, y) + else: + point_a = (x, y) + + + # Scan vertically + for x in range(width): + point_a = None + for y in range(height): + if original_pixels[x, y] == 255: + if point_a is None: + point_a = (x, y) + else: + if y - point_a[1] < high_threshold and y - point_a[1] > low_threshold: + draw_line(new_pixels, point_a[0], point_a[1], x, y) + point_a = (x, y) + else: + point_a = (x, y) + + # Scan diagonally (top-left to bottom-right) + for diag in range(-height + 1, width): + point_a = None + for y in range(max(-diag, 0), min(width - diag, height)): + x = y + diag + if original_pixels[x, y] == 255: + if point_a is None: + point_a = (x, y) + else: + if max(abs(x - point_a[0]), abs(y - point_a[1])) < high_threshold and max(abs(x - point_a[0]), abs(y - point_a[1])) > low_threshold: + draw_line(new_pixels, point_a[0], point_a[1], x, y) + point_a = (x, y) + else: + point_a = (x, y) + + # Scan diagonally (top-right to bottom-left) + for diag in range(0, width + height): + point_a = None + for y in range(max(diag - width + 1, 0), min(diag + 1, height)): + x = diag - y + if original_pixels[x, y] == 255: + if point_a is None: + point_a = (x, y) + else: + if max(abs(x - point_a[0]), abs(y - point_a[1])) < high_threshold and max(abs(x - point_a[0]), abs(y - point_a[1])) > low_threshold: + draw_line(new_pixels, point_a[0], point_a[1], x, y) + point_a = (x, y) + else: + point_a = (x, y) + + # Save the new image with only the drawn lines + return new_image + +def find_largest_white_component(image): + width, height = image.size + visited = set() + largest_component = [] + largest_size = 0 + + def bfs(x, y): + queue = deque([(x, y)]) + local_visited = set() + while queue: + x, y = queue.popleft() + if (x, y) not in visited and 0 <= x < width and 0 <= y < height and image.getpixel((x, y)) == 255: + visited.add((x, y)) + local_visited.add((x, y)) + queue.extend([(x+1, y), (x-1, y), (x, y+1), (x, y-1)]) + return local_visited + + for y in range(height): + for x in range(width): + if image.getpixel((x, y)) == 255 and (x, y) not in visited: + component = bfs(x, y) + if len(component) > largest_size: + largest_size = len(component) + largest_component = component + + # 创建一个新的图像来绘制最大的白色像素点整体 + new_image = Image.new('1', image.size) + for x, y in largest_component: + new_image.putpixel((x, y), 255) + + return new_image