move node, refractor

This commit is contained in:
City
2024-04-03 15:01:51 +02:00
parent 079639608b
commit df5d1ca4f0
4 changed files with 119 additions and 50 deletions
+6 -4
View File
@@ -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"]
+61 -30
View File
@@ -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,
}
+6 -1
View File
@@ -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"):
+46 -15
View File
@@ -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)
}
}
})