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:
facok
2026-03-17 02:26:21 +08:00
parent 534bc19033
commit 0973fd0fd8
7 changed files with 69 additions and 1 deletions
+4
View File
@@ -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
View File
@@ -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)
+2
View File
@@ -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]
+6
View File
@@ -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]
+17
View File
@@ -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":
+20
View File
@@ -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
+18
View File
@@ -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])