Patch console, add WASSaveLUT
This commit is contained in:
+14
-30
@@ -10,16 +10,7 @@ try:
|
||||
from rich.panel import Panel
|
||||
from rich.progress import Progress, SpinnerColumn, TimeElapsedColumn, TextColumn
|
||||
from rich.traceback import install as rich_traceback_install
|
||||
from rich.theme import Theme
|
||||
THEME = Theme({
|
||||
"was.blue": "#7cc7ff",
|
||||
"was.orange": "orange1",
|
||||
"was.muted": "grey70",
|
||||
"was.good": "#7cc7ff",
|
||||
"was.bad": "orange1",
|
||||
"rule.line": "#7cc7ff",
|
||||
})
|
||||
console = Console(theme=THEME)
|
||||
console = Console(force_terminal=True, no_color=False)
|
||||
rich_traceback_install(show_locals=False, word_wrap=True)
|
||||
except Exception:
|
||||
console = None
|
||||
@@ -68,19 +59,18 @@ class NodeLoader:
|
||||
return mod, ok
|
||||
|
||||
def print_intro(self) -> None:
|
||||
msg = (
|
||||
"[was.muted]Nodes in this repo solve specific problems and may not fit every workflow.[/was.muted]"
|
||||
)
|
||||
msg = "Nodes in this repo solve specific problems and may not fit every workflow."
|
||||
if console:
|
||||
console.rule(f"[was.blue]{self.prefix.strip()}[/was.blue]", style="rule.line")
|
||||
console.print(Panel(msg, title=f"[was.orange]{self.prefix.strip()}[/was.orange]", border_style="was.blue"))
|
||||
# Single rule and a single-titled panel
|
||||
console.rule(style="cyan")
|
||||
console.print(Panel(msg, title="WAS Extras", border_style="cyan"))
|
||||
else:
|
||||
print(f"{self.prefix}Nodes in this repo solve specific problems and may not fit every workflow.")
|
||||
|
||||
def print_no_nodes_pkg(self) -> None:
|
||||
msg = "No ./nodes package found or import failed."
|
||||
if console:
|
||||
console.print(f"[was.orange]{self.prefix}{msg}[/was.orange]")
|
||||
console.print(f"{self.prefix}[bold yellow]{msg}[/bold yellow]")
|
||||
else:
|
||||
print(f"{self.prefix}{msg}")
|
||||
|
||||
@@ -90,26 +80,20 @@ class NodeLoader:
|
||||
fail_count = total - ok_count
|
||||
if console:
|
||||
table = Table(
|
||||
title=f"[was.blue]{self.prefix}Import Summary[/was.blue]",
|
||||
expand=False,
|
||||
header_style="was.blue",
|
||||
border_style="was.blue",
|
||||
header_style="cyan",
|
||||
border_style="cyan",
|
||||
)
|
||||
table.add_column("Module/File", overflow="fold")
|
||||
table.add_column("Time (s)", justify="right")
|
||||
table.add_column("Status", justify="center")
|
||||
table.add_column("Error", overflow="fold")
|
||||
for path, (timing, success, err) in self.timings.items():
|
||||
status = "[was.good]OK[/was.good]" if success else "[was.bad]FAILED[/was.bad]"
|
||||
status = "[green]OK[/green]" if success else "[red]FAILED[/red]"
|
||||
err_text = "" if err is None else f"{type(err).__name__}: {err}"
|
||||
table.add_row(str(path), f"{timing:.2f}", status, err_text)
|
||||
console.print(table)
|
||||
console.rule(style="rule.line")
|
||||
console.print(
|
||||
f"[was.blue]{self.prefix}Totals:[/was.blue] "
|
||||
f"[was.good]{ok_count} ok[/was.good], [was.bad]{fail_count} failed[/was.bad], {total} modules."
|
||||
)
|
||||
console.rule(style="rule.line")
|
||||
console.print(f"Totals: [green]{ok_count} ok[/green], [red]{fail_count} failed[/red], {total} modules.")
|
||||
else:
|
||||
print(f"{self.prefix} Import times:")
|
||||
for path, (timing, success, err) in self.timings.items():
|
||||
@@ -127,18 +111,18 @@ class NodeLoader:
|
||||
return
|
||||
if console:
|
||||
with Progress(
|
||||
SpinnerColumn(style="was.blue"),
|
||||
TextColumn("[progress.description]{task.description}", style="was.muted"),
|
||||
SpinnerColumn(style="cyan"),
|
||||
TextColumn("[progress.description]{task.description}", style="bright_black"),
|
||||
TimeElapsedColumn(),
|
||||
console=console,
|
||||
transient=True,
|
||||
) as progress:
|
||||
task = progress.add_task(f"[was.orange]{self.prefix}Loading nodes from ./nodes ...[/was.orange]", total=None)
|
||||
task = progress.add_task("Loading nodes from ./nodes ...", total=None)
|
||||
for _, name, _ in pkgutil.walk_packages(nodes_pkg.__path__, prefix=nodes_pkg.__name__ + "."):
|
||||
self.import_module(name)
|
||||
progress.remove_task(task)
|
||||
else:
|
||||
print(f"{self.prefix}Loading nodes from ./nodes ...")
|
||||
print("Loading nodes from ./nodes ...")
|
||||
for _, name, _ in pkgutil.walk_packages(nodes_pkg.__path__, prefix=nodes_pkg.__name__ + "."):
|
||||
self.import_module(name)
|
||||
self.print_summary()
|
||||
|
||||
+109
-20
@@ -74,8 +74,13 @@ class LUTLoader:
|
||||
def discover_cube_files() -> list[Path]:
|
||||
out: list[Path] = []
|
||||
for d in LUTLoader.get_lut_dirs():
|
||||
out.extend(sorted(d.glob("*.cube")))
|
||||
return out
|
||||
try:
|
||||
for p in d.iterdir():
|
||||
if p.is_file() and p.suffix.lower() == ".cube":
|
||||
out.append(p)
|
||||
except Exception:
|
||||
continue
|
||||
return sorted(out, key=lambda p: p.name.lower())
|
||||
|
||||
@staticmethod
|
||||
def luts_signature() -> str:
|
||||
@@ -158,6 +163,27 @@ class LUTLoader:
|
||||
|
||||
raise ValueError(f"{path.name}: invalid .cube")
|
||||
|
||||
@staticmethod
|
||||
def save_cube(path: Path, lut: 'LUT') -> None:
|
||||
# Ensure 3D table
|
||||
if lut.table_3d is None:
|
||||
raise ValueError("save_cube expects a 3D LUT table")
|
||||
table = lut.table_3d
|
||||
N = int(table.shape[0])
|
||||
dom_min = np.asarray(lut.domain_min, dtype=np.float32).tolist()
|
||||
dom_max = np.asarray(lut.domain_max, dtype=np.float32).tolist()
|
||||
with path.open("w", encoding="utf-8") as f:
|
||||
f.write(f"TITLE \"{lut.title or path.stem}\"\n")
|
||||
f.write(f"LUT_3D_SIZE {N}\n")
|
||||
f.write(f"DOMAIN_MIN {dom_min[0]:.6f} {dom_min[1]:.6f} {dom_min[2]:.6f}\n")
|
||||
f.write(f"DOMAIN_MAX {dom_max[0]:.6f} {dom_max[1]:.6f} {dom_max[2]:.6f}\n")
|
||||
# Write values in r-fastest order matching our reader's reshape
|
||||
for r in range(N):
|
||||
for g in range(N):
|
||||
for b in range(N):
|
||||
R, G, B = table[r, g, b]
|
||||
f.write(f"{float(R):.6f} {float(G):.6f} {float(B):.6f}\n")
|
||||
|
||||
@staticmethod
|
||||
def synthesize_builtin_lut(name: str, size: int = 33) -> LUT:
|
||||
params = None
|
||||
@@ -518,14 +544,6 @@ class WASLoadLUT:
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs):
|
||||
sig = LUTLoader.luts_signature()
|
||||
if sig != cls._last_sig:
|
||||
cls._last_sig = sig
|
||||
return time.time()
|
||||
return None
|
||||
|
||||
RETURN_TYPES = ("LUT",)
|
||||
RETURN_NAMES = ("lut",)
|
||||
|
||||
@@ -709,7 +727,9 @@ class LUTBlender:
|
||||
@staticmethod
|
||||
def _linear_to_srgb(x: np.ndarray) -> np.ndarray:
|
||||
x = x.astype(np.float32)
|
||||
return np.where(x <= 0.0031308, x * 12.92, 1.055 * (np.clip(x, 0.0, None) ** (1/2.4)) - 0.055).astype(np.float32)
|
||||
xc = np.clip(x, 0.0, None)
|
||||
out = np.where(xc <= 0.0031308, xc * 12.92, 1.055 * (xc ** (1/2.4)) - 0.055)
|
||||
return out.astype(np.float32)
|
||||
|
||||
@staticmethod
|
||||
def _rgb_linear_to_xyz(rgb: np.ndarray) -> np.ndarray:
|
||||
@@ -718,7 +738,8 @@ class LUTBlender:
|
||||
[0.2126729, 0.7151522, 0.0721750],
|
||||
[0.0193339, 0.1191920, 0.9503041],
|
||||
], dtype=np.float32)
|
||||
return np.tensordot(rgb, M.T, axes=1).astype(np.float32)
|
||||
xyz = np.tensordot(rgb, M.T, axes=1).astype(np.float32)
|
||||
return np.nan_to_num(xyz, nan=0.0, posinf=1e6, neginf=-1e6)
|
||||
|
||||
@staticmethod
|
||||
def _xyz_to_rgb_linear(xyz: np.ndarray) -> np.ndarray:
|
||||
@@ -727,16 +748,38 @@ class LUTBlender:
|
||||
[-0.9692660, 1.8760108, 0.0415560],
|
||||
[ 0.0556434, -0.2040259, 1.0572252],
|
||||
], dtype=np.float32)
|
||||
rgb = np.tensordot(xyz, M.T, axes=1).astype(np.float32)
|
||||
return np.nan_to_num(rgb, nan=0.0, posinf=1e6, neginf=-1e6)
|
||||
|
||||
@staticmethod
|
||||
def _xyz_d65_to_d50(xyz: np.ndarray) -> np.ndarray:
|
||||
M = np.array([
|
||||
[ 0.9555766, -0.0230393, 0.0631636],
|
||||
[-0.0282895, 1.0099416, 0.0210077],
|
||||
[ 0.0122982, -0.0204830, 1.3299098],
|
||||
], dtype=np.float32)
|
||||
return np.tensordot(xyz, M.T, axes=1).astype(np.float32)
|
||||
|
||||
@staticmethod
|
||||
def _xyz_d50_to_d65(xyz: np.ndarray) -> np.ndarray:
|
||||
M = np.array([
|
||||
[ 1.0478112, 0.0228866, -0.0501270],
|
||||
[ 0.0295424, 0.9904844, -0.0170491],
|
||||
[-0.0092345, 0.0150436, 0.7521316],
|
||||
], dtype=np.float32)
|
||||
return np.tensordot(xyz, M.T, axes=1).astype(np.float32)
|
||||
|
||||
@staticmethod
|
||||
def _rgb_to_lab(rgb: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
lin = LUTBlender._srgb_to_linear(rgb)
|
||||
xyz = LUTBlender._rgb_linear_to_xyz(lin)
|
||||
Xn, Yn, Zn = 0.95047, 1.0, 1.08883
|
||||
x = xyz[..., 0] / Xn
|
||||
y = xyz[..., 1] / Yn
|
||||
z = xyz[..., 2] / Zn
|
||||
xyz_d65 = LUTBlender._rgb_linear_to_xyz(lin)
|
||||
xyz = LUTBlender._xyz_d65_to_d50(xyz_d65)
|
||||
# Avoid tiny negative values from numeric error before cbrt
|
||||
xyz = np.clip(xyz, 0.0, None)
|
||||
Xn, Yn, Zn = 0.96422, 1.0, 0.82521
|
||||
x = xyz[..., 0] / np.clip(Xn, 1e-8, None)
|
||||
y = xyz[..., 1] / np.clip(Yn, 1e-8, None)
|
||||
z = xyz[..., 2] / np.clip(Zn, 1e-8, None)
|
||||
e = (6/29) ** 3
|
||||
k = (29/6) ** 2 / 3
|
||||
f = lambda t: np.where(t > e, np.cbrt(t), k * t + 4/29)
|
||||
@@ -755,12 +798,15 @@ class LUTBlender:
|
||||
e3 = e ** 3
|
||||
k = 3 * (e ** 2)
|
||||
invf = lambda t: np.where(t > e, t ** 3, (t - 4/29) / k)
|
||||
Xn, Yn, Zn = 0.95047, 1.0, 1.08883
|
||||
|
||||
Xn, Yn, Zn = 0.96422, 1.0, 0.82521
|
||||
x = invf(fx) * Xn
|
||||
y = invf(fy) * Yn
|
||||
z = invf(fz) * Zn
|
||||
xyz = np.stack([x, y, z], axis=-1).astype(np.float32)
|
||||
lin = LUTBlender._xyz_to_rgb_linear(xyz)
|
||||
xyz_d50 = np.stack([x, y, z], axis=-1).astype(np.float32)
|
||||
|
||||
xyz_d65 = LUTBlender._xyz_d50_to_d65(xyz_d50)
|
||||
lin = LUTBlender._xyz_to_rgb_linear(xyz_d65)
|
||||
rgb = LUTBlender._linear_to_srgb(lin)
|
||||
return np.clip(rgb, 0.0, 1.0).astype(np.float32)
|
||||
|
||||
@@ -804,6 +850,10 @@ class LUTBlender:
|
||||
L = La * (1.0 - tt) + Lb * tt
|
||||
A = aa * (1.0 - tt) + ab * tt
|
||||
B = ba * (1.0 - tt) + bb * tt
|
||||
# Clamp to valid Lab ranges to reduce out-of-gamut artifacts
|
||||
L = np.clip(L, 0.0, 100.0)
|
||||
A = np.clip(A, -128.0, 128.0)
|
||||
B = np.clip(B, -128.0, 128.0)
|
||||
out = LUTBlender._lab_to_rgb(L, A, B)
|
||||
return np.clip(out, 0.0, 1.0).astype(np.float32)
|
||||
|
||||
@@ -929,6 +979,43 @@ class WASApplyLUT:
|
||||
y = image * (1.0 - strength) + y * strength
|
||||
return (y.clamp(0, 1),)
|
||||
|
||||
# Save LUT
|
||||
|
||||
class WASSaveLUT:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"lut": ("LUT",),
|
||||
"filename": ("STRING", {"default": "CustomLUT"}),
|
||||
"output_size": ("INT", {"default": 33, "min": 17, "max": 65, "step": 2}),
|
||||
"overwrite": ("BOOL", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LUT",)
|
||||
RETURN_NAMES = ("lut",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "WAS/Color/LUT"
|
||||
|
||||
def run(self, lut, filename, output_size, overwrite):
|
||||
lut_dirs = LUTLoader.get_lut_dirs()
|
||||
if not lut_dirs:
|
||||
raise RuntimeError("No models/LUT directory found. Please create one under your ComfyUI models folder.")
|
||||
dst_dir = lut_dirs[0]
|
||||
dst_dir.mkdir(parents=True, exist_ok=True)
|
||||
name = filename.strip()
|
||||
if not name.lower().endswith(".cube"):
|
||||
name += ".cube"
|
||||
path = dst_dir / name
|
||||
if path.exists() and not overwrite:
|
||||
raise FileExistsError(f"{path} exists. Enable overwrite to replace.")
|
||||
|
||||
lut3 = WASLUT.convert_to_3d(lut, output_size)
|
||||
LUTLoader.save_cube(path, lut3)
|
||||
|
||||
return (lut3,)
|
||||
|
||||
# RGB PARADE
|
||||
|
||||
def get_temp_dir() -> str:
|
||||
@@ -1022,6 +1109,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WASLoadLUT": WASLoadLUT,
|
||||
"WASCombineLUT": WASCombineLUT,
|
||||
"WASApplyLUT": WASApplyLUT,
|
||||
"WASSaveLUT": WASSaveLUT,
|
||||
"WASChannelWaveform": WASChannelWaveform,
|
||||
}
|
||||
|
||||
@@ -1029,5 +1117,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WASLoadLUT": "WAS Load LUT",
|
||||
"WASCombineLUT": "WAS LUT Blender",
|
||||
"WASApplyLUT": "WAS Apply LUT",
|
||||
"WASSaveLUT": "WAS Save LUT (.cube)",
|
||||
"WASChannelWaveform": "WAS Channel Waveform (Parade)",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user