Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
05667f6517 | ||
|
|
ab7cebcebc | ||
|
|
e4eabfefd5 | ||
|
|
bbe86060c6 | ||
|
|
be4bd9bb16 | ||
|
|
913cc7ba55 | ||
|
|
04dc35af00 | ||
|
|
7b1efef47e | ||
|
|
ec49b685db | ||
|
|
2ae545d64d |
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,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,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"
|
||||
|
||||
Generated
+2
-2
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -6,6 +6,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__version__ = "1.11.0"
|
||||
__version__ = "1.13.0"
|
||||
|
||||
__all__: list[str] = ["__version__"]
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
),
|
||||
)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.")
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)))
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
),
|
||||
)
|
||||
@@ -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))
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
+1095
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
+4099
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)
|
||||
|
||||
|
||||
+3
-1
@@ -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,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
+2
-1
@@ -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})
|
||||
|
||||
+2
-1
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
@@ -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
Reference in New Issue
Block a user