Added MiniBatchKMeans, color distance methods

This commit is contained in:
A,Y
2024-09-14 01:02:31 +03:00
parent b387f63534
commit fbeb993c46
2 changed files with 36 additions and 12 deletions
+1 -1
View File
@@ -8,4 +8,4 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"PaletteTransfer": "Palette Transfer",
"ColorPalette": "Color Palette",
}
}
+35 -11
View File
@@ -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), )
return (ast.literal_eval(colors), )