Compare commits

...
10 Commits
Author SHA1 Message Date
Daisy 05667f6517 chore(release): 1.13.0 [skip ci]
# [1.13.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.12.0...v1.13.0) (2026-10-02)

### Bug Fixes

* **negpip:** support Krea attention on ComfyUI 0.28 ([e4eabfe](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/e4eabfefd540f3a6066c29761cefb8f8b3c80f68))

### Features

* **sampling:** add noise inversion and composable sampler options ([be4bd9b](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/be4bd9bb162f6ed4252e0c4d4ef0f3efebe8c674))
* **sampling:** refine sampler options and inversion controls ([ab7cebc](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/ab7cebcebc4d5dfa95aa8d834456af1778fe56a6))
2026-10-02 17:16:58 +00:00
Artificial Sweetener ab7cebcebc feat(sampling): refine sampler options and inversion controls 2026-10-02 12:51:27 -04:00
Artificial Sweetener e4eabfefd5 fix(negpip): support Krea attention on ComfyUI 0.28 2026-09-30 21:48:07 -04:00
Artificial Sweetener bbe86060c6 chore(licensing): complete inversion source and test headers 2026-09-30 21:48:07 -04:00
Artificial Sweetener be4bd9bb16 feat(sampling): add noise inversion and composable sampler options 2026-09-30 20:50:48 -04:00
Daisy 913cc7ba55 chore(release): 1.12.0 [skip ci]
# [1.12.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.11.1...v1.12.0) (2026-09-29)

### Bug Fixes

* **release:** satisfy RES4LYF publication contracts ([04dc35a](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/04dc35af00163689509cace1541a8e7d33b5053e))

### Features

* **sampling:** add RES4LYF sampler methods and schedules ([7b1efef](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/7b1efef47e20b9bf6f032dd57971899d99881822))
2026-09-29 17:24:20 +00:00
Artificial Sweetener 04dc35af00 fix(release): satisfy RES4LYF publication contracts 2026-09-29 13:16:28 -04:00
Artificial Sweetener 7b1efef47e feat(sampling): add RES4LYF sampler methods and schedules 2026-09-29 13:03:55 -04:00
Daisy ec49b685db chore(release): 1.11.1 [skip ci]
## [1.11.1](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.11.0...v1.11.1) (2026-09-25)

### Bug Fixes

