diff --git a/__init__.py b/__init__.py index a774176..315c169 100644 --- a/__init__.py +++ b/__init__.py @@ -20,13 +20,15 @@ else: print("Can't find pypng! Please install to enable 16bit image support.") pass else: - from .highprec import SaveImageHighPrec, PreviewImageHighPrec + 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, }) diff --git a/highprec.py b/highprec.py index aba22f0..f15d2e5 100644 --- a/highprec.py +++ b/highprec.py @@ -1,6 +1,7 @@ import os import png import json +import torch import random import numpy as np from io import BytesIO @@ -9,7 +10,7 @@ from PIL.PngImagePlugin import PngInfo import folder_paths from comfy.cli_args import args -from nodes import SaveImage, PreviewImage +from nodes import SaveImage, PreviewImage, LoadImage def get_PIL_tEXt(image, prompt, extra_pnginfo): @@ -102,3 +103,32 @@ class PreviewImageHighPrec(SaveImageHighPrec): {"images": ("IMAGE", ), }, "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, } + +class LoadImageHighPrec(LoadImage): + TITLE = "Load Image (16 bit)" + FUNCTION = "load_image_high_precision" + + def load_image_high_precision(self, image): + if not image.endswith(".png"): + print("ColorMod: Only PNG files can be loaded in 16 bit color.") + return (self.load_image(image)) + + image_path = folder_paths.get_annotated_filepath(image) + reader = png.Reader(image_path) + + raw = reader.read() + image = np.vstack(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,] + + if image.shape[3] == 4: + mask = 1.0 - image[:,:,:,3] + image = torch.stack([image[:,:,:,x] for x in range(3)], dim = 3) + else: + mask = torch.zeros((64,64), dtype=torch.float32, device="cpu").unsqueeze(0) + return (image, mask)