diff --git a/simple_syrup/nodes/__init__.py b/simple_syrup/nodes/__init__.py index e9fe747..c06593d 100644 --- a/simple_syrup/nodes/__init__.py +++ b/simple_syrup/nodes/__init__.py @@ -32,6 +32,7 @@ from .seed import Seed from .simple_load_anima import SimpleLoadAnima from .simple_load_checkpoint import SimpleLoadCheckpoint from .tile_and_tag_segs import TileAndTagSEGS +from .vae_options import VAEDecodeOptions, VAEEncodeOptions from .vitmatte_model_loader import ViTMatteModelLoader from .wd14_tagger_loader import WD14TaggerLoader @@ -49,6 +50,8 @@ NODE_CLASS_MAPPINGS = { "SimpleSyrup.PromptSEGSWithSAM": PromptSEGSWithSAM, "SimpleSyrup.SimpleVAEEncode": SimpleVAEEncode, "SimpleSyrup.UpscaleLatentFromImage": UpscaleLatentFromImage, + "SimpleSyrup.VAEDecodeOptions": VAEDecodeOptions, + "SimpleSyrup.VAEEncodeOptions": VAEEncodeOptions, "SimpleSyrup.ResizeImageToTarget": ResizeImageToTarget, "SimpleSyrup.DetailSEGSAsRegions": DetailSEGSAsRegions, "SimpleSyrup.DetailSEGSByScaleFactor": DetailSEGSByScaleFactor, @@ -84,6 +87,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SimpleSyrup.PromptSEGSWithSAM": "Prompt SEGS w/ SAM", "SimpleSyrup.SimpleVAEEncode": "Simple VAE Encode", "SimpleSyrup.UpscaleLatentFromImage": "Upscale Latent From Image", + "SimpleSyrup.VAEDecodeOptions": "VAE Decode (Options)", + "SimpleSyrup.VAEEncodeOptions": "VAE Encode (Options)", "SimpleSyrup.ResizeImageToTarget": "Resize Image to Target", "SimpleSyrup.DetailSEGSAsRegions": "Detail SEGS as Regions", "SimpleSyrup.DetailSEGSByScaleFactor": "Detail SEGS by Scale Factor", @@ -132,6 +137,8 @@ __all__ = [ "SimpleVAEEncode", "TileAndTagSEGS", "UpscaleLatentFromImage", + "VAEDecodeOptions", + "VAEEncodeOptions", "ViTMatteModelLoader", "WD14TaggerLoader", ] diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index a7b6454..da145b4 100644 --- a/simple_syrup/nodes/tooltips.py +++ b/simple_syrup/nodes/tooltips.py @@ -23,6 +23,27 @@ MODEL_OUTPUT = "Loaded diffusion model for downstream MODEL inputs." CLIP_OUTPUT = "Loaded text encoder for downstream CLIP inputs." VAE_OUTPUT = "Loaded VAE used to encode images to latents and decode latents to images." +VAE_OPTIONS_USE_TILING = ( + "Use ComfyUI's tiled VAE node. Disabled uses normal ComfyUI VAE behavior, " + "including its automatic tiled retry after out-of-memory." +) +VAE_OPTIONS_ENCODE_PIXELS = "Image to encode into latent space." +VAE_OPTIONS_DECODE_SAMPLES = "Latent samples to decode into an image." +VAE_OPTIONS_VAE = "VAE used for the selected encode or decode operation." +VAE_OPTIONS_TILE_SIZE = ( + "Tile size in pixels. Larger tiles are faster but use more memory." +) +VAE_OPTIONS_OVERLAP = ( + "Overlap between tiles in pixels. Larger overlaps reduce seams but do more work." +) +VAE_OPTIONS_ENCODE_TEMPORAL_SIZE = "For video VAEs, number of frames to encode at once." +VAE_OPTIONS_DECODE_TEMPORAL_SIZE = "For video VAEs, number of frames to decode at once." +VAE_OPTIONS_TEMPORAL_OVERLAP = ( + "For video VAEs, number of overlapping frames between temporal tiles." +) +VAE_OPTIONS_LATENT_OUTPUT = "Latent produced by ComfyUI's selected VAE encode node." +VAE_OPTIONS_IMAGE_OUTPUT = "Image produced by ComfyUI's selected VAE decode node." + SAM_MODEL_OUTPUT = "Loaded SAM model for prompt-based mask and SEGS creation." GROUNDING_DINO_MODEL_OUTPUT = ( "Loaded GroundingDINO model for finding prompt-matched boxes in images." diff --git a/simple_syrup/nodes/vae_options.py b/simple_syrup/nodes/vae_options.py new file mode 100644 index 0000000..d37c20c --- /dev/null +++ b/simple_syrup/nodes/vae_options.py @@ -0,0 +1,248 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI nodes that switch between native normal and tiled VAE execution.""" + +from __future__ import annotations + +from typing import Any + +from ..runtime.vae_options_graph import ExpansionResult, VAEOptionsGraphBuilder +from . import tooltips + +VAE_TILE_SIZE_DEFAULT = 512 +VAE_TILE_SIZE_MIN = 64 +VAE_TILE_SIZE_MAX = 4096 +VAE_ENCODE_TILE_SIZE_STEP = 64 +VAE_DECODE_TILE_SIZE_STEP = 32 +VAE_OVERLAP_DEFAULT = 64 +VAE_OVERLAP_MIN = 0 +VAE_OVERLAP_MAX = 4096 +VAE_OVERLAP_STEP = 32 +VAE_TEMPORAL_SIZE_DEFAULT = 64 +VAE_TEMPORAL_SIZE_MIN = 8 +VAE_TEMPORAL_SIZE_MAX = 4096 +VAE_TEMPORAL_SIZE_STEP = 4 +VAE_TEMPORAL_OVERLAP_DEFAULT = 8 +VAE_TEMPORAL_OVERLAP_MIN = 4 +VAE_TEMPORAL_OVERLAP_MAX = 4096 +VAE_TEMPORAL_OVERLAP_STEP = 4 + + +class VAEEncodeOptions: + """Encode images through ComfyUI's normal or tiled VAE encode nodes.""" + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("latent",) + OUTPUT_TOOLTIPS = (tooltips.VAE_OPTIONS_LATENT_OUTPUT,) + FUNCTION = "encode" + CATEGORY = "SimpleSyrup/Latent" + DESCRIPTION = ( + "Encodes images to latent space with selectable normal or tiled VAE execution." + ) + SEARCH_ALIASES = [ + "vae encode", + "encode image", + "tiled vae encode", + "image to latent", + ] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare VAE encode inputs and tiled execution controls.""" + + return { + "required": { + "use_tiling": _use_tiling_input(), + "pixels": ( + "IMAGE", + { + "rawLink": True, + "tooltip": tooltips.VAE_OPTIONS_ENCODE_PIXELS, + }, + ), + "vae": _vae_input(), + "tile_size": _tile_size_input(VAE_ENCODE_TILE_SIZE_STEP), + "overlap": _overlap_input(), + "temporal_size": _temporal_size_input( + tooltips.VAE_OPTIONS_ENCODE_TEMPORAL_SIZE + ), + "temporal_overlap": _temporal_overlap_input(), + }, + } + + def encode( + self, + use_tiling: bool, + pixels: object, + vae: object, + tile_size: int, + overlap: int, + temporal_size: int = VAE_TEMPORAL_SIZE_DEFAULT, + temporal_overlap: int = VAE_TEMPORAL_OVERLAP_DEFAULT, + ) -> ExpansionResult: + """Expand to ComfyUI's selected native VAE encode node.""" + + return VAEOptionsGraphBuilder().build_encode( + pixels=pixels, + vae=vae, + use_tiling=use_tiling, + tile_size=tile_size, + overlap=overlap, + temporal_size=temporal_size, + temporal_overlap=temporal_overlap, + ) + + +class VAEDecodeOptions: + """Decode latents through ComfyUI's normal or tiled VAE decode nodes.""" + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + OUTPUT_TOOLTIPS = (tooltips.VAE_OPTIONS_IMAGE_OUTPUT,) + FUNCTION = "decode" + CATEGORY = "SimpleSyrup/Latent" + DESCRIPTION = ( + "Decodes latents to images with selectable normal or tiled VAE execution." + ) + SEARCH_ALIASES = [ + "vae decode", + "decode latent", + "tiled vae decode", + "latent to image", + ] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare VAE decode inputs and tiled execution controls.""" + + return { + "required": { + "use_tiling": _use_tiling_input(), + "samples": ( + "LATENT", + { + "rawLink": True, + "tooltip": tooltips.VAE_OPTIONS_DECODE_SAMPLES, + }, + ), + "vae": _vae_input(), + "tile_size": _tile_size_input(VAE_DECODE_TILE_SIZE_STEP), + "overlap": _overlap_input(), + "temporal_size": _temporal_size_input( + tooltips.VAE_OPTIONS_DECODE_TEMPORAL_SIZE + ), + "temporal_overlap": _temporal_overlap_input(), + }, + } + + def decode( + self, + use_tiling: bool, + samples: object, + vae: object, + tile_size: int, + overlap: int, + temporal_size: int = VAE_TEMPORAL_SIZE_DEFAULT, + temporal_overlap: int = VAE_TEMPORAL_OVERLAP_DEFAULT, + ) -> ExpansionResult: + """Expand to ComfyUI's selected native VAE decode node.""" + + return VAEOptionsGraphBuilder().build_decode( + samples=samples, + vae=vae, + use_tiling=use_tiling, + tile_size=tile_size, + overlap=overlap, + temporal_size=temporal_size, + temporal_overlap=temporal_overlap, + ) + + +def _use_tiling_input() -> tuple[str, dict[str, object]]: + """Return the shared tiling toggle declaration.""" + + return ( + "BOOLEAN", + { + "default": False, + "tooltip": tooltips.VAE_OPTIONS_USE_TILING, + }, + ) + + +def _vae_input() -> tuple[str, dict[str, object]]: + """Return the shared raw-link VAE input declaration.""" + + return ( + "VAE", + { + "rawLink": True, + "tooltip": tooltips.VAE_OPTIONS_VAE, + }, + ) + + +def _tile_size_input(step: int) -> tuple[str, dict[str, object]]: + """Return the tile-size input declaration for encode or decode.""" + + return ( + "INT", + { + "default": VAE_TILE_SIZE_DEFAULT, + "min": VAE_TILE_SIZE_MIN, + "max": VAE_TILE_SIZE_MAX, + "step": step, + "advanced": True, + "tooltip": tooltips.VAE_OPTIONS_TILE_SIZE, + }, + ) + + +def _overlap_input() -> tuple[str, dict[str, object]]: + """Return the shared tile-overlap input declaration.""" + + return ( + "INT", + { + "default": VAE_OVERLAP_DEFAULT, + "min": VAE_OVERLAP_MIN, + "max": VAE_OVERLAP_MAX, + "step": VAE_OVERLAP_STEP, + "advanced": True, + "tooltip": tooltips.VAE_OPTIONS_OVERLAP, + }, + ) + + +def _temporal_size_input(tooltip: str) -> tuple[str, dict[str, object]]: + """Return the shared temporal tile-size input declaration.""" + + return ( + "INT", + { + "default": VAE_TEMPORAL_SIZE_DEFAULT, + "min": VAE_TEMPORAL_SIZE_MIN, + "max": VAE_TEMPORAL_SIZE_MAX, + "step": VAE_TEMPORAL_SIZE_STEP, + "advanced": True, + "tooltip": tooltip, + }, + ) + + +def _temporal_overlap_input() -> tuple[str, dict[str, object]]: + """Return the shared temporal overlap input declaration.""" + + return ( + "INT", + { + "default": VAE_TEMPORAL_OVERLAP_DEFAULT, + "min": VAE_TEMPORAL_OVERLAP_MIN, + "max": VAE_TEMPORAL_OVERLAP_MAX, + "step": VAE_TEMPORAL_OVERLAP_STEP, + "advanced": True, + "tooltip": tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP, + }, + ) diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index 2733f56..fd6446e 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -15,6 +15,8 @@ def get_nodes() -> list[type[object]]: from .scale_factor import ScaleFactorV3 from .simple_load_checkpoint import SimpleLoadCheckpointV3 from .tile_and_tag_segs import TileAndTagSEGSV3 + from .vae_decode_options import VAEDecodeOptionsV3 + from .vae_encode_options import VAEEncodeOptionsV3 from .wd14_tagger_loader import WD14TaggerLoaderV3 if not prompt_control_is_available(): @@ -23,6 +25,8 @@ def get_nodes() -> list[type[object]]: TileAndTagSEGSV3, SimpleLoadCheckpointV3, ScaleFactorV3, + VAEDecodeOptionsV3, + VAEEncodeOptionsV3, ] from .encode_prompt_batch_with_prompt_control import ( @@ -34,6 +38,8 @@ def get_nodes() -> list[type[object]]: TileAndTagSEGSV3, SimpleLoadCheckpointV3, ScaleFactorV3, + VAEDecodeOptionsV3, + VAEEncodeOptionsV3, EncodePromptBatchWithPromptControl, ] diff --git a/simple_syrup/nodes_v3/vae_decode_options.py b/simple_syrup/nodes_v3/vae_decode_options.py new file mode 100644 index 0000000..71270e4 --- /dev/null +++ b/simple_syrup/nodes_v3/vae_decode_options.py @@ -0,0 +1,143 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node wrapper for VAE Decode (Options).""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..nodes import tooltips +from ..nodes.vae_options import ( + VAE_DECODE_TILE_SIZE_STEP, + VAE_OVERLAP_DEFAULT, + VAE_OVERLAP_MAX, + VAE_OVERLAP_MIN, + VAE_OVERLAP_STEP, + VAE_TEMPORAL_OVERLAP_DEFAULT, + VAE_TEMPORAL_OVERLAP_MAX, + VAE_TEMPORAL_OVERLAP_MIN, + VAE_TEMPORAL_OVERLAP_STEP, + VAE_TEMPORAL_SIZE_DEFAULT, + VAE_TEMPORAL_SIZE_MAX, + VAE_TEMPORAL_SIZE_MIN, + VAE_TEMPORAL_SIZE_STEP, + VAE_TILE_SIZE_DEFAULT, + VAE_TILE_SIZE_MAX, + VAE_TILE_SIZE_MIN, + VAEDecodeOptions, +) + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io + + +class VAEDecodeOptionsV3(_ComfyNodeBase): + """Expose VAE Decode (Options) through Comfy's v3 extension API.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the VAE Decode (Options) v3 schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.VAEDecodeOptions", + display_name="VAE Decode (Options)", + enable_expand=True, + category="SimpleSyrup/Latent", + description=VAEDecodeOptions.DESCRIPTION, + search_aliases=VAEDecodeOptions.SEARCH_ALIASES, + inputs=[ + _comfy_io.Boolean.Input( + "use_tiling", + default=False, + tooltip=tooltips.VAE_OPTIONS_USE_TILING, + ), + _comfy_io.Latent.Input( + "samples", + raw_link=True, + tooltip=tooltips.VAE_OPTIONS_DECODE_SAMPLES, + ), + _comfy_io.Vae.Input( + "vae", + raw_link=True, + tooltip=tooltips.VAE_OPTIONS_VAE, + ), + _comfy_io.Int.Input( + "tile_size", + default=VAE_TILE_SIZE_DEFAULT, + min=VAE_TILE_SIZE_MIN, + max=VAE_TILE_SIZE_MAX, + step=VAE_DECODE_TILE_SIZE_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_TILE_SIZE, + ), + _comfy_io.Int.Input( + "overlap", + default=VAE_OVERLAP_DEFAULT, + min=VAE_OVERLAP_MIN, + max=VAE_OVERLAP_MAX, + step=VAE_OVERLAP_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_OVERLAP, + ), + _comfy_io.Int.Input( + "temporal_size", + default=VAE_TEMPORAL_SIZE_DEFAULT, + min=VAE_TEMPORAL_SIZE_MIN, + max=VAE_TEMPORAL_SIZE_MAX, + step=VAE_TEMPORAL_SIZE_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_DECODE_TEMPORAL_SIZE, + ), + _comfy_io.Int.Input( + "temporal_overlap", + default=VAE_TEMPORAL_OVERLAP_DEFAULT, + min=VAE_TEMPORAL_OVERLAP_MIN, + max=VAE_TEMPORAL_OVERLAP_MAX, + step=VAE_TEMPORAL_OVERLAP_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP, + ), + ], + outputs=[ + _comfy_io.Image.Output( + "image", + tooltip=tooltips.VAE_OPTIONS_IMAGE_OUTPUT, + ), + ], + ) + + @classmethod + def execute( + cls, + use_tiling: bool, + samples: object, + vae: object, + tile_size: int, + overlap: int, + temporal_size: int, + temporal_overlap: int, + ) -> Any: + """Expand through the legacy VAE Decode (Options) implementation.""" + + return VAEDecodeOptions().decode( + use_tiling=use_tiling, + samples=samples, + vae=vae, + tile_size=tile_size, + overlap=overlap, + temporal_size=temporal_size, + temporal_overlap=temporal_overlap, + ) diff --git a/simple_syrup/nodes_v3/vae_encode_options.py b/simple_syrup/nodes_v3/vae_encode_options.py new file mode 100644 index 0000000..1c7714c --- /dev/null +++ b/simple_syrup/nodes_v3/vae_encode_options.py @@ -0,0 +1,143 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node wrapper for VAE Encode (Options).""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..nodes import tooltips +from ..nodes.vae_options import ( + VAE_ENCODE_TILE_SIZE_STEP, + VAE_OVERLAP_DEFAULT, + VAE_OVERLAP_MAX, + VAE_OVERLAP_MIN, + VAE_OVERLAP_STEP, + VAE_TEMPORAL_OVERLAP_DEFAULT, + VAE_TEMPORAL_OVERLAP_MAX, + VAE_TEMPORAL_OVERLAP_MIN, + VAE_TEMPORAL_OVERLAP_STEP, + VAE_TEMPORAL_SIZE_DEFAULT, + VAE_TEMPORAL_SIZE_MAX, + VAE_TEMPORAL_SIZE_MIN, + VAE_TEMPORAL_SIZE_STEP, + VAE_TILE_SIZE_DEFAULT, + VAE_TILE_SIZE_MAX, + VAE_TILE_SIZE_MIN, + VAEEncodeOptions, +) + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io + + +class VAEEncodeOptionsV3(_ComfyNodeBase): + """Expose VAE Encode (Options) through Comfy's v3 extension API.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the VAE Encode (Options) v3 schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.VAEEncodeOptions", + display_name="VAE Encode (Options)", + enable_expand=True, + category="SimpleSyrup/Latent", + description=VAEEncodeOptions.DESCRIPTION, + search_aliases=VAEEncodeOptions.SEARCH_ALIASES, + inputs=[ + _comfy_io.Boolean.Input( + "use_tiling", + default=False, + tooltip=tooltips.VAE_OPTIONS_USE_TILING, + ), + _comfy_io.Image.Input( + "pixels", + raw_link=True, + tooltip=tooltips.VAE_OPTIONS_ENCODE_PIXELS, + ), + _comfy_io.Vae.Input( + "vae", + raw_link=True, + tooltip=tooltips.VAE_OPTIONS_VAE, + ), + _comfy_io.Int.Input( + "tile_size", + default=VAE_TILE_SIZE_DEFAULT, + min=VAE_TILE_SIZE_MIN, + max=VAE_TILE_SIZE_MAX, + step=VAE_ENCODE_TILE_SIZE_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_TILE_SIZE, + ), + _comfy_io.Int.Input( + "overlap", + default=VAE_OVERLAP_DEFAULT, + min=VAE_OVERLAP_MIN, + max=VAE_OVERLAP_MAX, + step=VAE_OVERLAP_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_OVERLAP, + ), + _comfy_io.Int.Input( + "temporal_size", + default=VAE_TEMPORAL_SIZE_DEFAULT, + min=VAE_TEMPORAL_SIZE_MIN, + max=VAE_TEMPORAL_SIZE_MAX, + step=VAE_TEMPORAL_SIZE_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_ENCODE_TEMPORAL_SIZE, + ), + _comfy_io.Int.Input( + "temporal_overlap", + default=VAE_TEMPORAL_OVERLAP_DEFAULT, + min=VAE_TEMPORAL_OVERLAP_MIN, + max=VAE_TEMPORAL_OVERLAP_MAX, + step=VAE_TEMPORAL_OVERLAP_STEP, + advanced=True, + tooltip=tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP, + ), + ], + outputs=[ + _comfy_io.Latent.Output( + "latent", + tooltip=tooltips.VAE_OPTIONS_LATENT_OUTPUT, + ), + ], + ) + + @classmethod + def execute( + cls, + use_tiling: bool, + pixels: object, + vae: object, + tile_size: int, + overlap: int, + temporal_size: int, + temporal_overlap: int, + ) -> Any: + """Expand through the legacy VAE Encode (Options) implementation.""" + + return VAEEncodeOptions().encode( + use_tiling=use_tiling, + pixels=pixels, + vae=vae, + tile_size=tile_size, + overlap=overlap, + temporal_size=temporal_size, + temporal_overlap=temporal_overlap, + ) diff --git a/simple_syrup/runtime/comfy_graph_provenance.py b/simple_syrup/runtime/comfy_graph_provenance.py index 1828336..d2b46e4 100644 --- a/simple_syrup/runtime/comfy_graph_provenance.py +++ b/simple_syrup/runtime/comfy_graph_provenance.py @@ -18,6 +18,12 @@ from ..domain.graph_provenance import ( MAX_PROVENANCE_HOPS = 128 PASSTHROUGH_ATTRIBUTE = "GRAPH_PASSTHROUGH_OUTPUTS" +VAE_DECODE_SOURCE_CLASS_TYPES = frozenset( + { + "VAEDecode", + "SimpleSyrup.VAEDecodeOptions", + } +) PromptNode = Mapping[str, Any] PromptGraph = Mapping[str, Any] @@ -67,8 +73,8 @@ def trace_vae_decode_provenance( class_type=class_type, ) - if class_type == "VAEDecode": - return _trace_vae_decode(node_id, output_slot, inputs) + if class_type in VAE_DECODE_SOURCE_CLASS_TYPES: + return _trace_vae_decode(node_id, output_slot, inputs, class_type) class_def = node_registry.get(class_type) if class_def is None: @@ -135,15 +141,16 @@ def _trace_vae_decode( node_id: str, output_slot: int, inputs: Mapping[str, Any], + class_type: str, ) -> ProvenanceTrace: - """Resolve the latent and VAE links from a `VAEDecode` prompt node.""" + """Resolve latent and VAE links from a trusted VAE decode prompt node.""" image_output = (node_id, output_slot) if output_slot != 0: return BrokenProvenance( "VAEDecode output is not the image output", node_id=node_id, - class_type="VAEDecode", + class_type=class_type, ) samples_link = parse_graph_link(inputs.get("samples")) @@ -151,7 +158,7 @@ def _trace_vae_decode( return BrokenProvenance( "VAEDecode samples input is not a graph link", node_id=node_id, - class_type="VAEDecode", + class_type=class_type, ) return VaeDecodeProvenance( diff --git a/simple_syrup/runtime/detail_sampling.py b/simple_syrup/runtime/detail_sampling.py index 675df09..0044cc7 100644 --- a/simple_syrup/runtime/detail_sampling.py +++ b/simple_syrup/runtime/detail_sampling.py @@ -13,6 +13,7 @@ import torch from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback +from .differential_diffusion import clone_with_differential_diffusion Latent: TypeAlias = dict[str, Any] @@ -108,21 +109,7 @@ class DetailSampler: def apply_differential_diffusion(self, model: Any) -> Any: """Patch a model for feathered denoise masks when ComfyUI supports it.""" - options = getattr(model, "model_options", {}) - if ( - isinstance(options, dict) - and options.get("denoise_mask_function") is not None - ): - return model - - module = import_module("comfy_extras.nodes_differential_diffusion") - node = module.DifferentialDiffusion - output = node.execute(model, 1.0) - if hasattr(output, "result"): - return output.result[0] - if isinstance(output, tuple): - return output[0] - return output[0] + return clone_with_differential_diffusion(model) def _nodes() -> Any: diff --git a/simple_syrup/runtime/differential_diffusion.py b/simple_syrup/runtime/differential_diffusion.py new file mode 100644 index 0000000..063a423 --- /dev/null +++ b/simple_syrup/runtime/differential_diffusion.py @@ -0,0 +1,50 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Adapters for ComfyUI differential denoise-mask behavior.""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + + +def has_denoise_mask_function(model: Any) -> bool: + """Return whether a model patcher already has denoise-mask behavior.""" + + options = getattr(model, "model_options", {}) + return ( + isinstance(options, dict) and options.get("denoise_mask_function") is not None + ) + + +def clone_with_differential_diffusion(model: Any, strength: float = 1.0) -> Any: + """Return a clone patched with ComfyUI differential denoise masks.""" + + if has_denoise_mask_function(model): + return model + cloned_model = model.clone() + install_differential_diffusion(cloned_model, strength=strength) + return cloned_model + + +def install_differential_diffusion(model: Any, strength: float = 1.0) -> Any: + """Install ComfyUI differential denoise-mask behavior on a model patcher.""" + + if has_denoise_mask_function(model): + return model + set_mask_function = getattr(model, "set_model_denoise_mask_function", None) + if not callable(set_mask_function): + raise ValueError("Model does not support differential diffusion denoise masks.") + differential_diffusion = import_module( + "comfy_extras.nodes_differential_diffusion" + ).DifferentialDiffusion + set_mask_function( + lambda *args, **kwargs: differential_diffusion.forward( + *args, + **kwargs, + strength=strength, + ) + ) + return model diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index 75633ef..dc35eb5 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -26,6 +26,7 @@ from ..domain.tiled_diffusion import ( from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback +from .differential_diffusion import install_differential_diffusion from .tiled_sampling import ( ApplyModel, Latent, @@ -60,6 +61,7 @@ def sample_mixture_of_diffusers( latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample a latent with a cloned model patched for Mixture of Diffusers.""" @@ -105,6 +107,7 @@ def sample_mixture_of_diffusers( tile_height=latent_tile_height, overlap=latent_tile_overlap, tile_batch_size=latent_tile_batch_size, + differential_diffusion=differential_diffusion, ) batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None @@ -160,6 +163,7 @@ def clone_model_with_mixture_of_diffusers( tile_height: int, overlap: int, tile_batch_size: int, + differential_diffusion: bool = False, ) -> tuple[Any, TiledDiffusionPlan]: """Return a model clone patched with a pre-CFG Mixture wrapper.""" @@ -172,6 +176,8 @@ def clone_model_with_mixture_of_diffusers( tile_batch_size=tile_batch_size, ) cloned_model = model.clone() + if differential_diffusion: + install_differential_diffusion(cloned_model) old_wrapper = cloned_model.model_options.get("model_function_wrapper") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing model_function_wrapper is not callable.") diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index 9956722..472b213 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -25,6 +25,7 @@ from ..domain.tiled_diffusion import ( from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback +from .differential_diffusion import install_differential_diffusion from .tiled_sampling import ( ApplyModel, Latent, @@ -60,6 +61,7 @@ def sample_multidiffusion( latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample a latent with a cloned model patched for MultiDiffusion.""" @@ -106,6 +108,7 @@ def sample_multidiffusion( tile_height=latent_tile_height, overlap=latent_tile_overlap, tile_batch_size=latent_tile_batch_size, + differential_diffusion=differential_diffusion, ) batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None @@ -162,6 +165,7 @@ def clone_model_with_multidiffusion( tile_height: int, overlap: int, tile_batch_size: int, + differential_diffusion: bool = False, ) -> tuple[Any, TiledDiffusionPlan]: """Return a model clone patched with a pre-CFG MultiDiffusion wrapper.""" @@ -174,6 +178,8 @@ def clone_model_with_multidiffusion( tile_batch_size=tile_batch_size, ) cloned_model = model.clone() + if differential_diffusion: + install_differential_diffusion(cloned_model) old_wrapper = cloned_model.model_options.get("model_function_wrapper") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing model_function_wrapper is not callable.") diff --git a/simple_syrup/runtime/regional_multidiffusion_sampling.py b/simple_syrup/runtime/regional_multidiffusion_sampling.py index 83f5325..fdd0b8c 100644 --- a/simple_syrup/runtime/regional_multidiffusion_sampling.py +++ b/simple_syrup/runtime/regional_multidiffusion_sampling.py @@ -22,6 +22,7 @@ from ..domain.regional_detailing import LatentRegion from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback +from .differential_diffusion import install_differential_diffusion from .tiled_sampling import ( Latent, reject_unsupported_conditioning, @@ -62,6 +63,7 @@ def sample_regional_multidiffusion( denoise: float, global_prompt_weight: float, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample a latent with regional MultiDiffusion prompt blending.""" @@ -109,6 +111,7 @@ def sample_regional_multidiffusion( latent_ndim=latent_samples.ndim, regions=regions, global_prompt_weight=global_prompt_weight, + differential_diffusion=differential_diffusion, ) batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None @@ -163,6 +166,7 @@ def clone_model_with_regional_multidiffusion( latent_ndim: int, regions: tuple[LatentRegion, ...], global_prompt_weight: float = 0.0, + differential_diffusion: bool = False, ) -> tuple[Any, RegionalMultiDiffusionSummary]: """Return a model clone patched with regional calc-cond-batch blending.""" @@ -177,6 +181,8 @@ def clone_model_with_regional_multidiffusion( regions=regions, ) cloned_model = model.clone() + if differential_diffusion: + install_differential_diffusion(cloned_model) old_wrapper = cloned_model.model_options.get("sampler_calc_cond_batch_function") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing sampler_calc_cond_batch_function is not callable.") diff --git a/simple_syrup/runtime/vae_options_graph.py b/simple_syrup/runtime/vae_options_graph.py new file mode 100644 index 0000000..ce5513d --- /dev/null +++ b/simple_syrup/runtime/vae_options_graph.py @@ -0,0 +1,115 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Build ComfyUI VAE option node expansions without owning VAE execution.""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any, TypedDict, cast + + +class ExpansionResult(TypedDict): + """ComfyUI dynamic expansion result with one output link.""" + + expand: dict[str, dict[str, Any]] + result: tuple[list[object], ...] + + +class VAEOptionsGraphBuilder: + """Expand VAE option choices to ComfyUI's native VAE nodes.""" + + def build_encode( + self, + *, + pixels: object, + vae: object, + use_tiling: bool, + tile_size: int, + overlap: int, + temporal_size: int, + temporal_overlap: int, + ) -> ExpansionResult: + """Build a native VAE encode or tiled VAE encode expansion.""" + + inputs: dict[str, object] = { + "pixels": _graph_value(pixels), + "vae": _graph_value(vae), + } + class_type = "VAEEncode" + if use_tiling: + class_type = "VAEEncodeTiled" + inputs.update( + { + "tile_size": int(tile_size), + "overlap": int(overlap), + "temporal_size": int(temporal_size), + "temporal_overlap": int(temporal_overlap), + } + ) + + return _single_node_expansion(class_type, inputs) + + def build_decode( + self, + *, + samples: object, + vae: object, + use_tiling: bool, + tile_size: int, + overlap: int, + temporal_size: int, + temporal_overlap: int, + ) -> ExpansionResult: + """Build a native VAE decode or tiled VAE decode expansion.""" + + inputs: dict[str, object] = { + "samples": _graph_value(samples), + "vae": _graph_value(vae), + } + class_type = "VAEDecode" + if use_tiling: + class_type = "VAEDecodeTiled" + inputs.update( + { + "tile_size": int(tile_size), + "overlap": int(overlap), + "temporal_size": int(temporal_size), + "temporal_overlap": int(temporal_overlap), + } + ) + + return _single_node_expansion(class_type, inputs) + + +def _single_node_expansion( + class_type: str, + inputs: dict[str, object], +) -> ExpansionResult: + """Return a one-node dynamic expansion for a native ComfyUI node.""" + + builder = _graph_builder() + node = builder.node(class_type, **inputs) + return { + "expand": cast(dict[str, dict[str, Any]], builder.finalize()), + "result": (cast(list[object], node.out(0)),), + } + + +def _graph_value(value: object) -> object: + """Normalize tuple graph links to ComfyUI's serialized list shape.""" + + if isinstance(value, tuple) and len(value) == 2: + node_id, output_slot = value + if isinstance(node_id, str) and isinstance(output_slot, int): + return [node_id, output_slot] + return value + + +def _graph_builder() -> Any: + """Return ComfyUI's graph builder without import-time Comfy coupling.""" + + graph_utils = import_module("comfy_execution.graph_utils") + graph_builder = cast(Any, graph_utils.GraphBuilder) + return graph_builder() diff --git a/simple_syrup/services/detail_segs_as_regions_service.py b/simple_syrup/services/detail_segs_as_regions_service.py index ccf6237..9276b0b 100644 --- a/simple_syrup/services/detail_segs_as_regions_service.py +++ b/simple_syrup/services/detail_segs_as_regions_service.py @@ -64,12 +64,10 @@ class RegionalDetailSamplingBoundary(Protocol): denoise: float, global_prompt_weight: float, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample one full latent with paired regional conditioning.""" - def apply_differential_diffusion(self, model: Any) -> Any: - """Return a model patched for feathered denoise masks.""" - class RegionalDetailResizeBoundary(Protocol): """Image resize boundary used by regional detailing.""" @@ -133,6 +131,7 @@ class RegionalDetailSampler: denoise: float, global_prompt_weight: float, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample one full latent with regional MultiDiffusion.""" @@ -150,13 +149,9 @@ class RegionalDetailSampler: denoise=denoise, global_prompt_weight=global_prompt_weight, preview_context=preview_context, + differential_diffusion=differential_diffusion, ) - def apply_differential_diffusion(self, model: Any) -> Any: - """Patch a model for feathered denoise masks when ComfyUI supports it.""" - - return self._detail_sampler.apply_differential_diffusion(model) - class DetailSEGSAsRegionsService: """Detail provided SEGS through one regional MultiDiffusion pass.""" @@ -273,12 +268,10 @@ class DetailSEGSAsRegionsService: latent_for_sampling = ( self._with_noise_mask(latent, latent_regions) if noise_mask else latent ) - sampling_model = model - if noise_mask and noise_mask_feather > 0: - sampling_model = self._sampler.apply_differential_diffusion(model) + differential_diffusion = noise_mask and noise_mask_feather > 0 sampled = self._sampler.sample_regions( - model=sampling_model, + model=model, seed=seed, steps=steps, cfg=cfg, @@ -296,6 +289,7 @@ class DetailSEGSAsRegionsService: work_mask=image_union_mask, sampled_region=CropRegion(0, 0, image_width, image_height), ), + differential_diffusion=differential_diffusion, ) decoded = self._sampler.decode(vae, sampled, tiled_decode) if decoded.shape[1:3] != image_tensor.shape[1:3]: diff --git a/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py b/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py index 673ef87..1670c22 100644 --- a/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py +++ b/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py @@ -52,6 +52,7 @@ class TiledDiffusionLatentSamplingBoundary(Protocol): latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample a latent using the selected tiled diffusion mode.""" @@ -84,12 +85,10 @@ class TiledDetailSamplingBoundary(Protocol): latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample one latent crop with the requested tiled diffusion mode.""" - def apply_differential_diffusion(self, model: Any) -> Any: - """Return a model patched for feathered denoise masks.""" - class TiledDetailResizeBoundary(Protocol): """Image resize boundary used by tiled scale-factor detailing.""" @@ -163,6 +162,7 @@ class TiledDetailSampler: latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample one latent crop with the selected tiled diffusion runtime.""" @@ -183,13 +183,9 @@ class TiledDetailSampler: latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, + differential_diffusion=differential_diffusion, ) - def apply_differential_diffusion(self, model: Any) -> Any: - """Patch a model for feathered denoise masks when ComfyUI supports it.""" - - return self._detail_sampler.apply_differential_diffusion(model) - class DetailSEGSByScaleFactorTiledDiffusionService: """Detail SEGS crops with tiled diffusion latent sampling.""" @@ -251,9 +247,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService: return TiledDetailerResult(image=image_tensor.clone()) working_image = image_tensor.clone() - sampling_model = model - if noise_mask and noise_mask_feather > 0: - sampling_model = self._sampler.apply_differential_diffusion(model) + differential_diffusion = noise_mask and noise_mask_feather > 0 for index, segment in enumerate(segments): plan = build_detail_scale_plan( @@ -278,7 +272,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService: working_image=working_image, segment=segment, plan=plan, - model=sampling_model, + model=model, vae=vae, positive=select_conditioning(positive, index), negative=select_conditioning(negative, index), @@ -299,6 +293,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService: latent_tile_height=latent_tile_height, latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, + differential_diffusion=differential_diffusion, ) LOGGER.info( @@ -346,6 +341,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService: latent_tile_height: int, latent_tile_overlap: int, latent_tile_batch_size: int, + differential_diffusion: bool, ) -> torch.Tensor: """Detail one segment with tiled diffusion and return the updated image.""" @@ -380,6 +376,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService: latent_tile_height=latent_tile_height, latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, + differential_diffusion=differential_diffusion, preview_context=DetailPreviewContext( image=working_image, work_region=segment.crop_region, diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index a467cd8..2688779 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -37,6 +37,7 @@ class TiledDiffusionSamplingService: latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Sample a latent with the selected tiled diffusion method.""" @@ -58,6 +59,7 @@ class TiledDiffusionSamplingService: latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, + differential_diffusion=differential_diffusion, ) return mixture_of_diffusers_sampling.sample_mixture_of_diffusers( model=model, @@ -75,4 +77,5 @@ class TiledDiffusionSamplingService: latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, + differential_diffusion=differential_diffusion, ) diff --git a/tests/test_detail_segs_as_regions_service.py b/tests/test_detail_segs_as_regions_service.py index dfb5cff..b7c8d14 100644 --- a/tests/test_detail_segs_as_regions_service.py +++ b/tests/test_detail_segs_as_regions_service.py @@ -297,8 +297,8 @@ def test_noise_mask_false_omits_latent_mask_but_composites_pixels() -> None: assert torch.all(result.image[:, 4:, :, :] == 0.0) -def test_noise_mask_feather_applies_differential_diffusion() -> None: - """Feathered denoise masks request differential diffusion patching.""" +def test_noise_mask_feather_requests_single_clone_differential_diffusion() -> None: + """Feathered denoise masks are composed inside the regional runtime clone.""" sampler = _FakeRegionalSampler() service = _service(sampler) @@ -314,8 +314,9 @@ def test_noise_mask_feather_applies_differential_diffusion() -> None: **(_settings() | {"noise_mask": True, "noise_mask_feather": 2}), ) - assert sampler.patch_count == 1 - assert sampler.sample_calls[0]["model"] == "patched model" + assert sampler.patch_count == 0 + assert sampler.sample_calls[0]["model"] == "model" + assert sampler.sample_calls[0]["differential_diffusion"] is True def test_encode_decode_and_sampler_controls_are_forwarded() -> None: @@ -479,6 +480,7 @@ class _FakeRegionalSampler: denoise: float, global_prompt_weight: float, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Record regional sample options and return the latent unchanged.""" @@ -497,6 +499,7 @@ class _FakeRegionalSampler: "denoise": denoise, "global_prompt_weight": global_prompt_weight, "preview_context": preview_context, + "differential_diffusion": differential_diffusion, } ) return latent_image diff --git a/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py b/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py index b73b63a..73be82f 100644 --- a/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py +++ b/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py @@ -211,6 +211,29 @@ def test_tiled_noise_mask_feather_keeps_sampled_crop_geometry() -> None: assert float(noise_mask[0, 1, 1]) > float(noise_mask[0, 0, 1]) +def test_tiled_noise_mask_feather_requests_single_clone_differential_diffusion() -> ( + None +): + """Feathered denoise masks are composed inside the tiled runtime clone.""" + + sampler = _FakeTiledSampler() + service = _service(sampler) + + service.detail( + _image(), + _segs(_segment()), + "model", + "vae", + [], + [], + **(_settings() | {"noise_mask": True, "noise_mask_feather": 2}), + ) + + assert sampler.patch_count == 0 + assert sampler.sample_calls[0]["model"] == "model" + assert sampler.sample_calls[0]["differential_diffusion"] is True + + def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None: """The detailer adapter uses the shared tiled diffusion dispatch service.""" @@ -241,6 +264,7 @@ def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None: latent_tile_overlap=12, latent_tile_batch_size=3, preview_context=preview_context, + differential_diffusion=True, ) assert result is latent @@ -248,6 +272,7 @@ def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None: assert call["diffusion_mode"] == "mixture_of_diffusers" assert call["latent_image"] is latent assert call["preview_context"] is preview_context + assert call["differential_diffusion"] is True class _FakeTiledSamplingService: @@ -277,6 +302,7 @@ class _FakeTiledSamplingService: latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Record tiled sampling arguments and return the latent unchanged.""" @@ -298,6 +324,7 @@ class _FakeTiledSamplingService: "latent_tile_overlap": latent_tile_overlap, "latent_tile_batch_size": latent_tile_batch_size, "preview_context": preview_context, + "differential_diffusion": differential_diffusion, } ) return latent_image @@ -356,6 +383,7 @@ class _FakeTiledSampler: latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, ) -> Latent: """Record tiled sample options and return the latent unchanged.""" @@ -377,6 +405,7 @@ class _FakeTiledSampler: "latent_tile_overlap": latent_tile_overlap, "latent_tile_batch_size": latent_tile_batch_size, "preview_context": preview_context, + "differential_diffusion": differential_diffusion, } ) return latent_image diff --git a/tests/test_graph_provenance.py b/tests/test_graph_provenance.py index 5906ede..ddbf023 100644 --- a/tests/test_graph_provenance.py +++ b/tests/test_graph_provenance.py @@ -47,6 +47,33 @@ def test_direct_vae_decode_resolves_samples_and_vae_links() -> None: assert result.vae_link == ("loader", 2) +def test_vae_decode_options_resolves_samples_and_vae_links() -> None: + """VAE Decode (Options) output resolves to the source latent and VAE.""" + + prompt = { + "decode": { + "class_type": "SimpleSyrup.VAEDecodeOptions", + "inputs": { + "use_tiling": True, + "samples": ["latent", 0], + "vae": ["loader", 2], + "tile_size": 512, + "overlap": 64, + "temporal_size": 64, + "temporal_overlap": 8, + }, + } + } + + result = trace_vae_decode_provenance(prompt, ["decode", 0], {}) + + assert isinstance(result, VaeDecodeProvenance) + assert result.decode_node_id == "decode" + assert result.image_output == ("decode", 0) + assert result.samples_link == ("latent", 0) + assert result.vae_link == ("loader", 2) + + def test_transparent_node_resolves_to_upstream_vae_decode() -> None: """A declared pass-through node is traversed to its source link.""" @@ -137,6 +164,25 @@ def test_vae_decode_non_image_output_breaks_provenance() -> None: assert result.reason == "VAEDecode output is not the image output" +def test_vae_decode_options_non_image_output_breaks_provenance() -> None: + """Only VAE Decode (Options) output slot 0 is trusted as image provenance.""" + + result = trace_vae_decode_provenance( + { + "decode": { + "class_type": "SimpleSyrup.VAEDecodeOptions", + "inputs": {"samples": ["latent", 0], "vae": ["loader", 2]}, + } + }, + ["decode", 1], + {}, + ) + + assert isinstance(result, BrokenProvenance) + assert result.reason == "VAEDecode output is not the image output" + assert result.class_type == "SimpleSyrup.VAEDecodeOptions" + + def test_missing_samples_link_breaks_provenance() -> None: """VAEDecode without a graph-linked samples input cannot supply provenance.""" @@ -150,6 +196,25 @@ def test_missing_samples_link_breaks_provenance() -> None: assert result.reason == "VAEDecode samples input is not a graph link" +def test_vae_decode_options_missing_samples_link_breaks_provenance() -> None: + """VAE Decode (Options) requires a graph-linked samples input.""" + + result = trace_vae_decode_provenance( + { + "decode": { + "class_type": "SimpleSyrup.VAEDecodeOptions", + "inputs": {"samples": "latent"}, + } + }, + ["decode", 0], + {}, + ) + + assert isinstance(result, BrokenProvenance) + assert result.reason == "VAEDecode samples input is not a graph link" + assert result.class_type == "SimpleSyrup.VAEDecodeOptions" + + def test_missing_node_breaks_provenance() -> None: """Missing source nodes produce broken provenance.""" diff --git a/tests/test_mixture_of_diffusers_sampling.py b/tests/test_mixture_of_diffusers_sampling.py index 0047775..19c8c60 100644 --- a/tests/test_mixture_of_diffusers_sampling.py +++ b/tests/test_mixture_of_diffusers_sampling.py @@ -83,6 +83,7 @@ class FakeModel: def __init__( self, model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, ) -> None: """Create a fake model patcher.""" @@ -90,11 +91,14 @@ class FakeModel: self.model_options = {} if model_options is None else model_options self.wrapper: Any = None self.model_sampling = object() + self.parent = parent + self.clone_count = 0 def clone(self) -> FakeModel: """Return a cloned model with copied options.""" - return FakeModel(self.model_options.copy()) + self.clone_count += 1 + return FakeModel(self.model_options.copy(), parent=self) def set_model_unet_function_wrapper(self, wrapper: object) -> None: """Capture the installed model function wrapper.""" @@ -102,6 +106,11 @@ class FakeModel: self.wrapper = wrapper self.model_options["model_function_wrapper"] = wrapper + def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None: + """Capture the installed denoise-mask function.""" + + self.model_options["denoise_mask_function"] = denoise_mask_function + def get_model_object(self, name: str) -> object: """Return the requested fake model object.""" @@ -119,6 +128,31 @@ class FakeSampler: return None +def test_clone_model_with_mixture_composes_differential_on_same_clone() -> None: + """Differential diffusion is installed without cloning a temporary parent.""" + + model = FakeModel() + + wrapped_model, _plan = mod_sampling.clone_model_with_mixture_of_diffusers( + model, + latent_width=8, + latent_height=4, + tile_width=4, + tile_height=4, + overlap=0, + tile_batch_size=2, + differential_diffusion=True, + ) + + assert model.clone_count == 1 + assert wrapped_model.parent is model + assert callable(wrapped_model.model_options["denoise_mask_function"]) + assert isinstance( + wrapped_model.wrapper, + mod_sampling.MixtureOfDiffusersModelWrapper, + ) + + def test_model_wrapper_tiles_input_conditioning_and_transformer_options() -> None: """The wrapper tiles latents, conditioning tensors, timesteps, and metadata.""" diff --git a/tests/test_multidiffusion_sampling.py b/tests/test_multidiffusion_sampling.py index 7562840..8f6ac87 100644 --- a/tests/test_multidiffusion_sampling.py +++ b/tests/test_multidiffusion_sampling.py @@ -91,6 +91,7 @@ class FakeModel: def __init__( self, model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, ) -> None: """Create a fake model patcher.""" @@ -98,11 +99,14 @@ class FakeModel: self.model_options = {} if model_options is None else model_options self.wrapper: Any = None self.model_sampling = object() + self.parent = parent + self.clone_count = 0 def clone(self) -> FakeModel: """Return a cloned model with copied options.""" - return FakeModel(self.model_options.copy()) + self.clone_count += 1 + return FakeModel(self.model_options.copy(), parent=self) def set_model_unet_function_wrapper(self, wrapper: object) -> None: """Capture the installed model function wrapper.""" @@ -110,6 +114,11 @@ class FakeModel: self.wrapper = wrapper self.model_options["model_function_wrapper"] = wrapper + def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None: + """Capture the installed denoise-mask function.""" + + self.model_options["denoise_mask_function"] = denoise_mask_function + def get_model_object(self, name: str) -> object: """Return the requested fake model object.""" @@ -150,6 +159,31 @@ def test_clone_model_with_multidiffusion_installs_wrapper() -> None: assert plan.tile_batch_size == 2 +def test_clone_model_with_multidiffusion_composes_differential_on_same_clone() -> None: + """Differential diffusion is installed without cloning a temporary parent.""" + + model = FakeModel() + + wrapped_model, _plan = multidiffusion_sampling.clone_model_with_multidiffusion( + model, + latent_width=8, + latent_height=4, + tile_width=4, + tile_height=4, + overlap=0, + tile_batch_size=2, + differential_diffusion=True, + ) + + assert model.clone_count == 1 + assert wrapped_model.parent is model + assert callable(wrapped_model.model_options["denoise_mask_function"]) + assert isinstance( + wrapped_model.wrapper, + multidiffusion_sampling.MultiDiffusionModelWrapper, + ) + + def test_clone_model_rejects_non_callable_existing_wrapper() -> None: """Existing wrapper metadata must be callable.""" diff --git a/tests/test_node_tooltips.py b/tests/test_node_tooltips.py index 2c0a7cd..3988dfb 100644 --- a/tests/test_node_tooltips.py +++ b/tests/test_node_tooltips.py @@ -20,6 +20,8 @@ from simple_syrup.nodes_v3.encode_prompt_batch_with_prompt_control import ( from simple_syrup.nodes_v3.scale_factor import ScaleFactorV3 from simple_syrup.nodes_v3.simple_load_checkpoint import SimpleLoadCheckpointV3 from simple_syrup.nodes_v3.tile_and_tag_segs import TileAndTagSEGSV3 +from simple_syrup.nodes_v3.vae_decode_options import VAEDecodeOptionsV3 +from simple_syrup.nodes_v3.vae_encode_options import VAEEncodeOptionsV3 from simple_syrup.nodes_v3.wd14_tagger_loader import WD14TaggerLoaderV3 @@ -116,6 +118,8 @@ def test_legacy_named_outputs_provide_tooltips() -> None: SimpleLoadCheckpointV3, ScaleFactorV3, TileAndTagSEGSV3, + VAEDecodeOptionsV3, + VAEEncodeOptionsV3, WD14TaggerLoaderV3, EncodePromptBatchWithPromptControl, ], diff --git a/tests/test_regional_multidiffusion_sampling.py b/tests/test_regional_multidiffusion_sampling.py index 766bdf4..486180d 100644 --- a/tests/test_regional_multidiffusion_sampling.py +++ b/tests/test_regional_multidiffusion_sampling.py @@ -94,6 +94,7 @@ class FakeModel: def __init__( self, model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, ) -> None: """Create a fake model patcher.""" @@ -101,11 +102,14 @@ class FakeModel: self.model_options = {} if model_options is None else model_options self.calc_wrapper: Any = None self.model_sampling = object() + self.parent = parent + self.clone_count = 0 def clone(self) -> FakeModel: """Return a cloned model with copied options.""" - return FakeModel(self.model_options.copy()) + self.clone_count += 1 + return FakeModel(self.model_options.copy(), parent=self) def set_model_sampler_calc_cond_batch_function(self, wrapper: object) -> None: """Capture the installed calc-cond-batch wrapper.""" @@ -113,6 +117,11 @@ class FakeModel: self.calc_wrapper = wrapper self.model_options["sampler_calc_cond_batch_function"] = wrapper + def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None: + """Capture the installed denoise-mask function.""" + + self.model_options["denoise_mask_function"] = denoise_mask_function + def get_model_object(self, name: str) -> object: """Return the requested fake model object.""" @@ -153,6 +162,31 @@ def test_clone_model_installs_regional_calc_cond_batch_wrapper() -> None: assert summary.region_count == 1 +def test_clone_model_composes_differential_on_same_clone() -> None: + """Differential diffusion is installed without cloning a temporary parent.""" + + model = FakeModel() + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + model, + latent_width=8, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "positive", latent_width=8),), + differential_diffusion=True, + ) + ) + + assert model.clone_count == 1 + assert wrapped_model.parent is model + assert callable(wrapped_model.model_options["denoise_mask_function"]) + assert isinstance( + wrapped_model.calc_wrapper, + regional_multidiffusion_sampling.RegionalMultiDiffusionCalcCondBatch, + ) + + def test_clone_model_rejects_non_callable_existing_calc_wrapper() -> None: """Existing calc-cond-batch metadata must be callable.""" diff --git a/tests/test_registration.py b/tests/test_registration.py index 501e47e..d550fb5 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -476,6 +476,29 @@ def test_provenance_latent_nodes_are_registered() -> None: assert "UpscaleLatentFromImage" in nodes_package.__all__ +def test_vae_options_nodes_are_registered() -> None: + """VAE options nodes map to their classes and display names.""" + + package = importlib.import_module("SimpleSyrup") + nodes_package = importlib.import_module("SimpleSyrup.simple_syrup.nodes") + + encode = package.NODE_CLASS_MAPPINGS["SimpleSyrup.VAEEncodeOptions"] + decode = package.NODE_CLASS_MAPPINGS["SimpleSyrup.VAEDecodeOptions"] + + assert encode.__name__ == "VAEEncodeOptions" + assert decode.__name__ == "VAEDecodeOptions" + assert ( + package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.VAEEncodeOptions"] + == "VAE Encode (Options)" + ) + assert ( + package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.VAEDecodeOptions"] + == "VAE Decode (Options)" + ) + assert "VAEEncodeOptions" in nodes_package.__all__ + assert "VAEDecodeOptions" in nodes_package.__all__ + + def test_registration_import_does_not_require_torchlanc() -> None: """Importing registration does not eagerly import TorchLanc.""" @@ -504,6 +527,8 @@ def test_v3_entrypoint_registers_tile_and_prompt_control_batch_nodes( "TileAndTagSEGSV3", "SimpleLoadCheckpointV3", "ScaleFactorV3", + "VAEDecodeOptionsV3", + "VAEEncodeOptionsV3", "EncodePromptBatchWithPromptControl", ] assert "prompt_control.nodes_lazy" not in sys.modules @@ -527,5 +552,7 @@ def test_v3_entrypoint_keeps_tile_node_when_prompt_control_unavailable( "TileAndTagSEGSV3", "SimpleLoadCheckpointV3", "ScaleFactorV3", + "VAEDecodeOptionsV3", + "VAEEncodeOptionsV3", ] assert "prompt_control.nodes_lazy" not in sys.modules diff --git a/tests/test_tiled_diffusion_sampling_service.py b/tests/test_tiled_diffusion_sampling_service.py index d492b7d..fa616d1 100644 --- a/tests/test_tiled_diffusion_sampling_service.py +++ b/tests/test_tiled_diffusion_sampling_service.py @@ -129,6 +129,35 @@ def test_service_forwards_sampling_arguments_unchanged( } +def test_service_forwards_differential_diffusion_request( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The dispatcher preserves differential-denoise-mask composition requests.""" + + calls: dict[str, Any] = {} + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record forwarded arguments.""" + + calls.update(kwargs) + return {"samples": torch.ones((1, 4, 4, 4))} + + monkeypatch.setattr( + "simple_syrup.services.tiled_diffusion_sampling_service." + "multidiffusion_sampling.sample_multidiffusion", + fake_multidiffusion, + ) + + TiledDiffusionSamplingService().sample( + **( + _sample_kwargs(diffusion_mode="multidiffusion") + | {"differential_diffusion": True} + ) + ) + + assert calls["differential_diffusion"] is True + + def test_invalid_mode_fails_before_runtime_call( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -181,4 +210,5 @@ def _sample_kwargs( "latent_tile_overlap": 24, "latent_tile_batch_size": 3, "preview_context": preview_context, + "differential_diffusion": False, } diff --git a/tests/test_vae_decode_options_v3_node.py b/tests/test_vae_decode_options_v3_node.py new file mode 100644 index 0000000..13391f5 --- /dev/null +++ b/tests/test_vae_decode_options_v3_node.py @@ -0,0 +1,78 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the VAE Decode (Options) Comfy v3 wrapper.""" + +from __future__ import annotations + +from importlib import import_module + +from simple_syrup.nodes import tooltips +from simple_syrup.nodes.vae_options import ( + VAE_DECODE_TILE_SIZE_STEP, + VAE_OVERLAP_DEFAULT, + VAE_TEMPORAL_OVERLAP_DEFAULT, + VAE_TEMPORAL_SIZE_DEFAULT, + VAE_TILE_SIZE_DEFAULT, + VAEDecodeOptions, +) +from simple_syrup.nodes_v3.vae_decode_options import VAEDecodeOptionsV3 + + +def test_vae_decode_options_v3_schema() -> None: + """VAE Decode (Options) v3 schema mirrors the legacy node contract.""" + + schema = VAEDecodeOptionsV3.define_schema() + + assert schema.node_id == "SimpleSyrup.VAEDecodeOptions" + assert schema.display_name == "VAE Decode (Options)" + assert schema.category == "SimpleSyrup/Latent" + assert schema.description == VAEDecodeOptions.DESCRIPTION + assert schema.enable_expand is True + assert [input_item.id for input_item in schema.inputs] == [ + "use_tiling", + "samples", + "vae", + "tile_size", + "overlap", + "temporal_size", + "temporal_overlap", + ] + assert schema.inputs[0].io_type == "BOOLEAN" + assert schema.inputs[0].default is False + assert schema.inputs[0].tooltip == tooltips.VAE_OPTIONS_USE_TILING + assert schema.inputs[1].io_type == "LATENT" + assert schema.inputs[1].rawLink is True + assert schema.inputs[2].io_type == "VAE" + assert schema.inputs[2].rawLink is True + assert schema.inputs[3].default == VAE_TILE_SIZE_DEFAULT + assert schema.inputs[3].step == VAE_DECODE_TILE_SIZE_STEP + assert schema.inputs[4].default == VAE_OVERLAP_DEFAULT + assert schema.inputs[5].default == VAE_TEMPORAL_SIZE_DEFAULT + assert schema.inputs[6].default == VAE_TEMPORAL_OVERLAP_DEFAULT + assert [output.id for output in schema.outputs] == ["image"] + assert schema.outputs[0].io_type == "IMAGE" + assert schema.outputs[0].tooltip == tooltips.VAE_OPTIONS_IMAGE_OUTPUT + + +def test_vae_decode_options_v3_execute_delegates_to_legacy_node() -> None: + """VAE Decode (Options) v3 execution returns legacy expansion behavior.""" + + graph_utils = import_module("comfy_execution.graph_utils") + graph_utils.GraphBuilder.set_default_prefix("V3_DECODE", 0, 0) + + result = VAEDecodeOptionsV3.execute( + True, + ["latent", 0], + ["loader", 2], + 1024, + 128, + 96, + 16, + ) + node = next(iter(result["expand"].values())) + + assert node["class_type"] == "VAEDecodeTiled" + assert node["inputs"]["tile_size"] == 1024 + assert result["result"][0][1] == 0 diff --git a/tests/test_vae_encode_options_v3_node.py b/tests/test_vae_encode_options_v3_node.py new file mode 100644 index 0000000..8c3dcb2 --- /dev/null +++ b/tests/test_vae_encode_options_v3_node.py @@ -0,0 +1,78 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the VAE Encode (Options) Comfy v3 wrapper.""" + +from __future__ import annotations + +from importlib import import_module + +from simple_syrup.nodes import tooltips +from simple_syrup.nodes.vae_options import ( + VAE_ENCODE_TILE_SIZE_STEP, + VAE_OVERLAP_DEFAULT, + VAE_TEMPORAL_OVERLAP_DEFAULT, + VAE_TEMPORAL_SIZE_DEFAULT, + VAE_TILE_SIZE_DEFAULT, + VAEEncodeOptions, +) +from simple_syrup.nodes_v3.vae_encode_options import VAEEncodeOptionsV3 + + +def test_vae_encode_options_v3_schema() -> None: + """VAE Encode (Options) v3 schema mirrors the legacy node contract.""" + + schema = VAEEncodeOptionsV3.define_schema() + + assert schema.node_id == "SimpleSyrup.VAEEncodeOptions" + assert schema.display_name == "VAE Encode (Options)" + assert schema.category == "SimpleSyrup/Latent" + assert schema.description == VAEEncodeOptions.DESCRIPTION + assert schema.enable_expand is True + assert [input_item.id for input_item in schema.inputs] == [ + "use_tiling", + "pixels", + "vae", + "tile_size", + "overlap", + "temporal_size", + "temporal_overlap", + ] + assert schema.inputs[0].io_type == "BOOLEAN" + assert schema.inputs[0].default is False + assert schema.inputs[0].tooltip == tooltips.VAE_OPTIONS_USE_TILING + assert schema.inputs[1].io_type == "IMAGE" + assert schema.inputs[1].rawLink is True + assert schema.inputs[2].io_type == "VAE" + assert schema.inputs[2].rawLink is True + assert schema.inputs[3].default == VAE_TILE_SIZE_DEFAULT + assert schema.inputs[3].step == VAE_ENCODE_TILE_SIZE_STEP + assert schema.inputs[4].default == VAE_OVERLAP_DEFAULT + assert schema.inputs[5].default == VAE_TEMPORAL_SIZE_DEFAULT + assert schema.inputs[6].default == VAE_TEMPORAL_OVERLAP_DEFAULT + assert [output.id for output in schema.outputs] == ["latent"] + assert schema.outputs[0].io_type == "LATENT" + assert schema.outputs[0].tooltip == tooltips.VAE_OPTIONS_LATENT_OUTPUT + + +def test_vae_encode_options_v3_execute_delegates_to_legacy_node() -> None: + """VAE Encode (Options) v3 execution returns legacy expansion behavior.""" + + graph_utils = import_module("comfy_execution.graph_utils") + graph_utils.GraphBuilder.set_default_prefix("V3_ENCODE", 0, 0) + + result = VAEEncodeOptionsV3.execute( + True, + ["image", 0], + ["loader", 2], + 768, + 96, + 48, + 12, + ) + node = next(iter(result["expand"].values())) + + assert node["class_type"] == "VAEEncodeTiled" + assert node["inputs"]["tile_size"] == 768 + assert result["result"][0][1] == 0 diff --git a/tests/test_vae_options_node.py b/tests/test_vae_options_node.py new file mode 100644 index 0000000..25c4ce0 --- /dev/null +++ b/tests/test_vae_options_node.py @@ -0,0 +1,226 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for VAE Encode/Decode Options nodes.""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any + +from simple_syrup.nodes.vae_options import ( + VAE_DECODE_TILE_SIZE_STEP, + VAE_ENCODE_TILE_SIZE_STEP, + VAE_OVERLAP_DEFAULT, + VAE_OVERLAP_MAX, + VAE_OVERLAP_MIN, + VAE_OVERLAP_STEP, + VAE_TEMPORAL_OVERLAP_DEFAULT, + VAE_TEMPORAL_OVERLAP_MAX, + VAE_TEMPORAL_OVERLAP_MIN, + VAE_TEMPORAL_OVERLAP_STEP, + VAE_TEMPORAL_SIZE_DEFAULT, + VAE_TEMPORAL_SIZE_MAX, + VAE_TEMPORAL_SIZE_MIN, + VAE_TEMPORAL_SIZE_STEP, + VAE_TILE_SIZE_DEFAULT, + VAE_TILE_SIZE_MAX, + VAE_TILE_SIZE_MIN, + VAEDecodeOptions, + VAEEncodeOptions, +) + + +def test_vae_encode_options_declares_inputs() -> None: + """VAE Encode (Options) exposes normal/tiled controls.""" + + inputs = VAEEncodeOptions.INPUT_TYPES()["required"] + + assert VAEEncodeOptions.RETURN_TYPES == ("LATENT",) + assert VAEEncodeOptions.RETURN_NAMES == ("latent",) + assert list(inputs) == [ + "use_tiling", + "pixels", + "vae", + "tile_size", + "overlap", + "temporal_size", + "temporal_overlap", + ] + assert inputs["use_tiling"][0] == "BOOLEAN" + assert inputs["use_tiling"][1]["default"] is False + assert inputs["pixels"][0] == "IMAGE" + assert inputs["pixels"][1]["rawLink"] is True + assert inputs["vae"][0] == "VAE" + assert inputs["vae"][1]["rawLink"] is True + _assert_tile_controls(inputs, tile_size_step=VAE_ENCODE_TILE_SIZE_STEP) + + +def test_vae_decode_options_declares_inputs() -> None: + """VAE Decode (Options) exposes normal/tiled controls.""" + + inputs = VAEDecodeOptions.INPUT_TYPES()["required"] + + assert VAEDecodeOptions.RETURN_TYPES == ("IMAGE",) + assert VAEDecodeOptions.RETURN_NAMES == ("image",) + assert list(inputs) == [ + "use_tiling", + "samples", + "vae", + "tile_size", + "overlap", + "temporal_size", + "temporal_overlap", + ] + assert inputs["use_tiling"][0] == "BOOLEAN" + assert inputs["use_tiling"][1]["default"] is False + assert inputs["samples"][0] == "LATENT" + assert inputs["samples"][1]["rawLink"] is True + assert inputs["vae"][0] == "VAE" + assert inputs["vae"][1]["rawLink"] is True + _assert_tile_controls(inputs, tile_size_step=VAE_DECODE_TILE_SIZE_STEP) + + +def test_vae_encode_options_expands_to_native_encode() -> None: + """Normal encode mode expands to ComfyUI's native VAEEncode.""" + + _set_graph_prefix("ENCODE_NORMAL") + result = VAEEncodeOptions().encode( + use_tiling=False, + pixels=("image", 0), + vae=("loader", 2), + tile_size=512, + overlap=64, + temporal_size=64, + temporal_overlap=8, + ) + + node = _single_node(result["expand"]) + assert node["class_type"] == "VAEEncode" + assert node["inputs"] == {"pixels": ["image", 0], "vae": ["loader", 2]} + assert result["result"][0][0] in result["expand"] + assert result["result"][0][1] == 0 + + +def test_vae_encode_options_expands_to_native_tiled_encode() -> None: + """Tiled encode mode expands to ComfyUI's native VAEEncodeTiled.""" + + _set_graph_prefix("ENCODE_TILED") + result = VAEEncodeOptions().encode( + use_tiling=True, + pixels=["image", 0], + vae=["loader", 2], + tile_size=768, + overlap=96, + temporal_size=48, + temporal_overlap=12, + ) + + node = _single_node(result["expand"]) + assert node["class_type"] == "VAEEncodeTiled" + assert node["inputs"] == { + "pixels": ["image", 0], + "vae": ["loader", 2], + "tile_size": 768, + "overlap": 96, + "temporal_size": 48, + "temporal_overlap": 12, + } + + +def test_vae_decode_options_expands_to_native_decode() -> None: + """Normal decode mode expands to ComfyUI's native VAEDecode.""" + + _set_graph_prefix("DECODE_NORMAL") + result = VAEDecodeOptions().decode( + use_tiling=False, + samples=("latent", 0), + vae=("loader", 2), + tile_size=512, + overlap=64, + temporal_size=64, + temporal_overlap=8, + ) + + node = _single_node(result["expand"]) + assert node["class_type"] == "VAEDecode" + assert node["inputs"] == {"samples": ["latent", 0], "vae": ["loader", 2]} + assert result["result"][0][0] in result["expand"] + assert result["result"][0][1] == 0 + + +def test_vae_decode_options_expands_to_native_tiled_decode() -> None: + """Tiled decode mode expands to ComfyUI's native VAEDecodeTiled.""" + + _set_graph_prefix("DECODE_TILED") + result = VAEDecodeOptions().decode( + use_tiling=True, + samples=["latent", 0], + vae=["loader", 2], + tile_size=1024, + overlap=128, + temporal_size=96, + temporal_overlap=16, + ) + + node = _single_node(result["expand"]) + assert node["class_type"] == "VAEDecodeTiled" + assert node["inputs"] == { + "samples": ["latent", 0], + "vae": ["loader", 2], + "tile_size": 1024, + "overlap": 128, + "temporal_size": 96, + "temporal_overlap": 16, + } + + +def _assert_tile_controls( + inputs: dict[str, tuple[Any, ...]], + *, + tile_size_step: int, +) -> None: + """Assert shared VAE tile control metadata.""" + + tile_size = inputs["tile_size"][1] + assert tile_size["default"] == VAE_TILE_SIZE_DEFAULT + assert tile_size["min"] == VAE_TILE_SIZE_MIN + assert tile_size["max"] == VAE_TILE_SIZE_MAX + assert tile_size["step"] == tile_size_step + assert tile_size["advanced"] is True + + overlap = inputs["overlap"][1] + assert overlap["default"] == VAE_OVERLAP_DEFAULT + assert overlap["min"] == VAE_OVERLAP_MIN + assert overlap["max"] == VAE_OVERLAP_MAX + assert overlap["step"] == VAE_OVERLAP_STEP + assert overlap["advanced"] is True + + temporal_size = inputs["temporal_size"][1] + assert temporal_size["default"] == VAE_TEMPORAL_SIZE_DEFAULT + assert temporal_size["min"] == VAE_TEMPORAL_SIZE_MIN + assert temporal_size["max"] == VAE_TEMPORAL_SIZE_MAX + assert temporal_size["step"] == VAE_TEMPORAL_SIZE_STEP + assert temporal_size["advanced"] is True + + temporal_overlap = inputs["temporal_overlap"][1] + assert temporal_overlap["default"] == VAE_TEMPORAL_OVERLAP_DEFAULT + assert temporal_overlap["min"] == VAE_TEMPORAL_OVERLAP_MIN + assert temporal_overlap["max"] == VAE_TEMPORAL_OVERLAP_MAX + assert temporal_overlap["step"] == VAE_TEMPORAL_OVERLAP_STEP + assert temporal_overlap["advanced"] is True + + +def _single_node(graph: dict[str, dict[str, Any]]) -> dict[str, Any]: + """Return the only node from a dynamic expansion graph.""" + + assert len(graph) == 1 + return next(iter(graph.values())) + + +def _set_graph_prefix(prefix: str) -> None: + """Set a deterministic Comfy graph-builder prefix for assertions.""" + + graph_utils = import_module("comfy_execution.graph_utils") + graph_utils.GraphBuilder.set_default_prefix(prefix, 0, 0)