From c4f86c073a9ffad3e70e0b52d7b15855215e1cb0 Mon Sep 17 00:00:00 2001 From: Jordan Thompson Date: Sun, 28 Sep 2025 00:57:39 -0700 Subject: [PATCH] Patch console, add WASSaveLUT --- __init__.py | 44 ++++++----------- nodes/WASLUT.py | 129 ++++++++++++++++++++++++++++++++++++++++-------- 2 files changed, 123 insertions(+), 50 deletions(-) diff --git a/__init__.py b/__init__.py index 9f67e9a..b0bef50 100644 --- a/__init__.py +++ b/__init__.py @@ -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() diff --git a/nodes/WASLUT.py b/nodes/WASLUT.py index defad1f..eaa2ec6 100644 --- a/nodes/WASLUT.py +++ b/nodes/WASLUT.py @@ -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)", }