From a16c5c8bdf5376a105340e06769c9b5dd96ba6dd Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 28 Sep 2023 21:31:30 +0200 Subject: [PATCH] Initial commit --- .gitignore | 7 ++++ __init__.py | 32 ++++++++++++++++ colormod.py | 66 +++++++++++++++++++++++++++++++++ highprec.py | 104 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 209 insertions(+) create mode 100644 __init__.py create mode 100644 colormod.py create mode 100644 highprec.py diff --git a/.gitignore b/.gitignore index 68bc17f..32bd64e 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,10 @@ +research/ +other/ +test.py +*.png + +# default github .gitignore follows + # Byte-compiled / optimized / DLL files __pycache__/ *.py[cod] diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..a774176 --- /dev/null +++ b/__init__.py @@ -0,0 +1,32 @@ +# only import if running as a custom node +try: + import comfy.utils +except ImportError: + pass +else: + from .colormod import ColorModPivot, ColorModEdges + + NODE_CLASS_MAPPINGS = { + "ColorModPivot": ColorModPivot, + "ColorModEdges": ColorModEdges, + } + NODE_DISPLAY_NAME_MAPPINGS = { + "ColorModPivot": ColorModPivot.TITLE, + "ColorModEdges": ColorModEdges.TITLE, + } + try: + import png + except ImportError: + print("Can't find pypng! Please install to enable 16bit image support.") + pass + else: + from .highprec import SaveImageHighPrec, PreviewImageHighPrec + NODE_CLASS_MAPPINGS.update({ + "SaveImageHighPrec": SaveImageHighPrec, + "PreviewImageHighPrec": PreviewImageHighPrec, + }) + + NODE_DISPLAY_NAME_MAPPINGS.update({ + "SaveImageHighPrec": SaveImageHighPrec.TITLE, + "PreviewImageHighPrec": PreviewImageHighPrec.TITLE, + }) diff --git a/colormod.py b/colormod.py new file mode 100644 index 0000000..a6086b7 --- /dev/null +++ b/colormod.py @@ -0,0 +1,66 @@ +import os +import png +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/highprec.py b/highprec.py new file mode 100644 index 0000000..aba22f0 --- /dev/null +++ b/highprec.py @@ -0,0 +1,104 @@ +import os +import png +import json +import random +import numpy as np +from io import BytesIO +from PIL import Image +from PIL.PngImagePlugin import PngInfo + +import folder_paths +from comfy.cli_args import args +from nodes import SaveImage, PreviewImage + + +def get_PIL_tEXt(image, prompt, extra_pnginfo): + """This is extremely stupid""" + # prepare PIL image as normal + print(image.dtype) + i = image.cpu().numpy() + img = np.clip(255.0*i, 0, 255).astype(np.uint8) + img = Image.fromarray(img) + + metadata = None + if not args.disable_metadata: + metadata = PngInfo() + if prompt is not None: + metadata.add_text("prompt", json.dumps(prompt)) + if extra_pnginfo is not None: + for x in extra_pnginfo: + metadata.add_text(x, json.dumps(extra_pnginfo[x])) + + # write temp PIL image + tmp = BytesIO() + img.save(tmp, "png", pnginfo=metadata, compress_level=0) + tmp.seek(0) + + # read it back and PNG and get the tEXt chunks + img=png.Reader(tmp) + 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) + + writer = png.Writer( + size = (img.shape[1],img.shape[0]), + bitdepth = 16, + greyscale = False, + compression = 9, + ) + data = img.reshape(-1, img.shape[1]*img.shape[2]).tolist() + # default writer without metadata + if not extra_chunks: + with open(path, "wb") as f: + writer.write(f, data) + return + # jank in the tEXt chunks as well + tmp = BytesIO() + writer.write(tmp, data) + tmp.seek(0) + chunks = list(png.Reader(tmp).chunks()) + for k in extra_chunks: + chunks.insert(1, k) + with open(path, "wb") as f: + png.write_chunks(f, chunks) + +class SaveImageHighPrec(SaveImage): + TITLE = "Save Image (16 bit)" + 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]) + results = list() + for image in images: + metadata = get_PIL_tEXt(image, prompt, extra_pnginfo) + + file = f"{filename}_{counter:05}_.png" + path = os.path.join(full_output_folder, file) + save_png(image, metadata, path) + + results.append({ + "filename": file, + "subfolder": subfolder, + "type": self.type + }) + counter += 1 + + return { "ui": { "images": results } } + +# Directly copied from nodes.py +class PreviewImageHighPrec(SaveImageHighPrec): + TITLE = "Preview Image (16 bit)" + 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)) + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"images": ("IMAGE", ), }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + }