Added Adjustment node

This commit is contained in:
Fillip
2024-09-18 11:02:55 -07:00
parent f56831d1ca
commit 66e3b52cea
5 changed files with 226 additions and 26 deletions
+3
View File
@@ -68,6 +68,7 @@ from .nodes.FL_PDFImageExtractor import FL_PDFImageExtractor
from .nodes.FL_BulkPDFLoader import FL_BulkPDFLoader
from .nodes.FL_SaveAndDisplayImage import FL_SaveAndDisplayImage
from .nodes.FL_OllamaCaptioner import FL_OllamaCaptioner
from .nodes.FL_ImageAdjuster import FL_ImageAdjuster
@@ -143,6 +144,7 @@ NODE_CLASS_MAPPINGS = {
"FL_BulkPDFLoader": FL_BulkPDFLoader,
"FL_SaveAndDisplayImage": FL_SaveAndDisplayImage,
"FL_OllamaCaptioner": FL_OllamaCaptioner,
"FL_ImageAdjuster": FL_ImageAdjuster,
}
@@ -217,6 +219,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FL_BulkPDFLoader": "FL Bulk PDF Loader",
"FL_SaveAndDisplayImage": "FL Save And Display Image",
"FL_OllamaCaptioner": "FL Ollama Captioner by Cosmic",
"FL_ImageAdjuster": "FL_ImageAdjuster",
}
+9 -24
View File
@@ -7,6 +7,7 @@ import sys
from tqdm import tqdm
import base64
# removed api key from input for safer use
class FL_GPT_Vision:
@classmethod
def INPUT_TYPES(cls):
@@ -17,13 +18,12 @@ class FL_GPT_Vision:
"default": "You are a helpful assistant that describes images accurately and concisely.",
"multiline": True}),
"request_prompt": ("STRING", {"default": "Describe this image in detail.", "multiline": True}),
"output_directory": ("STRING", {"default": ""}),
"overwrite": ("BOOLEAN", {"default": False}),
"max_tokens": ("INT", {"default": 300, "min": 1, "max": 4096}),
"temperature": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.1}),
"detail": (["auto", "low", "high"],),
"batch_size": ("INT", {"default": 5, "min": 1, "max": 20}),
"upload_resolution": (["512", "768", "1024"],),
"output_directory": ("STRING", {"default": ""}),
},
"optional": {
"images": ("IMAGE",),
@@ -35,21 +35,9 @@ class FL_GPT_Vision:
RETURN_NAMES = ("message", "output_directory")
FUNCTION = "generate_captions"
CATEGORY = "🏵️Fill Nodes/GPT"
OUTPUT_NODE = True
def resize_image(self, img, target_size):
width, height = img.size
aspect_ratio = width / height
if width > height:
new_width = target_size
new_height = int(target_size / aspect_ratio)
else:
new_height = target_size
new_width = int(target_size * aspect_ratio)
return img.resize((new_width, new_height), Image.LANCZOS)
async def process_image(self, session, img, img_filename, output_directory, overwrite, api_key, model,
system_prompt, request_prompt, max_tokens, temperature, detail, upload_resolution):
system_prompt, request_prompt, max_tokens, temperature, detail):
caption_filename = os.path.splitext(img_filename)[0] + ".txt"
img_path = os.path.join(output_directory, img_filename)
caption_path = os.path.join(output_directory, caption_filename)
@@ -57,15 +45,12 @@ class FL_GPT_Vision:
if not overwrite and os.path.exists(caption_path):
return None
# Save the original image
# Save the image
img.save(img_path)
# Resize image for API upload
resized_img = self.resize_image(img, int(upload_resolution))
# Encode resized image to base64
# Encode image to base64
buffered = io.BytesIO()
resized_img.save(buffered, format="PNG")
img.save(buffered, format="PNG")
img_str = base64.b64encode(buffered.getvalue()).decode()
payload = {
@@ -116,8 +101,8 @@ class FL_GPT_Vision:
return await asyncio.gather(*tasks)
def generate_captions(self, model, system_prompt, request_prompt, output_directory, overwrite, max_tokens,
temperature, detail, batch_size, upload_resolution, images=None, input_directory=None):
api_key = os.getenv("OPENAI_API_KEY")
temperature, detail, batch_size, images=None, input_directory=None):
api_key = os.getenv("OPENAI_API_KEY") #looks for api key in env variable
try:
if not api_key:
raise ValueError("API key is not set as an environment variable")
@@ -155,7 +140,7 @@ class FL_GPT_Vision:
for batch in tqdm(batches, desc="Processing batches", file=sys.stdout):
batch_captions = await self.process_batch(batch, session, output_directory, overwrite, api_key,
model, system_prompt, request_prompt, max_tokens,
temperature, detail, upload_resolution)
temperature, detail)
all_captions.extend(batch_captions)
return all_captions
+83
View File
@@ -0,0 +1,83 @@
import torch
import numpy as np
from PIL import Image, ImageEnhance, ImageFilter
import base64
import io
from server import PromptServer
from .utils import tensor_to_pil, pil_to_tensor
class FL_ImageAdjuster:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"hue": ("FLOAT", {"default": 0, "min": -180, "max": 180, "step": 1}),
"saturation": ("FLOAT", {"default": 0, "min": -100, "max": 100, "step": 1}),
"brightness": ("FLOAT", {"default": 0, "min": -100, "max": 100, "step": 1}),
"contrast": ("FLOAT", {"default": 0, "min": -100, "max": 100, "step": 1}),
"sharpness": ("FLOAT", {"default": 0, "min": 0, "max": 100, "step": 1}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "adjust_image"
CATEGORY = "🏵️Fill Nodes/Image"
OUTPUT_NODE = True
def adjust_image(self, image, hue, saturation, brightness, contrast, sharpness):
# Convert tensor to PIL Image
pil_image = tensor_to_pil(image)
# Apply adjustments
adjusted_image = self.apply_adjustments(pil_image, hue, saturation, brightness, contrast, sharpness)
# Convert back to tensor
tensor_image = pil_to_tensor(adjusted_image)
# Prepare image for display
display_image = self.prepare_image_for_display(adjusted_image)
# Send the image to the frontend
PromptServer.instance.send_sync("fl_image_adjuster", {"image": display_image})
return (tensor_image,)
def apply_adjustments(self, image, hue, saturation, brightness, contrast, sharpness):
# Convert to HSV for hue and saturation adjustments
hsv_image = image.convert('HSV')
h, s, v = hsv_image.split()
# Hue adjustment
h = h.point(lambda x: (x + hue) % 256)
# Saturation adjustment
s = s.point(lambda x: max(0, min(255, x + saturation * 255 / 100)))
# Merge channels
hsv_image = Image.merge('HSV', (h, s, v))
rgb_image = hsv_image.convert('RGB')
# Brightness adjustment
enhancer = ImageEnhance.Brightness(rgb_image)
rgb_image = enhancer.enhance(1 + brightness / 100)
# Contrast adjustment
enhancer = ImageEnhance.Contrast(rgb_image)
rgb_image = enhancer.enhance(1 + contrast / 100)
# Sharpness adjustment
if sharpness > 0:
# Convert sharpness to an integer percentage between 100 and 200
sharpness_percent = int(100 + sharpness)
rgb_image = rgb_image.filter(ImageFilter.UnsharpMask(radius=2, percent=sharpness_percent, threshold=3))
return rgb_image
def prepare_image_for_display(self, pil_image):
# Convert PIL Image to base64 string
buffered = io.BytesIO()
pil_image.save(buffered, format="PNG")
img_str = base64.b64encode(buffered.getvalue()).decode()
return f"data:image/png;base64,{img_str}"
-2
View File
@@ -16,5 +16,3 @@ PyMuPDF
reportlab
PyPDF2
ollama
opencv-python
kornia
+131
View File
@@ -0,0 +1,131 @@
import { app } from "../../../scripts/app.js";
import { api } from "../../../scripts/api.js";
function hideWidgetForGood(node, widget) {
if (!widget) return;
widget.origType = widget.type;
widget.origComputeSize = widget.computeSize;
widget.origSerializeValue = widget.serializeValue;
widget.computeSize = () => [0, -4];
widget.type = "converted-widget";
}
app.registerExtension({
name: "Comfy.FL_ImageAdjuster",
async nodeCreated(node) {
if (node.comfyClass === "FL_ImageAdjuster") {
addImageAdjusterUI(node);
}
}
});
function addImageAdjusterUI(node) {
const MIN_WIDTH = 200;
const MIN_HEIGHT = 300;
const SLIDER_HEIGHT = 25;
const PADDING = 0;
// Default values for sliders
const DEFAULT_VALUES = {
hue: 0,
saturation: 0,
brightness: 0,
contrast: 0,
sharpness: 0
};
// Find and hide the original widgets
const originalWidgets = {};
["hue", "saturation", "brightness", "contrast", "sharpness"].forEach(name => {
const widget = node.widgets.find(w => w.name === name);
if (widget) {
originalWidgets[name] = widget;
hideWidgetForGood(node, widget);
}
});
// Create custom sliders
const sliders = [
createSlider("Hue", -180, 180, originalWidgets.hue),
createSlider("Saturation", -100, 100, originalWidgets.saturation),
createSlider("Brightness", -100, 100, originalWidgets.brightness),
createSlider("Contrast", -100, 100, originalWidgets.contrast),
createSlider("Sharpness", 0, 100, originalWidgets.sharpness)
];
function createSlider(name, min, max, originalWidget) {
const slider = node.addWidget("slider", name, originalWidget.value, (v) => {
originalWidget.value = v;
node.setDirtyCanvas(true);
}, { min: min, max: max, step: 1 });
return slider;
}
// Add reset button
const resetButton = node.addWidget("button", "Reset", null, () => {
resetSliders();
});
function resetSliders() {
sliders.forEach(slider => {
const defaultValue = DEFAULT_VALUES[slider.name.toLowerCase()];
slider.value = defaultValue;
slider.callback(defaultValue);
});
node.setDirtyCanvas(true);
node.triggerSlot(0);
}
// Add image preview
const img = new Image();
img.onload = () => node.setDirtyCanvas(true);
node.onDrawBackground = function(ctx) {
if (!this.flags.collapsed) {
const [w, h] = this.size;
// Calculate the Y position of the last widget
const lastWidget = node.widgets[node.widgets.length - 1];
const lastWidgetY = lastWidget.last_y || 0;
// Set the image Y offset to be just below the last widget
const IMAGE_Y_OFFSET = lastWidgetY + SLIDER_HEIGHT + PADDING;
const imageArea = h - IMAGE_Y_OFFSET - PADDING;
// Draw image
if (img.src) {
const aspectRatio = img.width / img.height;
let drawWidth = w - 2 * PADDING;
let drawHeight = imageArea;
if (drawWidth / drawHeight > aspectRatio) {
drawWidth = drawHeight * aspectRatio;
} else {
drawHeight = drawWidth / aspectRatio;
}
const x = PADDING + (w - 2 * PADDING - drawWidth) / 2;
const y = IMAGE_Y_OFFSET;
ctx.drawImage(img, x, y, drawWidth, drawHeight);
}
}
};
// Listen for the image from the backend
api.addEventListener("fl_image_adjuster", (event) => {
if (event.detail.image) {
img.src = event.detail.image;
}
});
function updateNodeSize() {
node.size[0] = Math.max(MIN_WIDTH, node.size[0]);
node.size[1] = Math.max(MIN_HEIGHT, node.size[1]);
}
node.onResize = updateNodeSize;
updateNodeSize();
}