diff --git a/__init__.py b/__init__.py index b75da80..80acd8c 100644 --- a/__init__.py +++ b/__init__.py @@ -8,4 +8,4 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "PaletteTransfer": "Palette Transfer", "ColorPalette": "Color Palette", -} +} \ No newline at end of file diff --git a/color_transfer.py b/color_transfer.py index 4b7eb87..c1f2473 100644 --- a/color_transfer.py +++ b/color_transfer.py @@ -1,22 +1,44 @@ import numpy as np -from sklearn.cluster import KMeans +from sklearn.cluster import KMeans, MiniBatchKMeans import torch import ast -def ColorClustering(image, k): +def EuclideanDistance(current_colors, target_colors): + return np.linalg.norm(current_colors - target_colors, axis=1) + + +def ManhattanDistance(current_colors, target_colors): + return np.sum(np.abs(current_colors - target_colors), axis=1) + + +def ColorClustering(image, k, cluster_method): img_array = image.reshape((image.shape[0] * image.shape[1], 3)) - kmeans = KMeans(n_clusters=k) + + cluster_methods = { + "Kmeans": KMeans, + "Mini batch Kmeans": MiniBatchKMeans + } + + kmeans = cluster_methods.get(cluster_method)(n_clusters=k) + kmeans.fit(img_array) main_colors = kmeans.cluster_centers_ - return image, main_colors.astype(int), kmeans -def SwitchColors(image, current_colors, target_colors, kmeans): +def SwitchColors(image, current_colors, target_colors, kmeans, distance_method): closest_colors = [] + + distance_methods = { + "Euclidean": EuclideanDistance, + "Manhattan": ManhattanDistance + } + + distance_method = distance_methods.get(distance_method) + for color in current_colors: - distances = np.linalg.norm(target_colors - color, axis=1) + distances = distance_method(color, target_colors) closest_color = target_colors[np.argmin(distances)] closest_colors.append(closest_color) closest_colors = np.array(closest_colors) @@ -34,7 +56,9 @@ class PaletteTransferNode: data_in = { "required": { "image": ("IMAGE",), - "colors": ("COLORS",) + "colors": ("COLORS",), + "cluster_method": (["Kmeans","Mini batch Kmeans"], {'default': 'Kmeans'}, ), + "distance_method": (["Euclidean", "Manhattan"], {'default': 'Euclidean'}, ) } } return data_in @@ -44,7 +68,7 @@ class PaletteTransferNode: CATEGORY = "Palette Transfer" - def color_transfer(self, image, colors): + def color_transfer(self, image, colors, cluster_method, distance_method): if len(colors) == 0: return (image,) @@ -54,8 +78,8 @@ class PaletteTransferNode: for image in image: img = 255. * image.cpu().numpy() - img, current_colors, kmeans = ColorClustering(img, len(colors)) - processed = SwitchColors(img, current_colors, colors, kmeans) + img, current_colors, kmeans = ColorClustering(img, len(colors), cluster_method) + processed = SwitchColors(img, current_colors, colors, kmeans, distance_method) processedImages.append(processed) output = torch.cat(processedImages, dim=0) @@ -76,4 +100,4 @@ class ColorPaletteNode: FUNCTION = "color_list" def color_list(self, colors): - return (ast.literal_eval(colors), ) \ No newline at end of file + return (ast.literal_eval(colors), )