Patch console, add WASSaveLUT

This commit is contained in:
Jordan Thompson
2025-09-28 00:57:39 -07:00
parent ba319e42f1
commit c4f86c073a
2 changed files with 123 additions and 50 deletions
+14 -30
View File
@@ -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
View File
@@ -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)",
}