diff --git a/__init__.py b/__init__.py index 27144e2..a4f9427 100644 --- a/__init__.py +++ b/__init__.py @@ -4,22 +4,24 @@ try: except ImportError: pass else: + WEB_DIRECTORY = "./web" NODE_CLASS_MAPPINGS = {} # main nodes (no deps) from .nodes.mod import NODE_CLASS_MAPPINGS as mod_nodes NODE_CLASS_MAPPINGS.update(mod_nodes) - # 10bit PNG nodes + # pypng dep try: import png except ImportError: print("ColorMod: Can't find pypng! Please install to enable 16bit image support.") else: - from .nodes.save import NODE_CLASS_MAPPINGS as save_nodes - NODE_CLASS_MAPPINGS.update(save_nodes) + # 10bit PNG nodes + from .nodes.save_png import NODE_CLASS_MAPPINGS as save_png_nodes + NODE_CLASS_MAPPINGS.update(save_png_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/nodes/mod.py b/nodes/mod.py index b531743..db2b80b 100644 --- a/nodes/mod.py +++ b/nodes/mod.py @@ -9,41 +9,67 @@ class ColorModCompress: return { "required": { "image": ("IMAGE",), + "mode": (["clip", "normalize", "compress"],) } } RETURN_TYPES = ("IMAGE",) FUNCTION = "mod_compress" - CATEGORY = "image/postprocessing" - TITLE = "ColorMod (compress to [0;1])" + CATEGORY = "ColorMod" + TITLE = "ColorMod (compress)" - 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 + def mod_compress(self, image, mode): + image = image.clone() + if mode == "clip": + out = torch.clip(image, 0.0, 1.0) + elif mode == "normalize": + out = [] + for img in image: + img_min = torch.min(image) + img_max = torch.max(image) + print(f"Normalizing [{img_min:6.4f}:{img_max:6.4f}] => [0.0;1.0]") + img = (img - img_min) / (img_max - img_min) + out.append(img) + out = torch.stack(out, dim=0) + elif mode == "compress": + out = [] + for img in image: + ll = torch.minimum(image, torch.zeros(image.shape)) + ll = torch.clip(torch.abs(ll), 0.0, 1.0) + hh = torch.maximum(image, torch.ones(image.shape)) + hh = torch.clip((torch.abs(hh)-1.0), 0.0, 1.0) + out.append(ll + hh) + out = torch.stack(out, dim=0) + else: + raise ValueError(f"Unknown mode '{mode}'") # self.check_range(image) # debug - out = torch.clip(image, 0.0, 1.0) # sanity + out = torch.clip(out, 0.0, 1.0) # sanity + return (out,) + +class ColorModMove: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "move": ("FLOAT", {"default": 0.0, "min": -1.000, "max": 1.000, "step": 0.01}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "mod_move" + CATEGORY = "ColorMod" + TITLE = "ColorMod (move)" + + def mod_move(self, image, move): + image = image.clone() + + move_map = torch.ones(image.shape) * move + out = torch.clip((image + move_map), -4.0, 4.0) return (out,) class ColorModPivot: @@ -62,10 +88,12 @@ class ColorModPivot: RETURN_TYPES = ("IMAGE",) FUNCTION = "mod_pivot" - CATEGORY = "image/postprocessing" + CATEGORY = "ColorMod" TITLE = "ColorMod (move pivot)" def mod_pivot(self, image, pivot, move): + image = image.clone() + pivot_map = torch.ones(image.shape) * pivot image_high = torch.maximum(image, pivot_map) - pivot image_low = torch.minimum(image, pivot_map) @@ -85,17 +113,19 @@ class ColorModEdges: "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}), + "high": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), } } RETURN_TYPES = ("IMAGE",) FUNCTION = "mod_edges" - CATEGORY = "image/postprocessing" + CATEGORY = "ColorMod" TITLE = "ColorMod (edges)" def mod_edges(self, image, low, pivot, high): + image = image.clone() + pivot_map = torch.ones(image.shape) * pivot image_high = torch.maximum(image, pivot_map) - pivot image_low = torch.minimum(image, pivot_map) @@ -107,6 +137,7 @@ class ColorModEdges: NODE_CLASS_MAPPINGS = { "ColorModCompress": ColorModCompress, + "ColorModMove" : ColorModMove, "ColorModPivot": ColorModPivot, "ColorModEdges": ColorModEdges, } diff --git a/nodes/save.py b/nodes/save_png.py similarity index 93% rename from nodes/save.py rename to nodes/save_png.py index 806fc2c..243089e 100644 --- a/nodes/save.py +++ b/nodes/save_png.py @@ -34,7 +34,7 @@ def get_PIL_tEXt(image, prompt, extra_pnginfo): img.save(tmp, "png", pnginfo=metadata, compress_level=0) tmp.seek(0) - # read it back and PNG and get the tEXt chunks + # read it back as PNG and get the tEXt chunks img=png.Reader(tmp) metadata = [x for x in img.chunks() if x[0] == b"tEXt"] return metadata @@ -67,6 +67,8 @@ def save_png(image, extra_chunks, path): class SaveImageHighPrec(SaveImage): TITLE = "Save Image (16 bit)" + CATEGORY = "ColorMod" + def save_images(self, images, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): filename_prefix += self.prefix_append full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0]) @@ -90,6 +92,8 @@ class SaveImageHighPrec(SaveImage): # Directly copied from nodes.py class PreviewImageHighPrec(SaveImageHighPrec): TITLE = "Preview Image (16 bit)" + CATEGORY = "ColorMod" + def __init__(self): self.output_dir = folder_paths.get_temp_directory() self.type = "temp" @@ -112,6 +116,7 @@ class PreviewImageHighPrec(SaveImageHighPrec): class LoadImageHighPrec(LoadImage): TITLE = "Load Image (16 bit)" FUNCTION = "load_image_high_precision" + CATEGORY = "ColorMod" def load_image_high_precision(self, image): if not image.endswith(".png"): diff --git a/web/ColorMod.js b/web/ColorMod.js index 7123e11..c360f15 100644 --- a/web/ColorMod.js +++ b/web/ColorMod.js @@ -23,19 +23,43 @@ function addCanvasWidget(node) { return widget; } -function drawCanvasInitial(ctx, width, height, pv_x, pv_y) { +function drawCanvasInitial(ctx, width, height, pv_x=null, pv_y=null) { // clear ctx.clearRect(0, 0, width, height) + if (pv_x && pv_y) { + // set style + ctx.beginPath() + ctx.lineWidth = 5 + ctx.strokeStyle = "#444" + // draw cross + ctx.moveTo(0, pv_y) + ctx.lineTo(width, pv_y) + ctx.moveTo(pv_x, 0) + ctx.lineTo(pv_x, height) + ctx.stroke() + } +} - // draw cross +function drawCanvasCMMove(node) { + let move = getNamedWidget(node, "move").value + let canvas = getNamedWidget(node, "canvas").canvas + + let ctx = canvas.getContext("2d") + + // Calc coords + let width = canvas.width + let height = canvas.height + let offset = height * -move + + // clear + drawCanvasInitial(ctx, width, height) + // set style 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.lineWidth = 15 + ctx.strokeStyle = canvas.stcolor + // draw lines + ctx.moveTo(0, height + offset) // start + ctx.lineTo(width, offset) // start -> end ctx.stroke() } @@ -54,12 +78,11 @@ function drawCanvasCMMovePivot(node) { // clear drawCanvasInitial(ctx, width, height, pv_x, pv_y) - - // draw lines + // set style ctx.beginPath() ctx.lineWidth = 15 ctx.strokeStyle = canvas.stcolor - + // draw lines ctx.moveTo(0, height) // start ctx.lineTo(pv_x, pv_y) // start -> pivot ctx.lineTo(width, 0) // pivot -> end @@ -84,12 +107,11 @@ function drawCanvasCMEdges(node) { // clear drawCanvasInitial(ctx, width, height, pv_x, pv_y) - - // draw lines + // set style ctx.beginPath() ctx.lineWidth = 15 ctx.strokeStyle = canvas.stcolor - + // draw lines ctx.moveTo(0, height * low) // start ctx.lineTo(pv_x, pv_y) // start -> pivot ctx.lineTo(width, height * (1.0-high)) // pivot -> end @@ -99,6 +121,12 @@ function drawCanvasCMEdges(node) { app.registerExtension({ name: "City96.ColorMod", nodeCreated(node, app) { + if (node.__proto__.comfyClass == "ColorModMove") { + var widget = addCanvasWidget(node) + var refresh = function(v=null) { drawCanvasCMMove(node) } + getNamedWidget(node, "move").callback = refresh + setTimeout(refresh, 100); + } if (node.__proto__.comfyClass == "ColorModPivot") { var widget = addCanvasWidget(node) var refresh = function(v=null) { drawCanvasCMMovePivot(node) } @@ -114,5 +142,8 @@ app.registerExtension({ getNamedWidget(node, "pivot").callback = refresh setTimeout(refresh, 100); } + if (node.__proto__.comfyClass == "ColorModExposureFusion") { + console.log(node) + } } })