Add docstrings to all public functions, classes, and methods
Coverage: 29% → 100% (41/41 public items documented). Keeps the existing concise style with inline shape annotations.
This commit is contained in:
@@ -11,7 +11,10 @@ from .nodes.observe import LCSPreviewColors, LCSStepObserver
|
||||
|
||||
|
||||
class LCSExtension(ComfyExtension):
|
||||
"""V3 ComfyExtension providing all LCS nodes to ComfyUI."""
|
||||
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
"""Return all 6 LCS node classes."""
|
||||
return [
|
||||
LCSCalibrate,
|
||||
LCSLoadData,
|
||||
@@ -23,6 +26,7 @@ class LCSExtension(ComfyExtension):
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> LCSExtension:
|
||||
"""V3 async entry point called by ComfyUI on startup."""
|
||||
return LCSExtension()
|
||||
|
||||
|
||||
|
||||
+2
-1
@@ -3,7 +3,8 @@ from .patchify import patchify, unpatchify
|
||||
from .timestep import sigma_to_paper_t, get_alpha_beta, normalize_to_t50, denormalize_from_t50
|
||||
from .color_space import decode_lcs_to_hsl, encode_hsl_to_lcs, hex_to_hsl, hsl_to_rgb
|
||||
|
||||
# Lazy import for calibration (depends on comfy.utils)
|
||||
|
||||
def calibrate(*args, **kwargs):
|
||||
"""Lazy wrapper for core.calibration.calibrate (avoids importing comfy.utils at module level)."""
|
||||
from .calibration import calibrate as _calibrate
|
||||
return _calibrate(*args, **kwargs)
|
||||
|
||||
@@ -50,6 +50,7 @@ _beta_tensor = None
|
||||
|
||||
|
||||
def get_alpha_table():
|
||||
"""Return α_t table as tensor [51, 3], cached after first call."""
|
||||
global _alpha_tensor
|
||||
if _alpha_tensor is None:
|
||||
_alpha_tensor = torch.tensor(ALPHA_T, dtype=torch.float32) # [51, 3]
|
||||
@@ -57,6 +58,7 @@ def get_alpha_table():
|
||||
|
||||
|
||||
def get_beta_table():
|
||||
"""Return β_t table as tensor [51, 3], cached after first call."""
|
||||
global _beta_tensor
|
||||
if _beta_tensor is None:
|
||||
_beta_tensor = torch.tensor(BETA_T, dtype=torch.float32) # [51, 3]
|
||||
|
||||
@@ -4,6 +4,12 @@ import torch
|
||||
|
||||
@dataclass
|
||||
class LCSData:
|
||||
"""Calibration data for the Latent Color Subspace.
|
||||
|
||||
Produced by PCA on FLUX VAE-encoded solid-color images. Flows between
|
||||
all LCS nodes as the shared LCS_DATA custom type.
|
||||
"""
|
||||
|
||||
basis: torch.Tensor # [64, 3] PCA basis B (orthonormal columns)
|
||||
mean: torch.Tensor # [64] PCA mean mu
|
||||
anchor_lcs: torch.Tensor # [8, 3] LCS coords of 8 anchor colors [R,B,G,M,C,Y,Black,White]
|
||||
|
||||
@@ -15,8 +15,16 @@ DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "data")
|
||||
|
||||
|
||||
class LCSCalibrate(io.ComfyNode):
|
||||
"""Compute LCS basis and anchors from a FLUX VAE via PCA on solid-color images.
|
||||
|
||||
Samples num_colors uniform HSV colors, encodes them through the VAE,
|
||||
runs PCA on the averaged 64-dim patch vectors, and extracts the 3D basis B,
|
||||
mean μ, and 8 anchor color positions. Auto-saves result to data/ directory.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (VAE, num_colors, image_size) and LCS_DATA output."""
|
||||
return io.Schema(
|
||||
node_id="LCSCalibrate",
|
||||
display_name="LCS Calibrate",
|
||||
@@ -36,6 +44,7 @@ class LCSCalibrate(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, vae, num_colors, image_size) -> io.NodeOutput:
|
||||
"""Run PCA calibration and save result as .safetensors. Returns LCS_DATA."""
|
||||
lcs_data = calibrate(vae, num_colors=num_colors, image_size=image_size)
|
||||
|
||||
# Auto-save to data/ directory
|
||||
@@ -52,8 +61,15 @@ class LCSCalibrate(io.ComfyNode):
|
||||
|
||||
|
||||
class LCSLoadData(io.ComfyNode):
|
||||
"""Load cached LCS calibration data from .safetensors, or auto-calibrate if VAE is provided.
|
||||
|
||||
Scans the data/ directory for available calibration files. When set to 'auto',
|
||||
loads the default file or runs calibration if no cache exists and a VAE is connected.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (calibration_file combo, optional VAE) and LCS_DATA output."""
|
||||
# Scan data/ dir for available calibration files
|
||||
files = ["auto"]
|
||||
if os.path.isdir(DATA_DIR):
|
||||
@@ -79,6 +95,7 @@ class LCSLoadData(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, calibration_file, vae=None) -> io.NodeOutput:
|
||||
"""Load or auto-generate calibration data. Returns LCS_DATA."""
|
||||
loaded = False
|
||||
|
||||
if calibration_file != "auto":
|
||||
|
||||
@@ -22,6 +22,7 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step,
|
||||
target_colors_hsl: list of (h, s, l) tuples, one per batch item (or one for all).
|
||||
"""
|
||||
def post_cfg_fn(args):
|
||||
"""Post-CFG hook: project to LCS, apply color intervention, reconstruct."""
|
||||
denoised = args["denoised"] # [B, 16, H, W] in process_in space
|
||||
sigma = args["sigma"]
|
||||
|
||||
@@ -186,8 +187,17 @@ def _hue_lerp(h1, h2, t):
|
||||
|
||||
|
||||
class LCSColorIntervene(io.ComfyNode):
|
||||
"""Steer colors during FLUX generation via the Latent Color Subspace.
|
||||
|
||||
Installs a post-CFG hook that projects the denoised prediction into the
|
||||
3D LCS, shifts it toward the target color (Type I, Type II, or interpolated),
|
||||
preserves the 61D residual, and writes the modified prediction back.
|
||||
Active only during [start_step, end_step].
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (MODEL, LCS_DATA, color, strength, mode, steps, mask) and MODEL output."""
|
||||
return io.Schema(
|
||||
node_id="LCSColorIntervene",
|
||||
display_name="LCS Color Intervene",
|
||||
@@ -217,6 +227,7 @@ class LCSColorIntervene(io.ComfyNode):
|
||||
@classmethod
|
||||
def execute(cls, model, lcs_data, color, strength, mode, start_step, end_step,
|
||||
mask=None) -> io.NodeOutput:
|
||||
"""Clone model, attach LCS color intervention hook. Returns patched MODEL."""
|
||||
m = model.clone()
|
||||
h, s, l = hex_to_hsl(color)
|
||||
hook = _build_post_cfg_fn(lcs_data, [(h, s, l)], strength, mode, start_step, end_step, mask)
|
||||
@@ -225,8 +236,16 @@ class LCSColorIntervene(io.ComfyNode):
|
||||
|
||||
|
||||
class LCSColorBatch(io.ComfyNode):
|
||||
"""Apply different target colors to each batch item for multi-color generation.
|
||||
|
||||
Parses comma-separated hex colors and installs a post-CFG hook that applies
|
||||
a distinct color target per batch index. Also outputs batch_size (INT) for
|
||||
connecting to EmptyLatentImage.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (MODEL, LCS_DATA, colors string, strength, mode, steps, mask) and (MODEL, INT) outputs."""
|
||||
return io.Schema(
|
||||
node_id="LCSColorBatch",
|
||||
display_name="LCS Color Batch",
|
||||
@@ -253,6 +272,7 @@ class LCSColorBatch(io.ComfyNode):
|
||||
@classmethod
|
||||
def execute(cls, model, lcs_data, colors, strength, mode, start_step, end_step,
|
||||
mask=None) -> io.NodeOutput:
|
||||
"""Clone model, attach per-batch color hooks. Returns (MODEL, batch_size INT)."""
|
||||
m = model.clone()
|
||||
|
||||
# Parse comma-separated hex colors
|
||||
|
||||
@@ -58,8 +58,15 @@ def _latent_to_color_preview(samples, lcs_data, sigma, upscale=8):
|
||||
|
||||
|
||||
class LCSPreviewColors(io.ComfyNode):
|
||||
"""Visualize latent colors without VAE decoding — pure math color preview from LCS.
|
||||
|
||||
Projects latent patches into the 3D LCS, normalizes to t=50, decodes to HSL,
|
||||
converts to RGB, and upscales 8x to pixel resolution. Produces a [B, H, W, 3] IMAGE.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (LATENT, LCS_DATA, sigma) and IMAGE output."""
|
||||
return io.Schema(
|
||||
node_id="LCSPreviewColors",
|
||||
display_name="LCS Preview Colors",
|
||||
@@ -78,14 +85,23 @@ class LCSPreviewColors(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, latent, lcs_data, sigma) -> io.NodeOutput:
|
||||
"""Decode latent to LCS color preview. Returns IMAGE [B, H, W, 3]."""
|
||||
samples = latent["samples"]
|
||||
result = _latent_to_color_preview(samples, lcs_data, sigma, upscale=8)
|
||||
return io.NodeOutput(result)
|
||||
|
||||
|
||||
class LCSStepObserver(io.ComfyNode):
|
||||
"""Patches model to save per-step LCS color previews to ComfyUI's temp directory.
|
||||
|
||||
Installs a post-CFG hook that generates a color preview image for the first
|
||||
batch item at each sampling step. Images are saved as lcs_step_NNN_sX.XXX.png.
|
||||
Does not modify the denoised prediction.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
"""Define inputs (MODEL, LCS_DATA) and MODEL output."""
|
||||
return io.Schema(
|
||||
node_id="LCSStepObserver",
|
||||
display_name="LCS Step Observer",
|
||||
@@ -102,10 +118,12 @@ class LCSStepObserver(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, lcs_data) -> io.NodeOutput:
|
||||
"""Clone model, attach step observer hook. Returns patched MODEL."""
|
||||
m = model.clone()
|
||||
step_counter = [0]
|
||||
|
||||
def observer_fn(args):
|
||||
"""Post-CFG hook: generate color preview and save to temp directory."""
|
||||
denoised = args["denoised"]
|
||||
sigma = args["sigma"]
|
||||
sigma_val = float(sigma.flatten()[0])
|
||||
|
||||
Reference in New Issue
Block a user