From 079639608b57893a2ec44218bcb1d945530d4567 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Sun, 31 Mar 2024 23:36:24 +0200 Subject: [PATCH] web widget, refractor --- __init__.py | 37 +++++------ colormod.py | 65 ------------------- nodes/mod.py | 112 +++++++++++++++++++++++++++++++++ highprec.py => nodes/save.py | 29 ++++++--- requirements.txt | 2 + web/ColorMod.js | 118 +++++++++++++++++++++++++++++++++++ 6 files changed, 266 insertions(+), 97 deletions(-) delete mode 100644 colormod.py create mode 100644 nodes/mod.py rename highprec.py => nodes/save.py (85%) create mode 100644 requirements.txt create mode 100644 web/ColorMod.js diff --git a/__init__.py b/__init__.py index 315c169..27144e2 100644 --- a/__init__.py +++ b/__init__.py @@ -4,31 +4,22 @@ try: except ImportError: pass else: - from .colormod import ColorModPivot, ColorModEdges + NODE_CLASS_MAPPINGS = {} - NODE_CLASS_MAPPINGS = { - "ColorModPivot": ColorModPivot, - "ColorModEdges": ColorModEdges, - } - NODE_DISPLAY_NAME_MAPPINGS = { - "ColorModPivot": ColorModPivot.TITLE, - "ColorModEdges": ColorModEdges.TITLE, - } + # main nodes (no deps) + from .nodes.mod import NODE_CLASS_MAPPINGS as mod_nodes + NODE_CLASS_MAPPINGS.update(mod_nodes) + + # 10bit PNG nodes try: import png except ImportError: - print("Can't find pypng! Please install to enable 16bit image support.") - pass + print("ColorMod: Can't find pypng! Please install to enable 16bit image support.") else: - from .highprec import SaveImageHighPrec, PreviewImageHighPrec, LoadImageHighPrec - NODE_CLASS_MAPPINGS.update({ - "SaveImageHighPrec": SaveImageHighPrec, - "PreviewImageHighPrec": PreviewImageHighPrec, - "LoadImageHighPrec": LoadImageHighPrec, - }) - - NODE_DISPLAY_NAME_MAPPINGS.update({ - "SaveImageHighPrec": SaveImageHighPrec.TITLE, - "PreviewImageHighPrec": PreviewImageHighPrec.TITLE, - "LoadImageHighPrec": LoadImageHighPrec.TITLE, - }) + from .nodes.save import NODE_CLASS_MAPPINGS as save_nodes + NODE_CLASS_MAPPINGS.update(save_nodes) + + # export + WEB_DIRECTORY = "./web" + NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} + __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/colormod.py b/colormod.py deleted file mode 100644 index 7894799..0000000 --- a/colormod.py +++ /dev/null @@ -1,65 +0,0 @@ -import os -import json -import torch -import numpy as np - - -class ColorModPivot: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image": ("IMAGE",), - "pivot": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), - "move": ("FLOAT", {"default": 0.0, "min": -2.0, "max": 2.0, "step": 0.01}), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "mod_pivot" - CATEGORY = "image/postprocessing" - TITLE = "ColorMod (move pivot)" - - def mod_pivot(self, image, pivot, move): - pivot_map = torch.ones(image.shape) * pivot - image_high = torch.maximum(image, pivot_map) - pivot - image_low = torch.minimum(image, pivot_map) - - image_high = image_high * (1/(1-pivot)) * (1-(pivot + move)) - image_low = image_low * (1/pivot) * (pivot + move) - out = torch.clip((image_high + image_low), 0.0, 1.0) - return (out,) - - -class ColorModEdges: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image": ("IMAGE",), - "low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), - "pivot": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), - "high": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "mod_edges" - CATEGORY = "image/postprocessing" - TITLE = "ColorMod (edges)" - - def mod_edges(self, image, low, pivot, high): - pivot_map = torch.ones(image.shape) * pivot - image_high = torch.maximum(image, pivot_map) - pivot - image_low = torch.minimum(image, pivot_map) - - image_low = image_low * low + pivot * (1-low) - image_high = image_high * high - out = torch.clip((image_high + image_low), 0.0, 1.0) - return (out,) diff --git a/nodes/mod.py b/nodes/mod.py new file mode 100644 index 0000000..b531743 --- /dev/null +++ b/nodes/mod.py @@ -0,0 +1,112 @@ +import torch + +class ColorModCompress: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "mod_compress" + CATEGORY = "image/postprocessing" + TITLE = "ColorMod (compress to [0;1])" + + def check_range(self, image): + image_min = torch.min(image) + image_max = torch.max(image) + if image_min < 0.0: + print("low trs hit") + image = image + image_min + + if image_max > 1.0: + print("high trs hit") + image = image * (1 / image_max) + + new_max = torch.max(image) + if new_max > 1.0: + print("high due to low trs hit") + image = image * (1 / new_max) + + print("MCX",image_min,image_max,new_max) + + def mod_compress(self, image): + ll = torch.minimum(image, torch.zeros(image.shape)) + ll = torch.clip(torch.abs(ll) * 8, 0.0, 1.0) + hh = torch.maximum(image, torch.ones(image.shape)) + hh = torch.clip((torch.abs(hh)-1.0) * 8, 0.0, 1.0) + image = ll + hh + + # self.check_range(image) # debug + out = torch.clip(image, 0.0, 1.0) # sanity + return (out,) + +class ColorModPivot: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "pivot": ("FLOAT", {"default": 0.5, "min": 0.001, "max": 0.999, "step": 0.01}), + "move": ("FLOAT", {"default": 0.0, "min": -2.000, "max": 2.000, "step": 0.01}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "mod_pivot" + CATEGORY = "image/postprocessing" + TITLE = "ColorMod (move pivot)" + + def mod_pivot(self, image, pivot, move): + pivot_map = torch.ones(image.shape) * pivot + image_high = torch.maximum(image, pivot_map) - pivot + image_low = torch.minimum(image, pivot_map) + + image_high = image_high * (1/(1-pivot)) * (1-(pivot + move)) + image_low = image_low * (1/pivot) * (pivot + move) + out = torch.clip((image_high + image_low), 0.0, 1.0) + return (out,) + +class ColorModEdges: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), + "high": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), + "pivot": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "mod_edges" + CATEGORY = "image/postprocessing" + TITLE = "ColorMod (edges)" + + def mod_edges(self, image, low, pivot, high): + pivot_map = torch.ones(image.shape) * pivot + image_high = torch.maximum(image, pivot_map) - pivot + image_low = torch.minimum(image, pivot_map) + + image_low = image_low * low + pivot * (1-low) + image_high = image_high * high + out = torch.clip((image_high + image_low), 0.0, 1.0) + return (out,) + +NODE_CLASS_MAPPINGS = { + "ColorModCompress": ColorModCompress, + "ColorModPivot": ColorModPivot, + "ColorModEdges": ColorModEdges, +} diff --git a/highprec.py b/nodes/save.py similarity index 85% rename from highprec.py rename to nodes/save.py index f15d2e5..806fc2c 100644 --- a/highprec.py +++ b/nodes/save.py @@ -12,7 +12,6 @@ import folder_paths from comfy.cli_args import args from nodes import SaveImage, PreviewImage, LoadImage - def get_PIL_tEXt(image, prompt, extra_pnginfo): """This is extremely stupid""" # prepare PIL image as normal @@ -40,7 +39,6 @@ def get_PIL_tEXt(image, prompt, extra_pnginfo): metadata = [x for x in img.chunks() if x[0] == b"tEXt"] return metadata - def save_png(image, extra_chunks, path): i = image.cpu().numpy() img = np.clip(65535.0*i, 0, 65535).astype(np.uint16) @@ -95,14 +93,21 @@ class PreviewImageHighPrec(SaveImageHighPrec): def __init__(self): self.output_dir = folder_paths.get_temp_directory() self.type = "temp" - self.prefix_append = "_temp_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + self.prefix_append = "_temp_" + ''.join( + random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5) + ) @classmethod def INPUT_TYPES(s): - return {"required": - {"images": ("IMAGE", ), }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, - } + return { + "required": { + "images": ("IMAGE",) + }, + "hidden": { + "prompt": "PROMPT", + "extra_pnginfo": "EXTRA_PNGINFO" + } + } class LoadImageHighPrec(LoadImage): TITLE = "Load Image (16 bit)" @@ -117,11 +122,11 @@ class LoadImageHighPrec(LoadImage): reader = png.Reader(image_path) raw = reader.read() - image = np.vstack(map(np.uint16, raw[2])) + image = np.vstack(list(map(np.uint16, raw[2]))) dim_rgb = image.shape[1] // raw[0] div_max = 1.0 if np.max(image) <= 255 else 256.0 - + image = np.reshape(image,(raw[1], raw[0], dim_rgb)) image = np.array(image).astype(np.float32) / (255.0 * div_max ) image = torch.from_numpy(np.clip(image, 0.0, 1.0))[None,] @@ -132,3 +137,9 @@ class LoadImageHighPrec(LoadImage): else: mask = torch.zeros((64,64), dtype=torch.float32, device="cpu").unsqueeze(0) return (image, mask) + +NODE_CLASS_MAPPINGS = { + "SaveImageHighPrec": SaveImageHighPrec, + "PreviewImageHighPrec": PreviewImageHighPrec, + "LoadImageHighPrec": LoadImageHighPrec, +} diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..4751b85 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +pypng +opencv-contrib-python diff --git a/web/ColorMod.js b/web/ColorMod.js new file mode 100644 index 0000000..7123e11 --- /dev/null +++ b/web/ColorMod.js @@ -0,0 +1,118 @@ +import { app } from "/scripts/app.js"; + +function getNamedWidget(node, name) { + return node.widgets.find(function(w){return w.name == name}) +} + +function addCanvasWidget(node) { + const canvas = document.createElement("canvas"); + canvas.width = 1000 + canvas.height = 400 + + canvas.style.pointerEvents = "none"; + canvas.style.border = "1px solid " + LiteGraph.WIDGET_OUTLINE_COLOR + canvas.style.backgroundColor = LiteGraph.WIDGET_BGCOLOR + canvas.stcolor = app.canvas.default_connection_color_byType.IMAGE + + let opts = { + getMinHeight() { return 100 }, + selectOn: [], + } + const widget = node.addDOMWidget("canvas", "CTYCanvas", canvas, opts); + widget.canvas = canvas; + return widget; +} + +function drawCanvasInitial(ctx, width, height, pv_x, pv_y) { + // clear + ctx.clearRect(0, 0, width, height) + + // draw cross + ctx.beginPath() + ctx.lineWidth = 5 + ctx.strokeStyle = "#444" + + ctx.moveTo(0, pv_y) + ctx.lineTo(width, pv_y) + ctx.moveTo(pv_x, 0) + ctx.lineTo(pv_x, height) + ctx.stroke() +} + +function drawCanvasCMMovePivot(node) { + let move = getNamedWidget(node, "move").value + let pivot = getNamedWidget(node, "pivot").value + let canvas = getNamedWidget(node, "canvas").canvas + + let ctx = canvas.getContext("2d") + + // Calc pivot coords + let width = canvas.width + let height = canvas.height + let pv_x = width * pivot + let pv_y = height * (1.0 - pivot) - (height * move) + + // clear + drawCanvasInitial(ctx, width, height, pv_x, pv_y) + + // draw lines + ctx.beginPath() + ctx.lineWidth = 15 + ctx.strokeStyle = canvas.stcolor + + ctx.moveTo(0, height) // start + ctx.lineTo(pv_x, pv_y) // start -> pivot + ctx.lineTo(width, 0) // pivot -> end + ctx.stroke() +} + +function drawCanvasCMEdges(node) { + let low = getNamedWidget(node, "low").value + let high = getNamedWidget(node, "high").value + let pivot = getNamedWidget(node, "pivot").value + let canvas = getNamedWidget(node, "canvas").canvas + + let ctx = canvas.getContext("2d") + ctx.lineWidth = 15 + ctx.strokeStyle = canvas.stcolor + + // Calc pivot coords + let width = canvas.width + let height = canvas.height + let pv_x = width * pivot + let pv_y = height * (1.0 - pivot) + + // clear + drawCanvasInitial(ctx, width, height, pv_x, pv_y) + + // draw lines + ctx.beginPath() + ctx.lineWidth = 15 + ctx.strokeStyle = canvas.stcolor + + ctx.moveTo(0, height * low) // start + ctx.lineTo(pv_x, pv_y) // start -> pivot + ctx.lineTo(width, height * (1.0-high)) // pivot -> end + ctx.stroke() +} + +app.registerExtension({ + name: "City96.ColorMod", + nodeCreated(node, app) { + if (node.__proto__.comfyClass == "ColorModPivot") { + var widget = addCanvasWidget(node) + var refresh = function(v=null) { drawCanvasCMMovePivot(node) } + getNamedWidget(node, "move").callback = refresh + getNamedWidget(node, "pivot").callback = refresh + setTimeout(refresh, 100); + } + if (node.__proto__.comfyClass == "ColorModEdges") { + var widget = addCanvasWidget(node) + var refresh = function(v) { drawCanvasCMEdges(node) } + getNamedWidget(node, "low").callback = refresh + getNamedWidget(node, "high").callback = refresh + getNamedWidget(node, "pivot").callback = refresh + setTimeout(refresh, 100); + } + } +})