web widget, refractor
This commit is contained in:
+14
-23
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
pypng
|
||||
opencv-contrib-python
|
||||
+118
@@ -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);
|
||||
}
|
||||
}
|
||||
})
|
||||
Reference in New Issue
Block a user