* **release:** attribute automation to Daisy ([2ae545d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/2ae545d64d70a454f635ee647fb7e6a1c3500b9b))
2026-09-25 05:07:53 +00:00
Artificial Sweetener 2ae545d64d fix(release): attribute automation to Daisy 2026-09-25 00:58:56 -04:00
142 changed files with 26993 additions and 686 deletions
+4
View File
@@ -122,6 +122,10 @@ jobs:
id: release
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
GIT_AUTHOR_NAME: Daisy
GIT_AUTHOR_EMAIL: daisy@artificialsweetener.ai
GIT_COMMITTER_NAME: Daisy
GIT_COMMITTER_EMAIL: daisy@artificialsweetener.ai
shell: bash
run: |
set -euo pipefail
+32
View File
@@ -1,3 +1,35 @@
# [1.13.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.12.0...v1.13.0) (2026-10-02)
### Bug Fixes
* **negpip:** support Krea attention on ComfyUI 0.28 ([e4eabfe](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/e4eabfefd540f3a6066c29761cefb8f8b3c80f68))
### Features
* **sampling:** add noise inversion and composable sampler options ([be4bd9b](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/be4bd9bb162f6ed4252e0c4d4ef0f3efebe8c674))
* **sampling:** refine sampler options and inversion controls ([ab7cebc](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/ab7cebcebc4d5dfa95aa8d834456af1778fe56a6))
# [1.12.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.11.1...v1.12.0) (2026-09-29)
### Bug Fixes
* **release:** satisfy RES4LYF publication contracts ([04dc35a](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/04dc35af00163689509cace1541a8e7d33b5053e))
### Features
* **sampling:** add RES4LYF sampler methods and schedules ([7b1efef](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/7b1efef47e20b9bf6f032dd57971899d99881822))
## [1.11.1](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.11.0...v1.11.1) (2026-09-25)
### Bug Fixes
* **release:** attribute automation to Daisy ([2ae545d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/2ae545d64d70a454f635ee647fb7e6a1c3500b9b))
# [1.11.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.10.1...v1.11.0) (2026-09-25)
+5 -3
View File
@@ -18,7 +18,7 @@ The pack now covers model loading, regional prompting and segmentation, high-res
- ADetailer-style `[SEP]` prompt batches, masked conditioning, and regional samplers, with optional Prompt Control scheduling and LoRA hooks.
- WD14 and external vision LLM tagging that stays aligned with the right regions.
- Ordered image and mask loading, GPU Lanczos resizing, tiled VAE options, and provenance-aware latent tools.
- WebUI-inspired sampling extras including seed variation, A1111 Euler ancestral behavior, AYS, GITS, `automatic_a1111`, and beta57.
- WebUI-inspired sampling extras including seed variation, A1111 Euler ancestral behavior, AYS, GITS, `automatic_a1111`, and RES4LYF sampler methods and schedules.
## Contents
@@ -152,7 +152,7 @@ The external LLM nodes use a configured OpenAI-compatible provider. **Tag SEGS w
**Simple VAE Encode** can reuse the source latent when the graph proves that its image came directly from an unmodified `VAEDecode`. **Upscale Latent From Image** uses the same provenance to find and resize the original latent. Loading, editing, cropping, detailing, or resizing the image breaks that provenance. These nodes follow the graph instead of trying to identify a latent from the finished tensor.
**KSampler (Extras)** adds the A1111/k-diffusion-style `euler_a_a1111` sampler, AYS SD1 and SDXL schedules, GITS, the `automatic_a1111` scheduler, and a local implementation of the RES4LYF beta57 preset. It keeps Comfy's regular seed handling, partial denoise behavior, progress callbacks, and conditioning inputs.
**KSampler (Extras)** adds the A1111/k-diffusion-style `euler_a_a1111` sampler, AYS SD1 and SDXL schedules, GITS, `automatic_a1111`, and the RES4LYF beta57 and `bong_tangent` schedules. Its sampler dropdown includes 118 RES4LYF methods, including `exponential/ddim`. These methods are also available in the contextual, tiled, and Attention Coupling KSamplers. RES4LYF methods use their upstream default initial noise; Comfy samplers keep Comfy's normal noise path.
**Seed Variation** patches a MODEL so Comfy-native samplers mix their normal initial noise toward a second deterministic seed. Strength `0` keeps the sampler seed unchanged, while strength `1` uses variation-seed initial noise. Ancestral and SDE samplers continue to use the sampler seed for additional noise introduced after initialization.
@@ -182,6 +182,8 @@ SimpleSyrup currently interoperates with:
AGPL-3.0-or-later is a strong copyleft license. If you convey SimpleSyrup or a modified version, you must provide the corresponding source. If users interact with a modified version over a network, you must offer those users the corresponding source for that version.
The vendored RES4LYF license copy includes its upstream commercial-service paragraph before the GNU AGPL v3 text. Read the [RES4LYF license copy](third_party/licenses/res4lyf.LICENSE.txt) and [third-party notices](third_party/NOTICE.md) for the terms and provenance recorded with that code.
SimpleSyrup owes a lot to other projects:
- [ComfyUI](https://github.com/Comfy-Org/ComfyUI) provides the engine and graph ecosystem this pack runs on.
@@ -190,7 +192,7 @@ SimpleSyrup owes a lot to other projects:
- [ComfyUI Prompt Control](https://github.com/asagi4/comfyui-prompt-control) provides the scheduled prompt and LoRA-hook behavior used by the optional integration.
- [ComfyUI Layer Style Advance](https://github.com/chflame163/ComfyUI_LayerStyle_Advance) provides the SAM model bundle SimpleSyrup can adapt.
- [Tiled Diffusion & VAE for AUTOMATIC1111](https://github.com/pkuliyi2015/multidiffusion-upscaler-for-automatic1111) informed the practical tiled diffusion and Mixture of Diffusers behavior reimplemented here.
- [RES4LYF](https://github.com/ClownsharkBatwing/RES4LYF) is the source of the beta57 scheduler preset reimplemented here.
- [RES4LYF](https://github.com/ClownsharkBatwing/RES4LYF) by ClownsharkBatwing and contributors provides the Runge-Kutta and exponential sampler methods included here, along with the `bong_tangent` schedule and the beta57 preset.
- [ComfyUI-ppm](https://github.com/pamparamm/ComfyUI-ppm) by pamparamm provides the ModelPatcher-based NegPiP behavior adapted here and builds on the [ComfyUI port](https://github.com/laksjdjf/cd-tuner_negpip-ComfyUI) by laksjdjf and the [original WebUI implementation](https://github.com/hako-mikan/sd-webui-negpip) by hako-mikan.
SimpleSyrup also vendors or reimplements selected third-party behavior for SAM-HQ, MobileSAM, GroundingDINO, AUTOMATIC1111 sampler behavior, k-diffusion, and tiled diffusion. See [third_party/NOTICE.md](third_party/NOTICE.md) for the complete notices.
+1 -2
View File
@@ -1,11 +1,10 @@
schema_version = 1
review_by = 2027-03-31
fingerprint = "sha256:4b78e0caf8bf4f90d5a7a46ff32b28b01c94c034efc3ddbc24cfa16b8e132ea0"
fingerprint = "sha256:a71fe87163eb7585fafbca08b585954be79abd884e3a1df38bc041c2bc934764"
cohesive_paths = [
"simple_syrup/masking/prompt_segs_with_sam_service.py",
"simple_syrup/nodes/prompt_segs_with_sam.py",
"simple_syrup/nodes_v3/legacy_node_wrappers.py",
"simple_syrup/runtime/attention_region_affinity.py",
"simple_syrup/runtime/attention_region_capture.py",
"simple_syrup/runtime/attention_sampler_lineage.py",
-11
View File
@@ -11,17 +11,6 @@ issue = "chore:SSY-WAIVER-S001"
review_by = 2027-03-31
max_lines = 687
[[waivers]]
id = "SSY-WAIVER-S002"
owner = "sampling scheduler policy"
rule = "STRUCT003"
path = "simple_syrup/runtime/sampling_schedulers.py"
kind = "structural"
justification = "This module owns the complete sigma-schedule policy exposed to every sampler: supported names, fixed published AYS/GITS tables, Comfy delegation, local schedule calculation, denoise truncation, and sampler-specific terminal handling. Roughly half the file is immutable numeric reference data, while the executable functions share one public calculation boundary and dependency direction."
issue = "chore:SSY-WAIVER-S002"
review_by = 2027-03-31
max_lines = 588
[[waivers]]
id = "SSY-WAIVER-S003"
owner = "Anima cross-attention patch contracts"
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "simple-syrup-comfyui",
"version": "1.11.0",
"version": "1.13.0",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "simple-syrup-comfyui",
"version": "1.11.0",
"version": "1.13.0",
"license": "AGPL-3.0-or-later",
"devDependencies": {
"@eslint/js": "^9.39.1",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "simple-syrup-comfyui",
"version": "1.11.0",
"version": "1.13.0",
"private": true,
"license": "AGPL-3.0-or-later",
"type": "module",
+7 -1
View File
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "SimpleSyrup"
description = "Workflow-focused ComfyUI extensions for image generation."
version = "1.11.0"
version = "1.13.0"
license = "AGPL-3.0-or-later"
license-files = ["LICENSE"]
requires-python = ">=3.11"
@@ -36,6 +36,7 @@ line-length = 88
target-version = "py311"
extend-exclude = [
"simple_syrup/third_party/groundingdino_runtime",
"simple_syrup/third_party/res4lyf_runtime",
"simple_syrup/third_party/sam_hq_runtime",
]
@@ -59,6 +60,7 @@ explicit_package_bases = true
mypy_path = ["tests"]
exclude = [
"simple_syrup/third_party/groundingdino_runtime",
"simple_syrup/third_party/res4lyf_runtime",
"simple_syrup/third_party/sam_hq_runtime",
]
@@ -66,6 +68,10 @@ exclude = [
module = ["comfy.*"]
ignore_missing_imports = true
[[tool.mypy.overrides]]
module = ["simple_syrup.third_party.res4lyf_runtime.*"]
follow_imports = "skip"
[tool.pytest.ini_options]
pythonpath = [".", "../.."]
testpaths = ["tests"]
+2
View File
@@ -7,3 +7,5 @@ addict>=2.4.0
yapf>=0.43.0
huggingface-hub>=0.34.0
keyring>=25.0.0
mpmath>=1.3.0
PyWavelets>=1.6.0
+1 -1
View File
@@ -6,6 +6,6 @@
from __future__ import annotations
__version__ = "1.11.0"
__version__ = "1.13.0"
__all__: list[str] = ["__version__"]
+32 -8
View File
@@ -27,16 +27,37 @@ class ContextualDiffusionControls:
global_weight: float
global_steps: int
global_decay: float
latent_tile_width: int | None = None
latent_tile_height: int | None = None
@property
def tile_width(self) -> int:
"""Use explicit local geometry or the convenience node's context size."""
return self.latent_tile_width or self.latent_context_size
@property
def tile_height(self) -> int:
"""Keep the global context independent of a rectangular local tile."""
return self.latent_tile_height or self.latent_context_size
def validate(self) -> None:
"""Reject controls that cannot produce a stable bounded context plan."""
if self.latent_context_size < 16:
raise ValueError("latent_context_size must be at least 16 latent pixels.")
if not 0 <= self.latent_context_overlap < self.latent_context_size:
for value in (self.latent_tile_width, self.latent_tile_height):
if value is not None and (type(value) is not int or value < 16):
raise ValueError(
"Local tile dimensions must be at least 16 latent pixels."
)
if (
not 0
<= self.latent_context_overlap
< min(self.tile_width, self.tile_height)
):
raise ValueError(
"latent_context_overlap must be non-negative and smaller than "
"latent_context_size."
"both local tile dimensions."
)
if self.latent_context_batch_size < 1:
raise ValueError("latent_context_batch_size must be at least 1.")
@@ -65,6 +86,7 @@ def build_contextual_diffusion_plan(
controls: ContextualDiffusionControls,
segs: NativeSegs | None,
region_masks: torch.Tensor | None = None,
segs_canvas: tuple[int, int] | None = None,
) -> ContextualDiffusionPlan:
"""Return a global context plus the regular or SEGS-guided context plan."""
@@ -89,27 +111,29 @@ def build_contextual_diffusion_plan(
segs=segs,
latent_width=latent_width,
latent_height=latent_height,
tile_width=controls.latent_context_size,
tile_height=controls.latent_context_size,
tile_width=controls.tile_width,
tile_height=controls.tile_height,
overlap=controls.latent_context_overlap,
tile_batch_size=controls.latent_context_batch_size,
segs_canvas=segs_canvas,
)
elif segs is not None:
tile_plan = build_segs_guided_tiled_diffusion_plan(
segs=segs,
latent_width=latent_width,
latent_height=latent_height,
tile_width=controls.latent_context_size,
tile_height=controls.latent_context_size,
tile_width=controls.tile_width,
tile_height=controls.tile_height,
overlap=controls.latent_context_overlap,
tile_batch_size=controls.latent_context_batch_size,
segs_canvas=segs_canvas,
)
else:
tile_plan = build_tiled_diffusion_plan(
latent_width=latent_width,
latent_height=latent_height,
tile_width=controls.latent_context_size,
tile_height=controls.latent_context_size,
tile_width=controls.tile_width,
tile_height=controls.tile_height,
overlap=controls.latent_context_overlap,
tile_batch_size=controls.latent_context_batch_size,
)
+89
View File
@@ -0,0 +1,89 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Integrate source-derived inversion states without ComfyUI dependencies."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
import torch
from .noise_inversion import INVERSION_METHODS, InversionMethod
InversionVelocity = Callable[[torch.Tensor, torch.Tensor, int], torch.Tensor]
SpatialResize = Callable[[torch.Tensor, int, int], torch.Tensor]
@dataclass(slots=True)
class InversionSolverEvidence:
"""Count the actual denoiser evaluations performed by an integration stage."""
evaluations: int = 0
def integrate_inversion(
source: torch.Tensor,
sigmas: torch.Tensor,
evaluate: InversionVelocity,
*,
method: InversionMethod,
evidence: InversionSolverEvidence | None = None,
) -> torch.Tensor:
"""Advance a finite state over strictly increasing positive inversion sigmas."""
if method not in INVERSION_METHODS:
raise ValueError("Inversion method must be euler or heun.")
if sigmas.ndim != 1 or len(sigmas) < 2 or not bool(torch.isfinite(sigmas).all()):
raise ValueError("A finite one-dimensional inversion schedule is required.")
if not bool(torch.all(sigmas > 0)) or not bool(torch.all(sigmas[1:] > sigmas[:-1])):
raise ValueError("Inversion sigmas must be positive and increasing.")
if not source.is_floating_point() or not bool(torch.isfinite(source).all()):
raise ValueError("Inversion source must contain finite floating-point values.")
record = evidence if evidence is not None else InversionSolverEvidence()
state = source.clone()
def velocity(x: torch.Tensor, sigma: torch.Tensor, index: int) -> torch.Tensor:
"""Count every denoiser evaluation and reject corrupted predictions."""
record.evaluations += 1
value = evaluate(x, sigma, index)
if value.shape != x.shape or not bool(torch.isfinite(value).all()):
raise FloatingPointError("Invalid inversion velocity shape or values.")
return value
for index in range(len(sigmas) - 1):
current, following = sigmas[index], sigmas[index + 1]
delta = following - current
estimate = velocity(state, current, index)
if method == "heun":
corrected = velocity(state + delta * estimate, following, index)
estimate = (estimate + corrected) / 2
state = state + delta * estimate
if not bool(torch.isfinite(state).all()):
raise FloatingPointError(f"Non-finite inversion state at step {index}.")
return state
def lift_inversion_displacement(
full_source: torch.Tensor,
coarse_source: torch.Tensor,
coarse_endpoint: torch.Tensor,
*,
resize: SpatialResize,
) -> torch.Tensor:
"""Lift only the inferred change so existing full-size detail survives transfer."""
if coarse_source.shape != coarse_endpoint.shape:
raise ValueError("Coarse inversion source and endpoint shapes must match.")
if full_source.shape[:-2] != coarse_source.shape[:-2]:
raise ValueError(
"Inversion transfer must preserve batch and channel dimensions."
)
height, width = full_source.shape[-2:]
lifted = resize(coarse_endpoint - coarse_source, height, width)
if lifted.shape != full_source.shape:
raise ValueError("Inversion displacement resize produced an invalid shape.")
endpoint = full_source + lifted
if not bool(torch.isfinite(endpoint).all()):
raise FloatingPointError("Inversion transfer produced non-finite values.")
return endpoint
+68
View File
@@ -0,0 +1,68 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Define validated source-preserving noise inversion configuration."""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Literal
InversionMethod = Literal["euler", "heun"]
INVERSION_METHODS: tuple[InversionMethod, ...] = ("euler", "heun")
@dataclass(frozen=True, slots=True)
class NoiseInversionOptions:
"""Configure reduced-resolution inversion and an optional full-size finish.
The transition is a fraction of the target inversion sigma, not the forward
denoise steps. A full-size inversion uses ``steps`` and needs no transfer.
"""
method: InversionMethod = "euler"
resolution_scale: float = 0.5
steps: int = 2
switch_fraction: float = 0.75
finishing_steps: int = 1
def __post_init__(self) -> None:
"""Reject invalid or internally incomplete inversion recipes."""
if self.method not in INVERSION_METHODS:
raise ValueError("Inversion method must be euler or heun.")
if (
not math.isfinite(self.resolution_scale)
or not 0 < self.resolution_scale <= 1
):
raise ValueError("Inversion resolution scale must be in (0, 1].")
if type(self.steps) is not int or not 1 <= self.steps <= 64:
raise ValueError("Inversion steps must be an integer between 1 and 64.")
if type(self.finishing_steps) is not int or not 0 <= self.finishing_steps <= 64:
raise ValueError("Inversion finishing steps must be between 0 and 64.")
if not math.isfinite(self.switch_fraction) or not 0 < self.switch_fraction <= 1:
raise ValueError("Inversion transition must be in (0, 1].")
if self.resolution_scale < 1 and self.finishing_steps > 0:
if self.switch_fraction == 1:
raise ValueError("A full-size finish requires a transition below 100%.")
@property
def coarse_target_fraction(self) -> float:
"""Reach the full target unless an enabled full-size stage follows transfer."""
return (
self.switch_fraction
if self.resolution_scale < 1 and self.finishing_steps
else 1.0
)
def coarse_shape(self, height: int, width: int) -> tuple[int, int]:
"""Preserve full dimensions or align reduced transformer grids to even sizes."""
if height < 1 or width < 1:
raise ValueError("Inversion source dimensions must be positive.")
if self.resolution_scale == 1:
return height, width
return (
max(2, round(height * self.resolution_scale / 2) * 2),
max(2, round(width * self.resolution_scale / 2) * 2),
)
@@ -0,0 +1,57 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Project semantic regional-detailing ownership into each inversion resolution."""
from __future__ import annotations
import math
import torch.nn.functional as functional
from .regional_detailing import LatentBox, LatentRegion
def project_inversion_regions(
regions: tuple[LatentRegion, ...],
*,
source_width: int,
source_height: int,
target_width: int,
target_height: int,
) -> tuple[LatentRegion, ...]:
"""Preserve region identity and conditioning while scaling masks and bounds."""
if min(source_width, source_height, target_width, target_height) < 1:
raise ValueError("Regional inversion canvases must have positive dimensions.")
if (source_width, source_height) == (target_width, target_height):
return regions
projected: list[LatentRegion] = []
for region in regions:
box = region.latent_box
if tuple(region.latent_mask.shape) != (source_height, source_width):
raise ValueError("Regional inversion mask must match its canonical canvas.")
if not (
0 <= box.x < box.x + box.width <= source_width
and 0 <= box.y < box.y + box.height <= source_height
):
raise ValueError("Regional inversion bounds must remain inside the canvas.")
left = math.floor(box.x * target_width / source_width)
top = math.floor(box.y * target_height / source_height)
right = math.ceil((box.x + box.width) * target_width / source_width)
bottom = math.ceil((box.y + box.height) * target_height / source_height)
mask = functional.interpolate(
region.latent_mask[None, None].float(),
size=(target_height, target_width),
mode="nearest",
)[0, 0].to(region.latent_mask)
projected.append(
LatentRegion(
region.index,
region.label,
LatentBox(left, top, right - left, bottom - top),
mask,
region.positive,
)
)
return tuple(projected)
@@ -26,6 +26,7 @@ def build_region_constrained_tiled_diffusion_plan(
tile_height: int,
overlap: int,
tile_batch_size: int,
segs_canvas: tuple[int, int] | None = None,
) -> TiledDiffusionPlan:
"""Build tiles split wherever regional composition or optional SEGS change."""
@@ -37,7 +38,8 @@ def build_region_constrained_tiled_diffusion_plan(
ownership_masks = region_ownership
if segs is not None:
native_segs = coerce_segs(segs)
validate_segs_aspect_ratio(native_segs, latent_height, latent_width)
canvas_height, canvas_width = segs_canvas or (latent_height, latent_width)
validate_segs_aspect_ratio(native_segs, canvas_height, canvas_width)
semantic_ownership = segs_ownership_masks(
native_segs,
latent_height=latent_height,
+162
View File
@@ -0,0 +1,162 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Assemble immutable, order-independent sampler capability configuration."""
from __future__ import annotations
import math
from dataclasses import dataclass, replace
from typing import TypeAlias
from .noise_inversion import NoiseInversionOptions
from .regional_prompting import validate_regional_prompt_weight
from .tiled_diffusion import validate_tiled_diffusion_mode
@dataclass(frozen=True, slots=True)
class TilingOptions:
"""Own the single local tile layout and blending configuration."""
diffusion_mode: str = "multidiffusion"
width: int = 128
height: int = 128
overlap: int = 32
batch_size: int = 4
differential_diffusion: bool = False
def __post_init__(self) -> None:
"""Reject tile settings that cannot form a bounded prediction plan."""
validate_tiled_diffusion_mode(self.diffusion_mode)
for name, value in (("width", self.width), ("height", self.height)):
if type(value) is not int or not 16 <= value <= 512:
raise ValueError(
f"Tile {name} must be between 16 and 512 latent pixels."
)
if type(self.overlap) is not int or not 0 <= self.overlap < min(
self.width, self.height
):
raise ValueError(
"Tile overlap must be nonnegative and smaller than both dimensions."
)
if type(self.batch_size) is not int or self.batch_size < 1:
raise ValueError("Tile batch size must be a positive integer.")
if type(self.differential_diffusion) is not bool:
raise TypeError("Differential diffusion must be a boolean.")
@dataclass(frozen=True, slots=True)
class ContextualDiffusionOptions:
"""Own global context and square local sampling geometry independently of tiling."""
context_size: int = 96
global_weight: float = 1.0
global_steps: int = 1
global_decay: float = 0.5
diffusion_mode: str = "multidiffusion"
overlap: int = 32
batch_size: int = 4
differential_diffusion: bool = False
def __post_init__(self) -> None:
"""Reject invalid global schedules and square local sampling settings."""
if type(self.context_size) is not int or not 16 <= self.context_size <= 512:
raise ValueError("Context size must be between 16 and 512 latent pixels.")
if not math.isfinite(self.global_weight) or not 0 <= self.global_weight <= 2:
raise ValueError("Global context weight must be between 0 and 2.")
if type(self.global_steps) is not int or self.global_steps < 0:
raise ValueError("Global context steps must be a nonnegative integer.")
if not math.isfinite(self.global_decay) or not 0 <= self.global_decay <= 1:
raise ValueError("Global context decay must be between 0 and 1.")
self.local_tiling()
def local_tiling(self) -> TilingOptions:
"""Use context size for both dimensions of the sole local sampling plan."""
return TilingOptions(
diffusion_mode=self.diffusion_mode,
width=self.context_size,
height=self.context_size,
overlap=self.overlap,
batch_size=self.batch_size,
differential_diffusion=self.differential_diffusion,
)
@dataclass(frozen=True, slots=True)
class AttentionCouplingOptions:
"""Configure regional attention strength without binding masks or a model."""
regional_prompt_weight: float = 1.0
region_mask_feather: int = 0
def __post_init__(self) -> None:
"""Require valid regional attention strengths and feathering controls."""
validate_regional_prompt_weight(self.regional_prompt_weight)
if type(self.region_mask_feather) is not int or self.region_mask_feather < 0:
raise ValueError("Region mask feather must be a nonnegative integer.")
SamplerCapability: TypeAlias = (
TilingOptions
| ContextualDiffusionOptions
| NoiseInversionOptions
| AttentionCouplingOptions
)
@dataclass(frozen=True, slots=True)
class SamplerOptions:
"""Own one immutable setting per capability, independent of graph order."""
tiling: TilingOptions | None = None
contextual_diffusion: ContextualDiffusionOptions | None = None
noise_inversion: NoiseInversionOptions | None = None
attention_coupling: AttentionCouplingOptions | None = None
def __post_init__(self) -> None:
"""Reject malformed connection payloads at the typed configuration boundary."""
for name, expected in (
("tiling", TilingOptions),
("contextual_diffusion", ContextualDiffusionOptions),
("noise_inversion", NoiseInversionOptions),
("attention_coupling", AttentionCouplingOptions),
):
value = getattr(self, name)
if value is not None and not isinstance(value, expected):
raise TypeError(f"Sampler option {name} must be {expected.__name__}.")
def with_capability(self, capability: SamplerCapability) -> SamplerOptions:
"""Return a fresh configuration or reject an ambiguous duplicate feature."""
names: dict[type[object], str] = {
TilingOptions: "tiling",
ContextualDiffusionOptions: "contextual_diffusion",
NoiseInversionOptions: "noise_inversion",
AttentionCouplingOptions: "attention_coupling",
}
name = names.get(type(capability))
if name is None:
raise TypeError("Unsupported sampler capability configuration.")
if getattr(self, name) is not None:
raise ValueError(
f"Duplicate sampler capability: {name}. Bypass or remove one node."
)
if isinstance(capability, TilingOptions):
return replace(self, tiling=capability)
if isinstance(capability, ContextualDiffusionOptions):
return replace(self, contextual_diffusion=capability)
if isinstance(capability, NoiseInversionOptions):
return replace(self, noise_inversion=capability)
return replace(self, attention_coupling=capability)
def append_sampler_capability(
options: SamplerOptions | None, capability: SamplerCapability | None
) -> SamplerOptions:
"""Append a capability or pass through a disabled contribution after validation."""
if options is not None and not isinstance(options, SamplerOptions):
raise TypeError(
"Options input must be a SimpleSyrup sampler options connection."
)
current = options if options is not None else SamplerOptions()
return current if capability is None else current.with_capability(capability)
+5 -1
View File
@@ -22,16 +22,20 @@ def build_segs_guided_tiled_diffusion_plan(
tile_height: int,
overlap: int,
tile_batch_size: int,
segs_canvas: tuple[int, int] | None = None,
) -> TiledDiffusionPlan:
"""Build bounded sampling windows whose irregular cores follow supplied SEGS.
Every latent pixel receives exactly one ownership core. Each core is sampled
through a rectangular window, while its local blend mask retains the irregular
boundary and shares a feathered overlap with neighboring cores.
A reduced inversion stage validates proportions against its original canvas
because rounding the reduced dimensions can change their aspect ratio.
"""
native_segs = coerce_segs(segs)
validate_segs_aspect_ratio(native_segs, latent_height, latent_width)
canvas_height, canvas_width = segs_canvas or (latent_height, latent_width)
validate_segs_aspect_ratio(native_segs, canvas_height, canvas_width)
ownership_masks = segs_ownership_masks(
native_segs,
latent_height=latent_height,
@@ -10,6 +10,7 @@ from typing import Any, ClassVar
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.segs import coerce_segs_group
from ..nodes import tooltips
from ..nodes.detailer_input_adapters import (
@@ -215,6 +216,7 @@ class DetailSEGSAsRegions:
noise_mask_feather: object = 20,
tiled_encode: object = False,
tiled_decode: object = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> tuple[object]:
"""Run regional detailing and return the detailed image."""
@@ -237,6 +239,7 @@ class DetailSEGSAsRegions:
strict=True,
):
result = service.detail(
noise_inversion=noise_inversion,
image=single_image,
segs=single_segs,
model=single_input(model, "model", list_mode, OPERATION),
@@ -10,6 +10,7 @@ from typing import Any, ClassVar
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.segs import coerce_segs_group
from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES
from ..nodes import tooltips
@@ -263,6 +264,7 @@ class DetailSEGSByScaleFactorTiledDiffusion:
latent_tile_height: object = 128,
latent_tile_overlap: object = 16,
latent_tile_batch_size: object = 4,
noise_inversion: NoiseInversionOptions | None = None,
) -> tuple[object]:
"""Run tiled diffusion scale-factor detailing and return the image."""
@@ -275,6 +277,7 @@ class DetailSEGSByScaleFactorTiledDiffusion:
outputs: list[torch.Tensor] = []
for single_image, single_segs in zip(images, segs_group, strict=True):
result = service.detail(
noise_inversion=noise_inversion,
image=single_image,
segs=single_segs,
model=single_input(model, "model", list_mode, OPERATION),
+10
View File
@@ -14,13 +14,16 @@ def get_nodes() -> list[type[object]]:
from .all_prompt_attention_segs import AllPromptAttentionSEGSV3
from .attention_capture_model import AttentionCaptureModelV3
from .attention_coupling_options import AttentionCouplingOptionsV3
from .attention_masked_conditioning import AttentionMaskedConditioningV3
from .attention_region_mask import AttentionRegionMaskV3
from .batch_region_conditioning import BatchRegionConditioningV3
from .batch_segs import BatchSEGSV3
from .compose_regional_conditioning import ComposeRegionalConditioningV3
from .concept_attention_segs import ConceptAttentionSEGSV3
from .contextual_diffusion_options import ContextualDiffusionOptionsV3
from .external_llm_prompt import ExternalLLMPromptV3
from .ksampler import KSamplerV3
from .ksampler_attention_coupling import KSamplerAttentionCouplingV3
from .ksampler_contextual_attention_coupling import (
KSamplerContextualAttentionCouplingV3,
@@ -62,6 +65,7 @@ def get_nodes() -> list[type[object]]:
from .load_image_list import LoadImageListV3
from .load_mask_batch import LoadMaskBatchV3
from .mask_to_segs import MaskToSEGSV3
from .noise_inversion_options import NoiseInversionOptionsV3
from .scale_factor import ScaleFactorV3
from .seed_variation import SeedVariationV3
from .simple_load_checkpoint import SimpleLoadCheckpointV3
@@ -71,12 +75,18 @@ def get_nodes() -> list[type[object]]:
from .tag_segs_with_external_llm import TagSEGSWithExternalLLMV3
from .tag_segs_with_wd14 import TagSEGSWithWD14V3
from .tile_and_tag_segs import TileAndTagSEGSV3
from .tiling_options import TilingOptionsV3
from .vae_decode_options import VAEDecodeOptionsV3
from .vae_encode_options import VAEEncodeOptionsV3
from .wd14_tagger_loader import WD14TaggerLoaderV3
nodes: list[type[object]] = [
AllPromptAttentionSEGSV3,
AttentionCouplingOptionsV3,
ContextualDiffusionOptionsV3,
NoiseInversionOptionsV3,
TilingOptionsV3,
KSamplerV3,
AttentionCaptureModelV3,
AttentionMaskedConditioningV3,
AttentionRegionMaskV3,
@@ -0,0 +1,67 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Configure regional attention without applying MODEL patches in the options graph."""
from __future__ import annotations
from typing import Any
from ..domain.sampler_options import (
AttentionCouplingOptions,
SamplerOptions,
append_sampler_capability,
)
from .ksampler_schema import attention_coupling_ksampler_inputs
from .sampler_options_schema import (
COMFY_IO,
OptionsNodeBase,
options_input,
options_output,
)
class AttentionCouplingOptionsV3(OptionsNodeBase):
"""Configure the sampler's regional attention strength and mask feathering."""
@classmethod
def define_schema(cls) -> Any:
"""Expose strength and feathering without binding region payloads."""
controls = attention_coupling_ksampler_inputs(COMFY_IO)
return COMFY_IO.Schema(
node_id="SimpleSyrup.AttentionCouplingOptions",
display_name="Attention Coupling Options",
category="SimpleSyrup/Sampling/Options",
description=(
"Routes global-first sampler conditioning to ordered image "
"regions through regional attention and LoRA hooks."
),
inputs=[
options_input(COMFY_IO),
*[
control
for control in controls
if control.id in {"regional_prompt_weight", "region_mask_feather"}
],
],
outputs=[options_output(COMFY_IO)],
)
@classmethod
def execute(
cls,
regional_prompt_weight: float = 1.0,
region_mask_feather: int = 0,
options: SamplerOptions | None = None,
) -> tuple[SamplerOptions]:
"""Append regional attention controls while deferring model preparation."""
return (
append_sampler_capability(
options,
AttentionCouplingOptions(
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
),
),
)
@@ -0,0 +1,91 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Configure complete contextual sampling with one square local context plan."""
from __future__ import annotations
from typing import Any
from ..domain.sampler_options import (
ContextualDiffusionOptions,
SamplerOptions,
append_sampler_capability,
)
from .ksampler_schema import contextual_diffusion_inputs
from .sampler_options_schema import (
COMFY_IO,
OptionsNodeBase,
options_input,
options_output,
)
class ContextualDiffusionOptionsV3(OptionsNodeBase):
"""Schedule global scene authority over one local tile prediction."""
@classmethod
def define_schema(cls) -> Any:
"""Append local controls after existing widgets to preserve saved values."""
controls = {
control.id: control for control in contextual_diffusion_inputs(COMFY_IO)
}
return COMFY_IO.Schema(
node_id="SimpleSyrup.ContextualDiffusionOptions",
display_name="Contextual Diffusion Options",
category="SimpleSyrup/Sampling/Options",
description=(
"Samples local contexts with global scene guidance; "
"takes precedence over connected Tiling Options."
),
inputs=[
options_input(COMFY_IO),
controls["latent_context_size"],
controls["global_weight"],
controls["global_steps"],
controls["global_decay"],
controls["diffusion_mode"],
controls["latent_context_overlap"],
controls["latent_context_batch_size"],
COMFY_IO.Boolean.Input(
"differential_diffusion",
default=False,
tooltip=(
"Uses the noise mask to vary denoising strength spatially; "
"preserves existing model mask behavior."
),
),
],
outputs=[options_output(COMFY_IO)],
)
@classmethod
def execute(
cls,
latent_context_size: int = 96,
global_weight: float = 1.0,
global_steps: int = 1,
global_decay: float = 0.5,
options: SamplerOptions | None = None,
diffusion_mode: str = "multidiffusion",
latent_context_overlap: int = 32,
latent_context_batch_size: int = 4,
differential_diffusion: bool = False,
) -> tuple[SamplerOptions]:
"""Append complete context settings without inheriting a Tiling contribution."""
return (
append_sampler_capability(
options,
ContextualDiffusionOptions(
context_size=latent_context_size,
global_weight=global_weight,
global_steps=global_steps,
global_decay=global_decay,
diffusion_mode=diffusion_mode,
overlap=latent_context_overlap,
batch_size=latent_context_batch_size,
differential_diffusion=differential_diffusion,
),
),
)
+100
View File
@@ -0,0 +1,100 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Expose one native KSampler consuming a composable capability configuration."""
from __future__ import annotations
from typing import Any, ClassVar
from ..domain.sampler_options import SamplerOptions
from ..nodes import tooltips
from ..services.sampler_options_sampling_service import SamplerOptionsSamplingService
from .ksampler_schema import ksampler_inputs
from .sampler_options_schema import COMFY_IO, OptionsNodeBase, options_input
class KSamplerV3(OptionsNodeBase):
"""Execute tiling, context, inversion and attention through shared authorities."""
service_class: ClassVar[type[SamplerOptionsSamplingService]] = (
SamplerOptionsSamplingService
)
@classmethod
def define_schema(cls) -> Any:
"""Declare sampling controls, capabilities and optional spatial regions."""
return COMFY_IO.Schema(
node_id="SimpleSyrup.KSampler",
display_name="KSampler (SimpleSyrup)",
category="SimpleSyrup/Sampling",
description=(
"Samples latents with connected sampler options for tiling, "
"Contextual Diffusion, noise inversion and Attention Coupling."
),
inputs=[
*ksampler_inputs(COMFY_IO, steps_default=20, cfg_default=8.0),
options_input(COMFY_IO),
COMFY_IO.SEGS.Input(
"segs",
optional=True,
tooltip=(
"Guides local sampling regions when Tiling or Contextual "
"Diffusion options are connected; ignored otherwise."
),
),
COMFY_IO.Mask.Input(
"region_masks",
optional=True,
tooltip=(
"Ordered masks paired with global-first conditioning batches "
"when Attention Coupling options are connected; "
"ignored otherwise."
),
),
],
outputs=[
COMFY_IO.Latent.Output(
"latent", tooltip=tooltips.DENOISED_LATENT_OUTPUT
)
],
)
@classmethod
def execute(
cls,
model: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
positive: object,
negative: object | None = None,
latent_image: dict[str, Any] | None = None,
denoise: float = 1.0,
options: SamplerOptions | None = None,
segs: object | None = None,
region_masks: object | None = None,
) -> tuple[dict[str, Any]]:
"""Delegate sampling without mutating capability configuration."""
if latent_image is None:
raise TypeError("KSampler requires latent_image.")
return (
cls.service_class().sample(
model=model,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
positive=positive,
negative=negative,
latent_image=latent_image,
denoise=denoise,
options=options,
segs=segs,
region_masks=region_masks,
),
)
@@ -17,6 +17,7 @@ from .ksampler_schema import (
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT,
attention_coupling_ksampler_inputs,
)
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -68,10 +69,13 @@ class KSamplerAttentionCouplingV3(_ComfyNodeBase):
"anima regional prompt",
"sdxl regional prompt",
],
inputs=attention_coupling_ksampler_inputs(
_comfy_io,
region_masks_optional=True,
),
inputs=[
*attention_coupling_ksampler_inputs(
_comfy_io,
region_masks_optional=True,
),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
"latent",
@@ -98,12 +102,24 @@ class KSamplerAttentionCouplingV3(_ComfyNodeBase):
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT
),
region_mask_feather: int = 0,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any]]:
"""Delegate ordinary or regional sampling to the routing service."""
if latent_image is None:
raise TypeError("KSampler Attention Coupling requires latent_image.")
output = cls.sampling_service_class().sample(
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
model=model,
seed=seed,
steps=steps,
@@ -17,6 +17,7 @@ from .ksampler_schema import (
attention_coupling_ksampler_inputs,
contextual_diffusion_inputs,
)
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -75,6 +76,7 @@ class KSamplerContextualAttentionCouplingV3(_ComfyNodeBase):
optional=True,
tooltip=tooltips.CONTEXTUAL_DIFFUSION_SEGS,
),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
@@ -112,6 +114,11 @@ class KSamplerContextualAttentionCouplingV3(_ComfyNodeBase):
global_steps: int = 1,
global_decay: float = 0.5,
segs: object | None = None,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any], object]:
"""Delegate the complete request to the combined application service."""
@@ -145,5 +152,12 @@ class KSamplerContextualAttentionCouplingV3(_ComfyNodeBase):
global_steps=global_steps,
global_decay=global_decay,
segs=segs,
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
)
return result.latent, result.contexts
@@ -18,6 +18,7 @@ from .ksampler_schema import (
ksampler_inputs,
optional_regional_sampling_inputs,
)
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -67,6 +68,7 @@ class KSamplerContextualDiffusionV3(_ComfyNodeBase):
_comfy_io,
segs_tooltip=tooltips.CONTEXTUAL_DIFFUSION_SEGS,
),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
@@ -104,6 +106,11 @@ class KSamplerContextualDiffusionV3(_ComfyNodeBase):
region_masks: object | None = None,
regional_prompt_weight: float = 0.5,
region_mask_feather: int = 0,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any], object]:
"""Delegate Contextual Diffusion sampling to its application service."""
@@ -131,5 +138,12 @@ class KSamplerContextualDiffusionV3(_ComfyNodeBase):
region_masks=region_masks,
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
)
return result.latent, result.contexts
@@ -13,6 +13,7 @@ from ..nodes import tooltips
from ..services.ksampler_sampling_service import KSamplerSamplingService
from ..services.regional_conditioning_service import RegionalConditioningService
from .ksampler_schema import regional_ksampler_inputs
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -51,7 +52,10 @@ class KSamplerPromptByRegionV3(_ComfyNodeBase):
"mask-bound regional prompts."
),
search_aliases=["ksampler", "regional prompt", "masked prompt"],
inputs=regional_ksampler_inputs(_comfy_io),
inputs=[
*regional_ksampler_inputs(_comfy_io),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
"latent",
@@ -76,6 +80,11 @@ class KSamplerPromptByRegionV3(_ComfyNodeBase):
region_mask_feather: int = 0,
latent_image: dict[str, Any] | None = None,
denoise: float = 1.0,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any]]:
"""Assemble regional conditioning and sample the full latent."""
@@ -93,6 +102,13 @@ class KSamplerPromptByRegionV3(_ComfyNodeBase):
)
)
output = cls.sampling_service_class().sample(
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
model=model,
seed=seed,
steps=steps,
@@ -14,6 +14,7 @@ from ..nodes import tooltips
from ..services.regional_conditioning_service import RegionalConditioningService
from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService
from .ksampler_schema import regional_ksampler_inputs, tiled_diffusion_inputs
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -60,6 +61,7 @@ class KSamplerPromptByTiledRegionV3(_ComfyNodeBase):
inputs=[
*regional_ksampler_inputs(_comfy_io),
*tiled_diffusion_inputs(_comfy_io),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
@@ -90,6 +92,11 @@ class KSamplerPromptByTiledRegionV3(_ComfyNodeBase):
latent_tile_height: int = 128,
latent_tile_overlap: int = 16,
latent_tile_batch_size: int = 4,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any]]:
"""Assemble regional conditioning and sample overlapping latent tiles."""
@@ -107,6 +114,13 @@ class KSamplerPromptByTiledRegionV3(_ComfyNodeBase):
)
)
output = cls.sampling_service_class().sample(
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
diffusion_mode=diffusion_mode,
model=model,
seed=seed,
@@ -18,6 +18,7 @@ from .ksampler_schema import (
attention_coupling_ksampler_inputs,
tiled_diffusion_inputs,
)
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -78,6 +79,7 @@ class KSamplerTiledAttentionCouplingV3(_ComfyNodeBase):
region_masks_optional=True,
),
*tiled_diffusion_inputs(_comfy_io),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
@@ -105,6 +107,11 @@ class KSamplerTiledAttentionCouplingV3(_ComfyNodeBase):
latent_tile_height: int = 128,
latent_tile_overlap: int = 16,
latent_tile_batch_size: int = 4,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
region_masks: object | None = None,
regional_prompt_weight: float = (
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT
@@ -116,6 +123,13 @@ class KSamplerTiledAttentionCouplingV3(_ComfyNodeBase):
if latent_image is None:
raise TypeError("KSampler Tiled Attention Coupling requires latent_image.")
output = cls.sampling_service_class().sample(
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
diffusion_mode=diffusion_mode,
model=model,
seed=seed,
@@ -16,6 +16,7 @@ from .ksampler_schema import (
optional_regional_sampling_inputs,
tiled_diffusion_inputs,
)
from .sampler_options_schema import inversion_from_controls, noise_inversion_inputs
if TYPE_CHECKING:
@@ -65,6 +66,7 @@ class KSamplerTiledDiffusionV3(_ComfyNodeBase):
"boundaries while preserving the configured overlap."
),
),
*noise_inversion_inputs(_comfy_io, convenience=True),
],
outputs=[
_comfy_io.Latent.Output(
@@ -96,6 +98,11 @@ class KSamplerTiledDiffusionV3(_ComfyNodeBase):
region_masks: object | None = None,
regional_prompt_weight: float = 0.5,
region_mask_feather: int = 0,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> tuple[dict[str, Any]]:
"""Delegate tiled diffusion sampling to its application service."""
@@ -122,5 +129,12 @@ class KSamplerTiledDiffusionV3(_ComfyNodeBase):
region_masks=region_masks,
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
noise_inversion=inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
),
)
return (output,)
@@ -0,0 +1,73 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Expose shared inversion controls on maintained implementation-backed samplers."""
from __future__ import annotations
from typing import Any
from ..nodes.detailer_input_adapters import (
float_input,
int_input,
str_input,
)
from .legacy_node_adapter import LegacyNodeV3Adapter
from .sampler_options_schema import (
COMFY_IO,
inversion_from_controls,
noise_inversion_inputs,
)
class LegacyInversionNodeV3Adapter(LegacyNodeV3Adapter):
"""Normalize direct or list-mode widgets into the shared inversion domain value."""
@classmethod
def define_schema(cls) -> Any:
"""Append the shared five inversion controls after the sampler inputs."""
schema = super().define_schema()
schema.inputs.extend(noise_inversion_inputs(COMFY_IO, convenience=True))
return schema
@classmethod
def execute(cls, **kwargs: object) -> Any:
"""Narrow inversion widgets before delegating normal implementation inputs."""
values = dict(kwargs)
list_mode = bool(getattr(cls.LEGACY_NODE_CLASS, "INPUT_IS_LIST", False))
operation = cls.DISPLAY_NAME
inversion = inversion_from_controls(
inversion_method=str_input(
values.pop("inversion_method", "euler"),
"inversion_method",
list_mode,
operation,
),
inversion_resolution_scale=float_input(
values.pop("inversion_resolution_scale", 0.5),
"inversion_resolution_scale",
list_mode,
operation,
),
inversion_steps=int_input(
values.pop("inversion_steps", 2),
"inversion_steps",
list_mode,
operation,
),
inversion_switch_fraction=float_input(
values.pop("inversion_switch_fraction", 0.75),
"inversion_switch_fraction",
list_mode,
operation,
),
inversion_finishing_steps=int_input(
values.pop("inversion_finishing_steps", 1),
"inversion_finishing_steps",
list_mode,
operation,
),
)
values["noise_inversion"] = inversion
return super().execute(**values)
@@ -0,0 +1,275 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Translate maintained implementation contracts into native Comfy v3 schemas."""
from __future__ import annotations
from collections.abc import Mapping
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for Comfy v3 nodes."""
hidden: ClassVar[Any]
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
_HIDDEN_INPUTS = {
"PROMPT": "prompt",
"DYNPROMPT": "dynprompt",
"EXTRA_PNGINFO": "extra_pnginfo",
"UNIQUE_ID": "unique_id",
}
class LegacyNodeV3Adapter(_ComfyNodeBase):
"""Build a v3 schema and execution bridge for a legacy implementation class."""
LEGACY_NODE_CLASS: ClassVar[type[Any]]
NODE_ID: ClassVar[str]
DISPLAY_NAME: ClassVar[str]
ENABLE_EXPAND: ClassVar[bool] = False
WORKFLOW_INPUT_ORDER: ClassVar[tuple[str, ...] | None] = None
@classmethod
def define_schema(cls) -> Any:
"""Declare a v3 schema from the implementation class contract."""
legacy = cls.LEGACY_NODE_CLASS
return _comfy_io.Schema(
node_id=cls.NODE_ID,
display_name=cls.DISPLAY_NAME,
category=str(getattr(legacy, "CATEGORY", "SimpleSyrup")),
description=str(getattr(legacy, "DESCRIPTION", "")),
search_aliases=list(getattr(legacy, "SEARCH_ALIASES", [])),
inputs=_v3_inputs(
legacy.INPUT_TYPES(),
workflow_order=cls.WORKFLOW_INPUT_ORDER,
),
outputs=_v3_outputs(legacy),
hidden=_v3_hidden_inputs(legacy.INPUT_TYPES()),
is_input_list=bool(getattr(legacy, "INPUT_IS_LIST", False)),
is_output_node=bool(getattr(legacy, "OUTPUT_NODE", False)),
enable_expand=cls.ENABLE_EXPAND,
)
@classmethod
def execute(cls, **kwargs: object) -> Any:
"""Run the wrapped implementation with v3-provided inputs."""
values = dict(kwargs)
for name, hidden_attr in _legacy_hidden_inputs(
cls.LEGACY_NODE_CLASS.INPUT_TYPES()
).items():
if name not in values:
values[name] = getattr(cls.hidden, hidden_attr)
function_name = str(cls.LEGACY_NODE_CLASS.FUNCTION)
implementation = cls.LEGACY_NODE_CLASS()
function = getattr(implementation, function_name)
return function(**values)
def _v3_inputs(
input_types: Mapping[str, Mapping[str, object]],
*,
workflow_order: tuple[str, ...] | None = None,
) -> list[Any]:
"""Return v3 inputs while preserving any explicit persisted socket order."""
declarations: dict[str, tuple[object, bool]] = {}
for section_name, optional in (("required", False), ("optional", True)):
section = input_types.get(section_name, {})
for name, declaration in section.items():
if name in declarations:
raise ValueError(f"legacy input {name} is declared more than once.")
declarations[name] = (declaration, optional)
order = tuple(declarations) if workflow_order is None else workflow_order
if len(order) != len(set(order)) or set(order) != set(declarations):
raise ValueError("legacy workflow input order must name every input once.")
return [
_v3_input(name, declarations[name][0], optional=declarations[name][1])
for name in order
]
def _v3_input(name: str, declaration: object, *, optional: bool) -> Any:
"""Return one v3 input declaration from a legacy field declaration."""
if not isinstance(declaration, tuple) or not declaration:
raise TypeError(f"legacy input {name} declaration must be a tuple.")
io_declaration = declaration[0]
options = _input_options(declaration)
tooltip = _string_option(options, "tooltip")
advanced = _bool_option(options, "advanced")
raw_link = _bool_option(options, "rawLink") or _bool_option(options, "raw_link")
force_input = _bool_option(options, "forceInput") or _bool_option(
options, "force_input"
)
if isinstance(io_declaration, (list, tuple)):
return _comfy_io.Combo.Input(
name,
options=list(io_declaration),
optional=optional,
default=options.get("default"),
control_after_generate=options.get("control_after_generate"),
tooltip=tooltip,
raw_link=raw_link,
advanced=advanced,
)
if not isinstance(io_declaration, str):
raise TypeError(f"legacy input {name} type must be a string or options list.")
input_type = io_declaration
input_class = _io_class(input_type)
common_options = {
"optional": optional,
"tooltip": tooltip,
"raw_link": raw_link,
"advanced": advanced,
}
if input_type == "INT":
return input_class.Input(
name,
default=options.get("default"),
min=options.get("min"),
max=options.get("max"),
step=options.get("step"),
control_after_generate=options.get("control_after_generate"),
**common_options,
)
if input_type == "FLOAT":
return input_class.Input(
name,
default=options.get("default"),
min=options.get("min"),
max=options.get("max"),
step=options.get("step"),
round=options.get("round"),
**common_options,
)
if input_type == "STRING":
return input_class.Input(
name,
default=options.get("default"),
multiline=bool(options.get("multiline", False)),
force_input=force_input,
**common_options,
)
if input_type == "BOOLEAN":
return input_class.Input(
name,
default=options.get("default"),
label_on=options.get("label_on"),
label_off=options.get("label_off"),
**common_options,
)
return input_class.Input(name, **common_options)
def _v3_outputs(legacy: type[Any]) -> list[Any]:
"""Return v3 output declarations from legacy return metadata."""
return_types = tuple(getattr(legacy, "RETURN_TYPES", ()))
return_names = getattr(legacy, "RETURN_NAMES", None)
output_tooltips = tuple(getattr(legacy, "OUTPUT_TOOLTIPS", ()))
output_is_list = tuple(
getattr(legacy, "OUTPUT_IS_LIST", (False,) * len(return_types))
)
outputs: list[Any] = []
for index, io_type in enumerate(return_types):
output_name = None
if isinstance(return_names, tuple) and index < len(return_names):
output_name = str(return_names[index])
tooltip = None
if index < len(output_tooltips):
tooltip = str(output_tooltips[index])
is_output_list = index < len(output_is_list) and bool(output_is_list[index])
outputs.append(
_io_class(str(io_type)).Output(
output_name,
tooltip=tooltip,
is_output_list=is_output_list,
)
)
return outputs
def _v3_hidden_inputs(input_types: Mapping[str, Mapping[str, object]]) -> list[Any]:
"""Return v3 hidden declarations requested by legacy hidden inputs."""
hidden_values = set(_legacy_hidden_inputs(input_types).values())
return [getattr(_comfy_io.Hidden, value) for value in sorted(hidden_values)]
def _legacy_hidden_inputs(
input_types: Mapping[str, Mapping[str, object]],
) -> dict[str, str]:
"""Return legacy hidden input names mapped to v3 hidden holder attributes."""
hidden_inputs: dict[str, str] = {}
for name, sentinel in input_types.get("hidden", {}).items():
if isinstance(sentinel, str) and sentinel in _HIDDEN_INPUTS:
hidden_inputs[name] = _HIDDEN_INPUTS[sentinel]
return hidden_inputs
def _io_class(io_type: str) -> Any:
"""Return the v3 IO class for a legacy Comfy type string."""
known_types = {
"BOOLEAN": _comfy_io.Boolean,
"INT": _comfy_io.Int,
"FLOAT": _comfy_io.Float,
"STRING": _comfy_io.String,
"IMAGE": _comfy_io.Image,
"MASK": _comfy_io.Mask,
"LATENT": _comfy_io.Latent,
"MODEL": _comfy_io.Model,
"CLIP": _comfy_io.Clip,
"VAE": _comfy_io.Vae,
"CONDITIONING": _comfy_io.Conditioning,
"SEGS": _comfy_io.SEGS,
}
return known_types.get(io_type, _comfy_io.Custom(io_type))
def _input_options(declaration: tuple[object, ...]) -> dict[str, object]:
"""Return an input options dictionary from a legacy declaration."""
if len(declaration) < 2 or not isinstance(declaration[1], dict):
return {}
return dict(declaration[1])
def _string_option(options: Mapping[str, object], name: str) -> str | None:
"""Return a string option when present."""
value = options.get(name)
if isinstance(value, str):
return value
return None
def _bool_option(options: Mapping[str, object], name: str) -> bool | None:
"""Return a boolean option when present."""
value = options.get(name)
if isinstance(value, bool):
return value
return None
+4 -270
View File
@@ -6,10 +6,6 @@
from __future__ import annotations
from collections.abc import Mapping
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from ..nodes.conditioning_batch_pack import (
ConditioningBatchAppend,
ConditioningBatchStart,
@@ -40,6 +36,8 @@ from ..nodes.segs_from_sam_output import SEGSFromSAMOutput
from ..nodes.simple_load_anima import SimpleLoadAnima
from ..nodes.simple_preview_segs import SimplePreviewSEGS
from ..nodes.vitmatte_model_loader import ViTMatteModelLoader
from .legacy_inversion_node_adapter import LegacyInversionNodeV3Adapter
from .legacy_node_adapter import LegacyNodeV3Adapter
from .legacy_workflow_input_order import (
DETAIL_SEGS_AS_REGIONS_INPUT_ORDER,
DETAIL_SEGS_BY_SCALE_FACTOR_INPUT_ORDER,
@@ -47,75 +45,6 @@ from .legacy_workflow_input_order import (
KSAMPLER_EXTRAS_INPUT_ORDER,
)
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for Comfy v3 nodes."""
hidden: ClassVar[Any]
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
_HIDDEN_INPUTS = {
"PROMPT": "prompt",
"DYNPROMPT": "dynprompt",
"EXTRA_PNGINFO": "extra_pnginfo",
"UNIQUE_ID": "unique_id",
}
class LegacyNodeV3Adapter(_ComfyNodeBase):
"""Build a v3 schema and execution bridge for a legacy implementation class."""
LEGACY_NODE_CLASS: ClassVar[type[Any]]
NODE_ID: ClassVar[str]
DISPLAY_NAME: ClassVar[str]
ENABLE_EXPAND: ClassVar[bool] = False
WORKFLOW_INPUT_ORDER: ClassVar[tuple[str, ...] | None] = None
@classmethod
def define_schema(cls) -> Any:
"""Declare a v3 schema from the implementation class contract."""
legacy = cls.LEGACY_NODE_CLASS
return _comfy_io.Schema(
node_id=cls.NODE_ID,
display_name=cls.DISPLAY_NAME,
category=str(getattr(legacy, "CATEGORY", "SimpleSyrup")),
description=str(getattr(legacy, "DESCRIPTION", "")),
search_aliases=list(getattr(legacy, "SEARCH_ALIASES", [])),
inputs=_v3_inputs(
legacy.INPUT_TYPES(),
workflow_order=cls.WORKFLOW_INPUT_ORDER,
),
outputs=_v3_outputs(legacy),
hidden=_v3_hidden_inputs(legacy.INPUT_TYPES()),
is_input_list=bool(getattr(legacy, "INPUT_IS_LIST", False)),
is_output_node=bool(getattr(legacy, "OUTPUT_NODE", False)),
enable_expand=cls.ENABLE_EXPAND,
)
@classmethod
def execute(cls, **kwargs: object) -> Any:
"""Run the wrapped implementation with v3-provided inputs."""
values = dict(kwargs)
for name, hidden_attr in _legacy_hidden_inputs(
cls.LEGACY_NODE_CLASS.INPUT_TYPES()
).items():
if name not in values:
values[name] = getattr(cls.hidden, hidden_attr)
function_name = str(cls.LEGACY_NODE_CLASS.FUNCTION)
implementation = cls.LEGACY_NODE_CLASS()
function = getattr(implementation, function_name)
return function(**values)
class ConditioningBatchStartV3(LegacyNodeV3Adapter):
"""Expose Conditioning Batch Start through Comfy v3 only."""
@@ -224,7 +153,7 @@ class ResizeImageToTargetV3(LegacyNodeV3Adapter):
DISPLAY_NAME = "Resize Image to Target"
class DetailSEGSAsRegionsV3(LegacyNodeV3Adapter):
class DetailSEGSAsRegionsV3(LegacyInversionNodeV3Adapter):
"""Expose Detail SEGS as Regions through Comfy v3 only."""
LEGACY_NODE_CLASS = DetailSEGSAsRegions
@@ -242,7 +171,7 @@ class DetailSEGSByScaleFactorV3(LegacyNodeV3Adapter):
WORKFLOW_INPUT_ORDER = DETAIL_SEGS_BY_SCALE_FACTOR_INPUT_ORDER
class DetailSEGSByScaleFactorTiledDiffusionV3(LegacyNodeV3Adapter):
class DetailSEGSByScaleFactorTiledDiffusionV3(LegacyInversionNodeV3Adapter):
"""Expose Detail SEGS by Scale Factor with Tiled Diffusion through Comfy v3."""
LEGACY_NODE_CLASS = DetailSEGSByScaleFactorTiledDiffusion
@@ -323,201 +252,6 @@ class ViTMatteModelLoaderV3(LegacyNodeV3Adapter):
DISPLAY_NAME = "ViTMatte Model Loader"
def _v3_inputs(
input_types: Mapping[str, Mapping[str, object]],
*,
workflow_order: tuple[str, ...] | None = None,
) -> list[Any]:
"""Return v3 inputs while preserving any explicit persisted socket order."""
declarations: dict[str, tuple[object, bool]] = {}
for section_name, optional in (("required", False), ("optional", True)):
section = input_types.get(section_name, {})
for name, declaration in section.items():
if name in declarations:
raise ValueError(f"legacy input {name} is declared more than once.")
declarations[name] = (declaration, optional)
order = tuple(declarations) if workflow_order is None else workflow_order
if len(order) != len(set(order)) or set(order) != set(declarations):
raise ValueError("legacy workflow input order must name every input once.")
return [
_v3_input(name, declarations[name][0], optional=declarations[name][1])
for name in order
]
def _v3_input(name: str, declaration: object, *, optional: bool) -> Any:
"""Return one v3 input declaration from a legacy field declaration."""
if not isinstance(declaration, tuple) or not declaration:
raise TypeError(f"legacy input {name} declaration must be a tuple.")
io_declaration = declaration[0]
options = _input_options(declaration)
tooltip = _string_option(options, "tooltip")
advanced = _bool_option(options, "advanced")
raw_link = _bool_option(options, "rawLink") or _bool_option(options, "raw_link")
force_input = _bool_option(options, "forceInput") or _bool_option(
options, "force_input"
)
if isinstance(io_declaration, (list, tuple)):
return _comfy_io.Combo.Input(
name,
options=list(io_declaration),
optional=optional,
default=options.get("default"),
control_after_generate=options.get("control_after_generate"),
tooltip=tooltip,
raw_link=raw_link,
advanced=advanced,
)
if not isinstance(io_declaration, str):
raise TypeError(f"legacy input {name} type must be a string or options list.")
input_type = io_declaration
input_class = _io_class(input_type)
common_options = {
"optional": optional,
"tooltip": tooltip,
"raw_link": raw_link,
"advanced": advanced,
}
if input_type == "INT":
return input_class.Input(
name,
default=options.get("default"),
min=options.get("min"),
max=options.get("max"),
step=options.get("step"),
control_after_generate=options.get("control_after_generate"),
**common_options,
)
if input_type == "FLOAT":
return input_class.Input(
name,
default=options.get("default"),
min=options.get("min"),
max=options.get("max"),
step=options.get("step"),
round=options.get("round"),
**common_options,
)
if input_type == "STRING":
return input_class.Input(
name,
default=options.get("default"),
multiline=bool(options.get("multiline", False)),
force_input=force_input,
**common_options,
)
if input_type == "BOOLEAN":
return input_class.Input(
name,
default=options.get("default"),
label_on=options.get("label_on"),
label_off=options.get("label_off"),
**common_options,
)
return input_class.Input(name, **common_options)
def _v3_outputs(legacy: type[Any]) -> list[Any]:
"""Return v3 output declarations from legacy return metadata."""
return_types = tuple(getattr(legacy, "RETURN_TYPES", ()))
return_names = getattr(legacy, "RETURN_NAMES", None)
output_tooltips = tuple(getattr(legacy, "OUTPUT_TOOLTIPS", ()))
output_is_list = tuple(
getattr(legacy, "OUTPUT_IS_LIST", (False,) * len(return_types))
)
outputs: list[Any] = []
for index, io_type in enumerate(return_types):
output_name = None
if isinstance(return_names, tuple) and index < len(return_names):
output_name = str(return_names[index])
tooltip = None
if index < len(output_tooltips):
tooltip = str(output_tooltips[index])
is_output_list = index < len(output_is_list) and bool(output_is_list[index])
outputs.append(
_io_class(str(io_type)).Output(
output_name,
tooltip=tooltip,
is_output_list=is_output_list,
)
)
return outputs
def _v3_hidden_inputs(input_types: Mapping[str, Mapping[str, object]]) -> list[Any]:
"""Return v3 hidden declarations requested by legacy hidden inputs."""
hidden_values = set(_legacy_hidden_inputs(input_types).values())
return [getattr(_comfy_io.Hidden, value) for value in sorted(hidden_values)]
def _legacy_hidden_inputs(
input_types: Mapping[str, Mapping[str, object]],
) -> dict[str, str]:
"""Return legacy hidden input names mapped to v3 hidden holder attributes."""
hidden_inputs: dict[str, str] = {}
for name, sentinel in input_types.get("hidden", {}).items():
if isinstance(sentinel, str) and sentinel in _HIDDEN_INPUTS:
hidden_inputs[name] = _HIDDEN_INPUTS[sentinel]
return hidden_inputs
def _io_class(io_type: str) -> Any:
"""Return the v3 IO class for a legacy Comfy type string."""
known_types = {
"BOOLEAN": _comfy_io.Boolean,
"INT": _comfy_io.Int,
"FLOAT": _comfy_io.Float,
"STRING": _comfy_io.String,
"IMAGE": _comfy_io.Image,
"MASK": _comfy_io.Mask,
"LATENT": _comfy_io.Latent,
"MODEL": _comfy_io.Model,
"CLIP": _comfy_io.Clip,
"VAE": _comfy_io.Vae,
"CONDITIONING": _comfy_io.Conditioning,
"SEGS": _comfy_io.SEGS,
}
return known_types.get(io_type, _comfy_io.Custom(io_type))
def _input_options(declaration: tuple[object, ...]) -> dict[str, object]:
"""Return an input options dictionary from a legacy declaration."""
if len(declaration) < 2 or not isinstance(declaration[1], dict):
return {}
return dict(declaration[1])
def _string_option(options: Mapping[str, object], name: str) -> str | None:
"""Return a string option when present."""
value = options.get(name)
if isinstance(value, str):
return value
return None
def _bool_option(options: Mapping[str, object], name: str) -> bool | None:
"""Return a boolean option when present."""
value = options.get(name)
if isinstance(value, bool):
return value
return None
__all__ = [
"ConditioningBatchAppendV3",
"ConditioningBatchStartV3",
@@ -0,0 +1,58 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Expose inversion resolution and integration controls as sampler options."""
from __future__ import annotations
from typing import Any
from ..domain.sampler_options import SamplerOptions, append_sampler_capability
from .sampler_options_schema import (
COMFY_IO,
OptionsNodeBase,
inversion_from_controls,
noise_inversion_inputs,
options_input,
options_output,
)
class NoiseInversionOptionsV3(OptionsNodeBase):
"""Add source-derived starting noise to an immutable sampler options chain."""
@classmethod
def define_schema(cls) -> Any:
"""Declare the accepted recipe with independently editable controls."""
return COMFY_IO.Schema(
node_id="SimpleSyrup.NoiseInversionOptions",
display_name="Noise Inversion Options",
category="SimpleSyrup/Sampling/Options",
description=(
"Derives starting noise from an input image before sampling; "
"control inversion quality and cost independently."
),
inputs=[options_input(COMFY_IO), *noise_inversion_inputs(COMFY_IO)],
outputs=[options_output(COMFY_IO)],
)
@classmethod
def execute(
cls,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
options: SamplerOptions | None = None,
) -> tuple[SamplerOptions]:
"""Append inversion or pass through at zero steps without preparing a model."""
inversion = inversion_from_controls(
inversion_method=inversion_method,
inversion_resolution_scale=inversion_resolution_scale,
inversion_steps=inversion_steps,
inversion_switch_fraction=inversion_switch_fraction,
inversion_finishing_steps=inversion_finishing_steps,
)
return (append_sampler_capability(options, inversion),)
@@ -0,0 +1,134 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Declare bypass-compatible sampler options sockets and inversion controls."""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar, cast
from ..domain.noise_inversion import (
INVERSION_METHODS,
InversionMethod,
NoiseInversionOptions,
)
if TYPE_CHECKING:
class OptionsNodeBase:
"""Describe Comfy's host-facing node metadata for strict type checking."""
RETURN_TYPES: ClassVar[list[str]]
RETURN_NAMES: ClassVar[list[str]]
else:
OptionsNodeBase = import_module("comfy_api.latest").io.ComfyNode
COMFY_IO: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
OPTIONS_TYPE = "SIMPLE_SYRUP_SAMPLER_OPTIONS"
def options_input(comfy_io: Any) -> Any:
"""Allow any capability to start a chain or consume a preceding capability."""
return comfy_io.Custom(OPTIONS_TYPE).Input(
"options",
optional=True,
tooltip=(
"Optional preceding sampler options; bypass this node "
"to omit its contribution."
),
)
def options_output(comfy_io: Any) -> Any:
"""Match the input type so Comfy can bypass capability nodes natively."""
return comfy_io.Custom(OPTIONS_TYPE).Output(
"options",
tooltip="Combined sampler options; connect another options node or KSampler.",
)
def noise_inversion_inputs(comfy_io: Any, *, convenience: bool = False) -> list[Any]:
"""Default to the accepted recipe and use zero steps to disable inversion."""
return [
comfy_io.Combo.Input(
"inversion_method",
options=list(INVERSION_METHODS),
default="euler",
optional=convenience,
tooltip=(
"Applies to both inversion stages; Euler uses one evaluation per step, "
"Heun uses two for greater accuracy."
),
),
comfy_io.Float.Input(
"inversion_resolution_scale",
default=0.5,
min=0.01,
max=1.0,
step=0.05,
optional=convenience,
tooltip=(
"Scales inversion width and height; "
"0.5 uses half-sized dimensions for lower cost."
),
),
comfy_io.Int.Input(
"inversion_steps",
default=2,
min=0,
max=64,
optional=convenience,
tooltip=(
"Steps at the selected inversion resolution; 0 disables all inversion, "
"including finishing. More steps cost more model evaluations."
),
),
comfy_io.Float.Input(
"inversion_switch_fraction",
default=0.75,
min=0.01,
max=1.0,
step=0.05,
optional=convenience,
tooltip=(
"Noise-level fraction reached before the full-resolution finish; "
"0.75 means 75%."
),
),
comfy_io.Int.Input(
"inversion_finishing_steps",
default=1,
min=0,
max=64,
optional=convenience,
tooltip=(
"Full-resolution inversion steps after a reduced stage; "
"0 finishes entirely at reduced size."
),
),
]
def inversion_from_controls(
*,
inversion_method: str = "euler",
inversion_resolution_scale: float = 0.5,
inversion_steps: int = 2,
inversion_switch_fraction: float = 0.75,
inversion_finishing_steps: int = 1,
) -> NoiseInversionOptions | None:
"""Disable all stages at zero steps or construct a shared-method recipe."""
if type(inversion_steps) is not int or not 0 <= inversion_steps <= 64:
raise ValueError("Inversion steps must be an integer between 0 and 64.")
if inversion_steps == 0:
return None
return NoiseInversionOptions(
method=cast(InversionMethod, inversion_method),
resolution_scale=inversion_resolution_scale,
steps=inversion_steps,
switch_fraction=inversion_switch_fraction,
finishing_steps=inversion_finishing_steps,
)
+78
View File
@@ -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
"""Configure the single local tiling authority for a sampler options chain."""
from __future__ import annotations
from typing import Any
from ..domain.sampler_options import (
SamplerOptions,
TilingOptions,
append_sampler_capability,
)
from .ksampler_schema import tiled_diffusion_inputs
from .sampler_options_schema import (
COMFY_IO,
OptionsNodeBase,
options_input,
options_output,
)
class TilingOptionsV3(OptionsNodeBase):
"""Add bounded local tiles, blend policy and mask-dependent denoising."""
@classmethod
def define_schema(cls) -> Any:
"""Share tiled controls and expose mask-dependent denoising."""
return COMFY_IO.Schema(
node_id="SimpleSyrup.TilingOptions",
display_name="Tiling Options",
category="SimpleSyrup/Sampling/Options",
description=(
"Samples bounded local tiles; ignored when "
"Contextual Diffusion Options is connected."
),
inputs=[
options_input(COMFY_IO),
*tiled_diffusion_inputs(COMFY_IO),
COMFY_IO.Boolean.Input(
"differential_diffusion",
default=False,
tooltip=(
"Uses the noise mask to vary denoising strength spatially; "
"preserves existing model mask behavior."
),
),
],
outputs=[options_output(COMFY_IO)],
)
@classmethod
def execute(
cls,
diffusion_mode: str = "multidiffusion",
latent_tile_width: int = 128,
latent_tile_height: int = 128,
latent_tile_overlap: int = 16,
latent_tile_batch_size: int = 4,
differential_diffusion: bool = False,
options: SamplerOptions | None = None,
) -> tuple[SamplerOptions]:
"""Append validated tiling without changing the incoming chain."""
return (
append_sampler_capability(
options,
TilingOptions(
diffusion_mode=diffusion_mode,
width=latent_tile_width,
height=latent_tile_height,
overlap=latent_tile_overlap,
batch_size=latent_tile_batch_size,
differential_diffusion=differential_diffusion,
),
),
)
@@ -16,16 +16,24 @@ from ..domain.contextual_diffusion import (
ContextualDiffusionControls,
ContextualDiffusionPlan,
)
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import (
EMPTY_REGIONAL_CAPABILITY_ADMISSION,
RegionalCapabilityAdmission,
)
from ..domain.sampler_options import TilingOptions
from ..domain.segs import NativeSegs
from ..shared.logging import get_logger
from . import sampling_samplers, sampling_schedulers
from . import sampling_noise, sampling_samplers, sampling_schedulers
from .contextual_model_wrapper import ContextualDiffusionModelWrapper
from .differential_diffusion import (
differential_diffusion_mutation,
has_denoise_mask_function,
)
from .guided_sampling import sample_with_optional_negative
from .inversion_model_factory import InversionModelFactory
from .model_patcher_mutations import ModelUnetWrapperMutation
from .patcher_lifecycle import PATCHER_LIFECYCLE
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
from .sampling_model_types import ModelFunctionWrapper
from .tiled_sampling_validation import (
Latent,
@@ -58,14 +66,18 @@ def sample_contextual_diffusion(
capability_admission: RegionalCapabilityAdmission = (
EMPTY_REGIONAL_CAPABILITY_ADMISSION
),
noise_inversion: NoiseInversionOptions | None = None,
inversion_segs: NativeSegs | None = None,
inversion_region_masks: torch.Tensor | None = None,
differential_diffusion: bool = False,
) -> Latent:
"""Sample one latent through global context and one tiled prediction plan."""
validate_sampling_controls(
steps=steps,
denoise=denoise,
latent_tile_width=controls.latent_context_size,
latent_tile_height=controls.latent_context_size,
latent_tile_width=controls.tile_width,
latent_tile_height=controls.tile_height,
latent_tile_batch_size=controls.latent_context_batch_size,
)
controls.validate()
@@ -90,8 +102,8 @@ def sample_contextual_diffusion(
steps=steps,
denoise=denoise,
view=sampling_schedulers.SchedulerView(
latent_width=controls.latent_context_size,
latent_height=controls.latent_context_size,
latent_width=controls.tile_width,
latent_height=controls.tile_height,
),
).to(model.load_device)
latent_samples = validate_latent_samples(latent_image, sampler_label=SAMPLER_LABEL)
@@ -114,9 +126,38 @@ def sample_contextual_diffusion(
controls=controls,
sigmas=sigmas,
diffusion_mode=diffusion_mode,
differential_diffusion=differential_diffusion,
)
batch_inds = latent_image.get("batch_index")
noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds)
inversion_factory = (
InversionModelFactory(
model=model,
canvas_width=plan.latent_width,
canvas_height=plan.latent_height,
tiling=TilingOptions(
diffusion_mode=diffusion_mode,
width=controls.tile_width,
height=controls.tile_height,
overlap=controls.latent_context_overlap,
batch_size=controls.latent_context_batch_size,
differential_diffusion=differential_diffusion,
),
context=controls,
forward_sigmas=sigmas,
segs=inversion_segs,
region_masks=inversion_region_masks,
)
if noise_inversion is not None
else None
)
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=batch_inds,
model=sampling_model,
)
callback = _latent_preview().prepare_callback(sampling_model, steps)
samples = sample_with_optional_negative(
comfy_sample=comfy_sample,
@@ -132,6 +173,8 @@ def sample_contextual_diffusion(
callback=callback,
disable_pbar=not comfy_utils.PROGRESS_BAR_ENABLED,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=inversion_factory,
)
LOGGER.info(
@@ -171,10 +214,12 @@ def clone_model_with_contextual_diffusion(
controls: ContextualDiffusionControls,
sigmas: torch.Tensor,
diffusion_mode: str,
differential_diffusion: bool = False,
existing_wrapper: ModelFunctionWrapper | None = None,
) -> Any:
"""Derive a model with one pre-CFG contextual prediction wrapper."""
old_wrapper = model.model_options.get("model_function_wrapper")
old_wrapper = existing_wrapper or 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.")
wrapper = ContextualDiffusionModelWrapper(
@@ -184,9 +229,13 @@ def clone_model_with_contextual_diffusion(
diffusion_mode=diffusion_mode,
existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper),
)
mutations: list[ModelMutation] = []
if differential_diffusion and not has_denoise_mask_function(model):
mutations.append(differential_diffusion_mutation())
mutations.append(ModelUnetWrapperMutation(wrapper))
return PATCHER_LIFECYCLE.derive_model(
model,
(ModelUnetWrapperMutation(wrapper),),
mutations,
operation="SimpleSyrup contextual diffusion",
)
+9 -2
View File
@@ -11,7 +11,7 @@ from typing import Any, TypeAlias, cast
import torch
from . import sampling_samplers, sampling_schedulers
from . import sampling_noise, sampling_samplers, sampling_schedulers
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
from .differential_diffusion import clone_with_differential_diffusion
from .guided_sampling import sample_with_optional_negative
@@ -81,7 +81,14 @@ class DetailSampler:
batch_inds = (
latent_image["batch_index"] if "batch_index" in latent_image else None
)
noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds)
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=batch_inds,
model=model,
)
noise_mask = latent_image.get("noise_mask", None)
if preview_context is None:
callback = _latent_preview().prepare_callback(model, steps)
+20 -1
View File
@@ -11,7 +11,9 @@ from typing import Any, cast
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..shared.logging import get_logger
from .noise_inversion import InversionModelFactory, invert_sampling_noise
LOGGER = get_logger(__name__)
@@ -31,8 +33,25 @@ def sample_with_optional_negative(
callback: Any = None,
disable_pbar: bool = False,
seed: int | None = None,
noise_inversion: NoiseInversionOptions | None = None,
inversion_model_factory: InversionModelFactory | None = None,
) -> torch.Tensor:
"""Sample with CFG when negative exists or Comfy's positive-only path otherwise."""
"""Prepare optional source-derived noise and select the actual Comfy guider."""
if noise_inversion is not None:
inversion = invert_sampling_noise(
model=model,
latent=latent_image,
forward_sigmas=sigmas,
positive=positive,
negative=negative,
cfg=cfg,
seed=seed,
options=noise_inversion,
model_factory=inversion_model_factory,
noise_mask=noise_mask,
)
noise = inversion.noise.to(noise)
if negative is not None:
return cast(
@@ -0,0 +1,164 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Prepare each inversion resolution from the original, spatially unwrapped MODEL."""
from __future__ import annotations
from typing import Any, cast
import torch
import torch.nn.functional as functional
from ..domain.contextual_diffusion import (
ContextualDiffusionControls,
build_contextual_diffusion_plan,
)
from ..domain.regional_tiled_diffusion import (
build_region_constrained_tiled_diffusion_plan,
)
from ..domain.sampler_options import TilingOptions
from ..domain.segs import NativeSegs
from ..domain.segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan
from ..domain.tiled_diffusion import TiledDiffusionPlan, build_tiled_diffusion_plan
from .inversion_spatial_context import InversionSpatialContextWrapper
from .model_patcher_mutations import ModelUnetWrapperMutation
from .patcher_lifecycle import PATCHER_LIFECYCLE
from .sampling_model_types import ModelFunctionWrapper
class InversionModelFactory:
"""Own stage planning while sharing the forward spatial wrapper authorities."""
def __init__(
self,
*,
model: Any,
canvas_width: int,
canvas_height: int,
tiling: TilingOptions | None = None,
context: ContextualDiffusionControls | None = None,
forward_sigmas: torch.Tensor | None = None,
segs: NativeSegs | None = None,
region_masks: torch.Tensor | None = None,
) -> None:
"""Retain canonical inputs, not an already wrapped full-resolution model."""
if context is not None and (tiling is None or forward_sigmas is None):
raise ValueError(
"Contextual inversion requires tile controls and forward sigmas."
)
self._model = model
self._width = canvas_width
self._height = canvas_height
self._tiling = tiling
self._context = context
self._sigmas = forward_sigmas
self._segs = segs
self._masks = region_masks
def __call__(self, latent: torch.Tensor) -> Any:
"""Replan one resolution with canonical regional-mask coordinates."""
from .contextual_diffusion_sampling import clone_model_with_contextual_diffusion
from .mixture_of_diffusers_sampling import clone_model_with_mixture_of_diffusers
from .multidiffusion_sampling import clone_model_with_multidiffusion
width, height = int(latent.shape[-1]), int(latent.shape[-2])
old_wrapper = self._model.model_options.get("model_function_wrapper")
if old_wrapper is not None and not callable(old_wrapper):
raise TypeError("Existing model_function_wrapper must be callable.")
wrapper = cast(ModelFunctionWrapper | None, old_wrapper)
if wrapper is not None and (width, height) != (self._width, self._height):
wrapper = InversionSpatialContextWrapper(
wrapper,
canvas_width=self._width,
canvas_height=self._height,
stage_width=width,
stage_height=height,
)
masks = self._stage_masks(height, width)
if self._context is not None:
assert self._tiling is not None and self._sigmas is not None
plan = build_contextual_diffusion_plan(
latent_width=width,
latent_height=height,
controls=self._context,
segs=self._segs,
region_masks=masks,
segs_canvas=(self._height, self._width),
)
return clone_model_with_contextual_diffusion(
self._model,
plan=plan,
controls=self._context,
sigmas=self._sigmas,
diffusion_mode=self._tiling.diffusion_mode,
differential_diffusion=self._tiling.differential_diffusion,
existing_wrapper=wrapper,
)
if self._tiling is not None:
tile_plan = self._tile_plan(width, height, masks)
clone = (
clone_model_with_multidiffusion
if self._tiling.diffusion_mode == "multidiffusion"
else clone_model_with_mixture_of_diffusers
)
derived, _ = clone(
self._model,
latent_width=width,
latent_height=height,
tile_width=self._tiling.width,
tile_height=self._tiling.height,
overlap=self._tiling.overlap,
tile_batch_size=self._tiling.batch_size,
differential_diffusion=self._tiling.differential_diffusion,
tiled_plan=tile_plan,
existing_wrapper=wrapper,
)
return derived
if wrapper is old_wrapper:
return self._model
assert wrapper is not None
return PATCHER_LIFECYCLE.derive_model(
self._model,
(ModelUnetWrapperMutation(wrapper),),
operation="SimpleSyrup inversion spatial context",
)
def _stage_masks(self, height: int, width: int) -> torch.Tensor | None:
"""Resize planning masks once; attention masks keep their canonical bank."""
if self._masks is None:
return None
if tuple(self._masks.shape[-2:]) == (height, width):
return self._masks
return functional.interpolate(
self._masks.unsqueeze(1).float(), size=(height, width), mode="nearest"
).squeeze(1)
def _tile_plan(
self, width: int, height: int, masks: torch.Tensor | None
) -> TiledDiffusionPlan:
"""Share regular, SEGS and regional ownership with forward sampling."""
assert self._tiling is not None
geometry = {
"latent_width": width,
"latent_height": height,
"tile_width": self._tiling.width,
"tile_height": self._tiling.height,
"overlap": self._tiling.overlap,
"tile_batch_size": self._tiling.batch_size,
}
if masks is not None:
return build_region_constrained_tiled_diffusion_plan(
region_masks=masks,
segs=self._segs,
segs_canvas=(self._height, self._width),
**geometry,
)
if self._segs is not None:
return build_segs_guided_tiled_diffusion_plan(
segs=self._segs,
segs_canvas=(self._height, self._width),
**geometry,
)
return build_tiled_diffusion_plan(**geometry)
@@ -0,0 +1,128 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Project reduced inversion views into the original regional-attention canvas."""
from __future__ import annotations
from typing import Any
import torch
from ..domain.spatial_views import SpatialBatchLayout, SpatialView, SpatialViewKind
from .sampling_model_types import ApplyModel, ModelFunctionWrapper
from .spatial_model_arguments import (
SIMPLE_SYRUP_TRANSFORMER_NAMESPACE,
SPATIAL_BATCH_LAYOUT_KEY,
)
def rebase_inversion_layout(
layout: SpatialBatchLayout, *, canvas_width: int, canvas_height: int
) -> SpatialBatchLayout:
"""Keep actual model dimensions while mapping source rectangles to full masks."""
scale_x = canvas_width / layout.canvas_width
scale_y = canvas_height / layout.canvas_height
views: list[SpatialView] = []
for view in layout.views:
left = round(view.source_x * scale_x)
top = round(view.source_y * scale_y)
right = round(view.source_right * scale_x)
bottom = round(view.source_bottom * scale_y)
views.append(
SpatialView(
kind=(
SpatialViewKind.CONTEXTUAL_GLOBAL
if view.kind is SpatialViewKind.FULL
else view.kind
),
source_x=left,
source_y=top,
source_width=right - left,
source_height=bottom - top,
model_width=view.model_width,
model_height=view.model_height,
)
)
return SpatialBatchLayout(
canvas_width, canvas_height, tuple(views), layout.input_batch_size
)
class InversionSpatialContextWrapper:
"""Preserve regional mask coordinates without resizing conditioning twice."""
def __init__(
self,
existing_wrapper: ModelFunctionWrapper,
*,
canvas_width: int,
canvas_height: int,
stage_width: int,
stage_height: int,
) -> None:
"""Bind one stage's geometry to the original attention-mask canvas."""
self._existing_wrapper = existing_wrapper
self._canvas_width = canvas_width
self._canvas_height = canvas_height
self._stage_width = stage_width
self._stage_height = stage_height
def __call__(self, apply_model: ApplyModel, args: dict[str, Any]) -> torch.Tensor:
"""Replace only layout metadata, preserving expanded CFG batch metadata."""
x = args.get("input")
c = args.get("c", {})
if not isinstance(x, torch.Tensor) or not isinstance(c, dict):
raise TypeError(
"Inversion spatial context requires tensor input and conditioning."
)
transformer_options = c.get("transformer_options", {})
if not isinstance(transformer_options, dict):
raise TypeError("Inversion transformer_options must be a dictionary.")
namespace = transformer_options.get(SIMPLE_SYRUP_TRANSFORMER_NAMESPACE, {})
if not isinstance(namespace, dict):
raise TypeError(
"Inversion SimpleSyrup transformer namespace must be a dictionary."
)
layout = namespace.get(SPATIAL_BATCH_LAYOUT_KEY)
if layout is None:
if tuple(x.shape[-2:]) != (self._stage_height, self._stage_width):
raise ValueError(
"Reduced inversion calls require explicit spatial layout."
)
layout = SpatialBatchLayout(
self._stage_width,
self._stage_height,
(
SpatialView(
SpatialViewKind.FULL,
0,
0,
self._stage_width,
self._stage_height,
self._stage_width,
self._stage_height,
),
),
int(x.shape[0]),
)
if not isinstance(layout, SpatialBatchLayout):
raise TypeError("Inversion spatial layout has an invalid type.")
if (layout.canvas_width, layout.canvas_height) != (
self._stage_width,
self._stage_height,
) or layout.expanded_batch_size != int(x.shape[0]):
raise ValueError(
"Inversion layout must describe the current stage and model batch."
)
rebased = rebase_inversion_layout(
layout, canvas_width=self._canvas_width, canvas_height=self._canvas_height
)
projected_namespace = {**namespace, SPATIAL_BATCH_LAYOUT_KEY: rebased}
projected_options = {
**transformer_options,
SIMPLE_SYRUP_TRANSFORMER_NAMESPACE: projected_namespace,
}
projected_args = {**args, "c": {**c, "transformer_options": projected_options}}
return self._existing_wrapper(apply_model, projected_args)
@@ -16,22 +16,26 @@ from typing import Any, cast
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import (
EMPTY_REGIONAL_CAPABILITY_ADMISSION,
RegionalCapabilityAdmission,
)
from ..domain.sampler_options import TilingOptions
from ..domain.segs import NativeSegs
from ..domain.tiled_diffusion import (
TiledDiffusionPlan,
build_tiled_diffusion_plan,
)
from ..shared.logging import get_logger
from . import sampling_samplers, sampling_schedulers
from . import sampling_noise, sampling_samplers, sampling_schedulers
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
from .differential_diffusion import (
differential_diffusion_mutation,
has_denoise_mask_function,
)
from .guided_sampling import sample_with_optional_negative
from .inversion_model_factory import InversionModelFactory
from .model_patcher_mutations import ModelUnetWrapperMutation
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
from .sampling_model_types import (
@@ -73,6 +77,9 @@ def sample_mixture_of_diffusers(
EMPTY_REGIONAL_CAPABILITY_ADMISSION
),
tiled_plan: TiledDiffusionPlan | None = None,
noise_inversion: NoiseInversionOptions | None = None,
inversion_segs: NativeSegs | None = None,
inversion_region_masks: torch.Tensor | None = None,
) -> Latent:
"""Sample a latent with a cloned model patched for Mixture of Diffusers."""
@@ -135,7 +142,14 @@ def sample_mixture_of_diffusers(
)
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds)
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=batch_inds,
model=sampling_model,
)
noise_mask = latent_image.get("noise_mask", None)
callback = _sampling_callback(sampling_model, steps, preview_context)
samples = sample_with_optional_negative(
@@ -152,6 +166,26 @@ def sample_mixture_of_diffusers(
callback=callback,
disable_pbar=not comfy_utils.PROGRESS_BAR_ENABLED,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=(
InversionModelFactory(
model=model,
canvas_width=latent_width,
canvas_height=latent_height,
tiling=TilingOptions(
diffusion_mode="mixture_of_diffusers",
width=latent_tile_width,
height=latent_tile_height,
overlap=latent_tile_overlap,
batch_size=latent_tile_batch_size,
differential_diffusion=differential_diffusion,
),
segs=inversion_segs,
region_masks=inversion_region_masks,
)
if noise_inversion is not None
else None
),
)
LOGGER.info(
@@ -190,6 +224,7 @@ def clone_model_with_mixture_of_diffusers(
tile_batch_size: int,
differential_diffusion: bool = False,
tiled_plan: TiledDiffusionPlan | None = None,
existing_wrapper: ModelFunctionWrapper | None = None,
) -> tuple[Any, TiledDiffusionPlan]:
"""Return a derived model patched with a pre-CFG Mixture wrapper."""
@@ -202,7 +237,7 @@ def clone_model_with_mixture_of_diffusers(
tile_batch_size=tile_batch_size,
)
_validate_supplied_plan(plan, latent_width, latent_height)
old_wrapper = model.model_options.get("model_function_wrapper")
old_wrapper = existing_wrapper or 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.")
@@ -16,22 +16,26 @@ from typing import Any, cast
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import (
EMPTY_REGIONAL_CAPABILITY_ADMISSION,
RegionalCapabilityAdmission,
)
from ..domain.sampler_options import TilingOptions
from ..domain.segs import NativeSegs
from ..domain.tiled_diffusion import (
TiledDiffusionPlan,
build_tiled_diffusion_plan,
)
from ..shared.logging import get_logger
from . import sampling_samplers, sampling_schedulers
from . import sampling_noise, sampling_samplers, sampling_schedulers
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
from .differential_diffusion import (
differential_diffusion_mutation,
has_denoise_mask_function,
)
from .guided_sampling import sample_with_optional_negative
from .inversion_model_factory import InversionModelFactory
from .model_patcher_mutations import ModelUnetWrapperMutation
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
from .sampling_model_types import (
@@ -74,6 +78,9 @@ def sample_multidiffusion(
EMPTY_REGIONAL_CAPABILITY_ADMISSION
),
tiled_plan: TiledDiffusionPlan | None = None,
noise_inversion: NoiseInversionOptions | None = None,
inversion_segs: NativeSegs | None = None,
inversion_region_masks: torch.Tensor | None = None,
) -> Latent:
"""Sample a latent with a cloned model patched for MultiDiffusion."""
@@ -137,7 +144,14 @@ def sample_multidiffusion(
)
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds)
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=batch_inds,
model=sampling_model,
)
noise_mask = latent_image.get("noise_mask", None)
callback = _sampling_callback(sampling_model, steps, preview_context)
samples = sample_with_optional_negative(
@@ -154,6 +168,25 @@ def sample_multidiffusion(
callback=callback,
disable_pbar=not comfy_utils.PROGRESS_BAR_ENABLED,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=(
InversionModelFactory(
model=model,
canvas_width=latent_width,
canvas_height=latent_height,
tiling=TilingOptions(
width=latent_tile_width,
height=latent_tile_height,
overlap=latent_tile_overlap,
batch_size=latent_tile_batch_size,
differential_diffusion=differential_diffusion,
),
segs=inversion_segs,
region_masks=inversion_region_masks,
)
if noise_inversion is not None
else None
),
)
LOGGER.info(
@@ -193,6 +226,7 @@ def clone_model_with_multidiffusion(
tile_batch_size: int,
differential_diffusion: bool = False,
tiled_plan: TiledDiffusionPlan | None = None,
existing_wrapper: ModelFunctionWrapper | None = None,
) -> tuple[Any, TiledDiffusionPlan]:
"""Return a derived model patched with a pre-CFG MultiDiffusion wrapper."""
@@ -205,7 +239,7 @@ def clone_model_with_multidiffusion(
tile_batch_size=tile_batch_size,
)
_validate_supplied_plan(plan, latent_width, latent_height)
old_wrapper = model.model_options.get("model_function_wrapper")
old_wrapper = existing_wrapper or 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.")
+11 -4
View File
@@ -185,9 +185,16 @@ def krea2_diffusion_negpip_wrapper(
*args: object,
**kwargs: object,
) -> object:
"""Move a processed Krea sign mask into call-local transformer options."""
"""Inject call-local signs at either supported Comfy Krea argument boundary."""
positional_options = args[5] if len(args) > 5 else None
options_index = (
5
if len(args) > 5
else 4
if len(args) == 5 and isinstance(args[4], dict)
else None
)
positional_options = args[options_index] if options_index is not None else None
transformer_options = (
positional_options
if positional_options is not None
@@ -201,9 +208,9 @@ def krea2_diffusion_negpip_wrapper(
if not isinstance(multiplier, torch.Tensor):
raise TypeError("Krea NegPiP processed mask must be a tensor.")
prepared[TRANSFORMER_MASK_KEY] = multiplier
if len(args) > 5:
if options_index is not None:
prepared_args = list(args)
prepared_args[5] = prepared
prepared_args[options_index] = prepared
return executor(*prepared_args, **kwargs)
kwargs["transformer_options"] = prepared
return executor(*args, **kwargs)
+144
View File
@@ -0,0 +1,144 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Supply model-local NegPiP attention hooks for Comfy's earlier Krea boundary."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from inspect import signature
from typing import Any, cast
import torch
from comfy.ldm.flux.math import apply_rope
from comfy.ldm.krea2.model import Attention
from comfy.ldm.modules.attention import optimized_attention_masked
from comfy.model_patcher import ModelPatcher
from einops import rearrange
from ..model_patcher_mutations import ModelCallableObjectPatchMutation
from .krea2 import TRANSFORMER_MASK_KEY
LOGGER = logging.getLogger(__name__)
def krea2_host_mutations(
model: ModelPatcher,
) -> tuple[ModelCallableObjectPatchMutation, ...]:
"""Adapt only the known five-argument host; retain native reference-capable Krea."""
diffusion = model.get_model_object("diffusion_model")
parameters = tuple(signature(diffusion._forward).parameters)
if "ref_latents" in parameters:
return ()
if parameters != (
"x",
"timesteps",
"context",
"attention_mask",
"transformer_options",
"kwargs",
):
raise ValueError(
f"Krea NegPiP does not support model signature {parameters!r}."
)
mutations: list[ModelCallableObjectPatchMutation] = []
for index, block in enumerate(diffusion.blocks):
attention = block.attn
if not isinstance(attention, Attention):
raise TypeError(f"Krea NegPiP requires host Attention at block {index}.")
mutations.append(
ModelCallableObjectPatchMutation(
f"diffusion_model.blocks.{index}.attn.forward",
Krea2HostAttention(attention, index, len(diffusion.blocks)),
)
)
if not mutations:
raise ValueError("Krea NegPiP requires at least one joint attention block.")
LOGGER.info(
"Krea NegPiP installed model-local host attention hooks",
extra={"blocks": len(mutations)},
)
return tuple(mutations)
@dataclass(frozen=True)
class Krea2HostAttention:
"""Preserve host attention math while exposing its missing pre-RoPE patch point.
Comfy 0.28's public model boundary lacks attention callbacks. Object patches
scope this adapter to the derived MODEL and Comfy restores them on unload.
Text-fusion attention stays untouched; only joint text/image blocks use it.
"""
attention: Any # Comfy Attention exposes dynamically constructed linear modules.
block_index: int
total_blocks: int
def __call__(
self,
x: torch.Tensor,
freqs: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
transformer_options: dict[str, Any] | None = None,
) -> torch.Tensor:
"""Run host QKV projections and attention with local patch metadata."""
options = {} if transformer_options is None else transformer_options.copy()
multiplier = options.get(TRANSFORMER_MASK_KEY)
if multiplier is None:
return cast(
torch.Tensor,
type(self.attention).forward(
self.attention,
x,
freqs,
mask,
transformer_options=options,
),
)
if not isinstance(multiplier, torch.Tensor) or multiplier.ndim != 3:
raise ValueError(
"Krea NegPiP host attention requires a processed token sign mask."
)
options.update(
block_index=self.block_index,
total_blocks=self.total_blocks,
block_type="single",
img_slice=[multiplier.shape[1], x.shape[1]],
)
attention = self.attention
q, k, v, gate = (
attention.wq(x),
attention.wk(x),
attention.wv(x),
attention.gate(x),
)
q = rearrange(q, "B L (H D) -> B H L D", H=attention.heads)
k = rearrange(k, "B L (H D) -> B H L D", H=attention.kvheads)
v = rearrange(v, "B L (H D) -> B H L D", H=attention.kvheads)
q, k = attention.qknorm(q, k)
for patch in options.get("patches", {}).get("attn1_patch", []):
result = patch(
q, k, v, pe=freqs, attn_mask=mask, extra_options=options.copy()
)
q, k, v = result.get("q", q), result.get("k", k), result.get("v", v)
freqs, mask = result.get("pe", freqs), result.get("attn_mask", mask)
if freqs is not None:
q, k = apply_rope(q, k, freqs)
if attention.kvheads != attention.heads:
repeats = attention.heads // attention.kvheads
k = k.repeat_interleave(repeats, dim=1)
v = v.repeat_interleave(repeats, dim=1)
out = optimized_attention_masked(
q,
k,
v,
attention.heads,
mask=mask,
skip_reshape=True,
transformer_options=options,
)
for patch in options.get("patches", {}).get("attn1_output_patch", []):
out = patch(out, options.copy())
return cast(torch.Tensor, attention.wo(out * torch.nn.functional.sigmoid(gate)))
+325
View File
@@ -0,0 +1,325 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Derive source-dependent sampling noise with measured, uncached inversion stages."""
from __future__ import annotations
import math
import time
from collections.abc import Callable
from dataclasses import dataclass
from importlib import import_module
from typing import Any, TypeAlias
import torch
from ..domain.inversion_solver import (
InversionSolverEvidence,
integrate_inversion,
lift_inversion_displacement,
)
from ..domain.noise_inversion import InversionMethod, NoiseInversionOptions
from ..shared.logging import get_logger
from .spatial_tensor_projection import resize_spatial_tensor
from .tiled_sampling_validation import validate_tensor_shape
LOGGER = get_logger(__name__)
INVERSION_START_SIGMA = 0.0001
InversionModelFactory: TypeAlias = Callable[[torch.Tensor], Any]
@dataclass(frozen=True, slots=True)
class InversionStageMeasurement:
"""Describe actual work and elapsed time for one inversion stage."""
name: str
seconds: float
latent_shape: tuple[int, ...]
steps: int
evaluations: int
@dataclass(frozen=True, slots=True)
class NoiseInversionResult:
"""Return inferred noise with paid inversion cost and reconstruction evidence."""
noise: torch.Tensor
seconds: float
stages: tuple[InversionStageMeasurement, ...]
reconstruction_max_error: float
def _clock(device: torch.device) -> float:
"""Measure completed CUDA work rather than asynchronous kernel submission."""
if device.type == "cuda":
torch.cuda.synchronize(device)
return time.perf_counter()
def _source_in_model_space(model: Any, latent: torch.Tensor) -> torch.Tensor:
"""Narrow host latent processing before arithmetic or model execution."""
source = model.model.process_latent_in(latent)
if not isinstance(source, torch.Tensor) or source.shape != latent.shape:
raise ValueError(
"Model latent processing must preserve the inversion source shape."
)
if not source.is_floating_point() or not bool(torch.isfinite(source).all()):
raise ValueError(
"Model latent processing must produce finite floating-point values."
)
return source.detach().float().cpu()
def validate_inversion_target(model: Any, sigmas: torch.Tensor) -> float:
"""Reject unsupported scaling and singular targets before model execution."""
sampling_types = import_module("comfy.model_sampling")
sampling = model.get_model_object("model_sampling")
if not isinstance(
sampling, (sampling_types.CONST, sampling_types.EPS)
) or isinstance(
sampling, (sampling_types.IMG_TO_IMG, sampling_types.IMG_TO_IMG_FLOW)
):
raise ValueError(
"Noise inversion requires a flow or EPS-compatible image model."
)
if sigmas.ndim != 1 or len(sigmas) < 2 or not bool(torch.isfinite(sigmas).all()):
raise ValueError("Noise inversion requires a finite forward sampling schedule.")
target = float(sigmas[0])
if target <= INVERSION_START_SIGMA:
raise ValueError(
"Noise inversion requires a positive partial-denoise start sigma."
)
if isinstance(sampling, sampling_types.CONST) and target >= 0.9999:
raise ValueError(
"Flow noise inversion requires denoise below the full-noise endpoint."
)
if isinstance(sampling, sampling_types.EPS):
maximum = float(sampling.sigma_max)
if target > maximum or math.isclose(target, maximum, rel_tol=1e-5):
raise ValueError(
"Noise inversion requires partial denoise, below the model's "
"maximum sigma."
)
return target
def invert_sampling_noise(
*,
model: Any,
latent: torch.Tensor,
forward_sigmas: torch.Tensor,
positive: Any,
negative: Any,
cfg: float,
seed: int | None,
options: NoiseInversionOptions,
model_factory: InversionModelFactory | None = None,
noise_mask: Any = None,
) -> NoiseInversionResult:
"""Invert a source latent through Comfy's actual CFG or positive-only guider.
A spatial sampler supplies a factory that replans each resolution from its
unwrapped prepared model. The final noise recreates the inferred endpoint
under the model's own affine noise scaling; no inference cache is used.
"""
from .guided_sampling import sample_with_optional_negative
if not isinstance(options, NoiseInversionOptions):
raise TypeError("Noise inversion requires validated NoiseInversionOptions.")
validate_tensor_shape(latent, sampler_label="Noise Inversion")
if not latent.is_floating_point() or not bool(torch.isfinite(latent).all()):
raise ValueError(
"Noise inversion source must contain finite floating-point values."
)
target = validate_inversion_target(model, forward_sigmas)
coarse_target = target * options.coarse_target_fraction
if coarse_target <= INVERSION_START_SIGMA:
raise ValueError(
"Noise inversion transition must exceed the initial inversion sigma."
)
device = torch.device(model.load_device)
started = _clock(device)
sampling = model.get_model_object("model_sampling")
source = _source_in_model_space(model, latent)
phases: list[InversionStageMeasurement] = []
comfy_sample = import_module("comfy.sample")
comfy_samplers = import_module("comfy.samplers")
def stage(
stage_latent: torch.Tensor,
begin: float,
end: float,
count: int,
initial: torch.Tensor | None,
name: str,
method: InversionMethod,
) -> torch.Tensor:
"""Capture the model-space endpoint before Comfy converts output latents."""
stage_started = _clock(device)
stage_model = (
model_factory(stage_latent) if model_factory is not None else model
)
schedule = torch.linspace(begin, end, count + 1, device=device)
evidence = InversionSolverEvidence()
endpoints: list[torch.Tensor] = []
def invert(
model_fn: Any,
state: torch.Tensor,
sigmas: torch.Tensor,
extra_args: dict[str, Any],
callback: Any,
disable: bool,
) -> torch.Tensor:
"""Use Comfy's denoised predictions as the inversion velocity field."""
if initial is not None:
state = initial.to(state)
def evaluate(
x: torch.Tensor, sigma: torch.Tensor, index: int
) -> torch.Tensor:
"""Narrow the dynamic Comfy model result before numeric integration."""
prediction = model_fn(
x, sigma * x.new_ones((x.shape[0],)), **extra_args
)
if not isinstance(prediction, torch.Tensor):
raise TypeError(
"Noise inversion model must return tensor predictions."
)
return (x - prediction) / sigma
endpoint = integrate_inversion(
state, sigmas, evaluate, method=method, evidence=evidence
)
endpoints.append(endpoint.detach().float().cpu())
return endpoint
sample_with_optional_negative(
comfy_sample=comfy_sample,
model=stage_model,
noise=torch.zeros_like(stage_latent),
cfg=cfg,
sampler=comfy_samplers.KSAMPLER(invert),
sigmas=schedule,
positive=positive,
negative=negative,
latent_image=stage_latent,
noise_mask=noise_mask,
seed=seed,
disable_pbar=True,
)
if len(endpoints) != 1:
raise RuntimeError(
"Noise inversion must produce exactly one endpoint per stage."
)
phases.append(
InversionStageMeasurement(
name,
_clock(device) - stage_started,
tuple(stage_latent.shape),
count,
evidence.evaluations,
)
)
return endpoints[0]
if options.resolution_scale == 1:
endpoint = stage(
latent,
INVERSION_START_SIGMA,
target,
options.steps,
None,
"full",
options.method,
)
else:
height, width = options.coarse_shape(
int(latent.shape[-2]), int(latent.shape[-1])
)
coarse = resize_spatial_tensor(latent, height=height, width=width, mode="area")
coarse_source = _source_in_model_space(model, coarse)
coarse_endpoint = stage(
coarse,
INVERSION_START_SIGMA,
coarse_target,
options.steps,
None,
"coarse",
options.method,
)
endpoint = lift_inversion_displacement(
source,
coarse_source,
coarse_endpoint,
resize=lambda x, h, w: resize_spatial_tensor(
x, height=h, width=w, mode="bilinear"
),
)
if options.finishing_steps:
endpoint = stage(
latent,
coarse_target,
target,
options.finishing_steps,
endpoint,
"full_finish",
options.method,
)
sigma = torch.tensor(target)
zero = torch.zeros_like(source)
base = sampling.noise_scaling(sigma, zero.clone(), source, max_denoise=False)
amplitude = sampling.noise_scaling(
sigma, torch.ones_like(source), zero, max_denoise=False
)
if not isinstance(base, torch.Tensor) or not isinstance(amplitude, torch.Tensor):
raise TypeError("Model noise scaling must return tensors.")
if not bool(torch.isfinite(amplitude).all()) or bool(torch.any(amplitude == 0)):
raise ValueError("Model noise scaling is not invertible at the target sigma.")
noise = (endpoint - base) / amplitude
if not bool(torch.isfinite(noise).all()):
raise FloatingPointError("Noise inversion produced non-finite sampling noise.")
reconstructed = sampling.noise_scaling(
sigma, noise.clone(), source, max_denoise=False
)
if (
not isinstance(reconstructed, torch.Tensor)
or reconstructed.shape != endpoint.shape
):
raise ValueError(
"Model noise scaling must preserve the inversion endpoint shape."
)
if not torch.allclose(reconstructed, endpoint, atol=1e-5, rtol=1e-5):
raise ValueError(
"Model noise scaling cannot reconstruct the inversion endpoint."
)
error = float((reconstructed - endpoint).abs().max())
elapsed = _clock(device) - started
LOGGER.info(
"Noise inversion completed in %.3f seconds",
elapsed,
extra={
"operation": "noise_inversion",
"method": options.method,
"resolution_scale": options.resolution_scale,
"steps": options.steps,
"finishing_steps": options.finishing_steps,
"inversion_seconds": elapsed,
"inversion_stages": [
{
"name": phase.name,
"seconds": phase.seconds,
"latent_shape": phase.latent_shape,
"steps": phase.steps,
"evaluations": phase.evaluations,
}
for phase in phases
],
"evaluations": sum(phase.evaluations for phase in phases),
"endpoint_reconstruction_max_error": error,
},
)
return NoiseInversionResult(noise, elapsed, tuple(phases), error)
@@ -10,6 +10,7 @@ from typing import Any
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_detailing import LatentRegion
from . import regional_multidiffusion_sampling
from .detail_previews import DetailPreviewContext
@@ -51,6 +52,7 @@ class RegionalDetailSampler:
global_prompt_weight: float,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample one full latent with regional MultiDiffusion."""
@@ -69,4 +71,5 @@ class RegionalDetailSampler:
global_prompt_weight=global_prompt_weight,
preview_context=preview_context,
differential_diffusion=differential_diffusion,
noise_inversion=noise_inversion,
)
@@ -0,0 +1,66 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Rebuild regional prediction ownership from canonical inputs for inversion stages."""
from __future__ import annotations
from typing import Any
import torch
from ..domain.regional_detailing import LatentRegion
from ..domain.regional_inversion_geometry import project_inversion_regions
from .inversion_model_factory import InversionModelFactory
class RegionalInversionModelFactory:
"""Retain one original MODEL and the full-resolution region bank."""
def __init__(
self,
*,
model: Any,
canvas_width: int,
canvas_height: int,
regions: tuple[LatentRegion, ...],
global_prompt_weight: float,
differential_diffusion: bool,
) -> None:
"""Keep regional conditioning unchanged across source-sized inversion views."""
self._base = InversionModelFactory(
model=model,
canvas_width=canvas_width,
canvas_height=canvas_height,
)
self._width = canvas_width
self._height = canvas_height
self._regions = regions
self._weight = global_prompt_weight
self._differential = differential_diffusion
def __call__(self, latent: torch.Tensor) -> Any:
"""Use the authoritative regional calc-cond-batch wrapper at each stage size."""
from .regional_multidiffusion_sampling import (
clone_model_with_regional_multidiffusion,
)
height, width = int(latent.shape[-2]), int(latent.shape[-1])
regions = project_inversion_regions(
self._regions,
source_width=self._width,
source_height=self._height,
target_width=width,
target_height=height,
)
derived, _ = clone_model_with_regional_multidiffusion(
self._base(latent),
latent_width=width,
latent_height=height,
latent_ndim=latent.ndim,
regions=regions,
global_prompt_weight=self._weight,
differential_diffusion=self._differential,
)
return derived
@@ -15,10 +15,11 @@ from importlib import import_module
from types import ModuleType
from typing import Any, cast
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_detailing import LatentRegion
from ..domain.regional_features import EMPTY_REGIONAL_CAPABILITY_ADMISSION
from ..shared.logging import get_logger
from . import sampling_samplers, sampling_schedulers
from . import sampling_noise, sampling_samplers, sampling_schedulers
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
from .differential_diffusion import (
differential_diffusion_mutation,
@@ -27,6 +28,7 @@ from .differential_diffusion import (
from .guided_sampling import sample_with_optional_negative
from .model_patcher_mutations import ModelCalcCondBatchMutation
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
from .regional_inversion_model_factory import RegionalInversionModelFactory
from .regional_multidiffusion_prediction import (
CalcCondBatchFunction,
RegionalMultiDiffusionCalcCondBatch,
@@ -72,6 +74,7 @@ def sample_regional_multidiffusion(
global_prompt_weight: float,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample a latent with regional MultiDiffusion prompt blending."""
@@ -135,7 +138,14 @@ def sample_regional_multidiffusion(
)
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds)
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=batch_inds,
model=sampling_model,
)
noise_mask = latent_image.get("noise_mask", None)
callback = _sampling_callback(sampling_model, steps, preview_context)
samples = sample_with_optional_negative(
@@ -152,6 +162,19 @@ def sample_regional_multidiffusion(
callback=callback,
disable_pbar=not comfy_utils.PROGRESS_BAR_ENABLED,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=(
RegionalInversionModelFactory(
model=model,
canvas_width=latent_width,
canvas_height=latent_height,
regions=regions,
global_prompt_weight=global_prompt_weight,
differential_diffusion=differential_diffusion,
)
if noise_inversion is not None
else None
),
)
LOGGER.info(
@@ -0,0 +1,129 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
#
# Registry values follow RES4LYF beta/rk_coefficients_beta.py at
# 3d1d69da69ee47f7647d59e1bd0967e472fccc41.
"""Expose pinned RES4LYF sampler names without loading its solver at import."""
RES4LYF_SAMPLER_NAMES: tuple[str, ...] = (
"multistep/res_2m",
"multistep/res_3m",
"multistep/dpmpp_2m",
"multistep/dpmpp_3m",
"multistep/abnorsett_2m",
"multistep/abnorsett_3m",
"multistep/abnorsett_4m",
"multistep/deis_2m",
"multistep/deis_3m",
"multistep/deis_4m",
"exponential/res_2s_rkmk2e",
"exponential/res_2s",
"exponential/res_2s_stable",
"exponential/res_3s",
"exponential/res_3s_non-monotonic",
"exponential/res_3s_alt",
"exponential/res_3s_cox_matthews",
"exponential/res_3s_lie",
"exponential/res_3s_sunstar",
"exponential/res_3s_strehmel_weiner",
"exponential/res_4s_krogstad",
"exponential/res_4s_krogstad_alt",
"exponential/res_4s_strehmel_weiner",
"exponential/res_4s_strehmel_weiner_alt",
"exponential/res_4s_cox_matthews",
"exponential/res_4s_cfree4",
"exponential/res_4s_friedli",
"exponential/res_4s_minchev",
"exponential/res_4s_munthe-kaas",
"exponential/res_5s",
"exponential/res_5s_hochbruck-ostermann",
"exponential/res_6s",
"exponential/res_8s",
"exponential/res_8s_alt",
"exponential/res_10s",
"exponential/res_15s",
"exponential/res_16s",
"exponential/etdrk2_2s",
"exponential/etdrk3_a_3s",
"exponential/etdrk3_b_3s",
"exponential/etdrk4_4s",
"exponential/etdrk4_4s_alt",
"exponential/dpmpp_2s",
"exponential/dpmpp_sde_2s",
"exponential/dpmpp_3s",
"exponential/lawson2a_2s",
"exponential/lawson2b_2s",
"exponential/lawson4_4s",
"exponential/lawson41-gen_4s",
"exponential/lawson41-gen-mod_4s",
"exponential/ddim",
"hybrid/pec423_2h2s",
"hybrid/pec433_2h3s",
"hybrid/abnorsett2_1h2s",
"hybrid/abnorsett3_2h2s",
"hybrid/abnorsett4_3h2s",
"hybrid/lawson42-gen-mod_1h4s",
"hybrid/lawson43-gen-mod_2h4s",
"hybrid/lawson44-gen-mod_3h4s",
"hybrid/lawson45-gen-mod_4h4s",
"linear/ralston_2s",
"linear/ralston_3s",
"linear/ralston_4s",
"linear/midpoint_2s",
"linear/heun_2s",
"linear/heun_3s",
"linear/houwen-wray_3s",
"linear/kutta_3s",
"linear/ssprk3_3s",
"linear/ssprk4_4s",
"linear/rk38_4s",
"linear/rk4_4s",
"linear/rk5_7s",
"linear/rk6_7s",
"linear/bogacki-shampine_4s",
"linear/bogacki-shampine_7s",
"linear/dormand-prince_6s",
"linear/dormand-prince_13s",
"linear/tsi_7s",
"linear/euler",
"diag_implicit/irk_exp_diag_2s",
"diag_implicit/kraaijevanger_spijker_2s",
"diag_implicit/qin_zhang_2s",
"diag_implicit/pareschi_russo_2s",
"diag_implicit/pareschi_russo_alt_2s",
"diag_implicit/crouzeix_2s",
"diag_implicit/crouzeix_3s",
"diag_implicit/crouzeix_3s_alt",
"fully_implicit/gauss-legendre_2s",
"fully_implicit/gauss-legendre_3s",
"fully_implicit/gauss-legendre_4s",
"fully_implicit/gauss-legendre_4s_alternating_a",
"fully_implicit/gauss-legendre_4s_ascending_a",
"fully_implicit/gauss-legendre_4s_alt",
"fully_implicit/gauss-legendre_5s",
"fully_implicit/gauss-legendre_5s_ascending",
"fully_implicit/radau_ia_2s",
"fully_implicit/radau_ia_3s",
"fully_implicit/radau_iia_2s",
"fully_implicit/radau_iia_3s",
"fully_implicit/radau_iia_3s_alt",
"fully_implicit/radau_iia_5s",
"fully_implicit/radau_iia_7s",
"fully_implicit/radau_iia_9s",
"fully_implicit/radau_iia_11s",
"fully_implicit/lobatto_iiia_2s",
"fully_implicit/lobatto_iiia_3s",
"fully_implicit/lobatto_iiia_4s",
"fully_implicit/lobatto_iiib_2s",
"fully_implicit/lobatto_iiib_3s",
"fully_implicit/lobatto_iiib_4s",
"fully_implicit/lobatto_iiic_2s",
"fully_implicit/lobatto_iiic_3s",
"fully_implicit/lobatto_iiic_4s",
"fully_implicit/lobatto_iiic_star_2s",
"fully_implicit/lobatto_iiic_star_3s",
"fully_implicit/lobatto_iiid_2s",
"fully_implicit/lobatto_iiid_3s",
)
+84
View File
@@ -0,0 +1,84 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Resolve pinned RES4LYF solver methods within ComfyUI's sampler boundary."""
from __future__ import annotations
from importlib import import_module
from typing import Any, cast
import torch
from .res4lyf_sampler_names import RES4LYF_SAMPLER_NAMES
from .sampling_samplers import SamplerObject
def resolve_res4lyf_sampler(sampler_name: str) -> SamplerObject:
"""Bind an upstream method name to the pinned RES4LYF solver."""
if sampler_name not in RES4LYF_SAMPLER_NAMES:
raise ValueError(f"Unsupported RES4LYF sampler '{sampler_name}'.")
method = sampler_name.rsplit("/", 1)[-1]
implicit = sampler_name.startswith(("fully_implicit/", "diag_implicit/"))
options = {
"rk_type": "euler" if implicit else method,
"implicit_sampler_name": method if implicit else "use_explicit",
"implicit_type": "bongmath",
"implicit_type_substeps": "bongmath",
"bongmath": sampler_name != "linear/rk5_7s",
}
comfy_samplers = import_module("comfy.samplers")
return cast(
SamplerObject,
comfy_samplers.KSAMPLER(_sample_res4lyf, extra_options=options),
)
def _sample_res4lyf(
model: Any,
x: torch.Tensor,
sigmas: torch.Tensor,
*,
extra_args: dict[str, Any],
callback: Any,
disable: bool,
rk_type: str,
implicit_sampler_name: str,
implicit_type: str,
implicit_type_substeps: str,
bongmath: bool,
) -> torch.Tensor:
"""Pass ComfyUI's seed to the same SDE stream used by RES4LYF's node."""
seed = extra_args.get("seed")
if not isinstance(seed, int):
raise ValueError("RES4LYF sampler requires an integer sampling seed.")
solver = import_module(
"simple_syrup.third_party.res4lyf_runtime.beta.rk_sampler_beta"
)
samples = cast(
torch.Tensor,
solver.sample_rk_beta(
model,
x,
sigmas,
extra_args=extra_args,
callback=callback,
disable=disable,
rk_type=rk_type,
implicit_sampler_name=implicit_sampler_name,
implicit_type=implicit_type,
implicit_type_substeps=implicit_type_substeps,
BONGMATH=bongmath,
noise_seed=seed + 1,
),
)
if not torch.isfinite(samples).all():
raise ValueError(
f"RES4LYF sampler '{rk_type}' produced non-finite latent values. "
"Try another sampler or scheduler for this model."
)
return samples
+68
View File
@@ -0,0 +1,68 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Own initial latent noise generation for core and RES4LYF samplers."""
from __future__ import annotations
from importlib import import_module
from typing import Any, Protocol, cast
import torch
from .res4lyf_sampler_names import RES4LYF_SAMPLER_NAMES
class SamplingModel(Protocol):
"""Expose the model sampling object used by RES4LYF noise generation."""
def get_model_object(self, name: str) -> object:
"""Return a named ComfyUI model object."""
class ModelSamplingBounds(Protocol):
"""Expose the sigma limits used by RES4LYF's noise generator."""
sigma_max: float | torch.Tensor
sigma_min: float | torch.Tensor
def prepare_sampling_noise(
*,
comfy_sample: Any,
sampler_name: str,
samples: torch.Tensor,
seed: int,
batch_indices: Any,
model: SamplingModel,
) -> torch.Tensor:
"""Generate RES4LYF's default noise or preserve ComfyUI's core path."""
if sampler_name not in RES4LYF_SAMPLER_NAMES:
return cast(
torch.Tensor, comfy_sample.prepare_noise(samples, seed, batch_indices)
)
model_sampling = cast(ModelSamplingBounds, model.get_model_object("model_sampling"))
sigma_max = model_sampling.sigma_max
sigma_min = model_sampling.sigma_min
noise_classes = import_module(
"simple_syrup.third_party.res4lyf_runtime.beta.noise_classes"
)
latents = import_module("simple_syrup.third_party.res4lyf_runtime.latents")
reference = samples.to(torch.float32)
generator = noise_classes.NOISE_GENERATOR_CLASSES_SIMPLE["gaussian"](
x=reference.to(torch.float64),
seed=seed,
sigma_max=sigma_max,
sigma_min=sigma_min,
)
noise = cast(torch.Tensor, generator(sigma=sigma_max, sigma_next=sigma_min))
if noise.std() > 0:
noise = cast(
torch.Tensor,
latents.normalize_zscore(noise, channelwise=True, inplace=True),
)
noise = noise - noise.mean()
return noise.to(reference.dtype)
@@ -0,0 +1,283 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Own reference noise-level tables used by local AYS and GITS schedules."""
from __future__ import annotations
AYS_NOISE_LEVELS: dict[str, tuple[float, ...]] = {
"SD1": (
14.6146412293,
6.4745760956,
3.8636745985,
2.6946151520,
1.8841921177,
1.3943805092,
0.9642583904,
0.6523686016,
0.3977456272,
0.1515232662,
0.0291671582,
),
"SDXL": (
14.6146412293,
6.3184485287,
3.7681790315,
2.1811480769,
1.3405244945,
0.8620721141,
0.5550693289,
0.3798540708,
0.2332364134,
0.1114188177,
0.0291671582,
),
}
GITS_DEFAULT_NOISE_LEVELS: tuple[tuple[float, ...], ...] = (
(14.61464119, 0.803307, 0.02916753),
(14.61464119, 1.56271636, 0.52423614, 0.02916753),
(14.61464119, 2.36326075, 0.92192322, 0.36617002, 0.02916753),
(14.61464119, 2.84484982, 1.24153244, 0.59516323, 0.25053367, 0.02916753),
(
14.61464119,
5.85520077,
2.05039096,
0.95350921,
0.45573691,
0.17026083,
0.02916753,
),
(
14.61464119,
5.85520077,
2.45070267,
1.24153244,
0.64427125,
0.29807833,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.45070267,
1.36964464,
0.803307,
0.45573691,
0.25053367,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.84484982,
1.61558151,
0.95350921,
0.59516323,
0.36617002,
0.19894916,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.84484982,
1.67050016,
1.08895338,
0.74807048,
0.50118381,
0.32104823,
0.19894916,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.95596409,
1.84880662,
1.24153244,
0.83188516,
0.59516323,
0.41087446,
0.27464288,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
3.07277966,
1.98035145,
1.36964464,
0.95350921,
0.69515091,
0.50118381,
0.36617002,
0.25053367,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
6.77309084,
3.46139455,
2.36326075,
1.56271636,
1.08895338,
0.803307,
0.59516323,
0.45573691,
0.34370604,
0.25053367,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
6.77309084,
3.46139455,
2.45070267,
1.61558151,
1.162866,
0.86115354,
0.64427125,
0.50118381,
0.38853383,
0.29807833,
0.22545385,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.38853383,
0.29807833,
0.22545385,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.41087446,
0.32104823,
0.25053367,
0.19894916,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.41087446,
0.34370604,
0.27464288,
0.22545385,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.20157266,
0.92192322,
0.72133851,
0.57119018,
0.45573691,
0.36617002,
0.29807833,
0.25053367,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.24153244,
0.95350921,
0.74807048,
0.59516323,
0.4783645,
0.38853383,
0.32104823,
0.27464288,
0.22545385,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.24153244,
0.95350921,
0.74807048,
0.59516323,
0.50118381,
0.41087446,
0.34370604,
0.29807833,
0.25053367,
0.22545385,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
)
+7 -1
View File
@@ -16,6 +16,7 @@ from typing import Protocol, cast
from ..shared.logging import get_logger
from .a1111_sampling import sample_euler_ancestral_a1111
from .res4lyf_sampler_names import RES4LYF_SAMPLER_NAMES
LOGGER = get_logger(__name__)
EXTRA_SAMPLERS = ("euler_a_a1111",)
@@ -33,7 +34,7 @@ def available_samplers() -> tuple[str, ...]:
comfy_samplers = _comfy_samplers()
core_samplers = tuple(str(name) for name in comfy_samplers.KSampler.SAMPLERS)
return _unique_sampler_names(core_samplers + EXTRA_SAMPLERS)
return _unique_sampler_names(core_samplers + EXTRA_SAMPLERS + RES4LYF_SAMPLER_NAMES)
def resolve_sampler(sampler_name: str) -> SamplerObject:
@@ -57,6 +58,11 @@ def resolve_sampler(sampler_name: str) -> SamplerObject:
if sampler_name in EXTRA_SAMPLERS:
return _resolve_extra_sampler(sampler_name)
if sampler_name in RES4LYF_SAMPLER_NAMES:
from .res4lyf_sampling import resolve_res4lyf_sampler
return resolve_res4lyf_sampler(sampler_name)
return cast(SamplerObject, _comfy_samplers().sampler_object(sampler_name))
+16 -276
View File
@@ -19,6 +19,7 @@ from typing import Protocol, cast
import torch
from ..shared.logging import get_logger
from .sampling_reference_schedules import AYS_NOISE_LEVELS, GITS_DEFAULT_NOISE_LEVELS
LOGGER = get_logger(__name__)
@@ -27,6 +28,7 @@ EXTRA_SCHEDULERS = (
"AYS SDXL",
"GITS",
"beta57",
"bong_tangent",
"automatic_a1111",
"Flux2",
)
@@ -34,282 +36,6 @@ GITS_DEFAULT_COEFF = 1.20
BETA57_ALPHA = 0.5
BETA57_BETA = 0.7
AYS_NOISE_LEVELS: dict[str, tuple[float, ...]] = {
"SD1": (
14.6146412293,
6.4745760956,
3.8636745985,
2.6946151520,
1.8841921177,
1.3943805092,
0.9642583904,
0.6523686016,
0.3977456272,
0.1515232662,
0.0291671582,
),
"SDXL": (
14.6146412293,
6.3184485287,
3.7681790315,
2.1811480769,
1.3405244945,
0.8620721141,
0.5550693289,
0.3798540708,
0.2332364134,
0.1114188177,
0.0291671582,
),
}
GITS_DEFAULT_NOISE_LEVELS: tuple[tuple[float, ...], ...] = (
(14.61464119, 0.803307, 0.02916753),
(14.61464119, 1.56271636, 0.52423614, 0.02916753),
(14.61464119, 2.36326075, 0.92192322, 0.36617002, 0.02916753),
(14.61464119, 2.84484982, 1.24153244, 0.59516323, 0.25053367, 0.02916753),
(
14.61464119,
5.85520077,
2.05039096,
0.95350921,
0.45573691,
0.17026083,
0.02916753,
),
(
14.61464119,
5.85520077,
2.45070267,
1.24153244,
0.64427125,
0.29807833,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.45070267,
1.36964464,
0.803307,
0.45573691,
0.25053367,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.84484982,
1.61558151,
0.95350921,
0.59516323,
0.36617002,
0.19894916,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.84484982,
1.67050016,
1.08895338,
0.74807048,
0.50118381,
0.32104823,
0.19894916,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
2.95596409,
1.84880662,
1.24153244,
0.83188516,
0.59516323,
0.41087446,
0.27464288,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
5.85520077,
3.07277966,
1.98035145,
1.36964464,
0.95350921,
0.69515091,
0.50118381,
0.36617002,
0.25053367,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
6.77309084,
3.46139455,
2.36326075,
1.56271636,
1.08895338,
0.803307,
0.59516323,
0.45573691,
0.34370604,
0.25053367,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
6.77309084,
3.46139455,
2.45070267,
1.61558151,
1.162866,
0.86115354,
0.64427125,
0.50118381,
0.38853383,
0.29807833,
0.22545385,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.38853383,
0.29807833,
0.22545385,
0.17026083,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.41087446,
0.32104823,
0.25053367,
0.19894916,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.12350607,
1.51179266,
1.08895338,
0.83188516,
0.64427125,
0.50118381,
0.41087446,
0.34370604,
0.27464288,
0.22545385,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.20157266,
0.92192322,
0.72133851,
0.57119018,
0.45573691,
0.36617002,
0.29807833,
0.25053367,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.24153244,
0.95350921,
0.74807048,
0.59516323,
0.4783645,
0.38853383,
0.32104823,
0.27464288,
0.22545385,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
(
14.61464119,
7.49001646,
4.65472794,
3.07277966,
2.19988537,
1.61558151,
1.24153244,
0.95350921,
0.74807048,
0.59516323,
0.50118381,
0.41087446,
0.34370604,
0.29807833,
0.25053367,
0.22545385,
0.19894916,
0.17026083,
0.13792117,
0.09824532,
0.02916753,
),
)
class SamplingModel(Protocol):
"""Expose the ComfyUI model sampling object needed for core schedulers."""
@@ -513,6 +239,8 @@ def _calculate_extra_schedule(
return _calculate_gits_schedule(steps)
if scheduler_name == "beta57":
return _calculate_beta57_schedule(model, steps)
if scheduler_name == "bong_tangent":
return _calculate_bong_tangent_schedule(model, steps)
if scheduler_name == "automatic_a1111":
return _calculate_automatic_a1111_schedule(model, steps)
if scheduler_name == "Flux2":
@@ -592,6 +320,18 @@ def _calculate_beta57_schedule(model: SamplingModel, steps: int) -> torch.Tensor
)
def _calculate_bong_tangent_schedule(model: SamplingModel, steps: int) -> torch.Tensor:
"""Use RES4LYF's pinned tangent schedule with its default controls."""
res4lyf_sigmas = import_module("simple_syrup.third_party.res4lyf_runtime.sigmas")
return cast(
torch.Tensor,
res4lyf_sigmas.bong_tangent_scheduler(
model.get_model_object("model_sampling"), steps
),
)
def _calculate_automatic_a1111_schedule(
model: SamplingModel,
steps: int,
@@ -12,6 +12,7 @@ from ..domain.attention_coupling_request import (
AttentionCouplingRequestMode,
classify_attention_coupling_request,
)
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_attention_execution import RegionalAttentionExecutionMode
from .attention_coupling_model_preparation_service import (
AttentionCouplingModelPreparationService,
@@ -45,6 +46,7 @@ class AttentionCouplingSamplingService:
region_mask_feather: int,
latent_image: dict[str, Any],
denoise: float,
noise_inversion: NoiseInversionOptions | None = None,
) -> dict[str, Any]:
"""Bypass ordinary requests or prepare one complete regional request."""
@@ -65,6 +67,7 @@ class AttentionCouplingSamplingService:
negative=negative,
latent_image=latent_image,
denoise=denoise,
noise_inversion=noise_inversion,
)
prepared = self.model_preparation_service_class().prepare(
@@ -88,6 +91,7 @@ class AttentionCouplingSamplingService:
negative=prepared.negative,
latent_image=latent_image,
denoise=denoise,
noise_inversion=noise_inversion,
)
@@ -8,8 +8,10 @@ from __future__ import annotations
from typing import Any, ClassVar
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_attention_execution import RegionalAttentionExecutionMode
from ..domain.regional_features import RegionalFeature, RegionalFeatureRequest
from ..domain.sampler_options import TilingOptions
from .attention_coupling_model_preparation_service import (
AttentionCouplingModelPreparationService,
)
@@ -57,6 +59,8 @@ class ContextualAttentionCouplingSamplingService:
global_steps: int,
global_decay: float,
segs: object | None = None,
noise_inversion: NoiseInversionOptions | None = None,
tiling: TilingOptions | None = None,
) -> ContextualDiffusionSamplingResult:
"""Prepare once and invoke established Contextual local-view execution."""
@@ -94,6 +98,8 @@ class ContextualAttentionCouplingSamplingService:
region_mask_feather=0,
feature_request=_CONTEXTUAL_ATTENTION_REQUEST,
planning_region_masks=prepared.mask_bank.planning_masks,
noise_inversion=noise_inversion,
tiling=tiling,
)
@@ -21,6 +21,7 @@ from ..domain.contextual_diffusion import (
ContextualDiffusionControls,
build_contextual_diffusion_plan,
)
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import (
CONTEXTUAL_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES,
EMPTY_REGIONAL_FEATURE_REQUEST,
@@ -28,6 +29,7 @@ from ..domain.regional_features import (
RegionalFeature,
RegionalFeatureRequest,
)
from ..domain.sampler_options import TilingOptions
from ..domain.segs import NativeSegs, coerce_segs_group
from ..domain.tiled_diffusion import validate_tiled_diffusion_mode
from ..runtime.contextual_diffusion_sampling import sample_contextual_diffusion
@@ -91,10 +93,16 @@ class ContextualDiffusionSamplingService:
region_mask_feather: int = 0,
feature_request: RegionalFeatureRequest = EMPTY_REGIONAL_FEATURE_REQUEST,
planning_region_masks: torch.Tensor | None = None,
noise_inversion: NoiseInversionOptions | None = None,
tiling: TilingOptions | None = None,
) -> ContextualDiffusionSamplingResult:
"""Sample a latent with global and bounded detail contexts."""
"""Sample one local tile authority with global context and inversion."""
validate_tiled_diffusion_mode(diffusion_mode)
if tiling is not None:
diffusion_mode = tiling.diffusion_mode
latent_context_overlap = tiling.overlap
latent_context_batch_size = tiling.batch_size
controls = ContextualDiffusionControls(
latent_context_size=latent_context_size,
latent_context_overlap=latent_context_overlap,
@@ -102,6 +110,8 @@ class ContextualDiffusionSamplingService:
global_weight=global_weight,
global_steps=global_steps,
global_decay=global_decay,
latent_tile_width=tiling.width if tiling is not None else None,
latent_tile_height=tiling.height if tiling is not None else None,
)
controls.validate()
regional = self.regional_preparation_service_class().prepare(
@@ -172,6 +182,10 @@ class ContextualDiffusionSamplingService:
image_width=image_width,
region_masks=planning_masks,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
differential_diffusion=tiling.differential_diffusion
if tiling is not None
else False,
)
outputs: list[torch.Tensor] = []
@@ -207,6 +221,10 @@ class ContextualDiffusionSamplingService:
image_width=image_width,
region_masks=planning_masks,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
differential_diffusion=tiling.differential_diffusion
if tiling is not None
else False,
)
samples = item_result.latent.get("samples")
if not isinstance(samples, torch.Tensor):
@@ -240,8 +258,10 @@ class ContextualDiffusionSamplingService:
image_width: int,
region_masks: torch.Tensor | None,
capability_admission: RegionalCapabilityAdmission,
noise_inversion: NoiseInversionOptions | None,
differential_diffusion: bool,
) -> ContextualDiffusionSamplingResult:
"""Build one canvas plan and execute it through the runtime adapter."""
"""Build the forward plan and retain canonical ownership for inversion."""
samples = latent_image.get("samples")
if not isinstance(samples, torch.Tensor):
@@ -270,6 +290,10 @@ class ContextualDiffusionSamplingService:
controls=controls,
plan=plan,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
inversion_segs=segs,
inversion_region_masks=region_masks,
differential_diffusion=differential_diffusion,
)
return ContextualDiffusionSamplingResult(
latent=latent,
@@ -16,6 +16,7 @@ from typing import Any, Protocol
import torch
from ..domain.conditioning_batch import ConditioningBatch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_detailing import (
LatentRegion,
pair_segments_with_conditioning,
@@ -65,6 +66,7 @@ class RegionalDetailSamplingBoundary(Protocol):
global_prompt_weight: float,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample one full latent with paired regional conditioning."""
@@ -133,6 +135,7 @@ class DetailSEGSAsRegionsService:
tiled_encode: bool,
tiled_decode: bool,
global_prompt_weight: float,
noise_inversion: NoiseInversionOptions | None = None,
) -> DetailSEGSAsRegionsResult:
"""Run regional MultiDiffusion detailing for provided SEGS."""
@@ -234,6 +237,7 @@ class DetailSEGSAsRegionsService:
sampled_region=CropRegion(0, 0, image_width, image_height),
),
differential_diffusion=differential_diffusion,
noise_inversion=noise_inversion,
)
decoded = self._sampler.decode(vae, sampled, tiled_decode)
if decoded.shape[1:3] != image_tensor.shape[1:3]:
@@ -13,6 +13,7 @@ import torch
from ..domain.conditioning_batch import select_conditioning
from ..domain.detail_geometry import DetailScalePlan, build_detail_scale_plan
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.segs import Segment, coerce_segs
from ..domain.segs_mask_ops import (
crop_image,
@@ -59,6 +60,7 @@ class TiledDetailSamplingBoundary(Protocol):
latent_tile_batch_size: int,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample one latent crop with the requested tiled diffusion mode."""
@@ -131,6 +133,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
latent_tile_height: int,
latent_tile_overlap: int,
latent_tile_batch_size: int,
noise_inversion: NoiseInversionOptions | None = None,
) -> TiledDetailerResult:
"""Run crop sampling and composite-back detailing with tiled diffusion."""
@@ -198,6 +201,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
latent_tile_overlap=latent_tile_overlap,
latent_tile_batch_size=latent_tile_batch_size,
differential_diffusion=differential_diffusion,
noise_inversion=noise_inversion,
)
LOGGER.info(
@@ -246,6 +250,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
latent_tile_overlap: int,
latent_tile_batch_size: int,
differential_diffusion: bool,
noise_inversion: NoiseInversionOptions | None,
) -> torch.Tensor:
"""Detail one segment with tiled diffusion and return the updated image."""
@@ -281,6 +286,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
latent_tile_overlap=latent_tile_overlap,
latent_tile_batch_size=latent_tile_batch_size,
differential_diffusion=differential_diffusion,
noise_inversion=noise_inversion,
preview_context=DetailPreviewContext(
image=working_image,
work_region=segment.crop_region,
@@ -12,9 +12,11 @@ from typing import Any, TypeAlias
import torch
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
from ..runtime import sampling_samplers, sampling_schedulers
from ..domain.noise_inversion import NoiseInversionOptions
from ..runtime import sampling_noise, sampling_samplers, sampling_schedulers
from ..runtime.comfy_latent_normalization import COMFY_LATENT_NORMALIZER
from ..runtime.guided_sampling import sample_with_optional_negative
from ..runtime.inversion_model_factory import InversionModelFactory
from ..shared.logging import get_logger
Latent: TypeAlias = dict[str, Any]
@@ -37,8 +39,9 @@ class KSamplerSamplingService:
negative: Any,
latent_image: Latent,
denoise: float,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample a latent with configured SimpleSyrup sampler extensions."""
"""Sample full latents with optional inversion and per-item conditioning."""
sampler = sampling_samplers.resolve_sampler(sampler_name)
latent_samples = latent_image["samples"]
@@ -67,10 +70,13 @@ class KSamplerSamplingService:
denoise=denoise,
view=sampling_schedulers.SchedulerView.from_tensor(latent_samples),
).to(model.load_device)
noise = comfy_sample.prepare_noise(
latent_samples,
seed,
latent_image.get("batch_index"),
noise = sampling_noise.prepare_sampling_noise(
comfy_sample=comfy_sample,
sampler_name=sampler_name,
samples=latent_samples,
seed=seed,
batch_indices=latent_image.get("batch_index"),
model=model,
)
noise_mask = latent_image.get("noise_mask")
callback = import_module("latent_preview").prepare_callback(model, steps)
@@ -90,6 +96,7 @@ class KSamplerSamplingService:
callback=callback,
disable_pbar=disable_pbar,
seed=seed,
noise_inversion=noise_inversion,
)
else:
samples = sample_with_optional_negative(
@@ -106,6 +113,14 @@ class KSamplerSamplingService:
callback=callback,
disable_pbar=disable_pbar,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=InversionModelFactory(
model=model,
canvas_width=int(latent_samples.shape[-1]),
canvas_height=int(latent_samples.shape[-2]),
)
if noise_inversion is not None
else None,
)
if not isinstance(samples, torch.Tensor):
raise TypeError("KSampler output samples must be a torch.Tensor.")
@@ -144,6 +159,7 @@ class KSamplerSamplingService:
callback: Any,
disable_pbar: bool,
seed: int,
noise_inversion: NoiseInversionOptions | None,
) -> torch.Tensor:
"""Sample latent items with existing per-item batch selection."""
@@ -168,6 +184,14 @@ class KSamplerSamplingService:
callback=callback,
disable_pbar=disable_pbar,
seed=seed,
noise_inversion=noise_inversion,
inversion_model_factory=InversionModelFactory(
model=model,
canvas_width=int(latent_samples.shape[-1]),
canvas_height=int(latent_samples.shape[-2]),
)
if noise_inversion is not None
else None,
)
)
return torch.cat(sampled, dim=0)
@@ -48,6 +48,7 @@ from ..runtime.negpip.krea2 import (
from ..runtime.negpip.krea2 import (
WRAPPER_KEY as KREA_WRAPPER_KEY,
)
from ..runtime.negpip.krea2_host import krea2_host_mutations
from ..runtime.negpip.standard import (
encode_token_weights_negpip,
standard_attn2_negpip,
@@ -190,6 +191,7 @@ class NegpipModelService:
prepared_model = PATCHER_LIFECYCLE.derive_model(
model,
(
*krea2_host_mutations(model),
ModelCallableObjectPatchMutation(
"extra_conds",
krea2_extra_conds_negpip_wrapper(previous),
@@ -10,6 +10,7 @@ from typing import Any, TypeAlias
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import RegionalCapabilityAdmission
from ..domain.regional_tiled_diffusion import (
build_region_constrained_tiled_diffusion_plan,
@@ -49,6 +50,7 @@ class RegionalTiledDiffusionSamplingService:
preview_context: DetailPreviewContext | None,
differential_diffusion: bool,
capability_admission: RegionalCapabilityAdmission,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample every latent item with one shared regional composition."""
@@ -106,6 +108,13 @@ class RegionalTiledDiffusionSamplingService:
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
tiled_plan=plan,
noise_inversion=noise_inversion,
inversion_segs=(
segs_group[0 if len(segs_group) == 1 else index]
if segs_group
else None
),
inversion_region_masks=region_masks,
)
output_samples = output.get("samples")
if not isinstance(output_samples, torch.Tensor):
@@ -0,0 +1,236 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Compile capability configuration into one authoritative sampling execution."""
from __future__ import annotations
from typing import Any, TypedDict
import torch
from ..domain.attention_coupling_request import (
AttentionCouplingRequestMode,
classify_attention_coupling_request,
)
from ..domain.sampler_options import SamplerOptions, TilingOptions
from ..runtime import sampling_schedulers
from ..runtime.noise_inversion import validate_inversion_target
from ..runtime.tiled_sampling_validation import validate_sampling_controls
from ..shared.logging import get_logger
from .attention_coupling_sampling_service import AttentionCouplingSamplingService
from .contextual_attention_coupling_sampling_service import (
ContextualAttentionCouplingSamplingService,
)
from .contextual_diffusion_sampling_service import ContextualDiffusionSamplingService
from .ksampler_sampling_service import KSamplerSamplingService
from .tiled_attention_coupling_sampling_service import (
TiledAttentionCouplingSamplingService,
)
from .tiled_diffusion_sampling_service import TiledDiffusionSamplingService
LOGGER = get_logger(__name__)
class SamplingArguments(TypedDict):
"""Narrow controls while retaining dynamic host MODEL and tensor payloads."""
model: Any
seed: int
steps: int
cfg: float
sampler_name: str
scheduler: str
positive: object
negative: object
latent_image: dict[str, Any]
denoise: float
class ContextualArguments(TypedDict):
"""Describe global context and the sole local tile plan's execution controls."""
diffusion_mode: str
latent_context_size: int
latent_context_overlap: int
latent_context_batch_size: int
global_weight: float
global_steps: int
global_decay: float
class TiledArguments(TypedDict):
"""Describe local geometry and mask-dependent denoising policy."""
diffusion_mode: str
latent_tile_width: int
latent_tile_height: int
latent_tile_overlap: int
latent_tile_batch_size: int
differential_diffusion: bool
class SamplerOptionsSamplingService:
"""Route configuration independently of connection order or node placement."""
def sample(
self,
*,
model: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
positive: object,
negative: object,
latent_image: dict[str, Any],
denoise: float,
options: SamplerOptions | None = None,
segs: object | None = None,
region_masks: object | None = None,
) -> dict[str, Any]:
"""Admit connected region data only through its enabled sampling capability."""
if options is not None and not isinstance(options, SamplerOptions):
raise TypeError(
"KSampler options must come from SimpleSyrup options nodes."
)
configured = options if options is not None else SamplerOptions()
context = configured.contextual_diffusion
tiling = configured.tiling
if context is not None:
if tiling is not None:
LOGGER.warning(
"Contextual Diffusion takes precedence; Tiling Options ignored.",
extra={"node_id": "SimpleSyrup.KSampler"},
)
tiling = context.local_tiling()
arguments: SamplingArguments = {
"model": model,
"seed": seed,
"steps": steps,
"cfg": cfg,
"sampler_name": sampler_name,
"scheduler": scheduler,
"positive": positive,
"negative": negative,
"latent_image": latent_image,
"denoise": denoise,
}
self._preflight(arguments, configured, tiling)
attention = configured.attention_coupling
if (
attention is not None
and classify_attention_coupling_request(
positive=positive, negative=negative, region_masks=region_masks
)
is AttentionCouplingRequestMode.BYPASS
):
attention = None
inversion = configured.noise_inversion
if context is not None:
assert tiling is not None
contextual_arguments: ContextualArguments = {
"diffusion_mode": tiling.diffusion_mode,
"latent_context_size": context.context_size,
"latent_context_overlap": tiling.overlap,
"latent_context_batch_size": tiling.batch_size,
"global_weight": context.global_weight,
"global_steps": context.global_steps,
"global_decay": context.global_decay,
}
if attention is not None:
result = ContextualAttentionCouplingSamplingService().sample(
**arguments,
**contextual_arguments,
region_masks=region_masks,
regional_prompt_weight=attention.regional_prompt_weight,
region_mask_feather=attention.region_mask_feather,
segs=segs,
tiling=tiling,
noise_inversion=inversion,
)
else:
result = ContextualDiffusionSamplingService().sample(
**arguments,
**contextual_arguments,
segs=segs,
tiling=tiling,
noise_inversion=inversion,
)
return result.latent
if tiling is not None:
tiled_arguments: TiledArguments = {
"diffusion_mode": tiling.diffusion_mode,
"latent_tile_width": tiling.width,
"latent_tile_height": tiling.height,
"latent_tile_overlap": tiling.overlap,
"latent_tile_batch_size": tiling.batch_size,
"differential_diffusion": tiling.differential_diffusion,
}
if attention is not None:
return TiledAttentionCouplingSamplingService().sample(
**arguments,
**tiled_arguments,
region_masks=region_masks,
regional_prompt_weight=attention.regional_prompt_weight,
region_mask_feather=attention.region_mask_feather,
segs=segs,
noise_inversion=inversion,
)
return TiledDiffusionSamplingService().sample(
**arguments,
**tiled_arguments,
segs=segs,
noise_inversion=inversion,
)
if attention is not None:
return AttentionCouplingSamplingService().sample(
**arguments,
region_masks=region_masks,
regional_prompt_weight=attention.regional_prompt_weight,
region_mask_feather=attention.region_mask_feather,
noise_inversion=inversion,
)
return KSamplerSamplingService().sample(**arguments, noise_inversion=inversion)
def _preflight(
self,
arguments: SamplingArguments,
options: SamplerOptions,
tiling: TilingOptions | None,
) -> None:
"""Reject unsupported schedules and endpoints before preparing models."""
samples = arguments["latent_image"].get("samples")
if not isinstance(samples, torch.Tensor):
raise TypeError("KSampler latent samples must be a torch.Tensor.")
validate_sampling_controls(
steps=arguments["steps"],
denoise=arguments["denoise"],
latent_tile_width=tiling.width if tiling is not None else 16,
latent_tile_height=tiling.height if tiling is not None else 16,
latent_tile_batch_size=tiling.batch_size if tiling is not None else 1,
)
incompatible_unipc = options.contextual_diffusion is not None or (
tiling is not None and tiling.diffusion_mode == "multidiffusion"
)
if incompatible_unipc and arguments["sampler_name"] in {"uni_pc", "uni_pc_bh2"}:
raise ValueError(
"Tiling and Contextual Diffusion do not support UniPC samplers."
)
if options.noise_inversion is not None:
view = (
sampling_schedulers.SchedulerView(tiling.width, tiling.height)
if tiling is not None
else sampling_schedulers.SchedulerView.from_tensor(samples)
)
sigmas = sampling_schedulers.calculate_sigmas(
model=arguments["model"],
scheduler_name=arguments["scheduler"],
sampler_name=arguments["sampler_name"],
steps=arguments["steps"],
denoise=arguments["denoise"],
view=view,
)
validate_inversion_target(arguments["model"], sigmas)
@@ -11,6 +11,7 @@ from typing import Any, TypeAlias
import torch
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import RegionalCapabilityAdmission
from ..domain.segs import coerce_segs_group
from ..domain.segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan
@@ -47,6 +48,7 @@ class SEGSGuidedTiledDiffusionSamplingService:
preview_context: DetailPreviewContext | None,
differential_diffusion: bool,
capability_admission: RegionalCapabilityAdmission,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample every latent batch item using its connected SEGS guide."""
@@ -108,6 +110,8 @@ class SEGSGuidedTiledDiffusionSamplingService:
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
tiled_plan=plan,
noise_inversion=noise_inversion,
inversion_segs=segs_for_item,
)
output_samples = output["samples"]
if not isinstance(output_samples, torch.Tensor):
@@ -12,6 +12,7 @@ from ..domain.attention_coupling_request import (
AttentionCouplingRequestMode,
classify_attention_coupling_request,
)
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_attention_execution import RegionalAttentionExecutionMode
from ..domain.regional_features import (
EMPTY_REGIONAL_FEATURE_REQUEST,
@@ -62,6 +63,8 @@ class TiledAttentionCouplingSamplingService:
latent_tile_batch_size: int,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
segs: object | None = None,
noise_inversion: NoiseInversionOptions | None = None,
) -> dict[str, Any]:
"""Bypass ordinary requests or prepare one complete regional request."""
@@ -90,6 +93,8 @@ class TiledAttentionCouplingSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
feature_request=EMPTY_REGIONAL_FEATURE_REQUEST,
segs=segs,
noise_inversion=noise_inversion,
)
prepared = self.model_preparation_service_class().prepare(
@@ -121,6 +126,8 @@ class TiledAttentionCouplingSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
feature_request=_TILED_ATTENTION_REQUEST,
segs=segs,
noise_inversion=noise_inversion,
)
def _sample_tiled(
@@ -144,6 +151,8 @@ class TiledAttentionCouplingSamplingService:
preview_context: DetailPreviewContext | None,
differential_diffusion: bool,
feature_request: RegionalFeatureRequest,
segs: object | None,
noise_inversion: NoiseInversionOptions | None,
) -> dict[str, Any]:
"""Delegate one ordinary or Attention Coupling tiled request."""
@@ -166,10 +175,11 @@ class TiledAttentionCouplingSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
feature_request=feature_request,
segs=None,
segs=segs,
region_masks=None,
regional_prompt_weight=0.5,
region_mask_feather=0,
noise_inversion=noise_inversion,
)
@@ -10,6 +10,7 @@ from typing import Any, Protocol
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..runtime.detail_previews import DetailPreviewContext
from ..runtime.detail_sampling import DetailSampler, Latent
from .tiled_diffusion_sampling_service import TiledDiffusionSamplingService
@@ -38,6 +39,7 @@ class TiledDiffusionLatentSamplingBoundary(Protocol):
latent_tile_batch_size: int,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample a latent using the selected tiled diffusion mode."""
@@ -87,6 +89,7 @@ class TiledDetailSampler:
latent_tile_batch_size: int,
preview_context: DetailPreviewContext | None = None,
differential_diffusion: bool = False,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample one latent crop with the selected tiled diffusion runtime."""
@@ -108,4 +111,5 @@ class TiledDetailSampler:
latent_tile_batch_size=latent_tile_batch_size,
preview_context=preview_context,
differential_diffusion=differential_diffusion,
noise_inversion=noise_inversion,
)
@@ -11,6 +11,7 @@ from typing import Any, TypeAlias
import torch
from ..domain.conditioning_batch import select_conditioning
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import RegionalCapabilityAdmission
from ..runtime.detail_previews import DetailPreviewContext
from .sampling_batch import combine_latent_outputs, single_item_latent
@@ -44,6 +45,7 @@ class TiledDiffusionConditioningBatchService:
preview_context: DetailPreviewContext | None,
differential_diffusion: bool,
capability_admission: RegionalCapabilityAdmission,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample each latent item with its selected conditioning values."""
@@ -72,6 +74,7 @@ class TiledDiffusionConditioningBatchService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
)
output_samples = output.get("samples")
if not isinstance(output_samples, torch.Tensor):
@@ -8,7 +8,11 @@ from __future__ import annotations
from typing import Any, Protocol, TypeAlias
import torch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import RegionalCapabilityAdmission
from ..domain.segs import NativeSegs
from ..domain.tiled_diffusion import TiledDiffusionPlan
from ..runtime import mixture_of_diffusers_sampling, multidiffusion_sampling
from ..runtime.detail_previews import DetailPreviewContext
@@ -41,6 +45,9 @@ class TiledDiffusionItemSampler(Protocol):
differential_diffusion: bool,
capability_admission: RegionalCapabilityAdmission,
tiled_plan: TiledDiffusionPlan | None = None,
noise_inversion: NoiseInversionOptions | None = None,
inversion_segs: NativeSegs | None = None,
inversion_region_masks: torch.Tensor | None = None,
) -> Latent:
"""Return one sampled latent item."""
@@ -70,6 +77,9 @@ class TiledDiffusionItemSamplingService:
differential_diffusion: bool,
capability_admission: RegionalCapabilityAdmission,
tiled_plan: TiledDiffusionPlan | None = None,
noise_inversion: NoiseInversionOptions | None = None,
inversion_segs: NativeSegs | None = None,
inversion_region_masks: torch.Tensor | None = None,
) -> Latent:
"""Invoke exactly one runtime with the unchanged sampling request."""
@@ -97,4 +107,7 @@ class TiledDiffusionItemSamplingService:
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
tiled_plan=tiled_plan,
noise_inversion=noise_inversion,
inversion_segs=inversion_segs,
inversion_region_masks=inversion_region_masks,
)
@@ -9,6 +9,7 @@ from __future__ import annotations
from typing import Any, ClassVar
from ..domain.conditioning_batch import ConditioningBatch
from ..domain.noise_inversion import NoiseInversionOptions
from ..domain.regional_features import (
EMPTY_REGIONAL_FEATURE_REQUEST,
TILED_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES,
@@ -84,6 +85,7 @@ class TiledDiffusionSamplingService:
region_masks: object | None = None,
regional_prompt_weight: float = 0.5,
region_mask_feather: int = 0,
noise_inversion: NoiseInversionOptions | None = None,
) -> Latent:
"""Sample a latent with the selected tiled diffusion method."""
@@ -129,6 +131,7 @@ class TiledDiffusionSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
)
if segs is not None:
return self.segs_sampling_service_class().sample(
@@ -151,6 +154,7 @@ class TiledDiffusionSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
segs=segs,
)
if isinstance(positive, ConditioningBatch) or isinstance(
@@ -177,6 +181,7 @@ class TiledDiffusionSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
)
return self.item_sampling_service_class().sample(
diffusion_mode=diffusion_mode,
@@ -197,4 +202,5 @@ class TiledDiffusionSamplingService:
preview_context=preview_context,
differential_diffusion=differential_diffusion,
capability_admission=capability_admission,
noise_inversion=noise_inversion,
)
+1
View File
@@ -0,0 +1 @@
"""Preserve pinned RES4LYF solver code for local sampler execution."""
@@ -0,0 +1 @@
"""Contain the pinned RES4LYF Runge-Kutta solver implementation."""
@@ -0,0 +1,60 @@
MAX_STEPS = 10000
IMPLICIT_TYPE_NAMES = [
"rebound",
"retro-eta",
"bongmath",
"predictor-corrector",
]
GUIDE_MODE_NAMES_SIMPLE = [
"flow",
"sync",
"lure",
"data",
"epsilon",
"inversion",
"pseudoimplicit",
"fully_pseudoimplicit",
]
GUIDE_MODE_NAMES_SELF_REFINE = [
"self_refine_epsilon",
"self_refine_pseudoimplicit",
]
FRAME_WEIGHTS_CONFIG_NAMES = [
"frame_weights",
"frame_weights_inv",
"frame_targets"
]
FRAME_WEIGHTS_DYNAMICS_NAMES = [
"constant",
"linear",
"ease_out",
"ease_in",
"middle",
"trough",
]
FRAME_WEIGHTS_SCHEDULE_NAMES = [
"moderate_early",
"moderate_late",
"fast_early",
"fast_late",
"slow_early",
"slow_late",
]
GUIDE_MODE_NAMES_PSEUDOIMPLICIT = [
"pseudoimplicit",
"pseudoimplicit_cw",
"pseudoimplicit_projection",
"pseudoimplicit_projection_cw",
"fully_pseudoimplicit",
"fully_pseudoimplicit_projection",
"fully_pseudoimplicit_cw",
"fully_pseudoimplicit_projection_cw"
]
@@ -0,0 +1,123 @@
# Adapted from: https://github.com/zju-pi/diff-sampler/blob/main/gits-main/solver_utils.py
# fixed the calcs for "rhoab" which suffered from an off-by-one error and made some other minor corrections
import torch
import numpy as np
# A pytorch reimplementation of DEIS (https://github.com/qsh-zh/deis).
#############################
### Utils for DEIS solver ###
#############################
#----------------------------------------------------------------------------
# Transfer from the input time (sigma) used in EDM to that (t) used in DEIS.
def edm2t(edm_steps, epsilon_s=1e-3, sigma_min=0.002, sigma_max=80):
vp_sigma = lambda beta_d, beta_min: lambda t: (np.e ** (0.5 * beta_d * (t ** 2) + beta_min * t) - 1) ** 0.5
vp_sigma_inv = lambda beta_d, beta_min: lambda sigma: ((beta_min ** 2 + 2 * beta_d * (sigma ** 2 + 1).log()).sqrt() - beta_min) / beta_d
vp_beta_d = 2 * (np.log(torch.tensor(sigma_min).cpu() ** 2 + 1) / epsilon_s - np.log(torch.tensor(sigma_max).cpu() ** 2 + 1)) / (epsilon_s - 1)
vp_beta_min = np.log(torch.tensor(sigma_max).cpu() ** 2 + 1) - 0.5 * vp_beta_d
t_steps = vp_sigma_inv(vp_beta_d.clone().detach().cpu(), vp_beta_min.clone().detach().cpu())(edm_steps.clone().detach().cpu())
return t_steps, vp_beta_min, vp_beta_d + vp_beta_min
#----------------------------------------------------------------------------
def cal_poly(prev_t, j, taus):
poly = 1
for k in range(prev_t.shape[0]):
if k == j:
continue
poly *= (taus - prev_t[k]) / (prev_t[j] - prev_t[k])
return poly
#----------------------------------------------------------------------------
# Transfer from t to alpha_t.
def t2alpha_fn(beta_0, beta_1, t):
return torch.exp(-0.5 * t ** 2 * (beta_1 - beta_0) - t * beta_0)
#----------------------------------------------------------------------------
def cal_integrand(beta_0, beta_1, taus):
with torch.inference_mode(mode=False):
taus = taus.clone()
beta_0 = beta_0.clone()
beta_1 = beta_1.clone()
with torch.enable_grad():
taus.requires_grad_(True)
alpha = t2alpha_fn(beta_0, beta_1, taus)
log_alpha = alpha.log()
log_alpha.sum().backward()
d_log_alpha_dtau = taus.grad
integrand = -0.5 * d_log_alpha_dtau / torch.sqrt(alpha * (1 - alpha))
return integrand
#----------------------------------------------------------------------------
def get_deis_coeff_list(t_steps, max_order, N=10000, deis_mode='tab'):
"""
Get the coefficient list for DEIS sampling.
Args:
t_steps: A pytorch tensor. The time steps for sampling.
max_order: A `int`. Maximum order of the solver. 1 <= max_order <= 4
N: A `int`. Use how many points to perform the numerical integration when deis_mode=='tab'.
deis_mode: A `str`. Select between 'tab' and 'rhoab'. Type of DEIS.
Returns:
A pytorch tensor. A batch of generated samples or sampling trajectories if return_inters=True.
"""
if deis_mode == 'tab':
t_steps, beta_0, beta_1 = edm2t(t_steps)
C = []
for i, (t_cur, t_next) in enumerate(zip(t_steps[:-1], t_steps[1:])):
order = min(i+1, max_order)
if order == 1:
C.append([])
else:
taus = torch.linspace(t_cur, t_next, N) # split the interval for integral approximation
dtau = (t_next - t_cur) / N
prev_t = t_steps[[i - k for k in range(order)]]
coeff_temp = []
integrand = cal_integrand(beta_0, beta_1, taus)
for j in range(order):
poly = cal_poly(prev_t, j, taus)
coeff_temp.append(torch.sum(integrand * poly) * dtau)
C.append(coeff_temp)
elif deis_mode == 'rhoab':
# Analytical solution, second order
def get_def_integral_2(a, b, start, end, c):
coeff = (end**3 - start**3) / 3 - (end**2 - start**2) * (a + b) / 2 + (end - start) * a * b
return coeff / ((c - a) * (c - b))
# Analytical solution, third order
def get_def_integral_3(a, b, c, start, end, d):
coeff = (end**4 - start**4) / 4 - (end**3 - start**3) * (a + b + c) / 3 \
+ (end**2 - start**2) * (a*b + a*c + b*c) / 2 - (end - start) * a * b * c
return coeff / ((d - a) * (d - b) * (d - c))
C = []
for i, (t_cur, t_next) in enumerate(zip(t_steps[:-1], t_steps[1:])):
order = min(i+1, max_order) #fixed order calcs
if order == 1:
C.append([])
else:
prev_t = t_steps[[i - k for k in range(order+1)]]
if order == 2:
coeff_cur = ((t_next - prev_t[1])**2 - (t_cur - prev_t[1])**2) / (2 * (t_cur - prev_t[1]))
coeff_prev1 = (t_next - t_cur)**2 / (2 * (prev_t[1] - t_cur))
coeff_temp = [coeff_cur, coeff_prev1]
elif order == 3:
coeff_cur = get_def_integral_2(prev_t[1], prev_t[2], t_cur, t_next, t_cur)
coeff_prev1 = get_def_integral_2(t_cur, prev_t[2], t_cur, t_next, prev_t[1])
coeff_prev2 = get_def_integral_2(t_cur, prev_t[1], t_cur, t_next, prev_t[2])
coeff_temp = [coeff_cur, coeff_prev1, coeff_prev2]
elif order == 4:
coeff_cur = get_def_integral_3(prev_t[1], prev_t[2], prev_t[3], t_cur, t_next, t_cur)
coeff_prev1 = get_def_integral_3(t_cur, prev_t[2], prev_t[3], t_cur, t_next, prev_t[1])
coeff_prev2 = get_def_integral_3(t_cur, prev_t[1], prev_t[3], t_cur, t_next, prev_t[2])
coeff_prev3 = get_def_integral_3(t_cur, prev_t[1], prev_t[2], t_cur, t_next, prev_t[3])
coeff_temp = [coeff_cur, coeff_prev1, coeff_prev2, coeff_prev3]
C.append(coeff_temp)
return C
@@ -0,0 +1,717 @@
import torch
import torch.nn.functional as F
from torch import nn, Tensor, Generator, lerp
from torch.nn.functional import unfold
from torch.distributions import StudentT, Laplace
import numpy as np
import pywt
import functools
from typing import Callable, Tuple
from math import pi
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
from ..res4lyf import RESplain
# Set this to "True" if you have installed OpenSimplex. Recommended to install without dependencies due to conflicting packages: pip3 install opensimplex --no-deps
OPENSIMPLEX_ENABLE = False
if OPENSIMPLEX_ENABLE:
from opensimplex import OpenSimplex
class PrecisionTool:
def __init__(self, cast_type='fp64'):
self.cast_type = cast_type
def cast_tensor(self, func):
@functools.wraps(func)
def wrapper(*args, **kwargs):
if self.cast_type not in ['fp64', 'fp32', 'fp16']:
return func(*args, **kwargs)
target_device = None
for arg in args:
if torch.is_tensor(arg):
target_device = arg.device
break
if target_device is None:
for v in kwargs.values():
if torch.is_tensor(v):
target_device = v.device
break
# recursively zs_recast tensors in nested dictionaries
def cast_and_move_to_device(data):
if torch.is_tensor(data):
if self.cast_type == 'fp64':
return data.to(torch.float64).to(target_device)
elif self.cast_type == 'fp32':
return data.to(torch.float32).to(target_device)
elif self.cast_type == 'fp16':
return data.to(torch.float16).to(target_device)
elif isinstance(data, dict):
return {k: cast_and_move_to_device(v) for k, v in data.items()}
return data
new_args = [cast_and_move_to_device(arg) for arg in args]
new_kwargs = {k: cast_and_move_to_device(v) for k, v in kwargs.items()}
return func(*new_args, **new_kwargs)
return wrapper
def set_cast_type(self, new_value):
if new_value in ['fp64', 'fp32', 'fp16']:
self.cast_type = new_value
else:
self.cast_type = 'fp64'
precision_tool = PrecisionTool(cast_type='fp64')
def noise_generator_factory(cls, **fixed_params):
def create_instance(**kwargs):
params = {**fixed_params, **kwargs}
return cls(**params)
return create_instance
def like(x):
return {'size': x.shape, 'dtype': x.dtype, 'layout': x.layout, 'device': x.device}
def scale_to_range(x, scaled_min = -1.73, scaled_max = 1.73): #1.73 is roughly the square root of 3
return scaled_min + (x - x.min()) * (scaled_max - scaled_min) / (x.max() - x.min())
def normalize(x):
return (x - x.mean())/ x.std()
def per_frame(noise_4d, size):
if len(size) == 5:
b, c, t, h, w = size
return torch.stack([noise_4d((b, c, h, w)) for _ in range(t)], dim=2)
return noise_4d(size)
class NoiseGenerator:
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None):
self.seed = seed
if x is not None:
self.x = x
self.size = x.shape
self.dtype = x.dtype
self.layout = x.layout
self.device = x.device
else:
self.x = torch.zeros(size, dtype=dtype, layout=layout, device=device)
# allow overriding parameters imported from latent 'x' if specified
if size is not None:
self.size = size
if dtype is not None:
self.dtype = dtype
if layout is not None:
self.layout = layout
if device is not None:
self.device = device
# Treat 1D latents as single-row images for 4D noise generators
self.out_size = tuple(self.size)
if len(self.size) == 3:
self.size = (self.size[0], self.size[1], 1, self.size[2])
self.sigma_max = sigma_max.to(device) if isinstance(sigma_max, torch.Tensor) else sigma_max
self.sigma_min = sigma_min.to(device) if isinstance(sigma_min, torch.Tensor) else sigma_min
self.last_seed = seed #- 1 #adapt for update being called during initialization, which increments last_seed
if generator is None:
self.generator = torch.Generator(device=self.device).manual_seed(seed)
else:
self.generator = generator
def __call__(self, **kwargs):
return self.generate(**kwargs).reshape(self.out_size)
def generate(self, **kwargs):
raise NotImplementedError("This method got clownsharked!")
def update(self, **kwargs):
#if not isinstance(self, BrownianNoiseGenerator):
# self.last_seed += 1
updated_values = []
for attribute_name, value in kwargs.items():
if value is not None:
setattr(self, attribute_name, value)
updated_values.append(getattr(self, attribute_name))
return tuple(updated_values)
class BrownianNoiseGenerator(NoiseGenerator):
def generate(self, *, sigma=None, sigma_next=None, **kwargs):
return BrownianTreeNoiseSampler(self.x, self.sigma_min, self.sigma_max, seed=self.seed, cpu = self.device.type=='cpu')(sigma, sigma_next)
class FractalNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
alpha=0.0, k=1.0, scale=0.1):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(alpha=alpha, k=k, scale=scale)
def generate(self, *, alpha=None, k=None, scale=None, **kwargs):
self.update(alpha=alpha, k=k, scale=scale)
self.last_seed += 1
if len(self.size) == 5:
b, c, t, h, w = self.size
else:
b, c, h, w = self.size
noise = torch.normal(mean=0.0, std=1.0, size=self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
y_freq = torch.fft.fftfreq(h, 1/h, device=self.device)
x_freq = torch.fft.fftfreq(w, 1/w, device=self.device)
if len(self.size) == 5:
t_freq = torch.fft.fftfreq(t, 1/t, device=self.device)
freq = torch.sqrt(t_freq[:, None, None]**2 + y_freq[None, :, None]**2 + x_freq[None, None, :]**2).clamp(min=1e-10)
else:
freq = torch.sqrt(y_freq[:, None]**2 + x_freq[None, :]**2).clamp(min=1e-10)
spectral_density = self.k / torch.pow(freq, self.alpha * self.scale)
spectral_density[0, 0] = 0
noise_fft = torch.fft.fftn(noise)
modified_fft = noise_fft * spectral_density
noise = torch.fft.ifftn(modified_fft).real
return noise / torch.std(noise)
class SimplexNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
scale=0.01):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.noise = OpenSimplex(seed=seed)
self.scale = scale
def generate(self, *, scale=None, **kwargs):
self.update(scale=scale)
self.last_seed += 1
if len(self.size) == 5:
b, c, t, h, w = self.size
else:
b, c, h, w = self.size
noise_array = self.noise.noise3array(np.arange(w),np.arange(h),np.arange(c))
self.noise = OpenSimplex(seed=self.noise.get_seed()+1)
noise_tensor = torch.from_numpy(noise_array).to(self.device)
noise_tensor = torch.unsqueeze(noise_tensor, dim=0)
if len(self.size) == 5:
noise_tensor = torch.unsqueeze(noise_tensor, dim=0)
return noise_tensor / noise_tensor.std()
#return normalize(scale_to_range(noise_tensor))
class HiresPyramidNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
discount=0.7, mode='nearest-exact'):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(discount=discount, mode=mode)
def generate(self, *, discount=None, mode=None, **kwargs):
self.update(discount=discount, mode=mode)
self.last_seed += 1
return per_frame(self._noise_4d, self.size)
def _noise_4d(self, size):
b, c, h, w = size
orig_h, orig_w = h, w
u = nn.Upsample(size=(orig_h, orig_w), mode=self.mode).to(self.device)
noise = ((torch.rand(size=size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator) - 0.5) * 2 * 1.73)
for i in range(4):
r = torch.rand(1, device=self.device, generator=self.generator).item() * 2 + 2
h, w = min(orig_h * 15, int(h * (r ** i))), min(orig_w * 15, int(w * (r ** i)))
new_noise = torch.randn((b, c, h, w), dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
upsampled_noise = u(new_noise)
noise += upsampled_noise * self.discount ** i
if h >= orig_h * 15 or w >= orig_w * 15:
break # if resolution is too high
return noise / noise.std()
class PyramidNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
discount=0.8, mode='nearest-exact'):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(discount=discount, mode=mode)
def generate(self, *, discount=None, mode=None, **kwargs):
self.update(discount=discount, mode=mode)
self.last_seed += 1
return per_frame(self._noise_4d, self.size)
def _noise_4d(self, size):
x = torch.zeros(size, dtype=self.dtype, layout=self.layout, device=self.device)
b, c, h, w = size
orig_h, orig_w = h, w
r = 1
for i in range(5):
r *= 2
scaledSize = (b, c, h * r, w * r)
origSize = (orig_h, orig_w)
x += torch.nn.functional.interpolate(
torch.normal(mean=0, std=0.5 ** i, size=scaledSize, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator),
size=origSize, mode=self.mode
) * self.discount ** i
return x / x.std()
class InterpolatedPyramidNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
discount=0.7, mode='nearest-exact'):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(discount=discount, mode=mode)
def generate(self, *, discount=None, mode=None, **kwargs):
self.update(discount=discount, mode=mode)
self.last_seed += 1
return per_frame(self._noise_4d, self.size)
def _noise_4d(self, size):
b, c, h, w = size
orig_h, orig_w = h, w
noise = ((torch.rand(size=size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator) - 0.5) * 2 * 1.73)
multipliers = [1]
for i in range(4):
r = torch.rand(1, device=self.device, generator=self.generator).item() * 2 + 2
h, w = min(orig_h * 15, int(h * (r ** i))), min(orig_w * 15, int(w * (r ** i)))
new_noise = torch.randn((b, c, h, w), dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
upsampled_noise = nn.functional.interpolate(new_noise, size=(orig_h, orig_w), mode=self.mode)
noise += upsampled_noise * self.discount ** i
multipliers.append( self.discount ** i)
if h >= orig_h * 15 or w >= orig_w * 15:
break # if resolution is too high
noise = noise / sum([m ** 2 for m in multipliers]) ** 0.5
return noise / noise.std()
class CascadeBPyramidNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
levels=10, mode='nearest', size_range=[1,16]):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(levels=levels, mode=mode, size_range=size_range)
def generate(self, *, levels=10, mode='nearest', size_range=[1,16], **kwargs):
self.update(levels=levels, mode=mode)
self.last_seed += 1
return per_frame(lambda size: self._noise_4d(size, size_range), self.size)
def _noise_4d(self, size, size_range):
epsilon = torch.randn(size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
multipliers = [1]
for i in range(1, self.levels):
m = 0.75 ** i
h, w = int(epsilon.size(-2) // (2 ** i)), int(epsilon.size(-1) // (2 ** i))
if size_range is None or (size_range[0] <= h <= size_range[1] or size_range[0] <= w <= size_range[1]):
offset = torch.randn(epsilon.size(0), epsilon.size(1), h, w, device=self.device, generator=self.generator)
epsilon = epsilon + torch.nn.functional.interpolate(offset, size=epsilon.shape[-2:], mode=self.mode) * m
multipliers.append(m)
if h <= 1 or w <= 1:
break
epsilon = epsilon / sum([m ** 2 for m in multipliers]) ** 0.5 #divides the epsilon tensor by the square root of the sum of the squared multipliers.
return epsilon
class UniformNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
mean=0.0, scale=1.73):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(mean=mean, scale=scale)
def generate(self, *, mean=None, scale=None, **kwargs):
self.update(mean=mean, scale=scale)
self.last_seed += 1
noise = torch.rand(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
return self.scale * 2 * (noise - 0.5) + self.mean
class GaussianNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
mean=0.0, std=1.0):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(mean=mean, std=std)
def generate(self, *, mean=None, std=None, **kwargs):
self.update(mean=mean, std=std)
self.last_seed += 1
noise = torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
return (noise - noise.mean()) / noise.std()
class GaussianBackwardsNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
mean=0.0, std=1.0):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(mean=mean, std=std)
def generate(self, *, mean=None, std=None, **kwargs):
self.update(mean=mean, std=std)
self.last_seed += 1
RESplain("GaussianBackwards last seed:", self.generator.initial_seed())
self.generator.manual_seed(self.generator.initial_seed() - 1)
noise = torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator)
return (noise - noise.mean()) / noise.std()
class LaplacianNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
loc=0, scale=1.0):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(loc=loc, scale=scale)
def generate(self, *, loc=None, scale=None, **kwargs):
self.update(loc=loc, scale=scale)
self.last_seed += 1
# b, c, h, w = self.size
# orig_h, orig_w = h, w
noise = torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator) / 4.0
rng_state = torch.random.get_rng_state()
torch.manual_seed(self.generator.initial_seed())
laplacian_noise = Laplace(loc=self.loc, scale=self.scale).rsample(self.size).to(self.device)
self.generator.manual_seed(self.generator.initial_seed() + 1)
torch.random.set_rng_state(rng_state)
noise += laplacian_noise
return noise / noise.std()
class StudentTNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
loc=0, scale=0.2, df=1):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(loc=loc, scale=scale, df=df)
def generate(self, *, loc=None, scale=None, df=None, **kwargs):
self.update(loc=loc, scale=scale, df=df)
self.last_seed += 1
# b, c, h, w = self.size
# orig_h, orig_w = h, w
rng_state = torch.random.get_rng_state()
torch.manual_seed(self.generator.initial_seed())
noise = StudentT(loc=self.loc, scale=self.scale, df=self.df).rsample(self.size)
if not isinstance(self, BrownianNoiseGenerator):
self.last_seed += 1
s = torch.quantile(noise.flatten(start_dim=1).abs(), 0.75, dim=-1)
s = s.reshape(-1, *([1] * (noise.dim() - 1)))
noise = noise.clamp(-s, s)
noise_latent = torch.copysign(torch.pow(torch.abs(noise), 0.5), noise).to(self.device)
self.generator.manual_seed(self.generator.initial_seed() + 1)
torch.random.set_rng_state(rng_state)
return (noise_latent - noise_latent.mean()) / noise_latent.std()
class WaveletNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
wavelet='haar'):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(wavelet=wavelet)
def generate(self, *, wavelet=None, **kwargs):
self.update(wavelet=wavelet)
self.last_seed += 1
# b, c, h, w = self.size
# orig_h, orig_w = h, w
# noise for spatial dimensions only
coeffs = pywt.wavedecn(torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator).to('cpu'), wavelet=self.wavelet, mode='periodization')
noise = pywt.waverecn(coeffs, wavelet=self.wavelet, mode='periodization')
noise_tensor = torch.tensor(noise, dtype=self.dtype, device=self.device)
noise_tensor = (noise_tensor - noise_tensor.mean()) / noise_tensor.std()
return noise_tensor
class PerlinNoiseGenerator(NoiseGenerator):
def __init__(self, x=None, size=None, dtype=None, layout=None, device=None, seed=42, generator=None, sigma_min=None, sigma_max=None,
detail=0.0):
super().__init__(x, size, dtype, layout, device, seed, generator, sigma_min, sigma_max)
self.update(detail=detail)
@staticmethod
def get_positions(block_shape: Tuple[int, int]) -> Tensor:
bh, bw = block_shape
positions = torch.stack(
torch.meshgrid(
[(torch.arange(b) + 0.5) / b for b in (bw, bh)],
indexing="xy",
),
-1,
).view(1, bh, bw, 1, 1, 2)
return positions
@staticmethod
def unfold_grid(vectors: Tensor) -> Tensor:
batch_size, _, gpy, gpx = vectors.shape
return (
unfold(vectors, (2, 2))
.view(batch_size, 2, 4, -1)
.permute(0, 2, 3, 1)
.view(batch_size, 4, gpy - 1, gpx - 1, 2)
)
@staticmethod
def smooth_step(t: Tensor) -> Tensor:
return t * t * (3.0 - 2.0 * t)
@staticmethod
def perlin_noise_tensor(
self,
vectors: Tensor, positions: Tensor, step: Callable = None
) -> Tensor:
if step is None:
step = self.smooth_step
batch_size = vectors.shape[0]
# grid height, grid width
gh, gw = vectors.shape[2:4]
# block height, block width
bh, bw = positions.shape[1:3]
for i in range(2):
if positions.shape[i + 3] not in (1, vectors.shape[i + 2]):
raise Exception(
f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})"
)
if positions.shape[0] not in (1, batch_size):
raise Exception(
f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})"
)
vectors = vectors.view(batch_size, 4, 1, gh * gw, 2)
positions = positions.view(positions.shape[0], bh * bw, -1, 2)
step_x = step(positions[..., 0])
step_y = step(positions[..., 1])
row0 = lerp(
(vectors[:, 0] * positions).sum(dim=-1),
(vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1),
step_x,
)
row1 = lerp(
(vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1),
(vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1),
step_x,
)
noise = lerp(row0, row1, step_y)
return (
noise.view(
batch_size,
bh,
bw,
gh,
gw,
)
.permute(0, 3, 1, 4, 2)
.reshape(batch_size, gh * bh, gw * bw)
)
def perlin_noise(
self,
grid_shape: Tuple[int, int],
out_shape: Tuple[int, int],
batch_size: int = 1,
generator: Generator = None,
*args,
**kwargs,
) -> Tensor:
gh, gw = grid_shape # grid height and width
oh, ow = out_shape # output height and width
bh, bw = oh // gh, ow // gw # block height and width
if oh != bh * gh:
raise Exception(f"Output height {oh} must be divisible by grid height {gh}")
if ow != bw * gw != 0:
raise Exception(f"Output width {ow} must be divisible by grid width {gw}")
angle = torch.empty(
[batch_size] + [s + 1 for s in grid_shape], device=self.device, *args, **kwargs
).uniform_(to=2.0 * pi, generator=self.generator)
# random vectors on grid points
vectors = self.unfold_grid(torch.stack((torch.cos(angle), torch.sin(angle)), dim=1))
# positions inside grid cells [0, 1)
positions = self.get_positions((bh, bw)).to(vectors)
return self.perlin_noise_tensor(self, vectors, positions).squeeze(0)
def generate(self, *, detail=None, **kwargs):
self.update(detail=detail) #currently unused
self.last_seed += 1
if len(self.size) == 5:
b, c, t, h, w = self.size
noise = torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator) / 2.0
for tt in range(t):
for i in range(2):
perlin_slice = self.perlin_noise((h, w), (h, w), batch_size=c, generator=self.generator).to(self.device)
perlin_expanded = perlin_slice.unsqueeze(0).unsqueeze(2)
time_slice = noise[:, :, tt:tt+1, :, :]
noise[:, :, tt:tt+1, :, :] += perlin_expanded
else:
b, c, h, w = self.size
#orig_h, orig_w = h, w
noise = torch.randn(self.size, dtype=self.dtype, layout=self.layout, device=self.device, generator=self.generator) / 2.0
for i in range(2):
noise += self.perlin_noise((h, w), (h, w), batch_size=c, generator=self.generator).to(self.device)
return noise / noise.std()
class PackedNoiseGenerator:
# one generator per stream of a packed multi-stream latent, outputs repacked to the flat [b, 1, n] layout
def __init__(self, cls, x, latent_shapes, seed=42, sigma_min=None, sigma_max=None, **kwargs):
self.x = x
self.size = x.shape
self.dtype = x.dtype
self.layout = x.layout
self.device = x.device
self.seed = seed
self.generator = torch.Generator(device=x.device).manual_seed(seed)
self.streams = []
for shape in latent_shapes:
stream_size = (x.shape[0], *shape[1:])
self.streams.append(cls(size=stream_size, dtype=x.dtype, layout=x.layout, device=x.device, seed=seed, generator=self.generator,
sigma_min=sigma_min, sigma_max=sigma_max, **kwargs))
def update(self, **kwargs):
for stream in self.streams:
stream.update(**kwargs)
def __call__(self, **kwargs):
noise = [stream(**kwargs).reshape(self.size[0], 1, -1) for stream in self.streams]
return torch.cat(noise, dim=-1)
from functools import partial
NOISE_GENERATOR_CLASSES = {
"fractal" : FractalNoiseGenerator,
"gaussian" : GaussianNoiseGenerator,
"gaussian_backwards" : GaussianBackwardsNoiseGenerator,
"uniform" : UniformNoiseGenerator,
"pyramid-cascade_B" : CascadeBPyramidNoiseGenerator,
"pyramid-interpolated" : InterpolatedPyramidNoiseGenerator,
"pyramid-bilinear" : noise_generator_factory(PyramidNoiseGenerator, mode='bilinear'),
"pyramid-bicubic" : noise_generator_factory(PyramidNoiseGenerator, mode='bicubic'),
"pyramid-nearest" : noise_generator_factory(PyramidNoiseGenerator, mode='nearest'),
"hires-pyramid-bilinear": noise_generator_factory(HiresPyramidNoiseGenerator, mode='bilinear'),
"hires-pyramid-bicubic" : noise_generator_factory(HiresPyramidNoiseGenerator, mode='bicubic'),
"hires-pyramid-nearest" : noise_generator_factory(HiresPyramidNoiseGenerator, mode='nearest'),
"brownian" : BrownianNoiseGenerator,
"laplacian" : LaplacianNoiseGenerator,
"studentt" : StudentTNoiseGenerator,
"wavelet" : WaveletNoiseGenerator,
"perlin" : PerlinNoiseGenerator,
}
NOISE_GENERATOR_CLASSES_SIMPLE = {
"none" : GaussianNoiseGenerator,
"brownian" : BrownianNoiseGenerator,
"gaussian" : GaussianNoiseGenerator,
"gaussian_backwards" : GaussianBackwardsNoiseGenerator,
"laplacian" : LaplacianNoiseGenerator,
"perlin" : PerlinNoiseGenerator,
"studentt" : StudentTNoiseGenerator,
"uniform" : UniformNoiseGenerator,
"wavelet" : WaveletNoiseGenerator,
"brown" : noise_generator_factory(FractalNoiseGenerator, alpha=2.0),
"pink" : noise_generator_factory(FractalNoiseGenerator, alpha=1.0),
"white" : noise_generator_factory(FractalNoiseGenerator, alpha=0.0),
"blue" : noise_generator_factory(FractalNoiseGenerator, alpha=-1.0),
"violet" : noise_generator_factory(FractalNoiseGenerator, alpha=-2.0),
"ultraviolet_A" : noise_generator_factory(FractalNoiseGenerator, alpha=-3.0),
"ultraviolet_B" : noise_generator_factory(FractalNoiseGenerator, alpha=-4.0),
"ultraviolet_C" : noise_generator_factory(FractalNoiseGenerator, alpha=-5.0),
"hires-pyramid-bicubic" : noise_generator_factory(HiresPyramidNoiseGenerator, mode='bicubic'),
"hires-pyramid-bilinear": noise_generator_factory(HiresPyramidNoiseGenerator, mode='bilinear'),
"hires-pyramid-nearest" : noise_generator_factory(HiresPyramidNoiseGenerator, mode='nearest'),
"pyramid-bicubic" : noise_generator_factory(PyramidNoiseGenerator, mode='bicubic'),
"pyramid-bilinear" : noise_generator_factory(PyramidNoiseGenerator, mode='bilinear'),
"pyramid-nearest" : noise_generator_factory(PyramidNoiseGenerator, mode='nearest'),
"pyramid-interpolated" : InterpolatedPyramidNoiseGenerator,
"pyramid-cascade_B" : CascadeBPyramidNoiseGenerator,
}
if OPENSIMPLEX_ENABLE:
NOISE_GENERATOR_CLASSES.update({
"simplex": SimplexNoiseGenerator,
})
NOISE_GENERATOR_NAMES = tuple(NOISE_GENERATOR_CLASSES.keys())
NOISE_GENERATOR_NAMES_SIMPLE = tuple(NOISE_GENERATOR_CLASSES_SIMPLE.keys())
@precision_tool.cast_tensor
def prepare_noise(latent_image, seed, noise_type, noise_inds=None, alpha=1.0, k=1.0): # adapted from comfy/sample.py: https://github.com/comfyanonymous/ComfyUI
#optional arg skip can be used to skip and discard x number of noise generations for a given seed
noise_func = NOISE_GENERATOR_CLASSES.get(noise_type)(x=latent_image, seed=seed, sigma_min=0.0291675, sigma_max=14.614642) # WARNING: HARDCODED SDXL SIGMA RANGE!
if noise_type == "fractal":
noise_func.alpha = alpha
noise_func.k = k
# from here until return is very similar to comfy/sample.py
if noise_inds is None:
return noise_func(sigma=14.614642, sigma_next=0.0291675)
unique_inds, inverse = np.unique(noise_inds, return_inverse=True)
noises = []
for i in range(unique_inds[-1]+1):
noise = noise_func(size = [1] + list(latent_image.size())[1:], dtype=latent_image.dtype, layout=latent_image.layout, device=latent_image.device)
if i in unique_inds:
noises.append(noise)
noises = [noises[i] for i in inverse]
noises = torch.cat(noises, axis=0)
return noises
@@ -0,0 +1,140 @@
import torch
import math
from typing import Optional
# Remainder solution
def _phi(j, neg_h):
remainder = torch.zeros_like(neg_h)
for k in range(j):
remainder += (neg_h)**k / math.factorial(k)
phi_j_h = ((neg_h).exp() - remainder) / (neg_h)**j
return phi_j_h
def calculate_gamma(c2, c3):
return (3*(c3**3) - 2*c3) / (c2*(2 - 3*c2))
# Exact analytic solution originally calculated by Clybius. https://github.com/Clybius/ComfyUI-Extra-Samplers/tree/main
def _gamma(n: int,) -> int:
"""
https://en.wikipedia.org/wiki/Gamma_function
for every positive integer n,
Γ(n) = (n-1)!
"""
return math.factorial(n-1)
def _incomplete_gamma(s: int, x: float, gamma_s: Optional[int] = None) -> float:
"""
https://en.wikipedia.org/wiki/Incomplete_gamma_function#Special_values
if s is a positive integer,
Γ(s, x) = (s-1)!*∑{k=0..s-1}(x^k/k!)
"""
if gamma_s is None:
gamma_s = _gamma(s)
sum_: float = 0
# {k=0..s-1} inclusive
for k in range(s):
numerator: float = x**k
denom: int = math.factorial(k)
quotient: float = numerator/denom
sum_ += quotient
incomplete_gamma_: float = sum_ * math.exp(-x) * gamma_s
return incomplete_gamma_
def phi(j: int, neg_h: float, ):
"""
For j={1,2,3}: you could alternatively use Kat's phi_1, phi_2, phi_3 which perform fewer steps
Lemma 1
https://arxiv.org/abs/2308.02157
ϕj(-h) = 1/h^j*∫{0..h}(e^(τ-h)*(τ^(j-1))/((j-1)!)dτ)
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84
= 1/h^j*[(e^(-h)*(-τ)^(-j)*τ(j))/((j-1)!)]{0..h}
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84+between+0+and+h
= 1/h^j*((e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h)))/(j-1)!)
= (e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h))/((j-1)!*h^j)
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/(j-1)!
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/Γ(j)
= (e^(-h)*(-h)^(-j)*(1-Γ(j,-h)/Γ(j))
requires j>0
"""
assert j > 0
gamma_: float = _gamma(j)
incomp_gamma_: float = _incomplete_gamma(j, neg_h, gamma_s=gamma_)
phi_: float = math.exp(neg_h) * neg_h**-j * (1-incomp_gamma_/gamma_)
return phi_
from mpmath import mp, mpf, factorial, exp
mp.dps = 80 # e.g. 80 decimal digits (~ float256)
def phi_mpmath_series(j: int, neg_h: float) -> float:
"""
Arbitrary‐precision phi_j(-h) via the remainder‐series definition,
using mpmath’s mpf and factorial.
"""
j = int(j)
z = mpf(float(neg_h))
S = mp.mpf('0') # S = sum_{k=0..j-1} z^k / k!
for k in range(j):
S += (z**k) / factorial(k)
phi_val = (exp(z) - S) / (z**j)
return float(phi_val)
class Phi:
def __init__(self, h, c, analytic_solution=False):
self.h = h
self.c = c
self.cache = {}
if analytic_solution:
#self.phi_f = superphi
self.phi_f = phi_mpmath_series
self.h = mpf(float(h))
self.c = [mpf(c_val) for c_val in c]
#self.c = c
#self.phi_f = phi
else:
self.phi_f = phi
#self.phi_f = _phi # remainder method
def __call__(self, j, i=-1):
if (j, i) in self.cache:
return self.cache[(j, i)]
if i < 0:
c = 1
else:
c = self.c[i - 1]
if c == 0:
self.cache[(j, i)] = 0
return 0
if j == 0 and type(c) in {float, torch.Tensor}:
result = math.exp(float(-self.h * c))
else:
result = self.phi_f(j, -self.h * c)
self.cache[(j, i)] = result
return result
from mpmath import mp, mpf, gamma, gammainc
def superphi(j: int, neg_h: float, ):
gamma_: float = gamma(j)
incomp_gamma_: float = gamma_ - gammainc(j, 0, float(neg_h))
phi_: float = float(math.exp(float(neg_h)) * neg_h**-j) * (1-incomp_gamma_/gamma_)
return float(phi_)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,969 @@
import math
import torch
from torch import Tensor
from typing import Optional, Callable, Tuple, Dict, Any, Union, TYPE_CHECKING, TypeVar
if TYPE_CHECKING:
from .rk_method_beta import RK_Method_Exponential, RK_Method_Linear
import comfy.model_patcher
import comfy.supported_models
from .noise_classes import NOISE_GENERATOR_CLASSES, NOISE_GENERATOR_CLASSES_SIMPLE, PackedNoiseGenerator
from .constants import MAX_STEPS
from ..helper import ExtraOptions, has_nested_attr
from ..latents import normalize_zscore, get_orthogonal, get_collinear, is_packed_latent
from ..res4lyf import RESplain
NOISE_MODE_NAMES = ["none",
#"hard_sq",
"hard",
"lorentzian",
"soft",
"soft-linear",
"softer",
"eps",
"sinusoidal",
"exp",
"vpsde",
"er4",
"hard_var",
]
def get_data_from_step(x, x_next, sigma, sigma_next): # assumes 100% linear trajectory
h = sigma_next - sigma
return (sigma_next * x - sigma * x_next) / h
def get_epsilon_from_step(x, x_next, sigma, sigma_next):
h = sigma_next - sigma
return (x - x_next) / h
class RK_NoiseSampler:
def __init__(self,
RK : Union["RK_Method_Exponential", "RK_Method_Linear"],
model,
step : int=0,
device : str='cuda',
dtype : torch.dtype=torch.float64,
extra_options : str=""
):
self.device = device
self.dtype = dtype
self.model = model
if has_nested_attr(model, "inner_model.inner_model.model_sampling"):
model_sampling = model.inner_model.inner_model.model_sampling
elif has_nested_attr(model, "model.model_sampling"):
model_sampling = model.model.model_sampling
self.sigma_max = model_sampling.sigma_max.to(dtype=self.dtype, device=self.device)
self.sigma_min = model_sampling.sigma_min.to(dtype=self.dtype, device=self.device)
self.sigma_fn = RK.sigma_fn
self.t_fn = RK.t_fn
self.h_fn = RK.h_fn
self.row_offset = 1 if not RK.IMPLICIT else 0
self.step = step
self.noise_sampler = None
self.noise_sampler2 = None
self.noise_mode_sde = None
self.noise_mode_sde_substep = None
self.LOCK_H_SCALE = True
self.CONST = isinstance(model_sampling, comfy.model_sampling.CONST)
self.VARIANCE_PRESERVING = isinstance(model_sampling, comfy.model_sampling.CONST)
self.extra_options = extra_options
self.EO = ExtraOptions(extra_options)
self.DOWN_SUBSTEP = self.EO("down_substep")
self.DOWN_STEP = self.EO("down_step")
self.init_noise = None
self.av_split = None
self.av_total = None
self.av_shift_video = None
self.av_shift_audio = None
self.av_audio_noise_scale = 1.0
self.av_audio_eta_scale = 1.0
self.latent_shapes = self._find_latent_shapes(model)
if not self.EO("av_disable"):
self._init_av_streams(model)
@staticmethod
def _find_latent_shapes(model):
conds = getattr(getattr(model, "inner_model", None), "conds", None)
if not isinstance(conds, dict):
return None
for cond_list in conds.values():
for cond in cond_list or []:
model_conds = cond.get('model_conds', {})
if 'latent_shapes' in model_conds:
return model_conds['latent_shapes'].cond
return None
def _init_av_streams(self, model) -> None:
# av_shift_audio stays None when both streams share one schedule (the column split and the audio noise knob still apply there)
guider = getattr(model, "inner_model", None)
inner_model = getattr(guider, "inner_model", None)
diffusion_model = getattr(inner_model, "diffusion_model", None)
latent_shapes = self.latent_shapes
if latent_shapes is None or len(latent_shapes) != 2:
return
self.av_split = int(math.prod(latent_shapes[0][1:]))
self.av_total = self.av_split + int(math.prod(latent_shapes[1][1:]))
self.av_audio_noise_scale = self.EO("av_audio_noise_scale", 1.0)
self.av_audio_eta_scale = self.EO("av_audio_eta_scale", 1.0)
shift_audio = getattr(diffusion_model, "sigma_shift_audio", None)
# when model_sampling uses audio_scale skip shifting the audio schedule
# todo: remove this shifting code eventually once audio is always pre-scaled properly
if hasattr(guider, "model_patcher"):
model_sampling = guider.model_patcher.get_model_object("model_sampling")
else:
model_sampling = getattr(inner_model, "model_sampling", None)
is_audio_scale_set = getattr(model_sampling, "audio_scale", 1.0) != 1.0
if shift_audio is not None and not is_audio_scale_set:
model_options = getattr(guider, "model_options", {})
transformer_options = model_options.get("transformer_options", {}) if isinstance(model_options, dict) else {}
self.av_shift_video = float(transformer_options.get("minimax_h3_sigma_shift_video", getattr(diffusion_model, "sigma_shift_video", 12.0)))
self.av_shift_audio = float(transformer_options.get("minimax_h3_sigma_shift_audio", shift_audio))
RESplain("AV stream shifts applied. shift_video:", self.av_shift_video, "shift_audio:", self.av_shift_audio, debug=True)
elif is_audio_scale_set:
RESplain("AV stream split active, shifts handled by model_sampling.audio_scale", debug=True)
def _av_sigma_audio(self, sigma:float) -> float:
base = sigma / (self.av_shift_video + sigma * (1.0 - self.av_shift_video))
return self.av_shift_audio * base / (1.0 + (self.av_shift_audio - 1.0) * base)
@staticmethod
def _av_renoise_var(sigma_from:float, sigma_to:float) -> float:
# variance a full-eta RF ancestral step injects stepping from sigma_from to sigma_to on one schedule
sigma_down = sigma_to * sigma_to / sigma_from
return sigma_to ** 2 - sigma_down ** 2 * (1.0 - sigma_to) ** 2 / (1.0 - sigma_down) ** 2
def scale_av_noise(self, noise:Tensor, sigma_from, sigma_to) -> Tensor:
# audio columns get the noise magnitude their own shifted schedule calls for over this step
# apply to unit-variance noise after any normalization, before the sigma_up multiply
if self.av_split is None or noise.shape[-1] != self.av_total:
return noise
ratio = self.av_audio_noise_scale
if self.av_shift_audio is not None:
s_from, s_to = float(sigma_from), float(sigma_to)
if s_to > 0.0 and s_from > s_to:
video_var = self._av_renoise_var(s_from, s_to)
if video_var > 0.0:
audio_var = self._av_renoise_var(self._av_sigma_audio(s_from), self._av_sigma_audio(s_to))
ratio *= (max(audio_var, 0.0) / video_var) ** 0.5
if ratio == 1.0:
return noise
noise[..., self.av_split:] *= ratio
return noise
def blend_av_eta(self, x_noised:Tensor, x_next:Tensor) -> Tensor:
# interpolate the audio columns between the deterministic landing (x_next) and the full eta result
# 0.0 gives audio a pure ODE step while video keeps its eta, 1.0 leaves the eta step untouched
if self.av_split is None or self.av_audio_eta_scale == 1.0 or x_noised.shape[-1] != self.av_total:
return x_noised
w = self.av_audio_eta_scale
x_noised[..., self.av_split:] = (1.0 - w) * x_next[..., self.av_split:] + w * x_noised[..., self.av_split:]
return x_noised
def init_noise_samplers(self,
x : Tensor,
noise_seed : int,
noise_seed_substep : int,
noise_sampler_type : str,
noise_sampler_type2 : str,
noise_mode_sde : str,
noise_mode_sde_substep : str,
overshoot_mode : str,
overshoot_mode_substep : str,
noise_boost_step : float,
noise_boost_substep : float,
alpha : float,
alpha2 : float,
k : float = 1.0,
k2 : float = 1.0,
scale : float = 0.1,
scale2 : float = 0.1,
last_rng = None,
last_rng_substep = None,
latent_shapes = None,
) -> None:
self.noise_sampler_type = noise_sampler_type
self.noise_sampler_type2 = noise_sampler_type2
self.noise_mode_sde = noise_mode_sde
self.noise_mode_sde_substep = noise_mode_sde_substep
self.overshoot_mode = overshoot_mode
self.overshoot_mode_substep = overshoot_mode_substep
self.noise_boost_step = noise_boost_step
self.noise_boost_substep = noise_boost_substep
self.s_in = x.new_ones([1], dtype=self.dtype, device=self.device)
# torch's RNG stream differs per dtype, so noise_dtype — not the math precision — decides
# which noise realization a seed produces; the float64 default keeps seeds stable across work_dtype
noise_dtype = self.EO("noise_dtype", self.dtype)
if x.dtype != noise_dtype:
x = x.to(noise_dtype)
if noise_seed >= 0:
seed = noise_seed
RESplain("SDE noise seed: ", seed, debug=True)
elif last_rng is not None:
seed = 0
RESplain("SDE noise seed: restoring from last_rng state", debug=True)
else:
seed = torch.initial_seed() + 1
RESplain("SDE noise seed: ", seed, " (set via torch.initial_seed()+1)", debug=True)
#seed2 = seed + MAX_STEPS #for substep noise generation. offset needed to ensure seeds are not reused
if latent_shapes is None:
latent_shapes = self.latent_shapes
if noise_sampler_type == "fractal":
self.noise_sampler = self._build_noise_sampler(NOISE_GENERATOR_CLASSES.get(noise_sampler_type), x, seed, latent_shapes)
self.noise_sampler.update(alpha=alpha, k=k, scale=scale)
else:
self.noise_sampler = self._build_noise_sampler(NOISE_GENERATOR_CLASSES_SIMPLE.get(noise_sampler_type), x, seed, latent_shapes)
if noise_sampler_type2 == "fractal":
self.noise_sampler2 = self._build_noise_sampler(NOISE_GENERATOR_CLASSES.get(noise_sampler_type2), x, noise_seed_substep, latent_shapes)
self.noise_sampler2.update(alpha=alpha2, k=k2, scale=scale2)
else:
self.noise_sampler2 = self._build_noise_sampler(NOISE_GENERATOR_CLASSES_SIMPLE.get(noise_sampler_type2), x, noise_seed_substep, latent_shapes)
if last_rng is not None:
self.noise_sampler .generator.set_state(last_rng)
self.noise_sampler2.generator.set_state(last_rng_substep)
def _build_noise_sampler(self, cls, x:Tensor, seed:int, latent_shapes):
# packed multi-stream latents get one generator per stream so structured noise sees each stream's real shape
if is_packed_latent(latent_shapes) and x.dim() == 3:
return PackedNoiseGenerator(cls, x=x, latent_shapes=latent_shapes, seed=seed, sigma_min=self.sigma_min, sigma_max=self.sigma_max)
return cls(x=x, seed=seed, sigma_min=self.sigma_min, sigma_max=self.sigma_max)
def set_substep_list(self, RK:Union["RK_Method_Exponential", "RK_Method_Linear"]) -> None:
self.multistep_stages = RK.multistep_stages
self.rows = RK.rows
self.C = RK.C
self.s_ = self.sigma_fn(self.t_fn(self.sigma) + self.h * self.C)
def get_substep_list(self, RK:Union["RK_Method_Exponential", "RK_Method_Linear"], sigma, h) -> None:
s_ = RK.sigma_fn(RK.t_fn(sigma) + h * RK.C)
return s_
def get_sde_coeff(self, sigma_next:Tensor, sigma_down:Tensor=None, sigma_up:Tensor=None, eta:float=0.0, VP_OVERRIDE=None) -> Tuple[Tensor,Tensor,Tensor]:
VARIANCE_PRESERVING = VP_OVERRIDE if VP_OVERRIDE is not None else self.VARIANCE_PRESERVING
if VARIANCE_PRESERVING:
if sigma_down is not None:
alpha_ratio = (1 - sigma_next) / (1 - sigma_down)
sigma_up = (sigma_next ** 2 - sigma_down ** 2 * alpha_ratio ** 2) ** 0.5
elif sigma_up is not None:
if sigma_up >= sigma_next:
RESplain("Maximum VPSDE noise level exceeded: falling back to hard noise mode.", debug=True)
if eta >= 1:
sigma_up = sigma_next * 0.9999 #avoid sqrt(neg_num) later
else:
sigma_up = sigma_next * eta
if VP_OVERRIDE is not None:
sigma_signal = 1 - sigma_next
else:
sigma_signal = self.sigma_max - sigma_next
sigma_residual = (sigma_next ** 2 - sigma_up ** 2) ** .5
alpha_ratio = sigma_signal + sigma_residual
sigma_down = sigma_residual / alpha_ratio
else:
alpha_ratio = torch.ones_like(sigma_next)
if sigma_down is not None:
sigma_up = (sigma_next ** 2 - sigma_down ** 2) ** .5 # not sure this is correct #TODO: CHECK THIS
elif sigma_up is not None:
sigma_down = (sigma_next ** 2 - sigma_up ** 2) ** .5
return alpha_ratio, sigma_down, sigma_up
def set_sde_step(self, sigma:Tensor, sigma_next:Tensor, eta:float, overshoot:float, s_noise:float) -> None:
self.sigma_0 = sigma
self.sigma_next = sigma_next
self.s_noise = s_noise
self.eta = eta
self.overshoot = overshoot
self.sigma_up_eta, self.sigma_eta, self.sigma_down_eta, self.alpha_ratio_eta \
= self.get_sde_step(sigma, sigma_next, eta, self.noise_mode_sde, self.DOWN_STEP, SUBSTEP=False)
self.sigma_up, self.sigma, self.sigma_down, self.alpha_ratio \
= self.get_sde_step(sigma, sigma_next, overshoot, self.overshoot_mode, self.DOWN_STEP, SUBSTEP=False)
self.h = self.h_fn(self.sigma_down, self.sigma)
self.h_no_eta = self.h_fn(self.sigma_next, self.sigma)
self.h = self.h + self.noise_boost_step * (self.h_no_eta - self.h)
def set_sde_substep(self,
row : int,
multistep_stages : int,
eta_substep : float,
overshoot_substep : float,
s_noise_substep : float,
full_iter : int = 0,
diag_iter : int = 0,
implicit_steps_full : int = 0,
implicit_steps_diag : int = 0
) -> None:
# start with stepsizes for no overshoot/noise addition/noise swapping
self.sub_sigma_up_eta = self.sub_sigma_up = 0.0
self.sub_sigma_eta = self.sub_sigma = self.s_[row]
self.sub_sigma_down_eta = self.sub_sigma_down = self.sub_sigma_next = self.s_[row+self.row_offset+multistep_stages]
self.sub_alpha_ratio_eta = self.sub_alpha_ratio = 1.0
self.s_noise_substep = s_noise_substep
self.eta_substep = eta_substep
self.overshoot_substep = overshoot_substep
if row < self.rows and self.s_[row+self.row_offset+multistep_stages] > 0:
if diag_iter > 0 and diag_iter == implicit_steps_diag and self.EO("implicit_substep_skip_final_eta"):
pass
elif diag_iter > 0 and self.EO("implicit_substep_only_first_eta"):
pass
elif full_iter > 0 and full_iter == implicit_steps_full and self.EO("implicit_step_skip_final_eta"):
pass
elif full_iter > 0 and self.EO("implicit_step_only_first_eta"):
pass
elif (full_iter > 0 or diag_iter > 0) and self.noise_sampler_type2 == "brownian":
pass # brownian noise does not increment its seed when generated, deactivate on implicit repeats to avoid burn
elif full_iter > 0 and self.EO("implicit_step_only_first_all_eta"):
self.sigma_down_eta = self.sigma_next
self.sigma_up_eta *= 0
self.alpha_ratio_eta /= self.alpha_ratio_eta
self.sigma_down = self.sigma_next
self.sigma_up *= 0
self.alpha_ratio /= self.alpha_ratio
self.h_new = self.h = self.h_no_eta
elif (row < self.rows-self.row_offset-multistep_stages or diag_iter < implicit_steps_diag) or self.EO("substep_eta_use_final"):
self.sub_sigma_up, self.sub_sigma, self.sub_sigma_down, self.sub_alpha_ratio = self.get_sde_substep(sigma = self.s_[row],
sigma_next = self.s_[row+self.row_offset+multistep_stages],
eta = overshoot_substep,
noise_mode_override = self.overshoot_mode_substep,
DOWN = self.DOWN_SUBSTEP)
self.sub_sigma_up_eta, self.sub_sigma_eta, self.sub_sigma_down_eta, self.sub_alpha_ratio_eta = self.get_sde_substep(sigma = self.s_[row],
sigma_next = self.s_[row+self.row_offset+multistep_stages],
eta = eta_substep,
noise_mode_override = self.noise_mode_sde_substep,
DOWN = self.DOWN_SUBSTEP)
if self.h_fn(self.sub_sigma_next, self.sigma) != 0:
self.h_new = self.h * self.h_fn(self.sub_sigma_down, self.sigma) / self.h_fn(self.sub_sigma_next, self.sigma)
self.h_eta = self.h * self.h_fn(self.sub_sigma_down_eta, self.sigma) / self.h_fn(self.sub_sigma_next, self.sigma)
self.h_new_orig = self.h_new.clone()
self.h_new = self.h_new + self.noise_boost_substep * (self.h - self.h_eta)
else:
self.h_new = self.h_eta = self.h
self.h_new_orig = self.h_new.clone()
def get_sde_substep(self,
sigma :Tensor,
sigma_next :Tensor,
eta :float = 0.0 ,
noise_mode_override :Optional[str] = None ,
DOWN :bool = False,
) -> Tuple[Tensor,Tensor,Tensor,Tensor]:
return self.get_sde_step(sigma=sigma, sigma_next=sigma_next, eta=eta, noise_mode_override=noise_mode_override, DOWN=DOWN, SUBSTEP=True,)
def get_sde_step(self,
sigma :Tensor,
sigma_next :Tensor,
eta :float = 0.0 ,
noise_mode_override :Optional[str] = None ,
DOWN :bool = False,
SUBSTEP :bool = False,
VP_OVERRIDE = None,
) -> Tuple[Tensor,Tensor,Tensor,Tensor]:
VARIANCE_PRESERVING = VP_OVERRIDE if VP_OVERRIDE is not None else self.VARIANCE_PRESERVING
if noise_mode_override is not None:
noise_mode = noise_mode_override
elif SUBSTEP:
noise_mode = self.noise_mode_sde_substep
else:
noise_mode = self.noise_mode_sde
if DOWN: #calculates noise level by first scaling sigma_down from sigma_next, instead of sigma_up from sigma_next
eta_fn = lambda eta_scale: 1-eta_scale
sud_fn = lambda sd: (sd, None)
else:
eta_fn = lambda eta_scale: eta_scale
sud_fn = lambda su: (None, su)
su, sd, sud = None, None, None
eta_ratio = None
sigma_base = sigma_next
sigmax = self.sigma_max if VP_OVERRIDE is None else 1
sigma_n = sigma / sigmax
sigma_next_n = sigma_next / sigmax
match noise_mode:
case "hard":
eta_ratio = eta
case "exp":
h = -(sigma_next_n/sigma_n).log()
eta_ratio = (1 - (-2*eta*h).exp())**.5
case "soft":
eta_ratio = 1-(1 - eta) + eta * (sigma_next_n / sigma_n)
case "softer":
eta_ratio = 1-torch.sqrt(1 - (eta**2 * (sigma_n**2 - sigma_next_n**2)) / sigma_n**2)
case "soft-linear":
eta_ratio = 1-eta * (sigma_next_n - sigma_n)
case "sinusoidal":
eta_ratio = eta * torch.sin(torch.pi * sigma_next_n) ** 2
case "eps":
eta_ratio = eta * torch.sqrt((sigma_next_n/sigma_n) ** 2 * (sigma_n ** 2 - sigma_next_n ** 2) )
case "lorentzian":
eta_ratio = eta
alpha = 1 / (sigma_next_n.to(sigma.dtype)**2 + 1)
sigma_base = (sigmax * (1 - alpha) ** 0.5).to(sigma.dtype)
case "hard_var":
sigma_var_n = (-1 + torch.sqrt(1 + 4 * sigma_n)) / 2
if sigma_next_n > sigma_var_n:
eta_ratio = 0
sigma_base = sigma_next
else:
eta_ratio = eta
sigma_base = torch.sqrt((sigma - sigma_next).abs() + 1e-10)
case "hard_sq":
sigma_hat = sigma * (1 + eta)
su = (sigma_hat ** 2 - sigma ** 2) ** .5 #su
if VARIANCE_PRESERVING:
alpha_ratio, sd, su = self.get_sde_coeff(sigma_next, None, su, eta, VARIANCE_PRESERVING)
else:
sd = sigma_next
sigma = sigma_hat
alpha_ratio = torch.ones_like(sigma)
case "vpsde":
alpha_ratio, sd, su = self.get_vpsde_step_RF(sigma, sigma_next, eta)
case "er4":
noise_scaler = lambda s: s * ((s ** eta).exp() + 10.0)
alpha_ratio = noise_scaler(sigma_next_n) / noise_scaler(sigma_n)
sigma_up = (sigma_next ** 2 - sigma ** 2 * alpha_ratio ** 2) ** 0.5
eta_ratio = sigma_up / sigma_next
if eta_ratio is not None:
sud = sigma_base * eta_fn(eta_ratio)
alpha_ratio, sd, su = self.get_sde_coeff(sigma_next, *sud_fn(sud), eta, VARIANCE_PRESERVING)
su = torch.nan_to_num(su, 0.0)
sd = torch.nan_to_num(sd, float(sigma_next))
alpha_ratio = torch.nan_to_num(alpha_ratio, 1.0)
return su, sigma, sd, alpha_ratio
def get_vpsde_step_RF(self, sigma:Tensor, sigma_next:Tensor, eta:float) -> Tuple[Tensor,Tensor,Tensor]:
dt = sigma - sigma_next
sigma_up = eta * sigma * dt**0.5
alpha_ratio = 1 - dt * (eta**2/4) * (1 + sigma)
sigma_down = sigma_next - (eta/4)*sigma*(1-sigma)*(sigma - sigma_next)
return sigma_up, sigma_down, alpha_ratio
def linear_noise_init(self, y:Tensor, sigma_curr:Tensor, x_base:Optional[Tensor]=None, x_curr:Optional[Tensor]=None, mask:Optional[Tensor]=None) -> Tensor:
y_noised = (self.sigma_max - sigma_curr) * y + sigma_curr * self.init_noise
if x_curr is not None:
x_curr = x_curr + sigma_curr * (self.init_noise - y)
x_base = x_base + self.sigma * (self.init_noise - y)
return y_noised, x_base, x_curr
if mask is not None:
y_noised = mask * y_noised + (1-mask) * y
return y_noised
def linear_noise_step(self, y:Tensor, sigma_curr:Optional[Tensor]=None, x_base:Optional[Tensor]=None, x_curr:Optional[Tensor]=None, brownian_sigma:Optional[Tensor]=None, brownian_sigma_next:Optional[Tensor]=None, mask:Optional[Tensor]=None) -> Tensor:
if self.sigma_up_eta == 0 or self.sigma_next == 0:
return y, x_base, x_curr
sigma_curr = self.sub_sigma if sigma_curr is None else sigma_curr
brownian_sigma = sigma_curr if brownian_sigma is None else brownian_sigma
brownian_sigma_next = self.sigma_next.clone() if brownian_sigma_next is None else brownian_sigma_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
if brownian_sigma_next > brownian_sigma and not self.EO("disable_brownian_swap"): # should this really be done?
brownian_sigma, brownian_sigma_next = brownian_sigma_next, brownian_sigma
noise = self.noise_sampler(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
y_noised = (self.sigma_max - sigma_curr) * y + sigma_curr * noise
if x_curr is not None:
x_curr = x_curr + sigma_curr * (noise - y)
x_base = x_base + self.sigma * (noise - y)
return y_noised, x_base, x_curr
if mask is not None:
y_noised = mask * y_noised + (1-mask) * y
return y_noised
def linear_noise_substep(self, y:Tensor, sigma_curr:Optional[Tensor]=None, x_base:Optional[Tensor]=None, x_curr:Optional[Tensor]=None, brownian_sigma:Optional[Tensor]=None, brownian_sigma_next:Optional[Tensor]=None, mask:Optional[Tensor]=None) -> Tensor:
if self.sub_sigma_up_eta == 0 or self.sub_sigma_next == 0:
return y, x_base, x_curr
sigma_curr = self.sub_sigma if sigma_curr is None else sigma_curr
brownian_sigma = sigma_curr if brownian_sigma is None else brownian_sigma
brownian_sigma_next = self.sub_sigma_next.clone() if brownian_sigma_next is None else brownian_sigma_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
if brownian_sigma_next > brownian_sigma and not self.EO("disable_brownian_swap"): # should this really be done?
brownian_sigma, brownian_sigma_next = brownian_sigma_next, brownian_sigma
noise = self.noise_sampler2(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
y_noised = (self.sigma_max - sigma_curr) * y + sigma_curr * noise
if x_curr is not None:
x_curr = x_curr + sigma_curr * (noise - y)
x_base = x_base + self.sigma * (noise - y)
return y_noised, x_base, x_curr
if mask is not None:
y_noised = mask * y_noised + (1-mask) * y
return y_noised
def swap_noise_step(self, x_0:Tensor, x_next:Tensor, brownian_sigma:Optional[Tensor]=None, brownian_sigma_next:Optional[Tensor]=None, mask:Optional[Tensor]=None) -> Tensor:
if self.sigma_up_eta == 0 or self.sigma_next == 0:
return x_next
brownian_sigma = self.sigma.clone() if brownian_sigma is None else brownian_sigma
brownian_sigma_next = self.sigma_next.clone() if brownian_sigma_next is None else brownian_sigma_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
eps_next = (x_0 - x_next) / (self.sigma - self.sigma_next)
denoised_next = x_0 - self.sigma * eps_next
if brownian_sigma_next > brownian_sigma and not self.EO("disable_brownian_swap"): # should this really be done?
brownian_sigma, brownian_sigma_next = brownian_sigma_next, brownian_sigma
noise = self.noise_sampler(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
noise = self.scale_av_noise(noise, self.sigma, self.sigma_next)
x_noised = self.alpha_ratio_eta * (denoised_next + self.sigma_down_eta * eps_next) + self.sigma_up_eta * noise * self.s_noise
x_noised = self.blend_av_eta(x_noised, x_next)
if mask is not None:
x = mask * x_noised + (1-mask) * x_next
else:
x = x_noised
return x
def swap_noise_substep(self, x_0:Tensor, x_next:Tensor, brownian_sigma:Optional[Tensor]=None, brownian_sigma_next:Optional[Tensor]=None, mask:Optional[Tensor]=None, guide:Optional[Tensor]=None) -> Tensor:
if self.sub_sigma_up_eta == 0 or self.sub_sigma_next == 0:
return x_next
brownian_sigma = self.sub_sigma.clone() if brownian_sigma is None else brownian_sigma
brownian_sigma_next = self.sub_sigma_next.clone() if brownian_sigma_next is None else brownian_sigma_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
eps_next = (x_0 - x_next) / (self.sigma - self.sub_sigma_next)
denoised_next = x_0 - self.sigma * eps_next
if brownian_sigma_next > brownian_sigma and not self.EO("disable_brownian_swap"): # should this really be done?
brownian_sigma, brownian_sigma_next = brownian_sigma_next, brownian_sigma
noise = self.noise_sampler2(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
noise = self.scale_av_noise(noise, self.sub_sigma, self.sub_sigma_next)
x_noised = self.sub_alpha_ratio_eta * (denoised_next + self.sub_sigma_down_eta * eps_next) + self.sub_sigma_up_eta * noise * self.s_noise_substep
x_noised = self.blend_av_eta(x_noised, x_next)
if mask is not None:
x = mask * x_noised + (1-mask) * x_next
else:
x = x_noised
return x
def swap_noise_inv_substep(self, x_0:Tensor, x_next:Tensor, eta_substep:float, row:int, row_offset_multistep_stages:int, brownian_sigma:Optional[Tensor]=None, brownian_sigma_next:Optional[Tensor]=None, mask:Optional[Tensor]=None, guide:Optional[Tensor]=None) -> Tensor:
if self.sub_sigma_up_eta == 0 or self.sub_sigma_next == 0:
return x_next
brownian_sigma = self.sub_sigma.clone() if brownian_sigma is None else brownian_sigma
brownian_sigma_next = self.sub_sigma_next.clone() if brownian_sigma_next is None else brownian_sigma_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
eps_next = (x_0 - x_next) / ((1-self.sigma) - (1-self.sub_sigma_next))
denoised_next = x_0 - (1-self.sigma) * eps_next
if brownian_sigma_next > brownian_sigma and not self.EO("disable_brownian_swap"): # should this really be done?
brownian_sigma, brownian_sigma_next = brownian_sigma_next, brownian_sigma
noise = self.noise_sampler2(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
# inverted-domain (unsampling) injection: audio columns are left unscaled for reverse steps
sub_sigma_up, sub_sigma, sub_sigma_down, sub_alpha_ratio = self.get_sde_substep(sigma = 1-self.s_[row],
sigma_next = 1-self.s_[row_offset_multistep_stages],
eta = eta_substep,
noise_mode_override = self.noise_mode_sde_substep,
DOWN = self.DOWN_SUBSTEP)
x_noised = sub_alpha_ratio * (denoised_next + sub_sigma_down * eps_next) + sub_sigma_up * noise * self.s_noise_substep
if mask is not None:
x = mask * x_noised + (1-mask) * x_next
else:
x = x_noised
return x
def swap_noise(self,
x_0 :Tensor,
x_next :Tensor,
sigma_0 :Tensor,
sigma :Tensor,
sigma_next :Tensor,
sigma_down :Tensor,
sigma_up :Tensor,
alpha_ratio :Tensor,
s_noise :float,
SUBSTEP :bool = False,
brownian_sigma :Optional[Tensor] = None,
brownian_sigma_next :Optional[Tensor] = None,
) -> Tensor:
if sigma_up == 0:
return x_next
if brownian_sigma is None:
brownian_sigma = sigma.clone()
if brownian_sigma_next is None:
brownian_sigma_next = sigma_next.clone()
if sigma_next == 0:
return x_next
if brownian_sigma == brownian_sigma_next:
brownian_sigma_next *= 0.999
eps_next = (x_0 - x_next) / (sigma_0 - sigma_next)
denoised_next = x_0 - sigma_0 * eps_next
if brownian_sigma_next > brownian_sigma:
s_tmp = brownian_sigma
brownian_sigma = brownian_sigma_next
brownian_sigma_next = s_tmp
if not SUBSTEP:
noise = self.noise_sampler(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
else:
noise = self.noise_sampler2(sigma=brownian_sigma, sigma_next=brownian_sigma_next)
noise = normalize_zscore(noise, channelwise=True, inplace=True)
noise = self.scale_av_noise(noise, sigma, sigma_next)
x = alpha_ratio * (denoised_next + sigma_down * eps_next) + sigma_up * noise * s_noise
x = self.blend_av_eta(x, x_next)
return x
# not used. WARNING: some parameters have a different order than swap_noise!
def add_noise_pre(self,
x_0 :Tensor,
x :Tensor,
sigma_up :Tensor,
sigma_0 :Tensor,
sigma :Tensor,
sigma_next :Tensor,
real_sigma_down :Tensor,
alpha_ratio :Tensor,
s_noise :float,
noise_mode :str,
SDE_NOISE_EXTERNAL :bool = False,
sde_noise_t :Optional[Tensor] = None,
SUBSTEP :bool = False,
) -> Tensor:
if not self.CONST and noise_mode == "hard_sq":
if self.LOCK_H_SCALE:
x = self.swap_noise(x_0 = x_0,
x = x,
sigma = sigma,
sigma_0 = sigma_0,
sigma_next = sigma_next,
real_sigma_down = real_sigma_down,
sigma_up = sigma_up,
alpha_ratio = alpha_ratio,
s_noise = s_noise,
SUBSTEP = SUBSTEP,
)
else:
x = self.add_noise( x = x,
sigma_up = sigma_up,
sigma = sigma,
sigma_next = sigma_next,
alpha_ratio = alpha_ratio,
s_noise = s_noise,
SDE_NOISE_EXTERNAL = SDE_NOISE_EXTERNAL,
sde_noise_t = sde_noise_t,
SUBSTEP = SUBSTEP,
)
return x
# only used for handle_tiled_etc_noise_steps() in rk_guide_func_beta.py
def add_noise_post(self,
x_0 :Tensor,
x :Tensor,
sigma_up :Tensor,
sigma_0 :Tensor,
sigma :Tensor,
sigma_next :Tensor,
real_sigma_down :Tensor,
alpha_ratio :Tensor,
s_noise :float,
noise_mode :str,
SDE_NOISE_EXTERNAL :bool = False,
sde_noise_t :Optional[Tensor] = None,
SUBSTEP :bool = False,
) -> Tensor:
if self.CONST or (not self.CONST and noise_mode != "hard_sq"):
if self.LOCK_H_SCALE:
x = self.swap_noise(x_0 = x_0,
x = x,
sigma = sigma,
sigma_0 = sigma_0,
sigma_next = sigma_next,
real_sigma_down = real_sigma_down,
sigma_up = sigma_up,
alpha_ratio = alpha_ratio,
s_noise = s_noise,
SUBSTEP = SUBSTEP,
)
else:
x = self.add_noise( x = x,
sigma_up = sigma_up,
sigma = sigma,
sigma_next = sigma_next,
alpha_ratio = alpha_ratio,
s_noise = s_noise,
SDE_NOISE_EXTERNAL = SDE_NOISE_EXTERNAL,
sde_noise_t = sde_noise_t,
SUBSTEP = SUBSTEP,
)
return x
def add_noise(self,
x :Tensor,
sigma_up :Tensor,
sigma :Tensor,
sigma_next :Tensor,
alpha_ratio :Tensor,
s_noise :float,
SDE_NOISE_EXTERNAL :bool = False,
sde_noise_t :Optional[Tensor] = None,
SUBSTEP :bool = False,
) -> Tensor:
if sigma_next > 0.0 and sigma_up > 0.0:
if sigma_next > sigma:
sigma, sigma_next = sigma_next, sigma
if sigma == sigma_next:
sigma_next = sigma * 0.9999
if not SUBSTEP:
noise = self.noise_sampler (sigma=sigma, sigma_next=sigma_next)
else:
noise = self.noise_sampler2(sigma=sigma, sigma_next=sigma_next)
#noise_ortho = get_orthogonal(noise, x)
#noise_ortho = noise_ortho / noise_ortho.std()model,
noise = normalize_zscore(noise, channelwise=True, inplace=True)
if SDE_NOISE_EXTERNAL:
noise = (1-s_noise) * noise + s_noise * sde_noise_t
noise = self.scale_av_noise(noise, sigma, sigma_next)
# av eta blend is not applied here, this path never sees the deterministic landing point
x_next = alpha_ratio * x + noise * sigma_up * s_noise
return x_next
else:
return x
def sigma_from_to(self,
x_0 : Tensor,
x_down : Tensor,
sigma : Tensor,
sigma_down : Tensor,
sigma_next : Tensor) -> Tensor: #sigma, sigma_from, sigma_to
eps = (x_0 - x_down) / (sigma - sigma_down)
denoised = x_0 - sigma * eps
x_next = denoised + sigma_next * eps # VESDE vs VPSDE equiv.?
return x_next
def rebound_overshoot_step(self, x_0:Tensor, x:Tensor) -> Tensor:
eps = (x_0 - x) / (self.sigma - self.sigma_down)
denoised = x_0 - self.sigma * eps
x = denoised + self.sigma_next * eps
return x
def rebound_overshoot_substep(self, x_0:Tensor, x:Tensor) -> Tensor:
if self.sigma - self.sub_sigma_down > 0:
sub_eps = (x_0 - x) / (self.sigma - self.sub_sigma_down)
sub_denoised = x_0 - self.sigma * sub_eps
x = sub_denoised + self.sub_sigma_next * sub_eps
return x
def prepare_sigmas(self,
sigmas : Tensor,
sigmas_override : Tensor,
d_noise : float,
d_noise_start_step : int,
sampler_mode : str) -> Tuple[Tensor,bool]:
#SIGMA_MIN = torch.full_like(self.sigma_min, 0.00227896) if self.sigma_min < 0.00227896 else self.sigma_min # prevent black image with unsampling flux, which has a sigma_min of 0.0002
SIGMA_MIN = self.sigma_min #torch.full_like(self.sigma_min, max(0.01, self.sigma_min.item()))
if sigmas_override is not None:
sigmas = sigmas_override.clone().to(sigmas.device).to(sigmas.dtype)
if d_noise_start_step == 0:
sigmas = sigmas.clone() * d_noise
UNSAMPLE_FROM_ZERO = False
if sigmas[0] == 0.0: #remove padding used to prevent comfy from adding noise to the latent (for unsampling, etc.)
UNSAMPLE = True
if sigmas[-1] == 0.0:
UNSAMPLE_FROM_ZERO = True
#sigmas = sigmas[1:-1] # was cleaving off 1.0 at the end when restart looping
sigmas = sigmas[1:]
if sigmas[-1] == 0.0:
sigmas = sigmas[:-1]
else:
UNSAMPLE = False
if hasattr(self.model, "sigmas"):
self.model.sigmas = sigmas
if sampler_mode == "standard":
UNSAMPLE = False
consecutive_duplicate_mask = torch.cat((torch.tensor([True], device=sigmas.device), torch.diff(sigmas) != 0))
sigmas = sigmas[consecutive_duplicate_mask]
if sigmas[-1] == 0:
if sigmas[-2] < SIGMA_MIN:
sigmas[-2] = SIGMA_MIN
elif (sigmas[-2] - SIGMA_MIN).abs() > 1e-4:
sigmas = torch.cat((sigmas[:-1], SIGMA_MIN.unsqueeze(0), sigmas[-1:]))
elif UNSAMPLE_FROM_ZERO and not torch.isclose(sigmas[0], SIGMA_MIN):
sigmas = torch.cat([SIGMA_MIN.unsqueeze(0), sigmas])
self.sigmas = sigmas
self.UNSAMPLE = UNSAMPLE
self.d_noise = d_noise
self.sampler_mode = sampler_mode
return sigmas, UNSAMPLE
def extract_latent_swap_noise(self, x:Tensor, x_noise_swapped:Tensor, sigma:Tensor, old_noise:Tensor) -> Tensor:
return (x - x_noise_swapped) / sigma + old_noise
def update_latent_swap_noise(self, x:Tensor, sigma:Tensor, old_noise:Tensor, new_noise:Tensor) -> Tensor:
return x + sigma * (new_noise - old_noise)
File diff suppressed because it is too large Load Diff
+887
View File
@@ -0,0 +1,887 @@
import torch
import torch.nn.functional as F
from typing import Optional, Callable, Tuple, Dict, Any, Union, TYPE_CHECKING, TypeVar, List
import re
import functools
import copy
from comfy.samplers import SCHEDULER_NAMES
from .res4lyf import RESplain
# EXTRA_OPTIONS OPS
class ExtraOptions():
def __init__(self, extra_options):
self.extra_options = extra_options
self.mute = False
# debugMode 0: Follow self.mute only
# debugMode 1: Print with debug flag if not muted
# debugMode 2: Never print
def __call__(self, option, default=None, ret_type=None, match_all_flags=False, debugMode=0):
if isinstance(option, (tuple, list)):
if match_all_flags:
return all(self(single_option, default, ret_type) for single_option in option)
else:
return any(self(single_option, default, ret_type) for single_option in option)
if default is None: # get flag
pattern = rf"^(?:{re.escape(option)}\s*$|{re.escape(option)}=)"
return bool(re.search(pattern, self.extra_options, flags=re.MULTILINE))
elif ret_type is None:
ret_type = type(default)
if ret_type.__module__ != "builtins":
mod = __import__(default.__module__)
ret_type = lambda v: getattr(mod, v, None)
if ret_type == list:
pattern = rf"^{re.escape(option)}\s*=\s*([a-zA-Z0-9_.,+-]+)\s*$"
match = re.search(pattern, self.extra_options, flags=re.MULTILINE)
if match:
value = match.group(1)
if not self.mute and debugMode != 2:
RESplain("Set extra_option: ", option, "=", value, debug=True)
else:
value = default
if type(value) == str:
value = value.split(',')
if type(default[0]) == type:
ret_type = default[0]
else:
ret_type = type(default[0])
value = [ret_type(value[_]) for _ in range(len(value))]
else:
pattern = rf"^{re.escape(option)}\s*=\s*([a-zA-Z0-9_.+-]+)\s*$"
match = re.search(pattern, self.extra_options, flags=re.MULTILINE)
if match:
if ret_type == bool:
value_str = match.group(1).lower()
value = value_str in ("true", "1", "yes", "on")
else:
value = ret_type(match.group(1))
if not self.mute and debugMode != 2:
RESplain("Set extra_option: ", option, "=", value, debug=True)
else:
value = default
# if "mute_EO" is in extra_options, set mute to True
if "mute_EO" in self.extra_options:
self.set_mute(True)
return value
def set_mute(self, mute=True):
self.mute = mute
return self
def extra_options_flag(flag, extra_options):
pattern = rf"^(?:{re.escape(flag)}\s*$|{re.escape(flag)}=)"
return bool(re.search(pattern, extra_options, flags=re.MULTILINE))
def get_extra_options_kv(key, default, extra_options, ret_type=None):
ret_type = type(default) if ret_type is None else ret_type
pattern = rf"^{re.escape(key)}\s*=\s*([a-zA-Z0-9_.+-]+)\s*$"
match = re.search(pattern, extra_options, flags=re.MULTILINE)
if match:
value = match.group(1)
else:
value = default
return ret_type(value)
def get_extra_options_list(key, default, extra_options, ret_type=None):
default = [default] if type(default) != list else default
#ret_type = type(default) if ret_type is None else ret_type
ret_type = type(default[0]) if ret_type is None else ret_type
pattern = rf"^{re.escape(key)}\s*=\s*([a-zA-Z0-9_.,+-]+)\s*$"
match = re.search(pattern, extra_options, flags=re.MULTILINE)
if match:
value = match.group(1)
else:
value = default
if type(value) == str:
value = value.split(',')
value = [ret_type(value[_]) for _ in range(len(value))]
return value
class OptionsManager:
APPEND_OPTIONS = {"extra_options"}
def __init__(self, options=None, options_group=None, **kwargs):
self.options_list = []
if options is not None:
self.options_list.append(options)
# v3 Autogrow delivers chained options as {"options0": dict, "options1": dict, ...}.
if options_group:
self.options_list.extend(
v for v in options_group.values() if v is not None
)
# Legacy-named chain inputs ("options", "options 2", ...) land here via **kwargs.
for key, value in kwargs.items():
if key.startswith('options') and value is not None:
self.options_list.append(value)
self._merged_dict = None
def add_option(self, option):
"""Add a single options dictionary"""
if option is not None:
self.options_list.append(option)
self._merged_dict = None # invalidate cached merged options
@property
def merged(self):
"""Get merged options with proper priority handling"""
if self._merged_dict is None:
self._merged_dict = {}
special_string_options = {
key: [] for key in self.APPEND_OPTIONS
}
for options_dict in self.options_list:
if options_dict is not None:
for key, value in options_dict.items():
if key in self.APPEND_OPTIONS and value:
special_string_options[key].append(value)
elif isinstance(value, dict):
# Deep merge dictionaries
if key not in self._merged_dict:
self._merged_dict[key] = {}
if isinstance(self._merged_dict[key], dict):
self._deep_update(self._merged_dict[key], value)
else:
self._merged_dict[key] = value.copy()
# Special case for FrameWeightsManager
elif key == "frame_weights_mgr" and hasattr(value, "_weight_configs"):
if key not in self._merged_dict:
self._merged_dict[key] = copy.deepcopy(value)
else:
existing_mgr = self._merged_dict[key]
if hasattr(value, "device") and value.device != torch.device('cpu'):
existing_mgr.device = value.device
if hasattr(value, "dtype") and value.dtype != torch.float64:
existing_mgr.dtype = value.dtype
# Merge all weight_configs
if hasattr(value, "_weight_configs"):
for name, config in value._weight_configs.items():
config_kwargs = config.copy()
existing_mgr.add_weight_config(name, **config_kwargs)
else:
self._merged_dict[key] = value
# append special case string options (e.g. extra_options)
for key, value in special_string_options.items():
if value:
self._merged_dict[key] = "\n".join(value)
return self._merged_dict
def update(self, key_or_dict, value=None, append=False):
"""Update options with a single key-value pair or a dictionary"""
if value is not None or isinstance(key_or_dict, (str, list)):
# single key-value update
key_path = key_or_dict
if isinstance(key_path, str):
key_path = key_path.split('.')
update_dict = {}
current = update_dict
for i, key in enumerate(key_path[:-1]):
current[key] = {}
current = current[key]
current[key_path[-1]] = value
self.add_option(update_dict)
else:
# dictionary update
flat_updates = {}
def _flatten_dict(d, prefix=""):
for key, value in d.items():
full_key = f"{prefix}.{key}" if prefix else key
if isinstance(value, dict):
_flatten_dict(value, full_key)
else:
flat_updates[full_key] = value
_flatten_dict(key_or_dict)
for key_path, value in flat_updates.items():
self.update(key_path, value) # Recursive call
return self
def get(self, key, default=None):
return self.merged.get(key, default)
def _deep_update(self, target_dict, source_dict):
for key, value in source_dict.items():
if isinstance(value, dict) and key in target_dict and isinstance(target_dict[key], dict):
# recursive dict update
self._deep_update(target_dict[key], value)
else:
target_dict[key] = value
def __getitem__(self, key):
"""Allow dictionary-like access to options"""
return self.merged[key]
def __contains__(self, key):
"""Allow 'in' operator for options"""
return key in self.merged
def as_dict(self):
"""Return the merged options as a dictionary"""
return self.merged.copy()
def __bool__(self):
"""Return True if there are any options"""
return len(self.options_list) > 0 and any(opt is not None for opt in self.options_list)
def debug_print_options(self):
for i, options_dict in enumerate(self.options_list):
RESplain(f"Options {i}:", debug=True)
if options_dict is not None:
for key, value in options_dict.items():
RESplain(f" {key}: {value}", debug=True)
else:
RESplain(" None", "\n", debug=True)
# MISCELLANEOUS OPS
def has_nested_attr(obj, attr_path):
attrs = attr_path.split('.')
for attr in attrs:
if not hasattr(obj, attr):
return False
obj = getattr(obj, attr)
return True
def safe_get_nested(d, keys, default=None):
for key in keys:
if isinstance(d, dict):
d = d.get(key, default)
else:
return default
return d
class AlwaysTrueList:
def __contains__(self, item):
return True
def __iter__(self):
while True:
yield True # kapow
def parse_range_string(s):
if "all" in s:
return AlwaysTrueList()
result = []
for part in s.split(','):
part = part.strip()
if not part:
continue
val = float(part) if '.' in part else int(part)
result.append(val)
return result
def parse_range_string_int(s):
if "all" in s:
return AlwaysTrueList()
result = []
for part in s.split(','):
if '-' in part:
start, end = part.split('-')
result.extend(range(int(start), int(end) + 1))
elif part.strip() != '':
result.append(int(part))
return result
def parse_tile_sizes(tile_sizes: str):
"""
Converts multiline string like:
"1024,1024\n768,1344\n1344,768"
into:
[(1024, 1024), (768, 1344), (1344, 768)]
"""
return [tuple(map(int, line.strip().split(',')))
for line in tile_sizes.strip().splitlines()
if line.strip()]
# COMFY OPS
def is_video_model(model):
is_video_model = False
try :
is_video_model = 'video' in model.inner_model.inner_model.model_config.unet_config['image_model'] or \
'cosmos' in model.inner_model.inner_model.model_config.unet_config['image_model'] or \
'wan2' in model.inner_model.inner_model.model_config.unet_config['image_model'] or \
'ltxv' in model.inner_model.inner_model.model_config.unet_config['image_model'] or \
'ltxav' in model.inner_model.inner_model.model_config.unet_config['image_model']
except:
pass
return is_video_model
def is_RF_model(model):
from comfy import model_sampling
modelsampling = model.inner_model.inner_model.model_sampling
return isinstance(modelsampling, model_sampling.CONST)
def get_res4lyf_scheduler_list():
scheduler_names = SCHEDULER_NAMES.copy()
if "beta57" not in scheduler_names:
scheduler_names.append("beta57")
return scheduler_names
def move_to_same_device(*tensors):
if not tensors:
return tensors
device = tensors[0].device
return tuple(tensor.to(device) for tensor in tensors)
def conditioning_set_values(conditioning, values={}):
c = []
for t in conditioning:
n = [t[0], t[1].copy()]
for k in values:
n[1][k] = values[k]
c.append(n)
return c
def extract_cond_from_guider(guider, cond_type):
"""Extract `cond_type` (e.g. 'positive' or 'negative') conditioning from a guider's
original_conds, converting from the guider's internal {cross_attn: tensor, ...} per-cond
dict format into the standard [[tensor, dict], ...] conditioning list format. Returns
None if guider is None / has no original_conds / doesn't contain cond_type."""
if guider is None:
return None
if not hasattr(guider, 'original_conds') or guider.original_conds is None:
return None
cond_list = guider.original_conds.get(cond_type)
if cond_list is None:
return None
return [
[cond.get('cross_attn'), {k: v for k, v in cond.items() if k != 'cross_attn'}]
for cond in cond_list
]
# MISC OPS
def initialize_or_scale(tensor, value, steps):
if tensor is None:
return torch.full((steps,), value)
else:
return value * tensor
def pad_tensor_list_to_max_len(tensors: List[torch.Tensor], dim: int = -2) -> List[torch.Tensor]:
"""Zero-pad each tensor in `tensors` along `dim` up to their common maximum length."""
max_len = max(t.shape[dim] for t in tensors)
padded = []
for t in tensors:
cur = t.shape[dim]
if cur < max_len:
pad_shape = list(t.shape)
pad_shape[dim] = max_len - cur
zeros = torch.zeros(*pad_shape, dtype=t.dtype, device=t.device)
t = torch.cat((t, zeros), dim=dim)
padded.append(t)
return padded
class PrecisionTool:
def __init__(self, cast_type='fp64'):
self.cast_type = cast_type
def cast_tensor(self, func):
@functools.wraps(func)
def wrapper(*args, **kwargs):
if self.cast_type not in ['fp64', 'fp32', 'fp16']:
return func(*args, **kwargs)
target_device = None
for arg in args:
if torch.is_tensor(arg):
target_device = arg.device
break
if target_device is None:
for v in kwargs.values():
if torch.is_tensor(v):
target_device = v.device
break
# recursively zs_recast tensors in nested dictionaries
def cast_and_move_to_device(data):
if torch.is_tensor(data):
if self.cast_type == 'fp64':
return data.to(torch.float64).to(target_device)
elif self.cast_type == 'fp32':
return data.to(torch.float32).to(target_device)
elif self.cast_type == 'fp16':
return data.to(torch.float16).to(target_device)
elif isinstance(data, dict):
return {k: cast_and_move_to_device(v) for k, v in data.items()}
return data
new_args = [cast_and_move_to_device(arg) for arg in args]
new_kwargs = {k: cast_and_move_to_device(v) for k, v in kwargs.items()}
return func(*new_args, **new_kwargs)
return wrapper
def set_cast_type(self, new_value):
if new_value in ['fp64', 'fp32', 'fp16']:
self.cast_type = new_value
else:
self.cast_type = 'fp64'
precision_tool = PrecisionTool(cast_type='fp64')
class FrameWeightsManager:
def __init__(self):
self._weight_configs = {}
self._default_config = {
"frame_weights": None, # Tensor of weights if directly specified
"dynamics": "linear", # Function type for dynamic period
"schedule": "moderate_early", # Schedule type
"scale": 0.5, # Amount of change
"is_reversed": False, # Whether to reverse weights
"custom_string": None, # Per-configuration custom string
}
self.dtype = torch.float64
self.device = torch.device('cpu')
def set_device_and_dtype(self, device=None, dtype=None):
"""Set the device and dtype for generated weights"""
if device is not None:
self.device = device
if dtype is not None:
self.dtype = dtype
return self
def set_custom_weights(self, config_name, weights):
"""Set custom weights for a specific configuration"""
if config_name not in self._weight_configs:
self._weight_configs[config_name] = self._default_config.copy()
self._weight_configs[config_name]["frame_weights"] = weights
return self
def add_weight_config(self, name, **kwargs):
if name not in self._weight_configs:
self._weight_configs[name] = self._default_config.copy()
for key, value in kwargs.items():
if key in self._default_config:
self._weight_configs[name][key] = value
# ignore unknown parameters
return self
def get_weight_config(self, name):
if name not in self._weight_configs:
return None
return self._weight_configs[name].copy()
def get_frame_weights_by_name(self, name, num_frames, step=None):
config = self.get_weight_config(name)
if config is None:
return None
weights_tensor = self._generate_frame_weights(
num_frames,
config["dynamics"],
config["schedule"],
config["scale"],
config["is_reversed"],
config["frame_weights"],
step=step,
custom_string=config["custom_string"]
)
if config["custom_string"] is not None and config["custom_string"].strip() != "" and weights_tensor is not None:
# ensure that the custom_string has more than just lines that begin with non-numeric characters
custom_string = config["custom_string"].strip()
custom_string = re.sub(r"^[^0-9].*", "", custom_string, flags=re.MULTILINE)
custom_string = re.sub(r"^\s*$", "", custom_string, flags=re.MULTILINE)
if custom_string.strip() != "":
# If the custom_string is not empty, show the custom weights
formatted_weights = [f"{w:.2f}" for w in weights_tensor.tolist()]
RESplain(f"Custom '{name}' for step {step}: {formatted_weights}", debug=True)
elif weights_tensor is None:
weights_tensor = torch.ones(num_frames, dtype=self.dtype, device=self.device)
return weights_tensor
def _generate_custom_weights(self, num_frames, custom_string, step=None):
"""
Generate custom weights based on the provided frame weights from a string with one line per step.
Args:
num_frames: Number of frames to generate weights for
custom_string: The custom weights string to parse
step: Specific step to use (0-indexed). If None, uses the last line.
Features:
- Each line represents weights for one step
- Add *[multiplier] at the end of a line to scale those weights (e.g., "1.0, 0.8, 0.6*1.5")
- Include "interpolate" on its own line to interpolate each line to match num_frames
- Prefix line with the steps to apply it to (e.g. "0-5: 1.0, 0.8, 0.6")
Example:
0-5:1.0, 0.8, 0.6, 0.4, 0.2, 0.0
6-10:0.0, 0.2, 0.4, 0.6, 0.8, 1.0*1.5
11-30:0.0, 0.5, 1.0, 0.5, 0.0, 0.0*0.8
interpolate
"""
if custom_string is not None:
interpolate_frames = "interpolate" in custom_string
lines = custom_string.strip().split('\n')
lines = [line for line in lines if line.strip() and not line.strip().startswith("interp")]
if not lines:
return None
if step is not None:
matching_line = None
for line in lines:
# Check if line has a step range prefix
step_range_match = re.match(r'^(\d+)-(\d+):(.*)', line.strip())
if step_range_match:
start_step = int(step_range_match.group(1))
end_step = int(step_range_match.group(2))
if start_step <= step <= end_step:
matching_line = step_range_match.group(3).strip()
if matching_line is not None:
weights_str = matching_line
else:
# if no matching line, try to use the step number line or the last line
if step < len(lines):
line_index = step
else:
line_index = len(lines) - 1
if line_index < 0:
return None
weights_str = lines[line_index].strip()
if ":" in weights_str:
weights_str = weights_str.split(":", 1)[1].strip()
else:
# When no specific step is provided, use the last line
line_index = len(lines) - 1
weights_str = lines[line_index].strip()
if ":" in weights_str:
weights_str = weights_str.split(":", 1)[1].strip()
if not weights_str:
return None
multiplier = 1.0
if "*" in weights_str:
parts = weights_str.rsplit("*", 1)
if len(parts) == 2:
weights_str = parts[0].strip()
try:
multiplier = float(parts[1].strip())
except ValueError as e:
RESplain(f"Invalid multiplier format: {parts[1]}")
try:
weights = [float(w.strip()) for w in weights_str.split(',')]
weights_tensor = torch.tensor(weights, dtype=self.dtype, device=self.device)
if multiplier != 1.0:
weights_tensor = weights_tensor * multiplier
if interpolate_frames and len(weights_tensor) != num_frames:
if len(weights_tensor) > 1:
orig_positions = torch.linspace(0, 1, len(weights_tensor), dtype=self.dtype, device=self.device)
new_positions = torch.linspace(0, 1, num_frames, dtype=self.dtype, device=self.device)
weights_tensor = torch.nn.functional.interpolate(
weights_tensor.view(1, 1, -1),
size=num_frames,
mode='linear',
align_corners=True
).squeeze()
else:
# If only one weight, repeat it for all frames
weights_tensor = weights_tensor.repeat(num_frames)
else:
if len(weights_tensor) < num_frames:
# If fewer weights than frames, repeat the last weight
weights_tensor = torch.cat([
weights_tensor,
torch.full((num_frames - len(weights_tensor),), weights_tensor[-1],
dtype=self.dtype, device=self.device)
])
# Trim if too many weights
if len(weights_tensor) > num_frames:
weights_tensor = weights_tensor[:num_frames]
return weights_tensor
except (ValueError, IndexError) as e:
RESplain(f"Error parsing custom frame weights: {e}")
return None
return None
def _generate_frame_weights(self, num_frames, dynamics, schedule, scale, is_reversed, frame_weights, step=None, custom_string=None):
# Look for the multiplier= parameter in the custom string and store it as a float value
multiplier = None
rate_factor = None
start_change_factor = None
if custom_string is not None:
if "multiplier" in custom_string:
multiplier_match = re.search(r"multiplier\s*=\s*([0-9.]+)", custom_string)
if multiplier_match:
multiplier = float(multiplier_match.group(1))
# Remove the multiplier= from the custom string
custom_string = re.sub(r"multiplier\s*=\s*[0-9.]+", "", custom_string).strip()
RESplain(f"Custom multiplier detected: {multiplier}", debug=True)
if "rate_factor" in custom_string:
rate_factor_match = re.search(r"rate_factor\s*=\s*([0-9.]+)", custom_string)
if rate_factor_match:
rate_factor = float(rate_factor_match.group(1))
# Remove the rate_factor= from the custom string
custom_string = re.sub(r"rate_factor\s*=\s*[0-9.]+", "", custom_string).strip()
RESplain(f"Custom rate factor detected: {rate_factor}", debug=True)
if "start_change_factor" in custom_string:
start_change_factor_match = re.search(r"start_change_factor\s*=\s*([0-9.]+)", custom_string)
if start_change_factor_match:
start_change_factor = float(start_change_factor_match.group(1))
# Remove the start_change_factor= from the custom string
custom_string = re.sub(r"start_change_factor\s*=\s*[0-9.]+", "", custom_string).strip()
RESplain(f"Custom start change factor detected: {start_change_factor}", debug=True)
if custom_string is not None and custom_string.strip() != "" and step is not None:
custom_weights = self._generate_custom_weights(num_frames, custom_string, step)
if custom_weights is not None:
weights = custom_weights
weights = torch.flip(weights, [0]) if is_reversed else weights
return weights
else:
RESplain("custom frame weights failed to parse, doing the normal thing...", debug=True)
if rate_factor is None:
if "fast" in schedule:
rate_factor = 0.25
elif "slow" in schedule:
rate_factor = 1.0
else: # moderate
rate_factor = 0.5
if start_change_factor is None:
if "early" in schedule:
start_change_factor = 0.0
elif "late" in schedule:
start_change_factor = 0.2
else:
start_change_factor = 0.0
change_frames = max(round(num_frames * rate_factor), 2)
change_start = round(num_frames * start_change_factor)
low_value = 1.0 - scale
if frame_weights is not None:
weights = torch.cat([frame_weights, torch.full((num_frames,), frame_weights[-1])])
weights = weights[:num_frames]
else:
if dynamics == "constant":
weights = self._generate_constant_schedule(change_start, change_frames, low_value, num_frames)
elif dynamics == "linear":
weights = self._generate_linear_schedule(change_start, change_frames, low_value, num_frames)
elif dynamics == "ease_out":
weights = self._generate_easeout_schedule(change_start, change_frames, low_value, num_frames)
elif dynamics == "ease_in":
weights = self._generate_easein_schedule(change_start, change_frames, low_value, num_frames)
elif dynamics == "middle":
weights = self._generate_middle_schedule(change_start, change_frames, low_value, num_frames)
elif dynamics == "trough":
weights = self._generate_trough_schedule(change_start, change_frames, low_value, num_frames)
else:
raise ValueError(f"Invalid schedule: {dynamics}")
if multiplier is None:
multiplier = 1.0
weights = torch.flip(weights, [0]) if is_reversed else weights
weights = weights * multiplier
weights = torch.clamp(weights, min=0.0, max=(max(1.0, multiplier)))
weights = weights.to(dtype=self.dtype, device=self.device)
return weights
def _generate_constant_schedule(self, change_start, change_frames, low_value, num_frames):
"""constant schedule with the scale as the low weight"""
return torch.ones(num_frames) * low_value
def _generate_linear_schedule(self, change_start, change_frames, low_value, num_frames):
"""linear schedule from 1 to the low weight"""
weights = torch.linspace(1, low_value, change_frames)
weights = torch.cat([torch.full((change_start,), 1.0), weights])
weights = torch.cat([weights, torch.full((num_frames,), weights[-1])])
weights = weights[:num_frames]
return weights
def _generate_easeout_schedule(self, change_start, change_frames, low_value, num_frames, k=4.0):
"""exponential schedule from 1 to the low weight"""
change_frames = max(change_frames, 4)
t = torch.linspace(0, 1, change_frames, dtype=self.dtype, device=self.device)
weights = 1.0 - (1.0 - low_value) * (1.0 - torch.exp(-k * t))
weights = torch.cat([torch.full((change_start,), 1.0), weights])
weights = torch.cat([weights, torch.full((num_frames,), weights[-1])])
weights = weights[:num_frames]
return weights
def _generate_easein_schedule(self, change_start, change_frames, low_value, num_frames):
"""a monomial power schedule from 1 to the low weight"""
change_frames = max(change_frames, 4)
t = torch.linspace(0, 1, change_frames, dtype=self.dtype, device=self.device)
weights = 1 - (1 - low_value) * torch.pow(t, 2)
# Prepend with change_start frames of 1.0
weights = torch.cat([torch.full((change_start,), 1.0), weights])
total_frames_to_pad = num_frames - len(weights)
if (total_frames_to_pad > 1):
mid_value_between_low_value_and_second_to_last_value = (weights[-2] + low_value) / 2.0
weights[-1] = mid_value_between_low_value_and_second_to_last_value
# Fill remaining with final value
weights = torch.cat([weights, torch.full((num_frames,), weights[-1])])
weights = weights[:num_frames]
return weights
def _generate_middle_schedule(self, change_start, change_frames, low_value, num_frames):
"""gaussian middle peaking schedule from 1 to the low weight"""
change_frames = max(change_frames, 4)
t = torch.linspace(0, 1, change_frames, dtype=self.dtype, device=self.device)
weights = torch.exp(-0.5 * ((t - 0.5) / 0.2) ** 2)
weights = weights / torch.max(weights)
weights = low_value + (1 - low_value) * weights
total_frames_to_pad = num_frames - len(weights)
pad_left = total_frames_to_pad // 2
pad_right = total_frames_to_pad - pad_left
weights = torch.cat([torch.full((pad_left,), low_value), weights, torch.full((pad_right,), low_value)])
if change_start > 0:
# Pad the beginning with the first value, and truncate to num_frames
weights = torch.cat([torch.full((change_start,), low_value), weights])
weights = weights[:num_frames]
return weights
def _generate_trough_schedule(self, change_start, change_frames, low_value, num_frames):
"""
Trough schedule with both ends at 1 and the middle at the low weight.
When change_start > 0, creates asymmetry with shorter decay at beginning and longer at end.
"""
change_frames = max(change_frames, 4)
# Calculate sigma based on change_frames - controls overall decay rate
sigma = max(0.2, change_frames / num_frames)
if change_start == 0:
t = torch.linspace(-1, 1, num_frames, dtype=self.dtype, device=self.device)
else:
asymmetry_factor = min(0.5, change_start / num_frames)
split_point = 0.5 - asymmetry_factor
first_size = int(split_point * num_frames)
first_size = max(1, first_size) # at least one frame
t1 = torch.linspace(-1, 0, first_size, dtype=self.dtype, device=self.device)
second_size = num_frames - first_size
t2 = torch.linspace(0, 1, second_size, dtype=self.dtype, device=self.device)
t = torch.cat([t1, t2])
# shape using Gaussian function
trough = 1.0 - torch.exp(-0.5 * (t / sigma) ** 2)
weights = low_value + (1.0 - low_value) * trough
return weights
def check_projection_consistency(x, W, b):
W_pinv = torch.linalg.pinv(W.T)
x_proj = (x - b) @ W_pinv
x_recon = x_proj @ W.T + b
error = torch.norm(x - x_recon)
in_subspace = error < 1e-3
return error, in_subspace
def get_max_dtype(device='cpu'):
if torch.backends.mps.is_available():
MAX_DTYPE = torch.float32
else:
try:
torch.tensor([0.0], dtype=torch.float64, device=device)
MAX_DTYPE = torch.float64
except (RuntimeError, TypeError):
MAX_DTYPE = torch.float32
return MAX_DTYPE
File diff suppressed because it is too large Load Diff
+26
View File
@@ -0,0 +1,26 @@
"""Provide RES4LYF solver logging without registering its ComfyUI nodes."""
from __future__ import annotations
import logging
LOGGER = logging.getLogger("SimpleSyrup.RES4LYF")
def RESplain(*values: object, debug: bool = False, **_kwargs: object) -> None:
"""Log solver diagnostics at the upstream-requested verbosity."""
message = " ".join(str(value) for value in values)
LOGGER.log(logging.DEBUG if debug else logging.INFO, message)
def is_debug_logging_enabled() -> bool:
"""Report whether solver debug diagnostics are enabled."""
return LOGGER.isEnabledFor(logging.DEBUG)
def get_display_sampler_category() -> bool:
"""Keep upstream sampler names stable without UI category mutation."""
return False
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -26,6 +26,14 @@ EXPECTED_RUNTIME_REQUIREMENTS = (
"yapf",
"huggingface-hub",
"keyring",
"mpmath",
"pywavelets",
)
EXPECTED_RELEASE_IDENTITY = (
"GIT_AUTHOR_NAME: Daisy",
"GIT_AUTHOR_EMAIL: daisy@artificialsweetener.ai",
"GIT_COMMITTER_NAME: Daisy",
"GIT_COMMITTER_EMAIL: daisy@artificialsweetener.ai",
)
@@ -124,3 +132,14 @@ def test_frontend_dist_bundle_is_tracked_for_comfy_serving() -> None:
dist_bundle = REPO_ROOT / "web" / "dist" / "simple-syrup.js"
assert dist_bundle.is_file()
def test_release_automation_uses_daisy_git_identity() -> None:
"""Keep automated release commits attributed to the project account."""
workflow = (REPO_ROOT / ".github" / "workflows" / "release.yml").read_text(
encoding="utf-8"
)
assert all(identity in workflow for identity in EXPECTED_RELEASE_IDENTITY)
assert "semantic-release-bot@martynus.net" not in workflow
@@ -18,6 +18,11 @@ from support.repository import REPOSITORY_ROOT
BASE_NODE_IDS = [
"SimpleSyrup.AllPromptAttentionSEGS",
"SimpleSyrup.AttentionCouplingOptions",
"SimpleSyrup.ContextualDiffusionOptions",
"SimpleSyrup.NoiseInversionOptions",
"SimpleSyrup.TilingOptions",
"SimpleSyrup.KSampler",
"SimpleSyrup.AttentionCaptureModel",
"SimpleSyrup.AttentionMaskedConditioning",
"SimpleSyrup.AttentionRegionMask",
@@ -32,6 +32,9 @@ FORBIDDEN_PATCHER_CALLS = frozenset(
FORBIDDEN_PATCHER_WRITES = frozenset({"forced_hooks", "use_clip_schedule"})
APPROVED_VALUE_CLONES = Counter(
{
("simple_syrup/domain/inversion_solver.py", "source"): 1,
("simple_syrup/runtime/noise_inversion.py", "zero"): 1,
("simple_syrup/runtime/noise_inversion.py", "noise"): 1,
("simple_syrup/domain/semantic_tiled_diffusion.py", "mask"): 2,
("simple_syrup/image/crop_composite.py", "image"): 1,
(
@@ -0,0 +1,98 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Verify optional inversion widgets on implementation-backed V3 sampler nodes."""
from __future__ import annotations
from typing import Any, ClassVar
import pytest
from simple_syrup.domain.noise_inversion import NoiseInversionOptions
from simple_syrup.nodes_v3.legacy_inversion_node_adapter import (
LegacyInversionNodeV3Adapter,
)
from simple_syrup.nodes_v3.legacy_node_wrappers import (
DetailSEGSAsRegionsV3,
DetailSEGSByScaleFactorTiledDiffusionV3,
)
class RecordingImplementation:
"""Expose a minimal maintained declaration and the execution boundary."""
RETURN_TYPES = ("STRING",)
OUTPUT_TOOLTIPS = ("Result.",)
FUNCTION = "run"
CATEGORY = "SimpleSyrup/Test"
DESCRIPTION = "Records sampling inputs."
INPUT_IS_LIST = False
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
"""Keep the persisted input contract independent of inversion widgets."""
return {"required": {"text": ("STRING", {"default": "", "tooltip": "Text."})}}
def run(
self, text: object, noise_inversion: NoiseInversionOptions | None = None
) -> tuple[object, NoiseInversionOptions | None]:
"""Return delegated inputs without invoking Comfy neural execution."""
return text, noise_inversion
class RecordingAdapter(LegacyInversionNodeV3Adapter):
"""Use the production schema and list-mode normalization implementation."""
LEGACY_NODE_CLASS: ClassVar[type[Any]] = RecordingImplementation
NODE_ID = "SimpleSyrup.Recording"
DISPLAY_NAME = "Recording"
@pytest.mark.parametrize("list_mode", [False, True])
def test_inversion_controls_normalize_and_default_on(
monkeypatch: pytest.MonkeyPatch,
list_mode: bool,
) -> None:
"""Default to the accepted recipe and normalize zero-step disable in list mode."""
monkeypatch.setattr(RecordingImplementation, "INPUT_IS_LIST", list_mode)
text = ["prompt"] if list_mode else "prompt"
output, enabled = RecordingAdapter.execute(text=text)
assert output == text and enabled == NoiseInversionOptions()
_, disabled = RecordingAdapter.execute(
text=text,
inversion_steps=[0] if list_mode else 0,
)
assert disabled is None
@pytest.mark.parametrize(
"controls",
[
{"inversion_steps": -1},
{"inversion_resolution_scale": 0},
{"inversion_method": "fireflow"},
{"inversion_switch_fraction": 1},
],
)
def test_invalid_selected_inversion_controls_fail(controls: dict[str, Any]) -> None:
"""Apply domain validation instead of handing malformed controls to a sampler."""
with pytest.raises(ValueError):
RecordingAdapter.execute(text="prompt", **controls)
@pytest.mark.parametrize(
"node",
[DetailSEGSAsRegionsV3, DetailSEGSByScaleFactorTiledDiffusionV3],
)
def test_existing_nodes_append_optional_inversion_without_reordering(node: Any) -> None:
"""Preserve sampler socket ordering before five optional inversion controls."""
schema = node.define_schema()
order = node.WORKFLOW_INPUT_ORDER
assert [item.id for item in schema.inputs[: len(order)]] == list(order)
new = schema.inputs[len(order) :]
assert len(new) == 5 and all(item.optional and item.tooltip for item in new)
assert new[0].id == "inversion_method" and new[0].default == "euler"
steps = next(item for item in new if item.id == "inversion_steps")
assert steps.default == 2 and steps.min == 0
@@ -8,7 +8,7 @@ from __future__ import annotations
from typing import Any, ClassVar
from simple_syrup.nodes_v3.legacy_node_wrappers import LegacyNodeV3Adapter
from simple_syrup.nodes_v3.legacy_node_adapter import LegacyNodeV3Adapter
class _FakeHidden:
@@ -0,0 +1,182 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Safeguard Krea NegPiP at both supported Comfy model call boundaries."""
from __future__ import annotations
from typing import cast
import pytest
import torch
from comfy.ldm.krea2.model import Attention, SingleStreamDiT
from comfy.model_patcher import ModelPatcher
from simple_syrup.runtime.negpip.krea2 import (
CONDITION_MASK_KEY,
TRANSFORMER_MASK_KEY,
krea2_attn1_negpip,
krea2_diffusion_negpip_wrapper,
)
from simple_syrup.runtime.negpip.krea2_host import (
Krea2HostAttention,
krea2_host_mutations,
)
from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE
@pytest.mark.parametrize("reference_argument", (False, True))
def test_diffusion_wrapper_preserves_both_host_signatures(
reference_argument: bool,
) -> None:
"""Inject one local mask without binding transformer options twice."""
source_options: dict[str, object] = {"existing": True}
mask = torch.tensor([[[-1.0], [1.0]]])
def earlier(
x: object,
timesteps: object,
context: object,
attention_mask: object,
transformer_options: dict[str, object],
**kwargs: object,
) -> dict[str, object]:
"""Expose Comfy 0.28's exact five-positional model boundary."""
del x, timesteps, context, attention_mask, kwargs
return transformer_options
def current(
x: object,
timesteps: object,
context: object,
attention_mask: object,
ref_latents: object,
transformer_options: dict[str, object],
**kwargs: object,
) -> dict[str, object]:
"""Expose Comfy's reference-capable six-positional model boundary."""
del x, timesteps, context, attention_mask, kwargs
assert ref_latents is None
return transformer_options
args = (
(None, None, None, None, None, source_options)
if reference_argument
else (None, None, None, None, source_options)
)
prepared = cast(
dict[str, object],
krea2_diffusion_negpip_wrapper(
current if reference_argument else earlier,
*args,
**{CONDITION_MASK_KEY: mask},
),
)
assert prepared is not source_options
assert prepared[TRANSFORMER_MASK_KEY] is mask
assert prepared["existing"] is True
assert TRANSFORMER_MASK_KEY not in source_options
@pytest.mark.parametrize("negative_sign", (False, True))
def test_host_attention_matches_native_math_and_keeps_call_options_local(
negative_sign: bool,
) -> None:
"""Earlier host hooks reproduce native grouped attention without source mutation."""
torch.manual_seed(4)
attention = Attention(128, 4, kvheads=2, operations=torch.nn)
with torch.no_grad():
attention.qknorm.qnorm.scale.zero_()
attention.qknorm.knorm.scale.zero_()
x = torch.randn(1, 7, 128)
mask = torch.tensor([[[-1.0 if negative_sign else 1.0], [1.0]]])
options = {
TRANSFORMER_MASK_KEY: mask,
"patches": {"attn1_patch": [krea2_attn1_negpip]},
}
native_options = {**options, "block_index": 0, "img_slice": [2, 7]}
with torch.no_grad():
expected = attention(x, transformer_options=native_options)
actual = Krea2HostAttention(attention, 0, 28)(x, transformer_options=options)
assert torch.equal(actual, expected)
assert "block_index" not in options
assert "img_slice" not in options
def test_host_attention_rejects_malformed_token_mask() -> None:
"""Malformed sign metadata must fail before executing attention projections."""
with pytest.raises(ValueError, match="processed token sign mask"):
Krea2HostAttention(None, 0, 28)(
torch.zeros(1, 3, 4), transformer_options={TRANSFORMER_MASK_KEY: "invalid"}
)
def test_host_attention_delegates_unsigned_calls() -> None:
"""Unmarked conditioning retains the host's ordinary attention operation."""
torch.manual_seed(4)
attention = Attention(128, 4, kvheads=2, operations=torch.nn)
with torch.no_grad():
attention.qknorm.qnorm.scale.zero_()
attention.qknorm.knorm.scale.zero_()
x = torch.randn(1, 7, 128)
expected = attention(x)
actual = Krea2HostAttention(attention, 0, 28)(x)
assert torch.equal(actual, expected)
class _EarlierModel(torch.nn.Module):
"""Expose the earlier host boundary around one real Comfy attention module."""
def __init__(self) -> None:
"""Retain a genuine patchable module hierarchy without loading weights."""
super().__init__()
block = torch.nn.Module()
block.attn = Attention(128, 4, kvheads=2, operations=torch.nn)
self.blocks = torch.nn.ModuleList([block])
def _forward(
self,
x: object,
timesteps: object,
context: object,
attention_mask: object = None,
transformer_options: object = None,
**kwargs: object,
) -> object:
"""Represent only the external host argument contract under test."""
del timesteps, context, attention_mask, transformer_options, kwargs
return x
def test_earlier_host_object_patches_restore_on_unload() -> None:
"""Patch only a derived MODEL and restore the shared host method on unload."""
root = torch.nn.Module()
root.diffusion_model = _EarlierModel()
device = torch.device("cpu")
source = ModelPatcher(root, load_device=device, offload_device=device)
attention = root.diffusion_model.blocks[0].attn
assert isinstance(attention, torch.nn.Module)
original = attention.forward
derived = PATCHER_LIFECYCLE.derive_model(
source, krea2_host_mutations(source), operation="test Krea host adaptation"
)
assert source.object_patches == {}
assert attention.forward == original
derived.patch_model(load_weights=False)
try:
assert isinstance(attention.forward, Krea2HostAttention)
finally:
derived.unpatch_model(unpatch_weights=False)
assert attention.forward == original
def test_reference_capable_host_needs_no_object_replacements() -> None:
"""Leave native Krea attention and its reference-image path entirely untouched."""
root = torch.nn.Module()
diffusion = object.__new__(SingleStreamDiT)
torch.nn.Module.__init__(diffusion)
root.diffusion_model = diffusion
device = torch.device("cpu")
model = ModelPatcher(root, load_device=device, offload_device=device)
assert krea2_host_mutations(model) == ()
@@ -11,6 +11,7 @@ from typing import Any, cast
import pytest
import torch
from comfy.ldm.krea2.model import SingleStreamDiT
from comfy.model_base import SDXL, Anima, BaseModel, Krea2, SDXLRefiner
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import WrappersMP
@@ -161,6 +162,10 @@ def _model_patcher(model_class: type[BaseModel]) -> ModelPatcher:
model = object.__new__(model_class)
torch.nn.Module.__init__(model)
if model_class is Krea2:
diffusion = object.__new__(SingleStreamDiT)
torch.nn.Module.__init__(diffusion)
model.diffusion_model = diffusion
device = torch.device("cpu")
return ModelPatcher(model, load_device=device, offload_device=device)
@@ -82,7 +82,7 @@ def test_full_context_service_delegates_prepared_model_to_ordinary_sampler() ->
seed=5,
steps=12,
cfg=1.0,
sampler_name="euler",
sampler_name="exponential/ddim",
scheduler="simple",
positive=positive,
negative=negative,
@@ -108,6 +108,7 @@ def test_full_context_service_delegates_prepared_model_to_ordinary_sampler() ->
assert sampler_call["positive"] == "base+"
assert sampler_call["negative"] == "base-"
assert sampler_call["latent_image"] is latent
assert sampler_call["sampler_name"] == "exponential/ddim"
def test_ordinary_request_bypasses_preparation_and_preserves_img2img_inputs(
@@ -157,6 +158,7 @@ def test_ordinary_request_bypasses_preparation_and_preserves_img2img_inputs(
"negative": "negative",
"latent_image": latent,
"denoise": 0.42,
"noise_inversion": None,
}
]
@@ -92,7 +92,7 @@ def test_contextual_attention_service_prepares_once_and_delegates_local_views(
seed=9,
steps=12,
cfg=1.0,
sampler_name="er_sde",
sampler_name="exponential/ddim",
scheduler="simple",
positive="regional-positive",
negative="regional-negative",
@@ -136,5 +136,6 @@ def test_contextual_attention_service_prepares_once_and_delegates_local_views(
assert call["latent_image"] is latent
assert call["segs"] == "segs"
assert call["diffusion_mode"] == diffusion_mode
assert call["sampler_name"] == "exponential/ddim"
request = cast(RegionalFeatureRequest, call["feature_request"])
assert request.features == frozenset({RegionalFeature.ATTENTION_COUPLING})
@@ -102,7 +102,7 @@ def test_tiled_attention_service_prepares_once_and_delegates_all_tiling(
seed=9,
steps=12,
cfg=1.0,
sampler_name="er_sde",
sampler_name="exponential/ddim",
scheduler="simple",
positive=positive,
negative=negative,
@@ -140,6 +140,7 @@ def test_tiled_attention_service_prepares_once_and_delegates_all_tiling(
assert request.features == frozenset({RegionalFeature.ATTENTION_COUPLING})
assert call["latent_tile_batch_size"] == 4
assert call["diffusion_mode"] == diffusion_mode
assert call["sampler_name"] == "exponential/ddim"
def test_ordinary_request_bypasses_preparation_and_preserves_tiled_img2img(
@@ -70,15 +70,21 @@ def test_schemas_expose_exact_names_and_shared_regional_contract() -> None:
"region_mask_feather",
"latent_image",
"denoise",
"inversion_method",
"inversion_resolution_scale",
"inversion_steps",
"inversion_switch_fraction",
"inversion_finishing_steps",
]
assert tiled_ids[: len(normal_ids)] == normal_ids
assert tiled_ids[len(normal_ids) :] == [
assert tiled_ids[:13] == normal_ids[:13]
assert tiled_ids[13:18] == [
"diffusion_mode",
"latent_tile_width",
"latent_tile_height",
"latent_tile_overlap",
"latent_tile_batch_size",
]
assert tiled_ids[18:] == normal_ids[13:]
regional_weight = normal.inputs[9]
assert regional_weight.default == 0.5
assert regional_weight.min == 0.0
@@ -95,6 +101,11 @@ def test_schemas_expose_exact_names_and_shared_regional_contract() -> None:
"regional_prompt_weight": 0.5,
"region_mask_feather": 0,
"denoise": 1.0,
"inversion_method": "euler",
"inversion_resolution_scale": 0.5,
"inversion_steps": 2,
"inversion_switch_fraction": 0.75,
"inversion_finishing_steps": 1,
}
tiled_defaults = {
value.id: value.default for value in tiled.inputs if hasattr(value, "default")
@@ -6,6 +6,8 @@
from __future__ import annotations
from dataclasses import replace
import pytest
import torch
@@ -88,6 +90,27 @@ def test_invalid_overlap_fails_before_planning() -> None:
)
def test_rectangular_tiles_do_not_resize_global_context() -> None:
"""Let Tiling configure the sole local plan independently of global context."""
controls = replace(_controls(), latent_tile_width=48, latent_tile_height=16)
plan = build_contextual_diffusion_plan(
latent_width=96, latent_height=64, controls=controls, segs=None
)
assert (plan.tile_plan.tile_width, plan.tile_plan.tile_height) == (48, 16)
assert (plan.global_view.model_width, plan.global_view.model_height) == (32, 22)
def test_overlap_is_bounded_by_both_explicit_tile_dimensions() -> None:
"""Reject an overlap that prevents advancing the shorter local dimension."""
controls = replace(
_controls(latent_context_overlap=16),
latent_tile_width=48,
latent_tile_height=16,
)
with pytest.raises(ValueError, match="both local tile dimensions"):
controls.validate()
def _controls(
*,
latent_context_overlap: int = 8,
+5
View File
@@ -0,0 +1,5 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Provide scoped sampling-boundary fixtures without replacing domain behavior."""
+40
View File
@@ -0,0 +1,40 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Represent the external MODEL patcher surface for real spatial wrapper tests."""
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import torch
from simple_syrup.runtime.sampling_model_types import ModelFunctionWrapper
class SpatialModel:
"""Keep independent model options and an explicit lifecycle parent."""
def __init__(self, options: dict[str, Any] | None = None) -> None:
"""Initialize the minimal host state needed by wrapper derivation."""
self.model_options = {} if options is None else options
self.load_device = torch.device("cpu")
self.parent: SpatialModel | None = None
def clone(self) -> SpatialModel:
"""Copy host options and retain the direct source identity."""
derived = SpatialModel(self.model_options.copy())
derived.parent = self
return derived
def set_model_unet_function_wrapper(self, wrapper: ModelFunctionWrapper) -> None:
"""Implement the public Comfy wrapper installation boundary."""
self.model_options["model_function_wrapper"] = wrapper
def set_model_sampler_calc_cond_batch_function(
self, function: Callable[..., object]
) -> None:
"""Implement the external Comfy regional prediction registration boundary."""
self.model_options["sampler_calc_cond_batch_function"] = function

Some files were not shown because too many files have changed in this diff Show More