added Color Transfer + image utils
This commit is contained in:
@@ -0,0 +1,91 @@
|
||||
import numpy as np
|
||||
import cv2
|
||||
from PIL import Image
|
||||
import torch
|
||||
|
||||
# Define the Lab Color Transfer Node
|
||||
class LabColorTransferNode:
|
||||
def __init__(self, device="cpu"):
|
||||
self.device = device
|
||||
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_image": ("IMAGE",),
|
||||
"hex_color": ("INT",), # Expecting a color input in hexadecimal integer format
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "apply_lab_color_transfer"
|
||||
CATEGORY = "Badman"
|
||||
|
||||
def hex_to_rgb(self, hex_value):
|
||||
"""
|
||||
Converts a hex integer (0xRRGGBB) to an (R, G, B) tuple.
|
||||
|
||||
Parameters:
|
||||
- hex_value: A hex integer representing the color.
|
||||
|
||||
Returns:
|
||||
- A tuple (R, G, B) with values in the range 0-255.
|
||||
"""
|
||||
# Extract R, G, B components from the hex integer
|
||||
r = (hex_value >> 16) & 0xFF
|
||||
g = (hex_value >> 8) & 0xFF
|
||||
b = hex_value & 0xFF
|
||||
print(r,g,b)
|
||||
return (r, g, b)
|
||||
|
||||
def apply_lab_color_transfer(self, input_image, hex_color):
|
||||
"""
|
||||
Transfers the input color (from hex) to the image using Lab color transfer,
|
||||
preserving luminance and applying color transformation to A and B channels.
|
||||
"""
|
||||
# Convert hex color to RGB tuple
|
||||
target_color = self.hex_to_rgb(hex_color)
|
||||
|
||||
# Ensure the input is a PyTorch tensor, and convert to NumPy
|
||||
if isinstance(input_image, torch.Tensor):
|
||||
input_image_np = input_image.cpu().numpy()
|
||||
else:
|
||||
raise TypeError("Input image must be a PyTorch tensor")
|
||||
|
||||
# Remove batch dimension if present (input shape is likely [1, H, W, 3])
|
||||
if input_image_np.shape[0] == 1:
|
||||
input_image_np = np.squeeze(input_image_np, axis=0) # Remove the batch dimension
|
||||
|
||||
# Now input_image_np should be in the format (H, W, 3) for RGB images
|
||||
# Ensure the image is in uint8 format (0-255 range)
|
||||
input_image_np = (input_image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert the NumPy array (input image) to Lab color space using OpenCV
|
||||
img_lab = cv2.cvtColor(input_image_np, cv2.COLOR_RGB2LAB)
|
||||
|
||||
# Split the image into L, A, and B channels
|
||||
L_channel, A_channel, B_channel = cv2.split(img_lab)
|
||||
|
||||
# Convert the target RGB color to Lab color space
|
||||
target_color_lab = cv2.cvtColor(np.uint8([[list(target_color)]]), cv2.COLOR_RGB2LAB)[0][0]
|
||||
target_A = target_color_lab[1] # A component of target color
|
||||
target_B = target_color_lab[2] # B component of target color
|
||||
|
||||
# Replace the A and B channels of the image with the target A and B values
|
||||
A_channel[:] = target_A
|
||||
B_channel[:] = target_B
|
||||
|
||||
# Merge the original L channel with the new A and B channels
|
||||
recolored_lab = cv2.merge([L_channel, A_channel, B_channel])
|
||||
|
||||
# Convert the recolored Lab image back to RGB
|
||||
recolored_rgb = cv2.cvtColor(recolored_lab, cv2.COLOR_LAB2RGB)
|
||||
|
||||
# Convert the result back to a PyTorch tensor
|
||||
# Convert the result back to a PyTorch tensor with the correct shape [batch, height, width, channels]
|
||||
recolored_rgb_tensor = torch.from_numpy(recolored_rgb / 255.0).float().unsqueeze(0) # Add back batch dimension
|
||||
print(recolored_rgb_tensor.shape)
|
||||
# Return only a single image (combined RGB image)
|
||||
return (recolored_rgb_tensor,)
|
||||
@@ -131,3 +131,71 @@ class HexGenerator:
|
||||
return (color_int,)
|
||||
|
||||
|
||||
import torch
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
|
||||
class RandomColorImageGrid:
|
||||
def __init__(self, device="cpu"):
|
||||
self.device = device
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"width": ("INT", {"default": 1024, "min": 1}),
|
||||
"height": ("INT", {"default": 1024, "min": 1}),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096}),
|
||||
"num_colors": ("INT", {"default": 4, "min": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate"
|
||||
|
||||
CATEGORY = "image"
|
||||
|
||||
def generate(self, width, height, batch_size=1, num_colors=4):
|
||||
# Seed the random number generator uniquely for each call
|
||||
random.seed(time.time() + random.randint(0, 10000))
|
||||
|
||||
# Calculate rows and columns based on number of colors
|
||||
rows = math.ceil(math.sqrt(num_colors))
|
||||
cols = math.ceil(num_colors / rows)
|
||||
|
||||
tile_width = width // cols
|
||||
tile_height = height // rows
|
||||
|
||||
# Create tensors for the R, G, B channels
|
||||
images = []
|
||||
for _ in range(batch_size):
|
||||
r = torch.zeros([height, width], dtype=torch.float32, device=self.device)
|
||||
g = torch.zeros([height, width], dtype=torch.float32, device=self.device)
|
||||
b = torch.zeros([height, width], dtype=torch.float32, device=self.device)
|
||||
|
||||
# Generate random colors and fill the tiles
|
||||
color_idx = 0
|
||||
for i in range(rows):
|
||||
for j in range(cols):
|
||||
if color_idx >= num_colors:
|
||||
break
|
||||
color_r = random.randint(0, 255) / 255.0
|
||||
color_g = random.randint(0, 255) / 255.0
|
||||
color_b = random.randint(0, 255) / 255.0
|
||||
|
||||
x_start, x_end = j * tile_width, (j + 1) * tile_width
|
||||
y_start, y_end = i * tile_height, (i + 1) * tile_height
|
||||
|
||||
r[y_start:y_end, x_start:x_end] = color_r
|
||||
g[y_start:y_end, x_start:x_end] = color_g
|
||||
b[y_start:y_end, x_start:x_end] = color_b
|
||||
|
||||
color_idx += 1
|
||||
|
||||
# Concatenate the R, G, B channels along the last dimension
|
||||
image = torch.stack([r, g, b], dim=-1)
|
||||
images.append(image)
|
||||
|
||||
# Return the batch of images
|
||||
return (torch.stack(images),)
|
||||
|
||||
Reference in New Issue
Block a user