diff --git a/color_transfer.py b/color_transfer.py index 7439d0f..27b72ea 100644 --- a/color_transfer.py +++ b/color_transfer.py @@ -3,8 +3,8 @@ from sklearn.cluster import KMeans, MiniBatchKMeans import torch import ast import cv2 -from .utils import EuclideanDistance, ManhattanDistance, CosineSimilarity, HSV_Color_Similarity, Blur - +from .utils import EuclideanDistance, ManhattanDistance, CosineSimilarity, Blur +from .utils import HSVColorSimilarity, RGBWeightedDistance, RGBWeightedSimilarity @@ -30,7 +30,9 @@ def SwitchColors(image, detected_colors, target_colors, clustering_model, distan "Euclidean": EuclideanDistance, "Manhattan": ManhattanDistance, "Cosine Similarity": CosineSimilarity, - "HSV Distance": HSV_Color_Similarity + "HSV Distance": HSVColorSimilarity, + "RGB Weighted Distance": RGBWeightedDistance, + "RGB Weighted Similarity": RGBWeightedSimilarity } distance_method = distance_methods.get(distance_method) @@ -56,7 +58,7 @@ class PaletteTransferNode: "target_colors": ("COLOR_LIST",), "color_space": (["RGB", "HSV", "LAB"], {'default': 'RGB'}), "cluster_method": (["Kmeans","Mini batch Kmeans"], {'default': 'Kmeans'}, ), - "distance_method": (["Euclidean", "Manhattan", "Cosine Similarity", "HSV Distance"], {'default': 'Euclidean'}, ), + "distance_method": (["Euclidean", "Manhattan", "Cosine Similarity", "HSV Distance", "RGB Weighted Distance", "RGB Weighted Similarity"], {'default': 'Euclidean'}, ), "gaussian_blur": ("INT", {'default': 3, 'min': 0, 'max': 27, 'step': 2}), } } diff --git a/utils.py b/utils.py index 4ec4763..7155ce5 100644 --- a/utils.py +++ b/utils.py @@ -14,7 +14,35 @@ def CosineSimilarity(detected_color, target_colors): return -np.dot(target_colors, detected_color) / (np.linalg.norm(detected_color) * np.linalg.norm(target_colors, axis=1)) -def HSV_Color_Similarity(detected_color, target_colors): +def RGBWeightedDistance(detected_color, target_colors): + detected_color = np.array(detected_color) + target_colors = np.array(target_colors) + + weights = np.array([0.299, 0.587, 0.114]) + + weighted_detected_color = np.dot(detected_color, weights) + weighted_target_colors = np.dot(target_colors, weights) + + return np.abs(weighted_detected_color - weighted_target_colors) + + +def RGBWeightedSimilarity(detected_color, target_colors): + detected_color = np.array(detected_color) + target_colors = np.array(target_colors) + + weights = np.array([0.299, 0.587, 0.114]) + + weighted_detected_color = np.dot(detected_color, weights) + weighted_target_colors = np.dot(target_colors, weights) + + dot_products = np.dot(weighted_detected_color, weighted_target_colors) + norm1 = np.linalg.norm(weighted_detected_color) + norm2 = np.linalg.norm(weighted_target_colors) + + return -dot_products / (norm1 * norm2) + + +def HSVColorSimilarity(detected_color, target_colors): detected_color = np.array(detected_color) target_colors = np.array(target_colors)