web widget, refractor

This commit is contained in:
City
2024-03-31 23:36:24 +02:00
parent 97f367bdfc
commit 079639608b
6 changed files with 266 additions and 97 deletions
+14 -23
View File
@@ -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"]
-65
View File
@@ -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,)
+112
View File
@@ -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,
}
+20 -9
View File
@@ -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,
}
+2
View File
@@ -0,0 +1,2 @@
pypng
opencv-contrib-python
+118
View File
@@ -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);
}
}
})