feat(sampling): add RES4LYF sampler methods and schedules
This commit is contained in:
@@ -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,6 +1,6 @@
|
||||
schema_version = 1
|
||||
review_by = 2027-03-31
|
||||
fingerprint = "sha256:4b78e0caf8bf4f90d5a7a46ff32b28b01c94c034efc3ddbc24cfa16b8e132ea0"
|
||||
fingerprint = "sha256:15ead72c26b1b5911aac191d07f36e5d8fd22d583447c2fcb62353be4ad1cd47"
|
||||
|
||||
cohesive_paths = [
|
||||
"simple_syrup/masking/prompt_segs_with_sam_service.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"
|
||||
|
||||
@@ -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,7 @@ addict>=2.4.0
|
||||
yapf>=0.43.0
|
||||
huggingface-hub>=0.34.0
|
||||
keyring>=25.0.0
|
||||
mpmath>=1.3.0
|
||||
einops>=0.8.0
|
||||
PyWavelets>=1.6.0
|
||||
scipy>=1.11.0
|
||||
|
||||
@@ -21,7 +21,7 @@ from ..domain.regional_features import (
|
||||
RegionalCapabilityAdmission,
|
||||
)
|
||||
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 .guided_sampling import sample_with_optional_negative
|
||||
from .model_patcher_mutations import ModelUnetWrapperMutation
|
||||
@@ -116,7 +116,14 @@ def sample_contextual_diffusion(
|
||||
diffusion_mode=diffusion_mode,
|
||||
)
|
||||
batch_inds = latent_image.get("batch_index")
|
||||
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,
|
||||
)
|
||||
callback = _latent_preview().prepare_callback(sampling_model, steps)
|
||||
samples = sample_with_optional_negative(
|
||||
comfy_sample=comfy_sample,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -25,7 +25,7 @@ from ..domain.tiled_diffusion import (
|
||||
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,
|
||||
@@ -135,7 +135,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(
|
||||
|
||||
@@ -25,7 +25,7 @@ from ..domain.tiled_diffusion import (
|
||||
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,
|
||||
@@ -137,7 +137,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(
|
||||
|
||||
@@ -18,7 +18,7 @@ from typing import Any, cast
|
||||
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,
|
||||
@@ -135,7 +135,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(
|
||||
|
||||
@@ -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,7 +12,7 @@ from typing import Any, TypeAlias
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
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 ..shared.logging import get_logger
|
||||
@@ -67,10 +67,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)
|
||||
|
||||
@@ -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
+2
-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(
|
||||
|
||||
+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(
|
||||
|
||||
@@ -102,6 +102,8 @@ def test_sampler_options_include_lcm() -> None:
|
||||
|
||||
assert "lcm" in sampler_options
|
||||
assert "euler_a_a1111" in sampler_options
|
||||
assert "exponential/ddim" in sampler_options
|
||||
assert "fully_implicit/radau_iia_3s" in sampler_options
|
||||
|
||||
|
||||
def test_scheduler_options_include_extras_and_exclude_svd() -> None:
|
||||
@@ -113,6 +115,7 @@ def test_scheduler_options_include_extras_and_exclude_svd() -> None:
|
||||
assert "AYS SDXL" in scheduler_options
|
||||
assert "GITS" in scheduler_options
|
||||
assert "beta57" in scheduler_options
|
||||
assert "bong_tangent" in scheduler_options
|
||||
assert "automatic_a1111" in scheduler_options
|
||||
assert "Flux2" in scheduler_options
|
||||
assert "AYS SVD" not in scheduler_options
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
"""Verify RES4LYF selection survives contextual and tiled model composition."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.contextual_diffusion import (
|
||||
ContextualDiffusionControls,
|
||||
build_contextual_diffusion_plan,
|
||||
)
|
||||
from simple_syrup.runtime import (
|
||||
contextual_diffusion_sampling,
|
||||
mixture_of_diffusers_sampling,
|
||||
multidiffusion_sampling,
|
||||
sampling_schedulers,
|
||||
)
|
||||
from simple_syrup.runtime.sampling_noise import prepare_sampling_noise
|
||||
|
||||
|
||||
class _Model:
|
||||
"""Provide the patcher and sigma bounds used at each composed boundary."""
|
||||
|
||||
def __init__(
|
||||
self, model_options: dict[str, Any] | None = None, parent: _Model | None = None
|
||||
) -> None:
|
||||
"""Create a CPU model with copyable wrapper options."""
|
||||
|
||||
self.load_device = torch.device("cpu")
|
||||
self.model_options = {} if model_options is None else model_options
|
||||
self.sampling = SimpleNamespace(sigma_max=1.0, sigma_min=0.0)
|
||||
self.parent = parent
|
||||
|
||||
def clone(self) -> _Model:
|
||||
"""Preserve model sampling and copy wrapper options."""
|
||||
|
||||
cloned = _Model(self.model_options.copy(), parent=self)
|
||||
cloned.sampling = self.sampling
|
||||
return cloned
|
||||
|
||||
def set_model_unet_function_wrapper(self, wrapper: object) -> None:
|
||||
"""Record the contextual or tiled prediction wrapper."""
|
||||
|
||||
self.model_options["model_function_wrapper"] = wrapper
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the same sigma bounds before and after model cloning."""
|
||||
|
||||
assert name == "model_sampling"
|
||||
return self.sampling
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mode",
|
||||
["contextual", "multidiffusion", "mixture_of_diffusers"],
|
||||
)
|
||||
def test_res4lyf_sampler_reaches_composed_sampling_boundary(
|
||||
monkeypatch: pytest.MonkeyPatch, mode: str
|
||||
) -> None:
|
||||
"""Pass the real solver, RES noise, and wrapped model into ComfyUI."""
|
||||
|
||||
model = _Model()
|
||||
latent_samples = torch.zeros((1, 4, 16, 32))
|
||||
latent = {"samples": latent_samples}
|
||||
runtime = {
|
||||
"contextual": contextual_diffusion_sampling,
|
||||
"multidiffusion": multidiffusion_sampling,
|
||||
"mixture_of_diffusers": mixture_of_diffusers_sampling,
|
||||
}[mode]
|
||||
comfy_sample = runtime._comfy_sample()
|
||||
preview = runtime._latent_preview()
|
||||
sigma_calls: list[dict[str, object]] = []
|
||||
sampled: list[dict[str, object]] = []
|
||||
|
||||
def calculate_sigmas(**kwargs: object) -> torch.Tensor:
|
||||
"""Capture the RES4LYF scheduler and return a short finite schedule."""
|
||||
|
||||
sigma_calls.append(kwargs)
|
||||
return torch.tensor([1.0, 0.5, 0.0])
|
||||
|
||||
def sample_custom(
|
||||
sampling_model: _Model,
|
||||
noise: torch.Tensor,
|
||||
cfg: float,
|
||||
sampler: object,
|
||||
sigmas: torch.Tensor,
|
||||
positive: object,
|
||||
negative: object,
|
||||
latent_image: torch.Tensor,
|
||||
**kwargs: object,
|
||||
) -> torch.Tensor:
|
||||
"""Inspect the arguments at ComfyUI's composed sampling boundary."""
|
||||
|
||||
del cfg, sigmas, positive, negative, kwargs
|
||||
sampled.append({"model": sampling_model, "noise": noise, "sampler": sampler})
|
||||
return latent_image + 1
|
||||
|
||||
monkeypatch.setattr(sampling_schedulers, "calculate_sigmas", calculate_sigmas)
|
||||
monkeypatch.setattr(
|
||||
comfy_sample,
|
||||
"fix_empty_latent_channels",
|
||||
lambda _model, samples, _ratio: samples,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
comfy_sample,
|
||||
"prepare_noise",
|
||||
lambda *_args: pytest.fail("RES4LYF must generate its own initial noise"),
|
||||
)
|
||||
monkeypatch.setattr(comfy_sample, "sample_custom", sample_custom)
|
||||
monkeypatch.setattr(preview, "prepare_callback", lambda _model, _steps: None)
|
||||
|
||||
controls = ContextualDiffusionControls(16, 0, 2, 1.0, 1, 0.5)
|
||||
plan = build_contextual_diffusion_plan(
|
||||
latent_width=32,
|
||||
latent_height=16,
|
||||
controls=controls,
|
||||
segs=None,
|
||||
)
|
||||
common = {
|
||||
"model": model,
|
||||
"seed": 42,
|
||||
"steps": 2,
|
||||
"cfg": 1.0,
|
||||
"sampler_name": "exponential/ddim",
|
||||
"scheduler": "bong_tangent",
|
||||
"positive": [],
|
||||
"negative": [],
|
||||
"latent_image": latent,
|
||||
"denoise": 1.0,
|
||||
}
|
||||
if mode == "contextual":
|
||||
sample_runtime: Callable[..., dict[str, Any]] = (
|
||||
contextual_diffusion_sampling.sample_contextual_diffusion
|
||||
)
|
||||
extra = {"diffusion_mode": "multidiffusion", "controls": controls, "plan": plan}
|
||||
else:
|
||||
sample_runtime = (
|
||||
multidiffusion_sampling.sample_multidiffusion
|
||||
if mode == "multidiffusion"
|
||||
else mixture_of_diffusers_sampling.sample_mixture_of_diffusers
|
||||
)
|
||||
extra = {
|
||||
"latent_tile_width": 16,
|
||||
"latent_tile_height": 16,
|
||||
"latent_tile_overlap": 0,
|
||||
"latent_tile_batch_size": 2,
|
||||
}
|
||||
output = sample_runtime(**common, **extra)
|
||||
|
||||
assert torch.equal(output["samples"], latent_samples + 1)
|
||||
assert sigma_calls[0]["scheduler_name"] == "bong_tangent"
|
||||
assert sigma_calls[0]["sampler_name"] == "exponential/ddim"
|
||||
assert len(sampled) == 1
|
||||
derived_model = sampled[0]["model"]
|
||||
assert isinstance(derived_model, _Model)
|
||||
assert derived_model is not model
|
||||
assert callable(derived_model.model_options["model_function_wrapper"])
|
||||
sampler = sampled[0]["sampler"]
|
||||
assert vars(sampler)["extra_options"]["rk_type"] == "ddim"
|
||||
expected_noise = prepare_sampling_noise(
|
||||
comfy_sample=comfy_sample,
|
||||
sampler_name="exponential/ddim",
|
||||
samples=latent_samples,
|
||||
seed=42,
|
||||
batch_indices=None,
|
||||
model=derived_model,
|
||||
)
|
||||
torch.testing.assert_close(sampled[0]["noise"], expected_noise, rtol=0, atol=0)
|
||||
@@ -0,0 +1,238 @@
|
||||
"""Prove pinned RES4LYF method registration and solver binding parity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime.res4lyf_sampler_names import RES4LYF_SAMPLER_NAMES
|
||||
from simple_syrup.runtime.sampling_noise import prepare_sampling_noise
|
||||
from simple_syrup.runtime.sampling_samplers import available_samplers, resolve_sampler
|
||||
from simple_syrup.runtime.sampling_schedulers import (
|
||||
available_schedulers,
|
||||
calculate_sigmas,
|
||||
)
|
||||
from simple_syrup.third_party.res4lyf_runtime import sigmas as upstream_sigmas
|
||||
from simple_syrup.third_party.res4lyf_runtime.beta.rk_coefficients_beta import (
|
||||
RK_SAMPLER_NAMES_BETA_FOLDERS,
|
||||
process_sampler_name,
|
||||
)
|
||||
|
||||
SOURCE_REVISION = "3d1d69da69ee47f7647d59e1bd0967e472fccc41"
|
||||
SOLVER_SOURCE_SHA256 = (
|
||||
"96c237a234da7486b597c12cf23f15229403698b6a6cb00d5d47ebdf249a5c51"
|
||||
)
|
||||
SOLVER_SOURCE_FILES = (
|
||||
"helper.py",
|
||||
"latents.py",
|
||||
"style_transfer.py",
|
||||
"sigmas.py",
|
||||
"beta/constants.py",
|
||||
"beta/deis_coefficients.py",
|
||||
"beta/phi_functions.py",
|
||||
"beta/rk_coefficients_beta.py",
|
||||
"beta/rk_method_beta.py",
|
||||
"beta/noise_classes.py",
|
||||
"beta/rk_noise_sampler_beta.py",
|
||||
"beta/rk_guide_func_beta.py",
|
||||
"beta/rk_sampler_beta.py",
|
||||
)
|
||||
|
||||
|
||||
def test_pinned_solver_source_matches_upstream_with_one_recorded_fix() -> None:
|
||||
"""Guard the solver copy and its single upstream error correction."""
|
||||
|
||||
root = (
|
||||
Path(__file__).resolve().parents[2]
|
||||
/ "simple_syrup"
|
||||
/ "third_party"
|
||||
/ "res4lyf_runtime"
|
||||
)
|
||||
digest = hashlib.sha256()
|
||||
for relative_path in SOLVER_SOURCE_FILES:
|
||||
digest.update(relative_path.encode())
|
||||
source = (root / relative_path).read_bytes()
|
||||
if relative_path == "beta/rk_coefficients_beta.py":
|
||||
branch = source.index(b'case "res_8s_alt"')
|
||||
corrected = b" ci = [c1, c2, c3, c4, c5, c6, c7, c8]"
|
||||
original = b" #ci = [c1, c2, c3, c4, c5, c6, c7, c8]"
|
||||
assert source[branch:].count(corrected) >= 1
|
||||
source = source[:branch] + source[branch:].replace(corrected, original, 1)
|
||||
digest.update(source)
|
||||
assert digest.hexdigest() == SOLVER_SOURCE_SHA256
|
||||
|
||||
|
||||
def test_res4lyf_menu_covers_every_upstream_solver_method() -> None:
|
||||
"""Expose each pinned upstream menu method exactly once."""
|
||||
|
||||
assert RES4LYF_SAMPLER_NAMES == tuple(RK_SAMPLER_NAMES_BETA_FOLDERS[1:])
|
||||
assert len(RES4LYF_SAMPLER_NAMES) == 118
|
||||
assert len(set(RES4LYF_SAMPLER_NAMES)) == 118
|
||||
assert Counter(name.split("/", 1)[0] for name in RES4LYF_SAMPLER_NAMES) == {
|
||||
"exponential": 41,
|
||||
"fully_implicit": 30,
|
||||
"linear": 20,
|
||||
"multistep": 10,
|
||||
"hybrid": 9,
|
||||
"diag_implicit": 8,
|
||||
}
|
||||
assert set(RES4LYF_SAMPLER_NAMES).issubset(available_samplers())
|
||||
|
||||
|
||||
def test_each_res4lyf_menu_name_binds_the_original_solver_and_method() -> None:
|
||||
"""Keep the method and implicit selection identical to RES4LYF."""
|
||||
|
||||
for name in RES4LYF_SAMPLER_NAMES:
|
||||
sampler = resolve_sampler(name)
|
||||
expected_method, expected_implicit = process_sampler_name(name)
|
||||
assert vars(sampler)["sampler_function"].__name__ == "_sample_res4lyf"
|
||||
assert vars(sampler)["extra_options"] == {
|
||||
"rk_type": expected_method,
|
||||
"implicit_sampler_name": expected_implicit,
|
||||
"implicit_type": "bongmath",
|
||||
"implicit_type_substeps": "bongmath",
|
||||
"bongmath": name != "linear/rk5_7s",
|
||||
}
|
||||
|
||||
|
||||
def test_bong_tangent_uses_original_res4lyf_sigma_function() -> None:
|
||||
"""Preserve RES4LYF's schedule values and local dropdown entry."""
|
||||
|
||||
class Model:
|
||||
"""Provide the scheduler's model-sampling reference."""
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the model-sampling object for sigma generation."""
|
||||
|
||||
assert name == "model_sampling"
|
||||
return SimpleNamespace()
|
||||
|
||||
model = Model()
|
||||
expected = upstream_sigmas.bong_tangent_scheduler(
|
||||
model.get_model_object("model_sampling"), 8
|
||||
)
|
||||
actual = calculate_sigmas(model, "bong_tangent", "euler", 8, 1.0)
|
||||
assert "bong_tangent" in available_schedulers()
|
||||
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
|
||||
|
||||
|
||||
def test_res4lyf_default_noise_matches_upstream_reference() -> None:
|
||||
"""Match upstream Gaussian generation and channel normalization exactly."""
|
||||
|
||||
class Model:
|
||||
"""Provide the same model sigma bounds as the upstream reference."""
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the model-sampling bounds."""
|
||||
|
||||
assert name == "model_sampling"
|
||||
return SimpleNamespace(sigma_max=1.0, sigma_min=0.0)
|
||||
|
||||
noise = prepare_sampling_noise(
|
||||
comfy_sample=SimpleNamespace(),
|
||||
sampler_name="exponential/ddim",
|
||||
samples=torch.zeros((1, 4, 8, 8)),
|
||||
seed=42,
|
||||
batch_indices=None,
|
||||
model=Model(),
|
||||
)
|
||||
assert hashlib.sha256(noise.numpy().tobytes()).hexdigest() == (
|
||||
"c4e80cfdfbbd88ad1dacfeacf19698188a4870ee2b17122d3e461001f6dce8d5"
|
||||
)
|
||||
|
||||
|
||||
def test_core_sampler_keeps_comfy_noise_path() -> None:
|
||||
"""Leave ComfyUI's seed and batch-index behavior intact for core methods."""
|
||||
|
||||
samples = torch.zeros((1, 4, 8, 8))
|
||||
expected = torch.ones_like(samples)
|
||||
calls: list[tuple[torch.Tensor, int, list[int]]] = []
|
||||
|
||||
class Model:
|
||||
"""Provide the sampling protocol without using it for core noise."""
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Reject model lookups on the core ComfyUI noise path."""
|
||||
|
||||
raise AssertionError(name)
|
||||
|
||||
def prepare_noise(
|
||||
input_samples: torch.Tensor, seed: int, batch_indices: list[int]
|
||||
) -> torch.Tensor:
|
||||
"""Record the arguments delegated to ComfyUI."""
|
||||
|
||||
calls.append((input_samples, seed, batch_indices))
|
||||
return expected
|
||||
|
||||
actual = prepare_sampling_noise(
|
||||
comfy_sample=SimpleNamespace(prepare_noise=prepare_noise),
|
||||
sampler_name="euler",
|
||||
samples=samples,
|
||||
seed=42,
|
||||
batch_indices=[3],
|
||||
model=Model(),
|
||||
)
|
||||
assert actual is expected
|
||||
assert calls == [(samples, 42, [3])]
|
||||
|
||||
|
||||
def test_every_res4lyf_method_executes_with_finite_latents() -> None:
|
||||
"""Exercise every advertised solver on a controlled two-step flow model."""
|
||||
|
||||
class ModelSampling:
|
||||
"""Provide the sigma bounds and denoising rule used by the solver."""
|
||||
|
||||
sigma_max = torch.tensor(1.0)
|
||||
sigma_min = torch.tensor(0.01)
|
||||
|
||||
def calculate_denoised(
|
||||
self, sigma: torch.Tensor, eps: torch.Tensor, x: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Return a linear test denoising prediction."""
|
||||
|
||||
return x - sigma * eps
|
||||
|
||||
class Model:
|
||||
"""Expose the ComfyUI model shape expected by the pinned solver."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create a small model surface without external weights."""
|
||||
|
||||
inner = SimpleNamespace(
|
||||
device=torch.device("cpu"),
|
||||
model_sampling=ModelSampling(),
|
||||
diffusion_model=SimpleNamespace(),
|
||||
)
|
||||
self.inner_model = SimpleNamespace(inner_model=inner)
|
||||
|
||||
def __call__(
|
||||
self, x: torch.Tensor, sigma: torch.Tensor, **kwargs: object
|
||||
) -> torch.Tensor:
|
||||
"""Return a deterministic prediction for every solver stage."""
|
||||
|
||||
return x * 0.5
|
||||
|
||||
model = Model()
|
||||
input_latent = torch.ones((1, 4, 8, 8))
|
||||
sigmas = torch.tensor([1.0, 0.5, 0.0])
|
||||
for name in RES4LYF_SAMPLER_NAMES:
|
||||
sampler = resolve_sampler(name)
|
||||
sample = vars(sampler)["sampler_function"]
|
||||
output = sample(
|
||||
model,
|
||||
input_latent.clone(),
|
||||
sigmas,
|
||||
extra_args={
|
||||
"seed": 42,
|
||||
"model_options": {"transformer_options": {}},
|
||||
},
|
||||
callback=None,
|
||||
disable=True,
|
||||
**vars(sampler)["extra_options"],
|
||||
)
|
||||
assert output.shape == input_latent.shape, name
|
||||
assert torch.isfinite(output).all(), name
|
||||
@@ -41,7 +41,7 @@ def test_available_samplers_includes_core_and_extras() -> None:
|
||||
assert samplers[: len(comfy.samplers.KSampler.SAMPLERS)] == tuple(
|
||||
comfy.samplers.KSampler.SAMPLERS
|
||||
)
|
||||
assert samplers[-1] == "euler_a_a1111"
|
||||
assert samplers[len(comfy.samplers.KSampler.SAMPLERS)] == "euler_a_a1111"
|
||||
|
||||
|
||||
def test_available_samplers_deduplicates_local_extra_when_globally_patched(
|
||||
|
||||
@@ -13,7 +13,7 @@ import comfy.samplers
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime import sampling_schedulers
|
||||
from simple_syrup.runtime import sampling_reference_schedules, sampling_schedulers
|
||||
from simple_syrup.runtime.sampling_schedulers import (
|
||||
calculate_sigmas,
|
||||
)
|
||||
@@ -153,7 +153,7 @@ def reference_full_extra_schedule(scheduler_name: str, steps: int) -> torch.Tens
|
||||
def reference_ays_schedule(model_type: str, steps: int) -> torch.Tensor:
|
||||
"""Calculate full AYS schedule with Comfy Extras formula semantics."""
|
||||
|
||||
sigmas = list(sampling_schedulers.AYS_NOISE_LEVELS[model_type])
|
||||
sigmas = list(sampling_reference_schedules.AYS_NOISE_LEVELS[model_type])
|
||||
if (steps + 1) != len(sigmas):
|
||||
sigmas = reference_loglinear_interpolate(sigmas, steps + 1)
|
||||
sigmas[-1] = 0.0
|
||||
@@ -164,10 +164,10 @@ def reference_gits_schedule(steps: int) -> torch.Tensor:
|
||||
"""Calculate full GITS schedule for the default coefficient."""
|
||||
|
||||
if steps <= 20:
|
||||
sigmas = list(sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[steps - 2])
|
||||
sigmas = list(sampling_reference_schedules.GITS_DEFAULT_NOISE_LEVELS[steps - 2])
|
||||
else:
|
||||
sigmas = reference_loglinear_interpolate(
|
||||
sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[-1],
|
||||
sampling_reference_schedules.GITS_DEFAULT_NOISE_LEVELS[-1],
|
||||
steps + 1,
|
||||
)
|
||||
sigmas[-1] = 0.0
|
||||
|
||||
@@ -15,7 +15,7 @@ import comfy.samplers
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime import sampling_schedulers
|
||||
from simple_syrup.runtime import sampling_reference_schedules, sampling_schedulers
|
||||
from simple_syrup.runtime.sampling_schedulers import (
|
||||
SchedulerView,
|
||||
available_schedulers,
|
||||
@@ -157,7 +157,7 @@ def reference_full_extra_schedule(scheduler_name: str, steps: int) -> torch.Tens
|
||||
def reference_ays_schedule(model_type: str, steps: int) -> torch.Tensor:
|
||||
"""Calculate full AYS schedule with Comfy Extras formula semantics."""
|
||||
|
||||
sigmas = list(sampling_schedulers.AYS_NOISE_LEVELS[model_type])
|
||||
sigmas = list(sampling_reference_schedules.AYS_NOISE_LEVELS[model_type])
|
||||
if (steps + 1) != len(sigmas):
|
||||
sigmas = reference_loglinear_interpolate(sigmas, steps + 1)
|
||||
sigmas[-1] = 0.0
|
||||
@@ -168,10 +168,10 @@ def reference_gits_schedule(steps: int) -> torch.Tensor:
|
||||
"""Calculate full GITS schedule for the default coefficient."""
|
||||
|
||||
if steps <= 20:
|
||||
sigmas = list(sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[steps - 2])
|
||||
sigmas = list(sampling_reference_schedules.GITS_DEFAULT_NOISE_LEVELS[steps - 2])
|
||||
else:
|
||||
sigmas = reference_loglinear_interpolate(
|
||||
sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[-1],
|
||||
sampling_reference_schedules.GITS_DEFAULT_NOISE_LEVELS[-1],
|
||||
steps + 1,
|
||||
)
|
||||
sigmas[-1] = 0.0
|
||||
@@ -227,11 +227,12 @@ def test_available_schedulers_includes_core_and_extras() -> None:
|
||||
|
||||
for scheduler in comfy.samplers.KSampler.SCHEDULERS:
|
||||
assert scheduler in schedulers
|
||||
assert schedulers[-6:] == (
|
||||
assert schedulers[-7:] == (
|
||||
"AYS SD1",
|
||||
"AYS SDXL",
|
||||
"GITS",
|
||||
"beta57",
|
||||
"bong_tangent",
|
||||
"automatic_a1111",
|
||||
"Flux2",
|
||||
)
|
||||
|
||||
Vendored
+13
-49
@@ -1,75 +1,39 @@
|
||||
# Third-Party Notices
|
||||
|
||||
This repository vendors selected third-party behavior for runtime use. Each
|
||||
vendored component is recorded in `third_party/manifest.toml`, and the
|
||||
corresponding license text is stored in `third_party/licenses/`.
|
||||
This repository vendors selected third-party behavior for runtime use. Each vendored component is recorded in `third_party/manifest.toml`, and the corresponding license text is stored in `third_party/licenses/`.
|
||||
|
||||
## SAM-HQ and MobileSAM runtime
|
||||
|
||||
SimpleSyrup vendors selected SAM-HQ and MobileSAM runtime files under
|
||||
Apache-2.0 for loading SAM-family segmentation models inside ComfyUI. The
|
||||
vendored runtime is kept under `simple_syrup/third_party/sam_hq_runtime/`
|
||||
and preserves upstream copyright notices where the source files carried them.
|
||||
SimpleSyrup vendors selected SAM-HQ and MobileSAM runtime files under Apache-2.0 for loading SAM-family segmentation models inside ComfyUI. The vendored runtime is kept under `simple_syrup/third_party/sam_hq_runtime/` and preserves upstream copyright notices where the source files carried them.
|
||||
|
||||
## GroundingDINO runtime
|
||||
|
||||
SimpleSyrup vendors selected GroundingDINO runtime files under Apache-2.0 for
|
||||
prompt-based box detection. The vendored runtime is kept under
|
||||
`simple_syrup/third_party/groundingdino_runtime/` and preserves upstream
|
||||
copyright notices where the source files carried them.
|
||||
SimpleSyrup vendors selected GroundingDINO runtime files under Apache-2.0 for prompt-based box detection. The vendored runtime is kept under `simple_syrup/third_party/groundingdino_runtime/` and preserves upstream copyright notices where the source files carried them.
|
||||
|
||||
## RES4LYF beta57 scheduler preset
|
||||
## RES4LYF samplers and schedulers
|
||||
|
||||
SimpleSyrup vendors the RES4LYF `beta57` scheduler preset under AGPL-3.0.
|
||||
The preset uses ComfyUI's beta scheduler with `alpha=0.5` and `beta=0.7`.
|
||||
SimpleSyrup resolves the preset locally for `KSampler (Extras)` and does not
|
||||
patch ComfyUI's global scheduler registry.
|
||||
ClownsharkBatwing and the RES4LYF contributors developed the Runge-Kutta and exponential solver methods included here. SimpleSyrup includes the solver modules and `bong_tangent` scheduler from revision `3d1d69da69ee47f7647d59e1bd0967e472fccc41` under AGPL-3.0. The `beta57` preset, recorded from revision `1c9bf61`, uses ComfyUI's beta scheduler with `alpha=0.5` and `beta=0.7`. SimpleSyrup exposes these methods and schedules through its KSampler inputs without registering RES4LYF nodes or changing ComfyUI's global sampler or scheduler registries. Source paths and copied files are recorded in the manifest.
|
||||
|
||||
The `res_8s_alt` coefficient branch initializes `ci` before using it; the pinned upstream branch leaves that assignment commented out and raises `UnboundLocalError`. The `linear/rk5_7s` menu option runs with Bongmath disabled because the pinned upstream default produces non-finite values on an ordinary two-step sigma grid. The local RES4LYF license copy includes the upstream LICENSE's commercial-service paragraph followed by the GNU AGPL version 3 text.
|
||||
|
||||
## AUTOMATIC1111 sampling integration
|
||||
|
||||
SimpleSyrup vendors selected AUTOMATIC1111 WebUI sampler integration behavior
|
||||
under AGPL-3.0. This provenance covers the `Euler a` sampler mapping, the
|
||||
`Automatic` scheduler fallback behavior, and seed-variation interpolation at
|
||||
ComfyUI's model-carried outer sampling boundary. SimpleSyrup does not port
|
||||
AUTOMATIC1111 ENSD or global RNG hijacking behavior.
|
||||
SimpleSyrup vendors selected AUTOMATIC1111 WebUI sampler integration behavior under AGPL-3.0. This provenance covers the `Euler a` sampler mapping, the `Automatic` scheduler fallback behavior, and seed-variation interpolation at ComfyUI's model-carried outer sampling boundary. SimpleSyrup does not port AUTOMATIC1111 ENSD or global RNG hijacking behavior.
|
||||
|
||||
## k-diffusion Euler ancestral sampler
|
||||
|
||||
SimpleSyrup vendors the Euler ancestral sampler loop from k-diffusion under
|
||||
the MIT license. The local sampler keeps the A1111/k-diffusion loop structure
|
||||
while running inside ComfyUI's deterministic seed system.
|
||||
SimpleSyrup vendors the Euler ancestral sampler loop from k-diffusion under the MIT license. The local sampler keeps the A1111/k-diffusion loop structure while running inside ComfyUI's deterministic seed system.
|
||||
|
||||
## Mixture of Diffusers and MultiDiffusion tiled diffusion behavior
|
||||
|
||||
SimpleSyrup reimplements tiled denoising behavior after
|
||||
inspecting the local `multidiffusion-upscaler-for-automatic1111` extension,
|
||||
which is licensed under CC-BY-NC-SA-4.0. The implementation preserves the
|
||||
extension's latent tile planning, Mixture Gaussian tile weighting,
|
||||
MultiDiffusion uniform tile averaging, regional prompt mask blending, and
|
||||
pre-CFG model prediction blending behavior without importing the extension at
|
||||
runtime.
|
||||
SimpleSyrup reimplements tiled denoising behavior after inspecting the local `multidiffusion-upscaler-for-automatic1111` extension, which is licensed under CC-BY-NC-SA-4.0. The implementation preserves the extension's latent tile planning, Mixture Gaussian tile weighting, MultiDiffusion uniform tile averaging, regional prompt mask blending, and pre-CFG model prediction blending behavior without importing the extension at runtime.
|
||||
|
||||
## SmilingWolf WD tagger models
|
||||
|
||||
SimpleSyrup's `Tile & Tag SEGS` node can download selected SmilingWolf WD
|
||||
tagger ONNX models and `selected_tags.csv` files at runtime from Hugging Face.
|
||||
These model files are not vendored in this repository. The runtime catalog
|
||||
points to the corresponding `SmilingWolf/*` repositories and stores downloaded
|
||||
files in the user's ComfyUI model directory.
|
||||
SimpleSyrup's `Tile & Tag SEGS` node can download selected SmilingWolf WD tagger ONNX models and `selected_tags.csv` files at runtime from Hugging Face. These model files are not vendored in this repository. The runtime catalog points to the corresponding `SmilingWolf/*` repositories and stores downloaded files in the user's ComfyUI model directory.
|
||||
|
||||
## NegPiP prompt weighting
|
||||
|
||||
SimpleSyrup adapts the AGPL-3.0 NegPiP implementation from
|
||||
`pamparamm/ComfyUI-ppm` at revision
|
||||
`6c6c360155cace9d7091306c1b8e26d9c7438620`. The standard SD1/SDXL and Anima
|
||||
paths preserve PPM's ModelPatcher-based magnitude-key and signed-value
|
||||
behavior. The Krea 2 path extends the same signed-value rule to Krea's layered
|
||||
Qwen conditioning and joint text/image attention while preserving its native
|
||||
conditioning shape.
|
||||
SimpleSyrup adapts the AGPL-3.0 NegPiP implementation from `pamparamm/ComfyUI-ppm` at revision `6c6c360155cace9d7091306c1b8e26d9c7438620`. The standard SD1/SDXL and Anima paths preserve PPM's ModelPatcher-based magnitude-key and signed-value behavior. The Krea 2 path extends the same signed-value rule to Krea's layered Qwen conditioning and joint text/image attention while preserving its native conditioning shape.
|
||||
|
||||
PPM credits the original ComfyUI port to
|
||||
`laksjdjf/cd-tuner_negpip-ComfyUI`; SimpleSyrup records revision
|
||||
`938b838546cf774dc8841000996552cef52cccf3`. That port credits the original
|
||||
Automatic1111 WebUI implementation in `hako-mikan/sd-webui-negpip`;
|
||||
SimpleSyrup records revision
|
||||
`fb7151f327ae56195f08b30b70d459493dadedbb`.
|
||||
PPM credits the original ComfyUI port to `laksjdjf/cd-tuner_negpip-ComfyUI`; SimpleSyrup records revision `938b838546cf774dc8841000996552cef52cccf3`. That port credits the original Automatic1111 WebUI implementation in `hako-mikan/sd-webui-negpip`; SimpleSyrup records revision `fb7151f327ae56195f08b30b70d459493dadedbb`.
|
||||
|
||||
Vendored
+37
@@ -75,6 +75,43 @@ vendored_files = [
|
||||
"simple_syrup/runtime/sampling_schedulers.py",
|
||||
]
|
||||
|
||||
[[component]]
|
||||
name = "RES4LYF sampler methods"
|
||||
license = "AGPL-3.0"
|
||||
license_file = "third_party/licenses/res4lyf.LICENSE.txt"
|
||||
source = "https://github.com/ClownsharkBatwing/RES4LYF"
|
||||
revision = "3d1d69da69ee47f7647d59e1bd0967e472fccc41"
|
||||
source_paths = [
|
||||
"beta/constants.py",
|
||||
"beta/deis_coefficients.py",
|
||||
"beta/phi_functions.py",
|
||||
"beta/rk_coefficients_beta.py",
|
||||
"beta/rk_method_beta.py",
|
||||
"beta/noise_classes.py",
|
||||
"beta/rk_noise_sampler_beta.py",
|
||||
"beta/rk_guide_func_beta.py",
|
||||
"beta/rk_sampler_beta.py",
|
||||
"helper.py",
|
||||
"latents.py",
|
||||
"style_transfer.py",
|
||||
"sigmas.py",
|
||||
]
|
||||
vendored_files = [
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/constants.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/deis_coefficients.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/phi_functions.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/rk_coefficients_beta.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/rk_method_beta.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/noise_classes.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/rk_noise_sampler_beta.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/rk_guide_func_beta.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/beta/rk_sampler_beta.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/helper.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/latents.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/style_transfer.py",
|
||||
"simple_syrup/third_party/res4lyf_runtime/sigmas.py",
|
||||
]
|
||||
|
||||
[[component]]
|
||||
name = "AUTOMATIC1111 sampling integration"
|
||||
license = "AGPL-3.0"
|
||||
|
||||
Reference in New Issue
Block a user