diff --git a/README.md b/README.md index 513c80e..2145e5f 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/governance/architecture/soft_reviews.toml b/governance/architecture/soft_reviews.toml index ad3c94d..4e53f17 100644 --- a/governance/architecture/soft_reviews.toml +++ b/governance/architecture/soft_reviews.toml @@ -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", diff --git a/governance/architecture/waivers.toml b/governance/architecture/waivers.toml index 8d06866..e54b72b 100644 --- a/governance/architecture/waivers.toml +++ b/governance/architecture/waivers.toml @@ -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" diff --git a/pyproject.toml b/pyproject.toml index 0e306fb..c40780b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] diff --git a/requirements.txt b/requirements.txt index 00b4d6c..588ae2f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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 diff --git a/simple_syrup/runtime/contextual_diffusion_sampling.py b/simple_syrup/runtime/contextual_diffusion_sampling.py index bd69ee0..88212bb 100644 --- a/simple_syrup/runtime/contextual_diffusion_sampling.py +++ b/simple_syrup/runtime/contextual_diffusion_sampling.py @@ -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, diff --git a/simple_syrup/runtime/detail_sampling.py b/simple_syrup/runtime/detail_sampling.py index 0bac7ed..5a8d6d2 100644 --- a/simple_syrup/runtime/detail_sampling.py +++ b/simple_syrup/runtime/detail_sampling.py @@ -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) diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index 86ed1fa..a79d363 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -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( diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index f777ef3..6053987 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -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( diff --git a/simple_syrup/runtime/regional_multidiffusion_sampling.py b/simple_syrup/runtime/regional_multidiffusion_sampling.py index b967901..765fa87 100644 --- a/simple_syrup/runtime/regional_multidiffusion_sampling.py +++ b/simple_syrup/runtime/regional_multidiffusion_sampling.py @@ -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( diff --git a/simple_syrup/runtime/res4lyf_sampler_names.py b/simple_syrup/runtime/res4lyf_sampler_names.py new file mode 100644 index 0000000..eef2d22 --- /dev/null +++ b/simple_syrup/runtime/res4lyf_sampler_names.py @@ -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", +) diff --git a/simple_syrup/runtime/res4lyf_sampling.py b/simple_syrup/runtime/res4lyf_sampling.py new file mode 100644 index 0000000..48a7136 --- /dev/null +++ b/simple_syrup/runtime/res4lyf_sampling.py @@ -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 diff --git a/simple_syrup/runtime/sampling_noise.py b/simple_syrup/runtime/sampling_noise.py new file mode 100644 index 0000000..c1b1fd2 --- /dev/null +++ b/simple_syrup/runtime/sampling_noise.py @@ -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) diff --git a/simple_syrup/runtime/sampling_reference_schedules.py b/simple_syrup/runtime/sampling_reference_schedules.py new file mode 100644 index 0000000..31886f3 --- /dev/null +++ b/simple_syrup/runtime/sampling_reference_schedules.py @@ -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, + ), +) diff --git a/simple_syrup/runtime/sampling_samplers.py b/simple_syrup/runtime/sampling_samplers.py index 803e92a..adc7b55 100644 --- a/simple_syrup/runtime/sampling_samplers.py +++ b/simple_syrup/runtime/sampling_samplers.py @@ -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)) diff --git a/simple_syrup/runtime/sampling_schedulers.py b/simple_syrup/runtime/sampling_schedulers.py index 56218ab..da0ccbc 100644 --- a/simple_syrup/runtime/sampling_schedulers.py +++ b/simple_syrup/runtime/sampling_schedulers.py @@ -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, diff --git a/simple_syrup/services/ksampler_sampling_service.py b/simple_syrup/services/ksampler_sampling_service.py index dda68b6..df1a176 100644 --- a/simple_syrup/services/ksampler_sampling_service.py +++ b/simple_syrup/services/ksampler_sampling_service.py @@ -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) diff --git a/simple_syrup/third_party/res4lyf_runtime/__init__.py b/simple_syrup/third_party/res4lyf_runtime/__init__.py new file mode 100644 index 0000000..e208994 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/__init__.py @@ -0,0 +1 @@ +"""Preserve pinned RES4LYF solver code for local sampler execution.""" diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/__init__.py b/simple_syrup/third_party/res4lyf_runtime/beta/__init__.py new file mode 100644 index 0000000..0d98b76 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/__init__.py @@ -0,0 +1 @@ +"""Contain the pinned RES4LYF Runge-Kutta solver implementation.""" diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/constants.py b/simple_syrup/third_party/res4lyf_runtime/beta/constants.py new file mode 100644 index 0000000..71178c8 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/constants.py @@ -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" +] diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/deis_coefficients.py b/simple_syrup/third_party/res4lyf_runtime/beta/deis_coefficients.py new file mode 100644 index 0000000..8e2e4c0 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/deis_coefficients.py @@ -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 + diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/noise_classes.py b/simple_syrup/third_party/res4lyf_runtime/beta/noise_classes.py new file mode 100644 index 0000000..6420ef1 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/noise_classes.py @@ -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 diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/phi_functions.py b/simple_syrup/third_party/res4lyf_runtime/beta/phi_functions.py new file mode 100644 index 0000000..8016425 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/phi_functions.py @@ -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_) + diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/rk_coefficients_beta.py b/simple_syrup/third_party/res4lyf_runtime/beta/rk_coefficients_beta.py new file mode 100644 index 0000000..92374b7 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/rk_coefficients_beta.py @@ -0,0 +1,3333 @@ +import torch +from torch import Tensor + +import copy +import math +from mpmath import mp, mpf, factorial, exp +mp.dps = 80 +from typing import Optional, Callable, Tuple, Dict, Any, Union, TYPE_CHECKING, TypeVar + +from .deis_coefficients import get_deis_coeff_list +from .phi_functions import phi, Phi, calculate_gamma + +from ..helper import ExtraOptions, get_extra_options_kv, extra_options_flag + + +from itertools import permutations, combinations +import random + +from einops import rearrange, einsum +from ..res4lyf import get_display_sampler_category, RESplain + +# Samplers with free parameters (c1, c2, c3) +# 1 2 3 +# X res_2s +# X X res_3s +# X res_3s_alt +# X res_3s_strehmel_weiner +# X dpmpp_2s (dpmpp_sde_2s has c2=1.0) +# X X dpmpp_3s +# X X irk_exp_diag_2s + +RK_EXPONENTIAL_PREFIXES = ( + "res", + "dpmpp", + "ddim", + "pec", + "etdrk", + "lawson", + "abnorsett", + ) + +def is_exponential(rk_type:str) -> bool: + return rk_type.startswith(RK_EXPONENTIAL_PREFIXES) + +RK_SAMPLER_NAMES_BETA_FOLDERS = ["none", + "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", + #"verner_robust_16s", + + "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", + #"gauss-legendre_diag_8s", + + + "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", + + ] + + + +RK_SAMPLER_NAMES_BETA_NO_FOLDERS = [] +for orig_sampler_name in RK_SAMPLER_NAMES_BETA_FOLDERS[1:]: + sampler_name = orig_sampler_name.split("/")[-1] if "/" in orig_sampler_name else orig_sampler_name + RK_SAMPLER_NAMES_BETA_NO_FOLDERS.append(sampler_name) + +IRK_SAMPLER_NAMES_BETA_FOLDERS = ["none", "use_explicit"] +for orig_sampler_name in RK_SAMPLER_NAMES_BETA_FOLDERS[1:]: + if "implicit" in orig_sampler_name and "/" in orig_sampler_name: + IRK_SAMPLER_NAMES_BETA_FOLDERS.append(orig_sampler_name) + +IRK_SAMPLER_NAMES_BETA_NO_FOLDERS = [] +for orig_sampler_name in IRK_SAMPLER_NAMES_BETA_FOLDERS[1:]: + sampler_name = orig_sampler_name.split("/")[-1] if "/" in orig_sampler_name else orig_sampler_name + IRK_SAMPLER_NAMES_BETA_NO_FOLDERS.append(sampler_name) + +RK_SAMPLER_FOLDER_MAP = {} +for orig_sampler_name in RK_SAMPLER_NAMES_BETA_FOLDERS: + if "/" in orig_sampler_name: + folder, sampler_name = orig_sampler_name.rsplit("/", 1) + else: + folder = "" + sampler_name = orig_sampler_name + RK_SAMPLER_FOLDER_MAP[sampler_name] = folder + +IRK_SAMPLER_FOLDER_MAP = {} +for orig_sampler_name in IRK_SAMPLER_NAMES_BETA_FOLDERS: + if "/" in orig_sampler_name: + folder, sampler_name = orig_sampler_name.rsplit("/", 1) + else: + folder = "" + sampler_name = orig_sampler_name + IRK_SAMPLER_FOLDER_MAP[sampler_name] = folder + +class DualFormatList(list): + """list that can match items with or without category prefixes.""" + def __contains__(self, item): + if super().__contains__(item): + return True + + if isinstance(item, str) and "/" in item: + base_name = item.split("/")[-1] + return any(name.endswith(base_name) for name in self) + + return any(isinstance(opt, str) and opt.endswith("/" + item) for opt in self) + +def get_sampler_name_list(nameOnly = False) -> list: + sampler_name_list = [] + for sampler_name in RK_SAMPLER_FOLDER_MAP: + if get_display_sampler_category() and not nameOnly: + folder_name = RK_SAMPLER_FOLDER_MAP[sampler_name] + full_sampler_name = f"{folder_name}/{sampler_name}" + else: + full_sampler_name = sampler_name + if full_sampler_name[0] == "/": + full_sampler_name = full_sampler_name[1:] + sampler_name_list.append(full_sampler_name) + return DualFormatList(sampler_name_list) + +def get_default_sampler_name(nameOnly = False) -> str: + default_sampler_name = "res_2m" + #find the key associated with the default value + for sampler_name in RK_SAMPLER_FOLDER_MAP: + if sampler_name == default_sampler_name: + if get_display_sampler_category() and not nameOnly: + folder_name = RK_SAMPLER_FOLDER_MAP[sampler_name] + return f"{folder_name}/{default_sampler_name}" + else: + return default_sampler_name + return default_sampler_name + +def get_implicit_sampler_name_list(nameOnly = False) -> list: + implicit_sampler_name_list = [] + for sampler_name in IRK_SAMPLER_FOLDER_MAP: + if get_display_sampler_category() and not nameOnly: + folder_name = IRK_SAMPLER_FOLDER_MAP[sampler_name] + full_sampler_name = f"{folder_name}/{sampler_name}" + else: + full_sampler_name = sampler_name + if full_sampler_name[0] == "/": + full_sampler_name = full_sampler_name[1:] + implicit_sampler_name_list.append(full_sampler_name) + return DualFormatList(implicit_sampler_name_list) + +def get_default_implicit_sampler_name(nameOnly = False) -> str: + default_sampler_value = "explicit_diagonal" + #find the key associated with the default value + for sampler_name in IRK_SAMPLER_FOLDER_MAP: + if sampler_name == default_sampler_value: + if get_display_sampler_category() and not nameOnly: + folder_name = IRK_SAMPLER_FOLDER_MAP[sampler_name] + return f"{folder_name}/{default_sampler_value}" + else: + return default_sampler_value + return default_sampler_value + +def get_full_sampler_name(sampler_name_in: str) -> str: + if "/" in sampler_name_in and sampler_name_in[0] != "/": + return sampler_name_in + for sampler_name in RK_SAMPLER_FOLDER_MAP: + if sampler_name == sampler_name_in: + folder_name = RK_SAMPLER_FOLDER_MAP[sampler_name] + return f"{folder_name}/{sampler_name}" + return sampler_name + +def process_sampler_name(sampler_name_in): + processed_name = sampler_name_in.split("/")[-1] if "/" in sampler_name_in else sampler_name_in + full_sampler_name = get_full_sampler_name(sampler_name_in) + + if sampler_name_in.startswith("fully_implicit") or sampler_name_in.startswith("diag_implicit"): + implicit_sampler_name = processed_name + sampler_name = "euler" + else: + sampler_name = processed_name + implicit_sampler_name = "use_explicit" + + return sampler_name, implicit_sampler_name + + + + +alpha_crouzeix = (2/(3**0.5)) * math.cos(math.pi / 18) +gamma_crouzeix = (1/(3**0.5)) * math.cos(math.pi / 18) + 1/2 # Crouzeix & Raviart 1980; A-stable; pg 100 in Solving Ordinary Differential Equations II +delta_crouzeix = 1 / (6 * (2 * gamma_crouzeix - 1)**2) # Crouzeix & Raviart 1980; A-stable; pg 100 in Solving Ordinary Differential Equations II + +rk_coeff = { + "gauss-legendre_diag_8s": ( # https://github.com/SciML/IRKGaussLegendre.jl/blob/master/src/IRKCoefficients.jl Antoñana, M., Makazaga, J., Murua, Ander. "Reducing and monitoring round-off error propagation for symplectic implicit Runge-Kutta schemes." Numerical Algorithms. 2017. + [ + [ + 0.5, + 0,0,0,0,0,0,0, + ], + [ + 1.0818949631055814971365081647359309e00, + 0.5, + 0,0,0,0,0,0, + ], + [ + 9.5995729622205494766003095439844678e-01, + 1.0869589243008327233290709646162480e00, + 0.5, + 0,0,0,0,0, + ], + [ + 1.0247213458032003748680445816450829e00, + 9.5505887369737431186016905653386876e-01, + 1.0880938387323083134422138713913203e00, + 0.5, + 0,0,0,0, + ], + [ + 9.8302382676362890697311829123888390e-01, + 1.0287597754747493109782305570410685e00, + 9.5383453518519996588326911440754302e-01, + 1.0883471611098277842507073806008045e00, + 0.5, + 0,0,0, + ], + [ + 1.0122259141132982060539425317219435e00, + 9.7998287236359129082628958290257329e-01, + 1.0296038730649779374630125982121223e00, + 9.5383453518519996588326911440754302e-01, + 1.0880938387323083134422138713913203e00, + 0.5, + 0,0, + ], + [ + 9.9125143323080263118822334698608777e-01, + 1.0140743558891669291459735166525994e00, + 9.7998287236359129082628958290257329e-01, + 1.0287597754747493109782305570410685e00, + 9.5505887369737431186016905653386876e-01, + 1.0869589243008327233290709646162480e00, + 0.5, + 0, + ], + [ + 1.0054828082532158826793409353214951e00, + 9.9125143323080263118822334698608777e-01, + 1.0122259141132982060539425317219435e00, + 9.8302382676362890697311829123888390e-01, + 1.0247213458032003748680445816450829e00, + 9.5995729622205494766003095439844678e-01, + 1.0818949631055814971365081647359309e00, + 0.5, + ], + ], + [ + [ + 5.0614268145188129576265677154981094e-02, + 1.1119051722668723527217799721312045e-01, + 1.5685332293894364366898110099330067e-01, + 1.8134189168918099148257522463859781e-01, + 1.8134189168918099148257522463859781e-01, + 1.5685332293894364366898110099330067e-01, + 1.1119051722668723527217799721312045e-01, + 5.0614268145188129576265677154981094e-02,] + ], + [ + 1.9855071751231884158219565715263505e-02, # 0.019855071751231884158219565715263505 + 1.0166676129318663020422303176208480e-01, + 2.3723379504183550709113047540537686e-01, + 4.0828267875217509753026192881990801e-01, + 5.9171732124782490246973807118009203e-01, + 7.6276620495816449290886952459462321e-01, + 8.9833323870681336979577696823791522e-01, + 9.8014492824876811584178043428473653e-01, + ] + ), + + + "gauss-legendre_5s": ( + [ + [4563950663 / 32115191526, + (310937500000000 / 2597974476091533 + 45156250000 * (739**0.5) / 8747388808389), + (310937500000000 / 2597974476091533 - 45156250000 * (739**0.5) / 8747388808389), + (5236016175 / 88357462711 + 709703235 * (739**0.5) / 353429850844), + (5236016175 / 88357462711 - 709703235 * (739**0.5) / 353429850844)], + + [(4563950663 / 32115191526 - 38339103 * (739**0.5) / 6250000000), + (310937500000000 / 2597974476091533 + 9557056475401 * (739**0.5) / 3498955523355600000), + (310937500000000 / 2597974476091533 - 14074198220719489 * (739**0.5) / 3498955523355600000), + (5236016175 / 88357462711 + 5601362553163918341 * (739**0.5) / 2208936567775000000000), + (5236016175 / 88357462711 - 5040458465159165409 * (739**0.5) / 2208936567775000000000)], + + [(4563950663 / 32115191526 + 38339103 * (739**0.5) / 6250000000), + (310937500000000 / 2597974476091533 + 14074198220719489 * (739**0.5) / 3498955523355600000), + (310937500000000 / 2597974476091533 - 9557056475401 * (739**0.5) / 3498955523355600000), + (5236016175 / 88357462711 + 5040458465159165409 * (739**0.5) / 2208936567775000000000), + (5236016175 / 88357462711 - 5601362553163918341 * (739**0.5) / 2208936567775000000000)], + + [(4563950663 / 32115191526 - 38209 * (739**0.5) / 7938810), + (310937500000000 / 2597974476091533 - 359369071093750 * (739**0.5) / 70145310854471391), + (310937500000000 / 2597974476091533 - 323282178906250 * (739**0.5) / 70145310854471391), + (5236016175 / 88357462711 - 470139 * (739**0.5) / 1413719403376), + (5236016175 / 88357462711 - 44986764863 * (739**0.5) / 21205791050640)], + + [(4563950663 / 32115191526 + 38209 * (739**0.5) / 7938810), + (310937500000000 / 2597974476091533 + 359369071093750 * (739**0.5) / 70145310854471391), + (310937500000000 / 2597974476091533 + 323282178906250 * (739**0.5) / 70145310854471391), + (5236016175 / 88357462711 + 44986764863 * (739**0.5) / 21205791050640), + (5236016175 / 88357462711 + 470139 * (739**0.5) / 1413719403376)], + ], + [ + + [ + 4563950663 / 16057595763, + 621875000000000 / 2597974476091533, + 621875000000000 / 2597974476091533, + 10472032350 / 88357462711, + 10472032350 / 88357462711] + ], + [ + 1 / 2, + 1 / 2 - 99 * (739**0.5) / 10000, # smallest # 0.06941899716778028758987101075583196 + 1 / 2 + 99 * (739**0.5) / 10000, # largest + 1 / 2 - (739**0.5) / 60, + 1 / 2 + (739**0.5) / 60 + ] + ), + + "gauss-legendre_5s_ascending": ( + [ + [(4563950663 / 32115191526 - 38339103 * (739**0.5) / 6250000000), + (310937500000000 / 2597974476091533 + 9557056475401 * (739**0.5) / 3498955523355600000), + (310937500000000 / 2597974476091533 - 14074198220719489 * (739**0.5) / 3498955523355600000), + (5236016175 / 88357462711 + 5601362553163918341 * (739**0.5) / 2208936567775000000000), + (5236016175 / 88357462711 - 5040458465159165409 * (739**0.5) / 2208936567775000000000)], + + + [(4563950663 / 32115191526 - 38209 * (739**0.5) / 7938810), + (310937500000000 / 2597974476091533 - 359369071093750 * (739**0.5) / 70145310854471391), + (310937500000000 / 2597974476091533 - 323282178906250 * (739**0.5) / 70145310854471391), + (5236016175 / 88357462711 - 470139 * (739**0.5) / 1413719403376), + (5236016175 / 88357462711 - 44986764863 * (739**0.5) / 21205791050640)], + + [4563950663 / 32115191526, + (310937500000000 / 2597974476091533 + 45156250000 * (739**0.5) / 8747388808389), + (310937500000000 / 2597974476091533 - 45156250000 * (739**0.5) / 8747388808389), + (5236016175 / 88357462711 + 709703235 * (739**0.5) / 353429850844), + (5236016175 / 88357462711 - 709703235 * (739**0.5) / 353429850844)], + + + [(4563950663 / 32115191526 + 38209 * (739**0.5) / 7938810), + (310937500000000 / 2597974476091533 + 359369071093750 * (739**0.5) / 70145310854471391), + (310937500000000 / 2597974476091533 + 323282178906250 * (739**0.5) / 70145310854471391), + (5236016175 / 88357462711 + 44986764863 * (739**0.5) / 21205791050640), + (5236016175 / 88357462711 + 470139 * (739**0.5) / 1413719403376)], + + + [(4563950663 / 32115191526 + 38339103 * (739**0.5) / 6250000000), + (310937500000000 / 2597974476091533 + 14074198220719489 * (739**0.5) / 3498955523355600000), + (310937500000000 / 2597974476091533 - 9557056475401 * (739**0.5) / 3498955523355600000), + (5236016175 / 88357462711 + 5040458465159165409 * (739**0.5) / 2208936567775000000000), + (5236016175 / 88357462711 - 5601362553163918341 * (739**0.5) / 2208936567775000000000)], + ], + [ + + [621875000000000 / 2597974476091533, + 10472032350 / 88357462711, + + 4563950663 / 16057595763, + + 10472032350 / 88357462711, + 621875000000000 / 2597974476091533,] + ], + [ + 1 / 2 - 99 * (739**0.5) / 10000, # smallest # 0.06941899716778028758987101075583196 + 1 / 2 - (739**0.5) / 60, + 1 / 2, + + + 1 / 2 + (739**0.5) / 60, + + 1 / 2 + 99 * (739**0.5) / 10000, # largest + ] + ), + "gauss-legendre_4s_alt": ( # https://ijstre.com/Publish/072016/371428231.pdf Four Point Gauss Quadrature Runge – Kuta Method Of Order 8 For Ordinary Differential Equations + [ + [1633/18780 - 71*206**0.5/96717000, + 134689/939000 - 927*206**0.5/78250, + 171511/939000 - 927*206**0.5/78250, + 1633/18780 - 121979*206**0.5/19343400,], + [7623/78250 - 1629507*206**0.5/257912000, + 347013/21284000, + -118701/4256800, + 7623/78250 + 1629507*206**0.5/257912000,], + [8978/117375 + 1629507*206**0.5/257912000, + 4520423/12770400, + 10410661/63852000, + 8978/117375 + 1629507*206**0.5/257912000,], + [1633/18780 + 121979*206**0.5/19343400, + 134689/939000 + 927*206**0.5/78250, + 171511/939000 + 927*206**0.5/78250, + 1633/18780 + 71*206**0.5/96717000,], + ], + [ + [1633/9390, + 1531/4695, + 1531/4695, + 1633/9390,] + ], + [ + 1/2 - 3*206**0.5 / 100, # 0.06941899716778028758987101075583196 + 33/100, + 67/100, + 1/2 + 3*206**0.5 / 100, + ] + ), + "gauss-legendre_4s": ( + [ + [1/4, 1/4 - 15**0.5 / 6, 1/4 + 15**0.5 / 6, 1/4], + [1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6, 1/4], + [1/4, 1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6], + [1/4 - 15**0.5 / 6, 1/4, 1/4 + 15**0.5 / 6, 1/4], + ], + [ + [ + 1/8, + 3/8, + 3/8, + 1/8,] + ], + [ + 1/2 - 15**0.5 / 10, # 0.11270166537925831148207346002176004 + 1/2 + 15**0.5 / 10, + 1/2 + 15**0.5 / 10, + 1/2 - 15**0.5 / 10 + ] + ), + "gauss-legendre_4s_alternating_a": ( + [ + [1/4, 1/4 - 15**0.5 / 6, 1/4 + 15**0.5 / 6, 1/4], + [1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6, 1/4], + [1/4 - 15**0.5 / 6, 1/4, 1/4 + 15**0.5 / 6, 1/4], + [1/4, 1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6], + ], + [ + [ + 1/8, + 3/8, + 1/8, + 3/8,] + ], + [ + 1/2 - 15**0.5 / 10, # 0.11270166537925831148207346002176004 + 1/2 + 15**0.5 / 10, + 1/2 - 15**0.5 / 10, + 1/2 + 15**0.5 / 10, + ] + ), + "gauss-legendre_4s_ascending_a": ( + [ + [1/4 - 15**0.5 / 6, 1/4, 1/4 + 15**0.5 / 6, 1/4], + [1/4, 1/4 - 15**0.5 / 6, 1/4 + 15**0.5 / 6, 1/4], + [1/4, 1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6], + [1/4 + 15**0.5 / 6, 1/4, 1/4 - 15**0.5 / 6, 1/4], + + ], + [ + [ + 1/8, + 3/8, + 1/8, + 3/8,] + ], + [ + 1/2 - 15**0.5 / 10, + 1/2 - 15**0.5 / 10, + 1/2 + 15**0.5 / 10, + 1/2 + 15**0.5 / 10, + ] + ), + + "gauss-legendre_3s": ( # Kunzmann-Butcher, IRK, order 6 https://www.math.umd.edu/~mariakc/SymplecticMethods.pdf + [ + [5/36, 2/9 - 15**0.5 / 15, 5/36 - 15**0.5 / 30], + [5/36 + 15**0.5 / 24, 2/9, 5/36 - 15**0.5 / 24], + [5/36 + 15**0.5 / 30, 2/9 + 15**0.5 / 15, 5/36], + ], + [ + [5/18, 4/9, 5/18] + ], + [1/2 - 15**0.5 / 10, 1/2, 1/2 + 15**0.5 / 10] # 0.11270166537925831148207346002176004 + ), + "gauss-legendre_2s": ( # Hammer-Hollingsworth, IRK, order 4 https://www.math.umd.edu/~mariakc/SymplecticMethods.pdf + [ + [1/4, 1/4 - 3**0.5 / 6], + [1/4 + 3**0.5 / 6, 1/4], + ], + [ + [1/2, 1/2], + ], + [1/2 - 3**0.5 / 6, 1/2 + 3**0.5 / 6] # 0.21132486540518711774542560974902127 # 1/2 - (1/2)*(1/3**0.5) 1/2 + (1/2)*(1/3**0.5) + ), + + "radau_iia_4s": ( + [ + [], + [], + [], + [], + ], + [ + [1/4, 1/4, 1/4, 1/4], + ], + [(1/11)*(4-6**0.5), (1/11)*(4+6**0.5), 1/2, 1] + ), + + "radau_iia_11s": ( # https://github.com/ryanelandt/Radau.jl + [ + [0.015280520789530369, -0.0057824996781311875, 0.00438010324638053, -0.0036210375473319026, 0.003092977042211754, -0.0026728314041491816, 0.0023050911672361017, -0.001955651803123845, 0.001593873849612843, -0.0011728625554916522, 0.00046993032567176855], + [0.03288397668119629, 0.03451351173940448, -0.009285420023734383, 0.00641324617083941, -0.005095455838865143, 0.0042460913690415955, -0.0035876743372353984, 0.003006834900018004, -0.0024326697483255453, 0.0017827773828584467, -0.0007131464180496306], + [0.029332502147155125, 0.0741624250777296, 0.0511486756872502, -0.012005023334430185, 0.00777794727524923, -0.005944695307870806, 0.004802655736401176, -0.003923600687657003, 0.003127328539609814, -0.0022731432208609507, 0.0009063777304940358], + [0.03111455337650569, 0.06578995121943092, 0.10929962691877611, 0.06381051663919307, -0.013853591907177828, 0.008557435524870741, -0.0063076358492939275, 0.004913357548166058, -0.0038139969541068734, 0.0027334306074068546, -0.0010839711153145738], + [0.03005269275666326, 0.07011284530154153, 0.09714692306747527, 0.1353916024839275, 0.07147107644479529, -0.014710238851905252, 0.008733191499420551, -0.00619941303527863, 0.004591640852897801, -0.003213330884490774, 0.001262857250740274], + [0.030728073929609766, 0.06751925856657341, 0.10334060375222286, 0.12083525997663601, 0.1503267876654705, 0.07350931976920085, -0.014512880052768446, 0.008296645645701008, -0.0056128275038367864, 0.003766229774466616, -0.001457705807615146], + [0.030292022376401242, 0.06914472100762357, 0.09972096441656238, 0.12801064060853223, 0.13493180383303127, 0.15289670039157693, 0.06975993047996924, -0.013274545709987746, 0.007258767272883859, -0.0044843888202694155, 0.0016878458203415244], + [0.03056654381836576, 0.06813851028407998, 0.10188107030389015, 0.12403361149690655, 0.14211431622263265, 0.13829395377418516, 0.14289135336320447, 0.06052636121446275, -0.011077739682117822, 0.005598667203856668, -0.0019877269625674446], + [0.030406629901865028, 0.06871880785022819, 0.10066095698900927, 0.12619527453091425, 0.13848875677027936, 0.14450773783254642, 0.13065188915037962, 0.1211140113707743, 0.046555483263607714, -0.008026200095719123, 0.002437640226261747], + [0.030484119381553945, 0.06843924691254653, 0.10124184869598654, 0.1251873187759311, 0.14011843430039864, 0.14190386755377057, 0.13500342651951197, 0.11262869537051934, 0.08930604389562254, 0.028969664972192485, -0.0033116985395201413], + [0.03046254890606557, 0.06851684106660112, 0.10108155427001221, 0.1254626888485642, 0.13968066655169153, 0.14258278197050367, 0.1339335430948421, 0.11443306192448831, 0.08565880960332992, 0.04992304095398403, 0.008264462809917356], + ], + [ + [0.03046254890606557, 0.06851684106660112, 0.10108155427001221, 0.1254626888485642, 0.13968066655169153, 0.14258278197050367, 0.1339335430948421, 0.11443306192448831, 0.08565880960332992, 0.04992304095398403, 0.008264462809917356], + ], + [0.011917613432415597, 0.061732071877148124, 0.14711144964307024, 0.26115967600845624, 0.39463984688578685, 0.5367387657156606, 0.6759444616766651, 0.8009789210368988, 0.9017109877901468, 0.9699709678385136, 1.0] + ), + + "radau_iia_9s": ( # https://github.com/ryanelandt/Radau.jl + [ + [0.022788378793458776, -0.008589639752938945, 0.0064510291769951465, -0.00525752869975012, 0.004388833809361376, -0.0036512155536904674, 0.0029404882137526148, -0.002149274163882554, 0.0008588433240576261], + [0.04890795244749932, 0.05070205048082808, -0.013523807196021316, 0.009209373774305071, -0.0071557133175369604, 0.005747246699432309, -0.004542582976394536, 0.003288161681791406, -0.0013090736941094112], + [0.04374276009157137, 0.10830189290274023, 0.07291956593742897, -0.016879877210016055, 0.010704551844802781, -0.007901946479238777, 0.005991406942179993, -0.0042480244399873135, 0.0016781498061495626], + [0.04624923745394712, 0.09656073072680009, 0.1542987697900386, 0.0867193693031384, -0.018451639643617873, 0.011036658729835513, -0.007673280940281649, 0.005228224999889903, -0.00203590583647778], + [0.044834436586910234, 0.10230684968594175, 0.13821763419236816, 0.18126393468214014, 0.09043360059943564, -0.018085063366782478, 0.010193387903855565, -0.006405265418866323, 0.0024271699384239612], + [0.045658755719323395, 0.09914547048938806, 0.14574704049699233, 0.16364828123387398, 0.18594458734451902, 0.08361326023153276, -0.015809936146309538, 0.00813825269404473, -0.002910469207795258], + [0.045200600187797244, 0.10085370671832047, 0.1419422367945749, 0.17118947183876332, 0.1697833861700019, 0.16776829117327952, 0.06707903432249304, -0.011792230536025322, 0.0036092462886493657], + [0.045416516657427734, 0.10006040244594375, 0.143652840987038, 0.16801908098069296, 0.17556076841841367, 0.15588627045003361, 0.12889391351650395, 0.04281082602522101, -0.004934574771244536], + [0.04535725246164146, 0.10027664901227598, 0.1431933481786156, 0.16884698348796479, 0.1741365013864833, 0.158421887835219, 0.12359468910229653, 0.0738270095231577, 0.012345679012345678], + ], + [ + [0.04535725246164146, 0.10027664901227598, 0.1431933481786156, 0.16884698348796479, 0.1741365013864833, 0.158421887835219, 0.12359468910229653, 0.0738270095231577, 0.012345679012345678], + ], + [0.01777991514736345, 0.09132360789979396, 0.21430847939563075, 0.37193216458327233, 0.5451866848034267, 0.7131752428555694, 0.8556337429578544, 0.9553660447100302, 1.0] + ), + + "radau_iia_7s": ( # https://github.com/ryanelandt/Radau.jl + [ + [0.03754626499392133, -0.0140393345564604, 0.0103527896007423, -0.008158322540275011, 0.006388413879534685, -0.004602326779148656, 0.0018289425614706437], + [0.08014759651561897, 0.08106206398589154, -0.021237992120711036, 0.014000291238817119, -0.010234185730090163, 0.0071534651513645905, -0.0028126393724067235], + [0.0720638469418819, 0.17106835498388662, 0.10961456404007211, -0.024619871728984055, 0.014760377043950817, -0.009575259396791401, 0.0036726783971383057], + [0.07570512581982441, 0.15409015514217114, 0.2271077366732024, 0.11747818703702478, -0.023810827153044174, 0.012709985533661206, -0.004608844281289633], + [0.07391234216319184, 0.16135560761594242, 0.2068672415521042, 0.23700711534269422, 0.10308679353381345, -0.018854139152580447, 0.0058589009748887914], + [0.07470556205979623, 0.1583072238724687, 0.21415342326720002, 0.21987784703186003, 0.19875212168063527, 0.06926550160550914, -0.00811600819772829], + [0.07449423555601031, 0.15910211573365074, 0.21235188950297781, 0.22355491450728324, 0.19047493682211558, 0.1196137446126562, 0.02040816326530612], + ], + [ + [0.07449423555601031, 0.15910211573365074, 0.21235188950297781, 0.22355491450728324, 0.19047493682211558, 0.1196137446126562, 0.02040816326530612], + ], + [0.029316427159784893, 0.1480785996684843, 0.3369846902811543, 0.5586715187715501, 0.7692338620300545, 0.9269456713197411, 1.0] + ), + + "radau_iia_5s": ( # https://github.com/ryanelandt/Radau.jl + [ + [0.07299886431790333, -0.02673533110794557, 0.018676929763984353, -0.01287910609330644, 0.005042839233882015], + [0.15377523147918246, 0.14621486784749352, -0.03644456890512809, 0.02123306311930472, -0.007935579902728777], + [0.14006304568480987, 0.29896712949128346, 0.16758507013524895, -0.03396910168661774, 0.010944288744192253], + [0.14489430810953477, 0.2765000687601592, 0.32579792291042103, 0.12875675325490976, -0.015708917378805327], + [0.14371356079122594, 0.28135601514946207, 0.31182652297574126, 0.22310390108357075, 0.04], + ], + [ + [0.14371356079122594, 0.28135601514946207, 0.31182652297574126, 0.22310390108357075, 0.04], + ], + [0.05710419611451768, 0.2768430136381238, 0.5835904323689168, 0.8602401356562195, 1.0] + ), + "radau_iia_3s": ( + [ + [11/45 - 7*6**0.5 / 360, 37/225 - 169*6**0.5 / 1800, -2/225 + 6**0.5 / 75], + [37/225 + 169*6**0.5 / 1800, 11/45 + 7*6**0.5 / 360, -2/225 - 6**0.5 / 75], + [4/9 - 6**0.5 / 36, 4/9 + 6**0.5 / 36, 1/9], + ], + [ + [4/9 - 6**0.5 / 36, 4/9 + 6**0.5 / 36, 1/9], + ], + [2/5 - 6**0.5 / 10, 2/5 + 6**0.5 / 10, 1.] + ), + "radau_iia_3s_alt": ( # https://www.unige.ch/~hairer/preprints/coimbra.pdf (page 7) Ehle [Eh69] and Axelsson [Ax69] + [ + [(88 - 7*6**0.5) / 360, (296 - 169*6**0.5) / 1800, (-2 + 3 * 6**0.5) / 225], + [(296 + 169*6**0.5) / 1800, (88 + 7*6**0.5) / 360, (-2 - 3*6**0.5) / 225], + [(16 - 6**0.5) / 36, (16 + 6**0.5) / 36, 1/9], + ], + [ + [ + (16 - 6**0.5) / 36, + (16 + 6**0.5) / 36, + 1/9], + ], + [ + (4 - 6**0.5) / 10, + (4 + 6**0.5) / 10, + 1.] + ), + "radau_iia_2s": ( + [ + [5/12, -1/12], + [3/4, 1/4], + ], + [ + [3/4, 1/4], + ], + [1/3, 1] + ), + "radau_ia_3s": ( + [ + [1/9, (-1-6**0.5)/18, (-1+6**0.5)/18], + [1/9, 11/45 + 7*6**0.5/360, 11/45-43*6**0.5/360], + [1/9, 11/45-43*6**0.5/360, 11/45 + 7*6**0.5/360], + ], + [ + [1/9, 4/9 + 6**0.5/36, 4/9 - 6**0.5/36], + ], + [0, 3/5-6**0.5/10, 3/5+6**0.5/10] + ), + "radau_ia_2s": ( + [ + [1/4, -1/4], + [1/4, 5/12], + ], + [ + [1/4, 3/4], + ], + [0, 2/3] + ), + "lobatto_iiia_4s": ( #6th order + [ + [0, 0, 0, 0], + [(11+5**0.5)/120, (25-5**0.5)/120, (25-13*5**0.5)/120, (-1+5**0.5)/120], + [(11-5**0.5)/120, (25+13*5**0.5)/120, (25+5**0.5)/120, (-1-5**0.5)/120], + [1/12, 5/12, 5/12, 1/12], + ], + [ + [1/12, 5/12, 5/12, 1/12], + ], + [0, (5-5**0.5)/10, (5+5**0.5)/10, 1] + ), + "lobatto_iiib_4s": ( #6th order + [ + [1/12, (-1-5**0.5)/24, (-1+5**0.5)/24, 0], + [1/12, (25+5**0.5)/120, (25-13*5**0.5)/120, 0], + [1/12, (25+13*5**0.5)/120, (25-5**0.5)/120, 0], + [1/12, (11-5**0.5)/24, (11+5**0.5)/24, 0], + ], + [ + [1/12, 5/12, 5/12, 1/12], + ], + [0, (5-5**0.5)/10, (5+5**0.5)/10, 1] + ), + "lobatto_iiic_4s": ( #6th order + [ + [1/12, (-5**0.5)/12, (5**0.5)/12, -1/12], + [1/12, 1/4, (10-7*5**0.5)/60, (5**0.5)/60], + [1/12, (10+7*5**0.5)/60, 1/4, (-5**0.5)/60], + [1/12, 5/12, 5/12, 1/12], + ], + [ + [1/12, 5/12, 5/12, 1/12], + ], + [0, (5-5**0.5)/10, (5+5**0.5)/10, 1] + ), + "lobatto_iiia_3s": ( + [ + [0, 0, 0], + [5/24, 1/3, -1/24], + [1/6, 2/3, 1/6], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1] + ), + "lobatto_iiia_2s": ( + [ + [0, 0], + [1/2, 1/2], + ], + [ + [1/2, 1/2], + ], + [0, 1] + ), + + + + "lobatto_iiib_3s": ( + [ + [1/6, -1/6, 0], + [1/6, 1/3, 0], + [1/6, 5/6, 0], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1] + ), + "lobatto_iiib_2s": ( + [ + [1/2, 0], + [1/2, 0], + ], + [ + [1/2, 1/2], + ], + [0, 1] + ), + + "lobatto_iiic_3s": ( + [ + [1/6, -1/3, 1/6], + [1/6, 5/12, -1/12], + [1/6, 2/3, 1/6], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1] + ), + "lobatto_iiic_2s": ( + [ + [1/2, -1/2], + [1/2, 1/2], + ], + [ + [1/2, 1/2], + ], + [0, 1] + ), + + + "lobatto_iiic_star_3s": ( + [ + [0, 0, 0], + [1/4, 1/4, 0], + [0, 1, 0], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1] + ), + "lobatto_iiic_star_2s": ( + [ + [0, 0], + [1, 0], + ], + [ + [1/2, 1/2], + ], + [0, 1] + ), + + "lobatto_iiid_3s": ( + [ + [1/6, 0, -1/6], + [1/12, 5/12, 0], + [1/2, 1/3, 1/6], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1] + ), + "lobatto_iiid_2s": ( + [ + [1/2, 1/2], + [-1/2, 1/2], + ], + [ + [1/2, 1/2], + ], + [0, 1] + ), + + "kraaijevanger_spijker_2s": ( #overshoots step + [ + [1/2, 0], + [-1/2, 2], + ], + [ + [-1/2, 3/2], + ], + [1/2, 3/2] + ), + + "qin_zhang_2s": ( + [ + [1/4, 0], + [1/2, 1/4], + ], + [ + [1/2, 1/2], + ], + [1/4, 3/4] + ), + + "pareschi_russo_2s": ( + [ + [(1-2**0.5/2), 0], + [1-2*(1-2**0.5/2), (1-2**0.5/2)], + ], + [ + [1/2, 1/2], + ], + [(1-2**0.5/2), 1-(1-2**0.5/2)] + ), + + "pareschi_russo_alt_2s": ( + [ + [(1-2**0.5/2), 0], + [1-(1-2**0.5/2), (1-2**0.5/2)], + ], + [ + [1-(1-2**0.5/2), (1-2**0.5/2)], + ], + [(1-2**0.5/2), 1] + ), + + "crouzeix_3s_alt": ( # Crouzeix & Raviart 1980; A-stable; pg 100 in Solving Ordinary Differential Equations II + [ + [gamma_crouzeix, 0, 0], + [1/2 - gamma_crouzeix, gamma_crouzeix, 0], + [2*gamma_crouzeix, 1-4*gamma_crouzeix, gamma_crouzeix], + ], + [ + [delta_crouzeix, 1-2*delta_crouzeix, delta_crouzeix], + ], + [gamma_crouzeix, 1/2, 1-gamma_crouzeix], + ), + + "crouzeix_3s": ( + [ + [(1+alpha_crouzeix)/2, 0, 0], + [-alpha_crouzeix/2, (1+alpha_crouzeix)/2, 0], + [1+alpha_crouzeix, -(1+2*alpha_crouzeix), (1+alpha_crouzeix)/2], + ], + [ + [1/(6*alpha_crouzeix**2), 1-(1/(3*alpha_crouzeix**2)), 1/(6*alpha_crouzeix**2)], + ], + [(1+alpha_crouzeix)/2, 1/2, (1-alpha_crouzeix)/2], + ), + + "crouzeix_2s": ( + [ + [1/2 + 3**0.5 / 6, 0], + [-(3**0.5 / 3), 1/2 + 3**0.5 / 6] + ], + [ + [1/2, 1/2], + ], + [1/2 + 3**0.5 / 6, 1/2 - 3**0.5 / 6], + ), + "verner_13s": ( #verner9. some values are missing, need to revise + [ + [], + ], + [ + [], + ], + [ + 0.03462, + 0.09702435063878045, + 0.14553652595817068, + 0.561, + 0.22900791159048503, + 0.544992088409515, + 0.645, + 0.48375, + 0.06757, + 0.25, + 0.6590650618730999, + 0.8206, + 0.9012, + ] + ), + "verner_robust_16s": ( + [ + [], + [0.04], + [-0.01988527319182291, 0.11637263332969652], + [0.0361827600517026, 0, 0.10854828015510781], + [2.272114264290177, 0, -8.526886447976398, 6.830772183686221], + [0.050943855353893744, 0, 0, 0.1755865049809071, 0.007022961270757467], + [0.1424783668683285, 0, 0, -0.3541799434668684, 0.07595315450295101, 0.6765157656337123], + [0.07111111111111111, 0, 0, 0, 0, 0.3279909287605898, 0.24089796012829906], + [0.07125, 0, 0, 0, 0, 0.32688424515752457, 0.11561575484247544, -0.03375], + [0.0482267732246581, 0, 0, 0, 0, 0.039485599804954, 0.10588511619346581, -0.021520063204743093, -0.10453742601833482], + [-0.026091134357549235, 0, 0, 0, 0, 0.03333333333333333, -0.1652504006638105, 0.03434664118368617, 0.1595758283215209, 0.21408573218281934], + [-0.03628423396255658, 0, 0, 0, 0, -1.0961675974272087, 0.1826035504321331, 0.07082254444170683, -0.02313647018482431, 0.2711204726320933, 1.3081337494229808], + [-0.5074635056416975, 0, 0, 0, 0, -6.631342198657237, -0.2527480100908801, -0.49526123800360955, 0.2932525545253887, 1.440108693768281, 6.237934498647056, 0.7270192054526988], + [0.6130118256955932, 0, 0, 0, 0, 9.088803891640463, -0.40737881562934486, 1.7907333894903747, 0.714927166761755, -1.4385808578417227, -8.26332931206474, -1.537570570808865, 0.34538328275648716], + [-1.2116979103438739, 0, 0, 0, 0, -19.055818715595954, 1.263060675389875, -6.913916969178458, -0.6764622665094981, 3.367860445026608, 18.00675164312591, 6.83882892679428, -1.0315164519219504, 0.4129106232130623], + [2.1573890074940536, 0, 0, 0, 0, 23.807122198095804, 0.8862779249216555, 13.139130397598764, -2.604415709287715, -5.193859949783872, -20.412340711541507, -12.300856252505723, 1.5215530950085394], + ], + [ + 0.014588852784055396, 0, 0, 0, 0, 0, 0, 0.0020241978878893325, 0.21780470845697167, + 0.12748953408543898, 0.2244617745463132, 0.1787254491259903, 0.07594344758096556, + 0.12948458791975614, 0.029477447612619417, 0 + ], + [ + 0, 0.04, 0.09648736013787361, 0.1447310402068104, 0.576, 0.2272326564618766, + 0.5407673435381234, 0.64, 0.48, 0.06754, 0.25, 0.6770920153543243, 0.8115, + 0.906, 1, 1 + ], + ), + + "dormand-prince_13s": ( #non-monotonic + [ + [], + [1/18], + [1/48, 1/16], + [1/32, 0, 3/32], + [5/16, 0, -75/64, 75/64], + [3/80, 0, 0, 3/16, 3/20], + [29443841/614563906, 0, 0, 77736538/692538347, -28693883/1125000000, 23124283/1800000000], + [16016141/946692911, 0, 0, 61564180/158732637, 22789713/633445777, 545815736/2771057229, -180193667/1043307555], + [39632708/573591083, 0, 0, -433636366/683701615, -421739975/2616292301, 100302831/723423059, 790204164/839813087, 800635310/3783071287], + [246121993/1340847787, 0, 0, -37695042795/15268766246, -309121744/1061227803, -12992083/490766935, 6005943493/2108947869, 393006217/1396673457, 123872331/1001029789], + [-1028468189/846180014, 0, 0, 8478235783/508512852, 1311729495/1432422823, -10304129995/1701304382, -48777925059/3047939560, 15336726248/1032824649, -45442868181/3398467696, 3065993473/597172653], + [185892177/718116043, 0, 0, -3185094517/667107341, -477755414/1098053517, -703635378/230739211, 5731566787/1027545527, 5232866602/850066563, -4093664535/808688257, 3962137247/1805957418, 65686358/487910083], + [403863854/491063109, 0, 0, -5068492393/434740067, -411421997/543043805, 652783627/914296604, 11173962825/925320556, -13158990841/6184727034, 3936647629/1978049680, -160528059/685178525, 248638103/1413531060], + ], + [ + [14005451/335480064, 0, 0, 0, 0, -59238493/1068277825, 181606767/758867731, 561292985/797845732, -1041891430/1371343529, 760417239/1151165299, 118820643/751138087, -528747749/2220607170, 1/4], + ], + [0, 1/18, 1/12, 1/8, 5/16, 3/8, 59/400, 93/200, 5490023248 / 9719169821, 13/20, 1201146811 / 1299019798, 1, 1], + ), + "dormand-prince_6s": ( + [ + [], + [1/5], + [3/40, 9/40], + [44/45, -56/15, 32/9], + [19372/6561, -25360/2187, 64448/6561, -212/729], + [9017/3168, -355/33, 46732/5247, 49/176, -5103/18656], + ], + [ + [35/384, 0, 500/1113, 125/192, -2187/6784, 11/84], + ], + [0, 1/5, 3/10, 4/5, 8/9, 1], + ), + "bogacki-shampine_7s": ( #5th order + [ + [], + [1/6], + [2/27, 4/27], + [183/1372, -162/343, 1053/1372], + [68/297, -4/11, 42/143, 1960/3861], + [597/22528, 81/352, 63099/585728, 58653/366080, 4617/20480], + [174197/959244, -30942/79937, 8152137/19744439, 666106/1039181, -29421/29068, 482048/414219], + ], + [ + [587/8064, 0, 4440339/15491840, 24353/124800, 387/44800, 2152/5985, 7267/94080], + ], + [0, 1/6, 2/9, 3/7, 2/3, 3/4, 1] + ), + "bogacki-shampine_4s": ( #5th order + [ + [], + [1/2], + [0, 3/4], + [2/9, 1/3, 4/9], + ], + [ + [2/9, 1/3, 4/9, 0], + ], + [0, 1/2, 3/4, 1] + ), + "tsi_7s": ( #5th order + [ + [], + [0.161], + [-0.008480655492356989, 0.335480655492357], + [2.8971530571054935, -6.359448489975075, 4.3622954328695815], + [5.325864828439257, -11.748883564062828, 7.4955393428898365, -0.09249506636175525], + [5.86145544294642, -12.92096931784711, 8.159367898576159, -0.071584973281401, -0.02826905039406838], + [0.09646076681806523, 0.01, 0.4798896504144996, 1.379008574103742, -3.290069515436081, 2.324710524099774], + ], + [ + [0.09646076681806523, 0.01, 0.4798896504144996, 1.379008574103742, -3.290069515436081, 2.324710524099774, 0.0], + ], + [0.0, 0.161, 0.327, 0.9, 0.9800255409045097, 1.0, 1.0], + ), + "rk6_7s": ( #non-monotonic #5th order + [ + [], + [1/3], + [0, 2/3], + [1/12, 1/3, -1/12], + [-1/16, 9/8, -3/16, -3/8], + [0, 9/8, -3/8, -3/4, 1/2], + [9/44, -9/11, 63/44, 18/11, 0, -16/11], + ], + [ + [11/120, 0, 27/40, 27/40, -4/15, -4/15, 11/120], + ], + [0, 1/3, 2/3, 1/3, 1/2, 1/2, 1], + ), + "rk5_7s": ( #5th order + [ + [], + [1/5], + [3/40, 9/40], + [44/45, -56/15, 32/9], + [19372/6561, -25360/2187, 64448/6561, 212/729], #flipped 212 sign + [-9017/3168, -355/33, 46732/5247, 49/176, -5103/18656], + [35/384, 0, 500/1113, 125/192, -2187/6784, 11/84], + ], + [ + [5179/57600, 0, 7571/16695, 393/640, -92097/339200, 187/2100, 1/40], + ], + [0, 1/5, 3/10, 4/5, 8/9, 1, 1], + ), + "ssprk4_4s": ( #non-monotonic #https://link.springer.com/article/10.1007/s41980-022-00731-x + [ + [], + [1/2], + [1/2, 1/2], + [1/6, 1/6, 1/6], + ], + [ + [1/6, 1/6, 1/6, 1/2], + ], + [0, 1/2, 1, 1/2], + ), + "rk4_4s": ( + [ + [], + [1/2], + [0, 1/2], + [0, 0, 1], + ], + [ + [1/6, 1/3, 1/3, 1/6], + ], + [0, 1/2, 1/2, 1], + ), + "rk38_4s": ( + [ + [], + [1/3], + [-1/3, 1], + [1, -1, 1], + ], + [ + [1/8, 3/8, 3/8, 1/8], + ], + [0, 1/3, 2/3, 1], + ), + "ralston_4s": ( + [ + [], + [2/5], + [(-2889+1428 * 5**0.5)/1024, (3785-1620 * 5**0.5)/1024], + [(-3365+2094 * 5**0.5)/6040, (-975-3046 * 5**0.5)/2552, (467040+203968*5**0.5)/240845], + ], + [ + [(263+24*5**0.5)/1812, (125-1000*5**0.5)/3828, (3426304+1661952*5**0.5)/5924787, (30-4*5**0.5)/123], + ], + [0, 2/5, (14-3 * 5**0.5)/16, 1], + ), + "heun_3s": ( + [ + [], + [1/3], + [0, 2/3], + ], + [ + [1/4, 0, 3/4], + ], + [0, 1/3, 2/3], + ), + "kutta_3s": ( + [ + [], + [1/2], + [-1, 2], + ], + [ + [1/6, 2/3, 1/6], + ], + [0, 1/2, 1], + ), + "ralston_3s": ( + [ + [], + [1/2], + [0, 3/4], + ], + [ + [2/9, 1/3, 4/9], + ], + [0, 1/2, 3/4], + ), + "houwen-wray_3s": ( + [ + [], + [8/15], + [1/4, 5/12], + ], + [ + [1/4, 0, 3/4], + ], + [0, 8/15, 2/3], + ), + "ssprk3_3s": ( #non-monotonic + [ + [], + [1], + [1/4, 1/4], + ], + [ + [1/6, 1/6, 2/3], + ], + [0, 1, 1/2], + ), + "midpoint_2s": ( + [ + [], + [1/2], + ], + [ + [0, 1], + ], + [0, 1/2], + ), + "heun_2s": ( + [ + [], + [1], + ], + [ + [1/2, 1/2], + ], + [0, 1], + ), + "ralston_2s": ( + [ + [], + [2/3], + ], + [ + [1/4, 3/4], + ], + [0, 2/3], + ), + "euler": ( + [ + [], + ], + [ + [1], + ], + [0], + ), +} + + + +def get_rk_methods_beta(rk_type : str, + h : Tensor, + c1 : float = 0.0, + c2 : float = 0.5, + c3 : float = 1.0, + h_prev : Optional[Tensor] = None, + step : int = 0, + sigmas : Optional[Tensor] = None, + sigma : Optional[Tensor] = None, + sigma_next : Optional[Tensor] = None, + sigma_down : Optional[Tensor] = None, + extra_options : Optional[str] = None + ): + + FSAL = False + multistep_stages = 0 + hybrid_stages = 0 + u = None + v = None + primary_rk_type = rk_type + sampler_change_reason = None + + EO = ExtraOptions(extra_options) + use_analytic_solution = not EO("disable_analytic_solution", debugMode=1) + multistep_initial_sampler = EO("multistep_initial_sampler", "", debugMode=1) + multistep_fallback_sampler = EO("multistep_fallback_sampler", "", debugMode=1) + multistep_extra_initial_steps = EO("multistep_extra_initial_steps", 1, debugMode=1) + + #if RK_Method_Beta.is_exponential(rk_type): + if rk_type.startswith(("res", "dpmpp", "ddim", "pec", "etdrk", "lawson")): + h_no_eta = -torch.log(sigma_next/sigma) + h_prev1_no_eta = -torch.log(sigmas[step]/sigmas[step-1]) if step >= 1 else None + h_prev2_no_eta = -torch.log(sigmas[step]/sigmas[step-2]) if step >= 2 else None + h_prev3_no_eta = -torch.log(sigmas[step]/sigmas[step-3]) if step >= 3 else None + h_prev4_no_eta = -torch.log(sigmas[step]/sigmas[step-4]) if step >= 4 else None + + else: + h_no_eta = sigma_next - sigma + h_prev1_no_eta = sigmas[step] - sigmas[step-1] if step >= 1 else None + h_prev2_no_eta = sigmas[step] - sigmas[step-2] if step >= 2 else None + h_prev3_no_eta = sigmas[step] - sigmas[step-3] if step >= 3 else None + h_prev4_no_eta = sigmas[step] - sigmas[step-4] if step >= 4 else None + + if type(c1) == torch.Tensor: + c1 = c1.item() + if type(c2) == torch.Tensor: + c2 = c2.item() + if type(c3) == torch.Tensor: + c3 = c3.item() + + if c1 == -1: + c1 = random.uniform(0, 1) + if c2 == -1: + c2 = random.uniform(0, 1) + if c3 == -1: + c3 = random.uniform(0, 1) + + if rk_type[:4] == "deis": + order = int(rk_type[-2]) + if step < order + multistep_extra_initial_steps: + sampler_change_reason = "initial" + if order == 4: + #rk_type = "res_4s_strehmel_weiner" + rk_type = "ralston_4s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + order = 3 + elif order == 3: + #rk_type = "res_3s" + rk_type = "ralston_3s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + elif order == 2: + #rk_type = "res_2s" + rk_type = "ralston_2s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + else: + rk_type = "deis" + multistep_stages = order-1 + + if rk_type[-2:] == "2m": #multistep method + rk_type = rk_type[:-2] + "2s" + #if h_prev is not None and step >= 1: + if h_no_eta < 1.0: + if step >= 1 + multistep_extra_initial_steps: + multistep_stages = 1 + c2 = (-h_prev1_no_eta / h_no_eta).item() + else: + sampler_change_reason = "initial" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + if rk_type.startswith("abnorsett"): + rk_type = "res_2s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + else: + sampler_change_reason = "fallback" + #rk_type = "res_2s" + rk_type = "ddim" if sigma < 0.1 else "res_2s" + rk_type = multistep_fallback_sampler if multistep_fallback_sampler else rk_type + + if rk_type[-2:] == "3m": #multistep method + rk_type = rk_type[:-2] + "3s" + #if h_prev2 is not None and step >= 2: + if h_no_eta < 1.0: + if step >= 2 + multistep_extra_initial_steps: + multistep_stages = 2 + c2 = (-h_prev1_no_eta / h_no_eta).item() + c3 = (-h_prev2_no_eta / h_no_eta).item() + else: + sampler_change_reason = "initial" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + if rk_type.startswith("abnorsett"): + rk_type = "res_3s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + else: + sampler_change_reason = "fallback" + #rk_type = "res_3s" + rk_type = "ddim" if sigma < 0.1 else "res_3s" + rk_type = multistep_fallback_sampler if multistep_fallback_sampler else rk_type + + if rk_type[-2:] == "4m": #multistep method + rk_type = rk_type[:-2] + "4s" + #if h_prev2 is not None and step >= 2: + if h_no_eta < 1.0: + if step >= 3 + multistep_extra_initial_steps: + multistep_stages = 3 + c2 = (-h_prev1_no_eta / h_no_eta).item() + c3 = (-h_prev2_no_eta / h_no_eta).item() + # WOULD NEED A C4 (POW) TO IMPLEMENT RES_4M IF IT EXISTED + else: + sampler_change_reason = "initial" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + if rk_type == "res_4s": + rk_type = "res_4s_strehmel_weiner" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + if rk_type.startswith("abnorsett"): + rk_type = "res_4s_strehmel_weiner" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + else: + sampler_change_reason = "fallback" + #rk_type = "res_4s_strehmel_weiner" + rk_type = "ddim" if sigma < 0.1 else "res_4s_strehmel_weiner" + rk_type = multistep_fallback_sampler if multistep_fallback_sampler else rk_type + + if rk_type[-3] == "h" and rk_type[-1] == "s": #hybrid method + hybrid_order = int(rk_type[-4]) + if step < hybrid_order + multistep_extra_initial_steps: + sampler_change_reason = "initial" + rk_type = "res_" + rk_type[-2:] + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + else: + hybrid_stages = hybrid_order #+1 adjustment needed? + if rk_type == "res_4s": + rk_type = "res_4s_strehmel_weiner" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + if rk_type == "res_1s": + rk_type = "res_2s" + rk_type = multistep_initial_sampler if multistep_initial_sampler else rk_type + + if rk_type in rk_coeff: + a, b, ci = copy.deepcopy(rk_coeff[rk_type]) + + a = [row + [0] * (len(ci) - len(row)) for row in a] + + match rk_type: + case "deis": + coeff_list = get_deis_coeff_list(sigmas, multistep_stages+1, deis_mode="rhoab") + coeff_list = [[elem / h for elem in inner_list] for inner_list in coeff_list] + if multistep_stages == 1: + b1, b2 = coeff_list[step] + a = [ + [0, 0], + [0, 0], + ] + b = [ + [b1, b2], + ] + ci = [0, 0] + if multistep_stages == 2: + b1, b2, b3 = coeff_list[step] + a = [ + [0, 0, 0], + [0, 0, 0], + [0, 0, 0], + ] + b = [ + [b1, b2, b3], + ] + ci = [0, 0, 0] + if multistep_stages == 3: + b1, b2, b3, b4 = coeff_list[step] + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 0, 0, 0], + ] + b = [ + [b1, b2, b3, b4], + ] + ci = [0, 0, 0, 0] + if multistep_stages > 0: + for i in range(len(b[0])): + b[0][i] *= ((sigma_down - sigma) / (sigma_next - sigma)) + + case "dormand-prince_6s": + FSAL = True + + case "ddim": + b1 = phi(1, -h) + a = [ + [0], + ] + b = [ + [b1], + ] + ci = [0] + + case "res_2s": + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + + ci = [0, c2] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(1,2) + b2 = φ(2)/c2 + b1 = φ(1) - b2 + + a = [ + [0,0], + [a2_1, 0], + ] + b = [ + [b1, b2], + ] + + case "res_2s_stable": + c2 = 1.0 #float(get_extra_options_kv("c2", str(c2), extra_options)) + + ci = [0, c2] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(1,2) + b2 = φ(2)/c2 + b1 = φ(1) - b2 + + a = [ + [0,0], + [a2_1, 0], + ] + b = [ + [b1, b2], + ] + + case "res_2s_rkmk2e": + + ci = [0, 1] + φ = Phi(h, ci, use_analytic_solution) + + b2 = φ(2) + + a = [ + [0,0], + [0, 0], + ] + b = [ + [0, b2], + ] + + gen_first_col_exp(a, b, ci, φ) + + + + case "abnorsett2_1h2s": + + c1, c2 = 0, 1 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + b1 = φ(1) #+ φ(2) + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, 0], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci) + + u1 = -φ1(2) + v1 = -φ1(2) + + u = [ + [0, 0], + [u1, 0], + ] + v = [ + [v1, 0], + ] + + gen_first_col_exp_uv(a, b, ci, u, v, φ) + + + + case "abnorsett_2m": + + c1, c2 = 0, 1 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, -φ(2)], + ] + + gen_first_col_exp(a, b, ci, φ) + + + case "abnorsett_3m": + + c1, c2, c3 = 0, 0, 1 + ci = [c1, c2, c3] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, 0, 0], + ] + b = [ + [0, -2*φ(2) - 2*φ(3), (1/2)*φ(2) + φ(3)], + ] + + gen_first_col_exp(a, b, ci, φ) + + + + case "abnorsett_4m": + + c1, c2, c3, c4 = 0, 0, 0, 1 + ci = [c1, c2, c3, c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 0, 0, 0], + ] + b = [ + [0, + -3*φ(2) - 5*φ(3) - 3*φ(4), + (3/2)*φ(2) + 4*φ(3) + 3*φ(4), + (-1/3)*φ(2) - φ(3) - φ(4), + ], + ] + + gen_first_col_exp(a, b, ci, φ) + + + case "abnorsett3_2h2s": + + c1,c2 = 0,1 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + b2 = 0 + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, 0], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci) + φ2 = Phi(h_prev2_no_eta, ci) + + u2_1 = -2*φ1(2) - 2*φ1(3) + u2_2 = (1/2)*φ2(2) + φ2(3) + + v1 = u2_1 # -φ1(2) + φ1(3) + 3*φ1(4) + v2 = u2_2 # (1/6)*φ2(2) - φ2(4) + + u = [ + [ 0, 0], + [u2_1, u2_2], + ] + v = [ + [v1, v2], + ] + + gen_first_col_exp_uv(a, b, ci, u, v, φ) + + + + case "pec423_2h2s": #https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + + c1,c2 = 0,1 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + b2 = (1/3)*φ(2) + φ(3) + φ(4) + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, b2], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci) + φ2 = Phi(h_prev2_no_eta, ci) + + u2_1 = -2*φ1(2) - 2*φ1(3) + u2_2 = (1/2)*φ2(2) + φ2(3) + + v1 = -φ1(2) + φ1(3) + 3*φ1(4) + v2 = (1/6)*φ2(2) - φ2(4) + + u = [ + [ 0, 0], + [u2_1, u2_2], + ] + v = [ + [v1, v2], + ] + + gen_first_col_exp_uv(a, b, ci, u, v, φ) + + + + + case "pec433_2h3s": #https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + + c1,c2,c3 = 0, 1, 1 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = (1/3)*φ(2) + φ(3) + φ(4) + + b2 = 0 + b3 = (1/3)*φ(2) + φ(3) + φ(4) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, a3_2, 0], + ] + b = [ + [0, b2, b3], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci) + φ2 = Phi(h_prev2_no_eta, ci) + + u2_1 = -2*φ1(2) - 2*φ1(3) + u3_1 = -φ1(2) + φ1(3) + 3*φ1(4) + v1 = -φ1(2) + φ1(3) + 3*φ1(4) + + u2_2 = (1/2)*φ2(2) + φ2(3) + u3_2 = (1/6)*φ2(2) - φ2(4) + v2 = (1/6)*φ2(2) - φ2(4) + + + u = [ + [ 0, 0, 0], + [u2_1, u2_2, 0], + [u3_1, u3_2, 0], + ] + v = [ + [v1, v2, 0], + ] + + gen_first_col_exp_uv(a, b, ci, u, v, φ) + + + + case "res_3s": + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + c3 = float(get_extra_options_kv("c3", str(c3), extra_options)) + + ci = [0,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + gamma = calculate_gamma(c2, c3) + + a3_2 = gamma * c2 * φ(2,2) + (c3 ** 2 / c2) * φ(2, 3) + + b3 = (1 / (gamma * c2 + c3)) * φ(2) + b2 = gamma * b3 #simplified version of: b2 = (gamma / (gamma * c2 + c3)) * phi_2_h + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, a3_2, 0], + ] + b = [ + [0, b2, b3], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_3s_non-monotonic": + c2 = float(get_extra_options_kv("c2", "1.0", extra_options)) + c3 = float(get_extra_options_kv("c3", "0.5", extra_options)) + + ci = [0,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + gamma = calculate_gamma(c2, c3) + + a3_2 = gamma * c2 * φ(2,2) + (c3 ** 2 / c2) * φ(2, 3) + + b3 = (1 / (gamma * c2 + c3)) * φ(2) + b2 = gamma * b3 #simplified version of: b2 = (gamma / (gamma * c2 + c3)) * phi_2_h + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, a3_2, 0], + ] + b = [ + [0, b2, b3], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "res_3s_alt": + c2 = 1/3 + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + + c1,c2,c3 = 0, c2, 2/3 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, (4/(9*c2)) * φ(2,3), 0], + ] + b = [ + [0, 0, (1/c3)*φ(2)], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_3s_strehmel_weiner": # + c2 = 1/2 + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + + ci = [0,c2,1] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, (1/c2) * φ(2,3), 0], + ] + b = [ + [0, 0, φ(2)], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "res_3s_cox_matthews": # Cox & Matthews; known as ETD3RK + c2 = 1/2 # must be 1/2 + ci = [0,c2,1] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, (1/c2) * φ(1,3), 0], # paper said 2 * φ(1,3), but this is the same and more consistent with res_3s_strehmel_weiner + ] + b = [ + [0, + -8*φ(3) + 4*φ(2), + 4*φ(3) - φ(2)], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_3s_lie": # Lie; known as ETD2CF3 + c1,c2,c3 = 0, 1/3, 2/3 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, (4/3)*φ(2,3), 0], # paper said 2 * φ(1,3), but this is the same and more consistent with res_3s_strehmel_weiner + ] + b = [ + [0, + 6*φ(2) - 18*φ(3), + (-3/2)*φ(2) + 9*φ(3)], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_3s_sunstar": # https://arxiv.org/pdf/2410.00498 pg 5 (tableau 2.7) + c1,c2,c3 = 0, 1/3, 2/3 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, (8/9)*φ(2,3), 0], # paper said 2 * φ(1,3), but this is the same and more consistent with res_3s_strehmel_weiner + ] + b = [ + [0, + 0, + (3/2)*φ(2)], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "res_4s_cox_matthews": # weak 4th order, Cox & Matthews; unresolved issue, see below + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(1,2) + a3_2 = c3 * φ(1,3) + a4_1 = (1/2) * φ(1,3) * (φ(0,3) - 1) # φ(0,3) == torch.exp(-h*c3) + a4_3 = φ(1,3) + + b1 = φ(1) - 3*φ(2) + 4*φ(3) + + b2 = 2*φ(2) - 4*φ(3) + b3 = 2*φ(2) - 4*φ(3) + b4 = 4*φ(3) - φ(2) + + a = [ + [0, 0,0,0], + [a2_1, 0,0,0], + [0, a3_2,0,0], + [a4_1, 0, a4_3,0], + ] + b = [ + [b1, b2, b3, b4], + ] + + + case "res_4s_cfree4": # weak 4th order, Cox & Matthews; unresolved issue, see below + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(1,2) + a3_2 = c3 * φ(1,2) + a4_1 = (1/2) * φ(1,2) * (φ(0,2) - 1) # φ(0,3) == torch.exp(-h*c3) + a4_3 = φ(1,2) + + b1 = (1/2)*φ(1) - (1/3)*φ(1,2) + + b2 = (1/3)*φ(1) + b3 = (1/3)*φ(1) + b4 = -(1/6)*φ(1) + (1/3)*φ(1,2) + + a = [ + [0, 0,0,0], + [a2_1, 0,0,0], + [0, a3_2,0,0], + [a4_1, 0, a4_3,0], + ] + b = [ + [b1, b2, b3, b4], + ] + + case "res_4s_friedli": # https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 2*φ(2,2) + a4_2 = -(26/25)*φ(1) + (2/25)*φ(2) + a4_3 = (26/25)*φ(1) + (48/25)*φ(2) + + + b2 = 0 + b3 = 4*φ(2) - 8*φ(3) + b4 = -φ(2) + 4*φ(3) + + a = [ + [0, 0,0,0], + [0, 0,0,0], + [0, a3_2,0,0], + [0, a4_2, a4_3,0], + ] + b = [ + [0, b2, b3, b4], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_4s_munthe-kaas": # unstable RKMK4t + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [c2*φ(1,2), 0, 0, 0], + [(h/8)*φ(1,2), (1/2)*(1-h/4)*φ(1,2), 0, 0], + [0, 0, φ(1), 0], + ] + b = [ + [ + (1/6)*φ(1)*(1+h/2), + (1/3)*φ(1), + (1/3)*φ(1), + (1/6)*φ(1)*(1-h/2) + ], + ] + + case "res_4s_krogstad": # weak 4th order, Krogstad + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, φ(2,3), 0, 0], + [0, 0, 2*φ(2,4), 0], + ] + b = [ + [ + 0, + 2*φ(2) - 4*φ(3), + 2*φ(2) - 4*φ(3), + -φ(2) + 4*φ(3) + ], + ] + + #a = [row + [0] * (len(ci) - len(row)) for row in a] + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_4s_krogstad_alt": # weak 4th order, Krogstad https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 4*φ(2,2), 0, 0], + [0, 0, 2*φ(2), 0], + ] + b = [ + [ + 0, + 2*φ(2) - 4*φ(3), + 2*φ(2) - 4*φ(3), + -φ(2) + 4*φ(3) + ], + ] + + #a = [row + [0] * (len(ci) - len(row)) for row in a] + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_4s_minchev": # https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = (4/25)*φ(1,2) + (24/25)*φ(2,2) + a4_2 = (21/5)*φ(2) - (108/5)*φ(3) + a4_3 = (1/20)*φ(1) - (33/10)*φ(2) + (123/5)*φ(3) + + + b2 = -(1/10)*φ(1) + (1/5)*φ(2) - 4*φ(3) + 12*φ(4) + b3 = (1/30)*φ(1) + (23/5)*φ(2) - 8*φ(3) - 4*φ(4) + b4 = (1/30)*φ(1) - (7/5)*φ(2) + 6*φ(3) - 4*φ(4) + + a = [ + [0, 0,0,0], + [0, 0,0,0], + [0, a3_2,0,0], + [0, 0, a4_3,0], + ] + b = [ + [0, b2, b3, b4], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_4s_strehmel_weiner": # weak 4th order, Strehmel & Weiner + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, c3*φ(2,3), 0, 0], + [0, -2*φ(2,4), 4*φ(2,4), 0], + ] + b = [ + [ + 0, + 0, + 4*φ(2) - 8*φ(3), + -φ(2) + 4*φ(3) + ], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_4s_strehmel_weiner_alt": # weak 4th order, Strehmel & Weiner https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, 2*φ(2,2), 0, 0], + [0, -2*φ(2), 4*φ(2), 0], + ] + b = [ + [ + 0, + 0, + 4*φ(2) - 8*φ(3), + -φ(2) + 4*φ(3) + ], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "lawson2a_2s": # based on midpoint rule, stiff order 1 https://cds.cern.ch/record/848126/files/cer-002531460.pdf + c1,c2 = 0,1/2 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(0,2) + b2 = φ(0,2) + b1 = 0 + + a = [ + [0,0], + [a2_1, 0], + ] + b = [ + [b1, b2], + ] + + case "lawson2b_2s": # based on trapezoidal rule, stiff order 1 https://cds.cern.ch/record/848126/files/cer-002531460.pdf + c1,c2 = 0,1 + ci = [c1, c2] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = φ(0) + b2 = 1/2 + b1 = (1/2)*φ(0) + + a = [ + [0,0], + [a2_1, 0], + ] + b = [ + [b1, b2], + ] + + + case "lawson4_4s": + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(0,2) + a3_2 = 1/2 + a4_3 = φ(0,2) + + b1 = (1/6) * φ(0) + b2 = (1/3) * φ(0,2) + b3 = (1/3) * φ(0,2) + b4 = 1/6 + + a = [ + [0, 0, 0, 0], + [a2_1, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [b1,b2,b3,b4], + ] + + case "lawson41-gen_4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b2 = (1/3) * φ(0,2) + b3 = (1/3) * φ(0,2) + b4 = 1/6 + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, + b2, + b3, + b4,], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "lawson41-gen-mod_4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b2 = (1/3) * φ(0,2) + b3 = (1/3) * φ(0,2) + b4 = φ(2) - (1/3)*φ(0,2) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, + b2, + b3, + b4,], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + + case "lawson42-gen-mod_1h4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b2 = (1/3) * φ(0,2) + b3 = (1/3) * φ(0,2) + b4 = (1/2)*φ(2) + φ(3) - (1/4)*φ(0,2) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, b2, b3, b4,], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci, use_analytic_solution) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci, use_analytic_solution) + + u2_1 = -φ1(2,2) + u3_1 = -φ1(2,2) + 1/4 + u4_1 = -φ1(2) + (1/2)*φ1(0,2) + v1 = -(1/2)*φ1(2) + φ1(3) + (1/12)*φ1(0,2) + + u = [ + [ 0, 0, 0, 0], + [u2_1, 0, 0, 0], + [u3_1, 0, 0, 0], + [u4_1, 0, 0, 0], + ] + v = [ + [v1, 0, 0, 0,], + ] + + a, b = gen_first_col_exp_uv(a,b,ci,u,v,φ) + + + + case "lawson43-gen-mod_2h4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b3 = b2 = (1/3) * a4_3 + b4 = (1/3)*φ(2) + φ(3) + φ(4) - (5/24)*φ(0,2) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, b2, b3, b4,], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci, use_analytic_solution) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta, ci, use_analytic_solution) + + u2_1 = -2*φ1(2,2) - 2*φ1(3,2) + u3_1 = -2*φ1(2,2) - 2*φ1(3,2) + 5/8 + u4_1 = -2*φ1(2) - 2*φ1(3) + (5/4)*φ1(0,2) + v1 = -φ1(2) + φ1(3) + 3*φ1(4) + (5/24)*φ1(0,2) + + u2_2 = -(1/2)*φ2(2,2) + φ2(3,2) + u3_2 = (1/2)*φ2(2,2) + φ2(3,2) - 3/16 + u4_2 = (1/2)*φ2(2) + φ2(3) - (3/8)*φ2(0,2) + v2 = (1/6)*φ2(2) - φ2(4) - (1/24)*φ2(0,2) + + u = [ + [ 0, 0, 0, 0], + [u2_1, u2_2, 0, 0], + [u3_1, u3_2, 0, 0], + [u4_1, u4_2, 0, 0], + ] + v = [ + [v1, v2, 0, 0,], + ] + + a, b = gen_first_col_exp_uv(a,b,ci,u,v,φ) + + + case "lawson44-gen-mod_3h4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b3 = b2 = (1/3) * a4_3 + b4 = (1/4)*φ(2) + (11/12)*φ(3) + (3/2)*φ(4) + φ(5) - (35/192)*φ(0,2) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, b2, b3, b4,], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ3 = Phi(h_prev3_no_eta * h/h_no_eta, ci, use_analytic_solution) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + φ3 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta, ci, use_analytic_solution) + φ3 = Phi(h_prev3_no_eta, ci, use_analytic_solution) + + u2_1 = -3*φ1(2,2) - 5*φ1(3,2) - 3*φ1(4,2) + u3_1 = u2_1 + 35/32 + u4_1 = -3*φ1(2) - 5*φ1(3) - 3*φ1(4) + (35/16)*φ1(0,2) + v1 = -(3/2)*φ1(2) + (1/2)*φ1(3) + 6*φ1(4) + 6*φ1(5) + (35/96)*φ1(0,2) + + u2_2 = (3/2)*φ2(2,2) + 4*φ2(3,2) + 3*φ2(4,2) + u3_2 = u2_2 - 21/32 + u4_2 = (3/2)*φ2(2) + 4*φ2(3) + 3*φ2(4) - (21/16)*φ2(0,2) + v2 = (1/2)*φ2(2) + (1/3)*φ2(3) - 3*φ2(4) - 4*φ2(5) - (7/48)*φ2(0,2) + + u2_3 = (-1/3)*φ3(2,2) - φ3(3,2) - φ3(4,2) + u3_3 = u2_3 + 5/32 + u4_3 = -(1/3)*φ3(2) - φ3(3) - φ3(4) + (5/16)*φ3(0,2) + v3 = -(1/12)*φ3(2) - (1/12)*φ3(3) + (1/2)*φ3(4) + φ3(5) + (5/192)*φ3(0,2) + + u = [ + [ 0, 0, 0, 0], + [u2_1, u2_2, u2_3, 0], + [u3_1, u3_2, u3_3, 0], + [u4_1, u4_2, u4_3, 0], + ] + v = [ + [v1, v2, v3, 0,], + ] + + a, b = gen_first_col_exp_uv(a,b,ci,u,v,φ) + + + + case "lawson45-gen-mod_4h4s": # GenLawson4 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 1/2 + a4_3 = φ(0,2) + + b2 = (1/3) * φ(0,2) + b3 = (1/3) * φ(0,2) + b4 = (12/59)*φ(2) + (50/59)*φ(3) + (105/59)*φ(4) + (120/59)*φ(5) - (60/59)*φ(6) - (157/944)*φ(0,2) + + a = [ + [0, 0, 0, 0], + [0, 0, 0, 0], + [0, a3_2, 0, 0], + [0, 0, a4_3, 0], + ] + b = [ + [0, b2, b3, b4,], + ] + + if extra_options_flag("h_prev_h_h_no_eta", extra_options): + φ1 = Phi(h_prev1_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ3 = Phi(h_prev3_no_eta * h/h_no_eta, ci, use_analytic_solution) + φ4 = Phi(h_prev4_no_eta * h/h_no_eta, ci, use_analytic_solution) + elif extra_options_flag("h_only", extra_options): + φ1 = Phi(h, ci, use_analytic_solution) + φ2 = Phi(h, ci, use_analytic_solution) + φ3 = Phi(h, ci, use_analytic_solution) + φ4 = Phi(h, ci, use_analytic_solution) + else: + φ1 = Phi(h_prev1_no_eta, ci, use_analytic_solution) + φ2 = Phi(h_prev2_no_eta, ci, use_analytic_solution) + φ3 = Phi(h_prev3_no_eta, ci, use_analytic_solution) + φ4 = Phi(h_prev4_no_eta, ci, use_analytic_solution) + + u2_1 = -4*φ1(2,2) - (26/3)*φ1(3,2) - 9*φ1(4,2) - 4*φ1(5,2) + u3_1 = u2_1 + 105/64 + u4_1 = -4*φ1(2) - (26/3)*φ1(3) - 9*φ1(4) - 4*φ1(5) + (105/32)*φ1(0,2) + v1 = -(116/59)*φ1(2) - (34/177)*φ1(3) + (519/59)*φ1(4) + (964/59)*φ1(5) - (600/59)*φ1(6) + (495/944)*φ1(0,2) + + u2_2 = 3*φ2(2,2) + (19/2)*φ2(3,2) + 12*φ2(4,2) + 6*φ2(5,2) + u3_2 = u2_2 - 189/128 + u4_2 = 3*φ2(2) + (19/2)*φ2(3) + 12*φ2(4) + 6*φ2(5) - (189/64)*φ2(0,2) + v2 = (57/59)*φ2(2) + (121/118)*φ2(3) - (342/59)*φ2(4) - (846/59)*φ2(5) + (600/59)*φ2(6) - (577/1888)*φ2(0,2) + + u2_3 = -(4/3)*φ3(2,2) - (14/3)*φ3(3,2) - 7*φ3(4,2) - 4*φ3(5,2) + u3_3 = u2_3 + 45/64 + u4_3 = -(4/3)*φ3(2) - (14/3)*φ3(3) - 7*φ3(4) - 4*φ3(5) +(45/32)*φ3(0,2) + v3 = -(56/177)*φ3(2) - (76/177)*φ3(3) + (112/59)*φ3(4) + (364/59)*φ3(5) - (300/59)*φ3(6) + (25/236)*φ3(0,2) + + u2_4 = (1/4)*φ4(2,2) + (88/96)*φ4(3,2) + (3/2)*φ4(4,2) + φ4(5,2) + u3_4 = u2_4 - 35/256 + u4_4 = (1/4)*φ4(2) + (11/12)*φ4(3) + (3/2)*φ4(4) + φ4(5) - (35/128)*φ4(0,2) + v4 = (11/236)*φ4(2) + (49/708)*φ4(3) - (33/118)*φ4(4) - (61/59)*φ4(5) + ( 60/59)*φ4(6) - (181/11328)*φ4(0,2) + + u = [ + [ 0, 0, 0, 0], + [u2_1, u2_2, u2_3, u2_4], + [u3_1, u3_2, u3_3, u3_4], + [u4_1, u4_2, u4_3, u4_4], + ] + v = [ + [v1, v2, v3, v4,], + ] + + a, b = gen_first_col_exp_uv(a,b,ci,u,v,φ) + + + + case "etdrk2_2s": # https://arxiv.org/pdf/2402.15142v1 + c1,c2 = 0, 1 + ci = [c1,c2] + φ = Phi(h, ci, use_analytic_solution) + + a = [ + [0, 0], + [φ(1), 0], + ] + b = [ + [φ(1)-φ(2), φ(2)], + ] + + case "etdrk3_a_3s": #non-monotonic # https://arxiv.org/pdf/2402.15142v1 + c1,c2,c3 = 0, 1, 2/3 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2*φ(1) + a3_2 = (4/9)*φ(2,3) + a3_1 = c3*φ(1,3) - a3_2 + + b2 = φ(2) - (1/2)*φ(1) + b3 = (3/4) * φ(1) + b1 = φ(1) - b2 - b3 + + a = [ + [0, 0, 0], + [a2_1, 0, 0], + [a3_1, a3_2, 0 ] + ] + b = [ + [b1, b2, b3], + ] + + case "etdrk3_b_3s": # https://arxiv.org/pdf/2402.15142v1 + c1,c2,c3 = 0, 4/9, 2/3 + ci = [c1,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2*φ(1,2) + a3_2 = φ(2,3) + a3_1 = c3*φ(1,3) - a3_2 + + b2 = 0 + b3 = (3/2) * φ(2) + b1 = φ(1) - b2 - b3 + + a = [ + [0, 0, 0], + [a2_1, 0, 0], + [a3_1, a3_2, 0 ] + ] + b = [ + [b1, b2, b3], + ] + + case "etdrk4_4s": # https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = φ(1,2) + a4_3 = 2*φ(1,2) + + b2 = 2*φ(2) - 4*φ(3) + b3 = 2*φ(2) - 4*φ(3) + b4 = -φ(2) + 4*φ(3) + + a = [ + [0, 0,0,0], + [0, 0,0,0], + [0, a3_2,0,0], + [0, 0, a4_3,0], + ] + b = [ + [0, b2, b3, b4], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "etdrk4_4s_alt": # pg 70 col 1 computed with (4.9) https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + c1,c2,c3,c4 = 0, 1/2, 1/2, 1 + ci = [c1,c2,c3,c4] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = φ(1,2) #unsure about this, looks bad and is pretty different from col #1 implementations for everything else except the other 4s alt and 5s ostermann??? from the link + a3_1 = 0 + a4_1 = φ(1) - 2*φ(1,2) + + a3_2 = φ(1,2) + a4_3 = 2*φ(1,2) + + b1 = φ(1) - 3*φ(2) + 4*φ(3) + b2 = 2*φ(2) - 4*φ(3) + b3 = 2*φ(2) - 4*φ(3) + b4 = -φ(2) + 4*φ(3) + + a = [ + [ 0, 0, 0,0], + [a2_1, 0, 0,0], + [a3_1, a3_2, 0,0], + [a4_1, 0, a4_3,0], + ] + b = [ + [0, b2, b3, b4], + ] + + #a, b = gen_first_col_exp(a,b,ci,φ) + + + case "dpmpp_2s": + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + + ci = [0,c2] + φ = Phi(h, ci, use_analytic_solution) + + b2 = (1/(2*c2)) * φ(1) + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, b2], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "dpmpp_sde_2s": + c2 = 1.0 #hardcoded to 1.0 to more closely emulate the configuration for k-diffusion's implementation + + ci = [0,c2] + φ = Phi(h, ci, use_analytic_solution) + + b2 = (1/(2*c2)) * φ(1) + + a = [ + [0, 0], + [0, 0], + ] + b = [ + [0, b2], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "dpmpp_3s": + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + c3 = float(get_extra_options_kv("c3", str(c3), extra_options)) + + ci = [0,c2,c3] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = (c3**2 / c2) * φ(2,3) + b3 = (1/c3) * φ(2) + + a = [ + [0, 0, 0], + [0, 0, 0], + [0, a3_2, 0], + ] + b = [ + [0, 0, b3], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_5s": #non-monotonic #4th order + + c1, c2, c3, c4, c5 = 0, 1/2, 1/2, 1, 1/2 + ci = [c1,c2,c3,c4,c5] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = φ(2,3) + a4_2 = φ(2,4) + a5_2 = (1/2)*φ(2,5) - φ(3,4) + (1/4)*φ(2,4) - (1/2)*φ(3,5) + + a4_3 = a4_2 + a5_3 = a5_2 + + a5_4 = (1/4)*φ(2,5) - a5_2 + + b4 = -φ(2) + 4*φ(3) + b5 = 4*φ(2) - 8*φ(3) + + a = [ + [0, 0, 0, 0, 0], + [0, 0, 0, 0, 0], + [0, a3_2, 0, 0, 0], + [0, a4_2, a4_3, 0, 0], + [0, a5_2, a5_3, a5_4, 0], + ] + b = [ + [0, 0, 0, b4, b5], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_5s_hochbruck-ostermann": #non-monotonic #4th order + + c1, c2, c3, c4, c5 = 0, 1/2, 1/2, 1, 1/2 + ci = [c1,c2,c3,c4,c5] + φ = Phi(h, ci, use_analytic_solution) + + a3_2 = 4*φ(2,2) + a4_2 = φ(2) + a5_2 = (1/4)*φ(2) - φ(3) + 2*φ(2,2) - 4*φ(3,2) + + a4_3 = φ(2) + a5_3 = a5_2 + + a5_4 = φ(2,2) - a5_2 + + b4 = -φ(2) + 4*φ(3) + b5 = 4*φ(2) - 8*φ(3) + + a = [ + [0, 0 , 0 , 0 , 0], + [0, 0 , 0 , 0 , 0], + [0, a3_2, 0 , 0 , 0], + [0, a4_2, a4_3, 0 , 0], + [0, a5_2, a5_3, a5_4, 0], + ] + b = [ + [0, 0, 0, b4, b5], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + + case "res_6s": #non-monotonic #4th order + + c1, c2, c3, c4, c5, c6 = 0, 1/2, 1/2, 1/3, 1/3, 5/6 + ci = [c1, c2, c3, c4, c5, c6] + φ = Phi(h, ci, use_analytic_solution) + + a2_1 = c2 * φ(1,2) + + a3_1 = 0 + a3_2 = (c3**2 / c2) * φ(2,3) + + a4_1 = 0 + a4_2 = (c4**2 / c2) * φ(2,4) + a4_3 = (c4**2 * φ(2,4) - a4_2 * c2) / c3 + + a5_1 = 0 + a5_2 = 0 #zero + a5_3 = (-c4 * c5**2 * φ(2,5) + 2*c5**3 * φ(3,5)) / (c3 * (c3 - c4)) + a5_4 = (-c3 * c5**2 * φ(2,5) + 2*c5**3 * φ(3,5)) / (c4 * (c4 - c3)) + + a6_1 = 0 + a6_2 = 0 #zero + a6_3 = (-c4 * c6**2 * φ(2,6) + 2*c6**3 * φ(3,6)) / (c3 * (c3 - c4)) + a6_4 = (-c3 * c6**2 * φ(2,6) + 2*c6**3 * φ(3,6)) / (c4 * (c4 - c3)) + a6_5 = (c6**2 * φ(2,6) - a6_3*c3 - a6_4*c4) / c5 + #a6_5_alt = (2*c6**3 * φ(3,6) - a6_3*c3**2 - a6_4*c4**2) / c5**2 + + b1 = 0 + b2 = 0 + b3 = 0 + b4 = 0 + b5 = (-c6*φ(2) + 2*φ(3)) / (c5 * (c5 - c6)) + b6 = (-c5*φ(2) + 2*φ(3)) / (c6 * (c6 - c5)) + + a = [ + [0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0], + [0, a3_2, 0, 0, 0, 0], + [0, a4_2, a4_3, 0, 0, 0], + [0, a5_2, a5_3, a5_4, 0, 0], + [0, a6_2, a6_3, a6_4, a6_5, 0], + ] + b = [ + [0, b2, b3, b4, b5, b6], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + case "res_8s": #non-monotonic # this is not EXPRK5S8 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + + c1, c2, c3, c4, c5, c6, c7, c8 = 0, 1/2, 1/2, 1/4, 1/2, 1/5, 2/3, 1 + ci = [c1, c2, c3, c4, c5, c6, c7, c8] + #φ = Phi(h, ci, analytic_solution=use_analytic_solution) + + ci = [mpf(c_val) for c_val in ci] + c1, c2, c3, c4, c5, c6, c7, c8 = [c_val for c_val in ci] + + φ = Phi(mpf(h.item()), ci, analytic_solution=use_analytic_solution) + + a3_2 = (1/2) * φ(2,3) + + a4_3 = (1/8) * φ(2,4) + + a5_3 = (-1/2) * φ(2,5) + 2 * φ(3,5) + a5_4 = 2 * φ(2,5) - 4 * φ(3,5) + + a6_4 = (8/25) * φ(2,6) - (32/125) * φ(3,6) + a6_5 = (2/25) * φ(2,6) - (1/2) * a6_4 + + a7_4 = (-125/162) * a6_4 + a7_5 = (125/1944) * a6_4 - (16/27) * φ(2,7) + (320/81) * φ(3,7) + a7_6 = (3125/3888) * a6_4 + (100/27) * φ(2,7) - (800/81) * φ(3,7) + + Φ = (5/32)*a6_4 - (1/28)*φ(2,6) + (36/175)*φ(2,7) - (48/25)*φ(3,7) + (6/175)*φ(4,6) + (192/35)*φ(4,7) + 6*φ(4,8) + + a8_5 = (208/3)*φ(3,8) - (16/3) *φ(2,8) - 40*Φ + a8_6 = (-250/3)*φ(3,8) + (250/21)*φ(2,8) + (250/7)*Φ + a8_7 = -27*φ(3,8) + (27/14)*φ(2,8) + (135/7)*Φ + + b6 = (125/14)*φ(2) - (625/14)*φ(3) + (1125/14)*φ(4) + b7 = (-27/14)*φ(2) + (162/7) *φ(3) - (405/7) *φ(4) + b8 = (1/2) *φ(2) - (13/2) *φ(3) + (45/2) *φ(4) + + b1 = φ(1) - b6 - b7 - b8 + + a = [ + [0 , 0 , 0 , 0 , 0 , 0 , 0 , 0], + [0 , 0 , 0 , 0 , 0 , 0 , 0 , 0], + + [0 , a3_2, 0 , 0 , 0 , 0 , 0 , 0], + [0 , 0 , a4_3, 0 , 0 , 0 , 0 , 0], + + [0 , 0 , a5_3, a5_4, 0 , 0 , 0 , 0], + [0 , 0 , 0 , a6_4, a6_5, 0 , 0 , 0], + + [0 , 0 , 0 , a7_4, a7_5, a7_6, 0 , 0], + [0 , 0 , 0 , 0 , a8_5, a8_6, a8_7, 0], + ] + b = [ + [0, 0, 0, 0, 0, b6, b7, b8], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + a = [[float(val) for val in row] for row in a] + b = [[float(val) for val in row] for row in b] + ci = [c1, c2, c3, c4, c5, c6, c7, c8] + + + + case "res_8s_alt": # this is EXPRK5S8 https://ora.ox.ac.uk/objects/uuid:cc001282-4285-4ca2-ad06-31787b540c61/files/m611df1a355ca243beb09824b70e5e774 + + c1, c2, c3, c4, c5, c6, c7, c8 = 0, 1/2, 1/2, 1/4, 1/2, 1/5, 2/3, 1 + ci = [c1, c2, c3, c4, c5, c6, c7, c8] + #φ = Phi(h, ci, analytic_solution=use_analytic_solution) + + ci = [mpf(c_val) for c_val in ci] + c1, c2, c3, c4, c5, c6, c7, c8 = [c_val for c_val in ci] + + φ = Phi(mpf(h.item()), ci, analytic_solution=use_analytic_solution) + + a3_2 = 2*φ(2,2) + + a4_3 = 2*φ(2,4) + + a5_3 = -2*φ(2,2) + 16*φ(3,2) + a5_4 = 8*φ(2,2) - 32*φ(3,2) + + a6_4 = 8*φ(2,6) - 32*φ(3,6) + a6_5 = -2*φ(2,6) + 16*φ(3,6) + + a7_4 = (-125/162) * a6_4 + a7_5 = (125/1944) * a6_4 - (4/3) * φ(2,7) + (40/3)*φ(3,7) + a7_6 = (3125/3888) * a6_4 + (25/3) * φ(2,7) - (100/3)*φ(3,7) + + Φ = (5/32)*a6_4 - (25/28)*φ(2,6) + (81/175)*φ(2,7) - (162/25)*φ(3,7) + (150/7)*φ(4,6) + (972/35)*φ(4,7) + 6*φ(4) + + a8_5 = -(16/3)*φ(2) + (208/3)*φ(3) - 40*Φ + a8_6 = (250/21)*φ(2) - (250/3)*φ(3) + (250/7)*Φ + a8_7 = (27/14)*φ(2) - 27*φ(3) + (135/7)*Φ + + b6 = (125/14)*φ(2) - (625/14)*φ(3) + (1125/14)*φ(4) + b7 = (-27/14)*φ(2) + (162/7) *φ(3) - (405/7) *φ(4) + b8 = (1/2) *φ(2) - (13/2) *φ(3) + (45/2) *φ(4) + + a = [ + [0 , 0 , 0 , 0 , 0 , 0 , 0 , 0], + [0 , 0 , 0 , 0 , 0 , 0 , 0 , 0], + + [0 , a3_2, 0 , 0 , 0 , 0 , 0 , 0], + [0 , 0 , a4_3, 0 , 0 , 0 , 0 , 0], + + [0 , 0 , a5_3, a5_4, 0 , 0 , 0 , 0], + [0 , 0 , 0 , a6_4, a6_5, 0 , 0 , 0], + + [0 , 0 , 0 , a7_4, a7_5, a7_6, 0 , 0], + [0 , 0 , 0 , 0 , a8_5, a8_6, a8_7, 0], + ] + b = [ + [0, 0, 0, 0, 0, b6, b7, b8], + ] + + a, b = gen_first_col_exp(a,b,ci,φ) + + a = [[float(val) for val in row] for row in a] + b = [[float(val) for val in row] for row in b] + ci = [c1, c2, c3, c4, c5, c6, c7, c8] + + + case "res_10s": + + c1, c2, c3, c4, c5, c6, c7, c8, c9, c10 = 0, 1/2, 1/2, 1/3, 1/2, 1/3, 1/4, 3/10, 3/4, 1 + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10] + #φ = Phi(h, ci, analytic_solution=use_analytic_solution) + + ci = [mpf(c_val) for c_val in ci] + c1, c2, c3, c4, c5, c6, c7, c8, c9, c10 = [c_val for c_val in ci] + + φ = Phi(mpf(h.item()), ci, analytic_solution=use_analytic_solution) + + a3_2 = (c3**2 / c2) * φ(2,3) + a4_2 = (c4**2 / c2) * φ(2,4) + + b8 = (c9*c10*φ(2) - 2*(c9+c10)*φ(3) + 6*φ(4)) / (c8 * (c8-c9) * (c8-c10)) + b9 = (c8*c10*φ(2) - 2*(c8+c10)*φ(3) + 6*φ(4)) / (c9 * (c9-c8) * (c9-c10)) + + b10 = (c8*c9*φ(2) - 2*(c8+c9) *φ(3) + 6*φ(4)) / (c10 * (c10-c8) * (c10-c9)) + + a = [ + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, a3_2, 0, 0, 0, 0, 0, 0, 0, 0], + [0, a4_2, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], + ] + b = [ + [0, 0, 0, 0, 0, 0, 0, b8, b9, b10], + ] + + # a5_3, a5_4 + # a6_3, a6_4 + # a7_3, a7_4 + for i in range(5, 8): # i=5,6,7 j,k ∈ {3, 4}, j != k + jk = [(3, 4), (4, 3)] + jk = list(permutations([3, 4], 2)) + for j,k in jk: + a[i-1][j-1] = (-ci[i-1]**2 * ci[k-1] * φ(2,i) + 2*ci[i-1]**3 * φ(3,i)) / (ci[j-1] * (ci[j-1] - ci[k-1])) + + for i in range(8, 11): # i=8,9,10 j,k,l ∈ {5, 6, 7}, j != k != l [ (5, 6, 7), (5, 7, 6), (6, 5, 7), (6, 7, 5), (7, 5, 6), (7, 6, 5)] 6 total coeff + jkl = list(permutations([5, 6, 7], 3)) + for j,k,l in jkl: + a[i-1][j-1] = (ci[i-1]**2 * ci[k-1] * ci[l-1] * φ(2,i) - 2*ci[i-1]**3 * (ci[k-1] + ci[l-1]) * φ(3,i) + 6*ci[i-1]**4 * φ(4,i)) / (ci[j-1] * (ci[j-1] - ci[k-1]) * (ci[j-1] - ci[l-1])) + + gen_first_col_exp(a, b, ci, φ) + + a = [[float(val) for val in row] for row in a] + b = [[float(val) for val in row] for row in b] + c1, c2, c3, c4, c5, c6, c7, c8, c9, c10 = 0, 1/2, 1/2, 1/3, 1/2, 1/3, 1/4, 3/10, 3/4, 1 + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10] + + + case "res_15s": + + c1,c2,c3,c4,c5,c6,c7,c8,c9,c10,c11,c12,c13,c14,c15 = 0, 1/2, 1/2, 1/3, 1/2, 1/5, 1/4, 18/25, 1/3, 3/10, 1/6, 90/103, 1/3, 3/10, 1/5 + c1 = 0 + c2 = c3 = c5 = 1/2 + c4 = c9 = c13 = 1/3 + c6 = c15 = 1/5 + c7 = 1/4 + c8 = 18/25 + c10 = c14 = 3/10 + c11 = 1/6 + c12 = 90/103 + c15 = 1/5 + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10, c11, c12, c13, c14, c15] + ci = [mpf(c_val) for c_val in ci] + + φ = Phi(mpf(h.item()), ci, analytic_solution=use_analytic_solution) + + a = [[mpf(0) for _ in range(15)] for _ in range(15)] + b = [[mpf(0) for _ in range(15)]] + + for i in range(3, 5): # i=3,4 j=2 + j=2 + a[i-1][j-1] = (ci[i-1]**2 / ci[j-1]) * φ(j,i) + + + for i in range(5, 8): # i=5,6,7 j,k ∈ {3, 4}, j != k + jk = list(permutations([3, 4], 2)) + for j,k in jk: + a[i-1][j-1] = (-ci[i-1]**2 * ci[k-1] * φ(2,i) + 2*ci[i-1]**3 * φ(3,i)) / prod_diff(ci[j-1], ci[k-1]) + + for i in range(8, 12): # i=8,9,10,11 j,k,l ∈ {5, 6, 7}, j != k != l [ (5, 6, 7), (5, 7, 6), (6, 5, 7), (6, 7, 5), (7, 5, 6), (7, 6, 5)] 6 total coeff + jkl = list(permutations([5, 6, 7], 3)) + for j,k,l in jkl: + a[i-1][j-1] = (ci[i-1]**2 * ci[k-1] * ci[l-1] * φ(2,i) - 2*ci[i-1]**3 * (ci[k-1] + ci[l-1]) * φ(3,i) + 6*ci[i-1]**4 * φ(4,i)) / (ci[j-1] * (ci[j-1] - ci[k-1]) * (ci[j-1] - ci[l-1])) + + for i in range(12,16): # i=12,13,14,15 + jkld = list(permutations([8,9,10,11], 4)) + for j,k,l,d in jkld: + numerator = -ci[i-1]**2 * ci[d-1]*ci[k-1]*ci[l-1] * φ(2,i) + 2*ci[i-1]**3 * (ci[d-1]*ci[k-1] + ci[d-1]*ci[l-1] + ci[k-1]*ci[l-1]) * φ(3,i) - 6*ci[i-1]**4 * (ci[d-1] + ci[k-1] + ci[l-1]) * φ(4,i) + 24*ci[i-1]**5 * φ(5,i) + a[i-1][j-1] = numerator / prod_diff(ci[j-1], ci[k-1], ci[l-1], ci[d-1]) + + """ijkl = list(permutations([12,13,14,15], 4)) + for i,j,k,l in ijkl: + #numerator = -ci[j-1]*ci[k-1]*ci[l-1]*φ(2) + 2*(ci[j-1]*ci[k-1] + ci[j-1]*ci[l-1] + ci[k-1]*ci[l-1])*φ(3) - 6*(ci[j-1] + ci[k-1] + ci[l-1])*φ(4) + 24*φ(5) + #b[0][i-1] = numerator / prod_diff(ci[i-1], ci[j-1], ci[k-1], ci[l-1]) + for jjj in range (2, 6): # 2,3,4,5 + b[0][i-1] += mu_numerator(jjj, ci[j-1], ci[i-1], ci[k-1], ci[l-1]) * φ(jjj) + b[0][i-1] /= prod_diff(ci[i-1], ci[j-1], ci[k-1], ci[l-1])""" + + ijkl = list(permutations([12,13,14,15], 4)) + for i,j,k,l in ijkl: + numerator = 0 + for jjj in range(2, 6): # 2, 3, 4, 5 + numerator += mu_numerator(jjj, ci[j-1], ci[i-1], ci[k-1], ci[l-1]) * φ(jjj) + #print(i,j,k,l) + + b[0][i-1] = numerator / prod_diff(ci[i-1], ci[j-1], ci[k-1], ci[l-1]) + + + ijkl = list(permutations([12, 13, 14, 15], 4)) + selected_permutations = {} + sign = 1 + + for i in range(12, 16): + results = [] + for j, k, l, d in ijkl: + if i != j and i != k and i != l and i != d: + numerator = 0 + for jjj in range(2, 6): # 2, 3, 4, 5 + numerator += mu_numerator(jjj, ci[j-1], ci[i-1], ci[k-1], ci[l-1]) * φ(jjj) + theta_value = numerator / prod_diff(ci[i-1], ci[j-1], ci[k-1], ci[l-1]) + results.append((theta_value, (i, j, k, l, d))) + + results.sort(key=lambda x: abs(x[0])) + + for theta_value, permutation in results: + if sign == 1 and theta_value > 0: + selected_permutations[i] = (theta_value, permutation) + sign *= -1 + break + elif sign == -1 and theta_value < 0: + selected_permutations[i] = (theta_value, permutation) + sign *= -1 + break + + for i in range(12, 16): + if i in selected_permutations: + theta_value, (i, j, k, l, d) = selected_permutations[i] + b[0][i-1] = theta_value + + for i in selected_permutations: + theta_value, permutation = selected_permutations[i] + print(f"i={i}") + print(f" Selected Theta: {theta_value:.6f}, Permutation: {permutation}") + + + gen_first_col_exp(a, b, ci, φ) + + a = [[float(val) for val in row] for row in a] + b = [[float(val) for val in row] for row in b] + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10, c11, c12, c13, c14, c15] + + + case "res_16s": # 6th order without weakened order conditions + + c1 = 0 + c2 = c3 = c5 = c8 = c12 = 1/2 + c4 = c11 = c15 = 1/3 + c6 = c9 = c13 = 1/5 + c7 = c10 = c14 = 1/4 + c16 = 1 + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10, c11, c12, c13, c14, c15, c16] + ci = [mpf(c_val) for c_val in ci] + φ = Phi(mpf(h.item()), ci, analytic_solution=use_analytic_solution) + + a3_2 = (1/2) * φ(2,3) + + a = [[mpf(0) for _ in range(16)] for _ in range(16)] + b = [[mpf(0) for _ in range(16)]] + + for i in range(3, 5): # i=3,4 j=2 + j=2 + a[i-1][j-1] = (ci[i-1]**2 / ci[j-1]) * φ(j,i) + + for i in range(5, 8): # i=5,6,7 j,k ∈ {3, 4}, j != k + jk = list(permutations([3, 4], 2)) + for j,k in jk: + a[i-1][j-1] = (-ci[i-1]**2 * ci[k-1] * φ(2,i) + 2*ci[i-1]**3 * φ(3,i)) / prod_diff(ci[j-1], ci[k-1]) + + for i in range(8, 12): # i=8,9,10,11 j,k,l ∈ {5, 6, 7}, j != k != l [ (5, 6, 7), (5, 7, 6), (6, 5, 7), (6, 7, 5), (7, 5, 6), (7, 6, 5)] 6 total coeff + jkl = list(permutations([5, 6, 7], 3)) + for j,k,l in jkl: + a[i-1][j-1] = (ci[i-1]**2 * ci[k-1] * ci[l-1] * φ(2,i) - 2*ci[i-1]**3 * (ci[k-1] + ci[l-1]) * φ(3,i) + 6*ci[i-1]**4 * φ(4,i)) / (ci[j-1] * (ci[j-1] - ci[k-1]) * (ci[j-1] - ci[l-1])) + + for i in range(12,17): # i=12,13,14,15,16 + jkld = list(permutations([8,9,10,11], 4)) + for j,k,l,d in jkld: + numerator = -ci[i-1]**2 * ci[d-1]*ci[k-1]*ci[l-1] * φ(2,i) + 2*ci[i-1]**3 * (ci[d-1]*ci[k-1] + ci[d-1]*ci[l-1] + ci[k-1]*ci[l-1]) * φ(3,i) - 6*ci[i-1]**4 * (ci[d-1] + ci[k-1] + ci[l-1]) * φ(4,i) + 24*ci[i-1]**5 * φ(5,i) + a[i-1][j-1] = numerator / prod_diff(ci[j-1], ci[k-1], ci[l-1], ci[d-1]) + + """ijdkl = list(permutations([12,13,14,15,16], 5)) + for i,j,d,k,l in ijdkl: + #numerator = -ci[j-1]*ci[k-1]*ci[l-1]*φ(2) + 2*(ci[j-1]*ci[k-1] + ci[j-1]*ci[l-1] + ci[k-1]*ci[l-1])*φ(3) - 6*(ci[j-1] + ci[k-1] + ci[l-1])*φ(4) + 24*φ(5) + b[0][i-1] = theta(2, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(2) + theta(3, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(3) + theta(4, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(4) + theta(5, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(5) + theta(6, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(6) + #b[0][i-1] = numerator / prod_diff(ci[i-1], ci[j-1], ci[k-1], ci[l-1])""" + + + ijdkl = list(permutations([12,13,14,15,16], 5)) + for i,j,d,k,l in ijdkl: + #numerator = -ci[j-1]*ci[k-1]*ci[l-1]*φ(2) + 2*(ci[j-1]*ci[k-1] + ci[j-1]*ci[l-1] + ci[k-1]*ci[l-1])*φ(3) - 6*(ci[j-1] + ci[k-1] + ci[l-1])*φ(4) + 24*φ(5) + #numerator = theta_numerator(2, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(2) + theta_numerator(3, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(3) + theta_numerator(4, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(4) + theta_numerator(5, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(5) + theta_numerator(6, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(6) + #b[0][i-1] = numerator / (ci[i-1] *, ci[d-1], ci[j-1], ci[k-1], ci[l-1]) + #b[0][i-1] = numerator / denominator(ci[i-1], ci[d-1], ci[j-1], ci[k-1], ci[l-1]) + b[0][i-1] = theta(2, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(2) + theta(3, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(3) + theta(4, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(4) + theta(5, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1])*φ(5) + theta(6, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(6) + + + ijdkl = list(permutations([12,13,14,15,16], 5)) + for i,j,d,k,l in ijdkl: + numerator = 0 + for jjj in range(2, 7): # 2, 3, 4, 5, 6 + numerator += theta_numerator(jjj, ci[d-1], ci[i-1], ci[k-1], ci[j-1], ci[l-1]) * φ(jjj) + #print(i,j,d,k,l) + b[0][i-1] = numerator / (ci[i-1] * (ci[i-1] - ci[k-1]) * (ci[i-1] - ci[j-1] * (ci[i-1] - ci[d-1]) * (ci[i-1] - ci[l-1]))) + + gen_first_col_exp(a, b, ci, φ) + + a = [[float(val) for val in row] for row in a] + b = [[float(val) for val in row] for row in b] + ci = [c1, c2, c3, c4, c5, c6, c7, c8, c9, c10, c11, c12, c13, c14, c15, c16] + + case "irk_exp_diag_2s": + c1 = 1/3 + c2 = 2/3 + c1 = float(get_extra_options_kv("c1", str(c1), extra_options)) + c2 = float(get_extra_options_kv("c2", str(c2), extra_options)) + + lam = (1 - torch.exp(-c1 * h)) / h + a2_1 = ( torch.exp(c2*h) - torch.exp(c1*h)) / (h * torch.exp(2*c1*h)) + b1 = (1 + c2*h + torch.exp(h) * (-1 + h - c2*h)) / ((c1-c2) * h**2 * torch.exp(c1*h)) + b2 = -(1 + c1*h - torch.exp(h) * ( 1 - h + c1*h)) / ((c1-c2) * h**2 * torch.exp(c2*h)) + + a = [ + [lam, 0], + [a2_1, lam], + ] + b = [ + [b1, b2], + ] + ci = [c1, c2] + + ci = ci[:] + #if rk_type.startswith("lob") == False: + ci.append(1) + + if EO("exp2lin_override_coeff") and is_exponential(rk_type): + a = scale_all(a, -sigma.item()) + b = scale_all(b, -sigma.item()) + + # Log when sampler differs from primary + is_normal_multistep = ( + (primary_rk_type.endswith("2m") and rk_type.endswith("2s") and multistep_stages >= 1) or + (primary_rk_type.endswith("3m") and rk_type.endswith("3s") and multistep_stages >= 2) or + (primary_rk_type.endswith("4m") and rk_type.endswith("4s") and multistep_stages >= 3) or + (primary_rk_type.startswith("deis") and rk_type == "deis" and multistep_stages >= 1) + ) + if rk_type != primary_rk_type and not is_normal_multistep: + if sampler_change_reason == "fallback": + h_val = h_no_eta.item() if isinstance(h_no_eta, torch.Tensor) else h_no_eta + reason_str = f" (fallback, h={h_val:.4f})" + elif sampler_change_reason == "initial": + reason_str = " (initial warmup step)" + else: + reason_str = "" + RESplain(f"step {step}: {primary_rk_type} -> {rk_type}{reason_str}", debug='debug') + + return a, b, u, v, ci, multistep_stages, hybrid_stages, FSAL + + +def scale_all(data, scalar): + if isinstance(data, torch.Tensor): + return data * scalar + elif isinstance(data, list): + return [scale_all(x, scalar) for x in data] + elif isinstance(data, (float, int)): + return data * scalar + else: + return data # passthrough unscaled if unknown type... or None, etc + + +def gen_first_col_exp(a, b, c, φ): + for i in range(len(c)): + a[i][0] = c[i] * φ(1,i+1) - sum(a[i]) + for i in range(len(b)): + b[i][0] = φ(1) - sum(b[i]) + return a, b + +def gen_first_col_exp_uv(a, b, c, u, v, φ): + for i in range(len(c)): + a[i][0] = c[i] * φ(1,i+1) - sum(a[i]) - sum(u[i]) + for i in range(len(b)): + b[i][0] = φ(1) - sum(b[i]) - sum(v[i]) + return a, b + +def rho(j, ci, ck, cl): + if j == 2: + numerator = ck*cl + if j == 3: + numerator = (-2 * (ck + cl)) + if j == 4: + numerator = 6 + return numerator / denominator(ci, ck, cl) + + +def mu(j, cd, ci, ck, cl): + if j == 2: + numerator = -cd * ck * cl + if j == 3: + numerator = 2 * (cd * ck + cd * cl + ck * cl) + if j == 4: + numerator = -6 * (cd + ck + cl) + if j == 5: + numerator = 24 + return numerator / denominator(ci, cd, ck, cl) + +def mu_numerator(j, cd, ci, ck, cl): + if j == 2: + numerator = -cd * ck * cl + if j == 3: + numerator = 2 * (cd * ck + cd * cl + ck * cl) + if j == 4: + numerator = -6 * (cd + ck + cl) + if j == 5: + numerator = 24 + return numerator #/ denominator(ci, cd, ck, cl) + + + +def theta_numerator(j, cd, ci, ck, cj, cl): + if j == 2: + numerator = -cj * cd * ck * cl + if j == 3: + numerator = 2 * (cj * ck * cd + cj*ck*cl + ck*cd*cl + cd*cl*cj) + if j == 4: + numerator = -6*(cj*ck + cj*cd + cj*cl + ck*cd + ck*cl + cd*cl) + if j == 5: + numerator = 24 * (cj + ck + cl + cd) + if j == 6: + numerator = -120 + return numerator # / denominator(ci, cj, ck, cl, cd) + + +def theta(j, cd, ci, ck, cj, cl): + if j == 2: + numerator = -cj * cd * ck * cl + if j == 3: + numerator = 2 * (cj * ck * cd + cj*ck*cl + ck*cd*cl + cd*cl*cj) + if j == 4: + numerator = -6*(cj*ck + cj*cd + cj*cl + ck*cd + ck*cl + cd*cl) + if j == 5: + numerator = 24 * (cj + ck + cl + cd) + if j == 6: + numerator = -120 + return numerator / ( ci * (ci - cj) * (ci - ck) * (ci - cl) * (ci - cd)) + return numerator / denominator(ci, cj, ck, cl, cd) + + +def prod_diff(cj, ck, cl=None, cd=None): + if cl is None and cd is None: + return cj * (cj - ck) + if cd is None: + return cj * (cj - ck) * (cj - cl) + else: + return cj * (cj - ck) * (cj - cl) * (cj - cd) + +def denominator(ci, *args): + result = ci + for arg in args: + result *= (ci - arg) + return result + + + +def check_condition_4_2(nodes): + + c12, c13, c14, c15 = nodes + + term_1 = (1 / 5) * (c12 + c13 + c14 + c15) + term_2 = (1 / 4) * (c12 * c13 + c12 * c14 + c12 * c15 + c13 * c14 + c13 * c15 + c14 * c15) + term_3 = (1 / 3) * (c12 * c13 * c14 + c12 * c13 * c15 + c12 * c14 * c15 + c13 * c14 * c15) + term_4 = (1 / 2) * (c12 * c13 * c14 * c15) + + result = term_1 - term_2 + term_3 - term_4 + + return abs(result - (1 / 6)) < 1e-6 + diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/rk_guide_func_beta.py b/simple_syrup/third_party/res4lyf_runtime/beta/rk_guide_func_beta.py new file mode 100644 index 0000000..048b8f3 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/rk_guide_func_beta.py @@ -0,0 +1,2728 @@ +import torch +import torch.nn.functional as F +from torch import Tensor + +import itertools +import copy + +from typing import Optional, Callable, Tuple, Dict, Any, Union, TYPE_CHECKING, TypeVar + +if TYPE_CHECKING: + from .noise_classes import NoiseGenerator + NoiseGeneratorSubclass = TypeVar("NoiseGeneratorSubclass", bound="NoiseGenerator") + +from einops import rearrange + +from ..sigmas import get_sigmas +from ..helper import ExtraOptions, FrameWeightsManager, initialize_or_scale, is_video_model +from ..latents import normalize_zscore, get_collinear, get_orthogonal, get_cosine_similarity, get_pearson_similarity, \ + get_slerp_weight_for_cossim, normalize_latent, hard_light_blend, slerp_tensor, get_orthogonal_noise_from_channelwise, \ + get_edge_mask, is_packed_latent + +from .rk_method_beta import RK_Method_Beta +from .constants import MAX_STEPS +from ..res4lyf import RESplain, is_debug_logging_enabled + +import comfy.utils + + +def flatten_to_match(guide: Tensor, target: Tensor) -> Tensor: + """Flatten guide tensor to match target's shape when target is flat [1,1,N]. + + If target is not flat, returns guide unchanged. + If guide is NestedTensor, packs it to flat format. + If guide is regular tensor, reshapes to [1,1,N]. + """ + if target.ndim != 3: + return guide + + # Target is flat [1,1,N], flatten the guide to match + if hasattr(guide, 'is_nested') and guide.is_nested: + # NestedTensor - pack to flat + flat_guide, _ = comfy.utils.pack_latents(guide.unbind()) + return flat_guide + else: + # Regular tensor - reshape to flat + return guide.reshape(1, 1, -1) + + +class LatentGuide: + OFFLOADABLE_ATTRS = [ + # Guide targets (full latent size) + 'y0', 'y0_inv', 'y0_mean', 'y0_adain', 'y0_attninj', 'y0_style_pos', 'y0_style_neg', + # Spatial masks (full latent size) + 'mask', 'mask_inv', 'mask_sync', 'mask_drift_x', 'mask_drift_y', + 'mask_lure_x', 'mask_lure_y', 'mask_mean', 'mask_adain', 'mask_attninj', + 'mask_style_pos', 'mask_style_neg', + # Self-refine state (full latent size) + 'self_refine_epsilon_ref', '_self_refine_iter_prediction', + '_self_refine_certain_mask_accum', '_debug_certainty_mask', + # Frame weights (if populated) + 'frame_weights', 'frame_weights_inv', + ] + + def offload(self, device): + for attr in self.OFFLOADABLE_ATTRS: + val = getattr(self, attr, None) + if isinstance(val, Tensor): + setattr(self, attr, val.to(device)) + # Handle list-of-tensor attributes + for list_attr in ('x_lying_', 's_lying_'): + val = getattr(self, list_attr, None) + if isinstance(val, list): + setattr(self, list_attr, [t.to(device) if isinstance(t, Tensor) else t for t in val]) + + def restore(self): + self.offload(self.device) + + def __init__(self, + model, + sigmas : Tensor, + UNSAMPLE : bool, + VE_MODEL : bool, + LGW_MASK_RESCALE_MIN : bool, + extra_options : str, + device : str = 'cpu', + dtype : torch.dtype = torch.float64, + frame_weights_mgr : FrameWeightsManager = None, + latent_shapes : list = None, + ): + + self.dtype = dtype + self.device = device + self.model = model + self.latent_shapes = latent_shapes + + if hasattr(model, "model"): + model_sampling = model.model.model_sampling + elif hasattr(model, "inner_model"): + model_sampling = model.inner_model.inner_model.model_sampling + + self.sigma_min = model_sampling.sigma_min.to(dtype=dtype, device=device) + self.sigma_max = model_sampling.sigma_max.to(dtype=dtype, device=device) + self.sigmas = sigmas .to(dtype=dtype, device=device) + self.UNSAMPLE = UNSAMPLE + self.VE_MODEL = VE_MODEL + self.VIDEO = is_video_model(model) + self.SAMPLE = (sigmas[0] > sigmas[1]) # type torch.bool + self.y0 = None + self.y0_inv = None + self.y0_mean = None + self.y0_adain = None + self.y0_attninj = None + self.y0_style_pos = None + self.y0_style_neg = None + + # Original shape for pack/unpack (pack-first experiment) + self.y0_original_shape = None + + self.guide_mode = "" + self.max_steps = MAX_STEPS + self.mask = None + self.mask_inv = None + self.invert_mask = False + self.mask_sync = None + self.mask_drift_x = None + self.mask_drift_y = None + self.mask_lure_x = None + self.mask_lure_y = None + self.mask_mean = None + self.mask_adain = None + self.mask_attninj = None + self.mask_style_pos = None + self.mask_style_neg = None + self.x_lying_ = None + self.s_lying_ = None + + self.LGW_MASK_RESCALE_MIN = LGW_MASK_RESCALE_MIN + self.HAS_LATENT_GUIDE = False + self.HAS_LATENT_GUIDE_INV = False + self.HAS_LATENT_GUIDE_MEAN = False + self.HAS_LATENT_GUIDE_ADAIN = False + self.HAS_LATENT_GUIDE_ATTNINJ = False + self.HAS_LATENT_GUIDE_STYLE_POS= False + self.HAS_LATENT_GUIDE_STYLE_NEG= False + self.USE_DENOISED_AS_GUIDE = False + self.SELF_REFINE_EPSILON_MODE = False + self.self_refine_epsilon_ref = None + self.self_refine_epsilon_last_step = -1 + self.self_refine_epsilon_last_row = -1 + self.self_refine_epsilon_call_count = 0 # Track calls per (step, row) + self.self_refine_threshold = 0.25 + self.self_refine_cutoff = 0.99 + self.self_refine_metric = "l1" + self._self_refine_converged = False + + self.lgw = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_inv = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_mean = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_adain = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_attninj = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_style_pos = torch.full_like(sigmas, 0., dtype=dtype) + self.lgw_style_neg = torch.full_like(sigmas, 0., dtype=dtype) + + self.cossim_tgt = torch.full_like(sigmas, 0., dtype=dtype) + self.cossim_tgt_inv = torch.full_like(sigmas, 0., dtype=dtype) + + self.guide_cossim_cutoff_ = 1.0 + self.guide_bkg_cossim_cutoff_ = 1.0 + self.guide_mean_cossim_cutoff_ = 1.0 + self.guide_adain_cossim_cutoff_ = 1.0 + self.guide_attninj_cossim_cutoff_ = 1.0 + self.guide_style_pos_cossim_cutoff_= 1.0 + self.guide_style_neg_cossim_cutoff_= 1.0 + + self.frame_weights_mgr = frame_weights_mgr + self.frame_weights = None + self.frame_weights_inv = None + + #self.freqsep_lowpass_method = "none" + #self.freqsep_sigma = 0. + #self.freqsep_kernel_size = 0 + + self.extra_options = extra_options + self.EO = ExtraOptions(extra_options) + + + def init_guides(self, + x : Tensor, + RK_IMPLICIT : bool, + guides : Optional[Tensor] = None, + noise_sampler : Optional["NoiseGeneratorSubclass"] = None, + batch_num : int = 0, + sigma_init = None, + guide_inversion_y0 = None, + guide_inversion_y0_inv = None, + ) -> Tensor: + + latent_guide_weight = 0.0 + latent_guide_weight_inv = 0.0 + latent_guide_weight_sync = 0.0 + latent_guide_weight_sync_inv = 0.0 + latent_guide_weight_drift_x = 0.0 + latent_guide_weight_drift_x_inv = 0.0 + latent_guide_weight_drift_y = 0.0 + latent_guide_weight_drift_y_inv = 0.0 + latent_guide_weight_lure_x = 0.0 + latent_guide_weight_lure_x_inv = 0.0 + latent_guide_weight_lure_y = 0.0 + latent_guide_weight_lure_y_inv = 0.0 + + latent_guide_weight_mean = 0.0 + latent_guide_weight_adain = 0.0 + latent_guide_weight_attninj = 0.0 + latent_guide_weight_style_pos = 0.0 + latent_guide_weight_style_neg = 0.0 + + latent_guide_weights = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_sync = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_sync_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_x = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_x_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_y = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_y_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_x = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_x_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_y = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_y_inv = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_mean = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_adain = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_attninj = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_style_pos = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + latent_guide_weights_style_neg = torch.zeros_like(self.sigmas, dtype=self.dtype, device=self.device) + + latent_guide = None + latent_guide_inv = None + latent_guide_mean = None + latent_guide_adain = None + latent_guide_attninj = None + latent_guide_style_pos = None + latent_guide_style_neg = None + + self.drift_x_data = 0.0 + self.drift_x_sync = 0.0 + self.drift_y_data = 0.0 + self.drift_y_sync = 0.0 + self.drift_y_guide = 0.0 + + if guides is not None: + self.guide_mode = guides.get("guide_mode", "none") + + if self.guide_mode.startswith("inversion"): + self.guide_mode = self.guide_mode.replace("inversion", "epsilon", 1) + else: + self.SAMPLE = True + self.UNSAMPLE = False + + self.self_refine_threshold = guides.get("self_refine_threshold", self.EO("self_refine_threshold", 0.25)) + self.self_refine_cutoff = guides.get("self_refine_cutoff", self.EO("self_refine_cutoff", 0.99)) + self.self_refine_metric = guides.get("self_refine_metric", self.EO("self_refine_metric", "l1")).lower() + + latent_guide_weight = guides.get("weight_masked", 0.) + latent_guide_weight_inv = guides.get("weight_unmasked", 0.) + latent_guide_weight_sync = guides.get("weight_masked_sync", 0.) + latent_guide_weight_sync_inv = guides.get("weight_unmasked_sync", 0.) + latent_guide_weight_drift_x = guides.get("weight_masked_drift_x", 0.) + latent_guide_weight_drift_x_inv = guides.get("weight_unmasked_drift_x", 0.) + latent_guide_weight_drift_y = guides.get("weight_masked_drift_y", 0.) + latent_guide_weight_drift_y_inv = guides.get("weight_unmasked_drift_y", 0.) + latent_guide_weight_lure_x = guides.get("weight_masked_lure_x", 0.) + latent_guide_weight_lure_x_inv = guides.get("weight_unmasked_lure_x", 0.) + latent_guide_weight_lure_y = guides.get("weight_masked_lure_y", 0.) + latent_guide_weight_lure_y_inv = guides.get("weight_unmasked_lure_y", 0.) + latent_guide_weight_mean = guides.get("weight_mean", 0.) + latent_guide_weight_adain = guides.get("weight_adain", 0.) + latent_guide_weight_attninj = guides.get("weight_attninj", 0.) + latent_guide_weight_style_pos = guides.get("weight_style_pos", 0.) + latent_guide_weight_style_neg = guides.get("weight_style_neg", 0.) + #latent_guide_synweight_style_pos = guides.get("synweight_style_pos", 0.) + #latent_guide_synweight_style_neg = guides.get("synweight_style_neg", 0.) + + self.drift_x_data = guides.get("drift_x_data", 0.) + self.drift_x_sync = guides.get("drift_x_sync", 0.) + self.drift_y_data = guides.get("drift_y_data", 0.) + self.drift_y_sync = guides.get("drift_y_sync", 0.) + self.drift_y_guide = guides.get("drift_y_guide", 0.) + + latent_guide_weights = guides.get("weights_masked") + latent_guide_weights_inv = guides.get("weights_unmasked") + latent_guide_weights_sync = guides.get("weights_masked_sync") + latent_guide_weights_sync_inv = guides.get("weights_unmasked_sync") + latent_guide_weights_drift_x = guides.get("weights_masked_drift_x") + latent_guide_weights_drift_x_inv = guides.get("weights_unmasked_drift_x") + latent_guide_weights_drift_y = guides.get("weights_masked_drift_y") + latent_guide_weights_drift_y_inv = guides.get("weights_unmasked_drift_y") + latent_guide_weights_lure_x = guides.get("weights_masked_lure_x") + latent_guide_weights_lure_x_inv = guides.get("weights_unmasked_lure_x") + latent_guide_weights_lure_y = guides.get("weights_masked_lure_y") + latent_guide_weights_lure_y_inv = guides.get("weights_unmasked_lure_y") + latent_guide_weights_mean = guides.get("weights_mean") + latent_guide_weights_adain = guides.get("weights_adain") + latent_guide_weights_attninj = guides.get("weights_attninj") + latent_guide_weights_style_pos = guides.get("weights_style_pos") + latent_guide_weights_style_neg = guides.get("weights_style_neg") + #latent_guide_synweights_style_p os = guides.get("synweights_style_pos") + #latent_guide_synweights_style_neg = guides.get("synweights_style_neg") + + latent_guide = guides.get("guide_masked") + latent_guide_inv = guides.get("guide_unmasked") + latent_guide_mean = guides.get("guide_mean") + latent_guide_adain = guides.get("guide_adain") + latent_guide_attninj = guides.get("guide_attninj") + latent_guide_style_pos = guides.get("guide_style_pos") + latent_guide_style_neg = guides.get("guide_style_neg") + + self.mask = guides.get("mask") + self.mask_inv = guides.get("unmask") + self.invert_mask = guides.get("invert_mask", False) + self.mask_sync = guides.get("mask_sync") + self.mask_drift_x = guides.get("mask_drift_x") + self.mask_drift_y = guides.get("mask_drift_y") + self.mask_lure_x = guides.get("mask_lure_x") + self.mask_lure_y = guides.get("mask_lure_y") + self.mask_mean = guides.get("mask_mean") + self.mask_adain = guides.get("mask_adain") + self.mask_attninj = guides.get("mask_attninj") + self.mask_style_pos = guides.get("mask_style_pos") + self.mask_style_neg = guides.get("mask_style_neg") + + scheduler_ = guides.get("weight_scheduler_masked") + scheduler_inv_ = guides.get("weight_scheduler_unmasked") + scheduler_sync_ = guides.get("weight_scheduler_masked_sync") + scheduler_sync_inv_ = guides.get("weight_scheduler_unmasked_sync") + scheduler_drift_x_ = guides.get("weight_scheduler_masked_drift_x") + scheduler_drift_x_inv_ = guides.get("weight_scheduler_unmasked_drift_x") + scheduler_drift_y_ = guides.get("weight_scheduler_masked_drift_y") + scheduler_drift_y_inv_ = guides.get("weight_scheduler_unmasked_drift_y") + scheduler_lure_x_ = guides.get("weight_scheduler_masked_lure_x") + scheduler_lure_x_inv_ = guides.get("weight_scheduler_unmasked_lure_x") + scheduler_lure_y_ = guides.get("weight_scheduler_masked_lure_y") + scheduler_lure_y_inv_ = guides.get("weight_scheduler_unmasked_lure_y") + scheduler_mean_ = guides.get("weight_scheduler_mean") + scheduler_adain_ = guides.get("weight_scheduler_adain") + scheduler_attninj_ = guides.get("weight_scheduler_attninj") + scheduler_style_pos_ = guides.get("weight_scheduler_style_pos") + scheduler_style_neg_ = guides.get("weight_scheduler_style_neg") + + start_steps_ = guides.get("start_step_masked", 0) + start_steps_inv_ = guides.get("start_step_unmasked", 0) + start_steps_sync_ = guides.get("start_step_masked_sync", 0) + start_steps_sync_inv_ = guides.get("start_step_unmasked_sync", 0) + start_steps_drift_x_ = guides.get("start_step_masked_drift_x", 0) + start_steps_drift_x_inv_ = guides.get("start_step_unmasked_drift_x", 0) + start_steps_drift_y_ = guides.get("start_step_masked_drift_y", 0) + start_steps_drift_y_inv_ = guides.get("start_step_unmasked_drift_y", 0) + start_steps_lure_x_ = guides.get("start_step_masked_lure_x", 0) + start_steps_lure_x_inv_ = guides.get("start_step_unmasked_lure_x", 0) + start_steps_lure_y_ = guides.get("start_step_masked_lure_y", 0) + start_steps_lure_y_inv_ = guides.get("start_step_unmasked_lure_y", 0) + start_steps_mean_ = guides.get("start_step_mean", 0) + start_steps_adain_ = guides.get("start_step_adain", 0) + start_steps_attninj_ = guides.get("start_step_attninj", 0) + start_steps_style_pos_ = guides.get("start_step_style_pos", 0) + start_steps_style_neg_ = guides.get("start_step_style_neg", 0) + + steps_ = guides.get("end_step_masked", 1) + steps_inv_ = guides.get("end_step_unmasked", 1) + steps_sync_ = guides.get("end_step_masked_sync", 1) + steps_sync_inv_ = guides.get("end_step_unmasked_sync", 1) + steps_drift_x_ = guides.get("end_step_masked_drift_x", 1) + steps_drift_x_inv_ = guides.get("end_step_unmasked_drift_x", 1) + steps_drift_y_ = guides.get("end_step_masked_drift_y", 1) + steps_drift_y_inv_ = guides.get("end_step_unmasked_drift_y", 1) + steps_lure_x_ = guides.get("end_step_masked_lure_x", 1) + steps_lure_x_inv_ = guides.get("end_step_unmasked_lure_x", 1) + steps_lure_y_ = guides.get("end_step_masked_lure_y", 1) + steps_lure_y_inv_ = guides.get("end_step_unmasked_lure_y", 1) + + steps_mean_ = guides.get("end_step_mean", 1) + steps_adain_ = guides.get("end_step_adain", 1) + steps_attninj_ = guides.get("end_step_attninj", 1) + steps_style_pos_ = guides.get("end_step_style_pos", 1) + steps_style_neg_ = guides.get("end_step_style_neg", 1) + + self.guide_cossim_cutoff_ = guides.get("cutoff_masked", 1.) + self.guide_bkg_cossim_cutoff_ = guides.get("cutoff_unmasked", 1.) + self.guide_mean_cossim_cutoff_ = guides.get("cutoff_mean", 1.) + self.guide_adain_cossim_cutoff_ = guides.get("cutoff_adain", 1.) + self.guide_attninj_cossim_cutoff_ = guides.get("cutoff_attninj", 1.) + self.guide_style_pos_cossim_cutoff_ = guides.get("cutoff_style_pos", 1.) + self.guide_style_neg_cossim_cutoff_ = guides.get("cutoff_style_neg", 1.) + + self.sync_lure_iter = guides.get("sync_lure_iter", 0) + self.sync_lure_sequence = guides.get("sync_lure_sequence") + + #self.SYNC_SEPARATE = False + #if scheduler_sync_ is not None: + # self.SYNC_SEPARATE = True + self.SYNC_SEPARATE = True + if scheduler_sync_ is None and scheduler_ is not None: + + latent_guide_weight_sync = latent_guide_weight + latent_guide_weight_sync_inv = latent_guide_weight_inv + latent_guide_weights_sync = latent_guide_weights + latent_guide_weights_sync_inv = latent_guide_weights_inv + + scheduler_sync_ = scheduler_ + scheduler_sync_inv_ = scheduler_inv_ + + start_steps_sync_ = start_steps_ + start_steps_sync_inv_ = start_steps_inv_ + + steps_sync_ = steps_ + steps_sync_inv_ = steps_inv_ + + self.SYNC_drift_X = True + if scheduler_drift_x_ is None and scheduler_ is not None: + self.SYNC_drift_X = False + + latent_guide_weight_drift_x = latent_guide_weight + latent_guide_weight_drift_x_inv = latent_guide_weight_inv + latent_guide_weights_drift_x = latent_guide_weights + latent_guide_weights_drift_x_inv = latent_guide_weights_inv + + scheduler_drift_x_ = scheduler_ + scheduler_drift_x_inv_ = scheduler_inv_ + + start_steps_drift_x_ = start_steps_ + start_steps_drift_x_inv_ = start_steps_inv_ + + steps_drift_x_ = steps_ + steps_drift_x_inv_ = steps_inv_ + + self.SYNC_drift_Y = True + if scheduler_drift_y_ is None and scheduler_ is not None: + self.SYNC_drift_Y = False + + latent_guide_weight_drift_y = latent_guide_weight + latent_guide_weight_drift_y_inv = latent_guide_weight_inv + latent_guide_weights_drift_y = latent_guide_weights + latent_guide_weights_drift_y_inv = latent_guide_weights_inv + + scheduler_drift_y_ = scheduler_ + scheduler_drift_y_inv_ = scheduler_inv_ + + start_steps_drift_y_ = start_steps_ + start_steps_drift_y_inv_ = start_steps_inv_ + + steps_drift_y_ = steps_ + steps_drift_y_inv_ = steps_inv_ + + self.SYNC_LURE_X = True + if scheduler_lure_x_ is None and scheduler_ is not None: + self.SYNC_LURE_X = False + + latent_guide_weight_lure_x = latent_guide_weight + latent_guide_weight_lure_x_inv = latent_guide_weight_inv + latent_guide_weights_lure_x = latent_guide_weights + latent_guide_weights_lure_x_inv = latent_guide_weights_inv + + scheduler_lure_x_ = scheduler_ + scheduler_lure_x_inv_ = scheduler_inv_ + + start_steps_lure_x_ = start_steps_ + start_steps_lure_x_inv_ = start_steps_inv_ + + steps_lure_x_ = steps_ + steps_lure_x_inv_ = steps_inv_ + + self.SYNC_LURE_Y = True + if scheduler_lure_y_ is None and scheduler_ is not None: + self.SYNC_LURE_Y = False + + latent_guide_weight_lure_y = latent_guide_weight + latent_guide_weight_lure_y_inv = latent_guide_weight_inv + latent_guide_weights_lure_y = latent_guide_weights + latent_guide_weights_lure_y_inv = latent_guide_weights_inv + + scheduler_lure_y_ = scheduler_ + scheduler_lure_y_inv_ = scheduler_inv_ + + start_steps_lure_y_ = start_steps_ + start_steps_lure_y_inv_ = start_steps_inv_ + + steps_lure_y_ = steps_ + steps_lure_y_inv_ = steps_inv_ + + if self.guide_mode.startswith("fully_") and not RK_IMPLICIT: + raise ValueError("fully_pseudoimplicit is only supported for implicit RK samplers.") + #self.guide_mode = self.guide_mode[6:] # fully_pseudoimplicit is only supported for implicit samplers, default back to pseudoimplicit + + guide_sigma_shift = self.EO("guide_sigma_shift", 0.0) # effectively hardcoding shift to 0 !!!!!! + + if latent_guide_weights is None and scheduler_ is not None: + total_steps = steps_ - start_steps_ + latent_guide_weights = get_sigmas(self.model, scheduler_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_, dtype=self.dtype, device=self.device) + latent_guide_weights = torch.cat((prepend, latent_guide_weights.to(self.device)), dim=0) + + if latent_guide_weights_inv is None and scheduler_inv_ is not None: + total_steps = steps_inv_ - start_steps_inv_ + latent_guide_weights_inv = get_sigmas(self.model, scheduler_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_inv = torch.cat((prepend, latent_guide_weights_inv.to(self.device)), dim=0) + + if latent_guide_weights_sync is None and scheduler_sync_ is not None: + total_steps = steps_sync_ - start_steps_sync_ + latent_guide_weights_sync = get_sigmas(self.model, scheduler_sync_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_sync_, dtype=self.dtype, device=self.device) + latent_guide_weights_sync = torch.cat((prepend, latent_guide_weights_sync.to(self.device)), dim=0) + + if latent_guide_weights_sync_inv is None and scheduler_sync_inv_ is not None: + total_steps = steps_sync_inv_ - start_steps_sync_inv_ + latent_guide_weights_sync_inv = get_sigmas(self.model, scheduler_sync_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_sync_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_sync_inv = torch.cat((prepend, latent_guide_weights_sync_inv.to(self.device)), dim=0) + + if latent_guide_weights_drift_x is None and scheduler_drift_x_ is not None: + total_steps = steps_drift_x_ - start_steps_drift_x_ + latent_guide_weights_drift_x = get_sigmas(self.model, scheduler_drift_x_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_drift_x_, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_x = torch.cat((prepend, latent_guide_weights_drift_x.to(self.device)), dim=0) + + if latent_guide_weights_drift_x_inv is None and scheduler_drift_x_inv_ is not None: + total_steps = steps_drift_x_inv_ - start_steps_drift_x_inv_ + latent_guide_weights_drift_x_inv = get_sigmas(self.model, scheduler_drift_x_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_drift_x_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_x_inv = torch.cat((prepend, latent_guide_weights_drift_x_inv.to(self.device)), dim=0) + + if latent_guide_weights_drift_y is None and scheduler_drift_y_ is not None: + total_steps = steps_drift_y_ - start_steps_drift_y_ + latent_guide_weights_drift_y = get_sigmas(self.model, scheduler_drift_y_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_drift_y_, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_y = torch.cat((prepend, latent_guide_weights_drift_y.to(self.device)), dim=0) + + if latent_guide_weights_drift_y_inv is None and scheduler_drift_y_inv_ is not None: + total_steps = steps_drift_y_inv_ - start_steps_drift_y_inv_ + latent_guide_weights_drift_y_inv = get_sigmas(self.model, scheduler_drift_y_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_drift_y_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_drift_y_inv = torch.cat((prepend, latent_guide_weights_drift_y_inv.to(self.device)), dim=0) + + if latent_guide_weights_lure_x is None and scheduler_lure_x_ is not None: + total_steps = steps_lure_x_ - start_steps_lure_x_ + latent_guide_weights_lure_x = get_sigmas(self.model, scheduler_lure_x_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_lure_x_, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_x = torch.cat((prepend, latent_guide_weights_lure_x.to(self.device)), dim=0) + + if latent_guide_weights_lure_x_inv is None and scheduler_lure_x_inv_ is not None: + total_steps = steps_lure_x_inv_ - start_steps_lure_x_inv_ + latent_guide_weights_lure_x_inv = get_sigmas(self.model, scheduler_lure_x_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_lure_x_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_x_inv = torch.cat((prepend, latent_guide_weights_lure_x_inv.to(self.device)), dim=0) + + if latent_guide_weights_lure_y is None and scheduler_lure_y_ is not None: + total_steps = steps_lure_y_ - start_steps_lure_y_ + latent_guide_weights_lure_y = get_sigmas(self.model, scheduler_lure_y_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_lure_y_, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_y = torch.cat((prepend, latent_guide_weights_lure_y.to(self.device)), dim=0) + + if latent_guide_weights_lure_y_inv is None and scheduler_lure_y_inv_ is not None: + total_steps = steps_lure_y_inv_ - start_steps_lure_y_inv_ + latent_guide_weights_lure_y_inv = get_sigmas(self.model, scheduler_lure_y_inv_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_lure_y_inv_, dtype=self.dtype, device=self.device) + latent_guide_weights_lure_y_inv = torch.cat((prepend, latent_guide_weights_lure_y_inv.to(self.device)), dim=0) + + + if latent_guide_weights_mean is None and scheduler_mean_ is not None: + total_steps = steps_mean_ - start_steps_mean_ + latent_guide_weights_mean = get_sigmas(self.model, scheduler_mean_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_mean_, dtype=self.dtype, device=self.device) + latent_guide_weights_mean = torch.cat((prepend, latent_guide_weights_mean.to(self.device)), dim=0) + + if latent_guide_weights_adain is None and scheduler_adain_ is not None: + total_steps = steps_adain_ - start_steps_adain_ + latent_guide_weights_adain = get_sigmas(self.model, scheduler_adain_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_adain_, dtype=self.dtype, device=self.device) + latent_guide_weights_adain = torch.cat((prepend, latent_guide_weights_adain.to(self.device)), dim=0) + + if latent_guide_weights_attninj is None and scheduler_attninj_ is not None: + total_steps = steps_attninj_ - start_steps_attninj_ + latent_guide_weights_attninj = get_sigmas(self.model, scheduler_attninj_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_attninj_, dtype=self.dtype, device=self.device) + latent_guide_weights_attninj = torch.cat((prepend, latent_guide_weights_attninj.to(self.device)), dim=0) + + if latent_guide_weights_style_pos is None and scheduler_style_pos_ is not None: + total_steps = steps_style_pos_ - start_steps_style_pos_ + latent_guide_weights_style_pos = get_sigmas(self.model, scheduler_style_pos_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_style_pos_, dtype=self.dtype, device=self.device) + latent_guide_weights_style_pos = torch.cat((prepend, latent_guide_weights_style_pos.to(self.device)), dim=0) + + if latent_guide_weights_style_neg is None and scheduler_style_neg_ is not None: + total_steps = steps_style_neg_ - start_steps_style_neg_ + latent_guide_weights_style_neg = get_sigmas(self.model, scheduler_style_neg_, total_steps, 1.0, shift=guide_sigma_shift).to(dtype=self.dtype, device=self.device) / self.sigma_max + prepend = torch.zeros(start_steps_style_neg_, dtype=self.dtype, device=self.device) + latent_guide_weights_style_neg = torch.cat((prepend, latent_guide_weights_style_neg.to(self.device)), dim=0) + + if scheduler_ != "constant": + latent_guide_weights = initialize_or_scale(latent_guide_weights, latent_guide_weight, self.max_steps) + if scheduler_inv_ != "constant": + latent_guide_weights_inv = initialize_or_scale(latent_guide_weights_inv, latent_guide_weight_inv, self.max_steps) + if scheduler_sync_ != "constant": + latent_guide_weights_sync = initialize_or_scale(latent_guide_weights_sync, latent_guide_weight_sync, self.max_steps) + if scheduler_sync_inv_ != "constant": + latent_guide_weights_sync_inv = initialize_or_scale(latent_guide_weights_sync_inv, latent_guide_weight_sync_inv, self.max_steps) + + latent_guide_weights_sync = 1 - latent_guide_weights_sync if latent_guide_weights_sync is not None else latent_guide_weights + latent_guide_weights_sync_inv = 1 - latent_guide_weights_sync_inv if latent_guide_weights_sync_inv is not None else latent_guide_weights_inv + latent_guide_weight_sync = 1 - latent_guide_weight_sync + latent_guide_weight_sync_inv = 1 - latent_guide_weight_sync_inv# these are more intuitive to use if these are reversed... so that sync weight = 1.0 means "maximum guide strength" + + + if scheduler_drift_x_ != "constant": + latent_guide_weights_drift_x = initialize_or_scale(latent_guide_weights_drift_x, latent_guide_weight_drift_x, self.max_steps) + if scheduler_drift_x_inv_ != "constant": + latent_guide_weights_drift_x_inv = initialize_or_scale(latent_guide_weights_drift_x_inv, latent_guide_weight_drift_x_inv, self.max_steps) + if scheduler_drift_y_ != "constant": + latent_guide_weights_drift_y = initialize_or_scale(latent_guide_weights_drift_y, latent_guide_weight_drift_y, self.max_steps) + if scheduler_drift_y_inv_ != "constant": + latent_guide_weights_drift_y_inv = initialize_or_scale(latent_guide_weights_drift_y_inv, latent_guide_weight_drift_y_inv, self.max_steps) + if scheduler_lure_x_ != "constant": + latent_guide_weights_lure_x = initialize_or_scale(latent_guide_weights_lure_x, latent_guide_weight_lure_x, self.max_steps) + if scheduler_lure_x_inv_ != "constant": + latent_guide_weights_lure_x_inv = initialize_or_scale(latent_guide_weights_lure_x_inv, latent_guide_weight_lure_x_inv, self.max_steps) + if scheduler_lure_y_ != "constant": + latent_guide_weights_lure_y = initialize_or_scale(latent_guide_weights_lure_y, latent_guide_weight_lure_y, self.max_steps) + if scheduler_lure_y_inv_ != "constant": + latent_guide_weights_lure_y_inv = initialize_or_scale(latent_guide_weights_lure_y_inv, latent_guide_weight_lure_y_inv, self.max_steps) + if scheduler_mean_ != "constant": + latent_guide_weights_mean = initialize_or_scale(latent_guide_weights_mean, latent_guide_weight_mean, self.max_steps) + if scheduler_adain_ != "constant": + latent_guide_weights_adain = initialize_or_scale(latent_guide_weights_adain, latent_guide_weight_adain, self.max_steps) + if scheduler_attninj_ != "constant": + latent_guide_weights_attninj = initialize_or_scale(latent_guide_weights_attninj, latent_guide_weight_attninj, self.max_steps) + if scheduler_style_pos_ != "constant": + latent_guide_weights_style_pos = initialize_or_scale(latent_guide_weights_style_pos, latent_guide_weight_style_pos, self.max_steps) + if scheduler_style_neg_ != "constant": + latent_guide_weights_style_neg = initialize_or_scale(latent_guide_weights_style_neg, latent_guide_weight_style_neg, self.max_steps) + + latent_guide_weights [steps_ :] = 0 + latent_guide_weights_inv [steps_inv_ :] = 0 + latent_guide_weights_sync [steps_sync_ :] = 1 #one + latent_guide_weights_sync_inv [steps_sync_inv_ :] = 1 #one + latent_guide_weights_drift_x [steps_drift_x_ :] = 0 + latent_guide_weights_drift_x_inv[steps_drift_x_inv_:] = 0 + latent_guide_weights_drift_y [steps_drift_y_ :] = 0 + latent_guide_weights_drift_y_inv[steps_drift_y_inv_:] = 0 + latent_guide_weights_lure_x [steps_lure_x_ :] = 0 + latent_guide_weights_lure_x_inv [steps_lure_x_inv_ :] = 0 + latent_guide_weights_lure_y [steps_lure_y_ :] = 0 + latent_guide_weights_lure_y_inv [steps_lure_y_inv_ :] = 0 + latent_guide_weights_mean [steps_mean_ :] = 0 + latent_guide_weights_adain [steps_adain_ :] = 0 + latent_guide_weights_attninj [steps_attninj_ :] = 0 + latent_guide_weights_style_pos [steps_style_pos_ :] = 0 + latent_guide_weights_style_neg [steps_style_neg_ :] = 0 + + self.lgw = F.pad(latent_guide_weights, (0, self.max_steps), value=0.0) + self.lgw_inv = F.pad(latent_guide_weights_inv, (0, self.max_steps), value=0.0) + self.lgw_sync = F.pad(latent_guide_weights_sync, (0, self.max_steps), value=1.0) #one + self.lgw_sync_inv = F.pad(latent_guide_weights_sync_inv, (0, self.max_steps), value=1.0) #one + self.lgw_drift_x = F.pad(latent_guide_weights_drift_x, (0, self.max_steps), value=0.0) + self.lgw_drift_x_inv = F.pad(latent_guide_weights_drift_x_inv, (0, self.max_steps), value=0.0) + self.lgw_drift_y = F.pad(latent_guide_weights_drift_y, (0, self.max_steps), value=0.0) + self.lgw_drift_y_inv = F.pad(latent_guide_weights_drift_y_inv, (0, self.max_steps), value=0.0) + self.lgw_lure_x = F.pad(latent_guide_weights_lure_x, (0, self.max_steps), value=0.0) + self.lgw_lure_x_inv = F.pad(latent_guide_weights_lure_x_inv, (0, self.max_steps), value=0.0) + self.lgw_lure_y = F.pad(latent_guide_weights_lure_y, (0, self.max_steps), value=0.0) + self.lgw_lure_y_inv = F.pad(latent_guide_weights_lure_y_inv, (0, self.max_steps), value=0.0) + self.lgw_mean = F.pad(latent_guide_weights_mean, (0, self.max_steps), value=0.0) + self.lgw_adain = F.pad(latent_guide_weights_adain, (0, self.max_steps), value=0.0) + self.lgw_attninj = F.pad(latent_guide_weights_attninj, (0, self.max_steps), value=0.0) + self.lgw_style_pos = F.pad(latent_guide_weights_style_pos, (0, self.max_steps), value=0.0) + self.lgw_style_neg = F.pad(latent_guide_weights_style_neg, (0, self.max_steps), value=0.0) + + mask, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask, self.LGW_MASK_RESCALE_MIN) + self.mask = mask.to(dtype=self.dtype, device=self.device) + + if self.mask_inv is not None: + mask_inv, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_inv, self.LGW_MASK_RESCALE_MIN) + self.mask_inv = mask_inv.to(dtype=self.dtype, device=self.device) + else: + self.mask_inv = (1-self.mask) + + if self.mask_sync is not None: + mask_sync, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_sync, self.LGW_MASK_RESCALE_MIN) + self.mask_sync = mask_sync.to(dtype=self.dtype, device=self.device) + else: + self.mask_sync = self.mask + + if self.mask_drift_x is not None: + mask_drift_x, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_drift_x, self.LGW_MASK_RESCALE_MIN) + self.mask_drift_x = mask_drift_x.to(dtype=self.dtype, device=self.device) + else: + self.mask_drift_x = self.mask + + if self.mask_drift_y is not None: + mask_drift_y, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_drift_y, self.LGW_MASK_RESCALE_MIN) + self.mask_drift_y = mask_drift_y.to(dtype=self.dtype, device=self.device) + else: + self.mask_drift_y = self.mask + + if self.mask_lure_x is not None: + mask_lure_x, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_lure_x, self.LGW_MASK_RESCALE_MIN) + self.mask_lure_x = mask_lure_x.to(dtype=self.dtype, device=self.device) + else: + self.mask_lure_x = self.mask + + if self.mask_lure_y is not None: + mask_lure_y, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_lure_y, self.LGW_MASK_RESCALE_MIN) + self.mask_lure_y = mask_lure_y.to(dtype=self.dtype, device=self.device) + else: + self.mask_lure_y = self.mask + + mask_style_pos, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_style_pos, self.LGW_MASK_RESCALE_MIN) + self.mask_style_pos = mask_style_pos.to(dtype=self.dtype, device=self.device) + + + mask_style_neg, self.LGW_MASK_RESCALE_MIN = prepare_mask(x, self.mask_style_neg, self.LGW_MASK_RESCALE_MIN) + self.mask_style_neg = mask_style_neg.to(dtype=self.dtype, device=self.device) + + if latent_guide is not None: + self.HAS_LATENT_GUIDE = True + if type(latent_guide) is dict: + latent_guide_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide) is torch.Tensor: + latent_guide_samples = latent_guide.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide)}") + + latent_guide_samples = flatten_to_match(latent_guide_samples, x).clone() + + if self.SAMPLE: + self.y0 = latent_guide_samples + elif sigma_init != 0.0: + pass + elif self.UNSAMPLE: # and self.mask is not None: + mask = self.mask.to(x.device) + x = (1-mask) * x + mask * latent_guide_samples.to(x.device) + else: + x = latent_guide_samples.to(x.device) + else: + self.y0 = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + # Initialize self_refine_epsilon mode (including projection variant via _projection suffix) + self.SELF_REFINE_EPSILON_MODE = self.guide_mode.startswith("self_refine_epsilon") + if self.SELF_REFINE_EPSILON_MODE: + self.HAS_LATENT_GUIDE = True # Enable guide processing + self.y0 = torch.zeros_like(x, dtype=self.dtype, device=self.device) # y0 will be set dynamically to denoised_prev + self.self_refine_epsilon_ref = None # Reference for within-step refinement + self.self_refine_epsilon_last_step = -1 # Track which step we're on + self.self_refine_epsilon_last_row = -1 + self.self_refine_epsilon_call_count = 0 + # Per-iteration tracking (for self_refine_per_iteration mode) + self._self_refine_last_iter = -1 + self._self_refine_iter_prediction = None + self._self_refine_certain_mask_accum = None + self._debug_certainty_mask = None + + if latent_guide_inv is not None: + self.HAS_LATENT_GUIDE_INV = True + if type(latent_guide_inv) is dict: + latent_guide_inv_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_inv['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_inv) is torch.Tensor: + latent_guide_inv_samples = latent_guide_inv.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_inv)}") + + latent_guide_inv_samples = flatten_to_match(latent_guide_inv_samples, x).clone() + + if self.SAMPLE: + self.y0_inv = latent_guide_inv_samples + elif sigma_init != 0.0: + pass + elif self.UNSAMPLE: # and self.mask is not None: + mask_inv = self.mask_inv.to(x.device) + x = (1-mask_inv) * x + mask_inv * latent_guide_inv_samples.to(x.device) #fixed old approach, which was mask, (1-mask) + else: + x = latent_guide_inv_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_INV AFTER SETTING TO LG above! + else: + self.y0_inv = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + if latent_guide_mean is not None: + self.HAS_LATENT_GUIDE_MEAN = True + if type(latent_guide_mean) is dict: + latent_guide_mean_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_mean['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_mean) is torch.Tensor: + latent_guide_mean_samples = latent_guide_mean.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_mean)}") + + latent_guide_mean_samples = flatten_to_match(latent_guide_mean_samples, x).clone() + self.y0_mean = latent_guide_mean_samples + """if self.SAMPLE: + self.y0_mean = latent_guide_mean_samples + elif self.UNSAMPLE: # and self.mask is not None: + mask_mean = self.mask_mean.to(x.device) + x = (1-mask_mean) * x + mask_mean * latent_guide_mean_samples.to(x.device) #fixed old approach, which was mask, (1-mask) # NECESSARY? + else: + x = latent_guide_mean_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_MEAN AFTER SETTING TO LG above!""" + else: + self.y0_mean = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + if latent_guide_adain is not None: + self.HAS_LATENT_GUIDE_ADAIN = True + if type(latent_guide_adain) is dict: + latent_guide_adain_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_adain['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_adain) is torch.Tensor: + latent_guide_adain_samples = latent_guide_adain.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_adain)}") + + latent_guide_adain_samples = flatten_to_match(latent_guide_adain_samples, x).clone() + self.y0_adain = latent_guide_adain_samples + """if self.SAMPLE: + self.y0_adain = latent_guide_adain_samples + elif self.UNSAMPLE: # and self.mask is not None: + if self.mask_adain is not None: + mask_adain = self.mask_adain.to(x.device) + x = (1-mask_adain) * x + mask_adain * latent_guide_adain_samples.to(x.device) #fixed old approach, which was mask, (1-mask) # NECESSARY? + else: + x = latent_guide_adain_samples.to(x.device) + else: + x = latent_guide_adain_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_ADAIN AFTER SETTING TO LG above!""" + else: + self.y0_adain = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + if latent_guide_attninj is not None: + self.HAS_LATENT_GUIDE_ATTNINJ = True + if type(latent_guide_attninj) is dict: + latent_guide_attninj_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_attninj['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_attninj) is torch.Tensor: + latent_guide_attninj_samples = latent_guide_attninj.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_attninj)}") + + latent_guide_attninj_samples = flatten_to_match(latent_guide_attninj_samples, x).clone() + self.y0_attninj = latent_guide_attninj_samples + """if self.SAMPLE: + self.y0_attninj = latent_guide_attninj_samples + elif self.UNSAMPLE: # and self.mask is not None: + if self.mask_attninj is not None: + mask_attninj = self.mask_attninj.to(x.device) + x = (1-mask_attninj) * x + mask_attninj * latent_guide_attninj_samples.to(x.device) #fixed old approach, which was mask, (1-mask) # NECESSARY? + else: + x = latent_guide_attninj_samples.to(x.device) + else: + x = latent_guide_attninj_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_ADAIN AFTER SETTING TO LG above!""" + else: + self.y0_attninj = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + + if latent_guide_style_pos is not None: + self.HAS_LATENT_GUIDE_STYLE_POS = True + if type(latent_guide_style_pos) is dict: + latent_guide_style_pos_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_style_pos['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_style_pos) is torch.Tensor: + latent_guide_style_pos_samples = latent_guide_style_pos.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_style_pos)}") + + latent_guide_style_pos_samples = flatten_to_match(latent_guide_style_pos_samples, x).clone() + self.y0_style_pos = latent_guide_style_pos_samples + """if self.SAMPLE: + self.y0_style_pos = latent_guide_style_pos_samples + elif self.UNSAMPLE: # and self.mask is not None: + if self.mask_style_pos is not None: + mask_style_pos = self.mask_style_pos.to(x.device) + x = (1-mask_style_pos) * x + mask_style_pos * latent_guide_style_pos_samples.to(x.device) #fixed old approach, which was mask, (1-mask) # NECESSARY? + else: + x = latent_guide_style_pos_samples.to(x.device) + else: + x = latent_guide_style_pos_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_ADAIN AFTER SETTING TO LG above!""" + else: + self.y0_style_pos = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + + if latent_guide_style_neg is not None: + self.HAS_LATENT_GUIDE_STYLE_NEG = True + if type(latent_guide_style_neg) is dict: + latent_guide_style_neg_samples = self.model.inner_model.inner_model.process_latent_in(latent_guide_style_neg['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent_guide_style_neg) is torch.Tensor: + latent_guide_style_neg_samples = latent_guide_style_neg.to(dtype=self.dtype, device=self.device) + else: + raise ValueError(f"Invalid latent type: {type(latent_guide_style_neg)}") + + latent_guide_style_neg_samples = flatten_to_match(latent_guide_style_neg_samples, x).clone() + self.y0_style_neg = latent_guide_style_neg_samples + """if self.SAMPLE: + self.y0_style_neg = latent_guide_style_neg_samples + elif self.UNSAMPLE: # and self.mask is not None: + if self.mask_style_neg is not None: + mask_style_neg = self.mask_style_neg.to(x.device) + x = (1-mask_style_neg) * x + mask_style_neg * latent_guide_style_neg_samples.to(x.device) #fixed old approach, which was mask, (1-mask) # NECESSARY? + else: + x = latent_guide_style_neg_samples.to(x.device) + else: + x = latent_guide_style_neg_samples.to(x.device) #THIS COULD LEAD TO WEIRD BEHAVIOR! OVERWRITING X WITH LG_ADAIN AFTER SETTING TO LG above!""" + else: + self.y0_style_neg = torch.zeros_like(x, dtype=self.dtype, device=self.device) + + if self.UNSAMPLE and not self.SAMPLE: #sigma_next > sigma: # TODO: VERIFY APPROACH FOR INVERSION + if guide_inversion_y0 is not None: + self.y0 = guide_inversion_y0.clone() + else: + self.y0 = noise_sampler(sigma=self.sigma_max, sigma_next=self.sigma_min).to(dtype=self.dtype, device=self.device) + self.y0 = normalize_zscore(self.y0, channelwise=True, inplace=True) + self.y0 *= self.sigma_max + + if guide_inversion_y0_inv is not None: + self.y0_inv = guide_inversion_y0_inv.clone() + else: + self.y0_inv = noise_sampler(sigma=self.sigma_max, sigma_next=self.sigma_min).to(dtype=self.dtype, device=self.device) + self.y0_inv = normalize_zscore(self.y0_inv, channelwise=True, inplace=True) + self.y0_inv*= self.sigma_max + + + if self.frame_weights_mgr is not None and x.ndim == 5: + num_frames = x.shape[2] + self.frame_weights = self.frame_weights_mgr.get_frame_weights_by_name('frame_weights', num_frames) + self.frame_weights_inv = self.frame_weights_mgr.get_frame_weights_by_name('frame_weights_inv', num_frames) + + x, self.y0, self.y0_inv = self.normalize_inputs(x, self.y0, self.y0_inv) # ??? + + return x + + def prepare_weighted_masks(self, step:int, lgw_type="default") -> Tuple[Tensor, Tensor]: + if lgw_type == "sync": + lgw_ = self.lgw_sync [step] + lgw_inv_ = self.lgw_sync_inv[step] + mask = torch.ones_like (self.y0) if self.mask_sync is None else self.mask_sync + mask_inv = torch.zeros_like(self.y0) if self.mask_sync is None else 1-self.mask_sync + elif lgw_type == "drift_x": + lgw_ = self.lgw_drift_x [step] + lgw_inv_ = self.lgw_drift_x_inv[step] + mask = torch.ones_like (self.y0) if self.mask_drift_x is None else self.mask_drift_x + mask_inv = torch.zeros_like(self.y0) if self.mask_drift_x is None else 1-self.mask_drift_x + elif lgw_type == "drift_y": + lgw_ = self.lgw_drift_y [step] + lgw_inv_ = self.lgw_drift_y_inv[step] + mask = torch.ones_like (self.y0) if self.mask_drift_y is None else self.mask_drift_y + mask_inv = torch.zeros_like(self.y0) if self.mask_drift_y is None else 1-self.mask_drift_y + elif lgw_type == "lure_x": + lgw_ = self.lgw_lure_x [step] + lgw_inv_ = self.lgw_lure_x_inv[step] + mask = torch.ones_like (self.y0) if self.mask_lure_x is None else self.mask_lure_x + mask_inv = torch.zeros_like(self.y0) if self.mask_lure_x is None else 1-self.mask_lure_x + elif lgw_type == "lure_y": + lgw_ = self.lgw_lure_y [step] + lgw_inv_ = self.lgw_lure_y_inv[step] + mask = torch.ones_like (self.y0) if self.mask_lure_y is None else self.mask_lure_y + mask_inv = torch.zeros_like(self.y0) if self.mask_lure_y is None else 1-self.mask_lure_y + else: + lgw_ = self.lgw [step] + lgw_inv_ = self.lgw_inv[step] + mask = torch.ones_like (self.y0) if self.mask is None else self.mask + mask_inv = torch.zeros_like(self.y0) if self.mask_inv is None else self.mask_inv + + if self.LGW_MASK_RESCALE_MIN: + lgw_mask = mask * (1-lgw_) + lgw_ + lgw_mask_inv = (1-mask) * (1-lgw_inv_) + lgw_inv_ + else: + if self.HAS_LATENT_GUIDE: + lgw_mask = mask * lgw_ + else: + lgw_mask = torch.zeros_like(mask) + + if self.HAS_LATENT_GUIDE_INV: + if mask_inv is not None: + lgw_mask_inv = torch.minimum(mask_inv, (1-mask) * lgw_inv_) + #lgw_mask_inv = torch.minimum(1-mask_inv, (1-mask) * lgw_inv_) + else: + lgw_mask_inv = (1-mask) * lgw_inv_ + else: + lgw_mask_inv = torch.zeros_like(mask) + + return lgw_mask, lgw_mask_inv + + + def get_masks_for_step(self, step:int, lgw_type="default") -> Tuple[Tensor, Tensor]: + lgw_mask, lgw_mask_inv = self.prepare_weighted_masks(step, lgw_type=lgw_type) + normalize_frame_weights_per_step = self.EO("normalize_frame_weights_per_step") + normalize_frame_weights_per_step_inv = self.EO("normalize_frame_weights_per_step_inv") + + if self.VIDEO and self.frame_weights_mgr and lgw_mask.ndim >= 5: + num_frames = lgw_mask.shape[2] + if self.HAS_LATENT_GUIDE: + frame_weights = self.frame_weights_mgr.get_frame_weights_by_name('frame_weights', num_frames, step) + apply_frame_weights(lgw_mask, frame_weights, normalize_frame_weights_per_step) + if self.HAS_LATENT_GUIDE_INV: + frame_weights_inv = self.frame_weights_mgr.get_frame_weights_by_name('frame_weights_inv', num_frames, step) + apply_frame_weights(lgw_mask_inv, frame_weights_inv, normalize_frame_weights_per_step_inv) + + return lgw_mask.to(self.device), lgw_mask_inv.to(self.device) + + + + def get_cossim_adjusted_lgw_masks(self, data:Tensor, step:int) -> Tuple[Tensor, Tensor, Tensor, Tensor]: + # PACK-FIRST: Require flat [1,1,N] input - callers must pack first + data_for_cossim = data + + if self.HAS_LATENT_GUIDE: + y0 = self.y0.clone() + else: + y0 = torch.zeros_like(data_for_cossim) + + if self.HAS_LATENT_GUIDE_INV: + y0_inv = self.y0_inv.clone() + else: + y0_inv = torch.zeros_like(data_for_cossim) + + if y0.shape[0] > 1: # this is for changing the guide on a per-step basis + y0 = y0[min(step, y0.shape[0]-1)].unsqueeze(0) + + lgw_mask, lgw_mask_inv = self.get_masks_for_step(step) + + y0_cossim, y0_cossim_inv = 1.0, 1.0 + if self.HAS_LATENT_GUIDE: + y0_cossim = get_pearson_similarity(data_for_cossim, y0, mask=lgw_mask) + if self.HAS_LATENT_GUIDE_INV: + y0_cossim_inv = get_pearson_similarity(data_for_cossim, y0_inv, mask=lgw_mask_inv) + + #if y0_cossim < self.guide_cossim_cutoff_ or y0_cossim_inv < self.guide_bkg_cossim_cutoff_: + if y0_cossim >= self.guide_cossim_cutoff_: + lgw_mask *= 0 + if y0_cossim_inv >= self.guide_bkg_cossim_cutoff_: + lgw_mask_inv *= 0 + + return y0, y0_inv, lgw_mask, lgw_mask_inv + + + def get_self_refine_epsilon_mask(self, current: Tensor, previous: Tensor, step_sched: int) -> Tensor: + """ + Compute per-pixel certainty mask for self_refine_epsilon mode. + Returns mask where 1 = certain (guide), 0 = uncertain (no guide). + + When invert_mask=False (default): guide CERTAIN (low-diff/stable) regions + When invert_mask=True: guide UNCERTAIN (high-diff/changing) regions + """ + threshold = self.self_refine_threshold + metric = self.self_refine_metric + + if metric == "l2": + # Normalized L2 (Euclidean distance per pixel, normalized by channel count) + diff = current - previous + if diff.ndim >= 4 and diff.shape[1] > 1: + diff = torch.sqrt(torch.sum(diff ** 2, dim=1, keepdim=True)) / diff.shape[1] + else: + diff = torch.abs(diff) + else: + # L1 (absolute difference, averaged over channels) + diff = torch.abs(current - previous) + if diff.ndim >= 4 and diff.shape[1] > 1: + diff = diff.mean(dim=1, keepdim=True) + + # Certain = low diff (BELOW threshold) + # Uncertain = high diff (ABOVE threshold) + certain_mask = (diff < threshold).float() + + # Apply invert_mask: if True, guide uncertain regions instead of certain + if self.invert_mask: + certain_mask = 1.0 - certain_mask + + # Apply spatial mask from ClownGuides if provided (intersection) + if self.mask is not None: + spatial_mask = self.mask + if spatial_mask.shape != certain_mask.shape: + if spatial_mask.ndim == certain_mask.ndim: + if spatial_mask.shape[1] == 1 and certain_mask.shape[1] > 1: + spatial_mask = spatial_mask.expand_as(certain_mask) + elif certain_mask.shape[1] == 1 and spatial_mask.shape[1] > 1: + spatial_mask = spatial_mask.mean(dim=1, keepdim=True) + certain_mask = certain_mask * spatial_mask + + # Apply guide weight schedule + lgw = self.lgw[step_sched] if step_sched < len(self.lgw) else 0.0 + + if self.EO("self_refine_debug"): + coverage = certain_mask.mean().item() + mode = "UNCERTAIN (inverted)" if self.invert_mask else "CERTAIN" + spatial_info = " (with spatial mask)" if self.mask is not None else "" + RESplain(f"self_refine_epsilon step {step_sched}: guiding {mode} regions{spatial_info}, coverage={coverage:.2%}, metric={metric}, threshold={threshold}, lgw={lgw:.4f}") + + # Store raw mask for visualization (before lgw scaling) + self._debug_certainty_mask = certain_mask.clone() + + return certain_mask * lgw + + + + + + + + + + + @torch.no_grad + def process_pseudoimplicit_guides_substep(self, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + eps_prev_ : Tensor, + data_ : Tensor, + denoised_prev : Tensor, + row : int, + step : int, + step_sched : int, + sigmas : Tensor, + NS , + RK , + pseudoimplicit_row_weights : Tensor, + pseudoimplicit_step_weights : Tensor, + full_iter : int, + BONGMATH : bool, + ): + + # Check if this is a pseudoimplicit mode (including self_refine_pseudoimplicit variants) + is_pseudoimplicit_mode = "pseudoimplicit" in self.guide_mode or self.guide_mode.startswith("self_refine_pseudoimplicit") + if not is_pseudoimplicit_mode or (self.lgw[step_sched] == 0 and self.lgw_inv[step_sched] == 0): + return x_0, x_, eps_, None, None + + if x_0.ndim == 3: # packed NestedTensor + BLOCKED_PSEUDOIMPLICIT_MODES = {"pseudoimplicit_cw", "pseudoimplicit_projection_cw", + "fully_pseudoimplicit_cw", "fully_pseudoimplicit_projection_cw"} + if self.guide_mode in BLOCKED_PSEUDOIMPLICIT_MODES: + raise NotImplementedError(f"Mode '{self.guide_mode}' requires channel structure, incompatible with packed latents") + + sigma = sigmas[step] + + # Handle self_refine_pseudoimplicit modes + if self.guide_mode.startswith("self_refine_pseudoimplicit"): + # Per-iteration mode: track changes across implicit iterations + per_iteration_mode = not self.EO("self_refine_by_step") + + # Skip conditions + if per_iteration_mode: + if step == 0 and full_iter == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}, iter {full_iter}, row {row}: SKIPPED - no valid reference") + return x_0, x_, eps_, None, None + else: + if step == 0 or denoised_prev.abs().max() == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}, row {row}: SKIPPED - no valid denoised_prev") + return x_0, x_, eps_, None, None + + is_new_step = (step != self.self_refine_epsilon_last_step) + is_new_iter = (full_iter != self._self_refine_last_iter) if per_iteration_mode else False + + # Reference management + if per_iteration_mode: + if is_new_step: + # New step: reset everything + self.self_refine_epsilon_ref = denoised_prev.clone() + self.self_refine_epsilon_last_step = step + self._self_refine_last_iter = full_iter + self._self_refine_certain_mask_accum = None + self._self_refine_iter_prediction = None + self._self_refine_converged = False + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}: NEW STEP - initialized reference from denoised_prev") + + elif is_new_iter: + # New iteration: update reference to previous iteration's prediction + if self._self_refine_iter_prediction is not None: + self.self_refine_epsilon_ref = self._self_refine_iter_prediction.clone() + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}, iter {full_iter}: updated reference from iter {full_iter-1}") + self._self_refine_last_iter = full_iter + if self.EO("self_refine_dont_accumulate_certainty"): + self._self_refine_certain_mask_accum = None + + if row == 0: + self._self_refine_iter_prediction = data_[row].clone() + else: + # Non-iterative mode: update reference at each step + if is_new_step: + self.self_refine_epsilon_ref = denoised_prev.clone() + self.self_refine_epsilon_last_step = step + self._self_refine_converged = False + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}: initialized reference from denoised_prev") + + # Use reference as guide target + y0 = self.self_refine_epsilon_ref + + # Compute certainty mask + lgw_mask = self.get_self_refine_epsilon_mask(data_[row], y0, step_sched) + + # Accumulate certainty across iterations + if per_iteration_mode and not self.EO("self_refine_dont_accumulate_certainty"): + if self._self_refine_certain_mask_accum is not None: + binary_mask = (lgw_mask > 0).float() + binary_accum = (self._self_refine_certain_mask_accum > 0).float() + combined = torch.maximum(binary_mask, binary_accum) + lgw = self.lgw[step_sched] if step_sched < len(self.lgw) else 0.0 + lgw_mask = combined * lgw + self._self_refine_certain_mask_accum = lgw_mask.clone() + + if lgw_mask.max() == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_pseudoimplicit step {step}, row {row}: SKIPPED - no certain regions") + return x_0, x_, eps_, None, None + + # Check coverage against cutoff + coverage = (lgw_mask > 0).float().mean().item() + if coverage >= self.self_refine_cutoff: + self._self_refine_converged = True + if self.EO("self_refine_debug"): + iter_info = f", iter {full_iter}" if per_iteration_mode else "" + RESplain(f"self_refine_pseudoimplicit step {step}{iter_info}, row {row}: CONVERGED - coverage={coverage:.2%} >= cutoff={self.self_refine_cutoff:.2%}") + + # Compute guide epsilon + eps_substep_guide = RK.get_guide_epsilon(x_0, x_[row], y0, sigma, NS.s_[row], NS.sigma_down, None) + + # Pseudoimplicit sigma adjustment + maxmin_ratio = (NS.sub_sigma - RK.sigma_min) / NS.sub_sigma + sub_sigma_2 = NS.sub_sigma - maxmin_ratio * (NS.sub_sigma * pseudoimplicit_row_weights[row] * pseudoimplicit_step_weights[full_iter] * self.lgw[step_sched]) + + eps_tmp_ = eps_.clone() + eps_row = eps_[row] + + # Blend with certainty mask + if "_projection" in self.guide_mode: + # Projection variant: preserve magnitude, steer direction + eps_row_lerp = eps_row + lgw_mask * (eps_substep_guide - eps_row) + eps_collinear = get_collinear(eps_row, eps_row_lerp) + eps_ortho = get_orthogonal(eps_row_lerp, eps_row) + eps_sum = eps_collinear + eps_ortho + eps_row = eps_row + lgw_mask * (eps_sum - eps_row) + else: + # Standard lerp blending + eps_row = eps_row + lgw_mask * (eps_substep_guide - eps_row) + eps_[row] = eps_row + + # Compute pseudoimplicit x + x_row_pseudoimplicit = x_[row] + RK.h_fn(sub_sigma_2, NS.sub_sigma) * eps_[row] + sub_sigma_pseudoimplicit = sub_sigma_2 + + eps_ = eps_tmp_ + + if self.EO("self_refine_debug"): + coverage = (lgw_mask > 0).float().mean().item() + iter_info = f", iter {full_iter}" if per_iteration_mode else "" + RESplain(f"self_refine_pseudoimplicit step {step}{iter_info}, row {row}: APPLIED - certain_coverage={coverage:.2%}") + + # Apply bongmath if enabled + if RK.IMPLICIT and BONGMATH and step < sigmas.shape[0]-1 and not self.EO("disable_pseudobongmath"): + x_[row] = NS.sigma_from_to(x_0, x_row_pseudoimplicit, sigma, sub_sigma_pseudoimplicit, NS.s_[row]) + x_0, x_, eps_ = RK.bong_iter(x_0, x_, eps_, eps_prev_, data_, sigma, NS.s_, row, RK.row_offset, NS.h, step, step_sched) + + return x_0, x_, eps_, x_row_pseudoimplicit, sub_sigma_pseudoimplicit + + if self.s_lying_ is not None: + if row >= len(self.s_lying_): + return x_0, x_, eps_, None, None + + if self.guide_mode.startswith("fully_"): + data_cossim_test = denoised_prev + else: + data_cossim_test = data_[row] + + y0, y0_inv, lgw_mask, lgw_mask_inv = self.get_cossim_adjusted_lgw_masks(data_cossim_test, step_sched) + + if not (lgw_mask.any() != 0 or lgw_mask_inv.any() != 0): # cossim score too similar! deactivate guide for this step + return x_0, x_, eps_, None, None + + + if "fully_pseudoimplicit" in self.guide_mode: + if self.x_lying_ is None: + return x_0, x_, eps_, None, None + else: + x_row_pseudoimplicit = self.x_lying_[row] + sub_sigma_pseudoimplicit = self.s_lying_[row] + + + + if RK.IMPLICIT: + x_ = RK.update_substep(x_0, + x_, + eps_, + eps_prev_, + row, + RK.row_offset, + NS.h_new, + NS.h_new_orig, + ) + + x_[row] = NS.rebound_overshoot_substep(x_0, x_[row]) + + if row > 0: + x_[row] = NS.swap_noise_substep(x_0, x_[row]) + if BONGMATH and step < sigmas.shape[0]-1 and not self.EO("disable_pseudoimplicit_bongmath"): + x_0, x_, eps_ = RK.bong_iter(x_0, + x_, + eps_, + eps_prev_, + data_, + sigma, + NS.s_, + row, + RK.row_offset, + NS.h, + step, + step_sched, + ) + else: + eps_[row] = RK.get_epsilon(x_0, x_[row], denoised_prev, sigma, NS.s_[row]) + + if self.EO("pseudoimplicit_denoised_prev"): + eps_[row] = RK.get_epsilon(x_0, x_[row], denoised_prev, sigma, NS.s_[row]) + + eps_substep_guide = torch.zeros_like(x_0) + eps_substep_guide_inv = torch.zeros_like(x_0) + + if self.HAS_LATENT_GUIDE: + eps_substep_guide = RK.get_guide_epsilon(x_0, x_[row], y0, sigma, NS.s_[row], NS.sigma_down, None) + if self.HAS_LATENT_GUIDE_INV: + eps_substep_guide_inv = RK.get_guide_epsilon(x_0, x_[row], y0_inv, sigma, NS.s_[row], NS.sigma_down, None) + + if self.guide_mode in {"pseudoimplicit", "pseudoimplicit_cw", "pseudoimplicit_projection", "pseudoimplicit_projection_cw"}: + maxmin_ratio = (NS.sub_sigma - RK.sigma_min) / NS.sub_sigma + + if self.EO("guide_pseudoimplicit_power_substep_flip_maxmin_scaling"): + maxmin_ratio *= (RK.rows-row) / RK.rows + elif self.EO("guide_pseudoimplicit_power_substep_maxmin_scaling"): + maxmin_ratio *= row / RK.rows + + sub_sigma_2 = NS.sub_sigma - maxmin_ratio * (NS.sub_sigma * pseudoimplicit_row_weights[row] * pseudoimplicit_step_weights[full_iter] * self.lgw[step_sched]) + + eps_tmp_ = eps_.clone() + + eps_ = self.process_channelwise(x_0, + eps_, + data_, + row, + eps_substep_guide, + eps_substep_guide_inv, + y0, + y0_inv, + lgw_mask, + lgw_mask_inv, + use_projection = self.guide_mode in {"pseudoimplicit_projection", "pseudoimplicit_projection_cw"}, + channelwise = self.guide_mode in {"pseudoimplicit_cw", "pseudoimplicit_projection_cw"}, + ) + + if self.EO("debug_pseudoimplicit"): + RESplain( + f"Step {step}, Row {row}: eps_[row] post-blend mean/std=" + f"{eps_[row].mean().item():.6f}/{eps_[row].std().item():.6f}" + ) + + x_row_tmp = x_[row] + RK.h_fn(sub_sigma_2, NS.sub_sigma) * eps_[row] + + eps_ = eps_tmp_ + x_row_pseudoimplicit = x_row_tmp + sub_sigma_pseudoimplicit = sub_sigma_2 + + + if RK.IMPLICIT and BONGMATH and step < sigmas.shape[0]-1 and not self.EO("disable_pseudobongmath"): + x_[row] = NS.sigma_from_to(x_0, x_row_pseudoimplicit, sigma, sub_sigma_pseudoimplicit, NS.s_[row]) + + x_0, x_, eps_ = RK.bong_iter(x_0, + x_, + eps_, + eps_prev_, + data_, + sigma, + NS.s_, + row, + RK.row_offset, + NS.h, + step, + step_sched, + ) + + return x_0, x_, eps_, x_row_pseudoimplicit, sub_sigma_pseudoimplicit + + + + @torch.no_grad + def prepare_fully_pseudoimplicit_guides_substep(self, + x_0, + x_, + eps_, + eps_prev_, + data_, + denoised_prev, + row, + step, + step_sched, + sigmas, + eta_substep, + overshoot_substep, + s_noise_substep, + NS, + RK, + pseudoimplicit_row_weights, + pseudoimplicit_step_weights, + full_iter, + BONGMATH, + ): + + if "fully_pseudoimplicit" not in self.guide_mode or (self.lgw[step_sched] == 0 and self.lgw_inv[step_sched] == 0): + if self.EO("debug_pseudoimplicit"): + RESplain(f"prepare_fully: SKIPPED - mode={self.guide_mode}, lgw={self.lgw[step_sched]:.4f}, lgw_inv={self.lgw_inv[step_sched]:.4f}") + return x_0, x_, eps_ + + # PACK-FIRST EXPERIMENT: Block channelwise fully_pseudoimplicit modes + if x_0.ndim == 3: # packed NestedTensor + BLOCKED_FULLY_MODES = {"fully_pseudoimplicit_cw", "fully_pseudoimplicit_projection_cw"} + if self.guide_mode in BLOCKED_FULLY_MODES: + raise NotImplementedError(f"Mode '{self.guide_mode}' requires channel structure, incompatible with packed latents") + + sigma = sigmas[step] + + y0, y0_inv, lgw_mask, lgw_mask_inv = self.get_cossim_adjusted_lgw_masks(denoised_prev, step_sched) + + if not (lgw_mask.any() != 0 or lgw_mask_inv.any() != 0): # cossim score too similar! deactivate guide for this step + return x_0, x_, eps_ + + + # PREPARE FULLY PSEUDOIMPLICIT GUIDES + if self.guide_mode in {"fully_pseudoimplicit", "fully_pseudoimplicit_cw", "fully_pseudoimplicit_projection", "fully_pseudoimplicit_projection_cw"} and (self.lgw[step_sched] > 0 or self.lgw_inv[step_sched] > 0): + x_lying_ = x_.clone() + eps_lying_ = eps_.clone() + s_lying_ = [] + + for r in range(RK.rows): + + NS.set_sde_substep(r, RK.multistep_stages, eta_substep, overshoot_substep, s_noise_substep) + + maxmin_ratio = (NS.sub_sigma - RK.sigma_min) / NS.sub_sigma + fully_sub_sigma_2 = NS.sub_sigma - maxmin_ratio * (NS.sub_sigma * pseudoimplicit_row_weights[r] * pseudoimplicit_step_weights[full_iter] * self.lgw[step_sched]) + + s_lying_.append(fully_sub_sigma_2) + + if RK.IMPLICIT: + x_ = RK.update_substep(x_0, + x_, + eps_, + eps_prev_, + r, + RK.row_offset, + NS.h_new, + NS.h_new_orig, + ) + + x_[r] = NS.rebound_overshoot_substep(x_0, x_[r]) + + if r > 0: + x_[r] = NS.swap_noise_substep(x_0, x_[r]) + if BONGMATH and step < sigmas.shape[0]-1 and not self.EO("disable_fully_pseudoimplicit_bongmath"): + x_0, x_, eps_ = RK.bong_iter(x_0, + x_, + eps_, + eps_prev_, + data_, + sigma, + NS.s_, + r, + RK.row_offset, + NS.h, + step, + step_sched, + ) + + if self.EO("fully_pseudoimplicit_denoised_prev"): + eps_[r] = RK.get_epsilon(x_0, x_[r], denoised_prev, sigma, NS.s_[r]) + + eps_substep_guide = torch.zeros_like(x_0) + eps_substep_guide_inv = torch.zeros_like(x_0) + + if self.HAS_LATENT_GUIDE: + eps_substep_guide = RK.get_guide_epsilon(x_0, x_[r], y0, sigma, NS.s_[r], NS.sigma_down, None) + if self.HAS_LATENT_GUIDE_INV: + eps_substep_guide_inv = RK.get_guide_epsilon(x_0, x_[r], y0_inv, sigma, NS.s_[r], NS.sigma_down, None) + + eps_ = self.process_channelwise(x_0, + eps_, + data_, + r, + eps_substep_guide, + eps_substep_guide_inv, + y0, + y0_inv, + lgw_mask, + lgw_mask_inv, + use_projection = self.guide_mode in {"fully_pseudoimplicit_projection", "fully_pseudoimplicit_projection_cw"}, + channelwise = self.guide_mode in {"fully_pseudoimplicit_cw", "fully_pseudoimplicit_projection_cw"}, + ) + + x_lying_[r] = x_[r] + RK.h_fn(fully_sub_sigma_2, NS.sub_sigma) * eps_[r] + data_lying = x_[r] + RK.h_fn(0, NS.s_[r]) * eps_[r] + + eps_lying_[r] = RK.get_epsilon(x_0, x_[r], data_lying, sigma, NS.s_[r]) + + if not self.EO("pseudoimplicit_disable_eps_lying"): + eps_ = eps_lying_ + + if not self.EO("pseudoimplicit_disable_newton_iter"): + x_, eps_ = RK.newton_iter(x_0, + x_, + eps_, + eps_prev_, + data_, + NS.s_, + 0, + NS.h, + sigmas, + step, + "lying", + False, # SYNC_GUIDE_ACTIVE + ) + + self.x_lying_ = x_lying_ + self.s_lying_ = s_lying_ + + return x_0, x_, eps_ + + + + @torch.no_grad + def process_guides_data_substep(self, + x_row : Tensor, + data_row : Tensor, + step : int, + sigma_row : Tensor, + ): + if not self.HAS_LATENT_GUIDE and not self.HAS_LATENT_GUIDE_INV: + return x_row + + y0, y0_inv, lgw_mask, lgw_mask_inv = self.get_cossim_adjusted_lgw_masks(data_row, step) + + if not (lgw_mask.any() != 0 or lgw_mask_inv.any() != 0): + return x_row + + if self.guide_mode in {"data", "data_projection", "lure", "lure_projection"}: + x_row = self.get_data_substep(x_row, data_row, y0, y0_inv, lgw_mask, lgw_mask_inv, step, sigma_row) + + return x_row + + + + + @torch.no_grad + def get_data_substep(self, + x_row : Tensor, + data_row : Tensor, + y0 : Tensor, + y0_inv : Tensor, + lgw_mask : Tensor, + lgw_mask_inv : Tensor, + step : int, + sigma_row : Tensor, + frame_target : float = 1.0, + ): + + if not self.HAS_LATENT_GUIDE and not self.HAS_LATENT_GUIDE_INV: + return x_row + + if self.guide_mode in {"data", "data_projection", "lure", "lure_projection"}: + data_targets = self.EO("data_targets", [1.0]) + step_target = step if len(data_targets) > step else len(data_targets)-1 + + cossim_target = frame_target * data_targets[step_target] + + if self.HAS_LATENT_GUIDE: + if self.guide_mode.endswith("projection"): + d_collinear_d_lerp = get_collinear(data_row, y0) + d_lerp_ortho_d = get_orthogonal(y0, data_row) + y0 = d_collinear_d_lerp + d_lerp_ortho_d + + if cossim_target == 1.0: + d_slerped = y0 + elif cossim_target == 0.0: + d_slerped = data_row + else: + y0_pearsim = get_pearson_similarity(data_row, y0, mask=self.mask) + slerp_weight = get_slerp_weight_for_cossim(y0_pearsim.item(), cossim_target) + d_slerped = slerp_tensor(slerp_weight, data_row, y0) # lgw_mask * slerp_weight same as using mask below + + """if self.guide_mode == "data_projection": + d_collinear_d_lerp = get_collinear(data_row, d_slerped) + d_lerp_ortho_d = get_orthogonal(d_slerped, data_row) + d_slerped = d_collinear_d_lerp + d_lerp_ortho_d""" + + if self.VE_MODEL: + x_row = x_row + lgw_mask * (d_slerped - data_row) + else: + x_row = x_row + lgw_mask * (self.sigma_max - sigma_row) * (d_slerped - data_row) + + + if self.HAS_LATENT_GUIDE_INV: + if self.guide_mode.endswith("projection"): + d_collinear_d_lerp = get_collinear(data_row, y0_inv) + d_lerp_ortho_d = get_orthogonal(y0_inv, data_row) + y0_inv = d_collinear_d_lerp + d_lerp_ortho_d + + if cossim_target == 1.0: + d_slerped_inv = y0_inv + elif cossim_target == 0.0: + d_slerped_inv = data_row + else: + y0_pearsim = get_pearson_similarity(data_row, y0_inv, mask=self.mask_inv) + slerp_weight = get_slerp_weight_for_cossim(y0_pearsim.item(), cossim_target) + d_slerped_inv = slerp_tensor(slerp_weight, data_row, y0_inv) + + """if self.guide_mode == "data_projection": + d_collinear_d_lerp = get_collinear(data_row, d_slerped_inv) + d_lerp_ortho_d = get_orthogonal(d_slerped_inv, data_row) + d_slerped_inv = d_collinear_d_lerp + d_lerp_ortho_d""" + + if self.VE_MODEL: + x_row = x_row + lgw_mask_inv * (d_slerped_inv - data_row) + else: + x_row = x_row + lgw_mask_inv * (self.sigma_max - sigma_row) * (d_slerped_inv - data_row) + + + return x_row + + @torch.no_grad + def swap_data(self, + x : Tensor, + data : Tensor, + y : Tensor, + sigma : Tensor, + mask : Optional[Tensor] = None, + ): + mask = 1.0 if mask is None else mask + if self.VE_MODEL: + return x + mask * (y - data) + else: + return x + mask * (self.sigma_max - sigma) * (y - data) + + @torch.no_grad + def process_guides_eps_substep(self, + x_0 : Tensor, + x_row : Tensor, + data_row : Tensor, + eps_row : Tensor, + step : int, + sigma : Tensor, + sigma_down : Tensor, + sigma_row : Tensor, + RK=None, + ): + if not self.HAS_LATENT_GUIDE and not self.HAS_LATENT_GUIDE_INV: + return eps_row + + y0, y0_inv, lgw_mask, lgw_mask_inv = self.get_cossim_adjusted_lgw_masks(data_row, step) + + if not (lgw_mask.any() != 0 or lgw_mask_inv.any() != 0): + return eps_row + + eps_y0 = torch.zeros_like(x_0) + eps_y0_inv = torch.zeros_like(x_0) + + if self.HAS_LATENT_GUIDE: + eps_y0 = RK.get_guide_epsilon(x_0, x_row, y0, sigma, sigma_row, sigma_down, None) + + if self.HAS_LATENT_GUIDE_INV: + eps_y0_inv = RK.get_guide_epsilon(x_0, x_row, y0_inv, sigma, sigma_row, sigma_down, None) + + if self.guide_mode in {"epsilon", "epsilon_projection"}: + eps_row = self.get_eps_substep(eps_row, eps_y0, eps_y0_inv, lgw_mask, lgw_mask_inv, step, sigma_row) + + return eps_row + + + + @torch.no_grad + def get_eps_substep(self, + eps_row : Tensor, + eps_y0 : Tensor, + eps_y0_inv : Tensor, + lgw_mask : Tensor, + lgw_mask_inv : Tensor, + step : int, + sigma_row : Tensor, + frame_target : float = 1.0, + ): + + if not self.HAS_LATENT_GUIDE and not self.HAS_LATENT_GUIDE_INV: + return eps_row + + if self.guide_mode in {"epsilon", "epsilon_projection"}: + eps_targets = self.EO("eps_targets", [1.0]) + step_target = step if len(eps_targets) > step else len(eps_targets)-1 + + cossim_target = frame_target * eps_targets[step_target] + + if self.HAS_LATENT_GUIDE: + if self.guide_mode == "epsilon_projection": + d_collinear_d_lerp = get_collinear(eps_row, eps_y0) + d_lerp_ortho_d = get_orthogonal(eps_y0, eps_row) + eps_y0 = d_collinear_d_lerp + d_lerp_ortho_d + + if cossim_target == 1.0: + d_slerped = eps_y0 + elif cossim_target == 0.0: + d_slerped = eps_row + else: + y0_pearsim = get_pearson_similarity(eps_row, eps_y0, mask=self.mask) + slerp_weight = get_slerp_weight_for_cossim(y0_pearsim.item(), cossim_target) + d_slerped = slerp_tensor(slerp_weight, eps_row, eps_y0) # lgw_mask * slerp_weight same as using mask below + + """if self.guide_mode == "data_projection": + d_collinear_d_lerp = get_collinear(data_row, d_slerped) + d_lerp_ortho_d = get_orthogonal(d_slerped, data_row) + d_slerped = d_collinear_d_lerp + d_lerp_ortho_d""" + + eps_row = eps_row + lgw_mask * (d_slerped - eps_row) + + + if self.HAS_LATENT_GUIDE_INV: + if self.guide_mode == "epsilon_projection": + d_collinear_d_lerp = get_collinear(eps_row, eps_y0_inv) + d_lerp_ortho_d = get_orthogonal(eps_y0_inv, eps_row) + eps_y0_inv = d_collinear_d_lerp + d_lerp_ortho_d + + if cossim_target == 1.0: + d_slerped_inv = eps_y0_inv + elif cossim_target == 0.0: + d_slerped_inv = eps_row + else: + y0_pearsim = get_pearson_similarity(eps_row, eps_y0_inv, mask=self.mask_inv) + slerp_weight = get_slerp_weight_for_cossim(y0_pearsim.item(), cossim_target) + d_slerped_inv = slerp_tensor(slerp_weight, eps_row, eps_y0_inv) + + """if self.guide_mode == "data_projection": + d_collinear_d_lerp = get_collinear(data_row, d_slerped_inv) + d_lerp_ortho_d = get_orthogonal(d_slerped_inv, data_row) + d_slerped_inv = d_collinear_d_lerp + d_lerp_ortho_d""" + + eps_row = eps_row + lgw_mask_inv * (d_slerped_inv - eps_row) + + return eps_row + + + + + + @torch.no_grad + def process_guides_substep(self, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + data_ : Tensor, + denoised_prev : Tensor, + row : int, + step : int, + step_sched : int, + sigma : Tensor, + sigma_next : Tensor, + sigma_down : Tensor, + s_ : Tensor, + epsilon_scale : float, + RK, + full_iter : int = 0, + ): + + if not self.HAS_LATENT_GUIDE and not self.HAS_LATENT_GUIDE_INV: + return eps_, x_ + + is_flat = x_0.ndim == 3 # packed NestedTensor: [1,1,N] + + if is_flat: + BLOCKED_MODES = {"epsilon_cw", "epsilon_projection_cw"} + if self.guide_mode in BLOCKED_MODES: + raise NotImplementedError(f"Mode '{self.guide_mode}' requires spatial structure, incompatible with packed latents") + + if self.frame_weights_mgr is not None: + raise NotImplementedError("Frame weights require temporal structure, incompatible with pack-first experiment") + + # Handle self_refine_epsilon mode - uses denoised_prev as guide target + if self.SELF_REFINE_EPSILON_MODE: + lgw = self.lgw[step_sched] if step_sched < len(self.lgw) else 0.0 + if lgw == 0: + return eps_, x_ + + data_row = data_[row] + eps_row = eps_[row] + x_row = x_[row] + sigma_row = s_[row] + + # Per-iteration mode: track changes across implicit iterations + per_iteration_mode = not self.EO("self_refine_by_step") + + # Skip conditions + if per_iteration_mode: + # In per-iteration mode, skip step 0 iter 0 only + if step == 0 and full_iter == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}, iter {full_iter}, row {row}: SKIPPED - no valid reference") + return eps_, x_ + else: + # Non-iterative mode: skip entire step 0 + if step == 0 or denoised_prev.abs().max() == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}, row {row}: SKIPPED - no valid denoised_prev") + return eps_, x_ + + # Track state changes + is_new_step = (step != self.self_refine_epsilon_last_step) + is_new_iter = (full_iter != self._self_refine_last_iter) if per_iteration_mode else False + is_new_row = (row != self.self_refine_epsilon_last_row) + + # Track calls per (step, row) - function is called twice per row (for eps_ and eps_prev_) + if is_new_step or is_new_iter or is_new_row: + # First call for this (step, iter, row) + self.self_refine_epsilon_call_count = 1 + self.self_refine_epsilon_last_row = row + else: + # Second call for same (step, iter, row) - this is eps_prev_ + self.self_refine_epsilon_call_count += 1 + + if self.EO("self_refine_dont_guide_eps_prev"): + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}, iter {full_iter}, row {row}: SKIPPED eps_prev_") + return eps_, x_ + + # Reference management + if per_iteration_mode: + # Per-iteration mode: update reference based on iteration changes + if is_new_step: + # New step: reset everything, use denoised_prev as initial reference + self.self_refine_epsilon_ref = denoised_prev.clone() + self.self_refine_epsilon_last_step = step + self._self_refine_last_iter = full_iter + self._self_refine_certain_mask_accum = None + self._self_refine_iter_prediction = None + self._self_refine_converged = False + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}: NEW STEP - initialized reference from denoised_prev") + + elif is_new_iter: + # New iteration within same step: update reference to previous iteration's prediction + if self._self_refine_iter_prediction is not None: + self.self_refine_epsilon_ref = self._self_refine_iter_prediction.clone() + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}, iter {full_iter}: updated reference from iter {full_iter-1}") + self._self_refine_last_iter = full_iter + # Reset mask accumulator for new iteration if not using accumulation + if self.EO("self_refine_dont_accumulate_certainty"): + self._self_refine_certain_mask_accum = None + + # Capture current prediction for next iteration's reference (at row 0, first call only) + if row == 0 and self.self_refine_epsilon_call_count == 1: + self._self_refine_iter_prediction = data_row.clone() + else: + # Non-iterative mode: update reference only on new steps + if is_new_step: + self.self_refine_epsilon_ref = denoised_prev.clone() + self.self_refine_epsilon_last_step = step + self._self_refine_converged = False + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}: initialized reference from denoised_prev") + + # Compare current prediction against reference + y0 = self.self_refine_epsilon_ref + + # Compute certainty mask (guide certain regions, let uncertain evolve) + lgw_mask = self.get_self_refine_epsilon_mask(data_row, y0, step_sched) + + # Accumulate certainty across iterations + if per_iteration_mode and not self.EO("self_refine_dont_accumulate_certainty"): + if self._self_refine_certain_mask_accum is not None: + # Union with previous certain regions + binary_mask = (lgw_mask > 0).float() + binary_accum = (self._self_refine_certain_mask_accum > 0).float() + combined = torch.maximum(binary_mask, binary_accum) + lgw = self.lgw[step_sched] if step_sched < len(self.lgw) else 0.0 + lgw_mask = combined * lgw + self._self_refine_certain_mask_accum = lgw_mask.clone() + + if lgw_mask.max() == 0: + if self.EO("self_refine_debug"): + RESplain(f"self_refine_epsilon step {step}, iter {full_iter}, row {row}: SKIPPED - no certain regions") + return eps_, x_ + + # Check coverage against cutoff + coverage = (lgw_mask > 0).float().mean().item() + if coverage >= self.self_refine_cutoff: + self._self_refine_converged = True + if self.EO("self_refine_debug"): + iter_info = f", iter {full_iter}" if per_iteration_mode else "" + RESplain(f"self_refine_epsilon step {step}{iter_info}, row {row}: CONVERGED - coverage={coverage:.2%} >= cutoff={self.self_refine_cutoff:.2%}") + + # Compute guide epsilon (direction toward reference) + eps_y0 = RK.get_guide_epsilon(x_0, x_row, y0, sigma, sigma_row, sigma_down, None) + + # Blend: anchor certain regions toward reference + if "_projection" in self.guide_mode: + # Projection variant: preserve magnitude, steer direction + eps_row_lerp = eps_row + lgw_mask * (eps_y0 - eps_row) + eps_collinear = get_collinear(eps_row, eps_row_lerp) + eps_ortho = get_orthogonal(eps_row_lerp, eps_row) + eps_sum = eps_collinear + eps_ortho + eps_[row] = eps_row + lgw_mask * (eps_sum - eps_row) + else: + # Standard lerp blending + eps_[row] = eps_row + lgw_mask * (eps_y0 - eps_row) + + if self.EO("self_refine_debug"): + call_type = "eps_" if self.self_refine_epsilon_call_count == 1 else "eps_prev_" + coverage = (lgw_mask > 0).float().mean().item() + iter_info = f", iter {full_iter}" if per_iteration_mode else "" + RESplain(f"self_refine_epsilon step {step}{iter_info}, row {row} ({call_type}): APPLIED - certain_coverage={coverage:.2%}, eps mean/std={eps_[row].mean():.4f}/{eps_[row].std():.4f}") + + return eps_, x_ + + # Local references for the row tensors we'll modify + eps_row = eps_[row] + data_row = data_[row] + x_row = x_[row] + x_row1 = x_[row+1] + + y0, y0_inv, lgw_mask, lgw_mask_inv = self.get_cossim_adjusted_lgw_masks(data_row, step_sched) + + if not (lgw_mask.any() != 0 or lgw_mask_inv.any() != 0): # cossim score too similar! deactivate guide for this step + return eps_, x_ + + if self.EO(["substep_eps_ch_mean_std", "substep_eps_ch_mean", "substep_eps_ch_std", "substep_eps_mean_std", "substep_eps_mean", "substep_eps_std"]): + eps_row_orig = eps_row.clone() + + if self.EO("dynamic_guides_mean_std"): + y_shift, y_inv_shift = normalize_latent([y0, y0_inv], [data_row, data_row]) + y0 = y_shift + if self.EO("dynamic_guides_inv"): + y0_inv = y_inv_shift + + if self.EO("dynamic_guides_mean"): + y_shift, y_inv_shift = normalize_latent([y0, y0_inv], [data_row, data_row], std=False) + y0 = y_shift + if self.EO("dynamic_guides_inv"): + y0_inv = y_inv_shift + + + + if "data_old" == self.guide_mode: + y0_tmp = y0.clone() + if self.HAS_LATENT_GUIDE: + y0_tmp = (1-lgw_mask) * data_row + lgw_mask * y0 + y0_tmp = (1-lgw_mask_inv) * y0_tmp + lgw_mask_inv * y0_inv + x_row1 = y0_tmp + eps_row + + if self.guide_mode == "data_old_projection": + + d_lerp = data_row + lgw_mask * (y0-data_row) + lgw_mask_inv * (y0_inv-data_row) + + d_collinear_d_lerp = get_collinear(data_row, d_lerp) + d_lerp_ortho_d = get_orthogonal(d_lerp, data_row) + + data_row = d_collinear_d_lerp + d_lerp_ortho_d + + x_row1 = data_row + eps_row * sigma + + + + #elif (self.UNSAMPLE or self.guide_mode in {"epsilon", "epsilon_cw", "epsilon_projection", "epsilon_projection_cw"}) and (self.lgw[step] > 0 or self.lgw_inv[step] > 0): + elif self.guide_mode in {"epsilon", "epsilon_cw", "epsilon_projection", "epsilon_projection_cw"} and (self.lgw[step_sched] > 0 or self.lgw_inv[step_sched] > 0): + if sigma_down < sigma or s_[row] < RK.sigma_max: + + eps_substep_guide = torch.zeros_like(x_0) + eps_substep_guide_inv = torch.zeros_like(x_0) + + if self.HAS_LATENT_GUIDE: + eps_substep_guide = RK.get_guide_epsilon(x_0, x_row, y0, sigma, s_[row], sigma_down, epsilon_scale) + + if self.HAS_LATENT_GUIDE_INV: + eps_substep_guide_inv = RK.get_guide_epsilon(x_0, x_row, y0_inv, sigma, s_[row], sigma_down, epsilon_scale) + + tol_value = self.EO("tol", -1.0) + if tol_value >= 0: + if is_flat: + raise NotImplementedError("Tolerance mode requires batch/channel structure, incompatible with packed latents") + for b, c in itertools.product(range(x_0.shape[0]), range(x_0.shape[1])): + current_diff = torch.norm(data_[row][b][c] - y0 [b][c]) + current_diff_inv = torch.norm(data_[row][b][c] - y0_inv[b][c]) + + lgw_scaled = torch.nan_to_num(1-(tol_value/current_diff), 0) + lgw_scaled_inv = torch.nan_to_num(1-(tol_value/current_diff_inv), 0) + + lgw_tmp = min(self.lgw[step_sched] , lgw_scaled) + lgw_tmp_inv = min(self.lgw_inv[step_sched], lgw_scaled_inv) + + lgw_mask_clamp = torch.clamp(lgw_mask, max=lgw_tmp) + lgw_mask_clamp_inv = torch.clamp(lgw_mask_inv, max=lgw_tmp_inv) + + eps_[row][b][c] = eps_[row][b][c] + lgw_mask_clamp[b][0] * (eps_substep_guide[b][c] - eps_[row][b][c]) + lgw_mask_clamp_inv[b][0] * (eps_substep_guide_inv[b][c] - eps_[row][b][c]) + + elif self.guide_mode in {"epsilon"}: + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide.sum() != 0: + eps_row = slerp_tensor(lgw_mask, eps_row, eps_substep_guide) + if eps_substep_guide_inv.sum() != 0: + eps_row = slerp_tensor(lgw_mask_inv, eps_row, eps_substep_guide_inv) + else: + eps_row = eps_row + lgw_mask * (eps_substep_guide - eps_row) + lgw_mask_inv * (eps_substep_guide_inv - eps_row) + + elif self.guide_mode in {"epsilon_projection"}: + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide.sum() != 0: + eps_row_slerp = slerp_tensor(self.mask, eps_row, eps_substep_guide) + if eps_substep_guide_inv.sum() != 0: + eps_row_slerp = slerp_tensor((1-self.mask), eps_row_slerp, eps_substep_guide_inv) + + eps_collinear_eps_slerp = get_collinear(eps_row, eps_row_slerp) + eps_slerp_ortho_eps = get_orthogonal(eps_row_slerp, eps_row) + + eps_sum = eps_collinear_eps_slerp + eps_slerp_ortho_eps + + eps_row = slerp_tensor(lgw_mask, eps_row, eps_sum) + eps_row = slerp_tensor(lgw_mask_inv, eps_row, eps_sum) + else: + eps_row_lerp = eps_row + self.mask * (eps_substep_guide-eps_row) + (1-self.mask) * (eps_substep_guide_inv-eps_row) + + eps_collinear_eps_lerp = get_collinear(eps_row, eps_row_lerp) + eps_lerp_ortho_eps = get_orthogonal(eps_row_lerp, eps_row) + + eps_sum = eps_collinear_eps_lerp + eps_lerp_ortho_eps + + eps_row = eps_row + lgw_mask * (eps_sum - eps_row) + lgw_mask_inv * (eps_sum - eps_row) + + elif self.guide_mode in {"epsilon_cw", "epsilon_projection_cw"}: + eps_ = self.process_channelwise(x_0, + eps_, + data_, + row, + eps_substep_guide, + eps_substep_guide_inv, + y0, + y0_inv, + lgw_mask, + lgw_mask_inv, + use_projection = self.guide_mode == "epsilon_projection_cw", + channelwise = True + ) + + temporal_smoothing = self.EO("temporal_smoothing", 0.0) + if temporal_smoothing > 0: + if is_flat: + raise NotImplementedError("Temporal smoothing requires temporal structure, incompatible with packed latents") + eps_row = apply_temporal_smoothing(eps_row, temporal_smoothing) + + if self.EO("substep_eps_ch_mean_std"): + eps_row = normalize_latent(eps_row, eps_row_orig) + if self.EO("substep_eps_ch_mean"): + eps_row = normalize_latent(eps_row, eps_row_orig, std=False) + if self.EO("substep_eps_ch_std"): + eps_row = normalize_latent(eps_row, eps_row_orig, mean=False) + if self.EO("substep_eps_mean_std"): + eps_row = normalize_latent(eps_row, eps_row_orig, channelwise=False) + if self.EO("substep_eps_mean"): + eps_row = normalize_latent(eps_row, eps_row_orig, std=False, channelwise=False) + if self.EO("substep_eps_std"): + eps_row = normalize_latent(eps_row, eps_row_orig, mean=False, channelwise=False) + + # Write results back to tensors + eps_[row] = eps_row + if self.guide_mode in {"data_old", "data_old_projection"}: + x_[row+1] = x_row1 + if self.guide_mode == "data_old_projection": + data_[row] = data_row + + return eps_, x_ + + + def process_channelwise(self, + x_0 : Tensor, + eps_ : Tensor, + data_ : Tensor, + row : int, + eps_substep_guide : Tensor, + eps_substep_guide_inv : Tensor, + y0 : Tensor, + y0_inv : Tensor, + lgw_mask : Tensor, + lgw_mask_inv : Tensor, + use_projection : bool = False, + channelwise : bool = False + ): + + avg, avg_inv = 0, 0 + lgw_mask_channels = lgw_mask.shape[1] + lgw_mask_inv_channels = lgw_mask_inv.shape[1] + + for b, c in itertools.product(range(x_0.shape[0]), range(x_0.shape[1])): + if self.EO("lgw_cw_test") and (lgw_mask_channels > 1 or lgw_mask_inv_channels > 1): + c_ = c % (lgw_mask_channels if lgw_mask_channels > 1 else lgw_mask_inv_channels) + else: + c_ = 0 + avg += torch.norm(lgw_mask [b][c_] * data_[row][b][c] - lgw_mask [b][c_] * y0 [b][c]) + avg_inv += torch.norm(lgw_mask_inv[b][c_] * data_[row][b][c] - lgw_mask_inv[b][c_] * y0_inv[b][c]) + + avg /= x_0.shape[1] + avg_inv /= x_0.shape[1] + + for b, c in itertools.product(range(x_0.shape[0]), range(x_0.shape[1])): + if channelwise: + ratio = torch.nan_to_num(torch.norm(lgw_mask [b][c_] * data_[row][b][c] - lgw_mask [b][c_] * y0 [b][c]) / avg, 0) + ratio_inv = torch.nan_to_num(torch.norm(lgw_mask_inv[b][c_] * data_[row][b][c] - lgw_mask_inv[b][c_] * y0_inv[b][c]) / avg_inv, 0) + else: + ratio = 1. + ratio_inv = 1. + + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio * lgw_mask[b][0], eps_[row][b][c], eps_substep_guide[b][c]) + if eps_substep_guide_inv[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio_inv * lgw_mask_inv[b][0], eps_[row][b][c], eps_substep_guide_inv[b][c]) + else: + eps_[row][b][c] = eps_[row][b][c] + ratio * lgw_mask[b][0] * (eps_substep_guide[b][c] - eps_[row][b][c]) + ratio_inv * lgw_mask_inv[b][0] * (eps_substep_guide_inv[b][c] - eps_[row][b][c]) + + if use_projection: + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide[b][c].sum() != 0: + eps_row_lerp = slerp_tensor(self.mask[b][0], eps_[row][b][c], eps_substep_guide[b][c]) + if eps_substep_guide_inv[b][c].sum() != 0: + eps_row_lerp = slerp_tensor((1-self.mask[b][0]), eps_[row][b][c], eps_substep_guide_inv[b][c]) + else: + eps_row_lerp = eps_[row][b][c] + self.mask[b][0] * (eps_substep_guide[b][c] - eps_[row][b][c]) + (1-self.mask[b][0]) * (eps_substep_guide_inv[b][c] - eps_[row][b][c]) # should this ever be self.mask_inv? + + eps_collinear_eps_lerp = get_collinear (eps_[row][b][c], eps_row_lerp) + eps_lerp_ortho_eps = get_orthogonal(eps_row_lerp , eps_[row][b][c]) + + eps_sum = eps_collinear_eps_lerp + eps_lerp_ortho_eps + + + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio * lgw_mask[b][0], eps_[row][b][c], eps_sum) + if eps_substep_guide_inv[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio_inv * lgw_mask_inv[b][0], eps_[row][b][c], eps_sum) + else: + eps_[row][b][c] = eps_[row][b][c] + ratio * lgw_mask[b][0] * (eps_sum - eps_[row][b][c]) + ratio_inv * lgw_mask_inv[b][0] * (eps_sum - eps_[row][b][c]) + else: + if self.EO("slerp_epsilon_guide"): + if eps_substep_guide[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio * lgw_mask[b][0], eps_[row][b][c], eps_substep_guide[b][c]) + if eps_substep_guide_inv[b][c].sum() != 0: + eps_[row][b][c] = slerp_tensor(ratio_inv * lgw_mask_inv[b][0], eps_[row][b][c], eps_substep_guide_inv[b][c]) + else: + eps_[row][b][c] = eps_[row][b][c] + ratio * lgw_mask[b][0] * (eps_substep_guide[b][c] - eps_[row][b][c]) + ratio_inv * lgw_mask_inv[b][0] * (eps_substep_guide_inv[b][c] - eps_[row][b][c]) + + return eps_ + + + def normalize_inputs(self, x:Tensor, y0:Tensor, y0_inv:Tensor): + """ + Modifies and returns 'x' by matching its mean and/or std to y0 and/or y0_inv. + Controlled by extra_options. + + Returns: + - x (modified) + - y0 (may be modified to match mean and std from y0_inv) + - y0_inv (unchanged) + """ + if self.guide_mode == "epsilon_guide_mean_std_from_bkg": + y0 = normalize_latent(y0, y0_inv) + + input_norm = self.EO("input_norm", "") + input_std = self.EO("input_std", 1.0) + + if input_norm == "input_ch_mean_set_std_to": + x = normalize_latent(x, set_std=input_std) + + if input_norm == "input_ch_set_std_to": + x = normalize_latent(x, set_std=input_std, mean=False) + + if input_norm == "input_mean_set_std_to": + x = normalize_latent(x, set_std=input_std, channelwise=False) + + if input_norm == "input_std_set_std_to": + x = normalize_latent(x, set_std=input_std, mean=False, channelwise=False) + + return x, y0, y0_inv + + + +def apply_frame_weights(mask, frame_weights, normalize=False): + original_mask_mean = mask.mean() + if frame_weights is not None: + for f in range(mask.shape[2]): + frame_weight = frame_weights[f] + mask[..., f:f+1, :, :] *= frame_weight + if normalize: + mask_mean = mask.mean() + mask *= (original_mask_mean / mask_mean) + + + +def prepare_mask(x, mask, LGW_MASK_RESCALE_MIN) -> tuple[torch.Tensor, bool]: + if mask is None: + mask = torch.ones_like(x[:,0:1,...]) + LGW_MASK_RESCALE_MIN = False + return mask, LGW_MASK_RESCALE_MIN + + # For flat/packed tensors, expect mask to already be compatible or use ones + if x.ndim == 3: + if mask.numel() == x.numel() or mask.numel() == x.shape[-1]: + return mask.reshape(1, 1, -1).expand_as(x).to(x.dtype), LGW_MASK_RESCALE_MIN + else: + # Mask incompatible with flat tensor, use ones + RESplain(f"Warning: mask shape {mask.shape} incompatible with flat tensor shape {x.shape}, using ones mask instead.") + return torch.ones_like(x), False + + target_height = x.shape[-2] + target_width = x.shape[-1] + + spatial_mask = None + if x.ndim == 5 and mask.shape[0] > 1 and mask.ndim < 4: + target_frames = x.shape[-3] + spatial_mask = mask.unsqueeze(0).unsqueeze(0) # [B, H, W] -> [1, 1, B, H, W] + spatial_mask = F.interpolate(spatial_mask, + size=(target_frames, target_height, target_width), + mode='trilinear', + align_corners=False) # [1, 1, F, H, W] + repeat_shape = [1] # batch + for i in range(1, x.ndim - 3): + repeat_shape.append(x.shape[i]) + repeat_shape.extend([1, 1, 1]) # frames, height, width + elif mask.ndim == 4: #temporal mask batch + mask = F.interpolate(mask, size=(target_height, target_width), mode='bilinear', align_corners=False) + mask = mask.repeat(x.shape[-4],1,1,1) + mask.unsqueeze_(0) + + else: + spatial_mask = mask.unsqueeze(1) + spatial_mask = F.interpolate(spatial_mask, size=(target_height, target_width), mode='bilinear', align_corners=False) + + while spatial_mask.ndim < x.ndim: + spatial_mask = spatial_mask.unsqueeze(2) + + repeat_shape = [1] # batch + for i in range(1, x.ndim - 2): + repeat_shape.append(x.shape[i]) + repeat_shape.extend([1, 1]) # height and width + repeat_shape[1] = 1 # only need one channel for masks + + if spatial_mask is not None: + mask = spatial_mask.repeat(*repeat_shape).to(x.dtype) + + del spatial_mask + return mask, LGW_MASK_RESCALE_MIN + +def apply_temporal_smoothing(tensor, temporal_smoothing): + if temporal_smoothing <= 0 or tensor.ndim != 5: + return tensor + + kernel_size = 5 + padding = kernel_size // 2 + temporal_kernel = torch.tensor( + [0.1, 0.2, 0.4, 0.2, 0.1], + device=tensor.device, dtype=tensor.dtype + ) * temporal_smoothing + temporal_kernel[kernel_size//2] += (1 - temporal_smoothing) + temporal_kernel = temporal_kernel / temporal_kernel.sum() + + # resahpe for conv1d + b, c, f, h, w = tensor.shape + data_flat = tensor.permute(0, 1, 3, 4, 2).reshape(-1, f) + + # apply smoohting + data_smooth = F.conv1d( + data_flat.unsqueeze(1), + temporal_kernel.view(1, 1, -1), + padding=padding + ).squeeze(1) + + return data_smooth.view(b, c, h, w, f).permute(0, 1, 4, 2, 3) + +def get_guide_epsilon_substep(x_0, x_, y0, y0_inv, s_, row, row_offset, rk_type, b=None, c=None): + s_in = x_0.new_ones([x_0.shape[0]]) + + if b is not None and c is not None: + index = (b, c) + elif b is not None: + index = (b,) + else: + index = () + + if RK_Method_Beta.is_exponential(rk_type): + eps_row = y0 [index] - x_0[index] + eps_row_inv = y0_inv[index] - x_0[index] + else: + eps_row = (x_[row][index] - y0 [index]) / (s_[row] * s_in) # was row+row_offset before for x_!! not right... also? potential issues here with x_[row+1] being RK.rows+2 with gauss-legendre_2s 1 imp step 1 imp substep + eps_row_inv = (x_[row][index] - y0_inv[index]) / (s_[row] * s_in) + + return eps_row, eps_row_inv + +def get_guide_epsilon(x_0, x_, y0, sigma, rk_type, b=None, c=None): + s_in = x_0.new_ones([x_0.shape[0]]) + + if b is not None and c is not None: + index = (b, c) + elif b is not None: + index = (b,) + else: + index = () + + if RK_Method_Beta.is_exponential(rk_type): + eps = y0 [index] - x_0[index] + else: + eps = (x_[index] - y0 [index]) / (sigma * s_in) + + return eps + + + +@torch.no_grad +def noise_cossim_guide_tiled(x_list, guide, cossim_mode="forward", tile_size=2, step=0): + + guide_tiled = rearrange(guide, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size) + + x_tiled_list = [ + rearrange(x, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size) + for x in x_list + ] + x_tiled_stack = torch.stack([x_tiled[0] for x_tiled in x_tiled_list]) # [n_x, n_tiles, c, h, w] + + guide_flat = guide_tiled[0].view(guide_tiled.shape[1], -1).unsqueeze(0) # [1, n_tiles, c*h*w] + x_flat = x_tiled_stack.view(x_tiled_stack.size(0), x_tiled_stack.size(1), -1) # [n_x, n_tiles, c*h*w] + + cossim_tmp_all = F.cosine_similarity(x_flat, guide_flat, dim=-1) # [n_x, n_tiles] + + if cossim_mode == "forward": + indices = cossim_tmp_all.argmax(dim=0) + elif cossim_mode == "reverse": + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal": + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + elif cossim_mode == "forward_reverse": + if step % 2 == 0: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_forward": + if step % 2 == 1: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal_reverse": + if step % 2 == 0: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_orthogonal": + if step % 2 == 1: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + else: + target_value = float(cossim_mode) + indices = torch.abs(cossim_tmp_all - target_value).argmin(dim=0) + + x_tiled_out = x_tiled_stack[indices, torch.arange(indices.size(0))] # [n_tiles, c, h, w] + + x_tiled_out = x_tiled_out.unsqueeze(0) + x_detiled = rearrange(x_tiled_out, "b (t1 t2) c h w -> b c (h t1) (w t2)", t1=tile_size, t2=tile_size) + + return x_detiled + + +@torch.no_grad +def noise_cossim_eps_tiled(x_list, eps, noise_list, cossim_mode="forward", tile_size=2, step=0): + + eps_tiled = rearrange(eps, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size) + x_tiled_list = [ + rearrange(x, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size) + for x in x_list + ] + noise_tiled_list = [ + rearrange(noise, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size) + for noise in noise_list + ] + + noise_tiled_stack = torch.stack([noise_tiled[0] for noise_tiled in noise_tiled_list]) # [n_x, n_tiles, c, h, w] + eps_expanded = eps_tiled[0].view(eps_tiled.shape[1], -1).unsqueeze(0) # [1, n_tiles, c*h*w] + noise_flat = noise_tiled_stack.view(noise_tiled_stack.size(0), noise_tiled_stack.size(1), -1) # [n_x, n_tiles, c*h*w] + cossim_tmp_all = F.cosine_similarity(noise_flat, eps_expanded, dim=-1) # [n_x, n_tiles] + + if cossim_mode == "forward": + indices = cossim_tmp_all.argmax(dim=0) + elif cossim_mode == "reverse": + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal": + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + elif cossim_mode == "orthogonal_pos": + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + elif cossim_mode == "orthogonal_neg": + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "orthogonal_posneg": + if step % 2 == 0: + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + else: + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "orthogonal_negpos": + if step % 2 == 1: + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + else: + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "forward_reverse": + if step % 2 == 0: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_forward": + if step % 2 == 1: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal_reverse": + if step % 2 == 0: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_orthogonal": + if step % 2 == 1: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + else: + target_value = float(cossim_mode) + indices = torch.abs(cossim_tmp_all - target_value).argmin(dim=0) + #else: + # raise ValueError(f"Unknown cossim_mode: {cossim_mode}") + + x_tiled_stack = torch.stack([x_tiled[0] for x_tiled in x_tiled_list]) # [n_x, n_tiles, c, h, w] + x_tiled_out = x_tiled_stack[indices, torch.arange(indices.size(0))] # [n_tiles, c, h, w] + + x_tiled_out = x_tiled_out.unsqueeze(0) # restore batch dim + x_detiled = rearrange(x_tiled_out, "b (t1 t2) c h w -> b c (h t1) (w t2)", t1=tile_size, t2=tile_size) + return x_detiled + + + +@torch.no_grad +def noise_cossim_guide_eps_tiled(x_0, x_list, y0, noise_list, cossim_mode="forward", tile_size=2, step=0, sigma=None, rk_type=None): + + x_tiled_stack = torch.stack([ + rearrange(x, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size)[0] + for x in x_list + ]) # [n_x, n_tiles, c, h, w] + eps_guide_stack = torch.stack([ + rearrange(x - y0, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size)[0] + for x in x_list + ]) # [n_x, n_tiles, c, h, w] + del x_list + + noise_tiled_stack = torch.stack([ + rearrange(noise, "b c (h t1) (w t2) -> b (t1 t2) c h w", t1=tile_size, t2=tile_size)[0] + for noise in noise_list + ]) # [n_x, n_tiles, c, h, w] + del noise_list + + noise_flat = noise_tiled_stack.view(noise_tiled_stack.size(0), noise_tiled_stack.size(1), -1) # [n_x, n_tiles, c*h*w] + eps_guide_flat = eps_guide_stack.view(eps_guide_stack.size(0), eps_guide_stack.size(1), -1) # [n_x, n_tiles, c*h*w] + + cossim_tmp_all = F.cosine_similarity(noise_flat, eps_guide_flat, dim=-1) # [n_x, n_tiles] + del noise_tiled_stack, noise_flat, eps_guide_stack, eps_guide_flat + + if cossim_mode == "forward": + indices = cossim_tmp_all.argmax(dim=0) + elif cossim_mode == "reverse": + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal": + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + elif cossim_mode == "orthogonal_pos": + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + elif cossim_mode == "orthogonal_neg": + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "orthogonal_posneg": + if step % 2 == 0: + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + else: + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "orthogonal_negpos": + if step % 2 == 1: + positive_mask = cossim_tmp_all > 0 + positive_tmp = torch.where(positive_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('inf'))) + indices = positive_tmp.argmin(dim=0) + else: + negative_mask = cossim_tmp_all < 0 + negative_tmp = torch.where(negative_mask, cossim_tmp_all, torch.full_like(cossim_tmp_all, float('-inf'))) + indices = negative_tmp.argmax(dim=0) + elif cossim_mode == "forward_reverse": + if step % 2 == 0: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_forward": + if step % 2 == 1: + indices = cossim_tmp_all.argmax(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "orthogonal_reverse": + if step % 2 == 0: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + elif cossim_mode == "reverse_orthogonal": + if step % 2 == 1: + indices = torch.abs(cossim_tmp_all).argmin(dim=0) + else: + indices = cossim_tmp_all.argmin(dim=0) + else: + target_value = float(cossim_mode) + indices = torch.abs(cossim_tmp_all - target_value).argmin(dim=0) + + x_tiled_out = x_tiled_stack[indices, torch.arange(indices.size(0))] # [n_tiles, c, h, w] + del x_tiled_stack + + x_tiled_out = x_tiled_out.unsqueeze(0) + x_detiled = rearrange(x_tiled_out, "b (t1 t2) c h w -> b c (h t1) (w t2)", t1=tile_size, t2=tile_size) + + return x_detiled + + + + + + + +class NoiseStepHandlerOSDE: + def __init__(self, x, eps=None, data=None, x_init=None, guide=None, guide_bkg=None): + self.noise = None + self.x = x + self.eps = eps + self.data = data + self.x_init = x_init + self.guide = guide + self.guide_bkg = guide_bkg + + self.eps_list = None + + self.noise_cossim_map = { + "eps_orthogonal": [self.noise, self.eps], + "eps_data_orthogonal": [self.noise, self.eps, self.data], + + "data_orthogonal": [self.noise, self.data], + "xinit_orthogonal": [self.noise, self.x_init], + + "x_orthogonal": [self.noise, self.x], + "x_data_orthogonal": [self.noise, self.x, self.data], + "x_eps_orthogonal": [self.noise, self.x, self.eps], + + "x_eps_data_orthogonal": [self.noise, self.x, self.eps, self.data], + "x_eps_data_xinit_orthogonal": [self.noise, self.x, self.eps, self.data, self.x_init], + + "x_eps_guide_orthogonal": [self.noise, self.x, self.eps, self.guide], + "x_eps_guide_bkg_orthogonal": [self.noise, self.x, self.eps, self.guide_bkg], + + "noise_orthogonal": [self.noise, self.x_init], + + "guide_orthogonal": [self.noise, self.guide], + "guide_bkg_orthogonal": [self.noise, self.guide_bkg], + } + + def check_cossim_source(self, source): + return source in self.noise_cossim_map + + def get_ortho_noise(self, noise, prev_noises=None, max_iter=100, max_score=1e-7, NOISE_COSSIM_SOURCE="eps_orthogonal"): + + if NOISE_COSSIM_SOURCE not in self.noise_cossim_map: + raise ValueError(f"Invalid NOISE_COSSIM_SOURCE: {NOISE_COSSIM_SOURCE}") + + self.noise_cossim_map[NOISE_COSSIM_SOURCE][0] = noise + + params = self.noise_cossim_map[NOISE_COSSIM_SOURCE] + + noise = get_orthogonal_noise_from_channelwise(*params, max_iter=max_iter, max_score=max_score) + + return noise + + + + + +# NOTE: NS AND SUBSTEP ADDED! +def handle_tiled_etc_noise_steps( + x_0, + x, + x_prenoise, + x_init, + eps, + denoised, + y0, + y0_inv, + step, + rk_type, + RK, + NS, + SUBSTEP, + sigma_up, + sigma, + sigma_next, + alpha_ratio, + s_noise, + noise_mode, + SDE_NOISE_EXTERNAL, + sde_noise_t, + NOISE_COSSIM_SOURCE, + NOISE_COSSIM_MODE, + noise_cossim_tile_size, + noise_cossim_iterations, + extra_options): + + EO = ExtraOptions(extra_options) + + x_tmp = [] + cossim_tmp = [] + noise_tmp_list = [] + + if step > EO("noise_cossim_end_step", MAX_STEPS): + NOISE_COSSIM_SOURCE = EO("noise_cossim_takeover_source" , "eps") + NOISE_COSSIM_MODE = EO("noise_cossim_takeover_mode" , "forward" ) + noise_cossim_tile_size = EO("noise_cossim_takeover_tile" , noise_cossim_tile_size ) + noise_cossim_iterations = EO("noise_cossim_takeover_iterations", noise_cossim_iterations) + + for i in range(noise_cossim_iterations): + #x_tmp.append(NS.swap_noise(x_0, x, sigma, sigma, sigma_next, )) + x_tmp.append(NS.add_noise_post(x, sigma_up, sigma, sigma_next, alpha_ratio, s_noise, noise_mode, SDE_NOISE_EXTERNAL, sde_noise_t) )#y0, lgw, sigma_down are currently unused + noise_tmp = x_tmp[i] - x + if EO("noise_noise_zscore_norm"): + noise_tmp = normalize_zscore(noise_tmp, channelwise=False, inplace=True) + if EO("noise_noise_zscore_norm_cw"): + noise_tmp = normalize_zscore(noise_tmp, channelwise=True, inplace=True) + if EO("noise_eps_zscore_norm"): + eps = normalize_zscore(eps, channelwise=False, inplace=True) + if EO("noise_eps_zscore_norm_cw"): + eps = normalize_zscore(eps, channelwise=True, inplace=True) + + if NOISE_COSSIM_SOURCE in ("eps_tiled", "guide_epsilon_tiled", "guide_bkg_epsilon_tiled", "iig_tiled"): + noise_tmp_list.append(noise_tmp) + if NOISE_COSSIM_SOURCE == "eps": + cossim_tmp.append(get_cosine_similarity(eps, noise_tmp)) + if NOISE_COSSIM_SOURCE == "eps_ch": + cossim_total = torch.zeros_like(eps[0][0][0][0]) + for ch in range(eps.shape[1]): + cossim_total += get_cosine_similarity(eps[0][ch], noise_tmp[0][ch]) + cossim_tmp.append(cossim_total) + elif NOISE_COSSIM_SOURCE == "data": + cossim_tmp.append(get_cosine_similarity(denoised, noise_tmp)) + elif NOISE_COSSIM_SOURCE == "latent": + cossim_tmp.append(get_cosine_similarity(x_prenoise, noise_tmp)) + elif NOISE_COSSIM_SOURCE == "x_prenoise": + cossim_tmp.append(get_cosine_similarity(x_prenoise, x_tmp[i])) + elif NOISE_COSSIM_SOURCE == "x": + cossim_tmp.append(get_cosine_similarity(x, x_tmp[i])) + elif NOISE_COSSIM_SOURCE == "x_data": + cossim_tmp.append(get_cosine_similarity(denoised, x_tmp[i])) + elif NOISE_COSSIM_SOURCE == "x_init_vs_noise": + cossim_tmp.append(get_cosine_similarity(x_init, noise_tmp)) + elif NOISE_COSSIM_SOURCE == "mom": + cossim_tmp.append(get_cosine_similarity(denoised, x + sigma_next*noise_tmp)) + elif NOISE_COSSIM_SOURCE == "guide": + cossim_tmp.append(get_cosine_similarity(y0, x_tmp[i])) + elif NOISE_COSSIM_SOURCE == "guide_bkg": + cossim_tmp.append(get_cosine_similarity(y0_inv, x_tmp[i])) + + if step < EO("noise_cossim_start_step", 0): + x = x_tmp[0] + + elif (NOISE_COSSIM_SOURCE == "eps_tiled"): + x = noise_cossim_eps_tiled(x_tmp, eps, noise_tmp_list, cossim_mode=NOISE_COSSIM_MODE, tile_size=noise_cossim_tile_size, step=step) + elif (NOISE_COSSIM_SOURCE == "guide_epsilon_tiled"): + x = noise_cossim_guide_eps_tiled(x_0, x_tmp, y0, noise_tmp_list, cossim_mode=NOISE_COSSIM_MODE, tile_size=noise_cossim_tile_size, step=step, sigma=sigma, rk_type=rk_type) + elif (NOISE_COSSIM_SOURCE == "guide_bkg_epsilon_tiled"): + x = noise_cossim_guide_eps_tiled(x_0, x_tmp, y0_inv, noise_tmp_list, cossim_mode=NOISE_COSSIM_MODE, tile_size=noise_cossim_tile_size, step=step, sigma=sigma, rk_type=rk_type) + elif (NOISE_COSSIM_SOURCE == "guide_tiled"): + x = noise_cossim_guide_tiled(x_tmp, y0, cossim_mode=NOISE_COSSIM_MODE, tile_size=noise_cossim_tile_size, step=step) + elif (NOISE_COSSIM_SOURCE == "guide_bkg_tiled"): + x = noise_cossim_guide_tiled(x_tmp, y0_inv, cossim_mode=NOISE_COSSIM_MODE, tile_size=noise_cossim_tile_size) + else: + for i in range(len(x_tmp)): + if (NOISE_COSSIM_MODE == "forward") and (cossim_tmp[i] == max(cossim_tmp)): + x = x_tmp[i] + break + elif (NOISE_COSSIM_MODE == "reverse") and (cossim_tmp[i] == min(cossim_tmp)): + x = x_tmp[i] + break + elif (NOISE_COSSIM_MODE == "orthogonal") and (abs(cossim_tmp[i]) == min(abs(val) for val in cossim_tmp)): + x = x_tmp[i] + break + elif (NOISE_COSSIM_MODE != "forward") and (NOISE_COSSIM_MODE != "reverse") and (NOISE_COSSIM_MODE != "orthogonal"): + x = x_tmp[0] + break + return x + + + + + +def get_masked_epsilon_projection(x_0, x_, eps_, y0, y0_inv, s_, row, row_offset, rk_type, LG, step): + + eps_row, eps_row_inv = get_guide_epsilon_substep(x_0, x_, y0, y0_inv, s_, row, row_offset, rk_type) + eps_row_lerp = eps_[row] + LG.mask * (eps_row-eps_[row]) + (1-LG.mask) * (eps_row_inv-eps_[row]) + eps_collinear_eps_lerp = get_collinear(eps_[row], eps_row_lerp) + eps_lerp_ortho_eps = get_orthogonal(eps_row_lerp, eps_[row]) + eps_sum = eps_collinear_eps_lerp + eps_lerp_ortho_eps + lgw_mask, lgw_mask_inv = LG.get_masks_for_step(step) + eps_substep_guide = eps_[row] + lgw_mask * (eps_sum - eps_[row]) + lgw_mask_inv * (eps_sum - eps_[row]) + return eps_substep_guide + + + diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/rk_method_beta.py b/simple_syrup/third_party/res4lyf_runtime/beta/rk_method_beta.py new file mode 100644 index 0000000..37651bc --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/rk_method_beta.py @@ -0,0 +1,1218 @@ +import torch +from torch import Tensor +from typing import Optional, Callable, Tuple, List, Dict, Any, Union + +import comfy.model_patcher +import comfy.supported_models + +import itertools + +from .phi_functions import Phi +from .rk_coefficients_beta import get_implicit_sampler_name_list, get_rk_methods_beta +from ..helper import ExtraOptions +from ..latents import get_orthogonal, get_collinear, get_cosine_similarity, tile_latent, untile_latent + +from ..res4lyf import RESplain, is_debug_logging_enabled + +MAX_STEPS = 10000 + + +def get_data_from_step (x:Tensor, x_next:Tensor, sigma:Tensor, sigma_next:Tensor) -> Tensor: + h = sigma_next - sigma + return (sigma_next * x - sigma * x_next) / h + +def get_epsilon_from_step(x:Tensor, x_next:Tensor, sigma:Tensor, sigma_next:Tensor) -> Tensor: + h = sigma_next - sigma + return (x - x_next) / h + + + +class RK_Method_Beta: + def __init__(self, + model, + rk_type : str, + VE_MODEL : bool, + noise_anchor : float, + noise_boost_normalize : bool = True, + model_device : str = 'cuda', + work_device : str = 'cpu', + dtype : torch.dtype = torch.float64, + extra_options : str = "" + ): + + self.work_device = work_device + self.model_device = model_device + self.dtype : torch.dtype = dtype + + self.model = model + + if hasattr(model, "model"): + model_sampling = model.model.model_sampling + elif hasattr(model, "inner_model"): + model_sampling = model.inner_model.inner_model.model_sampling + + self.sigma_min : Tensor = model_sampling.sigma_min.to(dtype=dtype, device=work_device) + self.sigma_max : Tensor = model_sampling.sigma_max.to(dtype=dtype, device=work_device) + + self.rk_type : str = rk_type + + self.IMPLICIT : str = rk_type in get_implicit_sampler_name_list(nameOnly=True) + self.EXPONENTIAL : bool = RK_Method_Beta.is_exponential(rk_type) + self.VE_MODEL : bool = VE_MODEL + + self.SYNC_SUBSTEP_MEAN_CW : bool = noise_boost_normalize + + self.A : Optional[Tensor] = None + self.B : Optional[Tensor] = None + self.U : Optional[Tensor] = None + self.V : Optional[Tensor] = None + + self.rows : int = 0 + self.cols : int = 0 + + self.denoised : Optional[Tensor] = None + self.uncond : Optional[Tensor] = None + + self.y0 : Optional[Tensor] = None + self.y0_inv : Optional[Tensor] = None + + self.multistep_stages : int = 0 + self.row_offset : Optional[int] = None + + self.cfg_cw : float = 1.0 + self.extra_args : Optional[Dict[str, Any]] = None + + self.extra_options : str = extra_options + self.EO : ExtraOptions = ExtraOptions(extra_options) + + self.reorder_tableau_indices : list[int] = self.EO("reorder_tableau_indices", [-1]) + + # ComfyUI casts the latent to the model's compute dtype before the network runs, so precision + # above work_dtype never reaches the model — it only widens ComfyUI's cond/CFG temporaries. + self.work_dtype : torch.dtype = self.EO("work_dtype", torch.float32) + + self.LINEAR_ANCHOR_X_0 : float = noise_anchor + + self.tile_sizes : Optional[List[Tuple[int,int]]] = None + self.tile_cnt : int = 0 + self.latent_compression_ratio : int = 8 + # track model calls + self.model_calls_total : int = 0 + self.model_calls_denoised : int = 0 + self.model_calls_epsilon : int = 0 + + self.latent_guide = None + + @staticmethod + def is_exponential(rk_type:str) -> bool: + if rk_type.startswith(( "res", + "dpmpp", + "ddim", + "pec", + "etdrk", + "lawson", + "abnorsett", + )): + return True + else: + return False + + @staticmethod + def create(model, + rk_type : str, + VE_MODEL : bool, + noise_anchor : float = 1.0, + noise_boost_normalize : bool = True, + model_device : str = 'cuda', + work_device : str = 'cpu', + dtype : torch.dtype = torch.float64, + extra_options : str = "" + ) -> "Union[RK_Method_Exponential, RK_Method_Linear]": + + if RK_Method_Beta.is_exponential(rk_type): + return RK_Method_Exponential(model, rk_type, VE_MODEL, noise_anchor, noise_boost_normalize, model_device, work_device, dtype, extra_options) + else: + return RK_Method_Linear (model, rk_type, VE_MODEL, noise_anchor, noise_boost_normalize, model_device, work_device, dtype, extra_options) + + def __call__(self): + raise NotImplementedError("This method got clownsharked!") + + def _offload_peripherals(self): + if self.latent_guide is not None: + self.latent_guide.offload('cpu') + + def _restore_peripherals(self): + if self.latent_guide is not None: + self.latent_guide.restore() + + def model_epsilon(self, x:Tensor, sigma:Tensor, **extra_args) -> Tuple[Tensor, Tensor]: + if x.dtype != self.work_dtype: + x = x .to(self.work_dtype) + sigma = sigma.to(self.work_dtype) + s_in = x.new_ones([x.shape[0]]) + self._offload_peripherals() + denoised = self.model(x, sigma * s_in, **extra_args) + self._restore_peripherals() + # increment counters (single call path) + self.model_calls_total += 1 + self.model_calls_epsilon += 1 + denoised = self.calc_cfg_channelwise(denoised) + eps = (x - denoised) / (sigma * s_in).view(x.shape[0], *[1]*(x.ndim-1)) + return eps, denoised + + def model_denoised(self, x:Tensor, sigma:Tensor, **extra_args) -> Tensor: + if x.dtype != self.work_dtype: + x = x .to(self.work_dtype) + sigma = sigma.to(self.work_dtype) + s_in = x.new_ones([x.shape[0]]) + control_tiles = None + y0_style_pos = self.extra_args['model_options']['transformer_options'].get("y0_style_pos") + y0_style_neg = self.extra_args['model_options']['transformer_options'].get("y0_style_neg") + y0_style_pos_tile, sy0_style_neg_tiles = None, None + + self._offload_peripherals() + + if self.EO("tile_model_calls"): + tile_h = self.EO("tile_h", 128) + tile_w = self.EO("tile_w", 128) + + denoised_tiles = [] + + tiles, orig_shape, grid, strides = tile_latent(x, tile_size=(tile_h,tile_w)) + + for i in range(tiles.shape[0]): + tile = tiles[i].unsqueeze(0) + + denoised_tile = self.model(tile, sigma * s_in, **extra_args) + # increment counters per tile + self.model_calls_total += 1 + self.model_calls_denoised += 1 + denoised_tiles.append(denoised_tile) + + denoised_tiles = torch.cat(denoised_tiles, dim=0) + + denoised = untile_latent(denoised_tiles, orig_shape, grid, strides) + + elif self.tile_sizes is not None: + tile_h_full = self.tile_sizes[self.tile_cnt % len(self.tile_sizes)][0] + tile_w_full = self.tile_sizes[self.tile_cnt % len(self.tile_sizes)][1] + + if tile_h_full == -1: + tile_h = x.shape[-2] + tile_h_full = tile_h * self.latent_compression_ratio + else: + tile_h = tile_h_full // self.latent_compression_ratio + + if tile_w_full == -1: + tile_w = x.shape[-1] + tile_w_full = tile_w * self.latent_compression_ratio + else: + tile_w = tile_w_full // self.latent_compression_ratio + + #tile_h = tile_h_full // self.latent_compression_ratio + #tile_w = tile_w_full // self.latent_compression_ratio + + self.tile_cnt += 1 + + #if len(self.tile_sizes) == 1 and self.tile_cnt % 2 == 1: + # tile_h, tile_w = tile_w, tile_h + # tile_h_full, tile_w_full = tile_w_full, tile_h_full + + if (self.tile_cnt // len(self.tile_sizes)) % 2 == 1 and self.EO("tiles_autorotate"): + tile_h, tile_w = tile_w, tile_h + tile_h_full, tile_w_full = tile_w_full, tile_h_full + + xt_negative = self.model.inner_model.conds.get('xt_negative', self.model.inner_model.conds.get('negative')) + negative_control = xt_negative[0].get('control') + + if negative_control is not None and hasattr(negative_control, 'cond_hint_original'): + negative_cond_hint_init = negative_control.cond_hint.clone() if negative_control.cond_hint is not None else None + + xt_positive = self.model.inner_model.conds.get('xt_positive', self.model.inner_model.conds.get('positive')) + positive_control = xt_positive[0].get('control') + + if positive_control is not None and hasattr(positive_control, 'cond_hint_original'): + positive_cond_hint_init = positive_control.cond_hint.clone() if positive_control.cond_hint is not None else None + if positive_control.cond_hint_original.shape[-1] != x.shape[-2] * self.latent_compression_ratio or positive_control.cond_hint_original.shape[-2] != x.shape[-1] * self.latent_compression_ratio: + positive_control_pretile = comfy.utils.common_upscale(positive_control.cond_hint_original.clone().to(torch.float16).to('cuda'), x.shape[-1] * self.latent_compression_ratio, x.shape[-2] * self.latent_compression_ratio, "bislerp", "disabled") + positive_control.cond_hint_original = positive_control_pretile.to(positive_control.cond_hint_original) + positive_control_pretile = positive_control.cond_hint_original.clone().to(torch.float16).to('cuda') + control_tiles, control_orig_shape, control_grid, control_strides = tile_latent(positive_control_pretile, tile_size=(tile_h_full,tile_w_full)) + control_tiles = control_tiles + + denoised_tiles = [] + + tiles, orig_shape, grid, strides = tile_latent(x, tile_size=(tile_h,tile_w)) + + if y0_style_pos is not None: + y0_style_pos_tiles, _, _, _ = tile_latent(y0_style_pos, tile_size=(tile_h,tile_w)) + if y0_style_neg is not None: + y0_style_neg_tiles, _, _, _ = tile_latent(y0_style_neg, tile_size=(tile_h,tile_w)) + + for i in range(tiles.shape[0]): + tile = tiles[i].unsqueeze(0) + self.extra_args['model_options']['transformer_options']['x_tmp'] = tile + if control_tiles is not None: + positive_control.cond_hint = control_tiles[i].unsqueeze(0).to(positive_control.cond_hint) + if negative_control is not None: + negative_control.cond_hint = control_tiles[i].unsqueeze(0).to(positive_control.cond_hint) + + if y0_style_pos is not None: + self.extra_args['model_options']['transformer_options']['y0_style_pos'] = y0_style_pos_tiles[i].unsqueeze(0) + if y0_style_neg is not None: + self.extra_args['model_options']['transformer_options']['y0_style_neg'] = y0_style_neg_tiles[i].unsqueeze(0) + + denoised_tile = self.model(tile, sigma * s_in, **extra_args) + # increment counters per tile + self.model_calls_total += 1 + self.model_calls_denoised += 1 + denoised_tiles.append(denoised_tile) + + denoised_tiles = torch.cat(denoised_tiles, dim=0) + + denoised = untile_latent(denoised_tiles, orig_shape, grid, strides) + + else: + denoised = self.model(x, sigma * s_in, **extra_args) + # increment counters (single call path) + self.model_calls_total += 1 + self.model_calls_denoised += 1 + + self._restore_peripherals() + + if control_tiles is not None: + positive_control.cond_hint = positive_cond_hint_init + if negative_control is not None: + negative_control.cond_hint = negative_cond_hint_init + + if y0_style_pos is not None: + self.extra_args['model_options']['transformer_options']['y0_style_pos'] = y0_style_pos + if y0_style_neg is not None: + self.extra_args['model_options']['transformer_options']['y0_style_neg'] = y0_style_neg + + denoised = self.calc_cfg_channelwise(denoised) + return denoised + + def update_transformer_options(self, + transformer_options : Optional[dict] = None, + ): + + self.extra_args.setdefault("model_options", {}).setdefault("transformer_options", {}).update(transformer_options) + return + + # helper API to reset/read counters + def reset_model_call_counters(self) -> None: + self.model_calls_total = 0 + self.model_calls_denoised = 0 + self.model_calls_epsilon = 0 + + def get_model_call_counters(self) -> Dict[str, int]: + return { + "total": self.model_calls_total, + "denoised": self.model_calls_denoised, + "epsilon": self.model_calls_epsilon, + } + + def set_coeff(self, + rk_type : str, + h : Tensor, + c1 : float = 0.0, + c2 : float = 0.5, + c3 : float = 1.0, + step : int = 0, + sigmas : Optional[Tensor] = None, + sigma_down : Optional[Tensor] = None, + ) -> None: + + self.rk_type = rk_type + self.IMPLICIT = rk_type in get_implicit_sampler_name_list(nameOnly=True) + self.EXPONENTIAL = RK_Method_Beta.is_exponential(rk_type) + + sigma = sigmas[step] + sigma_next = sigmas[step+1] + + h_prev = [] + a, b, u, v, ci, multistep_stages, hybrid_stages, FSAL = get_rk_methods_beta(rk_type, + h, + c1, + c2, + c3, + h_prev, + step, + sigmas, + sigma, + sigma_next, + sigma_down, + self.extra_options, + ) + + self.multistep_stages = multistep_stages + self.hybrid_stages = hybrid_stages + + self.A = torch.tensor(a, dtype=h.dtype, device=h.device) + self.B = torch.tensor(b, dtype=h.dtype, device=h.device) + self.C = torch.tensor(ci, dtype=h.dtype, device=h.device) + + self.U = torch.tensor(u, dtype=h.dtype, device=h.device) if u is not None else None + self.V = torch.tensor(v, dtype=h.dtype, device=h.device) if v is not None else None + + self.rows = self.A.shape[0] + self.cols = self.A.shape[1] + + self.row_offset = 1 if not self.IMPLICIT and self.A[0].sum() == 0 else 0 + + if self.IMPLICIT and self.reorder_tableau_indices[0] != -1: + self.reorder_tableau(self.reorder_tableau_indices) + + + + def reorder_tableau(self, indices:list[int]) -> None: + #if indices[0]: + self.A = self.A [indices] + self.B[0] = self.B[0][indices] + self.C = self.C [indices] + self.C = torch.cat((self.C, self.C[-1:])) + return + + + + def update_substep(self, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + eps_prev_ : Tensor, + row : int, + row_offset : int, + h_new : Tensor, + h_new_orig : Tensor, + lying_eps_row_factor : float = 1.0, + sigma : Optional[Tensor] = None, + ) -> Tensor: + + if row < self.rows - row_offset and self.multistep_stages == 0: + row_tmp_offset = row + row_offset + + else: + row_tmp_offset = row + 1 + + #zr_base = self.zum(row+row_offset+self.multistep_stages, eps_, eps_prev_) # TODO: why unused? + + if self.SYNC_SUBSTEP_MEAN_CW and lying_eps_row_factor != 1.0: + zr_orig = self.zum(row+row_offset+self.multistep_stages, eps_, eps_prev_) + x_orig_row = x_0 + h_new * zr_orig + + #eps_row = eps_ [row].clone() + #eps_prev_row = eps_prev_[row].clone() + + eps_ [row] *= lying_eps_row_factor + eps_prev_[row] *= lying_eps_row_factor + + if self.EO("exp2lin_override"): + zr = self.zum2(row+row_offset+self.multistep_stages, eps_, eps_prev_, h_new, sigma) + x_[row_tmp_offset] = x_0 + zr + else: + zr = self.zum(row+row_offset+self.multistep_stages, eps_, eps_prev_) + + x_[row_tmp_offset] = x_0 + h_new * zr + + if self.SYNC_SUBSTEP_MEAN_CW and lying_eps_row_factor != 1.0: + x_[row_tmp_offset] = x_[row_tmp_offset] - x_[row_tmp_offset].mean(dim=(-2,-1), keepdim=True) + x_orig_row.mean(dim=(-2,-1), keepdim=True) + + #eps_ [row] = eps_row + #eps_prev_[row] = eps_prev_row + + if (self.SYNC_SUBSTEP_MEAN_CW and h_new != h_new_orig) or self.EO("sync_mean_noise"): + if not self.EO("disable_sync_mean_noise"): + x_row_down = x_0 + h_new_orig * zr + x_[row_tmp_offset] = x_[row_tmp_offset] - x_[row_tmp_offset].mean(dim=(-2,-1), keepdim=True) + x_row_down.mean(dim=(-2,-1), keepdim=True) + + return x_ + + + + def zum2(self, row:int, k:Tensor, k_prev:Tensor=None, h_new:Tensor=None, sigma:Tensor=None) -> Tensor: + if row < self.rows: + return self.a_k_einsum2(row, k, h_new, sigma) + else: + row = row - self.rows + return self.b_k_einsum2(row, k, h_new, sigma) + + # einsum does not type-promote: the float64 tableau operands must be cast to the buffer dtype + def a_k_einsum2(self, row:int, k:Tensor, h:Tensor, sigma:Tensor) -> Tensor: + return torch.einsum('i,j,k,i... -> ...', self.A[row].to(k.dtype), h.unsqueeze(0).to(k.dtype), -sigma.unsqueeze(0).to(k.dtype), k[:self.cols]) + + def b_k_einsum2(self, row:int, k:Tensor, h:Tensor, sigma:Tensor) -> Tensor: + return torch.einsum('i,j,k,i... -> ...', self.B[row].to(k.dtype), h.unsqueeze(0).to(k.dtype), -sigma.unsqueeze(0).to(k.dtype), k[:self.cols]) + + + def a_k_einsum(self, row:int, k :Tensor) -> Tensor: + return torch.einsum('i, i... -> ...', self.A[row].to(k.dtype), k[:self.cols]) + + def b_k_einsum(self, row:int, k :Tensor) -> Tensor: + return torch.einsum('i, i... -> ...', self.B[row].to(k.dtype), k[:self.cols]) + + def u_k_einsum(self, row:int, k_prev:Tensor) -> Tensor: + return torch.einsum('i, i... -> ...', self.U[row].to(k_prev.dtype), k_prev[:self.cols]) if (self.U is not None and k_prev is not None) else 0 + + def v_k_einsum(self, row:int, k_prev:Tensor) -> Tensor: + return torch.einsum('i, i... -> ...', self.V[row].to(k_prev.dtype), k_prev[:self.cols]) if (self.V is not None and k_prev is not None) else 0 + + + + def zum(self, row:int, k:Tensor, k_prev:Tensor=None,) -> Tensor: + if row < self.rows: + return self.a_k_einsum(row, k) + self.u_k_einsum(row, k_prev) + else: + row = row - self.rows + return self.b_k_einsum(row, k) + self.v_k_einsum(row, k_prev) + + def zum_tableau(self, k:Tensor, k_prev:Tensor=None,) -> Tensor: + a_k_sum = torch.einsum('ij, j... -> i...', self.A.to(k.dtype), k[:self.cols]) + u_k_sum = torch.einsum('ij, j... -> i...', self.U.to(k.dtype), k_prev[:self.cols]) if (self.U is not None and k_prev is not None) else 0 + return a_k_sum + u_k_sum + + def get_x(self, data:Tensor, noise:Tensor, sigma:Tensor): + if self.VE_MODEL: + return data + sigma * noise + else: + return (self.sigma_max - sigma) * data + sigma * noise + + def init_cfg_channelwise(self, x:Tensor, cfg_cw:float=1.0, **extra_args) -> Dict[str, Any]: + self.uncond = [torch.full_like(x, 0.0)] + self.cfg_cw = cfg_cw + if cfg_cw != 1.0: + def post_cfg_function(args): + self.uncond[0] = args["uncond_denoised"] + return args["denoised"] + model_options = extra_args.get("model_options", {}).copy() + extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True) + return extra_args + + + def calc_cfg_channelwise(self, denoised:Tensor) -> Tensor: + if self.cfg_cw != 1.0: + avg = 0 + for b, c in itertools.product(range(denoised.shape[0]), range(denoised.shape[1])): + avg += torch.norm(denoised[b][c] - self.uncond[0][b][c]) + avg /= denoised.shape[1] + + for b, c in itertools.product(range(denoised.shape[0]), range(denoised.shape[1])): + ratio = torch.nan_to_num(torch.norm(denoised[b][c] - self.uncond[0][b][c]) / avg, 0) + denoised_new = self.uncond[0] + ratio * self.cfg_cw * (denoised - self.uncond[0]) + return denoised_new + else: + return denoised + + + @staticmethod + def calculate_res_2m_step( + x_0 : Tensor, + denoised_ : Tensor, + sigma_down : Tensor, + sigmas : Tensor, + step : int, + ) -> Tuple[Tensor, Tensor]: + + if denoised_[2].sum() == 0: + return None, None + + sigma = sigmas[step] + sigma_prev = sigmas[step-1] + + h_prev = -torch.log(sigma/sigma_prev) + h = -torch.log(sigma_down/sigma) + + c1 = 0 + c2 = (-h_prev / h).item() + + ci = [c1,c2] + φ = Phi(h, ci, analytic_solution=True) + + b2 = φ(2)/c2 + b1 = φ(1) - b2 + + eps_2 = denoised_[1] - x_0 + eps_1 = denoised_[0] - x_0 + + h_a_k_sum = h * (b1 * eps_1 + b2 * eps_2) + + x = torch.exp(-h) * x_0 + h_a_k_sum + + denoised = x_0 + (sigma / (sigma - sigma_down)) * h_a_k_sum + + return x, denoised + + + @staticmethod + def calculate_res_3m_step( + x_0 : Tensor, + denoised_ : Tensor, + sigma_down : Tensor, + sigmas : Tensor, + step : int, + ) -> Tuple[Tensor, Tensor]: + + if denoised_[3].sum() == 0: + return None, None + + sigma = sigmas[step] + sigma_prev = sigmas[step-1] + sigma_prev2 = sigmas[step-2] + + h = -torch.log(sigma_down/sigma) + h_prev = -torch.log(sigma/sigma_prev) + h_prev2 = -torch.log(sigma/sigma_prev2) + + c1 = 0 + c2 = (-h_prev / h).item() + c3 = (-h_prev2 / h).item() + + ci = [c1,c2,c3] + φ = Phi(h, ci, analytic_solution=True) + + gamma = (3*(c3**3) - 2*c3) / (c2*(2 - 3*c2)) + + b3 = (1 / (gamma * c2 + c3)) * φ(2, -h) + b2 = gamma * b3 + b1 = φ(1, -h) - b2 - b3 + + eps_3 = denoised_[2] - x_0 + eps_2 = denoised_[1] - x_0 + eps_1 = denoised_[0] - x_0 + + h_a_k_sum = h * (b1 * eps_1 + b2 * eps_2 + b3 * eps_3) + + x = torch.exp(-h) * x_0 + h_a_k_sum + + denoised = x_0 + (sigma / (sigma - sigma_down)) * h_a_k_sum + + return x, denoised + + def swap_rk_type_at_step_or_threshold(self, + x_0 : Tensor, + data_prev_ : Tensor, + NS, + sigmas : Tensor, + step : int, + step_sched : int, + rk_swaps : list, + ): + if not rk_swaps: + return self.rk_type, False + + swap_at = {swap['step']: swap for swap in rk_swaps} + + target_swap = swap_at.get(step_sched) + + if target_swap is None: + for swap in sorted(rk_swaps, key=lambda s: s['step']): + if step_sched < swap['step'] and swap['threshold'] > 0: + threshold = swap['threshold'] + if x_0 is not None and step > 2 and sigmas[step+1] > 0 and self.rk_type != swap['type']: + x_res_2m, denoised_res_2m = self.calculate_res_2m_step(x_0, data_prev_, NS.sigma_down, sigmas, step) + x_res_3m, denoised_res_3m = self.calculate_res_3m_step(x_0, data_prev_, NS.sigma_down, sigmas, step) + if denoised_res_2m is not None: + if swap['print']: + RESplain("res_3m - res_2m:", torch.norm(denoised_res_3m - denoised_res_2m).item()) + if threshold > torch.norm(denoised_res_2m - denoised_res_3m): + target_swap = swap + break + + if target_swap is None: + return self.rk_type, False + + swap_type = target_swap['type'] + if swap_type == "": + swap_type = "res_3m" if self.EXPONENTIAL else "deis_3m" + + if self.rk_type == swap_type: + return self.rk_type, False + + RESplain("Switching rk_type to:", swap_type, "at step:", step_sched) + + self.rk_type = swap_type + + if RK_Method_Beta.is_exponential(swap_type): + self.__class__ = RK_Method_Exponential + else: + self.__class__ = RK_Method_Linear + + if swap_type in get_implicit_sampler_name_list(nameOnly=True): + self.IMPLICIT = True + self.row_offset = 0 + NS.row_offset = 0 + else: + self.IMPLICIT = False + self.row_offset = 1 + NS.row_offset = 1 + NS.h_fn = self.h_fn + NS.t_fn = self.t_fn + NS.sigma_fn = self.sigma_fn + + return self.rk_type, True + + + def bong_iter(self, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + eps_prev_ : Tensor, + data_ : Tensor, + sigma : Tensor, + s_ : Tensor, + row : int, + row_offset: int, + h : Tensor, + step : int, + step_sched: int, + BONGMATH_Y : bool = False, + y0_bongflow : Optional[Tensor] = None, + noise_sync: Optional[Tensor] = None, + eps_x_ : Optional[Tensor] = None, + eps_y_ : Optional[Tensor] = None, + #eps_x2y_ : Optional[Tensor] = None, + data_x_ : Optional[Tensor] = None, + data_y_ : Optional[Tensor] = None, + #yt_ : Optional[Tensor] = None, + #yt_0 : Optional[Tensor] = None, + LG = None, + ) -> Tuple[Tensor, Tensor, Tensor]: + + if x_0.ndim == 4: + norm_dim = (-2,-1) + elif x_0.ndim == 5: + norm_dim = (-4,-2,-1) + + if BONGMATH_Y: + lgw_mask_, lgw_mask_inv_ = LG.get_masks_for_step(step_sched) + lgw_mask_sync_, lgw_mask_sync_inv_ = LG.get_masks_for_step(step_sched, lgw_type="sync") + + weight_mask = lgw_mask_+lgw_mask_inv_ + if LG.SYNC_SEPARATE: + sync_mask = lgw_mask_sync_+lgw_mask_sync_inv_ + else: + sync_mask = 1. + + + if self.EO("bong_start_step", 0) > step or step > self.EO("bong_stop_step", 10000) or (self.unsample_bongmath == False and s_[-1] > s_[0]): + return x_0, x_, eps_ + + bong_iter_max_row = self.rows - row_offset + if self.EO("bong_iter_max_row_full"): + bong_iter_max_row = self.rows + + if self.EO("bong_iter_lock_x_0_ch_means"): + x_0_ch_means = x_0.mean(dim=norm_dim, keepdim=True) + + if self.EO("bong_iter_lock_x_row_ch_means"): + x_row_means = [] + for rr in range(row+row_offset): + x_row_mean = x_[rr].mean(dim=norm_dim, keepdim=True) + x_row_means.append(x_row_mean) + + if row < bong_iter_max_row and self.multistep_stages == 0: + bong_strength = self.EO("bong_strength", 1.0) + + if bong_strength != 1.0: + x_0_tmp = x_0 .clone() + x_tmp_ = x_ .clone() + eps_tmp_ = eps_.clone() + + for i in range(100): #bongmath for eps_prev_ not implemented? + x_0 = x_[row+row_offset] - h * self.zum(row+row_offset, eps_, eps_prev_) + + if self.EO("bong_iter_lock_x_0_ch_means"): + x_0 = x_0 - x_0.mean(dim=norm_dim, keepdim=True) + x_0_ch_means + + for rr in range(row+row_offset): + x_[rr] = x_0 + h * self.zum(rr, eps_, eps_prev_) + + if self.EO("bong_iter_lock_x_row_ch_means"): + for rr in range(row+row_offset): + x_[rr] = x_[rr] - x_[rr].mean(dim=norm_dim, keepdim=True) + x_row_means[rr] + + for rr in range(row+row_offset): + if self.EO("zonkytar"): + #eps_[rr] = self.get_unsample_epsilon(x_[rr], x_0, data_[rr], sigma, s_[rr]) + eps_[rr] = self.get_epsilon(x_[rr], x_0, data_[rr], sigma, s_[rr]) + else: + if BONGMATH_Y and not self.EO("disable_bongmath_y"): + if self.EXPONENTIAL: + eps_x_ = data_x_ - x_0 + eps_x2y_ = data_y_ - x_0 + if self.VE_MODEL: + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (-eps_y_+sigma*(-noise_sync)) + if self.EO("sync_x2y"): + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (-eps_x2y_+sigma*(-noise_sync)) + else: + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (-eps_y_+sigma*(y0_bongflow-noise_sync)) + if self.EO("sync_x2y"): + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (-eps_x2y_+sigma*(y0_bongflow-noise_sync)) + else: + eps_x_ [:s_.shape[0]] = (x_[:s_.shape[0]] - data_x_[:s_.shape[0]]) / s_.view(-1, *[1]*(x_.ndim-1)) + eps_x2y_ = torch.zeros_like(eps_x_) + eps_x2y_[:s_.shape[0]] = (x_[:s_.shape[0]] - data_y_[:s_.shape[0]]) / s_.view(-1, *[1]*(x_.ndim-1)) + + if self.VE_MODEL: + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (noise_sync-eps_y_) + if self.EO("sync_x2y"): + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (noise_sync-eps_x2y_) + else: + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (noise_sync-eps_y_-y0_bongflow) + if self.EO("sync_x2y"): + eps_ = sync_mask * eps_x_ + (1-sync_mask) * eps_x2y_ + weight_mask * (noise_sync-eps_x2y_-y0_bongflow) + + else: + eps_[rr] = self.get_epsilon(x_0, x_[rr], data_[rr], sigma, s_[rr]) + + if bong_strength != 1.0: + x_0 = x_0_tmp + bong_strength * (x_0 - x_0_tmp) + x_ = x_tmp_ + bong_strength * (x_ - x_tmp_) + eps_ = eps_tmp_ + bong_strength * (eps_ - eps_tmp_) + + return x_0, x_, eps_ #, yt_0, yt_ + + + def newton_iter(self, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + eps_prev_ : Tensor, + data_ : Tensor, + s_ : Tensor, + row : int, + h : Tensor, + sigmas : Tensor, + step : int, + newton_name: str, + SYNC_GUIDE_ACTIVE: bool, + ) -> Tuple[Tensor, Tensor]: + if SYNC_GUIDE_ACTIVE: + return x_, eps_ + newton_iter_name = "newton_iter_" + newton_name + + default_anchor_x_all = False + if newton_name == "lying": + default_anchor_x_all = True + + newton_iter = self.EO(newton_iter_name, 100) + newton_iter_skip_last_steps = self.EO(newton_iter_name + "_skip_last_steps", 0) + newton_iter_mixing_rate = self.EO(newton_iter_name + "_mixing_rate", 1.0) + + newton_iter_anchor = self.EO(newton_iter_name + "_anchor", 0) + newton_iter_anchor_x_all = self.EO(newton_iter_name + "_anchor_x_all", default_anchor_x_all) + newton_iter_type = self.EO(newton_iter_name + "_type", "from_epsilon") + newton_iter_sequence = self.EO(newton_iter_name + "_sequence", "double") + + row_b_offset = 0 + if self.EO(newton_iter_name + "_include_row_b"): + row_b_offset = 1 + + if step >= len(sigmas)-1-newton_iter_skip_last_steps or sigmas[step+1] == 0 or not self.IMPLICIT: + return x_, eps_ + + sigma = sigmas[step] + + start, stop = 0, self.rows+row_b_offset + if newton_name == "pre": + start = row + elif newton_name == "post": + start = row + 1 + + if newton_iter_anchor >= 0: + eps_anchor = eps_[newton_iter_anchor].clone() + + if newton_iter_anchor_x_all: + x_orig_ = x_.clone() + + for n_iter in range(newton_iter): + for r in range(start, stop): + if newton_iter_anchor >= 0: + eps_[newton_iter_anchor] = eps_anchor.clone() + if newton_iter_anchor_x_all: + x_ = x_orig_.clone() + x_tmp, eps_tmp = x_[r].clone(), eps_[r].clone() + + seq_start, seq_stop = r, r+1 + + if newton_iter_sequence == "double": + seq_start, seq_stop = start, stop + + for r_ in range(seq_start, seq_stop): + x_[r_] = x_0 + h * self.zum(r_, eps_, eps_prev_) + + for r_ in range(seq_start, seq_stop): + if newton_iter_type == "from_data": + data_[r_] = get_data_from_step(x_0, x_[r_], sigma, s_[r_]) + eps_ [r_] = self.get_epsilon(x_0, x_[r_], data_[r_], sigma, s_[r_]) + elif newton_iter_type == "from_step": + eps_ [r_] = get_epsilon_from_step(x_0, x_[r_], sigma, s_[r_]) + elif newton_iter_type == "from_alt": + eps_ [r_] = x_0/sigma - x_[r_]/s_[r_] + elif newton_iter_type == "from_epsilon": + eps_ [r_] = self.get_epsilon(x_0, x_[r_], data_[r_], sigma, s_[r_]) + + if self.EO(newton_iter_name + "_opt"): + opt_timing, opt_type, opt_subtype = self.EO(newton_iter_name+"_opt", [str]) + + opt_start, opt_stop = 0, self.rows+row_b_offset + if opt_timing == "early": + opt_stop = row + 1 + elif opt_timing == "late": + opt_start = row + 1 + + for r2 in range(opt_start, opt_stop): + if r_ != r2: + if opt_subtype == "a": + eps_a = eps_[r2] + eps_b = eps_[r_] + elif opt_subtype == "b": + eps_a = eps_[r_] + eps_b = eps_[r2] + + if opt_type == "ortho": + eps_ [r_] = get_orthogonal(eps_a, eps_b) + elif opt_type == "collin": + eps_ [r_] = get_collinear (eps_a, eps_b) + elif opt_type == "proj": + eps_ [r_] = get_collinear (eps_a, eps_b) + get_orthogonal(eps_b, eps_a) + + x_ [r_] = x_tmp + newton_iter_mixing_rate * (x_ [r_] - x_tmp) + eps_[r_] = eps_tmp + newton_iter_mixing_rate * (eps_[r_] - eps_tmp) + + if newton_iter_sequence == "double": + break + + return x_, eps_ + + + + +class RK_Method_Exponential(RK_Method_Beta): + def __init__(self, + model, + rk_type : str, + VE_MODEL : bool, + noise_anchor : float, + noise_boost_normalize : bool, + + model_device : str = 'cuda', + work_device : str = 'cpu', + dtype : torch.dtype = torch.float64, + extra_options : str = "", + ): + + super().__init__(model, + rk_type, + VE_MODEL, + noise_anchor, + noise_boost_normalize, + model_device = model_device, + work_device = work_device, + dtype = dtype, + extra_options = extra_options, + ) + + @staticmethod + def alpha_fn(neg_h:Tensor) -> Tensor: + return torch.exp(neg_h) + + @staticmethod + def sigma_fn(t:Tensor) -> Tensor: + #return 1/(torch.exp(-t)+1) + return t.neg().exp() + + @staticmethod + def t_fn(sigma:Tensor) -> Tensor: + #return -torch.log((1.-sigma)/sigma) + return sigma.log().neg() + + @staticmethod + def h_fn(sigma_down:Tensor, sigma:Tensor) -> Tensor: + #return (-torch.log((1.-sigma_down)/sigma_down)) - (-torch.log((1.-sigma)/sigma)) + return -torch.log(sigma_down/sigma) + + def __call__(self, + x : Tensor, + sub_sigma : Tensor, + x_0 : Optional[Tensor] = None, + sigma : Optional[Tensor] = None, + transformer_options : Optional[dict] = None, + ) -> Tuple[Tensor, Tensor]: + + x_0 = x if x_0 is None else x_0 + sigma = sub_sigma if sigma is None else sigma + + if transformer_options is not None: + self.extra_args.setdefault("model_options", {}).setdefault("transformer_options", {}).update(transformer_options) + + denoised = self.model_denoised(x.to(self.model_device), sub_sigma.to(self.model_device), **self.extra_args).to(sigma.device) + + eps_anchored = (x_0 - denoised) / sigma + eps_unmoored = (x - denoised) / sub_sigma + + eps = eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchored - eps_unmoored) + + denoised = x_0 - sigma * eps + + epsilon = denoised - x_0 + + #epsilon = denoised - x + + if self.EO("exp2lin_override"): + epsilon = (x_0 - denoised) / sigma + + return epsilon, denoised + + def get_eps(self, *args): + if len(args) == 3: + x, denoised, sigma = args + return denoised - x + elif len(args) == 5: + x_0, x, denoised, sigma, sub_sigma = args + eps_anchored = (x_0 - denoised) / sigma + eps_unmoored = (x - denoised) / sub_sigma + eps = eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchored - eps_unmoored) + denoised = x_0 - sigma * eps + eps_out = denoised - x_0 + if self.EO("exp2lin_override"): + eps_out = (x_0 - denoised) / sigma + return eps_out + + else: + raise ValueError(f"get_eps expected 3 or 5 arguments, got {len(args)}") + + def get_epsilon(self, + x_0 : Tensor, + x : Tensor, + denoised : Tensor, + sigma : Tensor, + sub_sigma : Tensor, + ) -> Tensor: + + eps_anchored = (x_0 - denoised) / sigma + eps_unmoored = (x - denoised) / sub_sigma + + eps = eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchored - eps_unmoored) + + denoised = x_0 - sigma * eps + if self.EO("exp2lin_override"): + return (x_0 - denoised) / sigma + else: + return denoised - x_0 + + + + def get_epsilon_anchored(self, x_0:Tensor, denoised:Tensor, sigma:Tensor) -> Tensor: + return denoised - x_0 + + + + def get_guide_epsilon(self, + x_0 : Tensor, + x : Tensor, + y : Tensor, + sigma : Tensor, + sigma_cur : Tensor, + sigma_down : Optional[Tensor] = None, + epsilon_scale : Optional[Tensor] = None, + ) -> Tensor: + + sigma_cur = epsilon_scale if epsilon_scale is not None else sigma_cur + + if sigma_down > sigma: + eps_unmoored = (sigma_cur/(self.sigma_max - sigma_cur)) * (x - y) + else: + eps_unmoored = y - x + + if self.EO("manually_anchor_unsampler"): + if sigma_down > sigma: + eps_anchored = (sigma /(self.sigma_max - sigma)) * (x_0 - y) + else: + eps_anchored = y - x_0 + eps_guide = eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchored - eps_unmoored) + else: + eps_guide = eps_unmoored + + return eps_guide + + + +class RK_Method_Linear(RK_Method_Beta): + def __init__(self, + model, + rk_type : str, + VE_MODEL : bool, + noise_anchor : float, + noise_boost_normalize : bool, + model_device : str = 'cuda', + work_device : str = 'cpu', + dtype : torch.dtype = torch.float64, + extra_options : str = "", + ): + + super().__init__(model, + rk_type, + VE_MODEL, + noise_anchor, + noise_boost_normalize, + model_device = model_device, + work_device = work_device, + dtype = dtype, + extra_options = extra_options, + ) + + @staticmethod + def alpha_fn(neg_h:Tensor) -> Tensor: + return torch.ones_like(neg_h) + + @staticmethod + def sigma_fn(t:Tensor) -> Tensor: + return t + + @staticmethod + def t_fn(sigma:Tensor) -> Tensor: + return sigma + + @staticmethod + def h_fn(sigma_down:Tensor, sigma:Tensor) -> Tensor: + return sigma_down - sigma + + def __call__(self, + x : Tensor, + sub_sigma : Tensor, + x_0 : Optional[Tensor] = None, + sigma : Optional[Tensor] = None, + transformer_options : Optional[dict] = None, + ) -> Tuple[Tensor, Tensor]: + + x_0 = x if x_0 is None else x_0 + sigma = sub_sigma if sigma is None else sigma + + if transformer_options is not None: + self.extra_args.setdefault("model_options", {}).setdefault("transformer_options", {}).update(transformer_options) + + denoised = self.model_denoised(x.to(self.model_device), sub_sigma.to(self.model_device), **self.extra_args).to(sigma.device) + + epsilon_anchor = (x_0 - denoised) / sigma + epsilon_unmoored = (x - denoised) / sub_sigma + + epsilon = epsilon_unmoored + self.LINEAR_ANCHOR_X_0 * (epsilon_anchor - epsilon_unmoored) + + return epsilon, denoised + + def get_eps(self, *args): + if len(args) == 3: + x, denoised, sigma = args + return (x - denoised) / sigma + elif len(args) == 5: + x_0, x, denoised, sigma, sub_sigma = args + eps_anchor = (x_0 - denoised) / sigma + eps_unmoored = (x - denoised) / sub_sigma + return eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchor - eps_unmoored) + else: + raise ValueError(f"get_eps expected 3 or 5 arguments, got {len(args)}") + + def get_epsilon(self, + x_0 : Tensor, + x : Tensor, + denoised : Tensor, + sigma : Tensor, + sub_sigma : Tensor, + ) -> Tensor: + + eps_anchor = (x_0 - denoised) / sigma + eps_unmoored = (x - denoised) / sub_sigma + + return eps_unmoored + self.LINEAR_ANCHOR_X_0 * (eps_anchor - eps_unmoored) + + + + def get_epsilon_anchored(self, x_0:Tensor, denoised:Tensor, sigma:Tensor) -> Tensor: + return (x_0 - denoised) / sigma + + + + def get_guide_epsilon(self, + x_0 : Tensor, + x : Tensor, + y : Tensor, + sigma : Tensor, + sigma_cur : Tensor, + sigma_down : Optional[Tensor] = None, + epsilon_scale : Optional[Tensor] = None, + ) -> Tensor: + + if sigma_down > sigma: + sigma_ratio = self.sigma_max - sigma_cur.clone() + else: + sigma_ratio = sigma_cur.clone() + sigma_ratio = epsilon_scale if epsilon_scale is not None else sigma_ratio + + if sigma_down is None: + return (x - y) / sigma_ratio + else: + if sigma_down > sigma: + return (y - x) / sigma_ratio + else: + return (x - y) / sigma_ratio + + + + +""" + + + + +if EO("bong2m") and RK.multistep_stages > 0 and step < len(sigmas)-4: + h_no_eta = -torch.log(sigmas[step+1]/sigmas[step]) + h_prev1_no_eta = -torch.log(sigmas[step] /sigmas[step-1]) + c2_prev = (-h_prev1_no_eta / h_no_eta).item() + eps_prev = denoised_data_prev - x_0 + + φ = Phi(h_prev, [0.,c2_prev]) + a2_1 = c2_prev * φ(1,2) + for i in range(100): + x_prev = x_0 - h_prev * (a2_1 * eps_prev) + eps_prev = denoised_data_prev - x_prev + + eps_[1] = eps_prev + +if EO("bong3m") and RK.multistep_stages > 0 and step < len(sigmas)-10: + h_no_eta = -torch.log(sigmas[step+1]/sigmas[step]) + h_prev1_no_eta = -torch.log(sigmas[step] /sigmas[step-1]) + h_prev2_no_eta = -torch.log(sigmas[step] /sigmas[step-2]) + c2_prev = (-h_prev1_no_eta / h_no_eta).item() + c3_prev = (-h_prev2_no_eta / h_no_eta).item() + + eps_prev2 = denoised_data_prev2 - x_0 + eps_prev = denoised_data_prev - x_0 + + φ = Phi(h_prev1_no_eta, [0.,c2_prev, c3_prev]) + a2_1 = c2_prev * φ(1,2) + for i in range(100): + x_prev = x_0 - h_prev1_no_eta * (a2_1 * eps_prev) + eps_prev = denoised_data_prev2 - x_prev + + eps_[1] = eps_prev + + φ = Phi(h_prev2_no_eta, [0.,c3_prev, c3_prev]) + + def calculate_gamma(c2_prev, c3_prev): + return (3*(c3_prev**3) - 2*c3_prev) / (c2_prev*(2 - 3*c2_prev)) + gamma = calculate_gamma(c2_prev, c3_prev) + + a2_1 = c2_prev * φ(1,2) + a3_2 = gamma * c2_prev * φ(2,2) + (c3_prev ** 2 / c2_prev) * φ(2, 3) + a3_1 = c3_prev * φ(1,3) - a3_2 + + for i in range(100): + x_prev2 = x_0 - h_prev2_no_eta * (a3_1 * eps_prev + a3_2 * eps_prev2) + x_prev = x_prev2 + h_prev2_no_eta * (a2_1 * eps_prev) + + eps_prev2 = denoised_data_prev - x_prev2 + eps_prev = denoised_data_prev2 - x_prev + + eps_[2] = eps_prev2 +""" \ No newline at end of file diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/rk_noise_sampler_beta.py b/simple_syrup/third_party/res4lyf_runtime/beta/rk_noise_sampler_beta.py new file mode 100644 index 0000000..db38ca9 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/rk_noise_sampler_beta.py @@ -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) + + + + diff --git a/simple_syrup/third_party/res4lyf_runtime/beta/rk_sampler_beta.py b/simple_syrup/third_party/res4lyf_runtime/beta/rk_sampler_beta.py new file mode 100644 index 0000000..3a6f68e --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/beta/rk_sampler_beta.py @@ -0,0 +1,2407 @@ +import torch +from torch import Tensor +import torch.nn.functional as F +from tqdm.auto import trange +import gc +from typing import Optional, Callable, Tuple, List, Dict, Any, Union +import math +import copy + +from comfy.model_sampling import EPS +import comfy + +from ..res4lyf import RESplain +from ..helper import ExtraOptions, FrameWeightsManager +from ..latents import lagrange_interpolation, get_collinear, get_orthogonal, get_cosine_similarity, get_pearson_similarity, \ + get_slerp_weight_for_cossim, get_slerp_ratio, slerp_tensor, get_edge_mask, normalize_zscore, \ + compute_slerp_ratio_for_target, find_slerp_ratio_grid, \ + is_packed_latent, get_latent, apply_per_step_latent_normalization, LatentHandler, \ + derive_old_latent_shapes, extract_video_tail, extend_state_info_tensors +from ..style_transfer import apply_scattersort_spatial, apply_adain_spatial + +from .rk_method_beta import RK_Method_Beta +from .rk_noise_sampler_beta import RK_NoiseSampler +from .rk_guide_func_beta import LatentGuide +from .phi_functions import Phi +from .constants import MAX_STEPS, GUIDE_MODE_NAMES_PSEUDOIMPLICIT + +_VRAM_CAP_SET = False + +def init_implicit_sampling( + RK : RK_Method_Beta, + x_0 : Tensor, + x_ : Tensor, + eps_ : Tensor, + eps_prev_ : Tensor, + data_ : Tensor, + eps : Tensor, + denoised : Tensor, + denoised_prev2 : Tensor, + step : int, + sigmas : Tensor, + h : Tensor, + s_ : Tensor, + EO : ExtraOptions, + SYNC_GUIDE_ACTIVE, + ): + + sigma = sigmas[step] + if EO("implicit_skip_model_call_at_start") and denoised.sum() + eps.sum() != 0: + if denoised_prev2.sum() == 0: + eps_ [0] = eps.clone() + data_[0] = denoised.clone() + eps_ [0] = RK.get_epsilon_anchored(x_0, denoised, sigma) + else: + sratio = sigma - s_[0] + data_[0] = denoised + sratio * (denoised - denoised_prev2) + + elif EO("implicit_full_skip_model_call_at_start") and denoised.sum() + eps.sum() != 0: + if denoised_prev2.sum() == 0: + eps_ [0] = eps.clone() + data_[0] = denoised.clone() + eps_ [0] = RK.get_epsilon_anchored(x_0, denoised, sigma) + else: + for r in range(RK.rows): + sratio = sigma - s_[r] + data_[r] = denoised + sratio * (denoised - denoised_prev2) + eps_ [r] = RK.get_epsilon_anchored(x_0, data_[r], s_[r]) + + elif EO("implicit_lagrange_skip_model_call_at_start") and denoised.sum() + eps.sum() != 0: + if denoised_prev2.sum() == 0: + eps_ [0] = eps.clone() + data_[0] = denoised.clone() + eps_ [0] = RK.get_epsilon_anchored(x_0, denoised, sigma) + else: + sigma_prev = sigmas[step-1] + h_prev = sigma - sigma_prev + w = h / h_prev + substeps_prev = len(RK.C[:-1]) + + for r in range(RK.rows): + sratio = sigma - s_[r] + data_[r] = lagrange_interpolation([0,1], [denoised_prev2, denoised], 1 + w*RK.C[r]).squeeze(0) + denoised_prev2 - denoised + eps_ [r] = RK.get_epsilon_anchored(x_0, data_[r], s_[r]) + + if EO("implicit_lagrange_skip_model_call_at_start_0_only"): + for r in range(RK.rows): + eps_ [r] = eps_ [0].clone() * s_[0] / s_[r] + data_[r] = denoised.clone() + + + elif EO("implicit_lagrange_init") and denoised.sum() + eps.sum() != 0: + sigma_prev = sigmas[step-1] + h_prev = sigma - sigma_prev + w = h / h_prev + substeps_prev = len(RK.C[:-1]) + + z_prev_ = eps_.clone() + for r in range (substeps_prev): + z_prev_[r] = h * RK.zum(r, eps_) # u,v not implemented for lagrange guess for implicit + zi_1 = lagrange_interpolation(RK.C[:-1], z_prev_[:substeps_prev], RK.C[0]).squeeze(0) # + x_prev - x_0""" + x_[0] = x_0 + zi_1 + + else: + + eps_[0], data_[0] = RK(x_[0], sigma, x_0, sigma) + + if not EO(("implicit_lagrange_init", "radaucycle", "implicit_full_skip_model_call_at_start", "implicit_lagrange_skip_model_call_at_start")): + for r in range(RK.rows): + eps_ [r] = eps_ [0].clone() * sigma / s_[r] + data_[r] = data_[0].clone() + + x_, eps_ = RK.newton_iter(x_0, x_, eps_, eps_prev_, data_, s_, 0, h, sigmas, step, "init", SYNC_GUIDE_ACTIVE) + return x_, eps_, data_ + + +@torch.no_grad() +def sample_rk_beta( + model, + x : Tensor, + sigmas : Tensor, + sigmas_override : Optional[Tensor] = None, + + extra_args : Optional[Tensor] = None, + callback : Optional[Callable] = None, + disable : bool = None, + + sampler_mode : str = "standard", + + rk_type : str = "res_2m", + implicit_sampler_name : str = "use_explicit", + + c1 : float = 0.0, + c2 : float = 0.5, + c3 : float = 1.0, + + noise_sampler_type : str = "gaussian", + noise_sampler_type_substep : str = "gaussian", + noise_mode_sde : str = "hard", + noise_mode_sde_substep : str = "hard", + + eta : float = 0.5, + eta_substep : float = 0.5, + + + + + noise_scaling_weight : float = 0.0, + noise_scaling_type : str = "sampler", + noise_scaling_mode : str = "linear", + noise_scaling_eta : float = 0.0, + noise_scaling_cycles : int = 1, + + noise_scaling_weights : Optional[Tensor] = None, + noise_scaling_etas : Optional[Tensor] = None, + + noise_boost_step : float = 0.0, + noise_boost_substep : float = 0.0, + noise_boost_normalize : bool = True, + noise_anchor : float = 1.0, + + s_noise : float = 1.0, + s_noise_substep : float = 1.0, + d_noise : float = 1.0, + d_noise_start_step : int = 0, + d_noise_inv : float = 1.0, + d_noise_inv_start_step : int = 0, + + + + alpha : float = -1.0, + alpha_substep : float = -1.0, + k : float = 1.0, + k_substep : float = 1.0, + + momentum : float = 0.0, + + + overshoot_mode : str = "hard", + overshoot_mode_substep : str = "hard", + overshoot : float = 0.0, + overshoot_substep : float = 0.0, + + implicit_type : str = "predictor-corrector", + implicit_type_substeps : str = "predictor-corrector", + + implicit_steps_diag : int = 0, + implicit_steps_full : int = 0, + + etas : Optional[Tensor] = None, + etas_substep : Optional[Tensor] = None, + s_noises : Optional[Tensor] = None, + s_noises_substep : Optional[Tensor] = None, + + momentums : Optional[Tensor] = None, + + regional_conditioning_weights : Optional[Tensor] = None, + regional_conditioning_floors : Optional[Tensor] = None, + narcissism_start_step : int = 0, + narcissism_end_step : int = 5, + + LGW_MASK_RESCALE_MIN : bool = True, + guides : Optional[Tuple[Any, ...]] = None, + epsilon_scales : Optional[Tensor] = None, + frame_weights_mgr : Optional[FrameWeightsManager] = None, + + sde_noise : list [Tensor] = [], + + noise_seed : int = -1, + noise_initial : Optional[Tensor] = None, + image_initial : Optional[Tensor] = None, + + cfgpp : float = 0.0, + cfg_cw : float = 1.0, + + BONGMATH : bool = True, + unsample_bongmath = None, + + state_info : Optional[dict[str, Any]] = None, + state_info_out : Optional[dict[str, Any]] = None, + + rk_swaps : list = [], + + steps_to_run : int = -1, + start_at_step : int = -1, + tile_sizes : Optional[List[Tuple[int,int]]] = None, + + flow_sync_eps : float = 0.0, + + sde_mask : Optional[Tensor] = None, + + batch_num : int = 0, + + extra_options : str = "", + + outer_sigmas_len : int = -1, + + latent_shapes : Optional[List[tuple]] = None, + latent_normalize_idx_0_steps : Optional[List[float]] = None, + latent_normalize_idx_1_steps : Optional[List[float]] = None, + + AttnMask = None, + RegContext = None, + RegParam = None, + + AttnMask_neg = None, + RegContext_neg = None, + RegParam_neg = None, + ): + + if sampler_mode == "NULL": + return x + + # Precision contract: + # default_dtype (float64): sigmas, schedules, tableau/phi coefficients, RK scalar math — + # the cancellation-sensitive operations. Also the default for noise_dtype. + # work_dtype (float32): x and all latent-sized buffers, and the model call boundary. + # work_dtype=float64 restores the classic full-float64 behavior. + # noise_dtype: the dtype noise is generated at — decides which noise realization a seed + # produces (torch's RNG stream differs per dtype), independent of the math precision. + EO = ExtraOptions(extra_options) + default_dtype = EO("default_dtype", torch.float64) + work_dtype = EO("work_dtype", torch.float32) + + REPORT_VRAM = EO("report_vram") and torch.cuda.is_available() + if REPORT_VRAM: + torch.cuda.reset_peak_memory_stats() + + # vram_cap_gb=N caps the CUDA allocator so OOM reproduces deterministically — headroom flags + # like --reserve-vram get absorbed by ComfyUI's weight paging instead of failing + global _VRAM_CAP_SET + vram_cap_gb = EO("vram_cap_gb", 0.0) + if torch.cuda.is_available(): + if vram_cap_gb > 0: + total_vram = torch.cuda.get_device_properties(0).total_memory + torch.cuda.set_per_process_memory_fraction(min(1.0, vram_cap_gb * 1024**3 / total_vram)) + RESplain(f"vram_cap_gb: capping allocator at {vram_cap_gb} GB of {total_vram / 1024**3:.1f} GB", debug=False) + _VRAM_CAP_SET = True + elif _VRAM_CAP_SET: + torch.cuda.set_per_process_memory_fraction(1.0) + _VRAM_CAP_SET = False + + extra_args = {} if extra_args is None else extra_args + model_device = model.inner_model.inner_model.device #x.device + work_device = 'cpu' if EO("work_device_cpu") else model_device + + state_info = {} if state_info is None else state_info + state_info_out = {} if state_info_out is None else state_info_out + + VE_MODEL = isinstance(model.inner_model.inner_model.model_sampling, EPS) + + RENOISE = False + if 'raw_x' in state_info and sampler_mode in {"resample", "unsample"}: + if x.shape == state_info['raw_x'].shape: + x = state_info['raw_x'].to(work_device) + RESplain("Continuing from raw latent from previous sampler.", debug=False) + else: + shapes_new = latent_shapes if latent_shapes is not None else [x.shape] + shapes_old = derive_old_latent_shapes(state_info['raw_x'], shapes_new) + can_extend_temporally = ( + shapes_old is not None + and shapes_new[0][-2:] == shapes_old[0][-2:] + and shapes_new[0][-3] > shapes_old[0][-3] + ) + + if can_extend_temporally: + extra_T = shapes_new[0][-3] - shapes_old[0][-3] + video_tail = extract_video_tail(x, shapes_new, extra_T) + state_info = extend_state_info_tensors(state_info, shapes_old, video_tail) + x = state_info['raw_x'].to(work_device) + RESplain(f"Continuing from raw latent with temporal extension (+{extra_T} video frames).", debug=False) + else: + x = (LatentHandler(x, latent_shapes) + .map_with(state_info['denoised'], lambda x_t, d_t: comfy.utils.common_upscale(d_t, x_t.shape[-1], x_t.shape[-2], "bislerp", "disabled").to(x_t)) + .tensor) + RENOISE = True + RESplain("Continuing from raw latent from previous sampler (spatial rescale).", debug=False) + + + + start_step = 0 + if 'end_step' in state_info and (sampler_mode == "resample" or sampler_mode == "unsample"): + + if state_info['completed'] != True and state_info['end_step'] != 0 and state_info['end_step'] != -1 and state_info['end_step'] < len(state_info['sigmas'])-1 : #incomplete run in previous sampler node + + if state_info['sampler_mode'] in {"standard","resample"} and sampler_mode == "unsample" and sigmas[2] < sigmas[1]: + sigmas = torch.flip(state_info['sigmas'], dims=[0]) + start_step = (len(sigmas)-1) - (state_info['end_step']) #-1) #removed -1 at the end here. correct? + + if state_info['sampler_mode'] == "unsample" and sampler_mode == "resample" and sigmas[2] > sigmas[1]: + sigmas = torch.flip(state_info['sigmas'], dims=[0]) + start_step = (len(sigmas)-1) - state_info['end_step'] #-1) + elif state_info['sampler_mode'] == "unsample" and sampler_mode == "resample": + start_step = 0 + + if state_info['sampler_mode'] in {"standard", "resample"} and sampler_mode == "resample": + start_step = state_info['end_step'] if state_info['end_step'] != -1 else 0 + if start_step > 0: + sigmas = state_info['sigmas'].clone() + + + + if sde_mask is not None: + from .rk_guide_func_beta import prepare_mask + sde_mask, _ = prepare_mask(get_latent(x, latent_shapes, 0), sde_mask, LGW_MASK_RESCALE_MIN) + sde_mask = sde_mask.to(x.device).to(x.dtype) + + + + x = x .to(dtype=work_dtype, device=work_device) + sigmas = sigmas.to(dtype=default_dtype, device=work_device) + + # sync sample_sigmas in model_options to the effective (unpadded) schedule. + if 'model_options' in extra_args: + transformer_options = extra_args['model_options'].setdefault('transformer_options', {}) + transformer_options['sample_sigmas'] = sigmas + + c1 = EO("c1" , c1) + c2 = EO("c2" , c2) + c3 = EO("c3" , c3) + + cfg_cw = EO("cfg_cw" , cfg_cw) + + noise_seed = EO("noise_seed" , noise_seed) + noise_seed_substep = EO("noise_seed_substep" , noise_seed + MAX_STEPS) + + pseudoimplicit_row_weights = EO("pseudoimplicit_row_weights" , [1. for _ in range(100)]) + pseudoimplicit_step_weights = EO("pseudoimplicit_step_weights", [1. for _ in range(max(implicit_steps_diag, implicit_steps_full)+1)]) + + noise_scaling_cycles = EO("noise_scaling_cycles", 1) + noise_boost_step = EO("noise_boost_step", 0.0) + noise_boost_substep = EO("noise_boost_substep", 0.0) + + # SETUP SAMPLER + if implicit_sampler_name not in ("use_explicit", "none"): + rk_type = implicit_sampler_name + RESplain("rk_type:", rk_type) + if implicit_sampler_name == "none": + implicit_steps_diag = implicit_steps_full = 0 + + RK = RK_Method_Beta.create(model, rk_type, VE_MODEL, noise_anchor, noise_boost_normalize, model_device=model_device, work_device=work_device, dtype=default_dtype, extra_options=extra_options) + RK.extra_args = RK.init_cfg_channelwise(x, cfg_cw, **extra_args) + RK.tile_sizes = tile_sizes + RK.reset_model_call_counters() + RK.extra_args['model_options']['transformer_options']['regional_conditioning_weight'] = 0.0 + RK.extra_args['model_options']['transformer_options']['regional_conditioning_floor'] = 0.0 + + RK.unsample_bongmath = BONGMATH if unsample_bongmath is None else unsample_bongmath # allow turning off bongmath for unsampling with cycles + + + # SETUP SIGMAS + sigmas_orig = sigmas.clone() + NS = RK_NoiseSampler(RK, model, device=work_device, dtype=default_dtype, extra_options=extra_options) + sigmas, UNSAMPLE = NS.prepare_sigmas(sigmas, sigmas_override, d_noise, d_noise_start_step, sampler_mode) + if UNSAMPLE and sigmas_orig[0] == 0.0 and sigmas_orig[0] != sigmas[0] and len(sigmas_orig) > 2 and sigmas_orig[1] < sigmas_orig[2]: + sigmas = torch.cat([torch.full_like(sigmas[0], 0.0).unsqueeze(0), sigmas]) + if start_step == 0: + start_step = 1 + else: + start_step -= 1 + + if sampler_mode in {"resample", "unsample"}: + start_step = resolve_start_step_from_sigma_next(sigmas, state_info.get('sigma_next', -1), start_step) + + start_step = start_at_step if start_at_step >= 0 else start_step + + + SDE_NOISE_EXTERNAL = False + if sde_noise is not None: + if len(sde_noise) > 0 and len(sigmas_orig) > 2 and sigmas_orig[1] > sigmas_orig[2]: + SDE_NOISE_EXTERNAL = True + sigma_up_total = torch.zeros_like(sigmas[0]) + for i in range(len(sde_noise)-1): + sigma_up_total += sigmas[i+1] + etas = torch.full_like(sigmas, eta / sigma_up_total) + + if 'last_rng' in state_info and sampler_mode in {"resample", "unsample"}: + last_rng = state_info['last_rng'].clone() + last_rng_substep = state_info['last_rng_substep'].clone() + else: + last_rng = None + last_rng_substep = None + + NS.init_noise_samplers(x, noise_seed, noise_seed_substep, noise_sampler_type, noise_sampler_type_substep, noise_mode_sde, noise_mode_sde_substep, \ + overshoot_mode, overshoot_mode_substep, noise_boost_step, noise_boost_substep, alpha, alpha_substep, k, k_substep, \ + last_rng=last_rng, last_rng_substep=last_rng_substep, latent_shapes=latent_shapes,) + + data_ = None + eps_ = None + eps = torch.zeros_like(x, dtype=work_dtype, device=work_device) + denoised = torch.zeros_like(x, dtype=work_dtype, device=work_device) + state_denoised = state_info.get('denoised') + if state_denoised is not None and state_denoised.shape == x.shape: + denoised_prev = state_denoised.to(dtype=work_dtype, device=work_device) + else: + denoised_prev = torch.zeros_like(x, dtype=work_dtype, device=work_device) + denoised_prev2 = torch.zeros_like(x, dtype=work_dtype, device=work_device) + x_ = None + eps_prev_ = None + denoised_data_prev = None + denoised_data_prev2 = None + h_prev = None + eps_y2x_ = None + eps_x2y_ = None + eps_y_ = None + eps_prev_y_ = None + data_y_ = None + yt_ = None + yt_0 = None + eps_yt_ = None + eps_x_ = None + data_y_ = None + data_x_ = None + z_ = None # for tracking residual noise for model scattersort/synchronized diffusion + + y0_bongflow = state_info.get('y0_bongflow') + y0_bongflow_orig = state_info.get('y0_bongflow_orig') + noise_bongflow = state_info.get('noise_bongflow') + y0_standard_guide = state_info.get('y0_standard_guide') + y0_inv_standard_guide = state_info.get('y0_inv_standard_guide') + + if y0_bongflow is not None: y0_bongflow = y0_bongflow.clone() + if y0_bongflow_orig is not None: y0_bongflow_orig = y0_bongflow_orig.clone() + if noise_bongflow is not None: noise_bongflow = noise_bongflow.clone() + if y0_standard_guide is not None: y0_standard_guide = y0_standard_guide.clone() + if y0_inv_standard_guide is not None: y0_inv_standard_guide = y0_inv_standard_guide.clone() + + data_prev_y_ = state_info.get('data_prev_y_') + data_prev_x_ = state_info.get('data_prev_x_') + data_prev_x2y_ = state_info.get('data_prev_x2y_') + if data_prev_y_ is not None: data_prev_y_ = data_prev_y_.clone() + if data_prev_x_ is not None: data_prev_x_ = data_prev_x_.clone() + if data_prev_x2y_ is not None: data_prev_x2y_ = data_prev_x2y_.clone() + + # BEGIN SAMPLING LOOP + try: + RESplain("Starting sampling loop. Model type: ", model.inner_model.model_patcher.model.diffusion_model._get_name(), debug=True) + except: + RESplain("Starting sampling loop.", debug=True) + + num_steps = len(sigmas[start_step:])-2 if sigmas[-1] == 0 else len(sigmas[start_step:])-1 + + if steps_to_run >= 0: + current_steps = min(num_steps, steps_to_run) + num_steps = start_step + min(num_steps, steps_to_run) + else: + current_steps = num_steps + num_steps = start_step + num_steps + #current_steps = current_steps + 1 if sigmas[-1] == 0 and steps_to_run < 0 and UNSAMPLE else current_steps + + INIT_SAMPLE_LOOP = True + step = start_step + sigma, sigma_next, data_prev_, x_0 = None, None, None, None + + if (num_steps-1) == len(sigmas)-2 and sigmas[-1] == 0 and sigmas[-2] == NS.sigma_min: + progress_bar = trange(current_steps+1, disable=disable) + else: + progress_bar = trange(current_steps, disable=disable) + + + # SETUP GUIDES + LG = LatentGuide(model, sigmas, UNSAMPLE, VE_MODEL, LGW_MASK_RESCALE_MIN, extra_options, device=work_device, dtype=work_dtype, frame_weights_mgr=frame_weights_mgr, latent_shapes=latent_shapes) + RK.latent_guide = LG + + guide_inversion_y0 = state_info.get('guide_inversion_y0') + guide_inversion_y0_inv = state_info.get('guide_inversion_y0_inv') + + x = LG.init_guides(x, RK.IMPLICIT, guides, NS.noise_sampler, batch_num, sigmas[step], guide_inversion_y0, guide_inversion_y0_inv) + LG.y0 = y0_standard_guide.clone() if y0_standard_guide is not None else LG.y0 + LG.y0_inv = y0_inv_standard_guide.clone() if y0_inv_standard_guide is not None else LG.y0_inv + if (LG.mask != 1.0).any() and ((LG.y0 == 0).all() or (LG.y0_inv == 0).all()) : # and not LG.guide_mode.startswith("flow"): # (LG.y0.sum() == 0 or LG.y0_inv.sum() == 0): + SKIP_PSEUDO = True + RESplain("skipping pseudo...") + if LG.y0 .sum() == 0: + SKIP_PSEUDO_Y = "y0" + elif LG.y0_inv.sum() == 0: + SKIP_PSEUDO_Y = "y0_inv" + else: + SKIP_PSEUDO = False + if guides is not None and guides.get('guide_mode', '') != "inversion" or sampler_mode != "unsample": #do not set denoised_prev to noise guide with inversion! + if LG.y0.sum() != 0 and LG.y0_inv.sum() != 0: + denoised_prev = LG.mask * LG.y0 + (1-LG.mask) * LG.y0_inv + elif LG.y0.sum() != 0: + denoised_prev = LG.y0 + elif LG.y0_inv.sum() != 0: + denoised_prev = LG.y0_inv + data_cached = None + + if EO("pseudo_mix_strength"): + orig_y0 = LG.y0.clone() + orig_y0_inv = LG.y0_inv.clone() + + #gc.collect() + BASE_STARTED = False + INV_STARTED = False + FLOW_STARTED = False + FLOW_STOPPED = False + noise_xt, noise_yt = None, None + FLOW_RESUMED = False + if state_info.get('FLOW_STARTED', False) and not state_info.get('FLOW_STOPPED', False): + FLOW_RESUMED = True + y0 = state_info['y0'].clone().to(work_device) + data_cached = state_info['data_cached'].clone().to(work_device) + data_x_prev_ = state_info['data_x_prev_'].clone().to(work_device) + + if noise_initial is not None: + x_init = noise_initial.to(x) + RK.update_transformer_options({'x_init': x_init}) + + #progress_bar = trange(len(sigmas)-1-start_step, disable=disable) + + #if EO("eps_adain") or EO("x_init_to_model"): + + if AttnMask is not None: + RK.update_transformer_options({'AttnMask' : AttnMask}) + RK.update_transformer_options({'RegContext': RegContext}) + + if AttnMask_neg is not None: + RK.update_transformer_options({'AttnMask_neg' : AttnMask_neg}) + RK.update_transformer_options({'RegContext_neg': RegContext_neg}) + + if EO("y0_to_transformer_options"): + RK.update_transformer_options({'y0': LG.y0.clone()}) + + if EO("y0_inv_to_transformer_options"): + RK.update_transformer_options({'y0_inv': LG.y0_inv.clone()}) + for block in model.inner_model.inner_model.diffusion_model.double_stream_blocks: + for attr in ["txt_q_cache", "txt_k_cache", "txt_v_cache", "img_q_cache", "img_k_cache", "img_v_cache"]: + if hasattr(block.block.attn1, attr): + delattr(block.block.attn1, attr) + + for block in model.inner_model.inner_model.diffusion_model.single_stream_blocks: + block.block.attn1.EO = EO + for attr in ["txt_q_cache", "txt_k_cache", "txt_v_cache", "img_q_cache", "img_k_cache", "img_v_cache"]: + if hasattr(block.block.attn1, attr): + delattr(block.block.attn1, attr) + + RK.update_transformer_options({'ExtraOptions': copy.deepcopy(EO)}) + if EO("update_cross_attn"): + update_cross_attn = { + 'src_llama_start': EO('src_llama_start', 0), + 'src_llama_end': EO('src_llama_end', 0), + 'src_t5_start': EO('src_t5_start', 0), + 'src_t5_end': EO('src_t5_end', 0), + + 'tgt_llama_start': EO('tgt_llama_start', 0), + 'tgt_llama_end': EO('tgt_llama_end', 0), + 'tgt_t5_start': EO('tgt_t5_start', 0), + 'tgt_t5_end': EO('tgt_t5_end', 0), + 'skip_cross_attn': EO('skip_cross_attn', False), + + 'update_q': EO('update_q', False), + 'update_k': EO('update_k', True), + 'update_v': EO('update_v', True), + + + 'lamb': EO('lamb', 0.01), + 'erase': EO('erase', 10.0), + } + RK.update_transformer_options({'update_cross_attn': update_cross_attn}) + else: + RK.update_transformer_options({'update_cross_attn': None}) + + if LG.HAS_LATENT_GUIDE_ADAIN: + RK.update_transformer_options({'blocks_adain_cache': []}) + if LG.HAS_LATENT_GUIDE_ATTNINJ: + RK.update_transformer_options({'blocks_attninj_cache': []}) + if LG.HAS_LATENT_GUIDE_STYLE_POS: + if LG.HAS_LATENT_GUIDE and y0_standard_guide is None: + y0_cache = LG.y0.clone().cpu() + RK.update_transformer_options({'y0_standard_guide': LG.y0}) + + sigmas_scheduled = sigmas.clone() # store for return in state_info_out + + if EO("sigma_restarts"): + sigma_restarts = 1 + EO("sigma_restarts", 0) + sigmas = sigmas[step:num_steps+1].repeat(sigma_restarts) + step = 0 + num_steps = 2 * sigma_restarts - 1 + + if RENOISE: # TODO: adapt for noise inversion somehow + if VE_MODEL: + x = x + sigmas[step] * NS.noise_sampler(sigma=sigmas[step], sigma_next=sigmas[step+1]) + else: + x = (1 - sigmas[step]) * x + sigmas[step] * NS.noise_sampler(sigma=sigmas[step], sigma_next=sigmas[step+1]) + LG.ADAIN_NOISE_MODE = "" + StyleMMDiT = None + if guides is not None: + RK.update_transformer_options({"freqsep_lowpass_method": guides.get("freqsep_lowpass_method")}) + RK.update_transformer_options({"freqsep_sigma": guides.get("freqsep_sigma")}) + RK.update_transformer_options({"freqsep_kernel_size": guides.get("freqsep_kernel_size")}) + RK.update_transformer_options({"freqsep_inner_kernel_size": guides.get("freqsep_inner_kernel_size")}) + RK.update_transformer_options({"freqsep_stride": guides.get("freqsep_stride")}) + + + RK.update_transformer_options({"freqsep_lowpass_weight": guides.get("freqsep_lowpass_weight")}) + RK.update_transformer_options({"freqsep_highpass_weight":guides.get("freqsep_highpass_weight")}) + RK.update_transformer_options({"freqsep_mask": guides.get("freqsep_mask")}) + + StyleMMDiT = guides.get('StyleMMDiT') + if StyleMMDiT is not None: + StyleMMDiT.init_guides(model) + LG.ADAIN_NOISE_MODE = StyleMMDiT.noise_mode + + if EO("mycoshock"): + StyleMMDiT.Retrojector = model.inner_model.inner_model.diffusion_model.Retrojector + image_initial_shock = StyleMMDiT.apply_data_shock(image_initial.to(x)) + if VE_MODEL: + x = image_initial_shock.to(x) + sigmas[0] * noise_initial.to(x) + else: + x = (1 - sigmas[0]) * image_initial_shock.to(x) + sigmas[0] * noise_initial.to(x) + + RK.update_transformer_options({"model_sampling": model.inner_model.inner_model.model_sampling}) + # BEGIN SAMPLING LOOP + + while step < num_steps: + sigma, sigma_next = sigmas[step], sigmas[step+1] + + # Apply per-step latent normalization for NestedTensor latents + if latent_shapes is not None and latent_normalize_idx_0_steps is not None: + x = apply_per_step_latent_normalization( + x, step, latent_shapes, + latent_normalize_idx_0_steps or [1.0], + latent_normalize_idx_1_steps or [1.0] + ) + + if sigma_next > sigma: + step_sched = torch.where(torch.flip(sigmas, dims=[0]) == sigma)[0][0].item() + else: + step_sched = step + + rk_type, swapped = RK.swap_rk_type_at_step_or_threshold(x_0, data_prev_, NS, sigmas, step, step_sched, rk_swaps) + if swapped: + implicit_steps_full = 0 + implicit_steps_diag = 0 + + SYNC_GUIDE_ACTIVE = LG.guide_mode.startswith("sync") and (LG.lgw[step_sched] != 0 or LG.lgw_inv[step_sched] != 0 or LG.lgw_sync[step_sched] != 0 or LG.lgw_sync_inv[step_sched] != 0) + + if StyleMMDiT is not None: + RK.update_transformer_options({'StyleMMDiT': StyleMMDiT}) + else: + if LG.HAS_LATENT_GUIDE_ADAIN: + if LG.lgw_adain[step_sched] == 0.0: + RK.update_transformer_options({'y0_adain': None}) + RK.update_transformer_options({'blocks_adain': {}}) + RK.update_transformer_options({'sort_and_scatter': {}}) + else: + RK.update_transformer_options({'y0_adain': LG.y0_adain.clone()}) + if 'blocks_adain_mmdit' in guides: + blocks_adain = { + "double_weights": [val * LG.lgw_adain[step_sched] for val in guides['blocks_adain_mmdit']['double_weights']], + "single_weights": [val * LG.lgw_adain[step_sched] for val in guides['blocks_adain_mmdit']['single_weights']], + "double_blocks" : guides['blocks_adain_mmdit']['double_blocks'], + "single_blocks" : guides['blocks_adain_mmdit']['single_blocks'], + } + RK.update_transformer_options({'blocks_adain': blocks_adain}) + RK.update_transformer_options({'sort_and_scatter': guides['sort_and_scatter']}) + RK.update_transformer_options({'noise_mode_adain': guides['sort_and_scatter']['noise_mode']}) + + + if LG.HAS_LATENT_GUIDE_ATTNINJ: + if LG.lgw_attninj[step_sched] == 0.0: + RK.update_transformer_options({'y0_attninj': None}) + RK.update_transformer_options({'blocks_attninj' : {}}) + RK.update_transformer_options({'blocks_attninj_qkv': {}}) + else: + RK.update_transformer_options({'y0_attninj': LG.y0_attninj.clone()}) + if 'blocks_attninj_mmdit' in guides: + blocks_attninj = { + "double_weights": [val * LG.lgw_attninj[step_sched] for val in guides['blocks_attninj_mmdit']['double_weights']], + "single_weights": [val * LG.lgw_attninj[step_sched] for val in guides['blocks_attninj_mmdit']['single_weights']], + "double_blocks" : guides['blocks_attninj_mmdit']['double_blocks'], + "single_blocks" : guides['blocks_attninj_mmdit']['single_blocks'], + } + RK.update_transformer_options({'blocks_attninj' : blocks_attninj}) + RK.update_transformer_options({'blocks_attninj_qkv': guides['blocks_attninj_qkv']}) + + if LG.HAS_LATENT_GUIDE_STYLE_POS: + if LG.lgw_style_pos[step_sched] == 0.0: + RK.update_transformer_options({'y0_style_pos': None}) + RK.update_transformer_options({'y0_style_pos_weight': 0.0}) + RK.update_transformer_options({'y0_style_pos_synweight': 0.0}) + RK.update_transformer_options({'y0_style_pos_mask': None}) + else: + RK.update_transformer_options({'y0_style_pos': LG.y0_style_pos.clone()}) + RK.update_transformer_options({'y0_style_pos_weight': LG.lgw_style_pos[step_sched]}) + RK.update_transformer_options({'y0_style_pos_synweight': guides['synweight_style_pos']}) + RK.update_transformer_options({'y0_style_pos_mask': LG.mask_style_pos.clone() if LG.mask_style_pos is not None else None}) + RK.update_transformer_options({'y0_style_pos_mask_edge': guides.get('mask_edge_style_pos')}) + RK.update_transformer_options({'y0_style_method': guides['style_method']}) + RK.update_transformer_options({'y0_style_tile_height': guides.get('style_tile_height')}) + RK.update_transformer_options({'y0_style_tile_width': guides.get('style_tile_width')}) + RK.update_transformer_options({'y0_style_tile_padding': guides.get('style_tile_padding')}) + + if EO("style_edge_width"): + RK.update_transformer + + #if LG.HAS_LATENT_GUIDE: + # y0_cache = LG.y0.clone().cpu() + # RK.update_transformer_options({'y0_standard_guide': LG.y0}) + + if LG.HAS_LATENT_GUIDE_INV and y0_inv_standard_guide is None: + y0_inv_cache = LG.y0_inv.clone().cpu() + RK.update_transformer_options({'y0_inv_standard_guide': LG.y0_inv}) + + + if LG.HAS_LATENT_GUIDE_STYLE_NEG: + if LG.lgw_style_neg[step_sched] == 0.0: + RK.update_transformer_options({'y0_style_neg': None}) + RK.update_transformer_options({'y0_style_neg_weight': 0.0}) + RK.update_transformer_options({'y0_style_neg_synweight': 0.0}) + RK.update_transformer_options({'y0_style_neg_mask': None}) + else: + RK.update_transformer_options({'y0_style_neg': LG.y0_style_neg.clone()}) + RK.update_transformer_options({'y0_style_neg_weight': LG.lgw_style_neg[step_sched]}) + RK.update_transformer_options({'y0_style_neg_synweight': guides['synweight_style_neg']}) + RK.update_transformer_options({'y0_style_neg_mask': LG.mask_style_neg.clone() if LG.mask_style_neg is not None else None}) + RK.update_transformer_options({'y0_style_neg_mask_edge': guides.get('mask_edge_style_neg')}) + RK.update_transformer_options({'y0_style_method': guides['style_method']}) + RK.update_transformer_options({'y0_style_tile_height': guides.get('style_tile_height')}) + RK.update_transformer_options({'y0_style_tile_width': guides.get('style_tile_width')}) + RK.update_transformer_options({'y0_style_tile_padding': guides.get('style_tile_padding')}) + + if AttnMask_neg is not None: + RK.update_transformer_options({'regional_conditioning_weight_neg': RegParam_neg.weights[step_sched]}) + RK.update_transformer_options({'regional_conditioning_floor_neg': RegParam_neg.floors[step_sched]}) + + if AttnMask is not None: + RK.update_transformer_options({'regional_conditioning_weight': RegParam.weights[step_sched]}) + RK.update_transformer_options({'regional_conditioning_floor': RegParam.floors[step_sched]}) + + elif regional_conditioning_weights is not None: + RK.extra_args['model_options']['transformer_options']['regional_conditioning_weight'] = regional_conditioning_weights[step_sched] + RK.extra_args['model_options']['transformer_options']['regional_conditioning_floor'] = regional_conditioning_floors [step_sched] + + epsilon_scale = float(epsilon_scales [step_sched]) if epsilon_scales is not None else None + eta = etas [step_sched].to(x) if etas is not None else eta + eta_substep = etas_substep [step_sched].to(x) if etas_substep is not None else eta_substep + s_noise = s_noises [step_sched].to(x) if s_noises is not None else s_noise + s_noise_substep = s_noises_substep [step_sched].to(x) if s_noises_substep is not None else s_noise_substep + noise_scaling_eta = noise_scaling_etas [step_sched].to(x) if noise_scaling_etas is not None else noise_scaling_eta + noise_scaling_weight = noise_scaling_weights[step_sched].to(x) if noise_scaling_weights is not None else noise_scaling_weight + + NS.set_sde_step(sigma, sigma_next, eta, overshoot, s_noise) + RK.set_coeff(rk_type, NS.h, c1, c2, c3, step, sigmas, NS.sigma_down) + NS.set_substep_list(RK) + + if (noise_scaling_eta > 0 or noise_scaling_weight != 0) and noise_scaling_type != "model_d": + if noise_scaling_type == "model_alpha": + VP_OVERRIDE=True + else: + VP_OVERRIDE=None + if noise_scaling_type in {"sampler", "model", "model_alpha"}: + if noise_scaling_type == "model_alpha": + sigma_divisor = NS.sigma_max + else: + sigma_divisor = 1.0 + + if RK.multistep_stages > 0: # hardcoded s_[1] for multistep samplers, which are never multistage + lying_su, lying_sigma, lying_sd, lying_alpha_ratio = NS.get_sde_step(NS.s_[1]/sigma_divisor, NS.s_[0]/sigma_divisor, noise_scaling_eta, noise_scaling_mode, VP_OVERRIDE=VP_OVERRIDE) + + else: + lying_su, lying_sigma, lying_sd, lying_alpha_ratio = NS.get_sde_step(sigma/sigma_divisor, NS.sigma_down/sigma_divisor, noise_scaling_eta, noise_scaling_mode, VP_OVERRIDE=VP_OVERRIDE) + for _ in range(noise_scaling_cycles-1): + lying_su, lying_sigma, lying_sd, lying_alpha_ratio = NS.get_sde_step(sigma/sigma_divisor, lying_sd/sigma_divisor, noise_scaling_eta, noise_scaling_mode, VP_OVERRIDE=VP_OVERRIDE) + lying_s_ = NS.get_substep_list(RK, sigma, RK.h_fn(lying_sd, lying_sigma)) + lying_s_ = NS.s_ + noise_scaling_weight * (lying_s_ - NS.s_) + else: + lying_s_ = NS.s_.clone() + + + rk_swap_stages = 3 if rk_swaps else 0 + data_prev_len = len(data_prev_)-1 if data_prev_ is not None else 3 + recycled_stages = max(rk_swap_stages, RK.multistep_stages, RK.hybrid_stages, data_prev_len) + + if INIT_SAMPLE_LOOP: + INIT_SAMPLE_LOOP = False + x_, data_, eps_, eps_prev_ = (torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) for _ in range(4)) + if LG.ADAIN_NOISE_MODE == "smart": + z_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + z_[0] = noise_initial.clone() + RK.update_transformer_options({'z_' : z_}) + + if sampler_mode in {"unsample", "resample"}: + data_prev_ = state_info.get('data_prev_') + if data_prev_ is not None: + if x.shape == state_info['raw_x'].shape: + data_prev_ = state_info['data_prev_'].clone().to(dtype=work_dtype, device=work_device) + else: + resized_items = [ + LatentHandler(prev_item, latent_shapes) + .map_with(x, lambda p_t, x_t: comfy.utils.common_upscale(p_t, x_t.shape[-1], x_t.shape[-2], "bislerp", "disabled").to(x_t)) + .tensor + for prev_item in state_info['data_prev_'] + ] + data_prev_ = torch.stack(resized_items).to(x) + else: + data_prev_ = torch.zeros(4, *x.shape, dtype=work_dtype, device=work_device) # multistep max is 4m... so 4 needed + else: + data_prev_ = torch.zeros(4, *x.shape, dtype=work_dtype, device=work_device) # multistep max is 4m... so 4 needed + + recycled_stages = len(data_prev_)-1 + + if RK.rows+2 > x_.shape[0]: + row_gap = RK.rows+2 - x_.shape[0] + x_gap_, data_gap_, eps_gap_, eps_prev_gap_ = (torch.zeros(row_gap, *x.shape, dtype=work_dtype, device=work_device) for _ in range(4)) + x_ = torch.cat((x_ ,x_gap_) , dim=0) + data_ = torch.cat((data_ ,data_gap_) , dim=0) + eps_ = torch.cat((eps_ ,eps_gap_) , dim=0) + eps_prev_ = torch.cat((eps_prev_,eps_prev_gap_), dim=0) + + if LG.ADAIN_NOISE_MODE == "smart": + z_gap_ = torch.zeros(row_gap, *x.shape, dtype=work_dtype, device=work_device) + z_ = torch.cat((z_ ,z_gap_) , dim=0) + RK.update_transformer_options({'z_' : z_}) + + sde_noise_t = None + if SDE_NOISE_EXTERNAL: + if step >= len(sde_noise): + SDE_NOISE_EXTERNAL=False + else: + sde_noise_t = sde_noise[step] + + x_[0] = x.clone() + # PRENOISE METHOD HERE! + x_0 = x_[0].clone() + if EO("guide_step_cutoff") or EO("guide_step_min"): + x_0_orig = x_0.clone() + + # RECYCLE STAGES FOR MULTISTEP + if RK.multistep_stages > 0 or RK.hybrid_stages > 0: + if SYNC_GUIDE_ACTIVE: + lgw_mask_, lgw_mask_inv_ = LG.get_masks_for_step(step) + lgw_mask_sync_, lgw_mask_sync_inv_ = LG.get_masks_for_step(step, lgw_type="sync") + + weight_mask = lgw_mask_+lgw_mask_inv_ + if LG.SYNC_SEPARATE: + sync_mask = lgw_mask_sync_+lgw_mask_sync_inv_ + else: + sync_mask = 1. + + if VE_MODEL: + yt_0 = y0_bongflow + sigma * noise_bongflow + else: + yt_0 = (1-sigma) * y0_bongflow + sigma * noise_bongflow + for ms in range(min(len(data_prev_), len(eps_))): + eps_x = RK.get_epsilon_anchored(x_0, data_prev_x_[ms], sigma) + eps_y = RK.get_epsilon_anchored(yt_0, data_prev_y_[ms], sigma) + eps_x2y = RK.get_epsilon_anchored(yt_0, data_prev_y_[ms], sigma) + + if RK.EXPONENTIAL: + if VE_MODEL: + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + sigma*(-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + sigma*(-noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + sigma*(y0_bongflow-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + sigma*(y0_bongflow-noise_bongflow)) + else: + if VE_MODEL: + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + (noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + (noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + (noise_bongflow-y0_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + (noise_bongflow-y0_bongflow)) + + #if RK.EXPONENTIAL: + # if VE_MODEL: + # eps_[ms] = sync_mask * weight_mask_inv * (eps_x - weight_mask * eps_y) + weight_mask * sigma*(-noise_bongflow) + # else: + # #eps_[ms] = (lgw_mask_sync_+lgw_mask_sync_inv_) * (1-(lgw_mask_+lgw_mask_inv_)) * (eps_x - (lgw_mask_+lgw_mask_inv_) * eps_y) + (lgw_mask_+lgw_mask_inv_) * sigma*(y0_bongflow-noise_bongflow) + # eps_[ms] = sync_mask * weight_mask_inv * (eps_x - weight_mask * eps_y) + weight_mask * sigma*(y0_bongflow-noise_bongflow) + #else: + # if VE_MODEL: + # eps_[ms] = sync_mask * weight_mask_inv * (eps_x - weight_mask * eps_y) + weight_mask * (noise_bongflow) + # else: + # #eps_[ms] = (lgw_mask_sync_+lgw_mask_sync_inv_) * (1-(lgw_mask_+lgw_mask_inv_)) * (eps_x - (lgw_mask_+lgw_mask_inv_) * eps_y) + (lgw_mask_+lgw_mask_inv_) * (noise_bongflow-y0_bongflow) + # eps_[ms] = sync_mask * weight_mask_inv * (eps_x - weight_mask * eps_y) + weight_mask * (noise_bongflow-y0_bongflow) + eps_prev_ = eps_.clone() + + else: + for ms in range(min(len(data_prev_), len(eps_))): + eps_[ms] = RK.get_epsilon_anchored(x_0, data_prev_[ms], sigma) + eps_prev_ = eps_.clone() + + + + # INITIALIZE IMPLICIT SAMPLING + if RK.IMPLICIT: + x_, eps_, data_ = init_implicit_sampling(RK, x_0, x_, eps_, eps_prev_, data_, eps, denoised, denoised_prev2, step, sigmas, NS.h, NS.s_, EO, SYNC_GUIDE_ACTIVE) + + implicit_steps_total = (implicit_steps_full + 1) * (implicit_steps_diag + 1) + + # BEGIN FULLY IMPLICIT LOOP + cossim_counter = 0 + adaptive_lgw = LG.lgw.clone() + full_iter = 0 + while full_iter < implicit_steps_full+1: + + if RK.IMPLICIT: + x_, eps_ = RK.newton_iter(x_0, x_, eps_, eps_prev_, data_, NS.s_, 0, NS.h, sigmas, step, "init", SYNC_GUIDE_ACTIVE) + + # PREPARE FULLY PSEUDOIMPLICIT GUIDES + if step > 0 or not SKIP_PSEUDO: + if full_iter > 0 and EO("fully_implicit_reupdate_x"): + x_[0] = NS.sigma_from_to(x_0, x, sigma, sigma_next, NS.s_[0]) + x_0 = NS.sigma_from_to(x_0, x, sigma, sigma_next, sigma) + + if EO("fully_pseudo_init") and full_iter == 0: + guide_mode_tmp = LG.guide_mode + LG.guide_mode = "fully_" + LG.guide_mode + x_0, x_, eps_ = LG.prepare_fully_pseudoimplicit_guides_substep(x_0, x_, eps_, eps_prev_, data_, denoised_prev, 0, step, step_sched, sigmas, eta_substep, overshoot_substep, s_noise_substep, \ + NS, RK, pseudoimplicit_row_weights, pseudoimplicit_step_weights, full_iter, BONGMATH) + if EO("fully_pseudo_init") and full_iter == 0: + LG.guide_mode = guide_mode_tmp + + # TABLEAU LOOP + for row in range(RK.rows - RK.multistep_stages - RK.row_offset + 1): + diag_iter = 0 + while diag_iter < implicit_steps_diag+1: + + + if noise_sampler_type_substep == "brownian" and (full_iter > 0 or diag_iter > 0): + eta_substep = 0. + + NS.set_sde_substep(row, RK.multistep_stages, eta_substep, overshoot_substep, s_noise_substep, full_iter, diag_iter, implicit_steps_full, implicit_steps_diag) + + # PRENOISE METHOD HERE! + + # A-TABLEAU + if row < RK.rows: + + # PREPARE PSEUDOIMPLICIT GUIDES + if step > 0 or not SKIP_PSEUDO: + x_0, x_, eps_, x_row_pseudoimplicit, sub_sigma_pseudoimplicit = LG.process_pseudoimplicit_guides_substep(x_0, x_, eps_, eps_prev_, data_, denoised_prev, row, step, step_sched, sigmas, NS, RK, \ + pseudoimplicit_row_weights, pseudoimplicit_step_weights, full_iter, BONGMATH) + + # PREPARE MODEL CALL + if LG.guide_mode in GUIDE_MODE_NAMES_PSEUDOIMPLICIT and (step > 0 or not SKIP_PSEUDO) and (LG.lgw[step_sched] > 0 or LG.lgw_inv[step_sched] > 0) and x_row_pseudoimplicit is not None: + + x_tmp = x_row_pseudoimplicit + s_tmp = sub_sigma_pseudoimplicit + + # Fully implicit iteration (explicit only) # or... Fully implicit iteration (implicit only... not standard) + elif (full_iter > 0 and RK.row_offset == 1 and row == 0) or (full_iter > 0 and RK.row_offset == 0 and row == 0 and EO("fully_implicit_update_x")): + if EO("fully_explicit_pogostick_eta"): + super_alpha_ratio, super_sigma_down, super_sigma_up = NS.get_sde_coeff(sigma, sigma_next, None, eta) + x = super_alpha_ratio * x + super_sigma_up * NS.noise_sampler(sigma=sigma_next, sigma_next=sigma) + + x_tmp = x + s_tmp = sigma + elif EO("enable_fully_explicit_lagrange_rebound1"): + substeps_prev = len(RK.C[:-1]) + x_tmp = lagrange_interpolation(RK.C[1:-1], x_[1:substeps_prev], RK.C[0]).squeeze(0) + + elif EO("enable_fully_explicit_lagrange_rebound2"): + substeps_prev = len(RK.C[:-1]) + x_tmp = lagrange_interpolation(RK.C[1:], x_[1:substeps_prev+1], RK.C[0]).squeeze(0) + + elif EO("enable_fully_explicit_rebound1"): # 17630, faded dots, just crap + eps_tmp, denoised_tmp = RK(x, sigma_next, x, sigma_next) + eps_tmp = (x - denoised_tmp) / sigma_next + x_[0] = denoised_tmp + sigma * eps_tmp + + x_0 = x_[0] + x_tmp = x_[0] + s_tmp = sigma + + elif implicit_type == "rebound": # TODO: ADAPT REBOUND IMPLICIT TO WORK WITH FLOW GUIDE MODE + eps_tmp, denoised_tmp = RK(x, sigma_next, x_0, sigma) + eps_tmp = (x - denoised_tmp) / sigma_next + x = denoised_tmp + sigma * eps_tmp + + x_tmp = x + s_tmp = sigma + + elif implicit_type == "retro-eta" and (NS.sub_sigma_up > 0 or NS.sub_sigma_up_eta > 0): + x_tmp = NS.sigma_from_to(x_0, x, sigma, sigma_next, sigma) + s_tmp = sigma + + elif implicit_type == "bongmath" and (NS.sub_sigma_up > 0 or NS.sub_sigma_up_eta > 0): + if BONGMATH: + x_tmp = x_[row] + s_tmp = NS.s_[row] + else: + x_tmp = NS.sigma_from_to(x_0, x, sigma, sigma_next, sigma) + s_tmp = sigma + + else: + x_tmp = x + s_tmp = sigma_next + + + + # All others + else: + # three potential toggle options: force rebound/model call, force PC style, force pogostick style + if diag_iter > 0: # Diagonally implicit iteration (explicit or implicit) + if EO("diag_explicit_pogostick_eta"): + super_alpha_ratio, super_sigma_down, super_sigma_up = NS.get_sde_coeff(NS.s_[row], NS.s_[row+RK.row_offset+RK.multistep_stages], None, eta) + x_[row+RK.row_offset] = super_alpha_ratio * x_[row+RK.row_offset] + super_sigma_up * NS.noise_sampler(sigma=NS.s_[row+RK.row_offset+RK.multistep_stages], sigma_next=NS.s_[row]) + + x_tmp = x_[row+RK.row_offset] + s_tmp = sigma + + elif implicit_type_substeps == "rebound": + eps_[row], data_[row] = RK(x_[row+RK.row_offset], NS.s_[row+RK.row_offset+RK.multistep_stages], x_0, sigma) + + x_ = RK.update_substep(x_0, x_, eps_, eps_prev_, row, RK.row_offset, NS.h_new, NS.h_new_orig) + x_[row+RK.row_offset] = NS.rebound_overshoot_substep(x_0, x_[row+RK.row_offset]) + + x_[row+RK.row_offset] = NS.sigma_from_to(x_0, x_[row+RK.row_offset], sigma, NS.s_[row+RK.row_offset+RK.multistep_stages], NS.s_[row]) + x_tmp = x_[row+RK.row_offset] + s_tmp = NS.s_[row] + + elif implicit_type_substeps == "retro-eta" and (NS.sub_sigma_up > 0 or NS.sub_sigma_up_eta > 0): + x_tmp = NS.sigma_from_to(x_0, x_[row+RK.row_offset], sigma, NS.s_[row+RK.row_offset+RK.multistep_stages], NS.s_[row]) + s_tmp = NS.s_[row] + + elif implicit_type_substeps == "bongmath" and (NS.sub_sigma_up > 0 or NS.sub_sigma_up_eta > 0) and not EO("disable_diag_explicit_bongmath_rebound"): + if BONGMATH: + x_tmp = x_[row] + s_tmp = NS.s_[row] + else: + x_tmp = NS.sigma_from_to(x_0, x_[row+RK.row_offset], sigma, NS.s_[row+RK.row_offset+RK.multistep_stages], NS.s_[row]) + s_tmp = NS.s_[row] + + else: + x_tmp = x_[row+RK.row_offset] + s_tmp = NS.s_[row+RK.row_offset+RK.multistep_stages] + else: + x_tmp = x_[row] + s_tmp = NS.sub_sigma + + + + if RK.IMPLICIT: + if not EO("disable_implicit_guide_preproc"): + eps_, x_ = LG.process_guides_substep(x_0, x_, eps_, data_, denoised_prev, row, step, step_sched, sigma, sigma_next, NS.sigma_down, NS.s_, epsilon_scale, RK, full_iter) + eps_prev_, x_ = LG.process_guides_substep(x_0, x_, eps_prev_, data_, denoised_prev, row, step, step_sched, sigma, sigma_next, NS.sigma_down, NS.s_, epsilon_scale, RK, full_iter) + if row == 0 and (EO("implicit_lagrange_init") or EO("radaucycle")): + pass + else: + x_[row+RK.row_offset] = x_0 + NS.h_new * RK.zum(row+RK.row_offset, eps_, eps_prev_) + x_[row+RK.row_offset] = NS.rebound_overshoot_substep(x_0, x_[row+RK.row_offset]) + if row > 0: + if not LG.guide_mode.startswith("flow") or (LG.lgw[step_sched] == 0 and LG.lgw[step+1] == 0 and LG.lgw_inv[step_sched] == 0 and LG.lgw_inv[step+1] == 0): + x_row_tmp = NS.swap_noise_substep(x_0, x_[row+RK.row_offset], mask=sde_mask, guide=LG.y0) + + if LG.ADAIN_NOISE_MODE == "smart": #_smartnoise_implicit"): + data_next = denoised + NS.h_new * RK.zum(row+RK.row_offset+RK.multistep_stages, data_, data_prev_) + if VE_MODEL: + z_[row+RK.row_offset] = (x_row_tmp - data_next) / s_tmp + else: + z_[row+RK.row_offset] = (x_row_tmp - (NS.sigma_max-s_tmp)*data_next) / s_tmp + RK.update_transformer_options({'z_' : z_}) + + if SYNC_GUIDE_ACTIVE: + noise_bongflow_new = (x_row_tmp - x_[row+RK.row_offset]) / s_tmp + noise_bongflow + yt_[row+RK.row_offset] += s_tmp * (noise_bongflow_new - noise_bongflow) + x_0 += sigma * (noise_bongflow_new - noise_bongflow) + if not EO("disable_i_bong"): + for i_bong in range(len(NS.s_)): + x_[i_bong] += NS.s_[i_bong] * (noise_bongflow_new - noise_bongflow) + noise_bongflow = noise_bongflow_new + + x_[row+RK.row_offset] = x_row_tmp + + if SYNC_GUIDE_ACTIVE: + if VE_MODEL: + yt_[:NS.s_.shape[0], 0] = y0_bongflow + NS.s_.view(-1, *[1]*(x.ndim-1)) * (noise_bongflow) + yt_0 = y0_bongflow + sigma * (noise_bongflow) + else: + yt_[:NS.s_.shape[0], 0] = y0_bongflow + NS.s_.view(-1, *[1]*(x.ndim-1)) * (noise_bongflow - y0_bongflow) + yt_0 = y0_bongflow + sigma * (noise_bongflow - y0_bongflow) + + if RK.EXPONENTIAL: + eps_y_ = data_y_ - yt_0 # yt_ # watch out for fuckery with size of tableau being smaller later in a chained sampler + else: + if BONGMATH: + eps_y_[:NS.s_.shape[0]] = (yt_[:NS.s_.shape[0]] - data_y_[:NS.s_.shape[0]]) / NS.s_.view(-1,*[1]*(x_.ndim-1)) + else: + eps_y_[:NS.s_.shape[0]] = (yt_0.repeat(NS.s_.shape[0], *[1]*(x_.ndim-1)) - data_y_[:NS.s_.shape[0]]) / sigma # calc exact to c0 node + if not BONGMATH: + if RK.EXPONENTIAL: + eps_x_ = data_x_ - x_0 + else: + eps_x_ = (x_0 - data_x_) / sigma + + weight_mask = lgw_mask_+lgw_mask_inv_ + if LG.SYNC_SEPARATE: + sync_mask = lgw_mask_sync_+lgw_mask_sync_inv_ + else: + sync_mask = 1. + + for ms in range(len(eps_)): + if RK.EXPONENTIAL: + if VE_MODEL: # ZERO IS THIS # ONE IS THIS + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + sigma*(-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + sigma*(-noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + sigma*(y0_bongflow-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + sigma*(y0_bongflow-noise_bongflow)) + else: + if VE_MODEL: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + (noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + (noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + (noise_bongflow-y0_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + (noise_bongflow-y0_bongflow)) + + + if BONGMATH and step < sigmas.shape[0]-1 and sigma > 0.03 and not EO("disable_implicit_prebong"): + BONGMATH_Y = SYNC_GUIDE_ACTIVE + + x_0, x_, eps_ = RK.bong_iter(x_0, x_, eps_, eps_prev_, data_, sigma, NS.s_, row, RK.row_offset, NS.h, step, step_sched, + BONGMATH_Y, y0_bongflow, noise_bongflow, eps_x_, eps_y_, data_x_, data_y_, LG) # TRY WITH h_new ?? + # BONGMATH_Y, y0_bongflow, noise_bongflow, eps_x_, eps_y_, eps_x2y_, data_x_, LG) # TRY WITH h_new ?? + + #if EO("eps_adain_smartnoise_bongmath"): + if LG.ADAIN_NOISE_MODE == "smart": + if VE_MODEL: + z_[:NS.s_.shape[0], ...] = (x_ - data_)[:NS.s_.shape[0], ...] / NS.s_.view(-1,*[1]*(x_.ndim-1)) + else: + z_[:NS.s_.shape[0], ...] = (x_[:NS.s_.shape[0], ...] - (NS.sigma_max - NS.s_.view(-1,*[1]*(x_.ndim-1)))*data_[:NS.s_.shape[0], ...])[:NS.s_.shape[0], ...] / NS.s_.view(-1,*[1]*(x_.ndim-1)) + RK.update_transformer_options({'z_' : z_}) + + x_tmp = x_[row+RK.row_offset] + + lying_eps_row_factor = 1.0 + # MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL + if RK.IMPLICIT and row == 0 and (EO("implicit_lazy_recycle_first_model_call_at_start") or EO("radaucycle") or RK.C[0] == 0.0): + pass + else: + if s_tmp == 0: + break + x_, eps_ = RK.newton_iter(x_0, x_, eps_, eps_prev_, data_, NS.s_, row, NS.h, sigmas, step, "pre", SYNC_GUIDE_ACTIVE) # will this do anything? not x_tmp + + # DETAIL BOOST + if noise_scaling_type == "model_alpha" and noise_scaling_weight != 0 and noise_scaling_eta > 0: + s_tmp = s_tmp + noise_scaling_weight * (s_tmp * lying_alpha_ratio - s_tmp) + if noise_scaling_type == "model" and noise_scaling_weight != 0 and noise_scaling_eta > 0: + s_tmp = lying_s_[row] + if RK.multistep_stages > 0: + s_tmp = lying_sd + + # SYNC GUIDE --------------------------- + if LG.guide_mode.startswith("sync") and (LG.lgw[step_sched] == 0 and LG.lgw_inv[step_sched] == 0 and LG.lgw_sync[step_sched] == 0 and LG.lgw_sync_inv[step_sched] == 0): + data_cached = None + elif SYNC_GUIDE_ACTIVE: + lgw_mask_, lgw_mask_inv_ = LG.get_masks_for_step(step_sched) + lgw_mask_sync_, lgw_mask_sync_inv_ = LG.get_masks_for_step(step_sched, lgw_type="sync") + lgw_mask_drift_x_, lgw_mask_drift_x_inv_ = LG.get_masks_for_step(step_sched, lgw_type="drift_x") + lgw_mask_drift_y_, lgw_mask_drift_y_inv_ = LG.get_masks_for_step(step_sched, lgw_type="drift_y") + lgw_mask_lure_x_, lgw_mask_lure_x_inv_ = LG.get_masks_for_step(step_sched, lgw_type="lure_x") + lgw_mask_lure_y_, lgw_mask_lure_y_inv_ = LG.get_masks_for_step(step_sched, lgw_type="lure_y") + + weight_mask = lgw_mask_ + lgw_mask_inv_ + sync_mask = lgw_mask_sync_ + lgw_mask_sync_inv_ + + + drift_x_mask = lgw_mask_drift_x_ + lgw_mask_drift_x_inv_ + drift_y_mask = lgw_mask_drift_y_ + lgw_mask_drift_y_inv_ + lure_x_mask = lgw_mask_lure_x_ + lgw_mask_lure_x_inv_ + lure_y_mask = lgw_mask_lure_y_ + lgw_mask_lure_y_inv_ + + if eps_x_ is None: + eps_x_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + data_x_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + eps_y2x_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + eps_x2y_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + eps_yt_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + eps_y_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + eps_prev_y_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + data_y_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + yt_ = torch.zeros(RK.rows+2, *x.shape, dtype=work_dtype, device=work_device) + + RUN_X_0_COPY = False + if noise_bongflow is None: + RUN_X_0_COPY = True + data_prev_x_ = torch.zeros(4, *x.shape, dtype=work_dtype, device=work_device) + data_prev_y_ = torch.zeros(4, *x.shape, dtype=work_dtype, device=work_device) + + noise_bongflow = normalize_zscore(NS.noise_sampler(sigma=sigma, sigma_next=NS.sigma_min), channelwise=True, inplace=True) + + _, _ = RK(noise_bongflow, s_tmp/s_tmp, noise_bongflow, sigma/sigma, transformer_options={'latent_type': 'xt'}) + + if RK.extra_args['model_options']['transformer_options'].get('y0_standard_guide') is not None: + if hasattr(model.inner_model.inner_model.diffusion_model, "y0_standard_guide"): + LG.y0 = y0_standard_guide = model.inner_model.inner_model.diffusion_model.y0_standard_guide.clone() + del model.inner_model.inner_model.diffusion_model.y0_standard_guide + RK.extra_args['model_options']['transformer_options']['y0_standard_guide'] = None + + if RK.extra_args['model_options']['transformer_options'].get('y0_inv_standard_guide') is not None: + if hasattr(model.inner_model.inner_model.diffusion_model, "y0_inv_standard_guide"): + LG.y0_inv = y0_inv_standard_guide = model.inner_model.inner_model.diffusion_model.y0_inv_standard_guide.clone() # RK.extra_args['model_options']['transformer_options'].get('y0_standard_guide') + del model.inner_model.inner_model.diffusion_model.y0_inv_standard_guide + RK.extra_args['model_options']['transformer_options']['y0_inv_standard_guide'] = None + + y0_bongflow = LG.HAS_LATENT_GUIDE * LG.mask * LG.y0 + LG.HAS_LATENT_GUIDE_INV * LG.mask_inv * LG.y0_inv #LG.y0.clone() + + if VE_MODEL: + yt_0 = y0_bongflow + sigma * noise_bongflow + yt = y0_bongflow + s_tmp * noise_bongflow + else: + yt_0 = (1-sigma) * y0_bongflow + sigma * noise_bongflow + yt = (1-s_tmp) * y0_bongflow + s_tmp * noise_bongflow + + yt_[row] = yt + + if RUN_X_0_COPY: + x_0 = yt_0.clone() + x_tmp = x_[row] = yt.clone() + else: + y0_bongflow_orig = y0_bongflow.clone() if y0_bongflow_orig is None else y0_bongflow_orig + y0_bongflow = y0_bongflow + LG.drift_x_data * drift_x_mask * (data_x - y0_bongflow) \ + + LG.drift_x_sync * drift_x_mask * (data_barf - y0_bongflow) \ + + LG.drift_y_data * drift_y_mask * (data_y - y0_bongflow) \ + + LG.drift_y_sync * drift_y_mask * (data_barf_y - y0_bongflow) \ + + LG.drift_y_guide * drift_y_mask * (y0_bongflow_orig - y0_bongflow) + + if torch.norm(y0_bongflow_orig - y0_bongflow) != 0 and EO("enable_y0_bongflow_update"): + RK.update_transformer_options({'y0_style_pos': y0_bongflow.clone()}) + + if not EO("skip_yt"): + yt_0 = RK.get_x(y0_bongflow, noise_bongflow, sigma) + yt = RK.get_x(y0_bongflow, noise_bongflow, s_tmp) + + yt_[row] = yt + + if ((LG.lgw[step_sched].item() in {1,0} and LG.lgw_inv[step_sched].item() in {1,0} and LG.lgw[step_sched] == 1-LG.lgw_sync[step_sched] and LG.lgw_inv[step_sched] == 1-LG.lgw_sync_inv[step_sched]) or EO("sync_speed_mode")) and not EO("disable_sync_speed_mode"): + data_y = y0_bongflow.clone() + eps_y = RK.get_eps(yt_0, yt_[row], data_y, sigma, s_tmp) + + else: + eps_y, data_y = RK(yt_[row], s_tmp, yt_0, sigma, transformer_options={'latent_type': 'yt'}) + + eps_x, data_x = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'xt', 'row': row, "x_tmp": x_tmp}) + #if hasattr(model.inner_model.inner_model.diffusion_model, "eps_out"): + + + for sync_lure_iter in range(LG.sync_lure_iter): + if LG.sync_lure_sequence == "x -> y": + + if lure_x_mask.abs().sum() > 0: + x_tmp = LG.swap_data(x_tmp, data_x, data_y, s_tmp, lure_x_mask) + eps_x_lure, data_x_lure = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'xt'}) + eps_x = eps_x + lure_x_mask * (eps_x_lure - eps_x) + data_x = data_x + lure_x_mask * (data_x_lure - data_x) + + if lure_y_mask.abs().sum() > 0: + y_tmp = yt_[row].clone() + y_tmp = LG.swap_data(y_tmp, data_y, data_x, s_tmp, lure_y_mask) + eps_y_lure, data_y_lure = RK(y_tmp, s_tmp, yt_0, sigma, transformer_options={'latent_type': 'yt'}) + eps_y = eps_y + lure_y_mask * (eps_y_lure - eps_y) + data_y = data_y + lure_y_mask * (data_y_lure - data_y) + + elif LG.sync_lure_sequence == "y -> x": + + if lure_y_mask.abs().sum() > 0: + y_tmp = yt_[row].clone() + y_tmp = LG.swap_data(y_tmp, data_y, data_x, s_tmp, lure_y_mask) + eps_y_lure, data_y_lure = RK(y_tmp, s_tmp, yt_0, sigma, transformer_options={'latent_type': 'yt'}) + eps_y = eps_y + lure_y_mask * (eps_y_lure - eps_y) + data_y = data_y + lure_y_mask * (data_y_lure - data_y) + + if lure_x_mask.abs().sum() > 0: + x_tmp = LG.swap_data(x_tmp, data_x, data_y, s_tmp, lure_x_mask) + eps_x_lure, data_x_lure = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'xt'}) + eps_x = eps_x + lure_x_mask * (eps_x_lure - eps_x) + data_x = data_x + lure_x_mask * (data_x_lure - data_x) + + elif LG.sync_lure_sequence == "xy -> xy": + data_x_orig, data_y_orig = data_x.clone(), data_y.clone() + + if lure_x_mask.abs().sum() > 0: + x_tmp = LG.swap_data(x_tmp, data_x_orig, data_y_orig, s_tmp, lure_x_mask) + eps_x_lure, data_x_lure = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'xt'}) + eps_x = eps_x + lure_x_mask * (eps_x_lure - eps_x) + data_x = data_x + lure_x_mask * (data_x_lure - data_x) + + if lure_y_mask.abs().sum() > 0: + y_tmp = yt_[row].clone() + y_tmp = LG.swap_data(y_tmp, data_y_orig, data_x_orig, s_tmp, lure_y_mask) + eps_y_lure, data_y_lure = RK(y_tmp, s_tmp, yt_0, sigma, transformer_options={'latent_type': 'yt'}) + eps_y = eps_y + lure_y_mask * (eps_y_lure - eps_y) + data_y = data_y + lure_y_mask * (data_y_lure - data_y) + + if EO("sync_proj_y"): + d_collinear_d_lerp = get_collinear(eps_x, eps_y) + d_lerp_ortho_d = get_orthogonal(eps_y, eps_x) + eps_y = d_collinear_d_lerp + d_lerp_ortho_d + + if EO("sync_proj_y2"): + d_collinear_d_lerp = get_collinear(eps_y, eps_x) + d_lerp_ortho_d = get_orthogonal(eps_x, eps_y) + eps_y = d_collinear_d_lerp + d_lerp_ortho_d + + if EO("sync_proj_x"): + d_collinear_d_lerp = get_collinear(eps_y, eps_x) + d_lerp_ortho_d = get_orthogonal(eps_x, eps_y) + eps_x = d_collinear_d_lerp + d_lerp_ortho_d + + if EO("sync_proj_x2"): + d_collinear_d_lerp = get_collinear(eps_x, eps_y) + d_lerp_ortho_d = get_orthogonal(eps_y, eps_x) + eps_x = d_collinear_d_lerp + d_lerp_ortho_d + + eps_x2y = RK.get_eps(x_0, x_[row], data_y, sigma, s_tmp) + eps_x2y_[row] = eps_x2y + + eps_y2x = RK.get_eps(x_0, x_[row], data_y, sigma, s_tmp) + eps_y2x_[row] = eps_y2x + + if RK.EXPONENTIAL: + if VE_MODEL: # ZERO IS THIS # ONE IS THIS + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + sigma*(-noise_bongflow)) + if EO("sync_x2y"): + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + sigma*(-noise_bongflow)) + else: + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (-eps_y + sigma*(y0_bongflow-noise_bongflow)) #+ lure_x_mask * sigma*(data_y - data_x) + if EO("sync_x2y"): + eps_[row] = sync_mask * eps_x - (1-sync_mask) * eps_x2y + weight_mask * (-eps_x2y + sigma*(y0_bongflow-noise_bongflow)) + eps_yt_[row] = sync_mask * eps_y + (1-sync_mask) * eps_y2x + weight_mask * (-eps_x + sigma*(y0_bongflow-noise_bongflow)) # differentiate guide as well toward the x pred? + else: + if VE_MODEL: + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (noise_bongflow - eps_y) + if EO("sync_x2y"): + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (noise_bongflow - eps_x2y) + else: + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (noise_bongflow - eps_y - y0_bongflow) + if EO("sync_x2y"): + eps_[row] = sync_mask * eps_x + (1-sync_mask) * eps_x2y + weight_mask * (noise_bongflow - eps_x2y - y0_bongflow) + eps_yt_[row] = sync_mask * eps_y + (1-sync_mask) * eps_y2x + weight_mask * (noise_bongflow - eps_x - y0_bongflow) # differentiate guide as well toward the x pred? + + if VE_MODEL: + data_[row] = x_0 + sync_mask * NS.h * eps_x + (1-sync_mask) * NS.h * eps_x2y - weight_mask * (sigma*(eps_y + noise_bongflow)) # - lure_x_mask * (sigma*(eps_y + eps_x)) + data_barf_y = yt_0 + sync_mask * NS.h * eps_y + (1-sync_mask) * NS.h * eps_y2x - weight_mask * (sigma*(eps_x + noise_bongflow)) + if EO("sync_x2y"): + data_[row] = x_0 + sync_mask * NS.h * eps_x + (1-sync_mask) * NS.h * eps_x2y - weight_mask * (sigma*(eps_x2y + noise_bongflow)) + + else: + + data_[row] = x_0 + sync_mask * NS.h * eps_x + (1-sync_mask) * NS.h * eps_x2y - weight_mask * (NS.h * eps_y + sigma*(noise_bongflow-y0_bongflow)) + data_barf_y = yt_0 + sync_mask * NS.h * eps_y + (1-sync_mask) * NS.h * eps_y2x - weight_mask * (NS.h * eps_x + sigma*(noise_bongflow-y0_bongflow)) + if EO("sync_x2y"): + data_[row] = x_0 + sync_mask * NS.h * eps_x + (1-sync_mask) * NS.h * eps_x2y - weight_mask * (NS.h * eps_x2y + sigma*(noise_bongflow-y0_bongflow)) + + if EO("data_is_y0_with_lure_x_mask"): + data_[row] = data_[row] + lure_x_mask * (y0_bongflow - data_[row]) + + if EO("eps_is_y0_with_lure_x_mask"): + if RK.EXPONENTIAL: + eps_[row] = eps_[row] + lure_x_mask * ((y0_bongflow - x_0) - eps_[row]) + else: + eps_[row] = eps_[row] + lure_x_mask * (((x_0 - y0_bongflow) / sigma) - eps_[row]) + data_barf = data_[row] + data_cached = data_x + + eps_x_ [row] = eps_x + data_x_[row] = data_x + + eps_y_ [row] = eps_y + data_y_[row] = data_y + + if EO("sync_use_fake_eps_y"): + if RK.EXPONENTIAL: + if VE_MODEL: + eps_y_ [row] = sigma * ( - noise_bongflow) + else: + eps_y_ [row] = sigma * (y0_bongflow - noise_bongflow) + else: + if VE_MODEL: + eps_y_ [row] = noise_bongflow + else: + eps_y_ [row] = noise_bongflow - y0_bongflow + if EO("sync_use_fake_data_y"): + data_y_[row] = y0_bongflow + + + + + elif LG.guide_mode.startswith("flow") and (LG.lgw[step_sched] > 0 or LG.lgw_inv[step_sched] > 0) and not FLOW_STOPPED and not EO("flow_sync") : + lgw_mask_, lgw_mask_inv_ = LG.get_masks_for_step(step) + if not FLOW_STARTED and not FLOW_RESUMED: + FLOW_STARTED = True + data_x_prev_ = torch.zeros_like(data_prev_) + + y0 = LG.HAS_LATENT_GUIDE * LG.mask * LG.y0 + LG.HAS_LATENT_GUIDE_INV * LG.mask_inv * LG.y0_inv + + yx0 = y0.clone() + + if EO("flow_slerp"): + y0_inv = LG.HAS_LATENT_GUIDE * LG.mask * LG.y0_inv + LG.HAS_LATENT_GUIDE_INV * LG.mask_inv * LG.y0 + y0 = LG.y0.clone() + y0_inv = LG.y0_inv.clone() + flow_slerp_guide_ratio = EO("flow_slerp_guide_ratio", 0.5) + y_slerp = slerp_tensor(flow_slerp_guide_ratio, y0, y0_inv) + yx0 = y_slerp.clone() + + x_[row], x_0 = yx0.clone(), yx0.clone() + if EO("guide_step_cutoff") or EO("guide_step_min"): + x_0_orig = yx0.clone() + + if EO("flow_yx0_init_y0_inv"): + yx0 = LG.HAS_LATENT_GUIDE * LG.mask * LG.y0_inv + LG.HAS_LATENT_GUIDE_INV * LG.mask_inv * LG.y0 + + if step > 0: + if EO("flow_manual_masks"): + y0 = (1 - (LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv)) * denoised + LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask * LG.y0 + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv * LG.y0_inv + else: + y0 = (1 - (lgw_mask_ + lgw_mask_inv_)) * denoised + lgw_mask_ * LG.y0 + lgw_mask_inv_ * LG.y0_inv + yx0 = y0.clone() + + if EO("flow_slerp"): + if EO("flow_manual_masks"): + y0_inv = (1 - (LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv)) * denoised + LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask * LG.y0_inv + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv * LG.y0 + else: + y0_inv = (1 - (lgw_mask_ + lgw_mask_inv_)) * denoised + lgw_mask_ * LG.y0_inv + lgw_mask_inv_ * LG.y0 + flow_slerp_guide_ratio = EO("flow_slerp_guide_ratio", 0.5) + y_slerp = slerp_tensor(flow_slerp_guide_ratio, y0, y0_inv) + yx0 = y_slerp.clone() + + else: + yx0_prev = data_cached + if EO("flow_manual_masks"): + yx0 = (1 - (LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv)) * yx0_prev + LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask * x_tmp + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv * x_tmp + else: + yx0 = (1 - (lgw_mask_ + lgw_mask_inv_)) * yx0_prev + (lgw_mask_ + lgw_mask_inv_) * x_tmp + + if not EO("flow_static_guides"): + if EO("flow_manual_masks"): + y0 = (1 - (LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv)) * yx0_prev + LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask * LG.y0 + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv * LG.y0_inv + else: + y0 = (1 - (lgw_mask_ + lgw_mask_inv_)) * yx0_prev + lgw_mask_ * LG.y0 + lgw_mask_inv_ * LG.y0_inv + + if EO("flow_slerp"): + if EO("flow_manual_masks"): + y0_inv = (1 - (LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv)) * yx0_prev + LG.HAS_LATENT_GUIDE * LG.lgw[step_sched] * LG.mask * LG.y0_inv + LG.HAS_LATENT_GUIDE_INV * LG.lgw_inv[step_sched] * LG.mask_inv * LG.y0 + else: + y0_inv = (1 - (lgw_mask_ + lgw_mask_inv_)) * yx0_prev + lgw_mask_ * LG.y0_inv + lgw_mask_inv_ * LG.y0 + + y0_orig = y0.clone() + if EO("flow_proj_xy"): + d_collinear_d_lerp = get_collinear(yx0, y0_orig) + d_lerp_ortho_d = get_orthogonal(y0_orig, yx0) + y0 = d_collinear_d_lerp + d_lerp_ortho_d + + if EO("flow_proj_yx"): + d_collinear_d_lerp = get_collinear(y0_orig, yx0) + d_lerp_ortho_d = get_orthogonal(yx0, y0_orig) + yx0 = d_collinear_d_lerp + d_lerp_ortho_d + + y0_inv_orig = None + if EO("flow_proj_xy_inv"): + y0_inv_orig = y0_inv.clone() + d_collinear_d_lerp = get_collinear(yx0, y0_inv) + d_lerp_ortho_d = get_orthogonal(y0_inv, yx0) + y0_inv = d_collinear_d_lerp + d_lerp_ortho_d + + if EO("flow_proj_yx_inv"): + y0_inv_orig = y0_inv if y0_inv_orig is None else y0_inv_orig + d_collinear_d_lerp = get_collinear(y0_inv_orig, yx0) + d_lerp_ortho_d = get_orthogonal(yx0, y0_inv_orig) + yx0 = d_collinear_d_lerp + d_lerp_ortho_d + del y0_orig + + flow_cossim_iter = EO("flow_cossim_iter", 1) + + if step == 0: + noise_yt = noise_fn(y0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) # normalize_zscore(NS.noise_sampler(sigma=sigma, sigma_next=sigma_next), channelwise=True, inplace=True) + if not EO("flow_disable_renoise_y0"): + if noise_yt is None: + noise_yt = noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + else: + noise_yt = (1-eta) * noise_yt + eta * noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + + if VE_MODEL: + yt = y0 + s_tmp * noise_yt + else: + yt = (NS.sigma_max-s_tmp) * y0 + (s_tmp/NS.sigma_max) * noise_yt + if not EO("flow_disable_doublenoise_y0"): + if noise_yt is None: + noise_yt = noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + else: + noise_yt = (1-eta) * noise_yt + eta * noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + + if VE_MODEL: + y0_noised = y0 + sigma * noise_yt + else: + y0_noised = (NS.sigma_max-sigma) * y0 + sigma * noise_yt + + if EO("flow_slerp"): + noise = noise_fn(y0_inv, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + yt_inv = (NS.sigma_max-s_tmp) * y0_inv + (s_tmp/NS.sigma_max) * noise + if not EO("flow_disable_doublenoise_y0_inv"): + noise = noise_fn(y0_inv, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + y0_noised_inv = (NS.sigma_max-sigma) * y0_inv + sigma * noise + + if step == 0: + noise_xt = noise_fn(yx0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + if EO("flow_slerp"): + xt = yx0 + (s_tmp/NS.sigma_max) * (noise - y_slerp) + if not EO("flow_disable_doublenoise_x_0"): + noise = noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + x_0_noised = x_0 + sigma * (noise - y_slerp) + else: + if not EO("flow_disable_renoise_x_0"): + if noise_xt is None: + noise_xt = noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + else: + noise_xt = (1-eta_substep) * noise_xt + eta_substep * noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + + if VE_MODEL: + xt = yx0 + (s_tmp) * yx0 + (s_tmp) * (noise_xt - y0) + else: + xt = yx0 + (s_tmp/NS.sigma_max) * (noise_xt - y0) + if not EO("flow_disable_doublenoise_x_0"): + if noise_xt is None: + noise_xt = noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + else: + noise_xt = (1-eta_substep) * noise_xt + eta_substep * noise_fn(x_0, sigma, sigma_next, NS.noise_sampler, flow_cossim_iter) + if VE_MODEL: + x_0_noised = x_0 + (sigma) * x_0 + (sigma) * (noise_xt - y0) + else: + x_0_noised = x_0 + (sigma/NS.sigma_max) * (noise_xt - y0) # just lerp noise add, (1-sigma)*y0 + sigma*noise assuming x_0 == y0, which is true initially... + + eps_y, data_y = RK(yt, s_tmp, y0_noised, sigma, transformer_options={'latent_type': 'yt'}) + eps_x, data_x = RK(xt, s_tmp, x_0_noised, sigma, transformer_options={'latent_type': 'xt'}) + + if EO("flow_slerp"): + eps_y_inv, data_y_inv = RK(yt_inv, s_tmp, y0_noised_inv, sigma, transformer_options={'latent_type': 'yt_inv'}) + + if LG.lgw[step+1] == 0 and LG.lgw_inv[step+1] == 0: # break out of differentiating x0 and return to differentiating eps/velocity field + if EO("flow_shit_out_yx0"): + eps_ [row] = eps_x - eps_y + data_[row] = yx0 + if row == 0: + x_[row] = x_0 = xt + else: + x_[row] = xt + if not EO("flow_shit_out_new"): + eps_ [row] = eps_x + data_[row] = data_x + if row == 0: + x_[row] = x_0 = xt + else: + x_[row] = xt + + else: + eps_ [row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * eps_x + (lgw_mask_ + lgw_mask_inv_) * eps_y + data_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * data_x + (lgw_mask_ + lgw_mask_inv_) * data_y + if row == 0: + x_[row] = x_0 = (1 - (lgw_mask_ + lgw_mask_inv_)) * xt + (lgw_mask_ + lgw_mask_inv_) * yt + else: + x_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * xt + (lgw_mask_ + lgw_mask_inv_) * yt + + FLOW_STOPPED = True + else: + if not EO("flow_slerp"): + if RK.EXPONENTIAL: + eps_y_alt = data_y - x_0 + eps_x_alt = data_x - x_0 + else: + eps_y_alt = (x_0 - data_y) / sigma + eps_x_alt = (x_0 - data_x) / sigma + + if EO("flow_y_zero"): + eps_y_alt *= LG.mask + + eps_[row] = eps_yx = (eps_y_alt - eps_x_alt) + eps_y_lin = (x_0 - data_y) / sigma + if EO("flow_y_zero"): + eps_y_lin *= LG.mask + eps_x_lin = (x_0 - data_x) / sigma + eps_yx_lin = (eps_y_lin - eps_x_lin) + + data_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * data_x + (lgw_mask_ + lgw_mask_inv_) * data_y + + if EO("flow_reverse_data_masks"): + data_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * data_y + (lgw_mask_ + lgw_mask_inv_) * data_x + + if flow_sync_eps != 0.0: + if RK.EXPONENTIAL: + eps_[row] = (1-flow_sync_eps) * eps_[row] + flow_sync_eps * (data_[row] - x_0) + else: + eps_[row] = (1-flow_sync_eps) * eps_[row] + flow_sync_eps * (x_0 - data_[row]) / sigma + + if EO("flow_sync_eps_mask"): + flow_sync_eps = EO("flow_sync_eps_mask", 1.0) + if RK.EXPONENTIAL: + eps_[row] = (lgw_mask_ + lgw_mask_inv_) * (1-flow_sync_eps) * eps_[row] + (1 - (lgw_mask_ + lgw_mask_inv_)) * flow_sync_eps * (data_[row] - x_0) + else: + eps_[row] = (lgw_mask_ + lgw_mask_inv_) * (1-flow_sync_eps) * eps_[row] + (1 - (lgw_mask_ + lgw_mask_inv_)) * flow_sync_eps * (x_0 - data_[row]) / sigma + + if EO("flow_sync_eps_revmask"): + flow_sync_eps = EO("flow_sync_eps_revmask", 1.0) + if RK.EXPONENTIAL: + eps_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * (1-flow_sync_eps) * eps_[row] + (lgw_mask_ + lgw_mask_inv_) * flow_sync_eps * (data_[row] - x_0) + else: + eps_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * (1-flow_sync_eps) * eps_[row] + (lgw_mask_ + lgw_mask_inv_) * flow_sync_eps * (x_0 - data_[row]) / sigma + + if EO("flow_sync_eps_maskonly"): + flow_sync_eps = EO("flow_sync_eps_maskonly", 1.0) + if RK.EXPONENTIAL: + eps_[row] = (lgw_mask_ + lgw_mask_inv_) * eps_[row] + (1 - (lgw_mask_ + lgw_mask_inv_)) * (data_[row] - x_0) + else: + eps_[row] = (lgw_mask_ + lgw_mask_inv_) * eps_[row] + (1 - (lgw_mask_ + lgw_mask_inv_)) * (x_0 - data_[row]) / sigma + + if EO("flow_sync_eps_revmaskonly"): + flow_sync_eps = EO("flow_sync_eps_revmaskonly", 1.0) + if RK.EXPONENTIAL: + eps_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * eps_[row] + (lgw_mask_ + lgw_mask_inv_) * (data_[row] - x_0) + else: + eps_[row] = (1 - (lgw_mask_ + lgw_mask_inv_)) * eps_[row] + (lgw_mask_ + lgw_mask_inv_) * (x_0 - data_[row]) / sigma + + if EO("flow_slerp"): + if RK.EXPONENTIAL: + eps_y_alt = data_y - x_0 + eps_y_alt_inv = data_y_inv - x_0 + eps_x_alt = data_x - x_0 + else: + eps_y_alt = (x_0 - data_y) / sigma + eps_y_alt_inv = (x_0 - data_y_inv) / sigma + eps_x_alt = (x_0 - data_x) / sigma + + flow_slerp_ratio2 = EO("flow_slerp_ratio2", 0.5) + + eps_yx = (eps_y_alt - eps_x_alt) + eps_y_lin = (x_0 - data_y) / sigma + eps_x_lin = (x_0 - data_x) / sigma + eps_yx_lin = (eps_y_lin - eps_x_lin) + + eps_yx_inv = (eps_y_alt_inv - eps_x_alt) + eps_y_lin_inv = (x_0 - data_y_inv) / sigma + eps_x_lin = (x_0 - data_x) / sigma + eps_yx_lin_inv = (eps_y_lin_inv - eps_x_lin) + + data_row = x_0 - sigma * eps_yx_lin + data_row_inv = x_0 - sigma * eps_yx_lin_inv + + if EO("flow_slerp_similarity_ratio"): + flow_slerp_similarity_ratio = EO("flow_slerp_similarity_ratio", 1.0) + flow_slerp_ratio2 = find_slerp_ratio_grid(data_row, data_row_inv, LG.y0.clone(), LG.y0_inv.clone(), flow_slerp_similarity_ratio) + + eps_ [row] = slerp_tensor(flow_slerp_ratio2, eps_yx, eps_yx_inv) + data_[row] = slerp_tensor(flow_slerp_ratio2, data_row, data_row_inv) + + if EO("flow_slerp_autoalter"): + data_row_slerp = slerp_tensor(0.5, data_row, data_row_inv) + y0_pearsim = get_pearson_similarity(data_row_slerp, y0) + y0_pearsim_inv = get_pearson_similarity(data_row_slerp, y0_inv) + + if y0_pearsim > y0_pearsim_inv: + data_[row] = data_row_inv + eps_ [row] = (eps_y_alt_inv - eps_x_alt) + else: + data_[row] = data_row + eps_ [row] = (eps_y_alt - eps_x_alt) + + if EO("flow_slerp_recalc_eps_row"): + if RK.EXPONENTIAL: + eps_[row] = data_[row] - x_0 + else: + eps_[row] = (x_0 - data_[row]) / sigma + + if EO("flow_slerp_recalc_data_row"): + if RK.EXPONENTIAL: + data_[row] = x_0 + eps_[row] + else: + data_[row] = x_0 - sigma * eps_[row] + + data_cached = data_x + + if step < EO("direct_pre_pseudo_guide", 0) and step > 0: + for i_pseudo in range(EO("direct_pre_pseudo_guide_iter", 1)): + x_tmp += LG.lgw[step_sched] * LG.mask * (NS.sigma_max - s_tmp) * (LG.y0 - denoised) + LG.lgw_inv[step_sched] * LG.mask_inv * (NS.sigma_max - s_tmp) * (LG.y0_inv - denoised) + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma) + + # MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL MODEL CALL + + if SYNC_GUIDE_ACTIVE: + pass + elif not ((not LG.guide_mode.startswith("flow")) or FLOW_STOPPED or (LG.guide_mode.startswith("flow") and LG.lgw[step_sched] == 0 and LG.lgw_inv[step_sched] == 0)): #(LG.guide_mode.startswith("flow") and (LG.lgw[step_sched] != 0 or LG.lgw_inv[step_sched] != 0)) or FLOW_STOPPED: + pass + elif LG.guide_mode.startswith("lure") and (LG.lgw[step_sched] > 0 or LG.lgw_inv[step_sched] > 0): + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'yt'}) + + else: + if EO("protoshock") and StyleMMDiT is not None and StyleMMDiT.data_shock_start_step <= step_sched < StyleMMDiT.data_shock_end_step: + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'row': row, 'x_tmp': x_tmp, 'sigma_next': sigma_next}) + data_wct = StyleMMDiT.apply_data_shock(data_[row]) + if VE_MODEL: + x_tmp = x_tmp + (data_wct - data_[row]) + else: + x_tmp = x_tmp + (NS.sigma_max-NS.s_[row]) * (data_wct - data_[row]) + #x_[row+RK.row_offset] = x_tmp + x_[row] = x_tmp + if row == 0: + x_0 = x_tmp + + if EO("preshock"): + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'row': row, 'x_tmp': x_tmp, 'sigma_next': sigma_next}) + if VE_MODEL: + x_tmp = x_tmp + (data_wct - data_[row]) + else: + x_tmp = x_tmp + (NS.sigma_max-NS.s_[row]) * (data_wct - data_[row]) + x_[row] = x_tmp + if row == 0: + x_0 = x_tmp + + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'row': row, 'x_tmp': x_tmp, 'sigma_next': sigma_next}) + + #if EO("yoloshock") and StyleMMDiT is not None and StyleMMDiT.data_shock_start_step <= step_sched < StyleMMDiT.data_shock_end_step: + if not EO("disable_yoloshock") and StyleMMDiT is not None and StyleMMDiT.data_shock_start_step <= step_sched < StyleMMDiT.data_shock_end_step: + data_wct = StyleMMDiT.apply_data_shock(data_[row]) + if VE_MODEL: + x_tmp = x_tmp + (data_wct - data_[row]) + else: + x_tmp = x_tmp + (NS.sigma_max-NS.s_[row]) * (data_wct - data_[row]) + #x_[row+RK.row_offset] = x_tmp + x_[row] = x_tmp + if row == 0: + x_0 = x_tmp + data_[row] = data_wct + if RK.EXPONENTIAL: + eps_[row] = data_[row] - x_0 + else: + eps_[row] = (x_0 - data_[row]) / sigma + + + if hasattr(model.inner_model.inner_model.diffusion_model, "eps_out"): # fp64 model out override, for testing only + eps_out = model.inner_model.inner_model.diffusion_model.eps_out + del model.inner_model.inner_model.diffusion_model.eps_out + if eps_out.shape[0] == 2: + data_cond = x_0 - sigma * eps_out[1] + data_uncond = x_0 - sigma * eps_out[0] + data_row = data_uncond + model.inner_model.cfg * (data_cond - data_uncond) + eps_row = (x_0 - data_row) / sigma + else: + data_row = x_0 - sigma * eps_out + if RK.EXPONENTIAL: + eps_row = data_row - x_0 + else: + eps_row = eps_out + if torch.norm(eps_row - eps_[row]) < 0.01 and torch.norm(data_row - data_[row]) < 0.01: # if some other cfg/post-cfg func was used, detect and ignore this + eps_[row] = eps_row + data_[row] = data_row + + + if RK.extra_args['model_options']['transformer_options'].get('y0_standard_guide') is not None: + if hasattr(model.inner_model.inner_model.diffusion_model, "y0_standard_guide"): + LG.y0 = model.inner_model.inner_model.diffusion_model.y0_standard_guide.clone() + del model.inner_model.inner_model.diffusion_model.y0_standard_guide + RK.extra_args['model_options']['transformer_options']['y0_standard_guide'] = None + + if RK.extra_args['model_options']['transformer_options'].get('y0_inv_standard_guide') is not None: + if hasattr(model.inner_model.inner_model.diffusion_model, "y0_inv_standard_guide"): + LG.y0_inv = model.inner_model.inner_model.diffusion_model.y0_inv_standard_guide.clone() # RK.extra_args['model_options']['transformer_options'].get('y0_standard_guide') + del model.inner_model.inner_model.diffusion_model.y0_inv_standard_guide + RK.extra_args['model_options']['transformer_options']['y0_inv_standard_guide'] = None + + if LG.guide_mode.startswith("lure") and (LG.lgw[step_sched] > 0 or LG.lgw_inv[step_sched] > 0): + x_tmp = LG.process_guides_data_substep(x_tmp, data_[row], step_sched, s_tmp) + eps_[row], data_[row] = RK(x_tmp, s_tmp, x_0, sigma, transformer_options={'latent_type': 'xt'}) + + if momentum != 0.0: + data_[row] = data_[row] - momentum * (data_prev_[0] - data_[row]) #negative! + eps_[row] = RK.get_epsilon(x_0, x_tmp, data_[row], sigma, s_tmp) # ... why was this here??? for momentum maybe? + + if row < RK.rows and noise_scaling_weight != 0 and noise_scaling_type in {"sampler", "sampler_substep"}: + if noise_scaling_type == "sampler_substep": + sub_lying_su, sub_lying_sigma, sub_lying_sd, sub_lying_alpha_ratio = NS.get_sde_substep(NS.s_[row], NS.s_[row+RK.row_offset+RK.multistep_stages], noise_scaling_eta, noise_scaling_mode) + for _ in range(noise_scaling_cycles-1): + sub_lying_su, sub_lying_sigma, sub_lying_sd, sub_lying_alpha_ratio = NS.get_sde_substep(NS.s_[row], sub_lying_sd, noise_scaling_eta, noise_scaling_mode) + lying_s_[row+1] = sub_lying_sd + substep_noise_scaling_ratio = NS.s_[row+1]/lying_s_[row+1] + if RK.multistep_stages > 0: + substep_noise_scaling_ratio = sigma_next/lying_sd #fails with resample? + + lying_eps_row_factor = (1 - noise_scaling_weight*(substep_noise_scaling_ratio-1)) + + # GUIDE + if not EO("disable_guides_eps_substep"): + eps_, x_ = LG.process_guides_substep(x_0, x_, eps_, data_, denoised_prev, row, step, step_sched, NS.sigma, NS.sigma_next, NS.sigma_down, NS.s_, epsilon_scale, RK, full_iter) + if not EO("disable_guides_eps_prev_substep"): + eps_prev_, x_ = LG.process_guides_substep(x_0, x_, eps_prev_, data_, denoised_prev, row, step, step_sched, NS.sigma, NS.sigma_next, NS.sigma_down, NS.s_, epsilon_scale, RK, full_iter) + + if LG.y0_mean is not None and LG.y0_mean.sum() != 0.0: + if x.ndim == 3: # packed NestedTensor + raise NotImplementedError("y0_mean guide requires spatial structure, incompatible with packed latents") + + if EO("guide_mean_scattersort"): + data_row_mean = apply_scattersort_spatial(data_[row], LG.y0_mean) + eps_row_mean = RK.get_eps(x_0, data_row_mean, s_tmp) + else: + eps_row_mean = eps_[row] - eps_[row].mean(dim=(-2,-1), keepdim=True) + (LG.y0_mean - x_0).mean(dim=(-2,-1), keepdim=True) + + if LG.mask_mean is not None: + eps_row_mean = LG.mask_mean * eps_row_mean + (1-LG.mask_mean) * eps_[row] + + eps_[row] = eps_[row] + LG.lgw_mean[step_sched] * (eps_row_mean - eps_[row]) + + if (full_iter == 0 and diag_iter == 0) or EO("newton_iter_post_use_on_implicit_steps"): + x_, eps_ = RK.newton_iter(x_0, x_, eps_, eps_prev_, data_, NS.s_, row, NS.h, sigmas, step, "post", SYNC_GUIDE_ACTIVE) + + # UPDATE #for row in range(RK.rows - RK.multistep_stages - RK.row_offset + 1): + if EO("exp2lin_override") and RK.EXPONENTIAL: + x_ = RK.update_substep(x_0, x_, eps_, eps_prev_, row, RK.row_offset, NS.h_new, NS.h_new_orig, lying_eps_row_factor=lying_eps_row_factor, sigma=sigma) #modifies eps_[row] if lying_eps_row_factor != 1.0 + #x_ = RK.update_substep(x_0, x_, eps_, eps_prev_, row, RK.row_offset, -sigma*NS.h_new, -sigma*NS.h_new_orig, lying_eps_row_factor=lying_eps_row_factor) #modifies eps_[row] if lying_eps_row_factor != 1.0 + else: + x_ = RK.update_substep(x_0, x_, eps_, eps_prev_, row, RK.row_offset, NS.h_new, NS.h_new_orig, lying_eps_row_factor=lying_eps_row_factor) #modifies eps_[row] if lying_eps_row_factor != 1.0 + + x_[row+RK.row_offset] = NS.rebound_overshoot_substep(x_0, x_[row+RK.row_offset]) + + if SYNC_GUIDE_ACTIVE: #yt_ is not None: + #yt_ = RK.update_substep(yt_0, yt_, eps_y_, eps_prev_y_, row, RK.row_offset, NS.h_new, NS.h_new_orig, lying_eps_row_factor=lying_eps_row_factor) #modifies eps_[row] if lying_eps_row_factor != 1.0 + yt_ = RK.update_substep(yt_0, yt_, eps_yt_, eps_prev_y_, row, RK.row_offset, NS.h_new, NS.h_new_orig, lying_eps_row_factor=lying_eps_row_factor, sigma=sigma) #modifies eps_[row] if lying_eps_row_factor != 1.0 + yt_[row+RK.row_offset] = NS.rebound_overshoot_substep(yt_0, yt_[row+RK.row_offset]) + + if not RK.IMPLICIT and NS.noise_mode_sde_substep != "hard_sq": + + if not LG.guide_mode.startswith("flow") or (LG.lgw[step_sched] == 0 and LG.lgw[step+1] == 0 and LG.lgw_inv[step_sched] == 0 and LG.lgw_inv[step+1] == 0): + #if LG.guide_mode.startswith("sync") and (LG.lgw[step_sched] != 0.0 or LG.lgw_inv[step_sched] != 0.0): + # x_row_tmp = x_[row+RK.row_offset].clone() + + #x_[row+RK.row_offset] = NS.swap_noise_substep(x_0, x_[row+RK.row_offset], mask=sde_mask, guide=LG.y0) + x_row_tmp = NS.swap_noise_substep(x_0, x_[row+RK.row_offset], mask=sde_mask, guide=LG.y0) + + #if EO("eps_adain_smartnoise_substep"): + if LG.ADAIN_NOISE_MODE == "smart": + #eps_row_next = (x_0 - x_[row+RK.row_offset]) / (sigma - NS.s_[row+RK.row_offset]) + #denoised_row_next = x_0 - sigma * eps_row_next + # + #eps_swapped = (x_row_tmp - denoised_row_next) / NS.s_[row+RK.row_offset] + # + #noise_row_next = eps_swapped + denoised_row_next + #z_[row+RK.row_offset] = noise_row_next + #RK.update_transformer_options({'z_' : z_}) + data_next = denoised + NS.h_new * RK.zum(row+RK.row_offset+RK.multistep_stages, data_, data_prev_) + if VE_MODEL: + z_[row+RK.row_offset] = (x_row_tmp - data_next) / NS.s_[row+RK.row_offset] + else: + z_[row+RK.row_offset] = (x_row_tmp - (NS.sigma_max-NS.s_[row+RK.row_offset])*data_next) / NS.s_[row+RK.row_offset] + RK.update_transformer_options({'z_' : z_}) + + elif LG.ADAIN_NOISE_MODE == "update": #EO("eps_adain"): + x_init_new = (x_row_tmp - x_[row+RK.row_offset]) / s_tmp + x_init + x_0 += sigma * (x_init_new - x_init) + x_init = x_init_new + RK.update_transformer_options({'x_init': x_init}) + + if SYNC_GUIDE_ACTIVE: + noise_bongflow_new = (x_row_tmp - x_[row+RK.row_offset]) / s_tmp + noise_bongflow + yt_[row+RK.row_offset] += s_tmp * (noise_bongflow_new - noise_bongflow) + x_0 += sigma * (noise_bongflow_new - noise_bongflow) + noise_bongflow = noise_bongflow_new + + x_[row+RK.row_offset] = x_row_tmp + + elif LG.guide_mode.startswith("flow"): + pass + + if not LG.guide_mode.startswith("lure"): + x_[row+RK.row_offset] = LG.process_guides_data_substep(x_[row+RK.row_offset], data_[row], step_sched, NS.s_[row]) + + if ((not EO("protoshock") and not EO("yoloshock")) or EO("fuckitshock")) and StyleMMDiT is not None and StyleMMDiT.data_shock_start_step <= step_sched < StyleMMDiT.data_shock_end_step: + data_wct = StyleMMDiT.apply_data_shock(data_[row]) + if VE_MODEL: + x_[row+RK.row_offset] = x_[row+RK.row_offset] + (data_wct - data_[row]) + else: + x_[row+RK.row_offset] = x_[row+RK.row_offset] + (NS.sigma_max-NS.s_[row]) * (data_wct - data_[row]) + + + if SYNC_GUIDE_ACTIVE: # # # # ## # # ## # YIIIIKES --------------------------------------------------------------------------------------------------------- + if VE_MODEL: + yt_[:NS.s_.shape[0], 0] = y0_bongflow + NS.s_.view(-1, *[1]*(x.ndim-1)) * (noise_bongflow) + yt_0 = y0_bongflow + sigma * (noise_bongflow) + else: + yt_[:NS.s_.shape[0], 0] = y0_bongflow + NS.s_.view(-1, *[1]*(x.ndim-1)) * (noise_bongflow - y0_bongflow) + yt_0 = y0_bongflow + sigma * (noise_bongflow - y0_bongflow) + if RK.EXPONENTIAL: + eps_y_ = data_y_ - yt_0 # yt_ # watch out for fuckery with size of tableau being smaller later in a chained sampler + else: + if BONGMATH: + eps_y_[:NS.s_.shape[0]] = (yt_[:NS.s_.shape[0]] - data_y_[:NS.s_.shape[0]]) / NS.s_.view(-1,*[1]*(x_.ndim-1)) + else: + eps_y_[:NS.s_.shape[0]] = (yt_0.repeat(NS.s_.shape[0], *[1]*(x_.ndim-1)) - data_y_[:NS.s_.shape[0]]) / sigma # calc exact to c0 node + if not BONGMATH and (eta != 0 or eta_substep != 0): + if RK.EXPONENTIAL: + eps_x_ = data_x_ - x_0 + else: + eps_x_ = (x_0 - data_x_) / sigma + + weight_mask = lgw_mask_+lgw_mask_inv_ + if LG.SYNC_SEPARATE: + sync_mask = lgw_mask_sync_+lgw_mask_sync_inv_ + else: + sync_mask = 1. + + for ms in range(len(eps_)): + if RK.EXPONENTIAL: + if VE_MODEL: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + sigma*(-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + sigma*(-noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + sigma*(y0_bongflow-noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + sigma*(y0_bongflow-noise_bongflow)) + else: + if VE_MODEL: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + (noise_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + (noise_bongflow)) + else: + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_y_[ms] + (noise_bongflow-y0_bongflow)) + if EO("sync_x2y"): + eps_[ms] = sync_mask * eps_x_[ms] + (1-sync_mask) * eps_x2y_[ms] + weight_mask * (-eps_x2y_[ms] + (noise_bongflow-y0_bongflow)) + + if BONGMATH and NS.s_[row] > RK.sigma_min and NS.h < RK.sigma_max/2 and (diag_iter == implicit_steps_diag or EO("enable_diag_explicit_bongmath_all")) and not EO("disable_terminal_bongmath"): + if step == 0 and UNSAMPLE: + pass + elif full_iter == implicit_steps_full or not EO("disable_fully_explicit_bongmath_except_final"): + if sigma > 0.03: + BONGMATH_Y = SYNC_GUIDE_ACTIVE + x_0, x_, eps_ = RK.bong_iter(x_0, x_, eps_, eps_prev_, data_, sigma, NS.s_, row, RK.row_offset, NS.h, step, step_sched, + BONGMATH_Y, y0_bongflow, noise_bongflow, eps_x_, eps_y_, data_x_, data_y_, LG) + # BONGMATH_Y, y0_bongflow, noise_bongflow, eps_x_, eps_y_, eps_x2y_, data_x_, LG) + #if EO("eps_adain_smartnoise_bongmath"): + if LG.ADAIN_NOISE_MODE == "smart": + if VE_MODEL: + z_[:NS.s_.shape[0], ...] = (x_ - data_)[:NS.s_.shape[0], ...] / NS.s_.view(-1,*[1]*(x_.ndim-1)) + else: + z_[:NS.s_.shape[0], ...] = (x_[:NS.s_.shape[0], ...] - (NS.sigma_max - NS.s_.view(-1,*[1]*(x_.ndim-1)))*data_[:NS.s_.shape[0], ...])[:NS.s_.shape[0], ...] / NS.s_.view(-1,*[1]*(x_.ndim-1)) + RK.update_transformer_options({'z_' : z_}) + diag_iter += 1 + + #progress_bar.update( round(1 / implicit_steps_total, 2) ) + + #step_update = round(1 / implicit_steps_total, 2) + #progress_bar.update(float(f"{step_update:.2f}")) + + x_next = x_[RK.rows - RK.multistep_stages - RK.row_offset + 1] + x_next = NS.rebound_overshoot_step(x_0, x_next) + + if SYNC_GUIDE_ACTIVE: # YT_NEXT UPDATE STEP -------------------------------------- + yt_next = yt_[RK.rows - RK.multistep_stages - RK.row_offset + 1] + yt_next = NS.rebound_overshoot_step(yt_0, yt_next) + + eps = (x_0 - x_next) / (sigma - sigma_next) + denoised = x_0 - sigma * eps + + if EO("postshock") and step < EO("postshock", 10): + eps_row, data_row = RK(x_next, sigma_next, x_next, sigma_next, transformer_options={'row': row, 'x_tmp': x_next, 'sigma_next': sigma_next}) + if VE_MODEL: + x_next = x_next + (data_row - denoised) + else: + x_next = x_next + (NS.sigma_max-sigma_next) * (data_row - denoised) + eps = (x_0 - x_next) / (sigma - sigma_next) + denoised = x_0 - sigma * eps + + if EO("data_sampler") and step > EO("data_sampler_start_step", 0) and step < EO("data_sampler_end_step", 5): + data_sampler_weight = EO("data_sampler_weight", 1.0) + denoised_step = RK.zum(row+RK.row_offset+RK.multistep_stages, data_, data_prev_) + x_next = LG.swap_data(x_next, denoised, denoised_step, data_sampler_weight * sigma_next) + eps = (x_0 - x_next) / (sigma - sigma_next) + denoised = x_0 - sigma * eps + + x_0_prev = x_0.clone() + + if eta == 0.0: + x = x_next + if SYNC_GUIDE_ACTIVE: + yt_0 = yt_[0] = yt_next + #elif LG.guide_mode.startswith("sync") and (LG.lgw[step_sched] != 0.0 or LG.lgw_inv[step_sched] != 0.0): + # noise_sync_new = NS.noise_sampler(sigma=sigma, sigma_next=sigma_next) + # x = x_next + sigma * eta * (noise_sync_new - noise_bongflow) + # noise_bongflow += eta * (noise_sync_new - noise_bongflow) + elif not LG.guide_mode.startswith("flow") or (LG.lgw[step_sched] == 0 and LG.lgw[step+1] == 0 and LG.lgw_inv[step_sched] == 0 and LG.lgw_inv[step+1] == 0): + x = NS.swap_noise_step(x_0, x_next, mask=sde_mask) + + #if EO("eps_adain_smartnoise"): + if LG.ADAIN_NOISE_MODE == "smart": + #noise_next = eps + denoised + #eps_swapped = (x - denoised) / sigma_next + # + #noise_next = eps_swapped + denoised + #z_[0] = noise_next + #RK.update_transformer_options({'z_' : z_}) + if full_iter+1 < implicit_steps_full+1: # are we to loop for full iter after this? + if VE_MODEL: + #z_[row+RK.row_offset] = (x - denoised) / sigma_next + z_[0] = (x_0 - denoised) / sigma + else: + #z_[row+RK.row_offset] = (x - (NS.sigma_max-sigma_next) * denoised) / sigma_next + z_[0] = (x_0 - (NS.sigma_max-sigma) * denoised) / sigma + else: #we're advancing to next step, x is x_next + if VE_MODEL: + #z_[row+RK.row_offset] = (x - denoised) / sigma_next + z_[0] = (x - denoised) / sigma_next + else: + #z_[row+RK.row_offset] = (x - (NS.sigma_max-sigma_next) * denoised) / sigma_next + z_[0] = (x - (NS.sigma_max-sigma_next) * denoised) / sigma_next + RK.update_transformer_options({'z_' : z_}) + + elif LG.ADAIN_NOISE_MODE == "update": #EO("eps_adain"): + x_init_new = (x - x_next) / sigma_next + x_init + x_0 += sigma * (x_init_new - x_init) + x_init = x_init_new + RK.update_transformer_options({'x_init': x_init}) + + if SYNC_GUIDE_ACTIVE: + noise_bongflow_new = (x - x_next) / sigma_next + noise_bongflow + yt_next += sigma_next * (noise_bongflow_new - noise_bongflow) + x_0 += sigma * (noise_bongflow_new - noise_bongflow) + if not EO("disable_i_bong"): + for i_bong in range(len(NS.s_)): + x_[i_bong] += NS.s_[i_bong] * (noise_bongflow_new - noise_bongflow) + #x_[0] += sigma * (noise_bongflow_new - noise_bongflow) + yt_0 = yt_[0] = yt_next + noise_bongflow = noise_bongflow_new + else: + x = x_next + + if EO("keep_step_means"): + if x.ndim == 3: # packed NestedTensor + raise NotImplementedError("keep_step_means requires spatial structure, incompatible with packed latents") + x = x - x.mean(dim=(-2,-1), keepdim=True) + x_means_per_step + + + callback_step = display_callback_step(step, sampler_mode, len(sigmas), outer_sigmas_len) + self_refine_mask = getattr(LG, '_debug_certainty_mask', None) + preview_callback(x, eps, denoised, x_, eps_, data_, callback_step, sigma, sigma_next, callback, EO, preview_override=data_cached, FLOW_STOPPED=FLOW_STOPPED, device=model_device, self_refine_mask=self_refine_mask, step_sched=step) + + h_prev = NS.h + x_prev = x_0 + + denoised_prev2 = denoised_prev + denoised_prev = denoised + + full_iter += 1 + + if getattr(LG, '_self_refine_converged', False): + break + + if LG.lgw[step_sched] > 0 and step >= EO("guide_cutoff_start_step", 0) and cossim_counter < EO("guide_cutoff_max_iter", 10) and (EO("guide_cutoff") or EO("guide_min")): + if x.ndim == 3: # packed NestedTensor + raise NotImplementedError("guide_cutoff/guide_min requires spatial structure, incompatible with packed latents") + guide_cutoff = EO("guide_cutoff", 1.0) + denoised_norm = data_[0] - data_[0].mean(dim=(-2,-1), keepdim=True) + y0_norm = LG.y0 - LG.y0 .mean(dim=(-2,-1), keepdim=True) + y0_cossim = get_cosine_similarity(denoised_norm, y0_norm) + if y0_cossim > guide_cutoff and LG.lgw[step_sched] > EO("guide_cutoff_floor", 0.0): + if not EO("guide_cutoff_fast"): + LG.lgw[step_sched] *= EO("guide_cutoff_factor", 0.9) + else: + LG.lgw *= EO("guide_cutoff_factor", 0.9) + full_iter -= 1 + if y0_cossim < EO("guide_min", 0.0) and LG.lgw[step_sched] < EO("guide_min_ceiling", 1.0): + if not EO("guide_cutoff_fast"): + LG.lgw[step_sched] *= EO("guide_min_factor", 1.1) + else: + LG.lgw *= EO("guide_min_factor", 1.1) + full_iter -= 1 + + #if EO("smartnoise"): #TODO: determine if this was useful + # z_[0] = z_next + + if FLOW_STARTED and FLOW_STOPPED: + data_prev_ = data_x_prev_ + if FLOW_STARTED and not FLOW_STOPPED: + data_x_prev_[0] = data_cached # data_cached is data_x from flow mode. this allows multistep to resume seamlessly. + for ms in range(recycled_stages): + data_x_prev_[recycled_stages - ms] = data_x_prev_[recycled_stages - ms - 1] + + #if LG.guide_mode.startswith("sync") and (LG.lgw[step_sched] != 0.0 or LG.lgw_inv[step_sched] != 0.0): + # data_prev_[0] = x_0 - sigma * eps_[0] + #else: + data_prev_[0] = data_[0] # with flow mode, this will be the differentiated guide/"denoised" + for ms in range(recycled_stages): + data_prev_[recycled_stages - ms] = data_prev_[recycled_stages - ms - 1] # TODO: verify that this does not run on every substep... + + if SYNC_GUIDE_ACTIVE: + data_prev_x_[0] = data_x + for ms in range(recycled_stages): + data_prev_x_[recycled_stages - ms] = data_prev_x_[recycled_stages - ms - 1] + + data_prev_y_[0] = data_y + for ms in range(recycled_stages): + data_prev_y_[recycled_stages - ms] = data_prev_y_[recycled_stages - ms - 1] + + if EO("bong2m") or EO("bong3m"): + denoised_data_prev2 = denoised_data_prev + denoised_data_prev = data_[0] + + if SKIP_PSEUDO and not LG.guide_mode.startswith("flow"): + if SKIP_PSEUDO_Y == "y0": + LG.y0 = denoised.clone() + LG.HAS_LATENT_GUIDE = True + else: + LG.y0_inv = denoised.clone() + LG.HAS_LATENT_GUIDE_INV = True + + if EO("pseudo_mix_strength"): + pseudo_mix_strength = EO("pseudo_mix_strength", 0.0) + LG.y0 = orig_y0 + pseudo_mix_strength * (denoised - orig_y0) + LG.y0_inv = orig_y0_inv + pseudo_mix_strength * (denoised - orig_y0_inv) + + #if sampler_mode == "unsample": + # progress_bar.n -= 1 + # progress_bar.refresh() + #else: + # progress_bar.update(1) + progress_bar.update(1) #THIS WAS HERE + step += 1 + + if EO("skip_step", -1) == step: + step += 1 + + if d_noise_start_step == step: + sigmas = sigmas.clone() * d_noise + if sigmas.max() > NS.sigma_max: + sigmas = sigmas / NS.sigma_max + if d_noise_inv_start_step == step: + sigmas = sigmas.clone() / d_noise_inv + if sigmas.max() > NS.sigma_max: + sigmas = sigmas / NS.sigma_max + + if LG.lgw[step_sched] > 0 and step >= EO("guide_step_cutoff_start_step", 0) and cossim_counter < EO("guide_step_cutoff_max_iter", 10) and (EO("guide_step_cutoff") or EO("guide_step_min")): + if x.ndim == 3: # packed NestedTensor + raise NotImplementedError("guide_step_cutoff/guide_step_min requires spatial structure, incompatible with packed latents") + guide_cutoff = EO("guide_step_cutoff", 1.0) + eps_trash, data_trash = RK(x, sigma_next, x_0, sigma) + denoised_norm = data_trash - data_trash.mean(dim=(-2,-1), keepdim=True) + y0_norm = LG.y0 - LG.y0 .mean(dim=(-2,-1), keepdim=True) + y0_cossim = get_cosine_similarity(denoised_norm, y0_norm) + if y0_cossim > guide_cutoff and LG.lgw[step_sched] > EO("guide_step_cutoff_floor", 0.0): + if not EO("guide_step_cutoff_fast"): + LG.lgw[step_sched] *= EO("guide_step_cutoff_factor", 0.9) + else: + LG.lgw *= EO("guide_step_cutoff_factor", 0.9) + step -= 1 + x_0 = x = x_[0] = x_0_orig.clone() + if y0_cossim < EO("guide_step_min", 0.0) and LG.lgw[step_sched] < EO("guide_step_min_ceiling", 1.0): + if not EO("guide_step_cutoff_fast"): + LG.lgw[step_sched] *= EO("guide_step_min_factor", 1.1) + else: + LG.lgw *= EO("guide_step_min_factor", 1.1) + step -= 1 + x_0 = x = x_[0] = x_0_orig.clone() + # END SAMPLING LOOP --------------------------------------------------------------------------------------------------- + + #progress_bar.close() + RK.update_transformer_options({'update_cross_attn': None}) + if step == len(sigmas)-2 and sigmas[-1] == 0 and sigmas[-2] == NS.sigma_min and not INIT_SAMPLE_LOOP: + if EO("enable_final_model_call"): + eps, denoised = RK(x, NS.sigma_min, x, NS.sigma_min) + x = denoised + else: + sigma_min = NS.sigma_min.view((1,) * x.ndim).to(x) + denoised = model.inner_model.inner_model.model_sampling.calculate_denoised(sigma_min, eps, x) + x = denoised + + eps = eps .to(model_device) + denoised = denoised.to(model_device) + x = x .to(model_device) + + progress_bar.close() + + state_info_out['model_call_counts'] = RK.get_model_call_counters() + RESplain("Model calls (total/denoised/epsilon):", state_info_out['model_call_counts'], debug=True) + + # re-report only at the schedule end, where it carries the progress bar to 100%; + # on partial runs it would just duplicate the loop's last report one step later + schedule_end_step = len(sigmas) - 2 if sigmas[-1] == 0 else len(sigmas) - 1 + if step >= schedule_end_step and not (UNSAMPLE and sigmas[1] > sigmas[0]) and not EO("preview_last_step_always") and sigma is not None and not (FLOW_STARTED and not FLOW_STOPPED): + callback_step = display_callback_step(step, sampler_mode, len(sigmas), outer_sigmas_len) + self_refine_mask = getattr(LG, '_debug_certainty_mask', None) + preview_callback(x, eps, denoised, x_, eps_, data_, callback_step, sigma, sigma_next, callback, EO, device=model_device, self_refine_mask=self_refine_mask, step_sched=step, final=True) + + # judged on the schedule the loop ran on — the same axis as step/end_step + RUN_COMPLETED = step == len(sigmas)-2 and sigmas[-1] == 0 and sigmas[-2] == NS.sigma_min + + if INIT_SAMPLE_LOOP: + state_info_out.update(state_info) + else: + if guides is not None and guides.get('guide_mode', "") == 'inversion': + guide_inversion_y0 = state_info.get('guide_inversion_y0') + guide_inversion_y0_inv = state_info.get('guide_inversion_y0_inv') + + if sampler_mode == "unsample" and guide_inversion_y0 is None: + guide_inversion_y0 = LG.y0.clone() + if sampler_mode == "unsample" and guide_inversion_y0_inv is None: + guide_inversion_y0_inv = LG.y0_inv.clone() + + if sampler_mode in {"standard", "resample"} and guide_inversion_y0 is None: + guide_inversion_y0 = NS.noise_sampler(sigma=NS.sigma_max, sigma_next=NS.sigma_min).to(x) + guide_inversion_y0 = normalize_zscore(guide_inversion_y0, channelwise=True, inplace=True) + if sampler_mode in {"standard", "resample"} and guide_inversion_y0_inv is None: + guide_inversion_y0_inv = NS.noise_sampler(sigma=NS.sigma_max, sigma_next=NS.sigma_min).to(x) + guide_inversion_y0_inv = normalize_zscore(guide_inversion_y0_inv, channelwise=True, inplace=True) + + state_info_out['guide_inversion_y0'] = guide_inversion_y0 + state_info_out['guide_inversion_y0_inv'] = guide_inversion_y0_inv + + state_info_out['raw_x'] = x.to('cpu') + state_info_out['denoised'] = denoised.to('cpu') + state_info_out['data_prev_'] = data_prev_.to('cpu') + state_info_out['end_step'] = step + state_info_out['sigma_next'] = sigma_next.clone() + state_info_out['sigmas'] = sigmas_scheduled.clone() + state_info_out['sampler_mode'] = sampler_mode + state_info_out['last_rng'] = NS.noise_sampler .generator.get_state().clone() + state_info_out['last_rng_substep'] = NS.noise_sampler2.generator.get_state().clone() + state_info_out['completed'] = RUN_COMPLETED + state_info_out['FLOW_STARTED'] = FLOW_STARTED + state_info_out['FLOW_STOPPED'] = FLOW_STOPPED + state_info_out['noise_bongflow'] = noise_bongflow.clone().cpu() if noise_bongflow is not None else None + state_info_out['y0_bongflow'] = y0_bongflow.clone().cpu() if y0_bongflow is not None else None + state_info_out['y0_bongflow_orig'] = y0_bongflow_orig.clone().cpu() if y0_bongflow_orig is not None else None + state_info_out['y0_standard_guide'] = y0_standard_guide.clone().cpu() if y0_standard_guide is not None else None + state_info_out['y0_inv_standard_guide'] = y0_inv_standard_guide.clone().cpu() if y0_inv_standard_guide is not None else None + state_info_out['data_prev_y_'] = data_prev_y_.clone().cpu() if data_prev_y_ is not None else None + state_info_out['data_prev_x_'] = data_prev_x_.clone().cpu() if data_prev_x_ is not None else None + + if noise_initial is not None: + state_info_out['noise_initial'] = noise_initial.to('cpu') + if image_initial is not None: + state_info_out['image_initial'] = image_initial.to('cpu') + + if FLOW_STARTED and not FLOW_STOPPED: + state_info_out['y0'] = y0.to('cpu') + #state_info_out['y0_inv'] = y0_inv.to('cpu') # TODO: implement this? + state_info_out['data_cached'] = data_cached.to('cpu') + state_info_out['data_x_prev_'] = data_x_prev_.to('cpu') + + if REPORT_VRAM: + RESplain(f"report_vram: peak allocated {torch.cuda.max_memory_allocated() / 1024**3:.2f} GB, " + f"peak reserved {torch.cuda.max_memory_reserved() / 1024**3:.2f} GB (torch allocator only; model weights under dynamic VRAM are not included)", debug=False) + + return x + +def noise_fn(x, sigma, sigma_next, noise_sampler, cossim_iter=1): + + noise = normalize_zscore(noise_sampler(sigma=sigma, sigma_next=sigma_next), channelwise=True, inplace=True) + cossim = get_pearson_similarity(x, noise) + + for i in range(cossim_iter): + noise_new = normalize_zscore(noise_sampler(sigma=sigma, sigma_next=sigma_next), channelwise=True, inplace=True) + cossim_new = get_pearson_similarity(x, noise_new) + + if cossim_new > cossim: + noise = noise_new + cossim = cossim_new + + return noise + + +def resolve_start_step_from_sigma_next(sigmas: Tensor, sigma_next, fallback: int) -> int: + """Resume index from matching the stored sigma_next on the schedule. Schedules can hold + duplicate sigmas (restarts): on multiple matches, take the closest to the fallback.""" + if isinstance(sigma_next, torch.Tensor): + matches = torch.isclose(sigmas, sigma_next.to(sigmas), rtol=1e-5, atol=1e-8).nonzero().flatten() + else: + matches = (sigmas == sigma_next).nonzero().flatten() + if matches.shape[0] == 0: + return fallback + if matches.shape[0] == 1: + return int(matches.item()) + return int(min(matches.tolist(), key=lambda idx: abs(idx - fallback))) + + +def display_callback_step(step: int, sampler_mode: str, sigmas_len: int, outer_sigmas_len: int) -> int: + """Map a loop step onto the axis the progress bar was sized on: unsample counts + backwards, resample stretches onto the longer padded outer axis (lossy rounding), + standard passes through. Display only — 'i_sched' carries the true index.""" + if sampler_mode == "unsample": + return sigmas_len - 1 - step + if sampler_mode == "resample" and outer_sigmas_len > sigmas_len: + return int(round(step / (sigmas_len - 1) * (outer_sigmas_len - 1))) + return step + + +def preview_callback( + x : Tensor, + eps : Tensor, + denoised : Tensor, + x_ : Tensor, + eps_ : Tensor, + data_ : Tensor, + step : int, + sigma : Tensor, + sigma_next : Tensor, + callback : Callable, + EO : ExtraOptions, + preview_override : Optional[Tensor] = None, + FLOW_STOPPED : bool = False, + device : Optional[torch.device] = None, + self_refine_mask : Optional[Tensor] = None, + step_sched : Optional[int] = None, + final : bool = False,): + + if EO("eps_substep_preview"): + row_callback = EO("eps_substep_preview", 0) + denoised_callback = eps_[row_callback] + + elif EO("denoised_substep_preview"): + row_callback = EO("denoised_substep_preview", 0) + denoised_callback = data_[row_callback] + + elif EO("x_substep_preview"): + row_callback = EO("x_substep_preview", 0) + denoised_callback = x_[row_callback] + + elif EO("eps_preview"): + denoised_callback = eps + + elif EO("denoised_preview"): + denoised_callback = denoised + + elif EO("x_preview"): + denoised_callback = x + + elif preview_override is not None and FLOW_STOPPED == False: + denoised_callback = preview_override + + else: + denoised_callback = data_[0] + + # Overlay self-refine certainty mask on preview + if EO("self_refine_mask_preview") and self_refine_mask is not None: + mask_mode = EO("self_refine_mask_preview_mode", "zero") + denoised_callback = denoised_callback.clone() + + # Expand mask to match data channels if needed + if self_refine_mask.shape != denoised_callback.shape: + if self_refine_mask.shape[1] == 1 and denoised_callback.shape[1] > 1: + mask_expanded = self_refine_mask.expand_as(denoised_callback) + else: + mask_expanded = self_refine_mask + else: + mask_expanded = self_refine_mask + + if mask_mode == "invert": + denoised_callback = denoised_callback * (1 - 2 * mask_expanded) + + elif mask_mode == "zero": + denoised_callback = denoised_callback * (1 - mask_expanded) + + + if device is not None: + denoised_callback = denoised_callback.to(device) + + # 'i' drives comfy's progress bar; 'i_sched' is the true schedule index (state_info axis); + # 'final' marks the post-loop re-report + if callback is not None: + callback({'x': x, 'i': step, 'i_sched': step if step_sched is None else step_sched, 'final': final, 'sigma': sigma, 'sigma_next': sigma_next, 'denoised': denoised_callback.to(torch.float32)}) + + return + diff --git a/simple_syrup/third_party/res4lyf_runtime/helper.py b/simple_syrup/third_party/res4lyf_runtime/helper.py new file mode 100644 index 0000000..b8ffdf0 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/helper.py @@ -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 + + diff --git a/simple_syrup/third_party/res4lyf_runtime/latents.py b/simple_syrup/third_party/res4lyf_runtime/latents.py new file mode 100644 index 0000000..4b2095a --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/latents.py @@ -0,0 +1,1095 @@ +import torch +import torch.nn.functional as F +from typing import Tuple, List, Union +import math +from .res4lyf import RESplain + +import comfy.utils + +# TENSOR PROJECTION OPS + +def get_cosine_similarity_manual(a, b): + return (a * b).sum() / (torch.norm(a) * torch.norm(b)) + +def get_cosine_similarity(a, b, mask=None, dim=0): + if a.ndim == 5 and b.ndim == 5 and b.shape[2] == 1: + b = b.expand(-1, -1, a.shape[2], -1, -1) + + if mask is not None: + return F.cosine_similarity((mask * a).flatten(), (mask * b).flatten(), dim=dim) + else: + return F.cosine_similarity(a.flatten(), b.flatten(), dim=dim) + +def get_pearson_similarity(a, b, mask=None, dim=0, norm_dim=None): + if a.ndim == 5 and b.ndim == 5 and b.shape[2] == 1: + b = b.expand(-1, -1, a.shape[2], -1, -1) + + if norm_dim is None: + if a.ndim == 3: # [1,1,N] flat tensor + norm_dim = -1 + elif a.ndim == 4: + norm_dim=(-2,-1) + elif a.ndim == 5: + norm_dim=(-4,-2,-1) + + a = a - a.mean(dim=norm_dim, keepdim=True) + b = b - b.mean(dim=norm_dim, keepdim=True) + + if mask is not None: + return F.cosine_similarity((mask * a).flatten(), (mask * b).flatten(), dim=dim) + else: + return F.cosine_similarity(a.flatten(), b.flatten(), dim=dim) + + + +def get_collinear(x, y): + return get_collinear_flat(x, y).reshape_as(x) + +def get_orthogonal(x, y): + x_flat = x.reshape(x.size(0), -1).clone() + x_ortho_y = x_flat - get_collinear_flat(x, y) + return x_ortho_y.view_as(x) + +def get_collinear_flat(x, y): + + y_flat = y.reshape(y.size(0), -1).clone() + x_flat = x.reshape(x.size(0), -1).clone() + + y_flat /= y_flat.norm(dim=-1, keepdim=True) + x_proj_y = torch.sum(x_flat * y_flat, dim=-1, keepdim=True) * y_flat + + return x_proj_y + + + +def get_orthogonal_noise_from_channelwise(*refs, max_iter=500, max_score=1e-15): + noise, *refs = refs + noise_tmp = noise.clone() + #b,c,h,w = noise.shape + if (noise.ndim == 4): + b,ch,h,w = noise.shape + elif (noise.ndim == 5): + b,ch,t,h,w = noise.shape + + for i in range(max_iter): + noise_tmp = gram_schmidt_channels_optimized(noise_tmp, *refs) + + cossim_scores = [] + for ref in refs: + #for c in range(noise.shape[-3]): + for c in range(ch): + cossim_scores.append(get_cosine_similarity(noise_tmp[0][c], ref[0][c]).abs()) + cossim_scores.append(get_cosine_similarity(noise_tmp[0], ref[0]).abs()) + + if max(cossim_scores) < max_score: + break + + return noise_tmp + + + +def gram_schmidt_channels_optimized(A, *refs): + if (A.ndim == 4): + b,c,h,w = A.shape + elif (A.ndim == 5): + b,c,t,h,w = A.shape + + A_flat = A.view(b, c, -1) + + for ref in refs: + ref_flat = ref.view(b, c, -1).clone() + + ref_flat /= ref_flat.norm(dim=-1, keepdim=True) + + proj_coeff = torch.sum(A_flat * ref_flat, dim=-1, keepdim=True) + projection = proj_coeff * ref_flat + + A_flat -= projection + + return A_flat.view_as(A) + + + +# Efficient implementation equivalent to the following: +def attention_weights( + query, + key, + attn_mask=None +) -> torch.Tensor: + L, S = query.size(-2), key.size(-2) + scale_factor = 1 / math.sqrt(query.size(-1)) + attn_bias = torch.zeros(L, S, dtype=query.dtype).to(query.device) + + if attn_mask is not None: + if attn_mask.dtype == torch.bool: + attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf")) + else: + attn_bias += attn_mask + + attn_weight = query @ key.transpose(-2, -1) * scale_factor + attn_weight += attn_bias + attn_weight = torch.softmax(attn_weight, dim=-1) + + return attn_weight + + +def attention_weights_orig(q, k): + # implementation of in-place softmax to reduce memory req + scores = torch.matmul(q, k.transpose(-2, -1)) + scores.div_(math.sqrt(q.size(-1))) + torch.exp(scores, out=scores) + summed = torch.sum(scores, dim=-1, keepdim=True) + scores /= summed + return scores.nan_to_num_(0.0, 65504., -65504.) + + +# calculate slerp ratio needed to hit a target cosine similarity score +def get_slerp_weight_for_cossim(cos_sim, target_cos): + # assumes unit vector matrices used for cossim + import math + c = cos_sim + T = target_cos + K = 1 - c + + A = K**2 - 2 * T**2 * K + B = 2 * (1 - c) * (c + T**2) + C = c**2 - T**2 + + if abs(A) < 1e-8: # nearly collinear + return 0.5 # just mix 50:50 + + disc = B**2 - 4*A*C + if disc < 0: + return None # no valid solution... blow up somewhere to get user's attention + + sqrt_disc = math.sqrt(disc) + w1 = (-B + sqrt_disc) / (2 * A) + w2 = (-B - sqrt_disc) / (2 * A) + + candidates = [w for w in [w1, w2] if 0 <= w <= 1] + if candidates: + return candidates[0] + else: + return max(0.0, min(1.0, w1)) + + + +def get_slerp_ratio(cos_sim_A, cos_sim_B, target_cos): + import math + alpha = math.acos(cos_sim_A) + beta = math.acos(cos_sim_B) + delta = math.acos(target_cos) + + if abs(beta - alpha) < 1e-6: + return 0.5 + + t = (delta - alpha) / (beta - alpha) + t = max(0.0, min(1.0, t)) + return t + +def find_slerp_ratio_grid(A: torch.Tensor, B: torch.Tensor, D: torch.Tensor, E: torch.Tensor, + target_ratio: float = 1.0, num_samples: int = 100) -> float: + """ + Finds the interpolation parameter t (in [0,1]) for which: + f(t) = cos(slerp(t, A, B), D) - target_ratio * cos(slerp(t, A, B), E) + is minimized in absolute value. + + Instead of requiring a sign change for bisection, we sample t values uniformly and pick the one that minimizes |f(t)|. + """ + ts = torch.linspace(0.0, 1.0, steps=num_samples, device=A.device, dtype=A.dtype) + best_t = 0.0 + best_val = float('inf') + for t_val in ts: + t_tensor = torch.tensor(t_val, dtype=A.dtype, device=A.device) + C = slerp_tensor(t_tensor, A, B) + diff = get_pearson_similarity(C, D) - target_ratio * get_pearson_similarity(C, E) + if abs(diff) < best_val: + best_val = abs(diff) + best_t = t_val + return best_t + + + +def compute_slerp_ratio_for_target(A: torch.Tensor, B: torch.Tensor, D: torch.Tensor, target: float) -> float: + """ + Given three unit vectors A, B, and D (all assumed to be coplanar) + and a target cosine similarity (target) for the slerp result C with D, + compute the interpolation parameter t such that: + C = slerp(t, A, B) + and cos(C, D) ≈ target. + + Args: + A: Tensor of shape (D,), starting vector. + B: Tensor of shape (D,), ending vector. + D: Tensor of shape (D,), the reference vector. + target: Desired cosine similarity between C and D. + + Returns: + t: A float between 0 and 1. + """ + A = A / (A.norm() + 1e-8) + B = B / (B.norm() + 1e-8) + D = D / (D.norm() + 1e-8) + + alpha = math.acos(max(-1.0, min(1.0, float(torch.dot(D, A))))) # angel between D and A + beta = math.acos(max(-1.0, min(1.0, float(torch.dot(D, B))))) # angle between D and B + + delta = math.acos(max(-1.0, min(1.0, target))) # target cosine similarity... angle etc... + + if abs(beta - alpha) < 1e-6: + return 0.5 + + t = (delta - alpha) / (beta - alpha) + t = max(0.0, min(1.0, t)) + return t + + + +# TENSOR NORMALIZATION OPS + +def normalize_zscore(x, channelwise=False, inplace=False): + if inplace: + if channelwise: + return x.sub_(x.mean(dim=(-2,-1), keepdim=True)).div_(x.std(dim=(-2,-1), keepdim=True)) + else: + return x.sub_(x.mean()).div_(x.std()) + else: + if channelwise: + return (x - x.mean(dim=(-2,-1), keepdim=True) / x.std(dim=(-2,-1), keepdim=True)) + else: + return (x - x.mean()) / x.std() + +def latent_normalize_channels(x): + mean = x.mean(dim=(-2, -1), keepdim=True) + std = x.std (dim=(-2, -1), keepdim=True) + return (x - mean) / std + +def latent_stdize_channels(x): + std = x.std (dim=(-2, -1), keepdim=True) + return x / std + +def latent_meancenter_channels(x): + mean = x.mean(dim=(-2, -1), keepdim=True) + return x - mean + + + +# TENSOR INTERPOLATION OPS + +def lagrange_interpolation(x_values, y_values, x_new): + + if not isinstance(x_values, torch.Tensor): + x_values = torch.tensor(x_values, dtype=torch.get_default_dtype()) + if x_values.ndim != 1: + raise ValueError("x_values must be a 1D tensor or a list of scalars.") + + if not isinstance(x_new, torch.Tensor): + x_new = torch.tensor(x_new, dtype=x_values.dtype, device=x_values.device) + if x_new.ndim == 0: + x_new = x_new.unsqueeze(0) + + if isinstance(y_values, list): + y_values = torch.stack(y_values, dim=0) + if y_values.ndim < 1: + raise ValueError("y_values must have at least one dimension (the sample dimension).") + + n = x_values.shape[0] + if y_values.shape[0] != n: + raise ValueError(f"Mismatch: x_values has length {n} but y_values has {y_values.shape[0]} samples.") + + m = x_new.shape[0] + result_shape = (m,) + y_values.shape[1:] + result = torch.zeros(result_shape, dtype=y_values.dtype, device=y_values.device) + + for i in range(n): + Li = torch.ones_like(x_new, dtype=y_values.dtype, device=y_values.device) + xi = x_values[i] + for j in range(n): + if i == j: + continue + xj = x_values[j] + Li = Li * ((x_new - xj) / (xi - xj)) + extra_dims = (1,) * (y_values.ndim - 1) + Li = Li.view(m, *extra_dims) + result = result + Li * y_values[i] + + return result + +def line_intersection(a: torch.Tensor, d1: torch.Tensor, b: torch.Tensor, d2: torch.Tensor, eps=1e-8) -> torch.Tensor: + """ + Computes the intersection (or closest point average) of two lines in R^D. + + The first line is defined by: L1: x = a + t * d1 + The second line is defined by: L2: x = b + s * d2 + + If the lines do not exactly intersect, this function returns the average of the closest points. + + a, d1, b, d2: Tensors of shape (D,) or with an extra batch dimension (B, D). + Returns: Tensor of shape (D,) or (B, D) representing the intersection (or midpoint of closest approach). + """ + # Compute dot products + d1d1 = (d1 * d1).sum(dim=-1, keepdim=True) # shape (B,1) or (1,) + d2d2 = (d2 * d2).sum(dim=-1, keepdim=True) + d1d2 = (d1 * d2).sum(dim=-1, keepdim=True) + + r = b - a # shape (B, D) or (D,) + r_d1 = (r * d1).sum(dim=-1, keepdim=True) + r_d2 = (r * d2).sum(dim=-1, keepdim=True) + + # Solve for t and s: + # t * d1d1 - s * d1d2 = r_d1 + # t * d1d2 - s * d2d2 = r_d2 + # Solve using determinants: + denom = d1d1 * d2d2 - d1d2 * d1d2 + # Avoid division by zero + denom = torch.where(denom.abs() < eps, torch.full_like(denom, eps), denom) + t = (r_d1 * d2d2 - r_d2 * d1d2) / denom + s = (r_d1 * d1d2 - r_d2 * d1d1) / denom + + point1 = a + t * d1 + point2 = b + s * d2 + # If they intersect exactly, point1 and point2 are identical. + # Otherwise, return the midpoint of the closest points. + return (point1 + point2) / 2 + +def slerp_direction(t: float, u0: torch.Tensor, u1: torch.Tensor, DOT_THRESHOLD=0.9995) -> torch.Tensor: + dot = (u0 * u1).sum(-1).clamp(-1.0, 1.0) #u0, u1 are unit vectors... should not be affected by clamp + if dot.item() > DOT_THRESHOLD: # u0, u1 nearly aligned, fallback to lerp + return torch.lerp(u0, u1, t) + theta_0 = torch.acos(dot) + sin_theta_0 = torch.sin(theta_0) + theta_t = theta_0 * t + sin_theta_t = torch.sin(theta_t) + s0 = torch.sin(theta_0 - theta_t) / sin_theta_0 + s1 = sin_theta_t / sin_theta_0 + return s0 * u0 + s1 * u1 + +def magnitude_aware_interpolation(t: float, v0: torch.Tensor, v1: torch.Tensor) -> torch.Tensor: + + m0 = v0.norm(dim=-1, keepdim=True) + m1 = v1.norm(dim=-1, keepdim=True) + + u0 = v0 / (m0 + 1e-8) + u1 = v1 / (m1 + 1e-8) + + u = slerp_direction(t, u0, u1) + + m = (1 - t) * m0 + t * m1 # tinerpolate magnitudes linearly + return m * u + + +def slerp_tensor(val: torch.Tensor, low: torch.Tensor, high: torch.Tensor, dim=-3) -> torch.Tensor: + if low.ndim == 3: # [1,1,N] flat tensor + dim = -1 + elif low.ndim == 4 and low.shape[-3] > 1: + dim=-3 + elif low.ndim == 5 and low.shape[-3] > 1: + dim=-4 + elif low.ndim == 2: + dim=(-2,-1) + + if type(val) == float: + val = torch.Tensor([val]).expand_as(low).to(low.dtype).to(low.device) + + if val.shape != low.shape: + val = val.expand_as(low) + + low_norm = low / (torch.norm(low, dim=dim, keepdim=True)) + high_norm = high / (torch.norm(high, dim=dim, keepdim=True)) + + dot = (low_norm * high_norm).sum(dim=dim, keepdim=True).clamp(-1.0, 1.0) + + #near = ~(-0.9995 < dot < 0.9995) #dot > 0.9995 or dot < -0.9995 + near = dot > 0.9995 + opposite = dot < -0.9995 + + condition = torch.logical_or(near, opposite) + + omega = torch.acos(dot) + so = torch.sin(omega) + + if val.ndim < low.ndim: + val = val.unsqueeze(dim) + + factor_low = torch.sin((1 - val) * omega) / so + factor_high = torch.sin(val * omega) / so + + res = factor_low * low + factor_high * high + res = torch.where(condition, low * (1 - val) + high * val, res) + return res + + + + +# pytorch slerp implementation from https://gist.github.com/Birch-san/230ac46f99ec411ed5907b0a3d728efa +from torch import FloatTensor, LongTensor, Tensor, Size, lerp, zeros_like +from torch.linalg import norm + +# adapted to PyTorch from: +# https://gist.github.com/dvschultz/3af50c40df002da3b751efab1daddf2c +# most of the extra complexity is to support: +# - many-dimensional vectors +# - v0 or v1 with last dim all zeroes, or v0 ~colinear with v1 +# - falls back to lerp() +# - conditional logic implemented with parallelism rather than Python loops +# - many-dimensional tensor for t +# - you can ask for batches of slerp outputs by making t more-dimensional than the vectors +# - slerp( +# v0: torch.Size([2,3]), +# v1: torch.Size([2,3]), +# t: torch.Size([4,1,1]), +# ) +# - this makes it interface-compatible with lerp() + +def slerp(v0: FloatTensor, v1: FloatTensor, t: float|FloatTensor, DOT_THRESHOLD=0.9995): + ''' + Spherical linear interpolation + Args: + v0: Starting vector + v1: Final vector + t: Float value between 0.0 and 1.0 + DOT_THRESHOLD: Threshold for considering the two vectors as + colinear. Not recommended to alter this. + Returns: + Interpolation vector between v0 and v1 + ''' + assert v0.shape == v1.shape, "shapes of v0 and v1 must match" + + # Normalize the vectors to get the directions and angles + v0_norm: FloatTensor = norm(v0, dim=-1) + v1_norm: FloatTensor = norm(v1, dim=-1) + + v0_normed: FloatTensor = v0 / v0_norm.unsqueeze(-1) + v1_normed: FloatTensor = v1 / v1_norm.unsqueeze(-1) + + # Dot product with the normalized vectors + dot: FloatTensor = (v0_normed * v1_normed).sum(-1) + dot_mag: FloatTensor = dot.abs() + + # if dp is NaN, it's because the v0 or v1 row was filled with 0s + # If absolute value of dot product is almost 1, vectors are ~colinear, so use lerp + gotta_lerp: LongTensor = dot_mag.isnan() | (dot_mag > DOT_THRESHOLD) + can_slerp: LongTensor = ~gotta_lerp + + t_batch_dim_count: int = max(0, t.ndim-v0.ndim) if isinstance(t, Tensor) else 0 + t_batch_dims: Size = t.shape[:t_batch_dim_count] if isinstance(t, Tensor) else Size([]) + out: FloatTensor = zeros_like(v0.expand(*t_batch_dims, *[-1]*v0.ndim)) + + # if no elements are lerpable, our vectors become 0-dimensional, preventing broadcasting + if gotta_lerp.any(): + lerped: FloatTensor = lerp(v0, v1, t) + + out: FloatTensor = lerped.where(gotta_lerp.unsqueeze(-1), out) + + # if no elements are slerpable, our vectors become 0-dimensional, preventing broadcasting + if can_slerp.any(): + + # Calculate initial angle between v0 and v1 + theta_0: FloatTensor = dot.arccos().unsqueeze(-1) + sin_theta_0: FloatTensor = theta_0.sin() + # Angle at timestep t + theta_t: FloatTensor = theta_0 * t + sin_theta_t: FloatTensor = theta_t.sin() + # Finish the slerp algorithm + s0: FloatTensor = (theta_0 - theta_t).sin() / sin_theta_0 + s1: FloatTensor = sin_theta_t / sin_theta_0 + slerped: FloatTensor = s0 * v0 + s1 * v1 + + out: FloatTensor = slerped.where(can_slerp.unsqueeze(-1), out) + + return out + + + +# this is silly... +def normalize_latent(target, source=None, mean=True, std=True, set_mean=None, set_std=None, channelwise=True): + target = target.clone() + source = source.clone() if source is not None else None + def normalize_single_latent(single_target, single_source=None): + y = torch.zeros_like(single_target) + for b in range(y.shape[0]): + if channelwise: + for c in range(y.shape[1]): + single_source_mean = single_source[b][c].mean() if set_mean is None else set_mean + single_source_std = single_source[b][c].std() if set_std is None else set_std + + if mean and std: + y[b][c] = (single_target[b][c] - single_target[b][c].mean()) / single_target[b][c].std() + if single_source is not None: + y[b][c] = y[b][c] * single_source_std + single_source_mean + elif mean: + y[b][c] = single_target[b][c] - single_target[b][c].mean() + if single_source is not None: + y[b][c] = y[b][c] + single_source_mean + elif std: + y[b][c] = single_target[b][c] / single_target[b][c].std() + if single_source is not None: + y[b][c] = y[b][c] * single_source_std + else: + single_source_mean = single_source[b].mean() if set_mean is None else set_mean + single_source_std = single_source[b].std() if set_std is None else set_std + + if mean and std: + y[b] = (single_target[b] - single_target[b].mean()) / single_target[b].std() + if single_source is not None: + y[b] = y[b] * single_source_std + single_source_mean + elif mean: + y[b] = single_target[b] - single_target[b].mean() + if single_source is not None: + y[b] = y[b] + single_source_mean + elif std: + y[b] = single_target[b] / single_target[b].std() + if single_source is not None: + y[b] = y[b] * single_source_std + return y + + if isinstance(target, (list, tuple)): + if source is not None: + assert isinstance(source, (list, tuple)) and len(source) == len(target), \ + "If target is a list/tuple, source must be a list/tuple of the same length." + return [normalize_single_latent(t, s) for t, s in zip(target, source)] + else: + return [normalize_single_latent(t) for t in target] + else: + return normalize_single_latent(target, source) + + + +def hard_light_blend(base_latent, blend_latent): + if base_latent.sum() == 0 and base_latent.std() == 0: + return base_latent + + blend_latent = (blend_latent - blend_latent.min()) / (blend_latent.max() - blend_latent.min()) + + positive_mask = base_latent >= 0 + negative_mask = base_latent < 0 + + positive_latent = base_latent * positive_mask.float() + negative_latent = base_latent * negative_mask.float() + + positive_result = torch.where(blend_latent < 0.5, + 2 * positive_latent * blend_latent, + 1 - 2 * (1 - positive_latent) * (1 - blend_latent)) + + negative_result = torch.where(blend_latent < 0.5, + 2 * negative_latent.abs() * blend_latent, + 1 - 2 * (1 - negative_latent.abs()) * (1 - blend_latent)) + + negative_result = -negative_result + + combined_result = positive_result * positive_mask.float() + negative_result * negative_mask.float() + + #combined_result *= base_latent.max() + + ks = combined_result + ks2 = torch.zeros_like(base_latent) + for n in range(base_latent.shape[1]): + ks2[0][n] = (ks[0][n]) / ks[0][n].std() + ks2[0][n] = (ks2[0][n] * base_latent[0][n].std()) + combined_result = ks2 + + return combined_result + + + + +def make_checkerboard(tile_size: int, num_tiles: int, dtype=torch.float16, device="cpu"): + pattern = torch.tensor([[0, 1], [1, 0]], dtype=dtype, device=device) + board = pattern.repeat(num_tiles // 2 + 1, num_tiles // 2 + 1)[:num_tiles, :num_tiles] + board_expanded = board.repeat_interleave(tile_size, dim=0).repeat_interleave(tile_size, dim=1) + return board_expanded + + + +def get_edge_mask_slug(mask: torch.Tensor, dilation: int = 3) -> torch.Tensor: + + mask = mask.float() + + eroded = -F.max_pool2d(-mask.unsqueeze(0).unsqueeze(0), kernel_size=3, stride=1, padding=1) + eroded = eroded.squeeze(0).squeeze(0) + + edge = mask - eroded + edge = (edge > 0).float() + + dilated_edge = F.max_pool2d(edge.unsqueeze(0).unsqueeze(0), kernel_size=dilation, stride=1, padding=dilation//2) + dilated_edge = dilated_edge.squeeze(0).squeeze(0) + + return dilated_edge + + + +def get_edge_mask(mask: torch.Tensor, dilation: int = 3) -> torch.Tensor: + if dilation == 0: # safeguard for zero kernel size... + return mask + mask_tmp = mask.squeeze().to('cuda') + mask_tmp = mask_tmp.float() + + eroded = -F.max_pool2d(-mask_tmp.unsqueeze(0).unsqueeze(0), kernel_size=3, stride=1, padding=1) + eroded = eroded.squeeze(0).squeeze(0) + + edge = mask_tmp - eroded + edge = (edge > 0).float() + + dilated_edge = F.max_pool2d(edge.unsqueeze(0).unsqueeze(0), kernel_size=dilation, stride=1, padding=dilation//2) + dilated_edge = dilated_edge.squeeze(0).squeeze(0) + + return dilated_edge[...,:mask.shape[-2], :mask.shape[-1]].view_as(mask).to(mask.device) + + + +def checkerboard_variable(widths, dtype=torch.float16, device='cpu'): + total = sum(widths) + mask = torch.zeros((total, total), dtype=dtype, device=device) + + x_start = 0 + for i, w_x in enumerate(widths): + y_start = 0 + for j, w_y in enumerate(widths): + if (i + j) % 2 == 0: # checkerboard logic + mask[x_start:x_start+w_x, y_start:y_start+w_y] = 1.0 + y_start += w_y + x_start += w_x + + return mask + + + + + +def interpolate_spd(cov1, cov2, t, eps=1e-5): + """ + Geodesic interpolation on the SPD manifold between cov1 and cov2. + + Args: + cov1, cov2: [D×D] symmetric positive-definite covariances (torch.Tensor). + t: interpolation factor in [0,1]. + eps: jitter added to diagonal for numerical stability. + + Returns: + cov_t: the SPD matrix at fraction t along the geodesic from cov1 to cov2. + """ + cov1 = cov1.double() + cov2 = cov2.double() + + M1 = cov1.clone() + M1.diagonal().add_(eps) + M2 = cov2.clone() + M2.diagonal().add_(eps) + + S1, U1 = torch.linalg.eigh(M1) + S1_clamped = S1.clamp(min=eps) + inv_sqrt_S1 = S1_clamped.rsqrt() + M1_inv_sqrt = U1 @ torch.diag(inv_sqrt_S1) @ U1.T + + middle = M1_inv_sqrt @ M2 @ M1_inv_sqrt + + Sm, Um = torch.linalg.eigh(middle) + Sm_clamped = Sm.clamp(min=eps) + + Sm_t = Sm_clamped.pow(t) + + middle_t = Um @ torch.diag(Sm_t) @ Um.T + + sqrt_S1 = S1_clamped.sqrt() + M1_sqrt = U1 @ torch.diag(sqrt_S1) @ U1.T + + cov_t = M1_sqrt @ middle_t @ M1_sqrt + + return cov_t.to(cov1.dtype) + + + + + +def tile_latent(latent: torch.Tensor, + tile_size: Tuple[int,int] + ) -> Tuple[torch.Tensor, + Tuple[int,...], + Tuple[int,int], + Tuple[List[int],List[int]]]: + """ + Split `latent` into spatial tiles of shape (t_h, t_w). + Works on either: + - 4D [B,C,H,W] + - 5D [B,C,T,H,W] + Returns: + tiles: [B*rows*cols, C, (T,), t_h, t_w] + orig_shape: the full shape of `latent` + tile_hw: (t_h, t_w) + positions: (pos_h, pos_w) lists of start y and x positions + """ + *lead, H, W = latent.shape + B, C = lead[0], lead[1] + has_time = (latent.ndim == 5) + if has_time: + T = lead[2] + t_h, t_w = tile_size + + rows = (H + t_h - 1) // t_h + cols = (W + t_w - 1) // t_w + + if rows == 1: + pos_h = [0] + else: + pos_h = [round(i*(H - t_h)/(rows-1)) for i in range(rows)] + if cols == 1: + pos_w = [0] + else: + pos_w = [round(j*(W - t_w)/(cols-1)) for j in range(cols)] + + tiles = [] + for y in pos_h: + for x in pos_w: + if has_time: + tile = latent[:, :, :, y:y+t_h, x:x+t_w] + else: + tile = latent[:, :, y:y+t_h, x:x+t_w] + tiles.append(tile) + + tiles = torch.cat(tiles, dim=0) + orig_shape = tuple(latent.shape) + return tiles, orig_shape, (t_h, t_w), (pos_h, pos_w) + + +def untile_latent(tiles: torch.Tensor, + orig_shape: Tuple[int,...], + tile_hw: Tuple[int,int], + positions: Tuple[List[int],List[int]] + ) -> torch.Tensor: + """ + Reconstruct latent from tiles + their start positions. + Works on either 4D or 5D original. + Args: + tiles: [B*rows*cols, C, (T,), t_h, t_w] + orig_shape: shape of original latent (B,C,H,W) or (B,C,T,H,W) + tile_hw: (t_h, t_w) + positions: (pos_h, pos_w) + Returns: + reconstructed latent of shape `orig_shape` + """ + *lead, H, W = orig_shape + B, C = lead[0], lead[1] + has_time = (len(orig_shape) == 5) + if has_time: + T = lead[2] + t_h, t_w = tile_hw + pos_h, pos_w = positions + rows, cols = len(pos_h), len(pos_w) + + if has_time: + out = torch.zeros(B, C, T, H, W, device=tiles.device, dtype=tiles.dtype) + count = torch.zeros_like(out) + tiles = tiles.view(B, rows, cols, C, T, t_h, t_w) + for bi in range(B): + for i, y in enumerate(pos_h): + for j, x in enumerate(pos_w): + tile = tiles[bi, i, j] + out[bi, :, :, y:y+t_h, x:x+t_w] += tile + count[bi, :, :, y:y+t_h, x:x+t_w] += 1 + else: + out = torch.zeros(B, C, H, W, device=tiles.device, dtype=tiles.dtype) + count = torch.zeros_like(out) + tiles = tiles.view(B, rows, cols, C, t_h, t_w) + for bi in range(B): + for i, y in enumerate(pos_h): + for j, x in enumerate(pos_w): + tile = tiles[bi, i, j] + out[bi, :, y:y+t_h, x:x+t_w] += tile + count[bi, :, y:y+t_h, x:x+t_w] += 1 + + valid = count > 0 + out[valid] = out[valid] / count[valid] + return out + + + +def upscale_to_match_spatial(tensor_5d, ref_4d, mode='bicubic'): + """ + Upscales a 5D tensor [B, C, T, H1, W1] to match the spatial size of a 4D tensor [1, C, H2, W2]. + + Args: + tensor_5d: Tensor of shape [B, C, T, H1, W1] + ref_4d: Tensor of shape [1, C, H2, W2] — used as spatial reference + mode: Interpolation mode ('bilinear' or 'bicubic') + + Returns: + Resized tensor of shape [B, C, T, H2, W2] + """ + b, c, t, _, _ = tensor_5d.shape + _, _, h_target, w_target = ref_4d.shape + + tensor_reshaped = tensor_5d.reshape(b * c, t, tensor_5d.shape[-2], tensor_5d.shape[-1]) + upscaled = F.interpolate(tensor_reshaped, size=(h_target, w_target), mode=mode, align_corners=False) + return upscaled.view(b, c, t, h_target, w_target) + + + + + +def gaussian_blur_2d(img: torch.Tensor, sigma: float, kernel_size: int = None) -> torch.Tensor: + B, C, H, W = img.shape + dtype = img.dtype + device = img.device + + if kernel_size is None: + kernel_size = int(2 * math.ceil(3 * sigma) + 1) + + if kernel_size % 2 == 0: + kernel_size += 1 + + coords = torch.arange(kernel_size, dtype=torch.float64) - kernel_size // 2 + g = torch.exp(-0.5 * (coords / sigma) ** 2) + g = g / g.sum() + + kernel_2d = g[:, None] * g[None, :] + kernel_2d = kernel_2d.to(dtype=dtype, device=device) + + kernel = kernel_2d.expand(C, 1, kernel_size, kernel_size) + + pad = kernel_size // 2 + img_padded = F.pad(img, (pad, pad, pad, pad), mode='reflect') + + return F.conv2d(img_padded, kernel, groups=C) + + +def median_blur_2d(img: torch.Tensor, kernel_size: int = 3) -> torch.Tensor: + if kernel_size % 2 == 0: + kernel_size += 1 + pad = kernel_size // 2 + + B, C, H, W = img.shape + img_padded = F.pad(img, (pad, pad, pad, pad), mode='reflect') + + unfolded = img_padded.unfold(2, kernel_size, 1).unfold(3, kernel_size, 1) + # unfolded: [B, C, H, W, kH, kW] → flatten to patches + patches = unfolded.contiguous().view(B, C, H, W, -1) + median = patches.median(dim=-1).values + return median + +def walk_state_info(obj, tensor_handler): + """ + Recurse through dict/list/tuple containers; call tensor_handler on every non-container + leaf (torch.Tensor, NestedTensor, or anything else). The handler is responsible for + deciding whether to transform or return the leaf unchanged. + """ + if isinstance(obj, dict): + changed = False + out = {} + for k, v in obj.items(): + nv = walk_state_info(v, tensor_handler) + changed |= (nv is not v) + out[k] = nv + return out if changed else obj + if isinstance(obj, list): + changed = False + out = [] + for v in obj: + nv = walk_state_info(v, tensor_handler) + changed |= (nv is not v) + out.append(nv) + return out if changed else obj + if isinstance(obj, tuple): + new_t = tuple(walk_state_info(v, tensor_handler) for v in obj) + if all(ov is nv for ov, nv in zip(obj, new_t)): + return obj + return new_t + return tensor_handler(obj) + + +def _is_packed_match(tensor, latent_shapes): + if not is_packed_latent(latent_shapes): + return False + if tensor.ndim < 3 or tensor.shape[-2] != 1: + return False + expected_flat = sum(math.prod(s[1:]) for s in latent_shapes) + return tensor.shape[-1] == expected_flat + + +def apply_to_state_info_tensors(obj, ref_shape, modify_func, *args, latent_shapes=None, **kwargs): + """ + Apply modify_func to every video-sized tensor in obj. Handles: + - NestedTensor: modify_func applied to modality 0 (video); others pass through. + - Packed 3D+ tensor (requires latent_shapes with len > 1): unpack, modify video, repack. + - Regular tensor with last 5 dims matching ref_shape[-5:]: modify_func applied directly. + + latent_shapes defaults to [ref_shape] (single-modality, no packing). + """ + from comfy.nested_tensor import NestedTensor + ref_last5 = ref_shape[-5:] if len(ref_shape) >= 5 else ref_shape + ls = latent_shapes if latent_shapes is not None else [ref_shape] + + def handler(t): + if isinstance(t, NestedTensor): + c = t.unbind() + if c[0].ndim >= 5 and c[0].shape[-5:] == ref_last5: + return NestedTensor([modify_func(c[0], *args, **kwargs)] + list(c[1:])) + return t + if isinstance(t, torch.Tensor): + if _is_packed_match(t, ls): + return LatentHandler(t, ls).map_first(lambda v: modify_func(v, *args, **kwargs)).tensor + if t.ndim >= 5 and t.shape[-5:] == ref_last5: + return modify_func(t, *args, **kwargs) + return t + return walk_state_info(obj, handler) + + +def derive_old_latent_shapes(raw_x, latent_shapes_new): + """ + Reconstruct per-modality shapes for a stored raw_x given the current latent_shapes. + Assumes only the video modality (index 0) changed its temporal length; all other + modalities are unchanged. Returns None if shapes can't be cleanly reconstructed + under that assumption (e.g. spatial dims also changed, or audio modality changed). + """ + if raw_x.ndim >= 5: + # Regular unpacked tensor — single modality + return [tuple(raw_x.shape)] + # Packed flat 3D [B, 1, flat_total] + audio_flat = sum(math.prod(s[1:]) for s in latent_shapes_new[1:]) + video_CHW = math.prod(latent_shapes_new[0][1:]) // latent_shapes_new[0][-3] + video_flat_old = raw_x.shape[-1] - audio_flat + if video_flat_old % video_CHW != 0: + return None + T_old = video_flat_old // video_CHW + video_shape_old = list(latent_shapes_new[0]) + video_shape_old[-3] = T_old + return [tuple(video_shape_old)] + list(latent_shapes_new[1:]) + + +def extract_video_tail(x, latent_shapes_new, extra_T): + """ + Extract the last `extra_T` video-temporal frames from x as a 5D tensor + [B, C, extra_T, H, W]. Handles NestedTensor, packed flat 3D, and regular 5D. + """ + from comfy.nested_tensor import NestedTensor + video = x.unbind()[0] if isinstance(x, NestedTensor) else get_latent(x, latent_shapes_new, 0) + return video[..., -extra_T:, :, :] + + +def extend_state_info_tensors(obj, latent_shapes_old, x_video_tail): + """ + For any tensor whose video dims match latent_shapes_old[0], extend along dim -3 + by concatenating a broadcast-adjusted copy of x_video_tail. Handles NestedTensor, + packed (3D+), and regular 5D/6D. + """ + from comfy.nested_tensor import NestedTensor + old_video_shape = latent_shapes_old[0] + + def _cat_tail(tensor): + tail = x_video_tail.to(device=tensor.device, dtype=tensor.dtype) + while tail.ndim < tensor.ndim: + tail = tail.unsqueeze(0) + expand_shape = list(tensor.shape[:-3]) + list(tail.shape[-3:]) + tail = tail.expand(expand_shape) + return torch.cat([tensor, tail], dim=-3) + + def handler(t): + if isinstance(t, NestedTensor): + c = t.unbind() + if c[0].shape[-3] == old_video_shape[-3] and c[0].shape[-5:] == old_video_shape[-5:]: + return NestedTensor([_cat_tail(c[0])] + list(c[1:])) + return t + if isinstance(t, torch.Tensor): + if _is_packed_match(t, latent_shapes_old): + return LatentHandler(t, latent_shapes_old).map_first(_cat_tail).tensor + if t.ndim >= 5 and t.shape[-5:] == old_video_shape[-5:]: + return _cat_tail(t) + return t + return walk_state_info(obj, handler) + + +# LATENT PACKING HELPERS +def is_packed_latent(latent_shapes): + return latent_shapes is not None and len(latent_shapes) > 1 + + +def get_latent(x, latent_shapes=None, index=0): + if latent_shapes is None: + return x + if not is_packed_latent(latent_shapes): + return x + return comfy.utils.unpack_latents(x, latent_shapes)[index] + + +def apply_per_step_latent_normalization(x, step, latent_shapes, factors_0_list, factors_1_list): + if latent_shapes is None or len(latent_shapes) <= 1: + return x + + # Get normalization factor for step + factor_0 = factors_0_list[min(step, len(factors_0_list) - 1)] + factor_1 = factors_1_list[min(step, len(factors_1_list) - 1)] + + if factor_0 == 1.0 and factor_1 == 1.0: + return x + + # Unpack + tensors = comfy.utils.unpack_latents(x, latent_shapes) + + factors = [factor_0, factor_1] + for idx, t in enumerate(tensors): + factor = factors[idx] if idx < len(factors) else 1.0 + if factor != 1.0: + tensors[idx] = t * factor + + # Repack + packed, _ = comfy.utils.pack_latents(tensors) + RESplain(f"Per-step latent normalize: step={step}, idx_0={factor_0}, idx_1={factor_1}", debug=True) + return packed + + +class LatentHandler: + """Interface for operations on packed/regular latents.""" + def __init__(self, x, latent_shapes=None): + self.x = x + self.latent_shapes = latent_shapes + + @property + def is_packed(self): + return is_packed_latent(self.latent_shapes) + + @property + def tensor(self): + return self.x + + def map(self, func): + """Apply func to each component (or x directly if not packed), repack.""" + if not self.is_packed: + self.x = func(self.x) + else: + tensors = comfy.utils.unpack_latents(self.x, self.latent_shapes) + self.x, _ = comfy.utils.pack_latents([func(t) for t in tensors]) + return self + + def map_first(self, func): + """Apply func only to first component (video), keep others unchanged, repack. + Handles 3D [B, 1, flat] and higher-dim [..., B, 1, flat] (e.g. stacked data_prev_).""" + if not self.is_packed: + self.x = func(self.x) + return self + leading_shape = self.x.shape[:-3] + flat_leading = int(math.prod(leading_shape)) + tensor_flat = self.x.reshape(flat_leading, *self.x.shape[-3:]) + results = [] + for i in range(flat_leading): + unpacked = comfy.utils.unpack_latents(tensor_flat[i], self.latent_shapes) + unpacked[0] = func(unpacked[0]) + repacked, _ = comfy.utils.pack_latents(unpacked) + results.append(repacked) + out = torch.stack(results, dim=0) + self.x = out.reshape(*leading_shape, *out.shape[-3:]) + return self + + def map_with(self, other, func): + """Apply func(self_component, other_component) pairwise, repack.""" + if not self.is_packed: + self.x = func(self.x, other) + else: + x_tensors = comfy.utils.unpack_latents(self.x, self.latent_shapes) + y_tensors = comfy.utils.unpack_latents(other, self.latent_shapes) + results = [func(x_t, y_t) for x_t, y_t in zip(x_tensors, y_tensors)] + self.x, _ = comfy.utils.pack_latents(results) + return self + + def get_first_tensor(self): + """Get first component tensor (for shape reference, mask prep, etc.).""" + return get_latent(self.x, self.latent_shapes, 0) + diff --git a/simple_syrup/third_party/res4lyf_runtime/res4lyf.py b/simple_syrup/third_party/res4lyf_runtime/res4lyf.py new file mode 100644 index 0000000..72a208d --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/res4lyf.py @@ -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 diff --git a/simple_syrup/third_party/res4lyf_runtime/sigmas.py b/simple_syrup/third_party/res4lyf_runtime/sigmas.py new file mode 100644 index 0000000..87105b0 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/sigmas.py @@ -0,0 +1,4099 @@ +import torch +import numpy as np +from math import * +import builtins +from scipy.interpolate import CubicSpline +from scipy import special, stats +import torch.nn.functional as F +import torch.nn as nn +import torch.optim as optim +import math + + +from comfy.k_diffusion.sampling import get_sigmas_polyexponential, get_sigmas_karras +import comfy.samplers + +from torch import Tensor, nn +from typing import Optional, Callable, Tuple, Dict, Any, Union, TYPE_CHECKING, TypeVar + +from .res4lyf import RESplain +from .helper import get_res4lyf_scheduler_list + + +def rescale_linear(input, input_min, input_max, output_min, output_max): + output = ((input - input_min) / (input_max - input_min)) * (output_max - output_min) + output_min; + return output + +class set_precision_sigmas: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", ), + "precision": (["16", "32", "64"], ), + "set_default": ("BOOLEAN", {"default": False}) + }, + } + + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("passthrough",) + CATEGORY = "RES4LYF/precision" + + FUNCTION = "main" + + def main(self, precision="32", sigmas=None, set_default=False): + match precision: + case "16": + if set_default is True: + torch.set_default_dtype(torch.float16) + sigmas = sigmas.to(torch.float16) + case "32": + if set_default is True: + torch.set_default_dtype(torch.float32) + sigmas = sigmas.to(torch.float32) + case "64": + if set_default is True: + torch.set_default_dtype(torch.float64) + sigmas = sigmas.to(torch.float64) + return (sigmas, ) + + +class SimpleInterpolator(nn.Module): + def __init__(self): + super(SimpleInterpolator, self).__init__() + self.net = nn.Sequential( + nn.Linear(1, 16), + nn.ReLU(), + nn.Linear(16, 32), + nn.ReLU(), + nn.Linear(32, 1) + ) + + def forward(self, x): + return self.net(x) + +def train_interpolator(model, sigma_schedule, steps, epochs=5000, lr=0.01): + with torch.inference_mode(False): + model = SimpleInterpolator() + sigma_schedule = sigma_schedule.clone() + + criterion = nn.MSELoss() + optimizer = optim.Adam(model.parameters(), lr=lr) + + x_train = torch.linspace(0, 1, steps=steps).unsqueeze(1) + y_train = sigma_schedule.unsqueeze(1) + + # disable inference mode for training + model.train() + for epoch in range(epochs): + optimizer.zero_grad() + + # fwd pass + outputs = model(x_train) + loss = criterion(outputs, y_train) + loss.backward() + optimizer.step() + + return model + +def interpolate_sigma_schedule_model(sigma_schedule, target_steps): + model = SimpleInterpolator() + sigma_schedule = sigma_schedule.float().detach() + + # train on original sigma schedule + trained_model = train_interpolator(model, sigma_schedule, len(sigma_schedule)) + + # generate target steps for interpolation + x_interpolated = torch.linspace(0, 1, target_steps).unsqueeze(1) + + # inference w/o gradients + trained_model.eval() + with torch.no_grad(): + interpolated_sigma = trained_model(x_interpolated).squeeze() + + return interpolated_sigma + + + + +class sigmas_interpolate: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_in": ("SIGMAS", {"forceInput": True}), + "output_length": ("INT", {"default": 0, "min": 0,"max": 10000,"step": 1}), + "mode": (["linear", "nearest", "polynomial", "exponential", "power", "model"],), + "order": ("INT", {"default": 8, "min": 1,"max": 64,"step": 1}), + "rescale_after": ("BOOLEAN", {"default": True, "tooltip": "Rescale the output to the original min/max range after interpolation."}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/sigmas" + DESCRIPTION = "Interpolate the sigmas schedule to a new length clamping the start and end values." + + def interpolate_sigma_schedule_poly(self, sigma_schedule, target_steps): + order = self.order + sigma_schedule_np = sigma_schedule.cpu().numpy() + + # orig steps (assuming even spacing) + original_steps = np.linspace(0, 1, len(sigma_schedule_np)) + + # fit polynomial of the given order + coefficients = np.polyfit(original_steps, sigma_schedule_np, deg=order) + + # generate new steps where we want to interpolate the data + target_steps_np = np.linspace(0, 1, target_steps) + + # eval polynomial at new steps + interpolated_sigma_np = np.polyval(coefficients, target_steps_np) + + interpolated_sigma = torch.tensor(interpolated_sigma_np, device=sigma_schedule.device, dtype=sigma_schedule.dtype) + return interpolated_sigma + + def interpolate_sigma_schedule_constrained(self, sigma_schedule, target_steps): + sigma_schedule_np = sigma_schedule.cpu().numpy() + + # orig steps + original_steps = np.linspace(0, 1, len(sigma_schedule_np)) + + # target steps for interpolation + target_steps_np = np.linspace(0, 1, target_steps) + + # fit cubic spline with fixed start and end values + cs = CubicSpline(original_steps, sigma_schedule_np, bc_type=((1, 0.0), (1, 0.0))) + + # eval spline at the target steps + interpolated_sigma_np = cs(target_steps_np) + + interpolated_sigma = torch.tensor(interpolated_sigma_np, device=sigma_schedule.device, dtype=sigma_schedule.dtype) + + return interpolated_sigma + + def interpolate_sigma_schedule_exp(self, sigma_schedule, target_steps): + # transform to log space + log_sigma_schedule = torch.log(sigma_schedule) + + # define the original and target step ranges + original_steps = torch.linspace(0, 1, steps=len(sigma_schedule)) + target_steps = torch.linspace(0, 1, steps=target_steps) + + # interpolate in log space + interpolated_log_sigma = F.interpolate( + log_sigma_schedule.unsqueeze(0).unsqueeze(0), # Add fake batch and channel dimensions + size=target_steps.shape[0], + mode='linear', + align_corners=True + ).squeeze() + + # transform back to exponential space + interpolated_sigma_schedule = torch.exp(interpolated_log_sigma) + + return interpolated_sigma_schedule + + def interpolate_sigma_schedule_power(self, sigma_schedule, target_steps): + sigma_schedule_np = sigma_schedule.cpu().numpy() + original_steps = np.linspace(1, len(sigma_schedule_np), len(sigma_schedule_np)) + + # power regression using a log-log transformation + log_x = np.log(original_steps) + log_y = np.log(sigma_schedule_np) + + # linear regression on log-log data + coefficients = np.polyfit(log_x, log_y, deg=1) # degree 1 for linear fit in log-log space + a = np.exp(coefficients[1]) # a = "b" = intercept (exp because of the log transform) + b = coefficients[0] # b = "m" = slope + + target_steps_np = np.linspace(1, len(sigma_schedule_np), target_steps) + + # power law prediction: y = a * x^b + interpolated_sigma_np = a * (target_steps_np ** b) + + interpolated_sigma = torch.tensor(interpolated_sigma_np, device=sigma_schedule.device, dtype=sigma_schedule.dtype) + + return interpolated_sigma + + def interpolate_sigma_schedule_linear(self, sigma_schedule, target_steps): + return F.interpolate(sigma_schedule.unsqueeze(0).unsqueeze(0), target_steps, mode='linear').squeeze(0).squeeze(0) + + def interpolate_sigma_schedule_nearest(self, sigma_schedule, target_steps): + return F.interpolate(sigma_schedule.unsqueeze(0).unsqueeze(0), target_steps, mode='nearest').squeeze(0).squeeze(0) + + def interpolate_nearest_neighbor(self, sigma_schedule, target_steps): + original_steps = torch.linspace(0, 1, steps=len(sigma_schedule)) + target_steps = torch.linspace(0, 1, steps=target_steps) + + # interpolate original -> target steps using nearest neighbor + indices = torch.searchsorted(original_steps, target_steps) + indices = torch.clamp(indices, 0, len(sigma_schedule) - 1) # clamp indices to valid range + + # set nearest neighbor via indices + interpolated_sigma = sigma_schedule[indices] + + return interpolated_sigma + + + def main(self, sigmas_in, output_length, mode, order, rescale_after=True): + + self.order = order + + sigmas_in = sigmas_in.clone().to(sigmas_in.dtype) + start = sigmas_in[0] + end = sigmas_in[-1] + + if mode == "linear": + interpolate = self.interpolate_sigma_schedule_linear + if mode == "nearest": + interpolate = self.interpolate_nearest_neighbor + elif mode == "polynomial": + interpolate = self.interpolate_sigma_schedule_poly + elif mode == "exponential": + interpolate = self.interpolate_sigma_schedule_exp + elif mode == "power": + interpolate = self.interpolate_sigma_schedule_power + elif mode == "model": + with torch.inference_mode(False): + interpolate = interpolate_sigma_schedule_model + + sigmas_interp = interpolate(sigmas_in, output_length) + if rescale_after: + sigmas_interp = ((sigmas_interp - sigmas_interp.min()) * (start - end)) / (sigmas_interp.max() - sigmas_interp.min()) + end + return (sigmas_interp,) + +class sigmas_noise_inversion: + # flip sigmas for unsampling, and pad both fwd/rev directions with null bytes to disable noise scaling, etc from the model. + # will cause model to return epsilon prediction instead of calculated denoised latent image. + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS","SIGMAS",) + RETURN_NAMES = ("sigmas_fwd","sigmas_rev",) + CATEGORY = "RES4LYF/sigmas" + DESCRIPTION = "For use with unsampling. Connect sigmas_fwd to the unsampling (first) node, and sigmas_rev to the sampling (second) node." + + def main(self, sigmas): + sigmas = sigmas.clone().to(sigmas.dtype) + + null = torch.tensor([0.0], device=sigmas.device, dtype=sigmas.dtype) + sigmas_fwd = torch.flip(sigmas, dims=[0]) + sigmas_fwd = torch.cat([sigmas_fwd, null]) + + sigmas_rev = torch.cat([null, sigmas]) + sigmas_rev = torch.cat([sigmas_rev, null]) + + return (sigmas_fwd, sigmas_rev,) + + +def compute_sigma_next_variance_floor(sigma): + return (-1 + torch.sqrt(1 + 4 * sigma)) / 2 + +class sigmas_variance_floor: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + DESCRIPTION = ("Process a sigma schedule so that any steps that are too large for variance-locked SDE sampling are replaced with the maximum permissible value." + "Will be very difficult to approach sigma = 0 due to the nature of the math, as steps become very small much below approximately sigma = 0.15 to 0.2.") + + def main(self, sigmas): + dtype = sigmas.dtype + sigmas = sigmas.clone().to(sigmas.dtype) + for i in range(len(sigmas) - 1): + sigma_next = (-1 + torch.sqrt(1 + 4 * sigmas[i])) / 2 + + if sigmas[i+1] < sigma_next and sigmas[i+1] > 0.0: + print("swapped i+1 with sigma_next+0.001: ", sigmas[i+1], sigma_next + 0.001) + sigmas[i+1] = sigma_next + 0.001 + return (sigmas.to(dtype),) + + +class sigmas_from_text: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"default": "", "multiline": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, text): + text_list = [float(val) for val in text.replace(",", " ").split()] + #text_list = [float(val.strip()) for val in text.split(",")] + + sigmas = torch.tensor(text_list) #.to('cuda').to(torch.float64) + + return (sigmas,) + + + +class sigmas_concatenate: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_1": ("SIGMAS", {"forceInput": True}), + "sigmas_2": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_1, sigmas_2): + return (torch.cat((sigmas_1, sigmas_2.to(sigmas_1))),) + +class sigmas_truncate: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "sigmas_until": ("INT", {"default": 10, "min": 0,"max": 1000,"step": 1}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, sigmas_until): + sigmas = sigmas.clone() + return (sigmas[:sigmas_until],) + +class sigmas_start: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "sigmas_until": ("INT", {"default": 10, "min": 0,"max": 1000,"step": 1}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, sigmas_until): + sigmas = sigmas.clone() + return (sigmas[sigmas_until:],) + +class sigmas_split: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "sigmas_start": ("INT", {"default": 0, "min": 0,"max": 1000,"step": 1}), + "sigmas_end": ("INT", {"default": 1000, "min": 0,"max": 1000,"step": 1}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, sigmas_start, sigmas_end): + sigmas = sigmas.clone() + return (sigmas[sigmas_start:sigmas_end],) + + sigmas_stop_step = sigmas_end - sigmas_start + return (sigmas[sigmas_start:][:sigmas_stop_step],) + +class sigmas_pad: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "value": ("FLOAT", {"default": 0.0, "min": -10000,"max": 10000,"step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, value): + sigmas = sigmas.clone() + return (torch.cat((sigmas, torch.tensor([value], dtype=sigmas.dtype))),) + +class sigmas_unpad: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas): + sigmas = sigmas.clone() + return (sigmas[:-1],) + +class sigmas_set_floor: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "floor": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + "new_floor": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}) + } + } + + RETURN_TYPES = ("SIGMAS",) + FUNCTION = "set_floor" + + CATEGORY = "RES4LYF/sigmas" + + def set_floor(self, sigmas, floor, new_floor): + sigmas = sigmas.clone() + sigmas[sigmas <= floor] = new_floor + return (sigmas,) + +class sigmas_delete_below_floor: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "floor": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}) + } + } + + RETURN_TYPES = ("SIGMAS",) + FUNCTION = "delete_below_floor" + + CATEGORY = "RES4LYF/sigmas" + + def delete_below_floor(self, sigmas, floor): + sigmas = sigmas.clone() + return (sigmas[sigmas >= floor],) + +class sigmas_delete_value: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "value": ("FLOAT", {"default": 0.0, "min": -1000,"max": 1000,"step": 0.01}) + } + } + + RETURN_TYPES = ("SIGMAS",) + FUNCTION = "delete_value" + + CATEGORY = "RES4LYF/sigmas" + + def delete_value(self, sigmas, value): + return (sigmas[sigmas != value],) + +class sigmas_delete_consecutive_duplicates: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_1": ("SIGMAS", {"forceInput": True}) + } + } + + RETURN_TYPES = ("SIGMAS",) + FUNCTION = "delete_consecutive_duplicates" + + CATEGORY = "RES4LYF/sigmas" + + def delete_consecutive_duplicates(self, sigmas_1): + mask = sigmas_1[:-1] != sigmas_1[1:] + mask = torch.cat((mask, torch.tensor([True]))) + return (sigmas_1[mask],) + +class sigmas_cleanup: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "sigmin": ("FLOAT", {"default": 0.0291675, "min": 0,"max": 1000,"step": 0.01}) + } + } + + RETURN_TYPES = ("SIGMAS",) + FUNCTION = "cleanup" + + CATEGORY = "RES4LYF/sigmas" + + def cleanup(self, sigmas, sigmin): + sigmas_culled = sigmas[sigmas >= sigmin] + + mask = sigmas_culled[:-1] != sigmas_culled[1:] + mask = torch.cat((mask, torch.tensor([True]))) + filtered_sigmas = sigmas_culled[mask] + return (torch.cat((filtered_sigmas,torch.tensor([0]))),) + +class sigmas_mult: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "multiplier": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}) + }, + "optional": { + "sigmas2": ("SIGMAS", {"forceInput": False}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, multiplier, sigmas2=None): + if sigmas2 is not None: + return (sigmas * sigmas2 * multiplier,) + else: + return (sigmas * multiplier,) + +class sigmas_modulus: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "divisor": ("FLOAT", {"default": 1, "min": -1000,"max": 1000,"step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, divisor): + return (sigmas % divisor,) + +class sigmas_quotient: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "divisor": ("FLOAT", {"default": 1, "min": -1000,"max": 1000,"step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, divisor): + return (sigmas // divisor,) + +class sigmas_add: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "addend": ("FLOAT", {"default": 1, "min": -1000,"max": 1000,"step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, addend): + return (sigmas + addend,) + +class sigmas_power: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "power": ("FLOAT", {"default": 1, "min": -100,"max": 100,"step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, power): + return (sigmas ** power,) + +class sigmas_abs: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas): + return (abs(sigmas),) + +class sigmas2_mult: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_1": ("SIGMAS", {"forceInput": True}), + "sigmas_2": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_1, sigmas_2): + return (sigmas_1 * sigmas_2,) + +class sigmas2_add: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_1": ("SIGMAS", {"forceInput": True}), + "sigmas_2": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_1, sigmas_2): + return (sigmas_1 + sigmas_2,) + +class sigmas_rescale: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start": ("FLOAT", {"default": 1.0, "min": -10000,"max": 10000,"step": 0.01}), + "end": ("FLOAT", {"default": 0.0, "min": -10000,"max": 10000,"step": 0.01}), + "sigmas": ("SIGMAS", ), + }, + "optional": { + } + } + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas_rescaled",) + CATEGORY = "RES4LYF/sigmas" + DESCRIPTION = ("Can be used to set denoise. Results are generally better than with the approach used by KSampler and most nodes with denoise values " + "(which slice the sigmas schedule according to step count, not the noise level). Will also flip the sigma schedule if the start and end values are reversed." + ) + + def main(self, start=0, end=-1, sigmas=None): + + s_out_1 = ((sigmas - sigmas.min()) * (start - end)) / (sigmas.max() - sigmas.min()) + end + + return (s_out_1,) + + +class sigmas_count: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", ), + } + } + FUNCTION = "main" + RETURN_TYPES = ("INT",) + RETURN_NAMES = ("count",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas=None): + return (len(sigmas),) + + +class sigmas_math1: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start": ("INT", {"default": 0, "min": 0,"max": 10000,"step": 1}), + "stop": ("INT", {"default": 0, "min": 0,"max": 10000,"step": 1}), + "trim": ("INT", {"default": 0, "min": -10000,"max": 0,"step": 1}), + "x": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "y": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "z": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "f1": ("STRING", {"default": "s", "multiline": True}), + "rescale" : ("BOOLEAN", {"default": False}), + "max1": ("FLOAT", {"default": 14.614642, "min": -10000,"max": 10000,"step": 0.01}), + "min1": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + }, + "optional": { + "a": ("SIGMAS", {"forceInput": False}), + "b": ("SIGMAS", {"forceInput": False}), + "c": ("SIGMAS", {"forceInput": False}), + } + } + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + def main(self, start=0, stop=0, trim=0, a=None, b=None, c=None, x=1.0, y=1.0, z=1.0, f1="s", rescale=False, min1=1.0, max1=1.0): + if stop == 0: + t_lens = [len(tensor) for tensor in [a, b, c] if tensor is not None] + t_len = stop = min(t_lens) if t_lens else 0 + else: + stop = stop + 1 + t_len = stop - start + + stop = stop + trim + t_len = t_len + trim + + t_a = t_b = t_c = None + if a is not None: + t_a = a[start:stop] + if b is not None: + t_b = b[start:stop] + if c is not None: + t_c = c[start:stop] + + t_s = torch.arange(0.0, t_len) + + t_x = torch.full((t_len,), x) + t_y = torch.full((t_len,), y) + t_z = torch.full((t_len,), z) + eval_namespace = {"__builtins__": None, "round": builtins.round, "np": np, "a": t_a, "b": t_b, "c": t_c, "x": t_x, "y": t_y, "z": t_z, "s": t_s, "torch": torch} + eval_namespace.update(np.__dict__) + + s_out_1 = eval(f1, eval_namespace) + + if rescale == True: + s_out_1 = ((s_out_1 - min(s_out_1)) * (max1 - min1)) / (max(s_out_1) - min(s_out_1)) + min1 + + return (s_out_1,) + +class sigmas_math3: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start": ("INT", {"default": 0, "min": 0,"max": 10000,"step": 1}), + "stop": ("INT", {"default": 0, "min": 0,"max": 10000,"step": 1}), + "trim": ("INT", {"default": 0, "min": -10000,"max": 0,"step": 1}), + }, + "optional": { + "a": ("SIGMAS", {"forceInput": False}), + "b": ("SIGMAS", {"forceInput": False}), + "c": ("SIGMAS", {"forceInput": False}), + "x": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "y": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "z": ("FLOAT", {"default": 1, "min": -10000,"max": 10000,"step": 0.01}), + "f1": ("STRING", {"default": "s", "multiline": True}), + "rescale1" : ("BOOLEAN", {"default": False}), + "max1": ("FLOAT", {"default": 14.614642, "min": -10000,"max": 10000,"step": 0.01}), + "min1": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + "f2": ("STRING", {"default": "s", "multiline": True}), + "rescale2" : ("BOOLEAN", {"default": False}), + "max2": ("FLOAT", {"default": 14.614642, "min": -10000,"max": 10000,"step": 0.01}), + "min2": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + "f3": ("STRING", {"default": "s", "multiline": True}), + "rescale3" : ("BOOLEAN", {"default": False}), + "max3": ("FLOAT", {"default": 14.614642, "min": -10000,"max": 10000,"step": 0.01}), + "min3": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + } + } + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS","SIGMAS","SIGMAS") + CATEGORY = "RES4LYF/sigmas" + def main(self, start=0, stop=0, trim=0, a=None, b=None, c=None, x=1.0, y=1.0, z=1.0, f1="s", f2="s", f3="s", rescale1=False, rescale2=False, rescale3=False, min1=1.0, max1=1.0, min2=1.0, max2=1.0, min3=1.0, max3=1.0): + if stop == 0: + t_lens = [len(tensor) for tensor in [a, b, c] if tensor is not None] + t_len = stop = min(t_lens) if t_lens else 0 + else: + stop = stop + 1 + t_len = stop - start + + stop = stop + trim + t_len = t_len + trim + + t_a = t_b = t_c = None + if a is not None: + t_a = a[start:stop] + if b is not None: + t_b = b[start:stop] + if c is not None: + t_c = c[start:stop] + + t_s = torch.arange(0.0, t_len) + + t_x = torch.full((t_len,), x) + t_y = torch.full((t_len,), y) + t_z = torch.full((t_len,), z) + eval_namespace = {"__builtins__": None, "np": np, "a": t_a, "b": t_b, "c": t_c, "x": t_x, "y": t_y, "z": t_z, "s": t_s, "torch": torch} + eval_namespace.update(np.__dict__) + + s_out_1 = eval(f1, eval_namespace) + s_out_2 = eval(f2, eval_namespace) + s_out_3 = eval(f3, eval_namespace) + + if rescale1 == True: + s_out_1 = ((s_out_1 - min(s_out_1)) * (max1 - min1)) / (max(s_out_1) - min(s_out_1)) + min1 + if rescale2 == True: + s_out_2 = ((s_out_2 - min(s_out_2)) * (max2 - min2)) / (max(s_out_2) - min(s_out_2)) + min2 + if rescale3 == True: + s_out_3 = ((s_out_3 - min(s_out_3)) * (max3 - min3)) / (max(s_out_3) - min(s_out_3)) + min3 + + return s_out_1, s_out_2, s_out_3 + +class sigmas_iteration_karras: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps_up": ("INT", {"default": 30, "min": 0,"max": 10000,"step": 1}), + "steps_down": ("INT", {"default": 30, "min": 0,"max": 10000,"step": 1}), + "rho_up": ("FLOAT", {"default": 3, "min": -10000,"max": 10000,"step": 0.01}), + "rho_down": ("FLOAT", {"default": 4, "min": -10000,"max": 10000,"step": 0.01}), + "s_min_start": ("FLOAT", {"default":0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + "s_max": ("FLOAT", {"default": 2, "min": -10000,"max": 10000,"step": 0.01}), + "s_min_end": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + }, + "optional": { + "momentums": ("SIGMAS", {"forceInput": False}), + "sigmas": ("SIGMAS", {"forceInput": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS","SIGMAS") + RETURN_NAMES = ("momentums","sigmas") + CATEGORY = "RES4LYF/schedulers" + + def main(self, steps_up, steps_down, rho_up, rho_down, s_min_start, s_max, s_min_end, sigmas=None, momentums=None): + s_up = get_sigmas_karras(steps_up, s_min_start, s_max, rho_up) + s_down = get_sigmas_karras(steps_down, s_min_end, s_max, rho_down) + s_up = s_up[:-1] + s_down = s_down[:-1] + s_up = torch.flip(s_up, dims=[0]) + sigmas_new = torch.cat((s_up, s_down), dim=0) + momentums_new = torch.cat((s_up, -1*s_down), dim=0) + + if sigmas is not None: + sigmas = torch.cat([sigmas, sigmas_new]) + else: + sigmas = sigmas_new + + if momentums is not None: + momentums = torch.cat([momentums, momentums_new]) + else: + momentums = momentums_new + + return (momentums,sigmas) + +class sigmas_iteration_polyexp: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps_up": ("INT", {"default": 30, "min": 0,"max": 10000,"step": 1}), + "steps_down": ("INT", {"default": 30, "min": 0,"max": 10000,"step": 1}), + "rho_up": ("FLOAT", {"default": 0.6, "min": -10000,"max": 10000,"step": 0.01}), + "rho_down": ("FLOAT", {"default": 0.8, "min": -10000,"max": 10000,"step": 0.01}), + "s_min_start": ("FLOAT", {"default":0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + "s_max": ("FLOAT", {"default": 2, "min": -10000,"max": 10000,"step": 0.01}), + "s_min_end": ("FLOAT", {"default": 0.0291675, "min": -10000,"max": 10000,"step": 0.01}), + }, + "optional": { + "momentums": ("SIGMAS", {"forceInput": False}), + "sigmas": ("SIGMAS", {"forceInput": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS","SIGMAS") + RETURN_NAMES = ("momentums","sigmas") + CATEGORY = "RES4LYF/schedulers" + + def main(self, steps_up, steps_down, rho_up, rho_down, s_min_start, s_max, s_min_end, sigmas=None, momentums=None): + s_up = get_sigmas_polyexponential(steps_up, s_min_start, s_max, rho_up) + s_down = get_sigmas_polyexponential(steps_down, s_min_end, s_max, rho_down) + s_up = s_up[:-1] + s_down = s_down[:-1] + s_up = torch.flip(s_up, dims=[0]) + sigmas_new = torch.cat((s_up, s_down), dim=0) + momentums_new = torch.cat((s_up, -1*s_down), dim=0) + + if sigmas is not None: + sigmas = torch.cat([sigmas, sigmas_new]) + else: + sigmas = sigmas_new + + if momentums is not None: + momentums = torch.cat([momentums, momentums_new]) + else: + momentums = momentums_new + + return (momentums,sigmas) + +class tan_scheduler: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 20, "min": 0,"max": 100000,"step": 1}), + "offset": ("FLOAT", {"default": 20, "min": 0,"max": 100000,"step": 0.1}), + "slope": ("FLOAT", {"default": 20, "min": -100000,"max": 100000,"step": 0.1}), + "start": ("FLOAT", {"default": 20, "min": -100000,"max": 100000,"step": 0.1}), + "end": ("FLOAT", {"default": 20, "min": -100000,"max": 100000,"step": 0.1}), + "sgm" : ("BOOLEAN", {"default": False}), + "pad" : ("BOOLEAN", {"default": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/schedulers" + + def main(self, steps, slope, offset, start, end, sgm, pad): + smax = ((2/pi)*atan(-slope*(0-offset))+1)/2 + smin = ((2/pi)*atan(-slope*((steps-1)-offset))+1)/2 + + srange = smax-smin + sscale = start - end + + if sgm: + steps+=1 + + sigmas = [ ( (((2/pi)*atan(-slope*(x-offset))+1)/2) - smin) * (1/srange) * sscale + end for x in range(steps)] + + if sgm: + sigmas = sigmas[:-1] + if pad: + sigmas = torch.tensor(sigmas+[0]) + else: + sigmas = torch.tensor(sigmas) + return (sigmas,) + +class tan_scheduler_2stage: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 40, "min": 0,"max": 100000,"step": 1}), + "midpoint": ("INT", {"default": 20, "min": 0,"max": 100000,"step": 1}), + "pivot_1": ("INT", {"default": 10, "min": 0,"max": 100000,"step": 1}), + "pivot_2": ("INT", {"default": 30, "min": 0,"max": 100000,"step": 1}), + "slope_1": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.1}), + "slope_2": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.1}), + "start": ("FLOAT", {"default": 1.0, "min": -100000,"max": 100000,"step": 0.1}), + "middle": ("FLOAT", {"default": 0.5, "min": -100000,"max": 100000,"step": 0.1}), + "end": ("FLOAT", {"default": 0.0, "min": -100000,"max": 100000,"step": 0.1}), + "pad" : ("BOOLEAN", {"default": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/schedulers" + + def get_tan_sigmas(self, steps, slope, pivot, start, end): + smax = ((2/pi)*atan(-slope*(0-pivot))+1)/2 + smin = ((2/pi)*atan(-slope*((steps-1)-pivot))+1)/2 + + srange = smax-smin + sscale = start - end + + sigmas = [ ( (((2/pi)*atan(-slope*(x-pivot))+1)/2) - smin) * (1/srange) * sscale + end for x in range(steps)] + + return sigmas + + def main(self, steps, midpoint, start, middle, end, pivot_1, pivot_2, slope_1, slope_2, pad): + steps += 2 + stage_2_len = steps - midpoint + stage_1_len = steps - stage_2_len + + tan_sigmas_1 = self.get_tan_sigmas(stage_1_len, slope_1, pivot_1, start, middle) + tan_sigmas_2 = self.get_tan_sigmas(stage_2_len, slope_2, pivot_2 - stage_1_len, middle, end) + + tan_sigmas_1 = tan_sigmas_1[:-1] + if pad: + tan_sigmas_2 = tan_sigmas_2+[0] + + tan_sigmas = torch.tensor(tan_sigmas_1 + tan_sigmas_2) + + return (tan_sigmas,) + +class tan_scheduler_2stage_simple: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 40, "min": 0,"max": 100000,"step": 1}), + "pivot_1": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.01}), + "pivot_2": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.01}), + "slope_1": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.01}), + "slope_2": ("FLOAT", {"default": 1, "min": -100000,"max": 100000,"step": 0.01}), + "start": ("FLOAT", {"default": 1.0, "min": -100000,"max": 100000,"step": 0.01}), + "middle": ("FLOAT", {"default": 0.5, "min": -100000,"max": 100000,"step": 0.01}), + "end": ("FLOAT", {"default": 0.0, "min": -100000,"max": 100000,"step": 0.01}), + "pad" : ("BOOLEAN", {"default": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/schedulers" + + def get_tan_sigmas(self, steps, slope, pivot, start, end): + smax = ((2/pi)*atan(-slope*(0-pivot))+1)/2 + smin = ((2/pi)*atan(-slope*((steps-1)-pivot))+1)/2 + + srange = smax-smin + sscale = start - end + + sigmas = [ ( (((2/pi)*atan(-slope*(x-pivot))+1)/2) - smin) * (1/srange) * sscale + end for x in range(steps)] + + return sigmas + + def main(self, steps, start=1.0, middle=0.5, end=0.0, pivot_1=0.6, pivot_2=0.6, slope_1=0.2, slope_2=0.2, pad=False, model_sampling=None): + steps += 2 + + midpoint = int( (steps*pivot_1 + steps*pivot_2) / 2 ) + pivot_1 = int(steps * pivot_1) + pivot_2 = int(steps * pivot_2) + + slope_1 = slope_1 / (steps/40) + slope_2 = slope_2 / (steps/40) + + stage_2_len = steps - midpoint + stage_1_len = steps - stage_2_len + + tan_sigmas_1 = self.get_tan_sigmas(stage_1_len, slope_1, pivot_1, start, middle) + tan_sigmas_2 = self.get_tan_sigmas(stage_2_len, slope_2, pivot_2 - stage_1_len, middle, end) + + tan_sigmas_1 = tan_sigmas_1[:-1] + if pad: + tan_sigmas_2 = tan_sigmas_2+[0] + + tan_sigmas = torch.tensor(tan_sigmas_1 + tan_sigmas_2) + + return (tan_sigmas,) + +class linear_quadratic_advanced: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "steps": ("INT", {"default": 40, "min": 0,"max": 100000,"step": 1}), + "denoise": ("FLOAT", {"default": 1.0, "min": -100000,"max": 100000,"step": 0.01}), + "inflection_percent": ("FLOAT", {"default": 0.5, "min": 0,"max": 1,"step": 0.01}), + "threshold_noise": ("FLOAT", {"default": 0.025, "min": 0.001,"max": 1.000,"step": 0.001}), + }, + # "optional": { + # } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/schedulers" + + def main(self, steps, denoise, inflection_percent, threshold_noise, model=None): + sigmas = get_sigmas(model, "linear_quadratic", steps, denoise, 0.0, inflection_percent, threshold_noise) + + return (sigmas, ) + + +class constant_scheduler: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 40, "min": 0,"max": 100000,"step": 1}), + "value_start": ("FLOAT", {"default": 1.0, "min": -100000,"max": 100000,"step": 0.01}), + "value_end": ("FLOAT", {"default": 0.0, "min": -100000,"max": 100000,"step": 0.01}), + "cutoff_percent": ("FLOAT", {"default": 1.0, "min": 0,"max": 1,"step": 0.01}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + CATEGORY = "RES4LYF/schedulers" + + def main(self, steps, value_start, value_end, cutoff_percent): + sigmas = torch.ones(steps + 1) * value_start + cutoff_step = int(round(steps * cutoff_percent)) + 1 + sigmas = torch.concat((sigmas[:cutoff_step], torch.ones(steps + 1 - cutoff_step) * value_end), dim=0) + + return (sigmas,) + + + + + + +class ClownScheduler: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "pad_start_value": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "start_value": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "end_value": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "pad_end_value": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "scheduler": (["constant"] + get_res4lyf_scheduler_list(), {"default": "beta57"},), + "scheduler_start_step": ("INT", {"default": 0, "min": 0, "max": 10000}), + "scheduler_end_step": ("INT", {"default": 30, "min": -1, "max": 10000}), + "total_steps": ("INT", {"default": 100, "min": -1, "max": 10000}), + "flip_schedule": ("BOOLEAN", {"default": False}), + }, + "optional": { + "model": ("MODEL", ), + } + } + + RETURN_TYPES = ("SIGMAS",) + RETURN_NAMES = ("sigmas",) + FUNCTION = "main" + CATEGORY = "RES4LYF/schedulers" + + def create_callback(self, **kwargs): + def callback(model): + kwargs["model"] = model + schedule, = self.prepare_schedule(**kwargs) + return schedule + return callback + + def main(self, + model = None, + pad_start_value : float = 1.0, + start_value : float = 0.0, + end_value : float = 1.0, + pad_end_value = None, + denoise : int = 1.0, + scheduler = None, + scheduler_start_step : int = 0, + scheduler_end_step : int = 30, + total_steps : int = 60, + flip_schedule = False, + ) -> Tuple[Tensor]: + + if model is None: + callback = self.create_callback(pad_start_value = pad_start_value, + start_value = start_value, + end_value = end_value, + pad_end_value = pad_end_value, + + scheduler = scheduler, + start_step = scheduler_start_step, + end_step = scheduler_end_step, + flip_schedule = flip_schedule, + ) + else: + default_dtype = torch.float64 + default_device = torch.device("cuda") + + if scheduler_end_step == -1: + scheduler_total_steps = total_steps - scheduler_start_step + else: + scheduler_total_steps = scheduler_end_step - scheduler_start_step + + if total_steps == -1: + total_steps = scheduler_start_step + scheduler_end_step + + end_pad_steps = total_steps - scheduler_end_step + + if scheduler != "constant": + values = get_sigmas(model, scheduler, scheduler_total_steps, denoise).to(dtype=default_dtype, device=default_device) + values = ((values - values.min()) * (start_value - end_value)) / (values.max() - values.min()) + end_value + else: + values = torch.linspace(start_value, end_value, scheduler_total_steps, dtype=default_dtype, device=default_device) + + if flip_schedule: + values = torch.flip(values, dims=[0]) + + prepend = torch.full((scheduler_start_step,), pad_start_value, dtype=default_dtype, device=default_device) + postpend = torch.full((end_pad_steps,), pad_end_value, dtype=default_dtype, device=default_device) + + values = torch.cat((prepend, values, postpend), dim=0) + + #ositive[0][1]['callback_regional'] = callback + + return (values,) + + + + def prepare_schedule(self, + model = None, + pad_start_value : float = 1.0, + start_value : float = 0.0, + end_value : float = 1.0, + pad_end_value = None, + weight_scheduler = None, + start_step : int = 0, + end_step : int = 30, + flip_schedule = False, + ) -> Tuple[Tensor]: + + default_dtype = torch.float64 + default_device = torch.device("cuda") + + return (None,) + + + + +def get_sigmas_simple_exponential(model, steps): + s = model.model_sampling + sigs = [] + ss = len(s.sigmas) / steps + for x in range(steps): + sigs += [float(s.sigmas[-(1 + int(x * ss))])] + sigs += [0.0] + sigs = torch.FloatTensor(sigs) + exp = torch.exp(torch.log(torch.linspace(1, 0, steps + 1))) + return sigs * exp + +extra_schedulers = { + "simple_exponential": get_sigmas_simple_exponential +} + + + +def get_sigmas(model, scheduler, steps, denoise, shift=0.0, lq_inflection_percent=0.5, lq_threshold_noise=0.025): #adapted from comfyui + total_steps = steps + if denoise < 1.0: + if denoise <= 0.0: + return (torch.FloatTensor([]),) + total_steps = int(steps/denoise) + + try: + model_sampling = model.get_model_object("model_sampling") + except: + if hasattr(model, "model"): + model_sampling = model.model.model_sampling + elif hasattr(model, "inner_model"): + model_sampling = model.inner_model.inner_model.model_sampling + else: + raise Exception("get_sigmas: Could not get model_sampling") + + if shift > 1e-6: + import copy + model_sampling = copy.deepcopy(model_sampling) + model_sampling.set_parameters(shift=shift) + RESplain("model_sampling shift manually set to " + str(shift), debug=True) + + if scheduler == "beta57": + sigmas = comfy.samplers.beta_scheduler(model_sampling, total_steps, alpha=0.5, beta=0.7).cpu() + elif scheduler == "linear_quadratic": + linear_steps = int(total_steps * lq_inflection_percent) + sigmas = comfy.samplers.linear_quadratic_schedule(model_sampling, total_steps, threshold_noise=lq_threshold_noise, linear_steps=linear_steps).cpu() + else: + sigmas = comfy.samplers.calculate_sigmas(model_sampling, scheduler, total_steps).cpu() + + sigmas = sigmas[-(steps + 1):] + return sigmas + +#/// Adam Kormendi /// Inspired from Unreal Engine Maths /// + + +# Sigmoid Function +class sigmas_sigmoid: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "variant": (["logistic", "tanh", "softsign", "hardswish", "mish", "swish"], {"default": "logistic"}), + "gain": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "offset": ("FLOAT", {"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, variant, gain, offset, normalize_output): + # Apply gain and offset + x = gain * (sigmas + offset) + + if variant == "logistic": + result = 1.0 / (1.0 + torch.exp(-x)) + elif variant == "tanh": + result = torch.tanh(x) + elif variant == "softsign": + result = x / (1.0 + torch.abs(x)) + elif variant == "hardswish": + result = x * torch.minimum(torch.maximum(x + 3, torch.zeros_like(x)), torch.tensor(6.0)) / 6.0 + elif variant == "mish": + result = x * torch.tanh(torch.log(1.0 + torch.exp(x))) + elif variant == "swish": + result = x * torch.sigmoid(x) + + if normalize_output: + # Normalize to [min(sigmas), max(sigmas)] + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Easing Function ----- +class sigmas_easing: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "easing_type": (["sine", "quad", "cubic", "quart", "quint", "expo", "circ", + "back", "elastic", "bounce"], {"default": "cubic"}), + "easing_mode": (["in", "out", "in_out"], {"default": "in_out"}), + "normalize_input": ("BOOLEAN", {"default": True}), + "normalize_output": ("BOOLEAN", {"default": True}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, easing_type, easing_mode, normalize_input, normalize_output, strength): + # Normalize input to [0, 1] if requested + if normalize_input: + t = (sigmas - sigmas.min()) / (sigmas.max() - sigmas.min()) + else: + t = torch.clamp(sigmas, 0.0, 1.0) + + # Apply strength + t_orig = t.clone() + t = t ** strength + + # Apply easing function based on type and mode + if easing_mode == "in": + result = self._ease_in(t, easing_type) + elif easing_mode == "out": + result = self._ease_out(t, easing_type) + else: # in_out + result = self._ease_in_out(t, easing_type) + + # Normalize output if requested + if normalize_output: + if normalize_input: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + else: + result = ((result - result.min()) / (result.max() - result.min())) + + return (result,) + + def _ease_in(self, t, easing_type): + if easing_type == "sine": + return 1 - torch.cos((t * math.pi) / 2) + elif easing_type == "quad": + return t * t + elif easing_type == "cubic": + return t * t * t + elif easing_type == "quart": + return t * t * t * t + elif easing_type == "quint": + return t * t * t * t * t + elif easing_type == "expo": + return torch.where(t == 0, torch.zeros_like(t), torch.pow(2, 10 * t - 10)) + elif easing_type == "circ": + return 1 - torch.sqrt(1 - torch.pow(t, 2)) + elif easing_type == "back": + c1 = 1.70158 + c3 = c1 + 1 + return c3 * t * t * t - c1 * t * t + elif easing_type == "elastic": + c4 = (2 * math.pi) / 3 + return torch.where( + t == 0, + torch.zeros_like(t), + torch.where( + t == 1, + torch.ones_like(t), + -torch.pow(2, 10 * t - 10) * torch.sin((t * 10 - 10.75) * c4) + ) + ) + elif easing_type == "bounce": + return 1 - self._ease_out_bounce(1 - t) + + def _ease_out(self, t, easing_type): + if easing_type == "sine": + return torch.sin((t * math.pi) / 2) + elif easing_type == "quad": + return 1 - (1 - t) * (1 - t) + elif easing_type == "cubic": + return 1 - torch.pow(1 - t, 3) + elif easing_type == "quart": + return 1 - torch.pow(1 - t, 4) + elif easing_type == "quint": + return 1 - torch.pow(1 - t, 5) + elif easing_type == "expo": + return torch.where(t == 1, torch.ones_like(t), 1 - torch.pow(2, -10 * t)) + elif easing_type == "circ": + return torch.sqrt(1 - torch.pow(t - 1, 2)) + elif easing_type == "back": + c1 = 1.70158 + c3 = c1 + 1 + return 1 + c3 * torch.pow(t - 1, 3) + c1 * torch.pow(t - 1, 2) + elif easing_type == "elastic": + c4 = (2 * math.pi) / 3 + return torch.where( + t == 0, + torch.zeros_like(t), + torch.where( + t == 1, + torch.ones_like(t), + torch.pow(2, -10 * t) * torch.sin((t * 10 - 0.75) * c4) + 1 + ) + ) + elif easing_type == "bounce": + return self._ease_out_bounce(t) + + def _ease_in_out(self, t, easing_type): + if easing_type == "sine": + return -(torch.cos(math.pi * t) - 1) / 2 + elif easing_type == "quad": + return torch.where(t < 0.5, 2 * t * t, 1 - torch.pow(-2 * t + 2, 2) / 2) + elif easing_type == "cubic": + return torch.where(t < 0.5, 4 * t * t * t, 1 - torch.pow(-2 * t + 2, 3) / 2) + elif easing_type == "quart": + return torch.where(t < 0.5, 8 * t * t * t * t, 1 - torch.pow(-2 * t + 2, 4) / 2) + elif easing_type == "quint": + return torch.where(t < 0.5, 16 * t * t * t * t * t, 1 - torch.pow(-2 * t + 2, 5) / 2) + elif easing_type == "expo": + return torch.where( + t < 0.5, + torch.pow(2, 20 * t - 10) / 2, + (2 - torch.pow(2, -20 * t + 10)) / 2 + ) + elif easing_type == "circ": + return torch.where( + t < 0.5, + (1 - torch.sqrt(1 - torch.pow(2 * t, 2))) / 2, + (torch.sqrt(1 - torch.pow(-2 * t + 2, 2)) + 1) / 2 + ) + elif easing_type == "back": + c1 = 1.70158 + c2 = c1 * 1.525 + return torch.where( + t < 0.5, + (torch.pow(2 * t, 2) * ((c2 + 1) * 2 * t - c2)) / 2, + (torch.pow(2 * t - 2, 2) * ((c2 + 1) * (t * 2 - 2) + c2) + 2) / 2 + ) + elif easing_type == "elastic": + c5 = (2 * math.pi) / 4.5 + return torch.where( + t < 0.5, + -(torch.pow(2, 20 * t - 10) * torch.sin((20 * t - 11.125) * c5)) / 2, + (torch.pow(2, -20 * t + 10) * torch.sin((20 * t - 11.125) * c5)) / 2 + 1 + ) + elif easing_type == "bounce": + return torch.where( + t < 0.5, + (1 - self._ease_out_bounce(1 - 2 * t)) / 2, + (1 + self._ease_out_bounce(2 * t - 1)) / 2 + ) + + def _ease_out_bounce(self, t): + n1 = 7.5625 + d1 = 2.75 + + mask1 = t < 1 / d1 + mask2 = t < 2 / d1 + mask3 = t < 2.5 / d1 + + result = torch.zeros_like(t) + result = torch.where(mask1, n1 * t * t, result) + result = torch.where(mask2 & ~mask1, n1 * (t - 1.5 / d1) * (t - 1.5 / d1) + 0.75, result) + result = torch.where(mask3 & ~mask2, n1 * (t - 2.25 / d1) * (t - 2.25 / d1) + 0.9375, result) + result = torch.where(~mask3, n1 * (t - 2.625 / d1) * (t - 2.625 / d1) + 0.984375, result) + + return result + +# ----- Hyperbolic Function ----- +class sigmas_hyperbolic: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "function": (["sinh", "cosh", "tanh", "asinh", "acosh", "atanh"], {"default": "tanh"}), + "scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, function, scale, normalize_output): + # Apply scaling + x = sigmas * scale + + if function == "sinh": + result = torch.sinh(x) + elif function == "cosh": + result = torch.cosh(x) + elif function == "tanh": + result = torch.tanh(x) + elif function == "asinh": + result = torch.asinh(x) + elif function == "acosh": + # Domain of acosh is [1, inf) + result = torch.acosh(torch.clamp(x, min=1.0)) + elif function == "atanh": + # Domain of atanh is (-1, 1) + result = torch.atanh(torch.clamp(x, min=-0.99, max=0.99)) + + if normalize_output: + # Normalize to [min(sigmas), max(sigmas)] + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Gaussian Distribution Function ----- +class sigmas_gaussian: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "mean": ("FLOAT", {"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.01}), + "std": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "operation": (["pdf", "cdf", "inverse_cdf", "transform", "modulate"], {"default": "transform"}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, mean, std, operation, normalize_output): + # Standardize values (z-score) + z = (sigmas - sigmas.mean()) / sigmas.std() + + if operation == "pdf": + # Probability density function + result = (1 / (std * math.sqrt(2 * math.pi))) * torch.exp(-0.5 * ((sigmas - mean) / std) ** 2) + elif operation == "cdf": + # Cumulative distribution function + result = 0.5 * (1 + torch.erf((sigmas - mean) / (std * math.sqrt(2)))) + elif operation == "inverse_cdf": + # Inverse CDF (quantile function) + # First normalize to [0.01, 0.99] to avoid numerical issues + normalized = ((sigmas - sigmas.min()) / (sigmas.max() - sigmas.min())) * 0.98 + 0.01 + result = mean + std * torch.sqrt(2) * torch.erfinv(2 * normalized - 1) + elif operation == "transform": + # Transform to Gaussian distribution with specified mean and std + result = z * std + mean + elif operation == "modulate": + # Modulate with a Gaussian curve centered at mean + result = sigmas * torch.exp(-0.5 * ((sigmas - mean) / std) ** 2) + + if normalize_output: + # Normalize to [min(sigmas), max(sigmas)] + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Percentile Function ----- +class sigmas_percentile: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "percentile_min": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 49.0, "step": 0.1}), + "percentile_max": ("FLOAT", {"default": 95.0, "min": 51.0, "max": 100.0, "step": 0.1}), + "target_min": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "target_max": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "clip_outliers": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, percentile_min, percentile_max, target_min, target_max, clip_outliers): + # Convert to numpy for percentile computation + sigmas_np = sigmas.cpu().numpy() + + # Compute percentiles + p_min = np.percentile(sigmas_np, percentile_min) + p_max = np.percentile(sigmas_np, percentile_max) + + # Convert back to tensor + p_min = torch.tensor(p_min, device=sigmas.device, dtype=sigmas.dtype) + p_max = torch.tensor(p_max, device=sigmas.device, dtype=sigmas.dtype) + + # Map values from [p_min, p_max] to [target_min, target_max] + if clip_outliers: + sigmas_clipped = torch.clamp(sigmas, p_min, p_max) + result = ((sigmas_clipped - p_min) / (p_max - p_min)) * (target_max - target_min) + target_min + else: + result = ((sigmas - p_min) / (p_max - p_min)) * (target_max - target_min) + target_min + + return (result,) + +# ----- Kernel Smooth Function ----- +class sigmas_kernel_smooth: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "kernel": (["gaussian", "box", "triangle", "epanechnikov", "cosine"], {"default": "gaussian"}), + "kernel_size": ("INT", {"default": 5, "min": 3, "max": 51, "step": 2}), # Must be odd + "sigma": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, kernel, kernel_size, sigma): + # Ensure kernel_size is odd + if kernel_size % 2 == 0: + kernel_size += 1 + + # Define kernel weights + if kernel == "gaussian": + # Gaussian kernel + kernel_1d = self._gaussian_kernel(kernel_size, sigma) + elif kernel == "box": + # Box (uniform) kernel + kernel_1d = torch.ones(kernel_size, device=sigmas.device, dtype=sigmas.dtype) / kernel_size + elif kernel == "triangle": + # Triangle kernel + x = torch.linspace(-(kernel_size//2), kernel_size//2, kernel_size, device=sigmas.device, dtype=sigmas.dtype) + kernel_1d = (1.0 - torch.abs(x) / (kernel_size//2)) + kernel_1d = kernel_1d / kernel_1d.sum() + elif kernel == "epanechnikov": + # Epanechnikov kernel + x = torch.linspace(-(kernel_size//2), kernel_size//2, kernel_size, device=sigmas.device, dtype=sigmas.dtype) + x = x / (kernel_size//2) # Scale to [-1, 1] + kernel_1d = 0.75 * (1 - x**2) + kernel_1d = kernel_1d / kernel_1d.sum() + elif kernel == "cosine": + # Cosine kernel + x = torch.linspace(-(kernel_size//2), kernel_size//2, kernel_size, device=sigmas.device, dtype=sigmas.dtype) + x = x / (kernel_size//2) * (math.pi/2) # Scale to [-π/2, π/2] + kernel_1d = torch.cos(x) + kernel_1d = kernel_1d / kernel_1d.sum() + + # Pad input to handle boundary conditions + pad_size = kernel_size // 2 + padded = F.pad(sigmas.unsqueeze(0).unsqueeze(0), (pad_size, pad_size), mode='reflect') + + # Apply convolution + smoothed = F.conv1d(padded, kernel_1d.unsqueeze(0).unsqueeze(0)) + + return (smoothed.squeeze(),) + + def _gaussian_kernel(self, kernel_size, sigma): + # Generate 1D Gaussian kernel + x = torch.linspace(-(kernel_size//2), kernel_size//2, kernel_size) + kernel = torch.exp(-x**2 / (2*sigma**2)) + return kernel / kernel.sum() + +# ----- Quantile Normalization ----- +class sigmas_quantile_norm: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "target_distribution": (["uniform", "normal", "exponential", "logistic", "custom"], {"default": "uniform"}), + "num_quantiles": ("INT", {"default": 100, "min": 10, "max": 1000, "step": 10}), + }, + "optional": { + "reference_sigmas": ("SIGMAS", {"forceInput": False}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, target_distribution, num_quantiles, reference_sigmas=None): + # Convert to numpy for processing + sigmas_np = sigmas.cpu().numpy() + + # Sort values + sorted_values = np.sort(sigmas_np) + + # Create rank for each value (fractional rank) + ranks = np.zeros_like(sigmas_np) + for i, val in enumerate(sigmas_np): + ranks[i] = np.searchsorted(sorted_values, val, side='right') / len(sorted_values) + + # Generate target distribution + if target_distribution == "uniform": + # Uniform distribution between min and max of sigmas + target_values = np.linspace(sigmas_np.min(), sigmas_np.max(), num_quantiles) + elif target_distribution == "normal": + # Normal distribution with same mean and std as sigmas + target_values = np.random.normal(sigmas_np.mean(), sigmas_np.std(), num_quantiles) + target_values.sort() + elif target_distribution == "exponential": + # Exponential distribution with lambda=1/mean + target_values = np.random.exponential(1/max(1e-6, sigmas_np.mean()), num_quantiles) + target_values.sort() + elif target_distribution == "logistic": + # Logistic distribution + target_values = np.random.logistic(0, 1, num_quantiles) + target_values.sort() + # Rescale to match sigmas range + target_values = (target_values - target_values.min()) / (target_values.max() - target_values.min()) + target_values = target_values * (sigmas_np.max() - sigmas_np.min()) + sigmas_np.min() + elif target_distribution == "custom" and reference_sigmas is not None: + # Use provided reference distribution + reference_np = reference_sigmas.cpu().numpy() + target_values = np.sort(reference_np) + if len(target_values) < num_quantiles: + # Interpolate if reference is smaller + old_indices = np.linspace(0, len(target_values)-1, len(target_values)) + new_indices = np.linspace(0, len(target_values)-1, num_quantiles) + target_values = np.interp(new_indices, old_indices, target_values) + else: + # Subsample if reference is larger + indices = np.linspace(0, len(target_values)-1, num_quantiles, dtype=int) + target_values = target_values[indices] + else: + # Default to uniform + target_values = np.linspace(sigmas_np.min(), sigmas_np.max(), num_quantiles) + + # Map each value to its corresponding quantile in the target distribution + result_np = np.interp(ranks, np.linspace(0, 1, len(target_values)), target_values) + + # Convert back to tensor + result = torch.tensor(result_np, device=sigmas.device, dtype=sigmas.dtype) + + return (result,) + +# ----- Adaptive Step Function ----- +class sigmas_adaptive_step: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "adaptation_type": (["gradient", "curvature", "importance", "density"], {"default": "gradient"}), + "sensitivity": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + "min_step": ("FLOAT", {"default": 0.01, "min": 0.0001, "max": 1.0, "step": 0.01}), + "max_step": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "target_steps": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, adaptation_type, sensitivity, min_step, max_step, target_steps): + if len(sigmas) <= 1: + return (sigmas,) + + # Compute step sizes based on chosen adaptation type + if adaptation_type == "gradient": + # Compute gradient (first difference) + grads = torch.abs(sigmas[1:] - sigmas[:-1]) + # Normalize gradients + if grads.max() > grads.min(): + norm_grads = (grads - grads.min()) / (grads.max() - grads.min()) + else: + norm_grads = torch.ones_like(grads) + + # Convert to step sizes: smaller steps where gradient is large + step_sizes = 1.0 / (1.0 + norm_grads * sensitivity) + + elif adaptation_type == "curvature": + # Compute second derivative approximation + if len(sigmas) >= 3: + # Second difference + second_diff = sigmas[2:] - 2*sigmas[1:-1] + sigmas[:-2] + # Pad to match length + second_diff = F.pad(second_diff, (0, 1), mode='replicate') + else: + second_diff = torch.zeros_like(sigmas[:-1]) + + # Normalize curvature + abs_curve = torch.abs(second_diff) + if abs_curve.max() > abs_curve.min(): + norm_curve = (abs_curve - abs_curve.min()) / (abs_curve.max() - abs_curve.min()) + else: + norm_curve = torch.ones_like(abs_curve) + + # Convert to step sizes: smaller steps where curvature is high + step_sizes = 1.0 / (1.0 + norm_curve * sensitivity) + + elif adaptation_type == "importance": + # Importance based on values: focus more on extremes + centered = torch.abs(sigmas - sigmas.mean()) + if centered.max() > centered.min(): + importance = (centered - centered.min()) / (centered.max() - centered.min()) + else: + importance = torch.ones_like(centered) + + # Steps are smaller for important regions + step_sizes = 1.0 / (1.0 + importance[:-1] * sensitivity) + + elif adaptation_type == "density": + # Density-based adaptation using kernel density estimation + # Use a simple histogram approximation + sigma_min, sigma_max = sigmas.min(), sigmas.max() + bins = 20 + hist = torch.histc(sigmas, bins=bins, min=sigma_min, max=sigma_max) + hist = hist / hist.sum() # Normalize + + # Map each sigma to its bin density + bin_indices = torch.floor((sigmas - sigma_min) / (sigma_max - sigma_min) * (bins-1)).long() + bin_indices = torch.clamp(bin_indices, 0, bins-1) + densities = hist[bin_indices] + + # Compute step sizes: smaller steps in high density regions + step_sizes = 1.0 / (1.0 + densities[:-1] * sensitivity) + + # Scale step sizes to [min_step, max_step] + if step_sizes.max() > step_sizes.min(): + step_sizes = (step_sizes - step_sizes.min()) / (step_sizes.max() - step_sizes.min()) + step_sizes = step_sizes * (max_step - min_step) + min_step + else: + step_sizes = torch.ones_like(step_sizes) * min_step + + # Cumulative sum to get positions + positions = torch.cat([torch.tensor([0.0], device=step_sizes.device), torch.cumsum(step_sizes, dim=0)]) + + # Normalize positions to match original range + positions = positions / positions[-1] * (sigmas[-1] - sigmas[0]) + sigmas[0] + + # Resample if target_steps is specified + if target_steps > 0: + new_positions = torch.linspace(sigmas[0], sigmas[-1], target_steps, device=sigmas.device) + # Interpolate to get new sigma values + new_sigmas = torch.zeros_like(new_positions) + + # Simple linear interpolation + for i, pos in enumerate(new_positions): + # Find enclosing original positions + idx = torch.searchsorted(positions, pos) + idx = torch.clamp(idx, 1, len(positions)-1) + + # Linear interpolation + t = (pos - positions[idx-1]) / (positions[idx] - positions[idx-1]) + new_sigmas[i] = sigmas[idx-1] * (1-t) + sigmas[idx-1] * t + + result = new_sigmas + else: + result = positions + + return (result,) + +# ----- Chaos Function ----- +class sigmas_chaos: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "system": (["logistic", "henon", "tent", "sine", "cubic"], {"default": "logistic"}), + "parameter": ("FLOAT", {"default": 3.9, "min": 0.1, "max": 5.0, "step": 0.01}), + "iterations": ("INT", {"default": 10, "min": 1, "max": 100, "step": 1}), + "normalize_output": ("BOOLEAN", {"default": True}), + "use_as_seed": ("BOOLEAN", {"default": False}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, system, parameter, iterations, normalize_output, use_as_seed): + # Normalize input to [0,1] for chaotic maps + if use_as_seed: + # Use input as initial seed + x = (sigmas - sigmas.min()) / (sigmas.max() - sigmas.min()) + else: + # Use single initial value and apply iterations + x = torch.zeros_like(sigmas) + for i in range(len(sigmas)): + # Use i/len as initial value for variety + x[i] = i / len(sigmas) + + # Apply chaos map iterations + for _ in range(iterations): + if system == "logistic": + # Logistic map: x_{n+1} = r * x_n * (1 - x_n) + x = parameter * x * (1 - x) + + elif system == "henon": + # Simplified 1D version of Henon map + x = 1 - parameter * x**2 + + elif system == "tent": + # Tent map + x = torch.where(x < 0.5, parameter * x, parameter * (1 - x)) + + elif system == "sine": + # Sine map: x_{n+1} = r * sin(pi * x_n) + x = parameter * torch.sin(math.pi * x) + + elif system == "cubic": + # Cubic map: x_{n+1} = r * x_n * (1 - x_n^2) + x = parameter * x * (1 - x**2) + + # Normalize output if requested + if normalize_output: + result = ((x - x.min()) / (x.max() - x.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + else: + result = x + + return (result,) + +# ----- Reaction Diffusion Function ----- +class sigmas_reaction_diffusion: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "system": (["gray_scott", "fitzhugh_nagumo", "brusselator"], {"default": "gray_scott"}), + "iterations": ("INT", {"default": 10, "min": 1, "max": 100, "step": 1}), + "dt": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}), + "param_a": ("FLOAT", {"default": 0.04, "min": 0.01, "max": 0.1, "step": 0.001}), + "param_b": ("FLOAT", {"default": 0.06, "min": 0.01, "max": 0.1, "step": 0.001}), + "diffusion_a": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}), + "diffusion_b": ("FLOAT", {"default": 0.05, "min": 0.01, "max": 1.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, system, iterations, dt, param_a, param_b, diffusion_a, diffusion_b, normalize_output): + # Initialize a and b based on sigmas + a = (sigmas - sigmas.min()) / (sigmas.max() - sigmas.min()) + b = 1.0 - a + + # Pad for diffusion calculation (periodic boundary) + a_pad = F.pad(a.unsqueeze(0).unsqueeze(0), (1, 1), mode='circular').squeeze() + b_pad = F.pad(b.unsqueeze(0).unsqueeze(0), (1, 1), mode='circular').squeeze() + + # Simple 1D reaction-diffusion + for _ in range(iterations): + # Compute Laplacian (diffusion term) as second derivative + laplacian_a = a_pad[:-2] + a_pad[2:] - 2 * a + laplacian_b = b_pad[:-2] + b_pad[2:] - 2 * b + + if system == "gray_scott": + # Gray-Scott model for pattern formation + # a is "U" (activator), b is "V" (inhibitor) + feed = 0.055 # feed rate + kill = 0.062 # kill rate + + # Update equations + a_new = a + dt * (diffusion_a * laplacian_a - a * b**2 + feed * (1 - a)) + b_new = b + dt * (diffusion_b * laplacian_b + a * b**2 - (feed + kill) * b) + + elif system == "fitzhugh_nagumo": + # FitzHugh-Nagumo model (simplified) + # a is the membrane potential, b is the recovery variable + + # Update equations + a_new = a + dt * (diffusion_a * laplacian_a + a - a**3 - b + param_a) + b_new = b + dt * (diffusion_b * laplacian_b + param_b * (a - b)) + + elif system == "brusselator": + # Brusselator model + # a is U, b is V + + # Update equations + a_new = a + dt * (diffusion_a * laplacian_a + 1 - (param_b + 1) * a + param_a * a**2 * b) + b_new = b + dt * (diffusion_b * laplacian_b + param_b * a - param_a * a**2 * b) + + # Update and repad + a, b = a_new, b_new + a_pad = F.pad(a.unsqueeze(0).unsqueeze(0), (1, 1), mode='circular').squeeze() + b_pad = F.pad(b.unsqueeze(0).unsqueeze(0), (1, 1), mode='circular').squeeze() + + # Use the activator component as the result + result = a + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Attractor Function ----- +class sigmas_attractor: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "attractor": (["lorenz", "rossler", "aizawa", "chen", "thomas"], {"default": "lorenz"}), + "iterations": ("INT", {"default": 5, "min": 1, "max": 50, "step": 1}), + "dt": ("FLOAT", {"default": 0.01, "min": 0.001, "max": 0.1, "step": 0.001}), + "component": (["x", "y", "z", "magnitude"], {"default": "x"}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, attractor, iterations, dt, component, normalize_output): + # Initialize 3D state from sigmas + n = len(sigmas) + + # Normalize sigmas to a reasonable range for the attractor + norm_sigmas = (sigmas - sigmas.min()) / (sigmas.max() - sigmas.min()) * 2.0 - 1.0 + + # Create initial state + x = norm_sigmas + y = torch.roll(norm_sigmas, 1) # Shifted version for variety + z = torch.roll(norm_sigmas, 2) # Another shifted version + + # Parameters for the attractors + if attractor == "lorenz": + sigma, rho, beta = 10.0, 28.0, 8.0/3.0 + elif attractor == "rossler": + a, b, c = 0.2, 0.2, 5.7 + elif attractor == "aizawa": + a, b, c, d, e, f = 0.95, 0.7, 0.6, 3.5, 0.25, 0.1 + elif attractor == "chen": + a, b, c = 5.0, -10.0, -0.38 + elif attractor == "thomas": + b = 0.208186 + + # Run the attractor dynamics + for _ in range(iterations): + if attractor == "lorenz": + # Lorenz attractor + dx = sigma * (y - x) + dy = x * (rho - z) - y + dz = x * y - beta * z + + elif attractor == "rossler": + # Rössler attractor + dx = -y - z + dy = x + a * y + dz = b + z * (x - c) + + elif attractor == "aizawa": + # Aizawa attractor + dx = (z - b) * x - d * y + dy = d * x + (z - b) * y + dz = c + a * z - z**3/3 - (x**2 + y**2) * (1 + e * z) + f * z * x**3 + + elif attractor == "chen": + # Chen attractor + dx = a * (y - x) + dy = (c - a) * x - x * z + c * y + dz = x * y - b * z + + elif attractor == "thomas": + # Thomas attractor + dx = -b * x + torch.sin(y) + dy = -b * y + torch.sin(z) + dz = -b * z + torch.sin(x) + + # Update state + x = x + dt * dx + y = y + dt * dy + z = z + dt * dz + + # Select component + if component == "x": + result = x + elif component == "y": + result = y + elif component == "z": + result = z + elif component == "magnitude": + result = torch.sqrt(x**2 + y**2 + z**2) + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Catmull-Rom Spline ----- +class sigmas_catmull_rom: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "tension": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "points": ("INT", {"default": 100, "min": 5, "max": 1000, "step": 5}), + "boundary_condition": (["repeat", "clamp", "mirror"], {"default": "clamp"}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, tension, points, boundary_condition): + n = len(sigmas) + + # Need at least 4 points for Catmull-Rom interpolation + if n < 4: + # If we have fewer, just use linear interpolation + t = torch.linspace(0, 1, points, device=sigmas.device) + result = torch.zeros(points, device=sigmas.device, dtype=sigmas.dtype) + + for i in range(points): + idx = min(int(i * (n - 1) / (points - 1)), n - 2) + alpha = (i * (n - 1) / (points - 1)) - idx + result[i] = (1 - alpha) * sigmas[idx] + alpha * sigmas[idx + 1] + + return (result,) + + # Handle boundary conditions for control points + if boundary_condition == "repeat": + # Repeat endpoints + p0 = sigmas[0] + p3 = sigmas[-1] + elif boundary_condition == "clamp": + # Extrapolate + p0 = 2 * sigmas[0] - sigmas[1] + p3 = 2 * sigmas[-1] - sigmas[-2] + elif boundary_condition == "mirror": + # Mirror + p0 = sigmas[1] + p3 = sigmas[-2] + + # Create extended control points + control_points = torch.cat([torch.tensor([p0], device=sigmas.device), sigmas, torch.tensor([p3], device=sigmas.device)]) + + # Compute spline + result = torch.zeros(points, device=sigmas.device, dtype=sigmas.dtype) + + # Parameter to adjust curve tension (0 = Catmull-Rom, 1 = Linear) + alpha = 1.0 - tension + + for i in range(points): + # Determine which segment we're in + t = i / (points - 1) * (n - 1) + idx = min(int(t), n - 2) + + # Normalized parameter within the segment [0, 1] + t_local = t - idx + + # Get control points for this segment + p0 = control_points[idx] + p1 = control_points[idx + 1] + p2 = control_points[idx + 2] + p3 = control_points[idx + 3] + + # Catmull-Rom basis functions + t2 = t_local * t_local + t3 = t2 * t_local + + # Compute spline point + result[i] = ( + (-alpha * t3 + 2 * alpha * t2 - alpha * t_local) * p0 + + ((2 - alpha) * t3 + (alpha - 3) * t2 + 1) * p1 + + ((alpha - 2) * t3 + (3 - 2 * alpha) * t2 + alpha * t_local) * p2 + + (alpha * t3 - alpha * t2) * p3 + ) * 0.5 + + return (result,) + +# ----- Lambert W-Function ----- +class sigmas_lambert_w: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "branch": (["principal", "secondary"], {"default": "principal"}), + "scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}), + "max_iterations": ("INT", {"default": 20, "min": 5, "max": 100, "step": 1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, branch, scale, normalize_output, max_iterations): + # Apply scaling + x = sigmas * scale + + # Lambert W function (numerically approximated) + result = torch.zeros_like(x) + + # Process each value separately (since Lambert W is non-vectorized) + for i in range(len(x)): + xi = x[i].item() + + # Initial guess varies by branch + if branch == "principal": + # Valid for x >= -1/e + if xi < -1/math.e: + xi = -1/math.e # Clamp to domain + + # Initial guess for W₀(x) + if xi < 0: + w = 0.0 + elif xi < 1: + w = xi * (1 - xi * (1 - 0.5 * xi)) + else: + w = math.log(xi) + + else: # secondary branch + # Valid for -1/e <= x < 0 + if xi < -1/math.e: + xi = -1/math.e # Clamp to lower bound + elif xi >= 0: + xi = -0.01 # Clamp to upper bound + + # Initial guess for W₋₁(x) + w = math.log(-xi) + + # Halley's method for numerical approximation + for _ in range(max_iterations): + ew = math.exp(w) + wew = w * ew + + # If we've converged, break + if abs(wew - xi) < 1e-10: + break + + # Halley's update + wpe = w + 1 # w plus 1 + div = ew * wpe - (ew * w - xi) * wpe / (2 * wpe * ew) + w_next = w - (wew - xi) / div + + # Check for convergence + if abs(w_next - w) < 1e-10: + w = w_next + break + + w = w_next + + result[i] = w + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Zeta & Eta Functions ----- +class sigmas_zeta_eta: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "function": (["riemann_zeta", "dirichlet_eta", "lerch_phi"], {"default": "riemann_zeta"}), + "offset": ("FLOAT", {"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.1}), + "scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}), + "approx_terms": ("INT", {"default": 100, "min": 10, "max": 1000, "step": 10}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, function, offset, scale, normalize_output, approx_terms): + # Apply offset and scaling + s = sigmas * scale + offset + + # Process based on function type + if function == "riemann_zeta": + # Riemann zeta function + # For Re(s) > 1, ζ(s) = sum(1/n^s, n=1 to infinity) + # For performance reasons, we'll use scipy's implementation for CPU + # and a truncated series approximation for GPU + + # Move to CPU for scipy + s_cpu = s.cpu().numpy() + + # Apply zeta function + result_np = np.zeros_like(s_cpu) + + for i, si in enumerate(s_cpu): + # Handle special values + if si == 1.0: + # ζ(1) is the harmonic series, which diverges to infinity + result_np[i] = float('inf') + elif si < 0 and si == int(si) and int(si) % 2 == 0: + # ζ(-2n) = 0 for n > 0 + result_np[i] = 0.0 + else: + try: + # Use scipy for computation + result_np[i] = float(special.zeta(si)) + except (ValueError, OverflowError): + # Fall back to approximation for problematic values + if si > 1: + # Truncated series for Re(s) > 1 + result_np[i] = sum(1.0 / np.power(n, si) for n in range(1, approx_terms)) + else: + # Use functional equation for Re(s) < 0 + if si < 0: + # ζ(s) = 2^s π^(s-1) sin(πs/2) Γ(1-s) ζ(1-s) + # Gamma function blows up at negative integers, so use the fact that + # ζ(-n) = -B_{n+1}/(n+1) for n > 0, where B is a Bernoulli number + # However, as this gets complex, we'll use a simpler approximation + result_np[i] = 0.0 # Default for problematic values + + # Convert back to tensor + result = torch.tensor(result_np, device=sigmas.device, dtype=sigmas.dtype) + + elif function == "dirichlet_eta": + # Dirichlet eta function (alternating zeta function) + # η(s) = sum((-1)^(n+1)/n^s, n=1 to infinity) + + # For GPU efficiency, compute directly using alternating series + result = torch.zeros_like(s) + + # Use a fixed number of terms for approximation + for i in range(1, approx_terms + 1): + term = torch.pow(i, -s) * (1 if i % 2 == 1 else -1) + result += term + + elif function == "lerch_phi": + # Lerch transcendent with fixed parameters + # Φ(z, s, a) = sum(z^n / (n+a)^s, n=0 to infinity) + # We'll use z=0.5, a=1 for simplicity + z, a = 0.5, 1.0 + + result = torch.zeros_like(s) + for i in range(approx_terms): + term = torch.pow(z, i) / torch.pow(i + a, s) + result += term + + # Replace infinities and NaNs with large or small values + result = torch.where(torch.isfinite(result), result, torch.sign(result) * 1e10) + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Gamma & Beta Functions ----- +class sigmas_gamma_beta: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "function": (["gamma", "beta", "incomplete_gamma", "incomplete_beta", "log_gamma"], {"default": "gamma"}), + "offset": ("FLOAT", {"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.1}), + "scale": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 10.0, "step": 0.01}), + "parameter_a": ("FLOAT", {"default": 0.5, "min": 0.1, "max": 10.0, "step": 0.1}), + "parameter_b": ("FLOAT", {"default": 0.5, "min": 0.1, "max": 10.0, "step": 0.1}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, function, offset, scale, parameter_a, parameter_b, normalize_output): + # Apply offset and scaling + x = sigmas * scale + offset + + # Convert to numpy for special functions + x_np = x.cpu().numpy() + + # Apply function + if function == "gamma": + # Gamma function Γ(x) + # For performance and stability, use scipy + result_np = np.zeros_like(x_np) + + for i, xi in enumerate(x_np): + # Handle special cases + if xi <= 0 and xi == int(xi): + # Gamma has poles at non-positive integers + result_np[i] = float('inf') + else: + try: + result_np[i] = float(special.gamma(xi)) + except (ValueError, OverflowError): + # Use approximation for large values + result_np[i] = float('inf') + + elif function == "log_gamma": + # Log Gamma function log(Γ(x)) + # More numerically stable for large values + result_np = np.zeros_like(x_np) + + for i, xi in enumerate(x_np): + # Handle special cases + if xi <= 0 and xi == int(xi): + # log(Γ(x)) is undefined for non-positive integers + result_np[i] = float('inf') + else: + try: + result_np[i] = float(special.gammaln(xi)) + except (ValueError, OverflowError): + # Use approximation for large values + result_np[i] = float('inf') + + elif function == "beta": + # Beta function B(a, x) + result_np = np.zeros_like(x_np) + + for i, xi in enumerate(x_np): + try: + result_np[i] = float(special.beta(parameter_a, xi)) + except (ValueError, OverflowError): + # Handle cases where beta is undefined + result_np[i] = float('inf') + + elif function == "incomplete_gamma": + # Regularized incomplete gamma function P(a, x) + result_np = np.zeros_like(x_np) + + for i, xi in enumerate(x_np): + if xi < 0: + # Undefined for negative x + result_np[i] = 0.0 + else: + try: + result_np[i] = float(special.gammainc(parameter_a, xi)) + except (ValueError, OverflowError): + result_np[i] = 1.0 # Approach 1 for large x + + elif function == "incomplete_beta": + # Regularized incomplete beta function I(x; a, b) + result_np = np.zeros_like(x_np) + + for i, xi in enumerate(x_np): + # Clamp to [0,1] for domain of incomplete beta + xi_clamped = min(max(xi, 0), 1) + + try: + result_np[i] = float(special.betainc(parameter_a, parameter_b, xi_clamped)) + except (ValueError, OverflowError): + result_np[i] = 0.5 # Default for errors + + # Convert back to tensor + result = torch.tensor(result_np, device=sigmas.device, dtype=sigmas.dtype) + + # Replace infinities and NaNs + result = torch.where(torch.isfinite(result), result, torch.sign(result) * 1e10) + + # Normalize output if requested + if normalize_output: + # Handle cases where result has infinities + if torch.isinf(result).any() or torch.isnan(result).any(): + # Replace inf/nan with max/min finite values + max_val = torch.max(result[torch.isfinite(result)]) if torch.any(torch.isfinite(result)) else 1e10 + min_val = torch.min(result[torch.isfinite(result)]) if torch.any(torch.isfinite(result)) else -1e10 + + result = torch.where(torch.isinf(result) & (result > 0), max_val, result) + result = torch.where(torch.isinf(result) & (result < 0), min_val, result) + result = torch.where(torch.isnan(result), (max_val + min_val) / 2, result) + + # Now normalize + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +# ----- Sigma Lerp ----- +class sigmas_lerp: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_a": ("SIGMAS", {"forceInput": True}), + "sigmas_b": ("SIGMAS", {"forceInput": True}), + "t": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "ensure_length": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_a, sigmas_b, t, ensure_length): + if ensure_length and len(sigmas_a) != len(sigmas_b): + # Resize the smaller one to match the larger one + if len(sigmas_a) < len(sigmas_b): + sigmas_a = torch.nn.functional.interpolate( + sigmas_a.unsqueeze(0).unsqueeze(0), + size=len(sigmas_b), + mode='linear' + ).squeeze(0).squeeze(0) + else: + sigmas_b = torch.nn.functional.interpolate( + sigmas_b.unsqueeze(0).unsqueeze(0), + size=len(sigmas_a), + mode='linear' + ).squeeze(0).squeeze(0) + + return ((1 - t) * sigmas_a + t * sigmas_b,) + +# ----- Sigma InvLerp ----- +class sigmas_invlerp: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "min_value": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "max_value": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, min_value, max_value): + # Clamp values to avoid division by zero + if min_value == max_value: + max_value = min_value + 1e-5 + + normalized = (sigmas - min_value) / (max_value - min_value) + # Clamp the values to be in [0, 1] + normalized = torch.clamp(normalized, 0.0, 1.0) + return (normalized,) + +# ----- Sigma ArcSine ----- +class sigmas_arcsine: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "normalize_input": ("BOOLEAN", {"default": True}), + "scale_output": ("BOOLEAN", {"default": True}), + "out_min": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "out_max": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, normalize_input, scale_output, out_min, out_max): + if normalize_input: + sigmas = torch.clamp(sigmas, -1.0, 1.0) + else: + # Ensure values are in valid arcsin domain + sigmas = torch.clamp(sigmas, -1.0, 1.0) + + result = torch.asin(sigmas) + + if scale_output: + # ArcSine output is in range [-π/2, π/2] + # Normalize to [0, 1] and then scale to [out_min, out_max] + result = (result + math.pi/2) / math.pi + result = result * (out_max - out_min) + out_min + + return (result,) + +# ----- Sigma LinearSine ----- +class sigmas_linearsine: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "amplitude": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 10.0, "step": 0.01}), + "frequency": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "phase": ("FLOAT", {"default": 0.0, "min": -6.28, "max": 6.28, "step": 0.01}), # -2π to 2π + "linear_weight": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, amplitude, frequency, phase, linear_weight): + # Create indices for the sine function + indices = torch.linspace(0, 1, len(sigmas), device=sigmas.device) + + # Calculate sine component + sine_component = amplitude * torch.sin(2 * math.pi * frequency * indices + phase) + + # Blend linear and sine components + step_indices = torch.linspace(0, 1, len(sigmas), device=sigmas.device) + result = linear_weight * sigmas + (1 - linear_weight) * (step_indices.unsqueeze(0) * sine_component) + + return (result.squeeze(0),) + +# ----- Sigmas Append ----- +class sigmas_append: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "value": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "count": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}) + }, + "optional": { + "additional_sigmas": ("SIGMAS", {"forceInput": False}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, value, count, additional_sigmas=None): + # Create tensor of the value to append + append_values = torch.full((count,), value, device=sigmas.device, dtype=sigmas.dtype) + + # Append the values + result = torch.cat([sigmas, append_values], dim=0) + + # If additional sigmas provided, append those as well + if additional_sigmas is not None: + result = torch.cat([result, additional_sigmas], dim=0) + + return (result,) + +# ----- Sigma Arccosine ----- +class sigmas_arccosine: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "normalize_input": ("BOOLEAN", {"default": True}), + "scale_output": ("BOOLEAN", {"default": True}), + "out_min": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "out_max": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, normalize_input, scale_output, out_min, out_max): + if normalize_input: + sigmas = torch.clamp(sigmas, -1.0, 1.0) + else: + # Ensure values are in valid arccos domain + sigmas = torch.clamp(sigmas, -1.0, 1.0) + + result = torch.acos(sigmas) + + if scale_output: + # ArcCosine output is in range [0, π] + # Normalize to [0, 1] and then scale to [out_min, out_max] + result = result / math.pi + result = result * (out_max - out_min) + out_min + + return (result,) + +# ----- Sigma Arctangent ----- +class sigmas_arctangent: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "scale_output": ("BOOLEAN", {"default": True}), + "out_min": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "out_max": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, scale_output, out_min, out_max): + result = torch.atan(sigmas) + + if scale_output: + # ArcTangent output is in range [-π/2, π/2] + # Normalize to [0, 1] and then scale to [out_min, out_max] + result = (result + math.pi/2) / math.pi + result = result * (out_max - out_min) + out_min + + return (result,) + +# ----- Sigma CrossProduct ----- +class sigmas_crossproduct: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_a": ("SIGMAS", {"forceInput": True}), + "sigmas_b": ("SIGMAS", {"forceInput": True}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_a, sigmas_b): + # Ensure we have at least 3 elements in each tensor + # If not, pad with zeros or truncate + if len(sigmas_a) < 3: + sigmas_a = torch.nn.functional.pad(sigmas_a, (0, 3 - len(sigmas_a))) + if len(sigmas_b) < 3: + sigmas_b = torch.nn.functional.pad(sigmas_b, (0, 3 - len(sigmas_b))) + + # Take the first 3 elements of each tensor + a = sigmas_a[:3] + b = sigmas_b[:3] + + # Compute cross product + c = torch.zeros(3, device=sigmas_a.device, dtype=sigmas_a.dtype) + c[0] = a[1] * b[2] - a[2] * b[1] + c[1] = a[2] * b[0] - a[0] * b[2] + c[2] = a[0] * b[1] - a[1] * b[0] + + return (c,) + +# ----- Sigma DotProduct ----- +class sigmas_dotproduct: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas_a": ("SIGMAS", {"forceInput": True}), + "sigmas_b": ("SIGMAS", {"forceInput": True}), + "normalize": ("BOOLEAN", {"default": False}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas_a, sigmas_b, normalize): + # Ensure equal lengths by taking the minimum + min_length = min(len(sigmas_a), len(sigmas_b)) + a = sigmas_a[:min_length] + b = sigmas_b[:min_length] + + if normalize: + a_norm = torch.norm(a) + b_norm = torch.norm(b) + # Avoid division by zero + if a_norm > 0 and b_norm > 0: + a = a / a_norm + b = b / b_norm + + # Compute dot product + result = torch.sum(a * b) + + # Return as a single-element tensor + return (torch.tensor([result], device=sigmas_a.device, dtype=sigmas_a.dtype),) + +# ----- Sigma Fmod ----- +class sigmas_fmod: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "divisor": ("FLOAT", {"default": 1.0, "min": 0.0001, "max": 10000.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, divisor): + # Ensure divisor is not zero + if divisor == 0: + divisor = 0.0001 + + result = torch.fmod(sigmas, divisor) + return (result,) + +# ----- Sigma Frac ----- +class sigmas_frac: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas): + # Get the fractional part (x - floor(x)) + result = sigmas - torch.floor(sigmas) + return (result,) + +# ----- Sigma If ----- +class sigmas_if: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "condition_sigmas": ("SIGMAS", {"forceInput": True}), + "true_sigmas": ("SIGMAS", {"forceInput": True}), + "false_sigmas": ("SIGMAS", {"forceInput": True}), + "threshold": ("FLOAT", {"default": 0.5, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "comp_type": (["greater", "less", "equal", "not_equal"], {"default": "greater"}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, condition_sigmas, true_sigmas, false_sigmas, threshold, comp_type): + # Make sure we have values to compare + max_length = max(len(condition_sigmas), len(true_sigmas), len(false_sigmas)) + + # Extend all tensors to the maximum length using interpolation + if len(condition_sigmas) != max_length: + condition_sigmas = torch.nn.functional.interpolate( + condition_sigmas.unsqueeze(0).unsqueeze(0), + size=max_length, + mode='linear' + ).squeeze(0).squeeze(0) + + if len(true_sigmas) != max_length: + true_sigmas = torch.nn.functional.interpolate( + true_sigmas.unsqueeze(0).unsqueeze(0), + size=max_length, + mode='linear' + ).squeeze(0).squeeze(0) + + if len(false_sigmas) != max_length: + false_sigmas = torch.nn.functional.interpolate( + false_sigmas.unsqueeze(0).unsqueeze(0), + size=max_length, + mode='linear' + ).squeeze(0).squeeze(0) + + # Create mask based on comparison type + if comp_type == "greater": + mask = condition_sigmas > threshold + elif comp_type == "less": + mask = condition_sigmas < threshold + elif comp_type == "equal": + mask = torch.isclose(condition_sigmas, torch.tensor(threshold, device=condition_sigmas.device)) + elif comp_type == "not_equal": + mask = ~torch.isclose(condition_sigmas, torch.tensor(threshold, device=condition_sigmas.device)) + + # Apply the mask to select values + result = torch.where(mask, true_sigmas, false_sigmas) + + return (result,) + +# ----- Sigma Logarithm2 ----- +class sigmas_logarithm2: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "handle_negative": ("BOOLEAN", {"default": True}), + "epsilon": ("FLOAT", {"default": 1e-10, "min": 1e-15, "max": 0.1, "step": 1e-10}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, handle_negative, epsilon): + if handle_negative: + # For negative values, compute -log2(-x) and negate the result + mask_negative = sigmas < 0 + mask_positive = ~mask_negative + + # Prepare positive and negative parts + pos_part = torch.log2(torch.clamp(sigmas[mask_positive], min=epsilon)) + neg_part = -torch.log2(torch.clamp(-sigmas[mask_negative], min=epsilon)) + + # Create result tensor + result = torch.zeros_like(sigmas) + result[mask_positive] = pos_part + result[mask_negative] = neg_part + else: + # Simply compute log2, clamping values to avoid log(0) + result = torch.log2(torch.clamp(sigmas, min=epsilon)) + + return (result,) + +# ----- Sigma SmoothStep ----- +class sigmas_smoothstep: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "edge0": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "edge1": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}), + "mode": (["smoothstep", "smootherstep"], {"default": "smoothstep"}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, edge0, edge1, mode): + # Normalize the values to the range [0, 1] + t = torch.clamp((sigmas - edge0) / (edge1 - edge0), 0.0, 1.0) + + if mode == "smoothstep": + # Smooth step: 3t^2 - 2t^3 + result = t * t * (3.0 - 2.0 * t) + else: # smootherstep + # Smoother step: 6t^5 - 15t^4 + 10t^3 + result = t * t * t * (t * (t * 6.0 - 15.0) + 10.0) + + # Scale back to the original range + result = result * (edge1 - edge0) + edge0 + + return (result,) + +# ----- Sigma SquareRoot ----- +class sigmas_squareroot: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "handle_negative": ("BOOLEAN", {"default": False}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, handle_negative): + if handle_negative: + # For negative values, compute sqrt(-x) and negate the result + mask_negative = sigmas < 0 + mask_positive = ~mask_negative + + # Prepare positive and negative parts + pos_part = torch.sqrt(sigmas[mask_positive]) + neg_part = -torch.sqrt(-sigmas[mask_negative]) + + # Create result tensor + result = torch.zeros_like(sigmas) + result[mask_positive] = pos_part + result[mask_negative] = neg_part + else: + # Only compute square root for non-negative values + # Negative values will be set to 0 + result = torch.sqrt(torch.clamp(sigmas, min=0)) + + return (result,) + +# ----- Sigma TimeStep ----- +class sigmas_timestep: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "dt": ("FLOAT", {"default": 0.1, "min": 0.0001, "max": 10.0, "step": 0.01}), + "scaling": (["linear", "quadratic", "sqrt", "log"], {"default": "linear"}), + "decay": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, dt, scaling, decay): + # Create time steps + timesteps = torch.arange(len(sigmas), device=sigmas.device, dtype=sigmas.dtype) * dt + + # Apply scaling + if scaling == "quadratic": + timesteps = timesteps ** 2 + elif scaling == "sqrt": + timesteps = torch.sqrt(timesteps) + elif scaling == "log": + # Add small epsilon to avoid log(0) + timesteps = torch.log(timesteps + 1e-10) + + # Apply decay + if decay > 0: + decay_factor = torch.exp(-decay * timesteps) + timesteps = timesteps * decay_factor + + # Normalize to match the range of sigmas + timesteps = ((timesteps - timesteps.min()) / + (timesteps.max() - timesteps.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (timesteps,) + +class sigmas_gaussian_cdf: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "mu": ("FLOAT", {"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.01}), + "sigma": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, mu, sigma, normalize_output): + # Apply Gaussian CDF transformation + result = 0.5 * (1 + torch.erf((sigmas - mu) / (sigma * math.sqrt(2)))) + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +class sigmas_stepwise_multirate: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 1, "max": 1000, "step": 1}), + "rates": ("STRING", {"default": "1.0,0.5,0.25", "multiline": False}), + "boundaries": ("STRING", {"default": "0.3,0.7", "multiline": False}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 100.0, "step": 0.01}), + "pad_end": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, rates, boundaries, start_value, end_value, pad_end): + # Parse rates and boundaries + rates_list = [float(r) for r in rates.split(',')] + if len(rates_list) < 1: + rates_list = [1.0] + + boundaries_list = [float(b) for b in boundaries.split(',')] + if len(boundaries_list) != len(rates_list) - 1: + # Create equal size segments if boundaries don't match rates + boundaries_list = [i / len(rates_list) for i in range(1, len(rates_list))] + + # Convert boundaries to step indices + boundary_indices = [int(b * steps) for b in boundaries_list] + + # Create steps array + result = torch.zeros(steps) + + # Fill segments with different rates + current_idx = 0 + for i, rate in enumerate(rates_list): + next_idx = boundary_indices[i] if i < len(boundary_indices) else steps + segment_length = next_idx - current_idx + if segment_length <= 0: + continue + + segment_start = start_value if i == 0 else result[current_idx-1] + segment_end = end_value if i == len(rates_list) - 1 else start_value * (1 - boundaries_list[i]) + + # Apply rate to the segment + t = torch.linspace(0, 1, segment_length) + segment = segment_start + (segment_end - segment_start) * (t ** rate) + + result[current_idx:next_idx] = segment + current_idx = next_idx + + # Add padding zero at the end if requested + if pad_end: + result = torch.cat([result, torch.tensor([0.0])]) + + return (result,) + +class sigmas_harmonic_decay: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 1, "max": 1000, "step": 1}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 100.0, "step": 0.01}), + "harmonic_offset": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "decay_rate": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + "pad_end": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, start_value, end_value, harmonic_offset, decay_rate, pad_end): + # Create harmonic series: 1/(n+offset)^rate + n = torch.arange(1, steps + 1, dtype=torch.float32) + harmonic_values = 1.0 / torch.pow(n + harmonic_offset, decay_rate) + + # Normalize to [0, 1] + normalized = (harmonic_values - harmonic_values.min()) / (harmonic_values.max() - harmonic_values.min()) + + # Scale to [end_value, start_value] and reverse (higher values first) + result = start_value - (start_value - end_value) * normalized + result = torch.flip(result, [0]) + + # Add padding zero at the end if requested + if pad_end: + result = torch.cat([result, torch.tensor([0.0])]) + + return (result,) + +class sigmas_adaptive_noise_floor: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "min_noise_level": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 1.0, "step": 0.001}), + "adaptation_factor": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "window_size": ("INT", {"default": 3, "min": 1, "max": 10, "step": 1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, min_noise_level, adaptation_factor, window_size): + # Initialize result with original sigmas + result = sigmas.clone() + + # Apply adaptive noise floor + for i in range(window_size, len(sigmas)): + # Calculate local statistics in the window + window = sigmas[i-window_size:i] + local_mean = torch.mean(window) + local_var = torch.var(window) + + # Adapt the noise floor based on local statistics + adaptive_floor = min_noise_level + adaptation_factor * local_var / (local_mean + 1e-6) + + # Apply the floor if needed + if result[i] < adaptive_floor: + result[i] = adaptive_floor + + return (result,) + +class sigmas_collatz_iteration: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "iterations": ("INT", {"default": 3, "min": 1, "max": 20, "step": 1}), + "scaling_factor": ("FLOAT", {"default": 0.1, "min": 0.0001, "max": 10.0, "step": 0.01}), + "normalize_output": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, iterations, scaling_factor, normalize_output): + # Scale input to reasonable range for Collatz + scaled_input = sigmas * scaling_factor + + # Apply Collatz iterations + result = scaled_input.clone() + + for _ in range(iterations): + # Create masks for even and odd values + even_mask = (result % 2 == 0) + odd_mask = ~even_mask + + # Apply Collatz function: n/2 for even, 3n+1 for odd + result[even_mask] = result[even_mask] / 2 + result[odd_mask] = 3 * result[odd_mask] + 1 + + # Normalize output if requested + if normalize_output: + result = ((result - result.min()) / (result.max() - result.min())) * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +class sigmas_conway_sequence: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 20, "min": 1, "max": 50, "step": 1}), + "sequence_type": (["look_and_say", "audioactive", "paperfolding", "thue_morse"], {"default": "look_and_say"}), + "normalize_range": ("BOOLEAN", {"default": True}), + "min_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}), + "max_value": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 50.0, "step": 0.1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, sequence_type, normalize_range, min_value, max_value): + if sequence_type == "look_and_say": + # Start with "1" + s = "1" + lengths = [1] # Length of first term is 1 + + # Generate look-and-say sequence + for _ in range(min(steps - 1, 25)): # Limit to prevent excessive computation + next_s = "" + i = 0 + while i < len(s): + count = 1 + while i + 1 < len(s) and s[i] == s[i + 1]: + i += 1 + count += 1 + next_s += str(count) + s[i] + i += 1 + s = next_s + lengths.append(len(s)) + + # Convert to tensor + result = torch.tensor(lengths, dtype=torch.float32) + + elif sequence_type == "audioactive": + # Audioactive sequence (similar to look-and-say but counts digits) + a = [1] + for _ in range(min(steps - 1, 30)): + b = [] + digit_count = {} + for digit in a: + digit_count[digit] = digit_count.get(digit, 0) + 1 + + for digit in sorted(digit_count.keys()): + b.append(digit_count[digit]) + b.append(digit) + a = b + + result = torch.tensor(a, dtype=torch.float32) + if len(result) > steps: + result = result[:steps] + + elif sequence_type == "paperfolding": + # Paper folding sequence (dragon curve) + sequence = [] + for i in range(min(steps, 30)): + sequence.append(1 if (i & (i + 1)) % 2 == 0 else 0) + + result = torch.tensor(sequence, dtype=torch.float32) + + elif sequence_type == "thue_morse": + # Thue-Morse sequence + sequence = [0] + while len(sequence) < steps: + sequence.extend([1 - x for x in sequence]) + + result = torch.tensor(sequence, dtype=torch.float32)[:steps] + + # Normalize to desired range + if normalize_range: + if result.max() > result.min(): + result = (result - result.min()) / (result.max() - result.min()) + result = result * (max_value - min_value) + min_value + else: + result = torch.ones_like(result) * min_value + + return (result,) + +class sigmas_gilbreath_sequence: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 10, "max": 100, "step": 1}), + "levels": ("INT", {"default": 3, "min": 1, "max": 10, "step": 1}), + "normalize_range": ("BOOLEAN", {"default": True}), + "min_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}), + "max_value": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 50.0, "step": 0.1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, levels, normalize_range, min_value, max_value): + # Generate first few prime numbers + def sieve_of_eratosthenes(limit): + sieve = [True] * (limit + 1) + sieve[0] = sieve[1] = False + for i in range(2, int(limit**0.5) + 1): + if sieve[i]: + for j in range(i*i, limit + 1, i): + sieve[j] = False + return [i for i in range(limit + 1) if sieve[i]] + + # Get primes + primes = sieve_of_eratosthenes(steps * 6) # Get enough primes + primes = primes[:steps] + + # Generate Gilbreath sequence levels + sequences = [primes] + for level in range(1, levels): + prev_seq = sequences[level-1] + new_seq = [abs(prev_seq[i] - prev_seq[i+1]) for i in range(len(prev_seq)-1)] + sequences.append(new_seq) + + # Select the requested level + selected_level = min(levels-1, len(sequences)-1) + result_list = sequences[selected_level] + + # Ensure we have enough values + while len(result_list) < steps: + result_list.append(1) # Gilbreath conjecture: eventually all 1s + + # Convert to tensor + result = torch.tensor(result_list[:steps], dtype=torch.float32) + + # Normalize to desired range + if normalize_range: + if result.max() > result.min(): + result = (result - result.min()) / (result.max() - result.min()) + result = result * (max_value - min_value) + min_value + else: + result = torch.ones_like(result) * min_value + + return (result,) + +class sigmas_cnf_inverse: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS", {"forceInput": True}), + "time_steps": ("INT", {"default": 20, "min": 5, "max": 100, "step": 1}), + "flow_type": (["linear", "quadratic", "sigmoid", "exponential"], {"default": "sigmoid"}), + "reverse": ("BOOLEAN", {"default": True}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, sigmas, time_steps, flow_type, reverse): + # Create normalized time steps + t = torch.linspace(0, 1, time_steps) + + # Apply CNF flow transformation + if flow_type == "linear": + flow = t + elif flow_type == "quadratic": + flow = t**2 + elif flow_type == "sigmoid": + flow = 1 / (1 + torch.exp(-10 * (t - 0.5))) + elif flow_type == "exponential": + flow = torch.exp(3 * t) - 1 + flow = flow / flow.max() # Normalize to [0,1] + + # Reverse flow if requested + if reverse: + flow = 1 - flow + + # Interpolate sigmas according to flow + # First normalize sigmas to [0,1] for interpolation + normalized_sigmas = (sigmas - sigmas.min()) / (sigmas.max() - sigmas.min()) + + # Create indices for interpolation + indices = flow * (len(sigmas) - 1) + + # Linear interpolation + result = torch.zeros(time_steps, device=sigmas.device, dtype=sigmas.dtype) + for i in range(time_steps): + idx_low = int(indices[i]) + idx_high = min(idx_low + 1, len(sigmas) - 1) + frac = indices[i] - idx_low + + result[i] = (1 - frac) * normalized_sigmas[idx_low] + frac * normalized_sigmas[idx_high] + + # Scale back to original sigma range + result = result * (sigmas.max() - sigmas.min()) + sigmas.min() + + return (result,) + +class sigmas_riemannian_flow: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 5, "max": 100, "step": 1}), + "metric_type": (["euclidean", "hyperbolic", "spherical", "lorentzian"], {"default": "hyperbolic"}), + "curvature": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.1, "max": 50.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, metric_type, curvature, start_value, end_value): + # Create parameter t in [0, 1] + t = torch.linspace(0, 1, steps) + + # Apply different Riemannian metrics + if metric_type == "euclidean": + # Simple linear interpolation in Euclidean space + result = start_value * (1 - t) + end_value * t + + elif metric_type == "hyperbolic": + # Hyperbolic space geodesic + K = -curvature # Negative curvature for hyperbolic space + + # Convert to hyperbolic coordinates (using Poincaré disk model) + x_start = torch.tanh(start_value / 2) + x_end = torch.tanh(end_value / 2) + + # Distance in hyperbolic space + d = torch.acosh(1 + 2 * ((x_start - x_end)**2) / ((1 - x_start**2) * (1 - x_end**2))) + + # Geodesic interpolation + lambda_t = torch.sinh(t * d) / torch.sinh(d) + result = 2 * torch.atanh((1 - lambda_t) * x_start + lambda_t * x_end) + + elif metric_type == "spherical": + # Spherical space geodesic (great circle) + K = curvature # Positive curvature for spherical space + + # Convert to angular coordinates + theta_start = start_value * torch.sqrt(K) + theta_end = end_value * torch.sqrt(K) + + # Geodesic interpolation along great circle + result = torch.sin((1 - t) * theta_start + t * theta_end) / torch.sqrt(K) + + elif metric_type == "lorentzian": + # Lorentzian spacetime-inspired metric (time dilation effect) + gamma = 1 / torch.sqrt(1 - curvature * t**2) # Lorentz factor + result = start_value * (1 - t) + end_value * t + result = result * gamma # Apply time dilation + + # Ensure the values are in the desired range + result = torch.clamp(result, min=min(start_value, end_value), max=max(start_value, end_value)) + + # Ensure result is decreasing if start_value > end_value + if start_value > end_value and result[0] < result[-1]: + result = torch.flip(result, [0]) + + return (result,) + +class sigmas_langevin_dynamics: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 5, "max": 100, "step": 1}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.1, "max": 50.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}), + "temperature": ("FLOAT", {"default": 0.5, "min": 0.01, "max": 10.0, "step": 0.01}), + "friction": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), + "seed": ("INT", {"default": 42, "min": 0, "max": 99999, "step": 1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, start_value, end_value, temperature, friction, seed): + # Set random seed for reproducibility + torch.manual_seed(seed) + + # Potential function (quadratic well centered at end_value) + def U(x): + return 0.5 * (x - end_value)**2 + + # Gradient of the potential + def grad_U(x): + return x - end_value + + # Initialize state + x = torch.tensor([start_value], dtype=torch.float32) + v = torch.zeros(1) # Initial velocity + + # Discretization parameters + dt = 1.0 / steps + sqrt_2dt = math.sqrt(2 * dt) + + # Storage for trajectory + trajectory = [start_value] + + # Langevin dynamics integration (velocity Verlet with Langevin thermostat) + for _ in range(steps - 1): + # Half step in velocity + v = v - dt * friction * v - dt * grad_U(x) / 2 + + # Full step in position + x = x + dt * v + + # Random force (thermal noise) + noise = torch.randn(1) * sqrt_2dt * temperature + + # Another half step in velocity with noise + v = v - dt * friction * v - dt * grad_U(x) / 2 + noise + + # Store current position + trajectory.append(x.item()) + + # Convert to tensor + result = torch.tensor(trajectory, dtype=torch.float32) + + # Ensure we reach the end value + result[-1] = end_value + + return (result,) + +class sigmas_persistent_homology: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 5, "max": 100, "step": 1}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.1, "max": 50.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}), + "persistence_type": (["linear", "exponential", "logarithmic", "sigmoidal"], {"default": "exponential"}), + "birth_density": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}), + "death_density": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0, "step": 0.01}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, start_value, end_value, persistence_type, birth_density, death_density): + # Basic filtration function (linear by default) + t = torch.linspace(0, 1, steps) + + # Persistence diagram simulation + # Create birth and death times + birth_points = int(steps * birth_density) + death_points = int(steps * death_density) + + # Filtration function based on selected type + if persistence_type == "linear": + filtration = t + elif persistence_type == "exponential": + filtration = 1 - torch.exp(-5 * t) + elif persistence_type == "logarithmic": + filtration = torch.log(1 + 9 * t) / torch.log(torch.tensor([10.0])) + elif persistence_type == "sigmoidal": + filtration = 1 / (1 + torch.exp(-10 * (t - 0.5))) + + # Generate birth-death pairs + birth_indices = torch.linspace(0, steps // 2, birth_points).long() + death_indices = torch.linspace(steps // 2, steps - 1, death_points).long() + + # Create persistence barcode + barcode = torch.zeros(steps) + for b_idx in birth_indices: + for d_idx in death_indices: + if b_idx < d_idx: + # Add a persistence feature from birth to death + barcode[b_idx:d_idx] += 1 + + # Normalize and weight the barcode + if barcode.max() > 0: + barcode = barcode / barcode.max() + + # Modulate the filtration function with the persistence barcode + result = filtration * (0.7 + 0.3 * barcode) + + # Scale to desired range + result = start_value + (end_value - start_value) * result + + return (result,) + +class sigmas_normalizing_flows: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "steps": ("INT", {"default": 30, "min": 5, "max": 100, "step": 1}), + "start_value": ("FLOAT", {"default": 10.0, "min": 0.1, "max": 50.0, "step": 0.1}), + "end_value": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0, "step": 0.01}), + "flow_type": (["affine", "planar", "radial", "realnvp"], {"default": "realnvp"}), + "num_transforms": ("INT", {"default": 3, "min": 1, "max": 10, "step": 1}), + "seed": ("INT", {"default": 42, "min": 0, "max": 99999, "step": 1}) + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS",) + CATEGORY = "RES4LYF/sigmas" + + def main(self, steps, start_value, end_value, flow_type, num_transforms, seed): + # Set random seed for reproducibility + torch.manual_seed(seed) + + # Create base linear schedule from start_value to end_value + base_schedule = torch.linspace(start_value, end_value, steps) + + # Apply different normalizing flow transformations + if flow_type == "affine": + # Affine transformation: f(x) = a*x + b + result = base_schedule.clone() + for _ in range(num_transforms): + a = torch.rand(1) * 0.5 + 0.75 # Scale in [0.75, 1.25] + b = (torch.rand(1) - 0.5) * 0.2 # Shift in [-0.1, 0.1] + result = a * result + b + + elif flow_type == "planar": + # Planar flow: f(x) = x + u * tanh(w * x + b) + result = base_schedule.clone() + for _ in range(num_transforms): + u = torch.rand(1) * 0.4 - 0.2 # in [-0.2, 0.2] + w = torch.rand(1) * 2 - 1 # in [-1, 1] + b = torch.rand(1) * 0.2 - 0.1 # in [-0.1, 0.1] + result = result + u * torch.tanh(w * result + b) + + elif flow_type == "radial": + # Radial flow: f(x) = x + beta * (x - x0) / (alpha + |x - x0|) + result = base_schedule.clone() + for _ in range(num_transforms): + # Pick a random reference point within the range + idx = torch.randint(0, steps, (1,)) + x0 = result[idx] + + alpha = torch.rand(1) * 0.5 + 0.5 # in [0.5, 1.0] + beta = torch.rand(1) * 0.4 - 0.2 # in [-0.2, 0.2] + + # Apply radial flow + diff = result - x0 + r = torch.abs(diff) + result = result + beta * diff / (alpha + r) + + elif flow_type == "realnvp": + # Simplified RealNVP-inspired flow with masking + result = base_schedule.clone() + + for _ in range(num_transforms): + # Create alternating mask + mask = torch.zeros(steps) + mask[::2] = 1 # Mask even indices + + # Generate scale and shift parameters + log_scale = torch.rand(steps) * 0.2 - 0.1 # in [-0.1, 0.1] + shift = torch.rand(steps) * 0.2 - 0.1 # in [-0.1, 0.1] + + # Apply affine coupling transformation + scale = torch.exp(log_scale * mask) + masked_shift = shift * mask + + # Transform + result = result * scale + masked_shift + + # Rescale to ensure we maintain start_value and end_value + if result[0] != start_value or result[-1] != end_value: + result = (result - result[0]) / (result[-1] - result[0]) * (end_value - start_value) + start_value + + return (result,) + + +class sigmas_split_value: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sigmas": ("SIGMAS",), + "split_value": ("FLOAT", {"default": 0.875, "min": 0.0, "max": 80085.0, "step": 0.001}), + "bias_split_up": ("BOOLEAN", {"default": False, "tooltip": "If True, split happens above the split value, so high_sigmas includes the split point."}), + } + } + + FUNCTION = "main" + RETURN_TYPES = ("SIGMAS", "SIGMAS") + RETURN_NAMES = ("high_sigmas", "low_sigmas") + CATEGORY = "RES4LYF/sigmas" + DESCRIPTION = ("Splits sigma schedule at a specific sigma value.") + + def main(self, sigmas, split_value, bias_split_up): + if len(sigmas) == 0: + return (sigmas, sigmas) + + # Find the split index + if bias_split_up: + # Find first sigma <= split_value + split_idx = None + for i, sigma in enumerate(sigmas): + if sigma <= split_value: + split_idx = i + break + + if split_idx is None: + # All sigmas are above split_value + return (sigmas, torch.tensor([], device=sigmas.device, dtype=sigmas.dtype)) + + # high_sigmas: from start to split_idx (inclusive) + # low_sigmas: from split_idx to end + high_sigmas = sigmas[:split_idx + 1] + low_sigmas = sigmas[split_idx:] + + else: + # Find first sigma < split_value + split_idx = None + for i, sigma in enumerate(sigmas): + if sigma < split_value: + split_idx = i + break + + if split_idx is None: + # All sigmas are >= split_value + return (torch.tensor([], device=sigmas.device, dtype=sigmas.dtype), sigmas) + + # high_sigmas: from start to split_idx (exclusive) + # low_sigmas: from split_idx-1 to end (includes the boundary point) + high_sigmas = sigmas[:split_idx] + low_sigmas = sigmas[split_idx - 1:] + + return (high_sigmas, low_sigmas) + + + + + + + + +def get_bong_tangent_sigmas(steps, slope, pivot, start, end): + smax = ((2/pi)*atan(-slope*(0-pivot))+1)/2 + smin = ((2/pi)*atan(-slope*((steps-1)-pivot))+1)/2 + + srange = smax-smin + sscale = start - end + + sigmas = [ ( (((2/pi)*atan(-slope*(x-pivot))+1)/2) - smin) * (1/srange) * sscale + end for x in range(steps)] + + return sigmas + +def bong_tangent_scheduler(model_sampling, steps, start=1.0, middle=0.5, end=0.0, pivot_1=0.6, pivot_2=0.6, slope_1=0.2, slope_2=0.2, pad=False): + steps += 2 + + midpoint = int( (steps*pivot_1 + steps*pivot_2) / 2 ) + pivot_1 = int(steps * pivot_1) + pivot_2 = int(steps * pivot_2) + + slope_1 = slope_1 / (steps/40) + slope_2 = slope_2 / (steps/40) + + stage_2_len = steps - midpoint + stage_1_len = steps - stage_2_len + + tan_sigmas_1 = get_bong_tangent_sigmas(stage_1_len, slope_1, pivot_1, start, middle) + tan_sigmas_2 = get_bong_tangent_sigmas(stage_2_len, slope_2, pivot_2 - stage_1_len, middle, end) + + tan_sigmas_1 = tan_sigmas_1[:-1] + if pad: + tan_sigmas_2 = tan_sigmas_2+[0] + + tan_sigmas = torch.tensor(tan_sigmas_1 + tan_sigmas_2) + + return tan_sigmas + diff --git a/simple_syrup/third_party/res4lyf_runtime/style_transfer.py b/simple_syrup/third_party/res4lyf_runtime/style_transfer.py new file mode 100644 index 0000000..27af5a9 --- /dev/null +++ b/simple_syrup/third_party/res4lyf_runtime/style_transfer.py @@ -0,0 +1,2182 @@ + +import torch +import torch.nn.functional as F +import torch.nn as nn +from torch import Tensor, FloatTensor +from typing import Optional, Callable, Tuple, Dict, List, Any, Union + +import einops +from einops import rearrange +import copy +import comfy + + +from .latents import gaussian_blur_2d, median_blur_2d + +# WIP... not yet in use... +class StyleTransfer: + def __init__(self, + style_method = "WCT", + embedder_method = None, + patch_size = 1, + pinv_dtype = torch.float64, + dtype = torch.float64, + ): + self.style_method = style_method + + self.embedder_method = None + self.unembedder_method = None + + if embedder_method is not None: + self.set_embedder_method(embedder_method) + + self.patch_size = patch_size + + #if embedder_type == "conv2d": + # self.unembedder = self.invert_conv2d + self.pinv_dtype = pinv_dtype + self.dtype = dtype + + self.patchify = None + self.unpatchify = None + + self.orig_shape = None + self.grid_sizes = None + + #self.x_embed_ndim = 0 + + + + def set_patchify_method(self, patchify_method=None): + self.patchify_method = patchify_method + + def set_unpatchify_method(self, unpatchify_method=None): + self.unpatchify_method = unpatchify_method + + def set_embedder_method(self, embedder_method): + self.embedder_method = copy.deepcopy(embedder_method).to(self.pinv_dtype) + self.W = self.embedder_method.weight + self.B = self.embedder_method.bias + + if isinstance(embedder_method, nn.Linear): + self.unembedder_method = self.invert_linear + + elif isinstance(embedder_method, nn.Conv2d): + self.unembedder_method = self.invert_conv2d + + elif isinstance(embedder_method, nn.Conv3d): + self.unembedder_method = self.invert_conv3d + + def set_patch_size(self, patch_size): + self.patch_size = patch_size + + def unpatchify(self, x: Tensor) -> List[Tensor]: + x_arr = [] + for i, img_size in enumerate(self.img_sizes): # [[64,64]] , img_sizes: List[Tuple[int, int]] + pH, pW = img_size + x_arr.append( + einops.rearrange(x[i, :pH*pW].reshape(1, pH, pW, -1), 'B H W (p1 p2 C) -> B C (H p1) (W p2)', + p1=self.patch_size, p2=self.patch_size) + ) + x = torch.cat(x_arr, dim=0) + return x + + def patchify(self, x: Tensor): + x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size)) + + pH, pW = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size + self.img_sizes = [[pH, pW]] * x.shape[0] + x = einops.rearrange(x, 'B C (H p1) (W p2) -> B (H W) (p1 p2 C)', p1=self.patch_size, p2=self.patch_size) + return x + + + def embedder(self, x): + if isinstance(self.embedder_method, nn.Linear): + x = self.patchify(x) + + self.orig_shape = x.shape + x = self.embedder_method(x) + self.grid_sizes = x.shape[2:] + + #self.x_embed_ndim = x.ndim + #if x.ndim > 3: + # x = einops.rearrange(x, "B C H W -> B (H W) C") + + return x + + def unembedder(self, x): + #if self.x_embed_ndim > 3: + # x = einops.rearrange(x, "B (H W) C -> B C H W", W=self.orig_shape[-1]) + + x = self.unembedder_method(x) + return x + + + def invert_linear(self, x : torch.Tensor,) -> torch.Tensor: + x = x.to(self.pinv_dtype) + #x = (x - self.B.to(self.dtype)) @ torch.linalg.pinv(self.W.to(self.pinv_dtype)).T.to(self.dtype) + x = (x - self.B) @ torch.linalg.pinv(self.W).T + + return x.to(self.dtype) + + + + def invert_conv2d(self, z: torch.Tensor,) -> torch.Tensor: + z = z.to(self.pinv_dtype) + conv = self.embedder_method + + B, C_in, H, W = self.orig_shape + C_out, _, kH, kW = conv.weight.shape + stride_h, stride_w = conv.stride + pad_h, pad_w = conv.padding + + b = conv.bias.view(1, C_out, 1, 1).to(z) + z_nobias = z - b + + W_flat = conv.weight.view(C_out, -1).to(z) + W_pinv = torch.linalg.pinv(W_flat) + + Bz, Co, Hp, Wp = z_nobias.shape + z_flat = z_nobias.reshape(Bz, Co, -1) + + x_patches = W_pinv @ z_flat + + x_sum = F.fold( + x_patches, + output_size=(H + 2*pad_h, W + 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + ones = torch.ones_like(x_patches) + count = F.fold( + ones, + output_size=(H + 2*pad_h, W + 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + + x_recon = x_sum / count.clamp(min=1e-6) + if pad_h > 0 or pad_w > 0: + x_recon = x_recon[..., pad_h:pad_h+H, pad_w:pad_w+W] + + return x_recon.to(self.dtype) + + + + def invert_conv3d(self, z: torch.Tensor, ) -> torch.Tensor: + z = z.to(self.pinv_dtype) + conv = self.embedder_method + grid_sizes = self.grid_sizes + + B, C_in, D, H, W = self.orig_shape + pD, pH, pW = self.patch_size + sD, sH, sW = pD, pH, pW + + if z.ndim == 3: + # [B, S, C_out] -> reshape to [B, C_out, D', H', W'] + S = z.shape[1] + if grid_sizes is None: + Dp = D // pD + Hp = H // pH # getting actual patchified dims + Wp = W // pW + else: + Dp, Hp, Wp = grid_sizes + C_out = z.shape[2] + z = z.transpose(1, 2).reshape(B, C_out, Dp, Hp, Wp) + else: + B2, C_out, Dp, Hp, Wp = z.shape + assert B2 == B, "Batch size mismatch... ya sharked it." + + b = conv.bias.view(1, C_out, 1, 1, 1) # need to kncokout bias to invert via weight + z_nobias = z - b + + # 2D filter -> pinv + w3 = conv.weight # [C_out, C_in, 1, pH, pW] + w2 = w3.squeeze(2) # [C_out, C_in, pH, pW] + out_ch, in_ch, kH, kW = w2.shape + W_flat = w2.view(out_ch, -1) # [C_out, in_ch*pH*pW] + W_pinv = torch.linalg.pinv(W_flat) # [in_ch*pH*pW, C_out] + + # merge depth for 2D unfold wackiness + z2 = z_nobias.permute(0,2,1,3,4).reshape(B*Dp, C_out, Hp, Wp) + + # apply pinv ... get patch vectors + z_flat = z2.reshape(B*Dp, C_out, -1) # [B*Dp, C_out, L] + x_patches = W_pinv @ z_flat # [B*Dp, in_ch*pH*pW, L] + + # fold -> restore spatial frames + x2 = F.fold( + x_patches, + output_size=(H, W), + kernel_size=(pH, pW), + stride=(sH, sW) + ) # → [B*Dp, C_in, H, W] + + # unmerge depth (de-depth charge) + x2 = x2.reshape(B, Dp, in_ch, H, W) # [B, Dp, C_in, H, W] + x_recon = x2.permute(0,2,1,3,4).contiguous() # [B, C_in, D, H, W] + return x_recon.to(self.dtype) + + + + def adain_seq_inplace(self, content: torch.Tensor, style: torch.Tensor, eps: float = 1e-7) -> torch.Tensor: + mean_c = content.mean(1, keepdim=True) + std_c = content.std (1, keepdim=True).add_(eps) + mean_s = style.mean (1, keepdim=True) + std_s = style.std (1, keepdim=True).add_(eps) + + content.sub_(mean_c).div_(std_c).mul_(std_s).add_(mean_s) + return content + + + + + + +class StyleWCT: + def __init__(self, dtype=torch.float64, use_svd=False,): + self.dtype = dtype + self.use_svd = use_svd + self.y0_adain_embed = None + self.mu_s = None + self.y0_color = None + self.spatial_shape = None + + def whiten(self, f_s_centered: torch.Tensor, set=False): + cov = (f_s_centered.T.double() @ f_s_centered.double()) / (f_s_centered.size(0) - 1) + + if self.use_svd: + U_svd, S_svd, Vh_svd = torch.linalg.svd(cov + 1e-5 * torch.eye(cov.size(0), dtype=cov.dtype, device=cov.device)) + S_eig = S_svd + U_eig = U_svd + else: + S_eig, U_eig = torch.linalg.eigh(cov + 1e-5 * torch.eye(cov.size(0), dtype=cov.dtype, device=cov.device)) + + if set: + S_eig_root = S_eig.clamp(min=0).sqrt() # eigenvalues -> singular values + else: + S_eig_root = S_eig.clamp(min=0).rsqrt() # inverse square root + + whiten = U_eig @ torch.diag(S_eig_root) @ U_eig.T + return whiten.to(f_s_centered) + + def set(self, y0_adain_embed: torch.Tensor, spatial_shape=None): + if self.y0_adain_embed is None or self.y0_adain_embed.shape != y0_adain_embed.shape or torch.norm(self.y0_adain_embed - y0_adain_embed) > 0: + self.y0_adain_embed = y0_adain_embed.clone() + if spatial_shape is not None: + self.spatial_shape = spatial_shape + + f_s = y0_adain_embed[0] # if y0_adain_embed.ndim > 4 else y0_adain_embed + self.mu_s = f_s.mean(dim=0, keepdim=True) + f_s_centered = f_s - self.mu_s + + self.y0_color = self.whiten(f_s_centered, set=True) + + def get(self, denoised_embed: torch.Tensor): + for wct_i in range(denoised_embed.shape[0]): + f_c = denoised_embed[wct_i] + mu_c = f_c.mean(dim=0, keepdim=True) + f_c_centered = f_c - mu_c + + whiten = self.whiten(f_c_centered) + + f_c_whitened = f_c_centered @ whiten.T + f_cs = f_c_whitened @ self.y0_color.T + self.mu_s + + denoised_embed[wct_i] = f_cs + + return denoised_embed + + + + +class WaveletStyleWCT(StyleWCT): + def set(self, y0_adain_embed: torch.Tensor, h_len, w_len): + if self.y0_adain_embed is None or self.y0_adain_embed.shape != y0_adain_embed.shape or torch.norm(self.y0_adain_embed - y0_adain_embed) > 0: + self.y0_adain_embed = y0_adain_embed.clone() + + B, HW, C = y0_adain_embed.shape + LL, _, _, _ = haar_wavelet_decompose(y0_adain_embed.contiguous().view(B, C, h_len, w_len)) + + B_LL, C_LL, H_LL, W_LL = LL.shape + #flat = rearrange(LL, 'b c h w -> b (h w) c') + flat = LL.contiguous().view(B_LL, H_LL * W_LL, C_LL) + + f_s = flat[0] # assuming batch size 1 or using only the first + self.mu_s = f_s.mean(dim=0, keepdim=True) + f_s_centered = f_s - self.mu_s + self.y0_color = self.whiten(f_s_centered, set=True) + #self.y0_adain_embed = flat # cache if needed + + def get(self, denoised_embed: torch.Tensor, h_len, w_len, stylize_highfreq=False): + + B, HW, C = denoised_embed.shape + + denoised_embed = denoised_embed.contiguous().view(B, C, h_len, w_len) + + for i in range(B): + x = denoised_embed[i:i+1] # [1, C, H, W] + LL, LH, HL, HH = haar_wavelet_decompose(x) + + def process_band(band): + Bc, Cc, Hc, Wc = band.shape + flat = band.contiguous().view(Bc, Hc * Wc, Cc) + + styled = super(WaveletStyleWCT, self).get(flat) + return styled.contiguous().view(Bc, Cc, Hc, Wc) + + LL_styled = process_band(LL) + + if stylize_highfreq: + LH_styled = process_band(LH) + HL_styled = process_band(HL) + HH_styled = process_band(HH) + else: + LH_styled, HL_styled, HH_styled = LH, HL, HH + + recon = haar_wavelet_reconstruct(LL_styled, LH_styled, HL_styled, HH_styled) + denoised_embed[i] = recon.squeeze(0) + + return denoised_embed.view(B, HW, C) + + + +def haar_wavelet_decompose(x): + """ + Orthonormal Haar decomposition. + Input: [B, C, H, W] + Output: LL, LH, HL, HH with shape [B, C, H//2, W//2] + """ + if x.dtype != torch.float32: + x = x.float() + + B, C, H, W = x.shape + assert H % 2 == 0 and W % 2 == 0, "Input must have even H, W" + + # Precompute + norm = 1 / 2**0.5 + + x00 = x[:, :, 0::2, 0::2] + x01 = x[:, :, 0::2, 1::2] + x10 = x[:, :, 1::2, 0::2] + x11 = x[:, :, 1::2, 1::2] + + LL = (x00 + x01 + x10 + x11) * norm * 0.5 + LH = (x00 - x01 + x10 - x11) * norm * 0.5 + HL = (x00 + x01 - x10 - x11) * norm * 0.5 + HH = (x00 - x01 - x10 + x11) * norm * 0.5 + + return LL, LH, HL, HH + +def haar_wavelet_reconstruct(LL, LH, HL, HH): + """ + Orthonormal inverse Haar reconstruction. + Input: LL, LH, HL, HH [B, C, H, W] + Output: Reconstructed [B, C, H*2, W*2] + """ + norm = 1 / 2**0.5 + B, C, H, W = LL.shape + + x00 = (LL + LH + HL + HH) * norm + x01 = (LL - LH + HL - HH) * norm + x10 = (LL + LH - HL - HH) * norm + x11 = (LL - LH - HL + HH) * norm + + out = torch.zeros(B, C, H * 2, W * 2, device=LL.device, dtype=LL.dtype) + out[:, :, 0::2, 0::2] = x00 + out[:, :, 0::2, 1::2] = x01 + out[:, :, 1::2, 0::2] = x10 + out[:, :, 1::2, 1::2] = x11 + + return out + + + + + + + + +""" + +class StyleFeatures: + def __init__(self, dtype=torch.float64,): + self.dtype = dtype + + def set(self, y0_adain_embed: torch.Tensor): + + def get(self, denoised_embed: torch.Tensor): + + return "Norpity McNerp" + +""" + + + + +class Retrojector: + def __init__(self, proj=None, patch_size=2, pinv_dtype=torch.float64, dtype=torch.float64, ENDO=False): + self.proj = proj + self.patch_size = patch_size + self.pinv_dtype = pinv_dtype + self.dtype = dtype + + self.LINEAR = isinstance(proj, nn.Linear) + self.CONV2D = isinstance(proj, nn.Conv2d) + self.CONV3D = isinstance(proj, nn.Conv3d) + self.ENDO = ENDO + self.W = proj.weight.data.to(dtype=pinv_dtype).cuda() + + if self.LINEAR: + self.W_inv = torch.linalg.pinv(self.W.cuda()) + elif self.CONV2D: + C_out, _, kH, kW = proj.weight.shape + W_flat = proj.weight.view(C_out, -1).to(dtype=pinv_dtype) + self.W_inv = torch.linalg.pinv(W_flat.cuda()) + + if proj.bias is None: + if self.LINEAR: + bias_size = proj.out_features + else: + bias_size = proj.out_channels + self.b = torch.zeros(bias_size, dtype=pinv_dtype, device=self.W_inv.device) + else: + self.b = proj.bias.data.to(dtype=pinv_dtype).to(self.W_inv.device) + + def embed(self, img: torch.Tensor): + self.h = img.shape[-2] // self.patch_size + self.w = img.shape[-1] // self.patch_size + + img = comfy.ldm.common_dit.pad_to_patch_size(img, (self.patch_size, self.patch_size)) + + if self.CONV2D: + self.orig_shape = img.shape # for unembed + img_embed = F.conv2d( + img.to(self.W), + weight=self.W, + bias=self.b, + stride=self.proj.stride, + padding=self.proj.padding + ) + #img_embed = rearrange(img_embed, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=self.patch_size, pw=self.patch_size) + img_embed = rearrange(img_embed, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=1, pw=1) + + elif self.LINEAR: + if img.ndim == 4: + img = rearrange(img, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=self.patch_size, pw=self.patch_size) + if self.ENDO: + img_embed = F.linear(img.to(self.b) - self.b, self.W_inv) + else: + img_embed = F.linear(img.to(self.W), self.W, self.b) + + return img_embed.to(img) + + def unembed(self, img_embed: torch.Tensor): + if self.CONV2D: + #img_embed = rearrange(img_embed, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=self.h, w=self.w, ph=self.patch_size, pw=self.patch_size) + img_embed = rearrange(img_embed, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=self.h, w=self.w, ph=1, pw=1) + img = self.invert_conv2d(img_embed) + + elif self.LINEAR: + if self.ENDO: + img = F.linear(img_embed.to(self.W), self.W, self.b) + else: + img = F.linear(img_embed.to(self.b) - self.b, self.W_inv) + if img.ndim == 3: + img = rearrange(img, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=self.h, w=self.w, ph=self.patch_size, pw=self.patch_size) + + return img.to(img_embed) + + def invert_conv2d(self, z: torch.Tensor,) -> torch.Tensor: + z_dtype = z.dtype + z = z.to(self.pinv_dtype) + conv = self.proj + + B, C_in, H, W = self.orig_shape + C_out, _, kH, kW = conv.weight.shape + stride_h, stride_w = conv.stride + pad_h, pad_w = conv.padding + + b = conv.bias.view(1, C_out, 1, 1).to(z) + z_nobias = z - b + + #W_flat = conv.weight.view(C_out, -1).to(z) + #W_pinv = torch.linalg.pinv(W_flat) + + Bz, Co, Hp, Wp = z_nobias.shape + z_flat = z_nobias.reshape(Bz, Co, -1) + + x_patches = self.W_inv @ z_flat + + x_sum = F.fold( + x_patches, + output_size=(H + 2*pad_h, W+ 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + ones = torch.ones_like(x_patches) + count = F.fold( + ones, + output_size=(H + 2*pad_h, W + 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + + x_recon = x_sum / count.clamp(min=1e-6) + if pad_h > 0 or pad_w > 0: + x_recon = x_recon[..., pad_h:pad_h+H, pad_w:pad_w+W] + + return x_recon.to(z_dtype) + + def invert_patch_embedding(self, z: torch.Tensor, original_shape: torch.Size, grid_sizes: Optional[Tuple[int,int,int]] = None) -> torch.Tensor: + + B, C_in, D, H, W = original_shape + pD, pH, pW = self.patch_size + sD, sH, sW = pD, pH, pW + + if z.ndim == 3: + # [B, S, C_out] -> reshape to [B, C_out, D', H', W'] + S = z.shape[1] + if grid_sizes is None: + Dp = D // pD + Hp = H // pH + Wp = W // pW + else: + Dp, Hp, Wp = grid_sizes + C_out = z.shape[2] + z = z.transpose(1, 2).reshape(B, C_out, Dp, Hp, Wp) + else: + B2, C_out, Dp, Hp, Wp = z.shape + assert B2 == B, "Batch size mismatch... ya sharked it." + + # kncokout bias + b = self.patch_embedding.bias.view(1, C_out, 1, 1, 1) + z_nobias = z - b + + # 2D filter -> pinv + w3 = self.patch_embedding.weight # [C_out, C_in, 1, pH, pW] + w2 = w3.squeeze(2) # [C_out, C_in, pH, pW] + out_ch, in_ch, kH, kW = w2.shape + W_flat = w2.view(out_ch, -1) # [C_out, in_ch*pH*pW] + W_pinv = torch.linalg.pinv(W_flat) # [in_ch*pH*pW, C_out] + + # merge depth for 2D unfold wackiness + z2 = z_nobias.permute(0,2,1,3,4).reshape(B*Dp, C_out, Hp, Wp) + + # apply pinv ... get patch vectors + z_flat = z2.reshape(B*Dp, C_out, -1) # [B*Dp, C_out, L] + x_patches = W_pinv @ z_flat # [B*Dp, in_ch*pH*pW, L] + + # fold -> spatial frames + x2 = F.fold( + x_patches, + output_size=(H, W), + kernel_size=(pH, pW), + stride=(sH, sW) + ) # → [B*Dp, C_in, H, W] + + # un-merge depth + x2 = x2.reshape(B, Dp, in_ch, H, W) # [B, Dp, C_in, H, W] + x_recon = x2.permute(0,2,1,3,4).contiguous() # [B, C_in, D, H, W] + return x_recon + + + + + + +def invert_conv2d( + conv: torch.nn.Conv2d, + z: torch.Tensor, + original_shape: torch.Size, +) -> torch.Tensor: + import torch.nn.functional as F + + B, C_in, H, W = original_shape + C_out, _, kH, kW = conv.weight.shape + stride_h, stride_w = conv.stride + pad_h, pad_w = conv.padding + + if conv.bias is not None: + b = conv.bias.view(1, C_out, 1, 1).to(z) + z_nobias = z - b + else: + z_nobias = z + + W_flat = conv.weight.view(C_out, -1).to(z) + W_pinv = torch.linalg.pinv(W_flat) + + Bz, Co, Hp, Wp = z_nobias.shape + z_flat = z_nobias.reshape(Bz, Co, -1) + + x_patches = W_pinv @ z_flat + + x_sum = F.fold( + x_patches, + output_size=(H + 2*pad_h, W + 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + ones = torch.ones_like(x_patches) + count = F.fold( + ones, + output_size=(H + 2*pad_h, W + 2*pad_w), + kernel_size=(kH, kW), + stride=(stride_h, stride_w), + ) + + x_recon = x_sum / count.clamp(min=1e-6) + if pad_h > 0 or pad_w > 0: + x_recon = x_recon[..., pad_h:pad_h+H, pad_w:pad_w+W] + + return x_recon + + + +def adain_seq_inplace(content: torch.Tensor, style: torch.Tensor, dim=1, eps: float = 1e-7) -> torch.Tensor: + mean_c = content.mean(dim, keepdim=True) + std_c = content.std (dim, keepdim=True).add_(eps) # in-place add + mean_s = style.mean (dim, keepdim=True) + std_s = style.std (dim, keepdim=True).add_(eps) + + content.sub_(mean_c).div_(std_c).mul_(std_s).add_(mean_s) # in-place chain + return content + +def adain_seq(content: torch.Tensor, style: torch.Tensor, eps: float = 1e-7) -> torch.Tensor: + return ((content - content.mean(1, keepdim=True)) / (content.std(1, keepdim=True) + eps)) * (style.std(1, keepdim=True) + eps) + style.mean(1, keepdim=True) + + + + + + + + + +def apply_scattersort_tiled( + denoised_spatial : torch.Tensor, + y0_adain_spatial : torch.Tensor, + tile_h : int, + tile_w : int, + pad : int, +): + """ + Apply spatial scattersort between denoised_spatial and y0_adain_spatial + using local tile-wise sorted value matching. + + Args: + denoised_spatial (Tensor): (B, C, H, W) tensor. + y0_adain_spatial (Tensor): (B, C, H, W) reference tensor. + tile_h (int): tile height. + tile_w (int): tile width. + pad (int): padding size to apply around tiles. + + Returns: + denoised_embed (Tensor): (B, H*W, C) tensor after sortmatch. + """ + denoised_padded = F.pad(denoised_spatial, (pad, pad, pad, pad), mode='reflect') + y0_padded = F.pad(y0_adain_spatial, (pad, pad, pad, pad), mode='reflect') + + denoised_padded_out = denoised_padded.clone() + _, _, h_len, w_len = denoised_spatial.shape + + for ix in range(pad, h_len, tile_h): + for jx in range(pad, w_len, tile_w): + tile = denoised_padded[:, :, ix - pad:ix + tile_h + pad, jx - pad:jx + tile_w + pad] + y0_tile = y0_padded[:, :, ix - pad:ix + tile_h + pad, jx - pad:jx + tile_w + pad] + + tile = rearrange(tile, "b c h w -> b c (h w)", h=tile_h + pad * 2, w=tile_w + pad * 2) + y0_tile = rearrange(y0_tile, "b c h w -> b c (h w)", h=tile_h + pad * 2, w=tile_w + pad * 2) + + src_sorted, src_idx = tile.sort(dim=-1) + ref_sorted, ref_idx = y0_tile.sort(dim=-1) + + new_tile = tile.scatter(dim=-1, index=src_idx, src=ref_sorted.expand(src_sorted.shape)) + new_tile = rearrange(new_tile, "b c (h w) -> b c h w", h=tile_h + pad * 2, w=tile_w + pad * 2) + + denoised_padded_out[:, :, ix:ix + tile_h, jx:jx + tile_w] = ( + new_tile if pad == 0 else new_tile[:, :, pad:-pad, pad:-pad] + ) + + denoised_padded_out = denoised_padded_out if pad == 0 else denoised_padded_out[:, :, pad:-pad, pad:-pad] + return denoised_padded_out + + + +def apply_scattersort_masked( + denoised_embed : torch.Tensor, + y0_adain_embed : torch.Tensor, + y0_style_pos_mask : torch.Tensor | None, + y0_style_pos_mask_edge : torch.Tensor | None, + h_len : int, + w_len : int +): + if y0_style_pos_mask is None: + flatmask = torch.ones((1,1,h_len,w_len)).bool().flatten().bool() + else: + flatmask = F.interpolate(y0_style_pos_mask, size=(h_len, w_len)).bool().flatten().cpu() + flatunmask = ~flatmask + + if y0_style_pos_mask_edge is not None: + edgemask = F.interpolate( + y0_style_pos_mask_edge.unsqueeze(0), size=(h_len, w_len) + ).bool().flatten() + flatmask = flatmask & (~edgemask) + flatunmask = flatunmask & (~edgemask) + + denoised_masked = denoised_embed[:, flatmask, :].clone() + y0_adain_masked = y0_adain_embed[:, flatmask, :].clone() + + src_sorted, src_idx = denoised_masked.sort(dim=-2) + ref_sorted, ref_idx = y0_adain_masked.sort(dim=-2) + + denoised_embed[:, flatmask, :] = src_sorted.scatter(dim=-2, index=src_idx, src=ref_sorted.expand(src_sorted.shape)) + + if (flatunmask == True).any(): + denoised_unmasked = denoised_embed[:, flatunmask, :].clone() + y0_adain_unmasked = y0_adain_embed[:, flatunmask, :].clone() + + src_sorted, src_idx = denoised_unmasked.sort(dim=-2) + ref_sorted, ref_idx = y0_adain_unmasked.sort(dim=-2) + + denoised_embed[:, flatunmask, :] = src_sorted.scatter(dim=-2, index=src_idx, src=ref_sorted.expand(src_sorted.shape)) + + if y0_style_pos_mask_edge is not None: + denoised_edgemasked = denoised_embed[:, edgemask, :].clone() + y0_adain_edgemasked = y0_adain_embed[:, edgemask, :].clone() + + src_sorted, src_idx = denoised_edgemasked.sort(dim=-2) + ref_sorted, ref_idx = y0_adain_edgemasked.sort(dim=-2) + + denoised_embed[:, edgemask, :] = src_sorted.scatter(dim=-2, index=src_idx, src=ref_sorted.expand(src_sorted.shape)) + + return denoised_embed + + + + +def apply_scattersort( + denoised_embed : torch.Tensor, + y0_adain_embed : torch.Tensor, +): + #src_sorted, src_idx = denoised_embed.cpu().sort(dim=-2) + src_idx = denoised_embed.argsort(dim=-2) + ref_sorted = y0_adain_embed.sort(dim=-2)[0] + + denoised_embed.scatter_(dim=-2, index=src_idx, src=ref_sorted.expand(ref_sorted.shape)) + + return denoised_embed + +def apply_scattersort_spatial( + denoised_spatial : torch.Tensor, + y0_adain_spatial : torch.Tensor, +): + denoised_embed = rearrange(denoised_spatial, "b c h w -> b (h w) c") + y0_adain_embed = rearrange(y0_adain_spatial, "b c h w -> b (h w) c") + src_sorted, src_idx = denoised_embed.sort(dim=-2) + ref_sorted, ref_idx = y0_adain_embed.sort(dim=-2) + + denoised_embed = src_sorted.scatter(dim=-2, index=src_idx, src=ref_sorted.expand(src_sorted.shape)) + + return rearrange(denoised_embed, "b (h w) c -> b c h w", h=denoised_spatial.shape[-2], w=denoised_spatial.shape[-1]) + + + + + +def apply_scattersort_spatial( + x_spatial : torch.Tensor, + y_spatial : torch.Tensor, +): + x_emb = rearrange(x_spatial, "b c h w -> b (h w) c") + y_emb = rearrange(y_spatial, "b c h w -> b (h w) c") + + x_sorted, x_idx = x_emb.sort(dim=-2) + y_sorted, y_idx = y_emb.sort(dim=-2) + + x_emb = x_sorted.scatter(dim=-2, index=x_idx, src=y_sorted.expand(x_sorted.shape)) + + return rearrange(x_emb, "b (h w) c -> b c h w", h=x_spatial.shape[-2], w=x_spatial.shape[-1]) + + + + +def apply_adain_spatial( + x_spatial : torch.Tensor, + y_spatial : torch.Tensor, +): + x_emb = rearrange(x_spatial, "b c h w -> b (h w) c") + y_emb = rearrange(y_spatial, "b c h w -> b (h w) c") + + x_mean = x_emb.mean(-2, keepdim=True) + x_std = x_emb.std (-2, keepdim=True) + y_mean = y_emb.mean(-2, keepdim=True) + y_std = y_emb.std (-2, keepdim=True) + + assert (x_std == 0).any() == 0, "Target tensor has no variance!" + assert (y_std == 0).any() == 0, "Reference tensor has no variance!" + + x_emb_adain = (x_emb - x_mean) / x_std + x_emb_adain = (x_emb_adain * y_std) + y_mean + + return x_emb_adain.reshape_as(x_spatial) + + + + + + + + + + + + + + + + + + + + + +def adain_patchwise(content: torch.Tensor, style: torch.Tensor, sigma: float = 1.0, kernel_size: int = None, eps: float = 1e-5) -> torch.Tensor: + # this one is really slow + B, C, H, W = content.shape + device = content.device + dtype = content.dtype + + if kernel_size is None: + kernel_size = int(2 * math.ceil(3 * sigma) + 1) + if kernel_size % 2 == 0: + kernel_size += 1 + + pad = kernel_size // 2 + coords = torch.arange(kernel_size, dtype=torch.float64, device=device) - pad + gauss = torch.exp(-0.5 * (coords / sigma) ** 2) + gauss /= gauss.sum() + kernel_2d = (gauss[:, None] * gauss[None, :]).to(dtype=dtype) + + weight = kernel_2d.view(1, 1, kernel_size, kernel_size) + + content_padded = F.pad(content, (pad, pad, pad, pad), mode='reflect') + style_padded = F.pad(style, (pad, pad, pad, pad), mode='reflect') + result = torch.zeros_like(content) + + for i in range(H): + for j in range(W): + c_patch = content_padded[:, :, i:i + kernel_size, j:j + kernel_size] + s_patch = style_padded[:, :, i:i + kernel_size, j:j + kernel_size] + w = weight.expand_as(c_patch) + + c_mean = (c_patch * w).sum(dim=(-1, -2), keepdim=True) + c_std = ((c_patch - c_mean)**2 * w).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + s_mean = (s_patch * w).sum(dim=(-1, -2), keepdim=True) + s_std = ((s_patch - s_mean)**2 * w).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + + normed = (c_patch[:, :, pad:pad+1, pad:pad+1] - c_mean) / c_std + stylized = normed * s_std + s_mean + result[:, :, i, j] = stylized.squeeze(-1).squeeze(-1) + + return result + + +def adain_patchwise_row_batch(content: torch.Tensor, style: torch.Tensor, sigma: float = 1.0, kernel_size: int = None, eps: float = 1e-5) -> torch.Tensor: + + B, C, H, W = content.shape + device, dtype = content.device, content.dtype + + if kernel_size is None: + kernel_size = int(2 * math.ceil(3 * sigma) + 1) + if kernel_size % 2 == 0: + kernel_size += 1 + + pad = kernel_size // 2 + coords = torch.arange(kernel_size, dtype=torch.float64, device=device) - pad + gauss = torch.exp(-0.5 * (coords / sigma) ** 2) + gauss = (gauss / gauss.sum()).to(dtype) + kernel_2d = (gauss[:, None] * gauss[None, :]) + + weight = kernel_2d.view(1, 1, kernel_size, kernel_size) + + content_padded = F.pad(content, (pad, pad, pad, pad), mode='reflect') + style_padded = F.pad(style, (pad, pad, pad, pad), mode='reflect') + result = torch.zeros_like(content) + + for i in range(H): + c_row_patches = torch.stack([ + content_padded[:, :, i:i+kernel_size, j:j+kernel_size] + for j in range(W) + ], dim=0) # [W, B, C, k, k] + + s_row_patches = torch.stack([ + style_padded[:, :, i:i+kernel_size, j:j+kernel_size] + for j in range(W) + ], dim=0) + + w = weight.expand_as(c_row_patches[0]) + + c_mean = (c_row_patches * w).sum(dim=(-1, -2), keepdim=True) + c_std = ((c_row_patches - c_mean) ** 2 * w).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + s_mean = (s_row_patches * w).sum(dim=(-1, -2), keepdim=True) + s_std = ((s_row_patches - s_mean) ** 2 * w).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + + center = kernel_size // 2 + central = c_row_patches[:, :, :, center:center+1, center:center+1] + normed = (central - c_mean) / c_std + stylized = normed * s_std + s_mean + + result[:, :, i, :] = stylized.squeeze(-1).squeeze(-1).permute(1, 2, 0) # [B,C,W] + + return result + + + +def adain_patchwise_row_batch_med(content: torch.Tensor, style: torch.Tensor, sigma: float = 1.0, kernel_size: int = None, eps: float = 1e-5, mask: torch.Tensor = None, use_median_blur: bool = False, lowpass_weight=1.0, highpass_weight=1.0) -> torch.Tensor: + B, C, H, W = content.shape + device, dtype = content.device, content.dtype + + if kernel_size is None: + kernel_size = int(2 * math.ceil(3 * abs(sigma)) + 1) + if kernel_size % 2 == 0: + kernel_size += 1 + + pad = kernel_size // 2 + + content_padded = F.pad(content, (pad, pad, pad, pad), mode='reflect') + style_padded = F.pad(style, (pad, pad, pad, pad), mode='reflect') + result = torch.zeros_like(content) + + scaling = torch.ones((B, 1, H, W), device=device, dtype=dtype) + sigma_scale = torch.ones((H, W), device=device, dtype=torch.float32) + if mask is not None: + with torch.no_grad(): + padded_mask = F.pad(mask.float(), (pad, pad, pad, pad), mode="reflect") + blurred_mask = F.avg_pool2d(padded_mask, kernel_size=kernel_size, stride=1, padding=pad) + blurred_mask = blurred_mask[..., pad:-pad, pad:-pad] + edge_proximity = blurred_mask * (1.0 - blurred_mask) + scaling = 1.0 - (edge_proximity / 0.25).clamp(0.0, 1.0) + sigma_scale = scaling[0, 0] # assuming single-channel mask broadcasted across B, C + + if not use_median_blur: + coords = torch.arange(kernel_size, dtype=torch.float64, device=device) - pad + base_gauss = torch.exp(-0.5 * (coords / sigma) ** 2) + base_gauss = (base_gauss / base_gauss.sum()).to(dtype) + gaussian_table = {} + for s in sigma_scale.unique(): + sig = float((sigma * s + eps).clamp(min=1e-3)) + gauss_local = torch.exp(-0.5 * (coords / sig) ** 2) + gauss_local = (gauss_local / gauss_local.sum()).to(dtype) + kernel_2d = gauss_local[:, None] * gauss_local[None, :] + gaussian_table[s.item()] = kernel_2d + + for i in range(H): + row_result = torch.zeros(B, C, W, dtype=dtype, device=device) + for j in range(W): + c_patch = content_padded[:, :, i:i+kernel_size, j:j+kernel_size] + s_patch = style_padded[:, :, i:i+kernel_size, j:j+kernel_size] + + if use_median_blur: + # Median blur with residual restoration + unfolded_c = c_patch.reshape(B, C, -1) + unfolded_s = s_patch.reshape(B, C, -1) + + c_median = unfolded_c.median(dim=-1, keepdim=True).values + s_median = unfolded_s.median(dim=-1, keepdim=True).values + + center = kernel_size // 2 + central = c_patch[:, :, center, center].view(B, C, 1) + residual = central - c_median + stylized = lowpass_weight * s_median + residual * highpass_weight + else: + k = gaussian_table[float(sigma_scale[i, j].item())] + local_weight = k.view(1, 1, kernel_size, kernel_size).expand(B, C, kernel_size, kernel_size) + + c_mean = (c_patch * local_weight).sum(dim=(-1, -2), keepdim=True) + c_std = ((c_patch - c_mean) ** 2 * local_weight).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + s_mean = (s_patch * local_weight).sum(dim=(-1, -2), keepdim=True) + s_std = ((s_patch - s_mean) ** 2 * local_weight).sum(dim=(-1, -2), keepdim=True).sqrt() + eps + + center = kernel_size // 2 + central = c_patch[:, :, center:center+1, center:center+1] + normed = (central - c_mean) / c_std + stylized = normed * s_std + s_mean + + local_scaling = scaling[:, :, i, j].view(B, 1, 1) + stylized = central * (1 - local_scaling) + stylized * local_scaling + + row_result[:, :, j] = stylized.squeeze(-1) + result[:, :, i, :] = row_result + + return result + + + + + + + +def weighted_mix_n(tensor_list, weight_list, dim=-1, offset=0): + assert all(t.shape == tensor_list[0].shape for t in tensor_list) + assert len(tensor_list) == len(weight_list) + + total_weight = sum(weight_list) + ratios = [w / total_weight for w in weight_list] + + length = tensor_list[0].shape[dim] + idx = torch.arange(length) + + # Create a bin index tensor based on weighted slots + float_bins = (idx + offset) * len(ratios) / length + bin_idx = torch.floor(float_bins).long() % len(ratios) + + # Allocate slots based on ratio using a cyclic pattern + counters = [0.0 for _ in ratios] + slots = torch.empty_like(idx) + + for i in range(length): + # Assign to the group that's most under-allocated + expected = [r * (i + 1) for r in ratios] + errors = [expected[j] - counters[j] for j in range(len(ratios))] + k = max(range(len(errors)), key=lambda j: errors[j]) + slots[i] = k + counters[k] += 1 + + # Create mask for each tensor + out = tensor_list[0].clone() + for i, tensor in enumerate(tensor_list): + mask = slots == i + while mask.dim() < tensor.dim(): + mask = mask.unsqueeze(0) + mask = mask.expand_as(tensor) + out = torch.where(mask, tensor, out) + + return out + + + + + + +from torch import vmap + +BLOCK_NAMES = {"double_blocks", "single_blocks", "up_blocks", "middle_blocks", "down_blocks", "input_blocks", "output_blocks"} + +DEFAULT_BLOCK_WEIGHTS_MMDIT = { + "attn_norm" : 0.0, + "attn_norm_mod": 0.0, + "attn" : 1.0, + "attn_gated" : 0.0, + "attn_res" : 1.0, + "ff_norm" : 0.0, + "ff_norm_mod" : 0.0, + "ff" : 1.0, + "ff_gated" : 0.0, + "ff_res" : 1.0, + + "h_tile" : 8, + "w_tile" : 8, +} + +DEFAULT_ATTN_WEIGHTS_MMDIT = { + "q_proj": 0.0, + "k_proj": 0.0, + "v_proj": 1.0, + "q_norm": 0.0, + "k_norm": 0.0, + "out" : 1.0, + + "h_tile": 8, + "w_tile": 8, +} + +DEFAULT_BASE_WEIGHTS_MMDIT = { + "proj_in" : 1.0, + "proj_out": 1.0, + + "h_tile" : 8, + "w_tile" : 8, +} + +class Stylizer: + buffer = {} + + CLS_WCT = StyleWCT() + + CLS_WCT2 = WaveletStyleWCT() + + def __init__(self, dtype=torch.float64, device=torch.device("cuda")): + self.dtype = dtype + self.device = device + self.mask = [None] + self.apply_to = [""] + self.method = ["passthrough"] + self.h_tile = [-1] + self.w_tile = [-1] + + self.w_len = 0 + self.h_len = 0 + self.img_len = 0 + + self.IMG_1ST = True + self.HEADS = 0 + self.KONTEXT = 0 + def set_mode(self, mode): + self.method = [mode] #[getattr(self, mode)] + + def set_weights(self, **kwargs): + for k, v in kwargs.items(): + if hasattr(self, k): + setattr(self, k, [v]) + + def set_weights_recursive(self, **kwargs): + for name, val in kwargs.items(): + if hasattr(self, name): + setattr(self, name, [val]) + + for attr_name, attr_val in vars(self).items(): + if isinstance(attr_val, Stylizer): + attr_val.set_weights_recursive(**kwargs) + + for list_name in BLOCK_NAMES: + lst = getattr(self, list_name, None) + if isinstance(lst, list): + for element in lst: + if isinstance(element, Stylizer): + element.set_weights_recursive(**kwargs) + + def merge_weights(self, other): + def recursive_merge(a, b, path): + if isinstance(a, list) and isinstance(b, list): + if path in BLOCK_NAMES: + out = [] + for i in range(max(len(a), len(b))): + if i < len(a) and i < len(b): + out.append(recursive_merge(a[i], b[i], path=None)) + elif i < len(a): + out.append(a[i]) + else: + out.append(b[i]) + return out + return a + b + + if isinstance(a, dict) and isinstance(b, dict): + merged = dict(a) + for k, v_b in b.items(): + if k in merged: + merged[k] = recursive_merge(merged[k], v_b, path=None) + else: + merged[k] = v_b + return merged + + if hasattr(a, "__dict__") and hasattr(b, "__dict__"): + for attr, val_b in vars(b).items(): + val_a = getattr(a, attr, None) + if val_a is not None: + setattr(a, attr, recursive_merge(val_a, val_b, path=attr)) + else: + setattr(a, attr, val_b) + return a + return b + + for attr in vars(self): + if attr in BLOCK_NAMES: + merged = recursive_merge(getattr(self, attr), getattr(other, attr, []), path=attr) + elif hasattr(other, attr): + merged = recursive_merge(getattr(self, attr), getattr(other, attr), path=attr) + else: + continue + setattr(self, attr, merged) + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + self.h_len = h_len + self.w_len = w_len + self.img_slice = img_slice + self.txt_slice = txt_slice + self.img_len = h_len * w_len + self.HEADS = HEADS + + @staticmethod + def middle_slice(length, weight): + """ + Returns a slice object that selects the middle `weight` fraction of a dimension. + Example: weight=1.0 → full slice; weight=0.5 → middle 50% + """ + if weight >= 1.0: + return slice(None) + wr = int((length * (1 - weight)) // 2) + return slice(wr, -wr if wr > 0 else None) + + @staticmethod + def get_outer_slice(x, weight): + if weight >= 0.0: + return x + length = x.shape[-2] + wr = int((length * (1 - (-weight))) // 2) + + return torch.cat([x[...,:wr,:], x[...,-wr:,:]], dim=-2) + + @staticmethod + def restore_outer_slice(x, x_outer, weight): + if weight >= 0.0: + return x + length = x.shape[-2] + wr = int((length * (1 - (-weight))) // 2) + + x[...,:wr,:] = x_outer[...,:wr,:] + x[...,-wr:,:] = x_outer[...,-wr:,:] + return x + + def __call__(self, x, attr): + if x.shape[0] == 1 and not self.KONTEXT: + return x + + weight_list = getattr(self, attr) + weights_all_zero = all(weight == 0.0 for weight in weight_list) + if weights_all_zero: + return x + + #self.HEADS=24 + #x_ndim = x.ndim + #if x_ndim == 3: + # B, HW, C = x.shape + # if x.shape[-2] != self.HEADS and self.HEADS != 0: + # x = x.reshape(B,self.HEADS,HW,-1) + + HEAD_DIM = x.shape[1] + if HEAD_DIM == self.HEADS: + B, HEAD_DIM, HW, C = x.shape + x = x.reshape(B, HW, C*HEAD_DIM) + + if hasattr(self, "KONTEXT") and self.KONTEXT == 1: + x = x.reshape(2, x.shape[1] // 2, x.shape[2]) + + txt_slice, img_slice, ktx_slice = self.txt_slice, self.img_slice, None + if hasattr(self, "KONTEXT") and self.KONTEXT == 2: + ktx_slice = self.img_slice # slice(2 * self.img_slice.start, None) + img_slice = slice(2 * self.img_slice.start, self.img_slice.start) + txt_slice = slice(None, 2 * self.txt_slice.stop) + + weights_all_one = all(weight == 1.0 for weight in weight_list) + methods_all_scattersort = all(name == "scattersort" for name in self.method) + masks_all_none = all(mask is None for mask in self.mask) + + if weights_all_one and methods_all_scattersort and len(weight_list) > 1 and masks_all_none: + buf = Stylizer.buffer + buf['src_idx'] = x[0:1].argsort(dim=-2) + buf['ref_sorted'], buf['ref_idx'] = x[1:].reshape(1, -1, x.shape[-1]).sort(dim=-2) + buf['src'] = buf['ref_sorted'][:,::len(weight_list)].expand_as(buf['src_idx']) # interleave_stride = len(weight_list) + + x[0:1] = x[0:1].scatter_(dim=-2, index=buf['src_idx'], src=buf['src'],) + + else: + for i, (weight, mask) in enumerate(zip(weight_list, self.mask)): + if mask is not None: + x01 = x[0:1].clone() + slc = Stylizer.middle_slice(x.shape[-2], weight) + #slc = slice(None) + + txt_method_name = self.method[i].removeprefix("tiled_") + txt_method = getattr(self, txt_method_name) + + method_name = self.method[i].removeprefix("tiled_") if self.img_len > x.shape[-2] or self.h_len < 0 else self.method[i] + method = getattr(self, method_name) + apply_to = self.apply_to[i] + if weight == 0.0: + continue + else: # if weight == 1.0: + if weight > 0 and weight < 1: + x_clone = x.clone() + if self.img_len == x.shape[-2] or apply_to == "img+txt" or self.h_len < 0: + x = method(x, idx=i+1, slc=slc) + elif self.img_len < x.shape[-2]: + if "img" in apply_to: + x[...,img_slice,:] = method(x[...,img_slice,:], idx=i+1, slc=slc) + #if ktx_slice is not None: + # x[...,ktx_slice,:] = method(x[...,ktx_slice,:], idx=i+1) + #x[:,:self.img_len,:] = method(x[:,:self.img_len,:], idx=i+1) + if "txt" in apply_to: + x[...,txt_slice,:] = txt_method(x[...,txt_slice,:], idx=i+1, slc=slc) + #x[:,self.img_len:,:] = method(x[:,self.img_len:,:], idx=i+1) + if not "img" in apply_to and not "txt" in apply_to: + pass + else: + x = method(x, idx=i+1, slc=slc) + if weight > 0 and weight < 1 and txt_method_name != "scattersort": + x = torch.lerp(x_clone, x, weight) + #else: + # x = torch.lerp(x, method(x.clone(), idx=i+1), weight) + + if mask is not None: + x[0:1,...,img_slice,:] = torch.lerp(x01[...,img_slice,:], x[0:1,...,img_slice,:], mask.view(1, -1, 1)) + if ktx_slice is not None: + x[0:1,...,ktx_slice,:] = torch.lerp(x01[...,ktx_slice,:], x[0:1,...,ktx_slice,:], mask.view(1, -1, 1)) + #x[0:1,:self.img_len] = torch.lerp(x01[:,:self.img_len], x[0:1,:self.img_len], mask.view(1, -1, 1)) + + #if x_ndim == 3: + # return x.view(B,HW,C) + if hasattr(self, "KONTEXT") and self.KONTEXT == 1: + x = x.reshape(1, x.shape[1] * 2, x.shape[2]) + + if HEAD_DIM == self.HEADS: + return x.reshape(B, HEAD_DIM, HW, C) + else: + return x + + + + def WCT(self, x, idx=1): + Stylizer.CLS_WCT.set(x[idx:idx+1]) + x[0:1] = Stylizer.CLS_WCT.get(x[0:1]) + return x + + def WCT2(self, x, idx=1): + Stylizer.CLS_WCT2.set(x[idx:idx+1], self.h_len, self.w_len) + x[0:1] = Stylizer.CLS_WCT2.get(x[0:1], self.h_len, self.w_len) + return x + + @staticmethod + def AdaIN_(x, y, eps: float = 1e-7) -> torch.Tensor: + mean_c = x.mean(-2, keepdim=True) + std_c = x.std (-2, keepdim=True).add_(eps) # in-place add + mean_s = y.mean (-2, keepdim=True) + std_s = y.std (-2, keepdim=True).add_(eps) + x.sub_(mean_c).div_(std_c).mul_(std_s).add_(mean_s) # in-place chain + return x + + def AdaIN(self, x, idx=1, eps: float = 1e-7) -> torch.Tensor: + mean_c = x[0:1].mean(-2, keepdim=True) + std_c = x[0:1].std (-2, keepdim=True).add_(eps) # in-place add + mean_s = x[idx:idx+1].mean (-2, keepdim=True) + std_s = x[idx:idx+1].std (-2, keepdim=True).add_(eps) + x[0:1].sub_(mean_c).div_(std_c).mul_(std_s).add_(mean_s) # in-place chain + return x + + def injection(self, x:torch.Tensor, idx=1) -> torch.Tensor: + x[0:1] = x[idx:idx+1] + return x + + @staticmethod + def injection_(x:torch.Tensor, y:torch.Tensor) -> torch.Tensor: + return y + + @staticmethod + def passthrough(x:torch.Tensor, idx=1) -> torch.Tensor: + return x + + @staticmethod + def decompose_magnitude_direction(x, dim=-1, eps=1e-8): + magnitude = x.norm(p=2, dim=dim, keepdim=True) + direction = x / (magnitude + eps) + return magnitude, direction + + @staticmethod + def scattersort_dir_(x, y, dim=-2): + #buf = Stylizer.buffer + #buf['src_sorted'], buf['src_idx'] = x.sort(dim=-2) + #buf['ref_sorted'], buf['ref_idx'] = y.sort(dim=-2) + #mag, _ = Stylizer.decompose_magnitude_direction(buf['src_sorted'], dim) + #_, dir = Stylizer.decompose_magnitude_direction(buf['ref_sorted'], dim) + mag, _ = Stylizer.decompose_magnitude_direction(x.to(torch.float64), dim) + + buf = Stylizer.buffer + buf['src_idx'] = x.argsort(dim=-2) + buf['ref_sorted'], buf['ref_idx'] = y .sort(dim=-2) + x.scatter_(dim=-2, index=buf['src_idx'], src=buf['ref_sorted'].expand_as(buf['src_idx'])) + + + _, dir = Stylizer.decompose_magnitude_direction(x.to(torch.float64), dim) + + return (mag * dir).to(x) + + + @staticmethod + def scattersort_dir2_(x, y, dim=-2): + #buf = Stylizer.buffer + #buf['src_sorted'], buf['src_idx'] = x.sort(dim=-2) + #buf['ref_sorted'], buf['ref_idx'] = y.sort(dim=-2) + #mag, _ = Stylizer.decompose_magnitude_direction(buf['src_sorted'], dim) + #_, dir = Stylizer.decompose_magnitude_direction(buf['ref_sorted'], dim) + + + buf = Stylizer.buffer + buf['src_sorted'], buf['src_idx'] = x.sort(dim=dim) + buf['ref_sorted'], buf['ref_idx'] = y.sort(dim=dim) + + + + + buf['x_sub'], buf['x_sub_idx'] = buf['src_sorted'].sort(dim=-1) + buf['y_sub'], buf['y_sub_idx'] = buf['ref_sorted'].sort(dim=-1) + + mag, _ = Stylizer.decompose_magnitude_direction(buf['x_sub'].to(torch.float64), -1) + _, dir = Stylizer.decompose_magnitude_direction(buf['y_sub'].to(torch.float64), -1) + + buf['y_sub'] = (mag * dir).to(x) + + buf['ref_sorted'].scatter_(dim=-1, index=buf['y_sub_idx'], src=buf['y_sub'].expand_as(buf['y_sub_idx'])) + + + + mag, _ = Stylizer.decompose_magnitude_direction(buf['src_sorted'].to(torch.float64), dim) + _, dir = Stylizer.decompose_magnitude_direction(buf['ref_sorted'].to(torch.float64), dim) + + buf['ref_sorted'] = (mag * dir).to(x) + + x.scatter_(dim=dim, index=buf['src_idx'], src=buf['ref_sorted'].expand_as(buf['src_idx'])) + + + return x + + + @staticmethod + def scattersort_dir(x, idx=1): + x[0:1] = Stylizer.scattersort_dir_(x[0:1], x[idx:idx+1]) + return x + + + @staticmethod + def scattersort_dir2(x, idx=1): + x[0:1] = Stylizer.scattersort_dir2_(x[0:1], x[idx:idx+1]) + return x + + @staticmethod + def scattersort_(x, y, slc=slice(None)): + buf = Stylizer.buffer + buf['src_idx'] = x.argsort(dim=-2) + buf['ref_sorted'], buf['ref_idx'] = y .sort(dim=-2) + + return x.scatter_(dim=-2, index=buf['src_idx'][...,slc,:], src=buf['ref_sorted'][...,slc,:].expand_as(buf['src_idx'][...,slc,:])) + + + @staticmethod + def scattersort_double(x, y): + buf = Stylizer.buffer + buf['src_sorted'], buf['src_idx'] = x.sort(dim=-2) + buf['ref_sorted'], buf['ref_idx'] = y.sort(dim=-2) + + buf['x_sub_idx'] = buf['src_sorted'].argsort(dim=-1) + buf['y_sub'], buf['y_sub_idx'] = buf['ref_sorted'].sort(dim=-1) + + x.scatter_(dim=-1, index=buf['x_sub_idx'], src=buf['y_sub'].expand_as(buf['x_sub_idx'])) + + return x.scatter_(dim=-2, index=buf['src_idx'], src=buf['ref_sorted'].expand_as(buf['src_idx'])) + + + def scattersort_aoeu(self, x, idx=1, slc=slice(None)): + x[0:1] = Stylizer.scattersort_(x[0:1], x[idx:idx+1], slc) + return x + + def scattersort(self, x, idx=1, slc=slice(None)): + if x.shape[0] != 2: + x[0:1] = Stylizer.scattersort_(x[0:1], x[idx:idx+1], slc) + return x + + buf = Stylizer.buffer + buf['sorted'], buf['idx'] = x.sort(dim=-2) + + return x.scatter_(dim=-2, index=buf['idx'][0:1][...,slc,:], src=buf['sorted'][1:2][...,slc,:].expand_as(buf['idx'][0:1][...,slc,:])) + + + + + def tiled_scattersort(self, x, idx=1): #, h_tile=None, w_tile=None): + #if HDModel.RECON_MODE: + # return denoised_embed + #den = x[0:1] [:,:self.img_len,:].view(-1, 2560, self.h_len, self.w_len) + #style = x[idx:idx+1][:,:self.img_len,:].view(-1, 2560, self.h_len, self.w_len) + #h_tile = self.h_tile[idx-1] if h_tile is None else h_tile + #w_tile = self.w_tile[idx-1] if w_tile is None else w_tile + + C = x.shape[-1] + den = x[0:1] [:,self.img_slice,:].reshape(-1, C, self.h_len, self.w_len) + style = x[idx:idx+1][:,self.img_slice,:].reshape(-1, C, self.h_len, self.w_len) + + tiles = Stylizer.get_tiles_as_strided(den, self.h_tile[idx-1], self.w_tile[idx-1]) + ref_tile = Stylizer.get_tiles_as_strided(style, self.h_tile[idx-1], self.w_tile[idx-1]) + + # rearrange for vmap to run on (nH, nW) ( as outer axes) + tiles_v = tiles .permute(2, 3, 0, 1, 4, 5) # (nH, nW, B, C, tile_h, tile_w) + ref_tile_v = ref_tile.permute(2, 3, 0, 1, 4, 5) # (nH, nW, B, C, tile_h, tile_w) + + # vmap over spatial dimms (nH, nW)... num of tiles high, num tiles wide + vmap2 = torch.vmap(torch.vmap(Stylizer.apply_scattersort_per_tile, in_dims=0), in_dims=0) + result = vmap2(tiles_v, ref_tile_v) # (nH, nW, B, C, tile_h, tile_w) + + # --> (B, C, nH, nW, tile_h, tile_w) + result = result.permute(2, 3, 0, 1, 4, 5) #( B, C, nH, nW, tile_h, tile_w) + + # in-place copy, werx if result has same shape/strides as tiles... overwrites same mem location "content" is using + tiles.copy_(result) + + return x + + + def tiled_AdaIN(self, x, idx=1): + #if HDModel.RECON_MODE: + # return denoised_embed + #den = x[0:1] [:,:self.img_len,:].view(-1, 2560, self.h_len, self.w_len) + #style = x[idx:idx+1][:,:self.img_len,:].view(-1, 2560, self.h_len, self.w_len) + C = x.shape[-1] + den = x[0:1] [:,self.img_slice,:].reshape(-1, C, self.h_len, self.w_len) + style = x[idx:idx+1][:,self.img_slice,:].reshape(-1, C, self.h_len, self.w_len) + + tiles = Stylizer.get_tiles_as_strided(den, self.h_tile[idx-1], self.w_tile[idx-1]) + ref_tile = Stylizer.get_tiles_as_strided(style, self.h_tile[idx-1], self.w_tile[idx-1]) + + # rearrange for vmap to run on (nH, nW) ( as outer axes) + tiles_v = tiles .permute(2, 3, 0, 1, 4, 5) # (nH, nW, B, C, tile_h, tile_w) + ref_tile_v = ref_tile.permute(2, 3, 0, 1, 4, 5) # (nH, nW, B, C, tile_h, tile_w) + + # vmap over spatial dimms (nH, nW)... num of tiles high, num tiles wide + vmap2 = torch.vmap(torch.vmap(Stylizer.apply_AdaIN_per_tile, in_dims=0), in_dims=0) + result = vmap2(tiles_v, ref_tile_v) # (nH, nW, B, C, tile_h, tile_w) + + # --> (B, C, nH, nW, tile_h, tile_w) + result = result.permute(2, 3, 0, 1, 4, 5) #( B, C, nH, nW, tile_h, tile_w) + + # in-place copy, werx if result has same shape/strides as tiles... overwrites same mem location "content" is using + tiles.copy_(result) + + return x + + + @staticmethod + def get_tiles_as_strided(x, tile_h, tile_w): + B, C, H, W = x.shape + stride = x.stride() + nH = H // tile_h + nW = W // tile_w + + tiles = x.as_strided( + size=(B, C, nH, nW, tile_h, tile_w), + stride=(stride[0], stride[1], stride[2] * tile_h, stride[3] * tile_w, stride[2], stride[3]) + ) + return tiles # shape: (B, C, nH, nW, tile_h, tile_w) + + @staticmethod + def apply_scattersort_per_tile(tile, ref_tile): + flat = tile .flatten(-2, -1) + ref_flat = ref_tile.flatten(-2, -1) + + sorted_ref, _ = ref_flat .sort(dim=-1) + src_sorted, src_idx = flat.sort(dim=-1) + + out = flat.scatter(dim=-1, index=src_idx, src=sorted_ref) + return out.view_as(tile) + + @staticmethod + def apply_AdaIN_per_tile(tile, ref_tile, eps: float = 1e-7): + mean_c = tile.mean(-2, keepdim=True) + std_c = tile.std (-2, keepdim=True).add_(eps) # in-place add + mean_s = ref_tile.mean (-2, keepdim=True) + std_s = ref_tile.std (-2, keepdim=True).add_(eps) + tile.sub_(mean_c).div_(std_c).mul_(std_s).add_(mean_s) # in-place chain + return tile + +class StyleMMDiT_Attn(Stylizer): + def __init__(self, mode): + super().__init__() + + self.q_proj = [0.0] + self.k_proj = [0.0] + self.v_proj = [0.0] + + self.q_norm = [0.0] + self.k_norm = [0.0] + + self.out = [0.0] + +class StyleMMDiT_FF(Stylizer): # these hit img or joint only, never txt + def __init__(self, mode): + super().__init__() + + self.ff_1 = [0.0] + self.ff_1_silu = [0.0] + self.ff_3 = [0.0] + self.ff_13 = [0.0] + self.ff_2 = [0.0] + +class StyleMMDiT_MoE(Stylizer): # these hit img or joint only, never txt + def __init__(self, mode): + super().__init__() + + self.FF_SHARED = StyleMMDiT_FF(mode) + self.FF_SEPARATE = StyleMMDiT_FF(mode) + + self.shared = [0.0] + self.gate = [False] + self.topk_weight = [0.0] + + self.separate = [0.0] + self.sum = [0.0] + self.out = [0.0] + + + + + +class StyleMMDiT_SubBlock(Stylizer): + def __init__(self, mode): + super().__init__() + + self.ATTN = StyleMMDiT_Attn(mode) # options for attn itself: qkv proj, qk norm, attn out + + self.attn_norm = [0.0] + self.attn_norm_mod = [0.0] + self.attn = [0.0] + self.attn_gated = [0.0] + self.attn_res = [0.0] + + self.ff_norm = [0.0] + self.ff_norm_mod = [0.0] + self.ff = [0.0] + self.ff_gated = [0.0] + self.ff_res = [0.0] + + self.mask = [None] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.ATTN.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + +class StyleMMDiT_IMG_Block(StyleMMDiT_SubBlock): # img or joint + def __init__(self, mode): + super().__init__(mode) + self.FF = StyleMMDiT_MoE(mode) # options for MoE if img or joint + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.FF.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + +class StyleMMDiT_TXT_Block(StyleMMDiT_SubBlock): # txt only + def __init__(self, mode): + super().__init__(mode) + self.FF = StyleMMDiT_FF(mode) # options for FF within MoE for img or joint; or for txt alone + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.FF.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + + + + + +class StyleMMDiT_BaseBlock: + def __init__(self, mode="passthrough"): + + self.img = StyleMMDiT_IMG_Block(mode) + self.txt = StyleMMDiT_TXT_Block(mode) + + self.mask = [None] + self.attn_mask = [None] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + self.h_len = h_len + self.w_len = w_len + self.img_len = h_len * w_len + + self.img_slice = img_slice + self.txt_slice = txt_slice + self.HEADS = HEADS + + self.img.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.txt.set_len(-1, -1, img_slice, txt_slice, HEADS) + + for i, mask in enumerate(self.mask): + if mask is not None and mask.ndim > 1: + self.mask[i] = F.interpolate(mask.unsqueeze(0), size=(h_len, w_len)).flatten().to(torch.bfloat16).cuda() + self.img.mask = self.mask + for i, mask in enumerate(self.attn_mask): + if mask is not None and mask.ndim > 1: + self.attn_mask[i] = F.interpolate(mask.unsqueeze(0), size=(h_len, w_len)).flatten().to(torch.bfloat16).cuda() + self.img.ATTN.mask = self.attn_mask + +class StyleMMDiT_DoubleBlock(StyleMMDiT_BaseBlock): + def __init__(self, mode="passthrough"): + super().__init__(mode) + self.txt = StyleMMDiT_TXT_Block(mode) + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.txt.set_len(-1, -1, img_slice, txt_slice, HEADS) + +class StyleMMDiT_SingleBlock(StyleMMDiT_BaseBlock): + def __init__(self, mode="passthrough"): + super().__init__(mode) + + + + + + + + + + + + + + + + + + + + + + + + + + + + +class StyleUNet_Resample(Stylizer): + def __init__(self, mode): + super().__init__() + self.conv = [0.0] + +class StyleUNet_Attn(Stylizer): + def __init__(self, mode): + super().__init__() + self.q_proj = [0.0] + self.k_proj = [0.0] + self.v_proj = [0.0] + self.out = [0.0] + +class StyleUNet_FF(Stylizer): + def __init__(self, mode): + super().__init__() + self.proj = [0.0] + self.geglu = [0.0] + self.linear = [0.0] + +class StyleUNet_TransformerBlock(Stylizer): + def __init__(self, mode): + super().__init__() + + self.ATTN1 = StyleUNet_Attn(mode) # self-attn + self.FF = StyleUNet_FF (mode) + self.ATTN2 = StyleUNet_Attn(mode) # cross-attn + + self.self_attn = [0.0] + self.ff = [0.0] + self.cross_attn = [0.0] + + self.self_attn_res = [0.0] + self.cross_attn_res = [0.0] + self.ff_res = [0.0] + + self.norm1 = [0.0] + self.norm2 = [0.0] + self.norm3 = [0.0] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.ATTN1.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.ATTN2.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + +class StyleUNet_SpatialTransformer(Stylizer): + def __init__(self, mode): + super().__init__() + + self.TFMR = StyleUNet_TransformerBlock(mode) + + self.spatial_norm_in = [0.0] + self.spatial_proj_in = [0.0] + self.spatial_transformer_block = [0.0] + self.spatial_transformer = [0.0] + self.spatial_proj_out = [0.0] + self.spatial_res = [0.0] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.TFMR.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + +class StyleUNet_ResBlock(Stylizer): + def __init__(self, mode): + super().__init__() + + self.in_norm = [0.0] + self.in_silu = [0.0] + self.in_conv = [0.0] + + self.emb_silu = [0.0] + self.emb_linear = [0.0] + self.emb_res = [0.0] + + self.out_norm = [0.0] + self.out_silu = [0.0] + self.out_conv = [0.0] + + self.residual = [0.0] + + +class StyleUNet_BaseBlock(Stylizer): + def __init__(self, mode="passthrough"): + + self.resample_block = StyleUNet_Resample(mode) + self.res_block = StyleUNet_ResBlock(mode) + self.spatial_block = StyleUNet_SpatialTransformer(mode) + + self.resample = [0.0] + self.res = [0.0] + self.spatial = [0.0] + + self.mask = [None] + self.attn_mask = [None] + + self.KONTEXT = 0 + + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + self.h_len = h_len + self.w_len = w_len + self.img_len = h_len * w_len + + self.img_slice = img_slice + self.txt_slice = txt_slice + self.HEADS = HEADS + + self.resample_block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.res_block .set_len(h_len, w_len, img_slice, txt_slice, HEADS) + self.spatial_block .set_len(h_len, w_len, img_slice, txt_slice, HEADS) + + for i, mask in enumerate(self.mask): + if mask is not None and mask.ndim > 1: + self.mask[i] = F.interpolate(mask.unsqueeze(0), size=(h_len, w_len)).flatten().to(torch.bfloat16).cuda() + self.resample_block.mask = self.mask + self.res_block.mask = self.mask + self.spatial_block.mask = self.mask + self.spatial_block.TFMR.mask = self.mask + + for i, mask in enumerate(self.attn_mask): + if mask is not None and mask.ndim > 1: + self.attn_mask[i] = F.interpolate(mask.unsqueeze(0), size=(h_len, w_len)).flatten().to(torch.bfloat16).cuda() + self.spatial_block.TFMR.ATTN1.mask = self.attn_mask + + def __call__(self, x, attr): + B, C, H, W = x.shape + x = super().__call__(x.reshape(B, H*W, C), attr) + return x.reshape(B,C,H,W) + + +class StyleUNet_InputBlock(StyleUNet_BaseBlock): + def __init__(self, mode="passthrough"): + super().__init__(mode) + +class StyleUNet_MiddleBlock(StyleUNet_BaseBlock): + def __init__(self, mode="passthrough"): + super().__init__(mode) + +class StyleUNet_OutputBlock(StyleUNet_BaseBlock): + def __init__(self, mode="passthrough"): + super().__init__(mode) + + + + + + + + + + + + + + + + +class Style_Model(Stylizer): + + def __init__(self, dtype=torch.float64, device=torch.device("cuda")): + super().__init__(dtype, device) + self.guides = [] + self.GUIDES_INITIALIZED = False + + #self.double_blocks = [StyleMMDiT_DoubleBlock() for _ in range(100)] + #self.single_blocks = [StyleMMDiT_SingleBlock() for _ in range(100)] + + self.h_len = -1 + self.w_len = -1 + self.img_len = -1 + self.h_tile = [-1] + self.w_tile = [-1] + + self.proj_in = [0.0] # these are for img only! not sliced + self.proj_out = [0.0] + + self.cond_pos = [None] + self.cond_neg = [None] + + self.noise_mode = "update" + self.recon_lure = "none" + self.data_shock = "none" + + self.data_shock_start_step = 0 + self.data_shock_end_step = 0 + + self.Retrojector = None + self.Endojector = None + + self.IMG_1ST = True + self.HEADS = 0 + self.KONTEXT = 0 + def __call__(self, x, attr): + if x.shape[0] == 1 and not self.KONTEXT: + return x + + weight_list = getattr(self, attr) + weights_all_zero = all(weight == 0.0 for weight in weight_list) + if weights_all_zero: + return x + + """x_ndim = x.ndim + if x_ndim == 4: + B, HEAD, HW, C = x.shape + + if x_ndim == 3: + B, HW, C = x.shape + if x.shape[-2] != self.HEADS and self.HEADS != 0: + x = x.reshape(B,self.HEADS,HW,-1)""" + + HEAD_DIM = x.shape[1] + if HEAD_DIM == self.HEADS: + B, HEAD_DIM, HW, C = x.shape + x = x.reshape(B, HW, C*HEAD_DIM) + + if self.KONTEXT == 1: + x = x.reshape(2, x.shape[1] // 2, x.shape[2]) + + weights_all_one = all(weight == 1.0 for weight in weight_list) + methods_all_scattersort = all(name == "scattersort" for name in self.method) + masks_all_none = all(mask is None for mask in self.mask) + + if weights_all_one and methods_all_scattersort and len(weight_list) > 1 and masks_all_none: + buf = Stylizer.buffer + buf['src_idx'] = x[0:1].argsort(dim=-2) + buf['ref_sorted'], buf['ref_idx'] = x[1:].reshape(1, -1, x.shape[-1]).sort(dim=-2) + buf['src'] = buf['ref_sorted'][:,::len(weight_list)].expand_as(buf['src_idx']) # interleave_stride = len(weight_list) + + x[0:1] = x[0:1].scatter_(dim=-2, index=buf['src_idx'], src=buf['src'],) + else: + for i, (weight, mask) in enumerate(zip(weight_list, self.mask)): + if weight > 0 and weight < 1: + x_clone = x.clone() + if mask is not None: + x01 = x[0:1].clone() + slc = Stylizer.middle_slice(x.shape[-2], weight) + + method = getattr(self, self.method[i]) + if weight == 0.0: + continue + elif weight == 1.0: + x = method(x, idx=i+1) + else: + x = method(x, idx=i+1, slc=slc) + if weight > 0 and weight < 1 and self.method[i] != "scattersort": + x = torch.lerp(x_clone, x, weight) + + #else: + # x = torch.lerp(x, method(x.clone(), idx=i), weight) + + if mask is not None: + x[0:1] = torch.lerp(x01, x[0:1], mask.view(1, -1, 1)) + + #if x_ndim == 3: + # return x.view(B,HW,C) + if self.KONTEXT == 1: + x = x.reshape(1, x.shape[1] * 2, x.shape[2]) + + if HEAD_DIM == self.HEADS: + return x.reshape(B, HEAD_DIM, HW, C) + else: + return x + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + self.h_len = h_len + self.w_len = w_len + self.img_len = h_len * w_len + + self.img_slice = img_slice + self.txt_slice = txt_slice + self.HEADS = HEADS + + #for block in self.double_blocks: + # block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + #for block in self.single_blocks: + # block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + + for i, mask in enumerate(self.mask): + if mask is not None and mask.ndim > 1: + self.mask[i] = F.interpolate(mask.unsqueeze(0), size=(h_len, w_len)).flatten().to(torch.bfloat16).cuda() + + def init_guides(self, model): + if not self.GUIDES_INITIALIZED: + if self.guides == []: + self.guides = None + elif self.guides is not None: + for i, latent in enumerate(self.guides): + if type(latent) is dict: + latent = model.inner_model.inner_model.process_latent_in(latent['samples']).to(dtype=self.dtype, device=self.device) + elif type(latent) is torch.Tensor: + latent = latent.to(dtype=self.dtype, device=self.device) + else: + latent = None + #raise ValueError(f"Invalid latent type: {type(latent)}") + + #if self.VIDEO and latent.shape[2] == 1: + # latent = latent.repeat(1, 1, x.shape[2], 1, 1) + + self.guides[i] = latent + if any(g is None for g in self.guides): + self.guides = None + print("Style guide nonetype set for Kontext.") + else: + self.guides = torch.cat(self.guides, dim=0) + self.GUIDES_INITIALIZED = True + + def set_conditioning(self, positive, negative): + self.cond_pos = [positive] + self.cond_neg = [negative] + + def apply_style_conditioning(self, UNCOND, base_context, base_y=None, base_llama3=None): + + def get_max_token_lengths(style_conditioning, base_context, base_y=None, base_llama3=None): + context_max_len = base_context.shape[-2] + llama3_max_len = base_llama3.shape[-2] if base_llama3 is not None else -1 + y_max_len = base_y.shape[-1] if base_y is not None else -1 + + for style_cond in style_conditioning: + if style_cond is None: + continue + context_max_len = max(context_max_len, style_cond[0][0].shape[-2]) + if base_llama3 is not None: + llama3_max_len = max(llama3_max_len, style_cond[0][1]['conditioning_llama3'].shape[-2]) + if base_y is not None: + y_max_len = max(y_max_len, style_cond[0][1]['pooled_output'].shape[-1]) + + return context_max_len, llama3_max_len, y_max_len + + def pad_to_len(x, target_len, pad_value=0.0, dim=1): + if target_len < 0: + return x + cur_len = x.shape[dim] + if cur_len == target_len: + return x + return F.pad(x, (0, 0, 0, target_len - cur_len), value=pad_value) + + style_conditioning = self.cond_pos if not UNCOND else self.cond_neg + + context_max_len, llama3_max_len, y_max_len = get_max_token_lengths( + style_conditioning = style_conditioning, + base_context = base_context, + base_y = base_y, + base_llama3 = base_llama3, + ) + + bsz_style = len(style_conditioning) + + context = base_context.repeat(bsz_style + 1, 1, 1) + y = base_y.repeat(bsz_style + 1, 1) if base_y is not None else None + llama3 = base_llama3.repeat(bsz_style + 1, 1, 1, 1) if base_llama3 is not None else None + + context = pad_to_len(context, context_max_len, dim=-2) + llama3 = pad_to_len(llama3, llama3_max_len, dim=-2) if base_llama3 is not None else None + y = pad_to_len(y, y_max_len, dim=-1) if base_y is not None else None + + for ci, style_cond in enumerate(style_conditioning): + if style_cond is None: + continue + context[ci+1:ci+2] = pad_to_len(style_cond[0][0], context_max_len, dim=-2).to(context) + if llama3 is not None: + llama3 [ci+1:ci+2] = pad_to_len(style_cond[0][1]['conditioning_llama3'], llama3_max_len, dim=-2).to(llama3) + if y is not None: + y [ci+1:ci+2] = pad_to_len(style_cond[0][1]['pooled_output'], y_max_len, dim=-1).to(y) + + return context, y, llama3 + + def WCT_data(self, denoised_embed, y0_style_embed): + Stylizer.CLS_WCT.set(y0_style_embed.to(denoised_embed)) + return Stylizer.CLS_WCT.get(denoised_embed) + + def WCT2_data(self, denoised_embed, y0_style_embed): + Stylizer.CLS_WCT2.set(y0_style_embed.to(denoised_embed)) + return Stylizer.CLS_WCT2.get(denoised_embed) + + def apply_to_data(self, denoised, y0_style=None, mode="none"): + if mode == "none": + return denoised + y0_style = self.guides if y0_style is None else y0_style + + y0_style_embed = self.Retrojector.embed(y0_style) + denoised_embed = self.Retrojector.embed(denoised) + B,HW,C = y0_style_embed.shape + embed = torch.cat([denoised_embed, y0_style_embed.view(1,B*HW,C)[:,::B,:]], dim=0) + method = getattr(self, mode) + if mode == "scattersort": + slc = Stylizer.middle_slice(embed.shape[-2], self.data_shock_weight) + embed = method(embed, slc=slc) + else: + embed = method(embed) + return self.Retrojector.unembed(embed[0:1]) + + def apply_recon_lure(self, denoised, y0_style): + if self.recon_lure == "none": + return denoised + for i in range(denoised.shape[0]): + denoised[i:i+1] = self.apply_to_data(denoised[i:i+1], y0_style, self.recon_lure) + return denoised + + def apply_data_shock(self, denoised): + if self.data_shock == "none": + return denoised + datashock_ref = getattr(self, "datashock_ref", None) + if self.data_shock == "scattersort": + return self.apply_to_data(denoised, datashock_ref, self.data_shock) + else: + return torch.lerp(denoised, self.apply_to_data(denoised, datashock_ref, self.data_shock), torch.Tensor([self.data_shock_weight]).double().cuda()) + + + + +class StyleMMDiT_Model(Style_Model): + + def __init__(self, dtype=torch.float64, device=torch.device("cuda")): + super().__init__(dtype, device) + self.double_blocks = [StyleMMDiT_DoubleBlock() for _ in range(100)] + self.single_blocks = [StyleMMDiT_SingleBlock() for _ in range(100)] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + for block in self.double_blocks: + block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + for block in self.single_blocks: + block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + + +class StyleUNet_Model(Style_Model): + + def __init__(self, dtype=torch.float64, device=torch.device("cuda")): + super().__init__(dtype, device) + self.input_blocks = [StyleUNet_InputBlock() for _ in range(100)] + self.middle_blocks = [StyleUNet_MiddleBlock() for _ in range(100)] + self.output_blocks = [StyleUNet_OutputBlock() for _ in range(100)] + + def set_len(self, h_len, w_len, img_slice, txt_slice, HEADS): + super().set_len(h_len, w_len, img_slice, txt_slice, HEADS) + for block in self.input_blocks: + block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + for block in self.middle_blocks: + block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + for block in self.output_blocks: + block.set_len(h_len, w_len, img_slice, txt_slice, HEADS) + + def __call__(self, x, attr): + B, C, H, W = x.shape + x = super().__call__(x.reshape(B, H*W, C), attr) + return x.reshape(B,C,H,W) + diff --git a/tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py index a98d277..f64df65 100644 --- a/tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py +++ b/tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py @@ -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( diff --git a/tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py index 38bdbb5..34abd1a 100644 --- a/tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py +++ b/tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py @@ -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}) diff --git a/tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py index 710b94c..3191a70 100644 --- a/tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py +++ b/tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py @@ -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( diff --git a/tests/sampling/test_ksampler_extras_node.py b/tests/sampling/test_ksampler_extras_node.py index 20d4253..d1e18ca 100644 --- a/tests/sampling/test_ksampler_extras_node.py +++ b/tests/sampling/test_ksampler_extras_node.py @@ -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 diff --git a/tests/sampling/test_res4lyf_composed_sampling.py b/tests/sampling/test_res4lyf_composed_sampling.py new file mode 100644 index 0000000..537c6ab --- /dev/null +++ b/tests/sampling/test_res4lyf_composed_sampling.py @@ -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) diff --git a/tests/sampling/test_res4lyf_sampling.py b/tests/sampling/test_res4lyf_sampling.py new file mode 100644 index 0000000..6b853a4 --- /dev/null +++ b/tests/sampling/test_res4lyf_sampling.py @@ -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 diff --git a/tests/sampling/test_sampling_samplers.py b/tests/sampling/test_sampling_samplers.py index 08343d6..c0fbf0d 100644 --- a/tests/sampling/test_sampling_samplers.py +++ b/tests/sampling/test_sampling_samplers.py @@ -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( diff --git a/tests/sampling/test_sampling_scheduler_references.py b/tests/sampling/test_sampling_scheduler_references.py index 90a1621..a4a9320 100644 --- a/tests/sampling/test_sampling_scheduler_references.py +++ b/tests/sampling/test_sampling_scheduler_references.py @@ -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 diff --git a/tests/sampling/test_sampling_schedulers.py b/tests/sampling/test_sampling_schedulers.py index 591a370..803e289 100644 --- a/tests/sampling/test_sampling_schedulers.py +++ b/tests/sampling/test_sampling_schedulers.py @@ -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", ) diff --git a/third_party/NOTICE.md b/third_party/NOTICE.md index 5909790..33d2d0c 100644 --- a/third_party/NOTICE.md +++ b/third_party/NOTICE.md @@ -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`. diff --git a/third_party/manifest.toml b/third_party/manifest.toml index b9f8383..ec0c9ae 100644 --- a/third_party/manifest.toml +++ b/third_party/manifest.toml @@ -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"