move node, refractor
This commit is contained in:
+6
-4
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user