Compare commits
76
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6153443986 | ||
|
|
9f1870e350 | ||
|
|
77b238a672 | ||
|
|
bf6c517a4c | ||
|
|
0f076a886f | ||
|
|
b166cd8741 | ||
|
|
751b8e771c | ||
|
|
3b5a6b7971 | ||
|
|
a8e94be4c4 | ||
|
|
a2b713a9ea | ||
|
|
0af4e4194a | ||
|
|
dae4b0906e | ||
|
|
e41982cdab | ||
|
|
e04f34897a | ||
|
|
132609e3d2 | ||
|
|
0d4e8d20e2 | ||
|
|
a31f9be2d0 | ||
|
|
25267e588a | ||
|
|
2d4fe5ea47 | ||
|
|
d015307e1a | ||
|
|
db6dfb5949 | ||
|
|
6667f62b75 | ||
|
|
189afe2e45 | ||
|
|
47aa1ddfbe | ||
|
|
e933378b43 | ||
|
|
420c8a490c | ||
|
|
1ffbb6fa56 | ||
|
|
3d64e1d3b1 | ||
|
|
1b3c11ff35 | ||
|
|
08d3c6adc0 | ||
|
|
a948909964 | ||
|
|
c168188617 | ||
|
|
8719d21b28 | ||
|
|
e4c296be32 | ||
|
|
a1dc924a7e | ||
|
|
2bcbcc68a2 | ||
|
|
0e1043fe29 | ||
|
|
1b85f680e6 | ||
|
|
ad283cf47e | ||
|
|
ede2b072b6 | ||
|
|
2c205acb77 | ||
|
|
885fc43eac | ||
|
|
e4f9712de8 | ||
|
|
e85d22b705 | ||
|
|
fb850de8b7 | ||
|
|
d2d998110b | ||
|
|
2aa4c18b20 | ||
|
|
d9ba0ed67e | ||
|
|
3b6728b587 | ||
|
|
0e39e0162e | ||
|
|
64251ed44d | ||
|
|
3e84dc9f9a | ||
|
|
65d485c540 | ||
|
|
bb79bc7a46 | ||
|
|
c444686b6e | ||
|
|
bbfbe2ec64 | ||
|
|
3042ef4de5 | ||
|
|
8bb7b28a62 | ||
|
|
2d02f2bdc1 | ||
|
|
a1a3edd656 | ||
|
|
65cb78ee25 | ||
|
|
7e5959cb97 | ||
|
|
93f413b6f5 | ||
|
|
064ff6b4a0 | ||
|
|
42995e8705 | ||
|
|
8894cdef35 | ||
|
|
fc261d0a08 | ||
|
|
d7c7a6df6a | ||
|
|
42164a494a | ||
|
|
410ae514ba | ||
|
|
7aa9223784 | ||
|
|
34c1bee399 | ||
|
|
1a94ac4e8a | ||
|
|
ca64d701b6 | ||
|
|
d9a26d49ec | ||
|
|
733d85b667 |
@@ -0,0 +1,28 @@
|
||||
name: Publish to Comfy registry
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -1,2 +1,3 @@
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
.vs/
|
||||
|
||||
Binary file not shown.
@@ -1,23 +0,0 @@
|
||||
{
|
||||
"Version": 1,
|
||||
"WorkspaceRootPath": "C:\\Users\\marre\\source\\repos\\ComfyUI-ScaleLockedResidualDiffusion\\",
|
||||
"Documents": [],
|
||||
"DocumentGroupContainers": [
|
||||
{
|
||||
"Orientation": 0,
|
||||
"VerticalTabListWidth": 256,
|
||||
"DocumentGroups": [
|
||||
{
|
||||
"DockedWidth": 200,
|
||||
"SelectedChildIndex": -1,
|
||||
"Children": [
|
||||
{
|
||||
"$type": "Bookmark",
|
||||
"Name": "ST:0:0:{aa2115a1-9712-457b-9047-dbb71ca2cdd2}"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
{
|
||||
"ExpandedNodes": [
|
||||
""
|
||||
],
|
||||
"SelectedNode": "\\C:\\Users\\marre\\Source\\Repos\\ComfyUI-ScaleLockedResidualDiffusion",
|
||||
"PreviewInSolutionExplorer": false
|
||||
}
|
||||
@@ -1,15 +1,15 @@
|
||||
# ComfyUI-ScaleLockedResidualDiffusion
|
||||
|
||||
A custom ComfyUI node pack implementing a practical MVP of Scale-Locked Residual Diffusion for the specific failure mode where a model behaves well around its native / comfortable resolution (for example ~1 MP) but drifts badly in composition, anatomy, or identity at much higher resolutions.
|
||||
A ComfyUI custom node pack implementing a practical MVP of Scale-Locked Residual Diffusion for the specific failure mode where a model behaves well around its native / comfortable resolution (for example ~1 MP) but drifts badly in composition, anatomy, or identity at much higher resolutions.
|
||||
|
||||
## What it does
|
||||
|
||||
Instead of letting the high-resolution branch freely re-plan the image, the node:
|
||||
Instead of letting the high-resolution branch freely re-plan the image, the nodes:
|
||||
|
||||
1. creates a low-resolution planner pass at a target megapixel level,
|
||||
2. records the planner's per-step denoised x0 trajectory,
|
||||
3. builds nested high-resolution noise so the high-res branch shares the same coarse stochastic layout,
|
||||
4. runs the final high-res sampling with a custom CFG guider that locks only the low-frequency denoised structure toward the planner trajectory while preserving the base model's high-frequency residual detail.
|
||||
1. create a low-resolution planner pass at a target megapixel level,
|
||||
2. record the planner's per-step denoised x0 trajectory,
|
||||
3. build nested high-resolution noise so the high-res branch shares the same coarse stochastic layout,
|
||||
4. run the final high-res sampling with a scale-lock correction that preserves the base model's high-frequency residual detail.
|
||||
|
||||
In practice, this is meant to reduce:
|
||||
|
||||
@@ -23,30 +23,80 @@ In practice, this is meant to reduce:
|
||||
|
||||
### 1. Scale-Locked Residual KSampler
|
||||
|
||||
Main all-in-one node.
|
||||
Main all-in-one node. It still owns the full SLRD runtime internally and is the easiest entry point.
|
||||
|
||||
**Outputs**
|
||||
- `output`: final high-res latent
|
||||
- `lowres_planner`: final low-res planner latent
|
||||
- `denoised_output`: final high-res denoised x0 latent when available
|
||||
|
||||
**Important controls**
|
||||
- `target_megapixels`: planner resolution in pixel-space megapixels (usually `0.8` to `1.5` for Flux-like native planning)
|
||||
- `lock_strength`: global multiplier for the scale lock
|
||||
- `lock_strength_start` / `lock_strength_end`: how strongly the lock applies early vs late in denoising
|
||||
- `lock_schedule`: linear / cosine / flat interpolation for the lock schedule
|
||||
- `coarse_cutoff`: retained spatial fraction for the strongest coarse lock band
|
||||
- `mid_band_cutoff`: retained spatial fraction for an additional mid-frequency lock band
|
||||
- `mid_band_strength`: how strongly the mid-band is pulled toward the planner relative to the low-band schedule
|
||||
- `nested_noise_strength`: amount of zero-mean high-frequency detail noise added on top of the lifted low-res noise
|
||||
- `lock_mask` (optional): spatial mask to strengthen the lock only in selected regions (for example body / face / hands)
|
||||
- `pin_anchors`: store planner anchors in pinned CPU memory when possible for faster non-blocking transfer during the high-res pass
|
||||
- `sampler_guard`: `warn` / `error` / `off` guard for samplers outside a conservative alignment-safe allowlist
|
||||
### 2. Scale-Locked Runtime Context
|
||||
|
||||
### 2. Scale-Locked Nested Noise Preview
|
||||
Builds the reusable SLRD runtime bundle for ComfyUI's modular custom-sampling path.
|
||||
|
||||
**Outputs**
|
||||
- `runtime`: internal SLRD planner/noise context
|
||||
- `prepared_noise`: a `NOISE` object containing the aligned nested high-res noise
|
||||
- `lowres_planner`: final low-res planner latent
|
||||
|
||||
Use this before `SamplerCustomAdvanced` when you want the modular graph equivalent of the all-in-one sampler.
|
||||
|
||||
### 3. Scale-Locked CFG Guider
|
||||
|
||||
Public guider node for the standard ComfyUI `GUIDER` contract. This applies the scale-lock denoiser correction using a `runtime` from `Scale-Locked Runtime Context` and returns a fresh patched guider instead of mutating the upstream guider object in place.
|
||||
|
||||
Important: this node is only the denoiser-side piece. Full SLRD behavior still depends on the paired runtime context so the high-res noise field is aligned with the low-res planner branch.
|
||||
|
||||
### 4. Scale-Locked Residual SamplerCustomAdvanced
|
||||
|
||||
One-node version of the modular custom-sampling workflow. Internally it now calls the same runtime builder and guider patcher that power the public guider node.
|
||||
|
||||
### 5. Scale-Locked Detailer Hook Provider
|
||||
|
||||
Experimental Impact Pack / FaceDetailer integration path.
|
||||
|
||||
This node returns a `DETAILER_HOOK` object that captures the FaceDetailer crop mask in `post_upscale(...)` and drives the masked residual/manifold correction through the hook's sampler path.
|
||||
|
||||
The live hook path is `pre_ksample(...)` plus the custom sampler/runtime integration, not the inert `post_encode(...)` / `pre_decode(...)` pair.
|
||||
|
||||
This provider no longer advertises sampler-runtime controls that do not participate in its live sampler-driven hook path. The exposed knobs are the ones that still affect the masked latent correction directly:
|
||||
|
||||
- residual lock strength and cutoffs,
|
||||
- optional manifold companding controls,
|
||||
- optional `lock_mask` / `manifold_mask`, which are intersected with the FaceDetailer support mask.
|
||||
|
||||
Compatibility note: this provider's input signature changed when the inert detailer-only sampler/runtime controls were removed. Older saved workflows that used the previous `Scale-Locked Detailer Hook Provider` input surface will need to be re-wired to the current node inputs.
|
||||
|
||||
This has been validated for import/compile in this repo, but not against a live current Impact Pack checkout in this environment.
|
||||
|
||||
### 6. Scale-Locked Nested Noise Preview
|
||||
|
||||
Utility/debug node to inspect the nested-noise construction separately.
|
||||
|
||||
## Suggested modular workflow
|
||||
|
||||
For standard ComfyUI custom sampling:
|
||||
|
||||
1. Build your base `GUIDER` normally.
|
||||
2. Build `sigmas` and choose your `sampler` normally.
|
||||
3. Run `Scale-Locked Runtime Context` with the base `noise`, `guider`, `sampler`, `sigmas`, and target latent.
|
||||
4. Run `Scale-Locked CFG Guider` on the base guider using the returned `runtime`; this returns a fresh patched guider and leaves the base guider untouched.
|
||||
5. Feed `prepared_noise` and the patched guider into `SamplerCustomAdvanced`.
|
||||
|
||||
That path uses the same SLRD runtime pieces as the all-in-one sampler instead of duplicating planner/noise logic.
|
||||
|
||||
## FaceDetailer / Impact Pack note
|
||||
|
||||
A public guider node alone does not make FaceDetailer use SLRD. `Scale-Locked Detailer Hook Provider` is the experimental integration point for applying the masked residual/manifold correction inside the FaceDetailer crop lifecycle.
|
||||
|
||||
The current hook implementation is duck-typed rather than source-verified against a live Impact Pack checkout:
|
||||
|
||||
- FaceDetailer mask capture via `post_upscale(...)`,
|
||||
- alias-tolerant request capture via `pre_ksample(...)`,
|
||||
- masked latent correction via the custom sampler/runtime path.
|
||||
|
||||
Because Impact Pack's internal contracts can move, this integration should be treated as unverified runtime glue until it is exercised against the current Impact Pack source.
|
||||
|
||||
## Installation
|
||||
|
||||
Clone or copy this directory into your ComfyUI `custom_nodes` folder:
|
||||
@@ -67,10 +117,13 @@ For a first test when your high-res target is around 4 MP:
|
||||
- `lock_strength = 0.85`
|
||||
- `lock_strength_start = 0.95`
|
||||
- `lock_strength_end = 0.25`
|
||||
- `lock_schedule = cosine`
|
||||
- `lock_schedule = hold_then_drop`
|
||||
- `lock_schedule_hold = 0.35`
|
||||
- `lock_schedule_power = 3.0`
|
||||
- `coarse_cutoff = 0.33`
|
||||
- `mid_band_cutoff = 0.60`
|
||||
- `mid_band_strength = 0.35`
|
||||
- `mid_band_schedule = linked`
|
||||
- `sampler_guard = warn`
|
||||
- `nested_noise_strength = 0.35`
|
||||
|
||||
@@ -86,19 +139,6 @@ If the result feels too constrained / too similar to the low-res planner:
|
||||
- raise `coarse_cutoff`
|
||||
- lower `lock_strength_end`
|
||||
|
||||
## Recommended workflow pattern
|
||||
|
||||
Use this node exactly where you would normally use a KSampler for the high-resolution generation pass.
|
||||
|
||||
Typical graph:
|
||||
|
||||
1. checkpoint / text encodes
|
||||
2. empty latent or incoming img2img latent at your final target resolution
|
||||
3. Scale-Locked Residual KSampler
|
||||
4. VAE decode / detailers / final upscaling if desired
|
||||
|
||||
The node internally creates the planner pass for you, so you do not need to build a separate 1 MP sampler branch unless you want to compare outputs.
|
||||
|
||||
## Current limitations
|
||||
|
||||
This is a carefully implemented MVP, not a mathematically complete research system.
|
||||
@@ -111,41 +151,22 @@ What is already implemented:
|
||||
- residual-preserving coarse-field replacement,
|
||||
- optional pinned-memory anchor staging,
|
||||
- conservative sampler-alignment safety gating,
|
||||
- optional spatial masking.
|
||||
- optional spatial masking,
|
||||
- modular `GUIDER` exposure,
|
||||
- Impact-facing detailer hook/provider path.
|
||||
|
||||
What is not implemented yet:
|
||||
- automatic anatomy / pose / segmentation mask extraction,
|
||||
- explicit residual-only tiled model execution,
|
||||
- sigma-perfect trajectory matching for samplers that perform unusual extra model evaluations,
|
||||
- scheduled cutoff animation for coarse or mid bands,
|
||||
- multi-stage 1 MP -> 2 MP -> 4 MP progressive ladder inside one node,
|
||||
- exact support tuning for every possible exotic custom sampler.
|
||||
|
||||
## Why this implementation is conservative
|
||||
|
||||
This node avoids invasive patching of ComfyUI's internal sampler code. Instead it uses:
|
||||
|
||||
- the standard Comfy custom-node registration path,
|
||||
- the standard custom-sampling guider path,
|
||||
- standard sigma generation,
|
||||
- standard sampler objects,
|
||||
- standard preview callback behavior.
|
||||
|
||||
That makes it much easier to maintain and much less likely to break when ComfyUI internals shift.
|
||||
- exact support tuning for every possible exotic custom sampler,
|
||||
- verified bindings for every historical Impact Pack hook/provider variant.
|
||||
|
||||
## Files
|
||||
|
||||
- `__init__.py` - node registration
|
||||
- `nodes.py` - ComfyUI node definitions and runtime integration
|
||||
- `nodes.py` - ComfyUI node definitions, public guider node, and Impact hook/provider nodes
|
||||
- `slrd_runtime.py` - shared ComfyUI runtime for planner capture, nested noise, guider patching, and final sampling
|
||||
- `slrd_core.py` - algorithm core, nested noise, latent resizing, residual locking, trajectory helpers
|
||||
|
||||
## Sampler safety note
|
||||
|
||||
The current implementation aligns planner anchors to the final pass using outer-step / sigma progression heuristics.
|
||||
That works best with a conservative subset of samplers whose effective evaluation pattern is close to one visible step <-> one anchor step.
|
||||
|
||||
Because Comfy's custom sampling system is flexible and some samplers can perform more complicated internal evaluations,
|
||||
the node exposes `sampler_guard`:
|
||||
|
||||
- `warn`: log a warning for samplers outside the conservative safe set
|
||||
- `error`: refuse to run those samplers
|
||||
- `off`: trust the sampler and run anyway
|
||||
|
||||
@@ -12,12 +12,17 @@ A native / low-resolution planner branch is sampled first. Its denoised trajecto
|
||||
- `seed`, `steps`, `cfg`, `sampler_name`, `scheduler`, `denoise`: standard sampler controls
|
||||
- `target_megapixels`: planner resolution in pixel-space MP
|
||||
- `lock_strength`: overall lock multiplier
|
||||
- `lock_strength_start`: early-step lock amount
|
||||
- `lock_strength_end`: late-step lock amount
|
||||
- `lock_schedule`: linear / cosine / flat
|
||||
- `lock_strength_start`: early low-band lock amount
|
||||
- `lock_strength_end`: late low-band lock amount
|
||||
- `lock_schedule`: low-band schedule shape, evaluated against normalized log-sigma position when planner sigmas are available, with raw step-index fallback otherwise; `flat` keeps the start value for the whole run and ignores the end value
|
||||
- `lock_schedule_hold`: hold region before `hold_then_drop` releases
|
||||
- `lock_schedule_power`: curvature control for power-based schedules
|
||||
- `coarse_cutoff`: strongest coarse-band resolution fraction
|
||||
- `mid_band_cutoff`: second, looser mid-band resolution fraction
|
||||
- `mid_band_strength`: relative strength of the mid-band lock
|
||||
- `mid_band_strength`: overall mid-band lock multiplier
|
||||
- `mid_band_strength_start` / `mid_band_strength_end`: independent mid-band envelope when `mid_band_schedule` is not `linked`; it multiplies the base `mid_band_strength`
|
||||
- `mid_band_schedule`: `linked` for legacy behavior, or an independent curve mode
|
||||
- `mid_band_schedule_hold` / `mid_band_schedule_power`: shape controls for the independent mid-band schedule
|
||||
- `nested_noise_strength`: amount of extra high-frequency detail noise
|
||||
- `pin_anchors`: pinned-memory staging for planner anchors when possible
|
||||
- `sampler_guard`: warn / error / off handling for samplers outside the conservative safe set
|
||||
@@ -33,5 +38,7 @@ A native / low-resolution planner branch is sampled first. Its denoised trajecto
|
||||
|
||||
- Lower `coarse_cutoff` = stronger global structure control.
|
||||
- Lower `mid_band_cutoff` and higher `mid_band_strength` = tighter control over medium-scale body/shape structure.
|
||||
- `hold_then_drop` with a `lock_schedule_hold` around `0.30` to `0.45` gives a stronger early anchor with a later release knee.
|
||||
- Use `mid_band_schedule = linked` to preserve the legacy shared curve, or switch it off to let medium structure release earlier than the coarse band.
|
||||
- Higher `nested_noise_strength` = more detail freedom, but also more chance of drift.
|
||||
- A `lock_mask` is recommended for body-heavy and anatomy-sensitive generations.
|
||||
|
||||
@@ -12,12 +12,17 @@ A native / low-resolution planner branch is sampled first. Its denoised trajecto
|
||||
- `seed`, `steps`, `cfg`, `sampler_name`, `scheduler`, `denoise`: standard sampler controls
|
||||
- `target_megapixels`: planner resolution in pixel-space MP
|
||||
- `lock_strength`: overall lock multiplier
|
||||
- `lock_strength_start`: early-step lock amount
|
||||
- `lock_strength_end`: late-step lock amount
|
||||
- `lock_schedule`: linear / cosine / flat
|
||||
- `lock_strength_start`: early low-band lock amount
|
||||
- `lock_strength_end`: late low-band lock amount
|
||||
- `lock_schedule`: low-band schedule shape, evaluated against normalized log-sigma position when planner sigmas are available, with raw step-index fallback otherwise; `flat` keeps the start value for the whole run and ignores the end value
|
||||
- `lock_schedule_hold`: hold region before `hold_then_drop` releases
|
||||
- `lock_schedule_power`: curvature control for power-based schedules
|
||||
- `coarse_cutoff`: strongest coarse-band resolution fraction
|
||||
- `mid_band_cutoff`: second, looser mid-band resolution fraction
|
||||
- `mid_band_strength`: relative strength of the mid-band lock
|
||||
- `mid_band_strength`: overall mid-band lock multiplier
|
||||
- `mid_band_strength_start` / `mid_band_strength_end`: independent mid-band envelope when `mid_band_schedule` is not `linked`; it multiplies the base `mid_band_strength`
|
||||
- `mid_band_schedule`: `linked` for legacy behavior, or an independent curve mode
|
||||
- `mid_band_schedule_hold` / `mid_band_schedule_power`: shape controls for the independent mid-band schedule
|
||||
- `nested_noise_strength`: amount of extra high-frequency detail noise
|
||||
- `pin_anchors`: pinned-memory staging for planner anchors when possible
|
||||
- `sampler_guard`: warn / error / off handling for samplers outside the conservative safe set
|
||||
@@ -33,5 +38,7 @@ A native / low-resolution planner branch is sampled first. Its denoised trajecto
|
||||
|
||||
- Lower `coarse_cutoff` = stronger global structure control.
|
||||
- Lower `mid_band_cutoff` and higher `mid_band_strength` = tighter control over medium-scale body/shape structure.
|
||||
- `hold_then_drop` with a `lock_schedule_hold` around `0.30` to `0.45` gives a stronger early anchor with a later release knee.
|
||||
- Use `mid_band_schedule = linked` to preserve the legacy shared curve, or switch it off to let medium structure release earlier than the coarse band.
|
||||
- Higher `nested_noise_strength` = more detail freedom, but also more chance of drift.
|
||||
- A `lock_mask` is recommended for body-heavy and anatomy-sensitive generations.
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
[project]
|
||||
name = "scale-locked-residual-diffusion"
|
||||
description = "A ComfyUI custom node pack implementing Scale-Locked Residual Diffusion for high-resolution composition and anatomy stability."
|
||||
version = "1.0.32"
|
||||
license = { file = "LICENSE" }
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/xmarre/ComfyUI-ScaleLockedResidualDiffusion"
|
||||
Documentation = "https://github.com/xmarre/ComfyUI-ScaleLockedResidualDiffusion/blob/main/README.md"
|
||||
"Bug Tracker" = "https://github.com/xmarre/ComfyUI-ScaleLockedResidualDiffusion/issues"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "xmarre"
|
||||
DisplayName = "ComfyUI-ScaleLockedResidualDiffusion"
|
||||
Icon = ""
|
||||
includes = []
|
||||
+644
-80
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
from typing import Iterable, Optional
|
||||
from typing import Iterable, Optional, Sequence
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -155,7 +155,7 @@ def build_nested_noise(
|
||||
hf = hf - hf_low_up
|
||||
|
||||
std = hf.std(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
hf = hf / std
|
||||
hf = (hf / std).to(device=device, non_blocking=True)
|
||||
|
||||
out = base + float(hf_strength) * hf
|
||||
out_std = out.std(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
@@ -210,14 +210,347 @@ def residual_lock_multiband(
|
||||
return base_high + torch.lerp(base_low, anchor_low, low_strength) + torch.lerp(base_mid, anchor_mid, mid_strength)
|
||||
|
||||
|
||||
def schedule_value(start: float, end: float, progress: float, mode: str) -> float:
|
||||
progress = float(max(0.0, min(1.0, progress)))
|
||||
if mode == "cosine":
|
||||
t = 0.5 - 0.5 * math.cos(math.pi * progress)
|
||||
elif mode == "flat":
|
||||
t = 0.0
|
||||
def _base_grid(batch: int, height: int, width: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
|
||||
yy, xx = torch.meshgrid(
|
||||
torch.linspace(-1.0, 1.0, height, device=device, dtype=dtype),
|
||||
torch.linspace(-1.0, 1.0, width, device=device, dtype=dtype),
|
||||
indexing="ij",
|
||||
)
|
||||
return torch.stack([xx, yy], dim=-1).unsqueeze(0).expand(batch, -1, -1, -1).contiguous()
|
||||
|
||||
|
||||
def warp_4d_tensor(x: torch.Tensor, flow_xy_norm: torch.Tensor, mode: str = "bilinear") -> torch.Tensor:
|
||||
if x.ndim != 4:
|
||||
raise ValueError(f"Expected BCHW tensor, got shape {tuple(x.shape)}")
|
||||
batch, _, height, width = x.shape
|
||||
expected = (batch, height, width, 2)
|
||||
if tuple(flow_xy_norm.shape) != expected:
|
||||
raise ValueError(f"Expected flow shape {expected}, got {tuple(flow_xy_norm.shape)}")
|
||||
base_grid = _base_grid(batch, height, width, x.device, x.dtype)
|
||||
sample_grid = (base_grid + flow_xy_norm).clamp(-1.25, 1.25)
|
||||
return F.grid_sample(
|
||||
x,
|
||||
sample_grid,
|
||||
mode=mode,
|
||||
padding_mode="border",
|
||||
align_corners=True,
|
||||
)
|
||||
|
||||
|
||||
def _expand_mask_channels(mask: torch.Tensor, like: torch.Tensor) -> torch.Tensor:
|
||||
if mask.ndim != 4:
|
||||
raise ValueError(f"Expected mask tensor BCHW, got shape {tuple(mask.shape)}")
|
||||
if mask.shape[0] < like.shape[0]:
|
||||
repeat = math.ceil(like.shape[0] / max(1, mask.shape[0]))
|
||||
mask = mask.repeat(repeat, 1, 1, 1)[: like.shape[0]]
|
||||
elif mask.shape[0] > like.shape[0]:
|
||||
mask = mask[: like.shape[0]]
|
||||
if mask.shape[1] == like.shape[1]:
|
||||
return mask
|
||||
if mask.shape[1] == 1:
|
||||
return mask.expand(like.shape[0], like.shape[1], like.shape[-2], like.shape[-1]).contiguous()
|
||||
return mask.mean(dim=1, keepdim=True).expand(like.shape[0], like.shape[1], like.shape[-2], like.shape[-1]).contiguous()
|
||||
|
||||
|
||||
def _latent_activity_map(x: torch.Tensor, mask_1ch: torch.Tensor) -> torch.Tensor:
|
||||
mass = mask_1ch.sum(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
mean = (x * mask_1ch).sum(dim=(-2, -1), keepdim=True) / mass
|
||||
centered = x - mean
|
||||
activity = centered.square().mean(dim=1, keepdim=True)
|
||||
activity = activity * mask_1ch
|
||||
activity = activity + mask_1ch * 1e-8
|
||||
return activity
|
||||
|
||||
|
||||
def _masked_spatial_stats(x: torch.Tensor, mask_1ch: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
batch, _, height, width = x.shape
|
||||
yy, xx = torch.meshgrid(
|
||||
torch.linspace(-1.0, 1.0, height, device=x.device, dtype=x.dtype),
|
||||
torch.linspace(-1.0, 1.0, width, device=x.device, dtype=x.dtype),
|
||||
indexing="ij",
|
||||
)
|
||||
xx = xx.view(1, 1, height, width).expand(batch, -1, -1, -1)
|
||||
yy = yy.view(1, 1, height, width).expand(batch, -1, -1, -1)
|
||||
|
||||
activity = _latent_activity_map(x, mask_1ch)
|
||||
mass = activity.sum(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
cx = (activity * xx).sum(dim=(-2, -1), keepdim=True) / mass
|
||||
cy = (activity * yy).sum(dim=(-2, -1), keepdim=True) / mass
|
||||
|
||||
dx = xx - cx
|
||||
dy = yy - cy
|
||||
rx = torch.sqrt((activity * dx.square()).sum(dim=(-2, -1), keepdim=True) / mass).clamp_min(1e-4)
|
||||
ry = torch.sqrt((activity * dy.square()).sum(dim=(-2, -1), keepdim=True) / mass).clamp_min(1e-4)
|
||||
return cx, cy, rx, ry
|
||||
|
||||
|
||||
def estimate_latent_compaction_flow(
|
||||
anchor_low: torch.Tensor,
|
||||
base_low: torch.Tensor,
|
||||
mask_1ch: torch.Tensor,
|
||||
strength: float,
|
||||
radial_strength: float,
|
||||
anisotropy: float,
|
||||
translation_strength: float,
|
||||
max_shift_px: float,
|
||||
) -> torch.Tensor:
|
||||
batch, _, height, width = base_low.shape
|
||||
anchor_cx, anchor_cy, anchor_rx, anchor_ry = _masked_spatial_stats(anchor_low, mask_1ch)
|
||||
base_cx, base_cy, base_rx, base_ry = _masked_spatial_stats(base_low, mask_1ch)
|
||||
|
||||
ratio_x = (base_rx / anchor_rx.clamp_min(1e-5)).clamp(0.85, 1.25)
|
||||
ratio_y = (base_ry / anchor_ry.clamp_min(1e-5)).clamp(0.85, 1.25)
|
||||
outward_x = (ratio_x - 1.0).clamp(min=0.0)
|
||||
outward_y = (ratio_y - 1.0).clamp(min=0.0)
|
||||
|
||||
yy, xx = torch.meshgrid(
|
||||
torch.linspace(-1.0, 1.0, height, device=base_low.device, dtype=base_low.dtype),
|
||||
torch.linspace(-1.0, 1.0, width, device=base_low.device, dtype=base_low.dtype),
|
||||
indexing="ij",
|
||||
)
|
||||
xx = xx.view(1, 1, height, width).expand(batch, -1, -1, -1)
|
||||
yy = yy.view(1, 1, height, width).expand(batch, -1, -1, -1)
|
||||
|
||||
dx = xx - base_cx
|
||||
dy = yy - base_cy
|
||||
ex = dx / base_rx.clamp_min(1e-5)
|
||||
ey = dy / base_ry.clamp_min(1e-5)
|
||||
radius = torch.sqrt(ex.square() + ey.square() + 1e-8)
|
||||
edge_envelope = torch.clamp(radius / 1.25, 0.0, 1.0)
|
||||
|
||||
smooth_mask = lowpass_latent(mask_1ch, 0.35).clamp(0.0, 1.0)
|
||||
mean_outward = 0.5 * (outward_x + outward_y)
|
||||
axis_x = mean_outward * float(radial_strength) + (outward_x - mean_outward) * float(anisotropy)
|
||||
axis_y = mean_outward * float(radial_strength) + (outward_y - mean_outward) * float(anisotropy)
|
||||
|
||||
shift_x = ex * edge_envelope * smooth_mask * float(strength) * axis_x
|
||||
shift_y = ey * edge_envelope * smooth_mask * float(strength) * axis_y
|
||||
|
||||
trans_x = (base_cx - anchor_cx) * smooth_mask * float(strength) * float(translation_strength)
|
||||
trans_y = (base_cy - anchor_cy) * smooth_mask * float(strength) * float(translation_strength)
|
||||
|
||||
max_shift_norm_x = 2.0 * float(max_shift_px) / max(width - 1, 1)
|
||||
max_shift_norm_y = 2.0 * float(max_shift_px) / max(height - 1, 1)
|
||||
|
||||
shift_x = (shift_x + trans_x).clamp(-max_shift_norm_x, max_shift_norm_x)
|
||||
shift_y = (shift_y + trans_y).clamp(-max_shift_norm_y, max_shift_norm_y)
|
||||
return torch.stack([shift_x[:, 0], shift_y[:, 0]], dim=-1)
|
||||
|
||||
|
||||
def _weighted_channel_stats(x: torch.Tensor, mask_1ch: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
mass = mask_1ch.sum(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
mean = (x * mask_1ch).sum(dim=(-2, -1), keepdim=True) / mass
|
||||
var = (((x - mean).square()) * mask_1ch).sum(dim=(-2, -1), keepdim=True) / mass
|
||||
std = torch.sqrt(var.clamp_min(1e-8))
|
||||
return mean, std
|
||||
|
||||
|
||||
def _weighted_global_std(x: torch.Tensor, mask_1ch: torch.Tensor, mean: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
if mean is None:
|
||||
mean, _ = _weighted_channel_stats(x, mask_1ch)
|
||||
mass = mask_1ch.sum(dim=(-2, -1), keepdim=True).clamp_min(1e-6)
|
||||
denom = mass * float(max(1, x.shape[1]))
|
||||
var = (((x - mean).square()) * mask_1ch).sum(dim=(1, 2, 3), keepdim=True) / denom
|
||||
return torch.sqrt(var.clamp_min(1e-8))
|
||||
|
||||
|
||||
def _bounded_ratio(target: torch.Tensor, source: torch.Tensor, gain_cap: float) -> torch.Tensor:
|
||||
gain_cap = float(max(1.0, gain_cap))
|
||||
lo = 1.0 / gain_cap
|
||||
hi = gain_cap
|
||||
return (target / source.clamp_min(1e-5)).clamp(lo, hi)
|
||||
|
||||
|
||||
def tether_latent_low_frequency_energy(
|
||||
base_low: torch.Tensor,
|
||||
anchor_low: torch.Tensor,
|
||||
mask_1ch: torch.Tensor,
|
||||
energy_tether: float,
|
||||
channel_tether: float,
|
||||
gain_cap: float,
|
||||
) -> torch.Tensor:
|
||||
energy_tether = float(max(0.0, min(1.0, energy_tether)))
|
||||
channel_tether = float(max(0.0, min(1.0, channel_tether)))
|
||||
if energy_tether <= 0.0 and channel_tether <= 0.0:
|
||||
return base_low
|
||||
|
||||
base_mean, _ = _weighted_channel_stats(base_low, mask_1ch)
|
||||
anchor_mean, anchor_std = _weighted_channel_stats(anchor_low, mask_1ch)
|
||||
|
||||
regulated = base_low
|
||||
if energy_tether > 0.0:
|
||||
base_global_std = _weighted_global_std(regulated, mask_1ch, mean=base_mean)
|
||||
anchor_global_std = _weighted_global_std(anchor_low, mask_1ch, mean=anchor_mean)
|
||||
global_gain = _bounded_ratio(anchor_global_std, base_global_std, gain_cap)
|
||||
global_gain = torch.lerp(torch.ones_like(global_gain), global_gain, energy_tether)
|
||||
regulated = base_mean + (regulated - base_mean) * global_gain
|
||||
|
||||
if channel_tether > 0.0:
|
||||
regulated_mean, regulated_std = _weighted_channel_stats(regulated, mask_1ch)
|
||||
channel_gain = _bounded_ratio(anchor_std, regulated_std, gain_cap)
|
||||
channel_gain = torch.lerp(torch.ones_like(channel_gain), channel_gain, channel_tether)
|
||||
regulated = regulated_mean + (regulated - regulated_mean) * channel_gain
|
||||
|
||||
return regulated
|
||||
|
||||
|
||||
def restore_latent_low_frequency_stats(
|
||||
warped_low: torch.Tensor,
|
||||
anchor_low: torch.Tensor,
|
||||
mask_1ch: torch.Tensor,
|
||||
anchor_mix: float,
|
||||
mean_anchor_mix: float,
|
||||
contrast_restore: float,
|
||||
) -> torch.Tensor:
|
||||
warped_mean, warped_std = _weighted_channel_stats(warped_low, mask_1ch)
|
||||
anchor_mean, anchor_std = _weighted_channel_stats(anchor_low, mask_1ch)
|
||||
|
||||
mean_target = torch.lerp(warped_mean, anchor_mean, float(mean_anchor_mix))
|
||||
contrast_gain = torch.lerp(
|
||||
torch.ones_like(anchor_std),
|
||||
anchor_std / warped_std.clamp_min(1e-5),
|
||||
float(contrast_restore),
|
||||
)
|
||||
restored = mean_target + (warped_low - warped_mean) * contrast_gain
|
||||
return torch.lerp(restored, anchor_low, float(anchor_mix))
|
||||
|
||||
|
||||
def latent_manifold_compand(
|
||||
base_denoised: torch.Tensor,
|
||||
anchor_denoised: torch.Tensor,
|
||||
mask: Optional[torch.Tensor],
|
||||
strength: float,
|
||||
cutoff: float,
|
||||
radial_strength: float,
|
||||
anisotropy: float,
|
||||
translation_strength: float,
|
||||
anchor_mix: float,
|
||||
mean_anchor_mix: float,
|
||||
contrast_restore: float,
|
||||
energy_tether: float,
|
||||
channel_tether: float,
|
||||
energy_gain_cap: float,
|
||||
max_shift_px: float,
|
||||
) -> torch.Tensor:
|
||||
strength = float(max(0.0, min(1.0, strength)))
|
||||
if strength <= 0.0:
|
||||
return base_denoised
|
||||
|
||||
cutoff = float(max(0.05, min(1.0, cutoff)))
|
||||
work_dtype = base_denoised.dtype
|
||||
compute_dtype = torch.float32
|
||||
|
||||
base = base_denoised.to(dtype=compute_dtype)
|
||||
anchor = anchor_denoised.to(device=base.device, dtype=compute_dtype)
|
||||
if tuple(anchor.shape[-2:]) != tuple(base.shape[-2:]):
|
||||
anchor = resize_4d_tensor(anchor, tuple(base.shape[-2:]))
|
||||
|
||||
if mask is None:
|
||||
mask_1ch = torch.ones((base.shape[0], 1, base.shape[-2], base.shape[-1]), device=base.device, dtype=compute_dtype)
|
||||
else:
|
||||
t = progress
|
||||
if mask.ndim == 3:
|
||||
mask = mask.unsqueeze(1)
|
||||
elif mask.ndim == 4 and mask.shape[1] != 1:
|
||||
mask = mask.mean(dim=1, keepdim=True)
|
||||
mask_1ch = _expand_mask_channels(mask.to(device=base.device, dtype=compute_dtype), base)[:, :1].clamp(0.0, 1.0)
|
||||
|
||||
base_low = lowpass_latent(base, cutoff)
|
||||
anchor_low = lowpass_latent(anchor, cutoff)
|
||||
base_high = base - base_low
|
||||
|
||||
regulated_low = tether_latent_low_frequency_energy(
|
||||
base_low,
|
||||
anchor_low,
|
||||
mask_1ch,
|
||||
energy_tether=float(max(0.0, min(1.0, energy_tether))) * strength,
|
||||
channel_tether=float(max(0.0, min(1.0, channel_tether))) * strength,
|
||||
gain_cap=float(max(1.0, energy_gain_cap)),
|
||||
)
|
||||
|
||||
flow = estimate_latent_compaction_flow(
|
||||
anchor_low=anchor_low,
|
||||
base_low=regulated_low,
|
||||
mask_1ch=mask_1ch,
|
||||
strength=strength,
|
||||
radial_strength=radial_strength,
|
||||
anisotropy=anisotropy,
|
||||
translation_strength=translation_strength,
|
||||
max_shift_px=max_shift_px,
|
||||
)
|
||||
warped_low = warp_4d_tensor(regulated_low, flow, mode="bilinear")
|
||||
restored_low = restore_latent_low_frequency_stats(
|
||||
warped_low,
|
||||
anchor_low,
|
||||
mask_1ch,
|
||||
anchor_mix=float(max(0.0, min(1.0, anchor_mix))) * strength,
|
||||
mean_anchor_mix=float(max(0.0, min(1.0, mean_anchor_mix))) * strength,
|
||||
contrast_restore=float(max(0.0, min(1.0, contrast_restore))) * strength,
|
||||
)
|
||||
corrected = base_high + restored_low
|
||||
return corrected.to(dtype=work_dtype)
|
||||
|
||||
|
||||
def clamp(value: float, lo: float, hi: float) -> float:
|
||||
return float(max(lo, min(hi, value)))
|
||||
|
||||
|
||||
def sigma_progress(planner_sigmas: Optional[Sequence[float]], idx: int) -> Optional[float]:
|
||||
if planner_sigmas is None:
|
||||
return None
|
||||
|
||||
if len(planner_sigmas) == 0:
|
||||
return None
|
||||
|
||||
idx = max(0, min(int(idx), len(planner_sigmas) - 1))
|
||||
sigma_hi = max(float(planner_sigmas[0]), 1e-6)
|
||||
sigma_lo = max(float(planner_sigmas[-1]), 1e-6)
|
||||
sigma = max(float(planner_sigmas[idx]), 1e-6)
|
||||
|
||||
hi = math.log(sigma_hi)
|
||||
lo = math.log(sigma_lo)
|
||||
denom = hi - lo
|
||||
if abs(denom) < 1e-6:
|
||||
return None
|
||||
|
||||
cur = math.log(sigma)
|
||||
return clamp((hi - cur) / denom, 0.0, 1.0)
|
||||
|
||||
|
||||
def schedule_curve(progress: float, mode: str, power: float = 2.0, hold: float = 0.0) -> float:
|
||||
p = clamp(progress, 0.0, 1.0)
|
||||
power = max(1e-6, float(power))
|
||||
hold = clamp(hold, 0.0, 0.95)
|
||||
|
||||
if mode == "flat":
|
||||
return 0.0
|
||||
if mode == "linear":
|
||||
return p
|
||||
if mode == "cosine":
|
||||
return 0.5 - 0.5 * math.cos(math.pi * p)
|
||||
if mode == "smoothstep":
|
||||
return p * p * (3.0 - 2.0 * p)
|
||||
if mode == "smootherstep":
|
||||
return p * p * p * (p * (p * 6.0 - 15.0) + 10.0)
|
||||
if mode == "ease_in":
|
||||
return p ** power
|
||||
if mode == "ease_out":
|
||||
return 1.0 - (1.0 - p) ** power
|
||||
if mode == "ease_in_out":
|
||||
if p < 0.5:
|
||||
return 0.5 * ((2.0 * p) ** power)
|
||||
return 1.0 - 0.5 * ((2.0 * (1.0 - p)) ** power)
|
||||
if mode == "hold_then_drop":
|
||||
if p <= hold:
|
||||
return 0.0
|
||||
u = (p - hold) / max(1e-6, 1.0 - hold)
|
||||
return u ** power
|
||||
if mode == "fast_drop":
|
||||
return p ** (1.0 / power)
|
||||
return p
|
||||
|
||||
|
||||
def schedule_value(start: float, end: float, progress: float, mode: str, power: float = 2.0, hold: float = 0.0) -> float:
|
||||
t = schedule_curve(progress, mode, power=power, hold=hold)
|
||||
return float(start + (end - start) * t)
|
||||
|
||||
|
||||
@@ -239,7 +572,7 @@ class TrajectoryRecorder:
|
||||
self.xt_steps.append(_stage_cpu_tensor(x, self.store_dtype, self.pin_memory))
|
||||
|
||||
|
||||
class ScaleLockedCFGGuider(torch.nn.Module):
|
||||
class ScaleLockedCFGGuider:
|
||||
"""
|
||||
A custom ComfyUI guider that applies the scale lock in denoised-latent space.
|
||||
|
||||
@@ -261,86 +594,317 @@ class ScaleLockedCFGGuider(torch.nn.Module):
|
||||
mid_cutoff: float,
|
||||
mid_strength: float,
|
||||
schedule: str,
|
||||
schedule_power: float,
|
||||
schedule_hold: float,
|
||||
mid_strength_start: float,
|
||||
mid_strength_end: float,
|
||||
mid_schedule: str,
|
||||
mid_schedule_power: float,
|
||||
mid_schedule_hold: float,
|
||||
spatial_mask: Optional[torch.Tensor],
|
||||
manifold_enabled: bool = False,
|
||||
manifold_strength: float = 0.0,
|
||||
manifold_strength_start: float = 1.0,
|
||||
manifold_strength_end: float = 0.0,
|
||||
manifold_schedule: str = "ease_out",
|
||||
manifold_schedule_power: float = 2.0,
|
||||
manifold_schedule_hold: float = 0.0,
|
||||
manifold_cutoff: float = 0.18,
|
||||
manifold_radial_strength: float = 1.0,
|
||||
manifold_anisotropy: float = 0.15,
|
||||
manifold_translation_strength: float = 1.0,
|
||||
manifold_anchor_mix: float = 0.18,
|
||||
manifold_mean_anchor_mix: float = 0.12,
|
||||
manifold_contrast_restore: float = 0.10,
|
||||
manifold_energy_tether: float = 0.0,
|
||||
manifold_channel_tether: float = 0.0,
|
||||
manifold_energy_gain_cap: float = 1.75,
|
||||
manifold_max_shift_px: float = 3.0,
|
||||
manifold_spatial_mask: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
self._slrd_model = model
|
||||
self._slrd_anchors_x0_cpu = list(anchors_x0_cpu)
|
||||
self._slrd_planner_sigmas = [float(x) for x in planner_sigmas] if planner_sigmas is not None else None
|
||||
self._slrd_lock_strength = float(lock_strength)
|
||||
self._slrd_lock_strength_start = float(lock_strength_start)
|
||||
self._slrd_lock_strength_end = float(lock_strength_end)
|
||||
self._slrd_cutoff = float(cutoff)
|
||||
self._slrd_mid_cutoff = float(max(cutoff, mid_cutoff))
|
||||
self._slrd_mid_strength = float(mid_strength)
|
||||
self._slrd_schedule = schedule
|
||||
self._slrd_seen_sigmas: list[float] = []
|
||||
self._slrd_prev_match_idx: int = 0
|
||||
self._slrd_spatial_mask = spatial_mask
|
||||
self._slrd_last_sigma: Optional[float] = None
|
||||
init_scale_lock_state(
|
||||
self,
|
||||
model=model,
|
||||
anchors_x0_cpu=anchors_x0_cpu,
|
||||
planner_sigmas=planner_sigmas,
|
||||
lock_strength=lock_strength,
|
||||
lock_strength_start=lock_strength_start,
|
||||
lock_strength_end=lock_strength_end,
|
||||
cutoff=cutoff,
|
||||
mid_cutoff=mid_cutoff,
|
||||
mid_strength=mid_strength,
|
||||
schedule=schedule,
|
||||
schedule_power=schedule_power,
|
||||
schedule_hold=schedule_hold,
|
||||
mid_strength_start=mid_strength_start,
|
||||
mid_strength_end=mid_strength_end,
|
||||
mid_schedule=mid_schedule,
|
||||
mid_schedule_power=mid_schedule_power,
|
||||
mid_schedule_hold=mid_schedule_hold,
|
||||
spatial_mask=spatial_mask,
|
||||
manifold_enabled=manifold_enabled,
|
||||
manifold_strength=manifold_strength,
|
||||
manifold_strength_start=manifold_strength_start,
|
||||
manifold_strength_end=manifold_strength_end,
|
||||
manifold_schedule=manifold_schedule,
|
||||
manifold_schedule_power=manifold_schedule_power,
|
||||
manifold_schedule_hold=manifold_schedule_hold,
|
||||
manifold_cutoff=manifold_cutoff,
|
||||
manifold_radial_strength=manifold_radial_strength,
|
||||
manifold_anisotropy=manifold_anisotropy,
|
||||
manifold_translation_strength=manifold_translation_strength,
|
||||
manifold_anchor_mix=manifold_anchor_mix,
|
||||
manifold_mean_anchor_mix=manifold_mean_anchor_mix,
|
||||
manifold_contrast_restore=manifold_contrast_restore,
|
||||
manifold_energy_tether=manifold_energy_tether,
|
||||
manifold_channel_tether=manifold_channel_tether,
|
||||
manifold_energy_gain_cap=manifold_energy_gain_cap,
|
||||
manifold_max_shift_px=manifold_max_shift_px,
|
||||
manifold_spatial_mask=manifold_spatial_mask,
|
||||
)
|
||||
|
||||
def _slrd_resolve_step_index(self, timestep: torch.Tensor | float | int) -> int:
|
||||
sigma = _sigma_scalar(timestep)
|
||||
|
||||
if self._slrd_planner_sigmas:
|
||||
best_idx = 0
|
||||
best_dist = float("inf")
|
||||
for i, planner_sigma in enumerate(self._slrd_planner_sigmas):
|
||||
dist = abs(planner_sigma - sigma)
|
||||
if dist < best_dist:
|
||||
best_idx = i
|
||||
best_dist = dist
|
||||
best_idx = max(best_idx, self._slrd_prev_match_idx)
|
||||
best_idx = min(best_idx, len(self._slrd_anchors_x0_cpu) - 1)
|
||||
self._slrd_prev_match_idx = best_idx
|
||||
return best_idx
|
||||
|
||||
if self._slrd_last_sigma is None:
|
||||
self._slrd_last_sigma = sigma
|
||||
step_index = 0
|
||||
self._slrd_seen_sigmas = [sigma]
|
||||
else:
|
||||
tol = 1e-6 * max(1.0, abs(self._slrd_last_sigma), abs(sigma))
|
||||
if abs(sigma - self._slrd_last_sigma) > tol:
|
||||
self._slrd_last_sigma = sigma
|
||||
self._slrd_seen_sigmas.append(sigma)
|
||||
unique_sigmas = []
|
||||
for seen_sigma in self._slrd_seen_sigmas:
|
||||
if all(
|
||||
abs(seen_sigma - unique_sigma) > (1e-6 * max(1.0, abs(seen_sigma), abs(unique_sigma)))
|
||||
for unique_sigma in unique_sigmas
|
||||
):
|
||||
unique_sigmas.append(seen_sigma)
|
||||
step_index = max(0, len(unique_sigmas) - 1)
|
||||
step_index = max(step_index, self._slrd_prev_match_idx)
|
||||
self._slrd_prev_match_idx = step_index
|
||||
if not self._slrd_anchors_x0_cpu:
|
||||
return 0
|
||||
return min(step_index, len(self._slrd_anchors_x0_cpu) - 1)
|
||||
return resolve_scale_lock_step_index(self, timestep)
|
||||
|
||||
def _slrd_strength_for_step(self, idx: int) -> float:
|
||||
total = max(1, len(self._slrd_anchors_x0_cpu) - 1)
|
||||
progress = idx / total
|
||||
scheduled = schedule_value(self._slrd_lock_strength_start, self._slrd_lock_strength_end, progress, self._slrd_schedule)
|
||||
return float(max(0.0, min(1.0, self._slrd_lock_strength * scheduled)))
|
||||
return scale_lock_strength_for_step(self, idx)
|
||||
|
||||
def _slrd_anchor_for(self, idx: int, like: torch.Tensor) -> torch.Tensor:
|
||||
anchor = self._slrd_anchors_x0_cpu[idx].to(device=like.device, dtype=like.dtype, non_blocking=True)
|
||||
if tuple(anchor.shape[-2:]) != tuple(like.shape[-2:]):
|
||||
anchor = resize_4d_tensor(anchor, tuple(like.shape[-2:]))
|
||||
return anchor
|
||||
return scale_lock_anchor_for(self, idx, like)
|
||||
|
||||
def _slrd_mask_for(self, like: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
if self._slrd_spatial_mask is None:
|
||||
return None
|
||||
mask = self._slrd_spatial_mask.to(device=like.device, dtype=like.dtype, non_blocking=True)
|
||||
if tuple(mask.shape[-2:]) != tuple(like.shape[-2:]):
|
||||
mask = resize_4d_tensor(mask, tuple(like.shape[-2:]))
|
||||
if mask.shape[0] < like.shape[0]:
|
||||
repeat = math.ceil(like.shape[0] / max(1, mask.shape[0]))
|
||||
mask = mask.repeat(repeat, 1, 1, 1)[: like.shape[0]]
|
||||
elif mask.shape[0] > like.shape[0]:
|
||||
mask = mask[: like.shape[0]]
|
||||
return mask
|
||||
return scale_lock_mask_for(self, like)
|
||||
|
||||
def _slrd_manifold_strength_for_step(self, idx: int) -> float:
|
||||
return scale_lock_manifold_strength_for_step(self, idx)
|
||||
|
||||
def _slrd_manifold_mask_for(self, like: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
return scale_lock_manifold_mask_for(self, like)
|
||||
|
||||
|
||||
def init_scale_lock_state(
|
||||
guider,
|
||||
*,
|
||||
model,
|
||||
anchors_x0_cpu: Iterable[torch.Tensor],
|
||||
planner_sigmas: Optional[Iterable[float]],
|
||||
lock_strength: float,
|
||||
lock_strength_start: float,
|
||||
lock_strength_end: float,
|
||||
cutoff: float,
|
||||
mid_cutoff: float,
|
||||
mid_strength: float,
|
||||
schedule: str,
|
||||
schedule_power: float,
|
||||
schedule_hold: float,
|
||||
mid_strength_start: float,
|
||||
mid_strength_end: float,
|
||||
mid_schedule: str,
|
||||
mid_schedule_power: float,
|
||||
mid_schedule_hold: float,
|
||||
spatial_mask: Optional[torch.Tensor],
|
||||
manifold_enabled: bool = False,
|
||||
manifold_strength: float = 0.0,
|
||||
manifold_strength_start: float = 1.0,
|
||||
manifold_strength_end: float = 0.0,
|
||||
manifold_schedule: str = "ease_out",
|
||||
manifold_schedule_power: float = 2.0,
|
||||
manifold_schedule_hold: float = 0.0,
|
||||
manifold_cutoff: float = 0.18,
|
||||
manifold_radial_strength: float = 1.0,
|
||||
manifold_anisotropy: float = 0.15,
|
||||
manifold_translation_strength: float = 1.0,
|
||||
manifold_anchor_mix: float = 0.18,
|
||||
manifold_mean_anchor_mix: float = 0.12,
|
||||
manifold_contrast_restore: float = 0.10,
|
||||
manifold_energy_tether: float = 0.0,
|
||||
manifold_channel_tether: float = 0.0,
|
||||
manifold_energy_gain_cap: float = 1.75,
|
||||
manifold_max_shift_px: float = 3.0,
|
||||
manifold_spatial_mask: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
guider._slrd_model = model
|
||||
guider._slrd_anchors_x0_cpu = list(anchors_x0_cpu)
|
||||
guider._slrd_planner_sigmas = [float(x) for x in planner_sigmas] if planner_sigmas is not None else None
|
||||
guider._slrd_lock_strength = float(lock_strength)
|
||||
guider._slrd_lock_strength_start = float(lock_strength_start)
|
||||
guider._slrd_lock_strength_end = float(lock_strength_end)
|
||||
guider._slrd_cutoff = float(cutoff)
|
||||
guider._slrd_mid_cutoff = float(max(cutoff, mid_cutoff))
|
||||
guider._slrd_mid_strength = float(mid_strength)
|
||||
guider._slrd_schedule = schedule
|
||||
guider._slrd_schedule_power = float(schedule_power)
|
||||
guider._slrd_schedule_hold = float(schedule_hold)
|
||||
guider._slrd_mid_strength_start = float(mid_strength_start)
|
||||
guider._slrd_mid_strength_end = float(mid_strength_end)
|
||||
guider._slrd_mid_schedule = mid_schedule
|
||||
guider._slrd_mid_schedule_power = float(mid_schedule_power)
|
||||
guider._slrd_mid_schedule_hold = float(mid_schedule_hold)
|
||||
guider._slrd_seen_sigmas = []
|
||||
guider._slrd_prev_match_idx = 0
|
||||
guider._slrd_spatial_mask = spatial_mask
|
||||
guider._slrd_last_sigma = None
|
||||
guider._slrd_manifold_enabled = bool(manifold_enabled)
|
||||
guider._slrd_manifold_strength = float(manifold_strength)
|
||||
guider._slrd_manifold_strength_start = float(manifold_strength_start)
|
||||
guider._slrd_manifold_strength_end = float(manifold_strength_end)
|
||||
guider._slrd_manifold_schedule = manifold_schedule
|
||||
guider._slrd_manifold_schedule_power = float(manifold_schedule_power)
|
||||
guider._slrd_manifold_schedule_hold = float(manifold_schedule_hold)
|
||||
guider._slrd_manifold_cutoff = float(max(0.05, min(1.0, manifold_cutoff)))
|
||||
guider._slrd_manifold_radial_strength = float(manifold_radial_strength)
|
||||
guider._slrd_manifold_anisotropy = float(manifold_anisotropy)
|
||||
guider._slrd_manifold_translation_strength = float(manifold_translation_strength)
|
||||
guider._slrd_manifold_anchor_mix = float(manifold_anchor_mix)
|
||||
guider._slrd_manifold_mean_anchor_mix = float(manifold_mean_anchor_mix)
|
||||
guider._slrd_manifold_contrast_restore = float(manifold_contrast_restore)
|
||||
guider._slrd_manifold_energy_tether = float(max(0.0, min(1.0, manifold_energy_tether)))
|
||||
guider._slrd_manifold_channel_tether = float(max(0.0, min(1.0, manifold_channel_tether)))
|
||||
guider._slrd_manifold_energy_gain_cap = float(max(1.0, manifold_energy_gain_cap))
|
||||
guider._slrd_manifold_max_shift_px = float(manifold_max_shift_px)
|
||||
guider._slrd_manifold_spatial_mask = manifold_spatial_mask if manifold_spatial_mask is not None else spatial_mask
|
||||
|
||||
|
||||
def resolve_scale_lock_step_index(guider, timestep: torch.Tensor | float | int) -> int:
|
||||
sigma = _sigma_scalar(timestep)
|
||||
|
||||
if guider._slrd_planner_sigmas:
|
||||
best_idx = 0
|
||||
best_dist = float("inf")
|
||||
for i, planner_sigma in enumerate(guider._slrd_planner_sigmas):
|
||||
dist = abs(planner_sigma - sigma)
|
||||
if dist < best_dist:
|
||||
best_idx = i
|
||||
best_dist = dist
|
||||
best_idx = max(best_idx, guider._slrd_prev_match_idx)
|
||||
best_idx = min(best_idx, len(guider._slrd_anchors_x0_cpu) - 1)
|
||||
guider._slrd_prev_match_idx = best_idx
|
||||
return best_idx
|
||||
|
||||
if guider._slrd_last_sigma is None:
|
||||
guider._slrd_last_sigma = sigma
|
||||
step_index = 0
|
||||
guider._slrd_seen_sigmas = [sigma]
|
||||
else:
|
||||
tol = 1e-6 * max(1.0, abs(guider._slrd_last_sigma), abs(sigma))
|
||||
if abs(sigma - guider._slrd_last_sigma) > tol:
|
||||
guider._slrd_last_sigma = sigma
|
||||
guider._slrd_seen_sigmas.append(sigma)
|
||||
unique_sigmas = []
|
||||
for seen_sigma in guider._slrd_seen_sigmas:
|
||||
if all(
|
||||
abs(seen_sigma - unique_sigma) > (1e-6 * max(1.0, abs(seen_sigma), abs(unique_sigma)))
|
||||
for unique_sigma in unique_sigmas
|
||||
):
|
||||
unique_sigmas.append(seen_sigma)
|
||||
step_index = max(0, len(unique_sigmas) - 1)
|
||||
step_index = max(step_index, guider._slrd_prev_match_idx)
|
||||
guider._slrd_prev_match_idx = step_index
|
||||
if not guider._slrd_anchors_x0_cpu:
|
||||
return 0
|
||||
return min(step_index, len(guider._slrd_anchors_x0_cpu) - 1)
|
||||
|
||||
|
||||
def scale_lock_strength_for_step(guider, idx: int) -> float:
|
||||
return scale_lock_strengths_for_step(guider, idx)[0]
|
||||
|
||||
|
||||
def scale_lock_progress_for_step(guider, idx: int) -> float:
|
||||
sigma_based = sigma_progress(getattr(guider, "_slrd_planner_sigmas", None), idx)
|
||||
if sigma_based is not None:
|
||||
return sigma_based
|
||||
total = max(1, len(guider._slrd_anchors_x0_cpu) - 1)
|
||||
return clamp(idx / total, 0.0, 1.0)
|
||||
|
||||
|
||||
def _scheduled_strength(
|
||||
base_strength: float,
|
||||
start: float,
|
||||
end: float,
|
||||
progress: float,
|
||||
mode: str,
|
||||
power: float,
|
||||
hold: float,
|
||||
) -> float:
|
||||
scheduled = schedule_value(start, end, progress, mode, power=power, hold=hold)
|
||||
return clamp(base_strength * scheduled, 0.0, 1.0)
|
||||
|
||||
|
||||
def scale_lock_strengths_for_step(guider, idx: int) -> tuple[float, float]:
|
||||
progress = scale_lock_progress_for_step(guider, idx)
|
||||
low_strength = _scheduled_strength(
|
||||
guider._slrd_lock_strength,
|
||||
guider._slrd_lock_strength_start,
|
||||
guider._slrd_lock_strength_end,
|
||||
progress,
|
||||
guider._slrd_schedule,
|
||||
guider._slrd_schedule_power,
|
||||
guider._slrd_schedule_hold,
|
||||
)
|
||||
if getattr(guider, "_slrd_mid_schedule", "linked") == "linked":
|
||||
mid_strength = clamp(low_strength * guider._slrd_mid_strength, 0.0, 1.0)
|
||||
else:
|
||||
mid_strength = _scheduled_strength(
|
||||
guider._slrd_mid_strength,
|
||||
guider._slrd_mid_strength_start,
|
||||
guider._slrd_mid_strength_end,
|
||||
progress,
|
||||
guider._slrd_mid_schedule,
|
||||
guider._slrd_mid_schedule_power,
|
||||
guider._slrd_mid_schedule_hold,
|
||||
)
|
||||
return low_strength, mid_strength
|
||||
|
||||
|
||||
def scale_lock_manifold_strength_for_step(guider, idx: int) -> float:
|
||||
if not getattr(guider, "_slrd_manifold_enabled", False):
|
||||
return 0.0
|
||||
progress = scale_lock_progress_for_step(guider, idx)
|
||||
return _scheduled_strength(
|
||||
guider._slrd_manifold_strength,
|
||||
guider._slrd_manifold_strength_start,
|
||||
guider._slrd_manifold_strength_end,
|
||||
progress,
|
||||
guider._slrd_manifold_schedule,
|
||||
guider._slrd_manifold_schedule_power,
|
||||
guider._slrd_manifold_schedule_hold,
|
||||
)
|
||||
|
||||
|
||||
def scale_lock_anchor_for(guider, idx: int, like: torch.Tensor) -> torch.Tensor:
|
||||
anchor = guider._slrd_anchors_x0_cpu[idx].to(device=like.device, dtype=like.dtype, non_blocking=True)
|
||||
if tuple(anchor.shape[-2:]) != tuple(like.shape[-2:]):
|
||||
anchor = resize_4d_tensor(anchor, tuple(like.shape[-2:]))
|
||||
return anchor
|
||||
|
||||
|
||||
def scale_lock_mask_for(guider, like: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
if guider._slrd_spatial_mask is None:
|
||||
return None
|
||||
mask = guider._slrd_spatial_mask.to(device=like.device, dtype=like.dtype, non_blocking=True)
|
||||
if tuple(mask.shape[-2:]) != tuple(like.shape[-2:]):
|
||||
mask = resize_4d_tensor(mask, tuple(like.shape[-2:]))
|
||||
if mask.shape[0] < like.shape[0]:
|
||||
repeat = math.ceil(like.shape[0] / max(1, mask.shape[0]))
|
||||
mask = mask.repeat(repeat, 1, 1, 1)[: like.shape[0]]
|
||||
elif mask.shape[0] > like.shape[0]:
|
||||
mask = mask[: like.shape[0]]
|
||||
return mask
|
||||
|
||||
|
||||
def scale_lock_manifold_mask_for(guider, like: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
manifold_mask = getattr(guider, "_slrd_manifold_spatial_mask", None)
|
||||
if manifold_mask is None:
|
||||
return None
|
||||
mask = manifold_mask.to(device=like.device, dtype=like.dtype, non_blocking=True)
|
||||
if tuple(mask.shape[-2:]) != tuple(like.shape[-2:]):
|
||||
mask = resize_4d_tensor(mask, tuple(like.shape[-2:]))
|
||||
if mask.shape[0] < like.shape[0]:
|
||||
repeat = math.ceil(like.shape[0] / max(1, mask.shape[0]))
|
||||
mask = mask.repeat(repeat, 1, 1, 1)[: like.shape[0]]
|
||||
elif mask.shape[0] > like.shape[0]:
|
||||
mask = mask[: like.shape[0]]
|
||||
return mask
|
||||
|
||||
|
||||
|
||||
|
||||
+784
@@ -0,0 +1,784 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
import types
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
import torch
|
||||
|
||||
import comfy.model_management
|
||||
import comfy.sample
|
||||
import comfy.samplers
|
||||
import comfy.utils
|
||||
import latent_preview
|
||||
|
||||
from .slrd_core import (
|
||||
TrajectoryRecorder,
|
||||
build_nested_noise,
|
||||
clone_latent,
|
||||
init_scale_lock_state,
|
||||
latent_manifold_compand,
|
||||
latent_target_hw_from_megapixels,
|
||||
resolve_scale_lock_step_index,
|
||||
resize_latent_dict,
|
||||
resize_mask,
|
||||
residual_lock_multiband,
|
||||
scale_lock_anchor_for,
|
||||
scale_lock_manifold_mask_for,
|
||||
scale_lock_manifold_strength_for_step,
|
||||
scale_lock_mask_for,
|
||||
scale_lock_strengths_for_step,
|
||||
)
|
||||
|
||||
|
||||
_LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
|
||||
class _NullPreviewCallback:
|
||||
def __call__(self, step, x0, x, total_steps):
|
||||
del step, x0, x, total_steps
|
||||
|
||||
|
||||
_CONSERVATIVE_SAFE_SAMPLERS = {
|
||||
"ddim",
|
||||
"euler",
|
||||
"euler_cfg_pp",
|
||||
"heun",
|
||||
"lcm",
|
||||
"dpmpp_2m",
|
||||
"dpmpp_2m_cfg_pp",
|
||||
}
|
||||
|
||||
LOCK_SCHEDULE_OPTIONS = [
|
||||
"linear",
|
||||
"cosine",
|
||||
"flat",
|
||||
"smoothstep",
|
||||
"smootherstep",
|
||||
"ease_in",
|
||||
"ease_out",
|
||||
"ease_in_out",
|
||||
"hold_then_drop",
|
||||
"fast_drop",
|
||||
]
|
||||
|
||||
MID_SCHEDULE_OPTIONS = ["linked", *LOCK_SCHEDULE_OPTIONS]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScaleLockConfig:
|
||||
lock_strength: float
|
||||
lock_strength_start: float
|
||||
lock_strength_end: float
|
||||
cutoff: float
|
||||
mid_cutoff: float
|
||||
mid_strength: float
|
||||
schedule: str
|
||||
schedule_power: float
|
||||
schedule_hold: float
|
||||
mid_strength_start: float
|
||||
mid_strength_end: float
|
||||
mid_schedule: str
|
||||
mid_schedule_power: float
|
||||
mid_schedule_hold: float
|
||||
spatial_mask: Optional[torch.Tensor] = None
|
||||
manifold_enabled: bool = False
|
||||
manifold_strength: float = 0.0
|
||||
manifold_strength_start: float = 1.0
|
||||
manifold_strength_end: float = 0.0
|
||||
manifold_schedule: str = "ease_out"
|
||||
manifold_schedule_power: float = 2.0
|
||||
manifold_schedule_hold: float = 0.0
|
||||
manifold_cutoff: float = 0.18
|
||||
manifold_radial_strength: float = 1.0
|
||||
manifold_anisotropy: float = 0.15
|
||||
manifold_translation_strength: float = 1.0
|
||||
manifold_anchor_mix: float = 0.18
|
||||
manifold_mean_anchor_mix: float = 0.12
|
||||
manifold_contrast_restore: float = 0.10
|
||||
manifold_energy_tether: float = 0.0
|
||||
manifold_channel_tether: float = 0.0
|
||||
manifold_energy_gain_cap: float = 1.75
|
||||
manifold_max_shift_px: float = 3.0
|
||||
manifold_spatial_mask: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScaleLockedRuntimeContext:
|
||||
model: Any
|
||||
highres_latent: dict
|
||||
lowres_latent: dict
|
||||
lowres_out: dict
|
||||
sigmas: torch.Tensor
|
||||
anchors_x0: list[torch.Tensor]
|
||||
planner_sigmas: list[float]
|
||||
highres_noise: torch.Tensor
|
||||
noise_seed: int
|
||||
|
||||
def prepared_noise(self) -> "ScaleLockedPreparedNoise":
|
||||
return ScaleLockedPreparedNoise(self.highres_noise, self.noise_seed)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScaleLockedSampleResult:
|
||||
output: dict
|
||||
lowres_planner: dict
|
||||
denoised_output: dict
|
||||
|
||||
|
||||
class ScaleLockedPreparedNoise:
|
||||
def __init__(self, noise_tensor: torch.Tensor, seed: int):
|
||||
self._noise_tensor = noise_tensor.detach().to(device="cpu").contiguous().clone()
|
||||
self.seed = int(seed)
|
||||
|
||||
def generate_noise(self, latent: dict) -> torch.Tensor:
|
||||
latent_samples = latent["samples"]
|
||||
if tuple(latent_samples.shape) != tuple(self._noise_tensor.shape):
|
||||
raise ValueError(
|
||||
"ScaleLockedPreparedNoise expected latent shape "
|
||||
f"{tuple(self._noise_tensor.shape)} but received {tuple(latent_samples.shape)}."
|
||||
)
|
||||
return self._noise_tensor.clone()
|
||||
|
||||
|
||||
def _normalize_sampler_name(name: str) -> str:
|
||||
return str(name).strip().lower()
|
||||
|
||||
|
||||
def guard_sampler_alignment(sampler_name: str, mode: str) -> None:
|
||||
mode = str(mode).strip().lower()
|
||||
if mode == "off":
|
||||
return
|
||||
|
||||
normalized = _normalize_sampler_name(sampler_name)
|
||||
if normalized in _CONSERVATIVE_SAFE_SAMPLERS:
|
||||
return
|
||||
|
||||
msg = (
|
||||
"ScaleLockedResidualKSampler: sampler "
|
||||
f"'{sampler_name}' is outside the conservative SLRD alignment-safe allowlist. "
|
||||
"The node will still work in many cases, but planner/high-res anchor matching is less trustworthy "
|
||||
"for samplers with more complex internal evaluation patterns."
|
||||
)
|
||||
if mode == "error":
|
||||
raise ValueError(msg)
|
||||
_LOGGER.warning(msg)
|
||||
|
||||
|
||||
def compat_sampler_names():
|
||||
return getattr(comfy.samplers, "SAMPLER_NAMES", comfy.samplers.KSampler.SAMPLERS)
|
||||
|
||||
|
||||
def compat_scheduler_names():
|
||||
return getattr(comfy.samplers, "SCHEDULER_NAMES", comfy.samplers.KSampler.SCHEDULERS)
|
||||
|
||||
|
||||
def clean_latent(latent: dict) -> dict:
|
||||
out = clone_latent(latent)
|
||||
out.pop("downscale_ratio_spacial", None)
|
||||
return out
|
||||
|
||||
|
||||
def _model_sampling_obj(model):
|
||||
if hasattr(model, "get_model_object"):
|
||||
return model.get_model_object("model_sampling")
|
||||
if hasattr(model, "model") and hasattr(model.model, "model_sampling"):
|
||||
return model.model.model_sampling
|
||||
raise AttributeError("Unable to resolve model_sampling from the ComfyUI model patcher.")
|
||||
|
||||
|
||||
def calculate_sigmas(model, scheduler: str, steps: int, denoise: float) -> torch.Tensor:
|
||||
total_steps = int(steps)
|
||||
if denoise < 1.0:
|
||||
if denoise <= 0.0:
|
||||
return torch.FloatTensor([])
|
||||
total_steps = int(steps / denoise)
|
||||
sigmas = comfy.samplers.calculate_sigmas(_model_sampling_obj(model), scheduler, total_steps).cpu()
|
||||
if denoise < 1.0:
|
||||
sigmas = sigmas[-(steps + 1) :]
|
||||
return sigmas
|
||||
|
||||
|
||||
def prepare_noise(latent_samples: torch.Tensor, seed: int, batch_inds=None, disable_noise: bool = False) -> torch.Tensor:
|
||||
if disable_noise:
|
||||
return torch.zeros(latent_samples.size(), dtype=latent_samples.dtype, layout=latent_samples.layout, device="cpu")
|
||||
return comfy.sample.prepare_noise(latent_samples, seed, batch_inds)
|
||||
|
||||
|
||||
def _resolve_sampler_device(model_or_wrap, fallback: torch.device) -> torch.device:
|
||||
for obj in (
|
||||
model_or_wrap,
|
||||
getattr(model_or_wrap, "inner_model", None),
|
||||
getattr(model_or_wrap, "model", None),
|
||||
getattr(model_or_wrap, "model_patcher", None),
|
||||
):
|
||||
if obj is None:
|
||||
continue
|
||||
device = getattr(obj, "load_device", None)
|
||||
if device is not None:
|
||||
return device
|
||||
return fallback
|
||||
|
||||
|
||||
def fix_latent_channels(model, latent_dict: dict) -> dict:
|
||||
out = clone_latent(latent_dict)
|
||||
ratio = out.get("downscale_ratio_spacial", None)
|
||||
out["samples"] = comfy.sample.fix_empty_latent_channels(model, out["samples"], ratio)
|
||||
return out
|
||||
|
||||
|
||||
def make_lowres_latent(latent: dict, target_megapixels: float) -> dict:
|
||||
low_hw = latent_target_hw_from_megapixels(latent["samples"], target_megapixels)
|
||||
return resize_latent_dict(latent, low_hw)
|
||||
|
||||
|
||||
def _store_dtype_for(x: torch.Tensor) -> torch.dtype:
|
||||
if x.dtype in (torch.float32, torch.float16, torch.bfloat16):
|
||||
return x.dtype
|
||||
return torch.float32
|
||||
|
||||
|
||||
def prepare_spatial_lock_mask(mask: Optional[torch.Tensor], latent_samples: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
if mask is None:
|
||||
return None
|
||||
return resize_mask(mask, tuple(latent_samples.shape[-2:]), latent_samples.shape[0], latent_samples.shape[1])
|
||||
|
||||
|
||||
def make_preview_callback(model, steps: int, x0_output: dict):
|
||||
try:
|
||||
return latent_preview.prepare_callback(model, max(0, steps), x0_output)
|
||||
except Exception:
|
||||
return _NullPreviewCallback()
|
||||
|
||||
|
||||
def _planner_sigmas_for_recorded_steps(sigmas: torch.Tensor, recorded_steps: int) -> list[float]:
|
||||
if recorded_steps <= 0:
|
||||
return []
|
||||
|
||||
sigma_values = [float(v) for v in sigmas.detach().flatten().cpu().tolist()]
|
||||
if not sigma_values:
|
||||
return []
|
||||
|
||||
visible_sigmas = sigma_values[:-1] if len(sigma_values) > 1 else sigma_values
|
||||
if not visible_sigmas:
|
||||
visible_sigmas = sigma_values
|
||||
|
||||
if recorded_steps <= len(visible_sigmas):
|
||||
return visible_sigmas[:recorded_steps]
|
||||
return visible_sigmas + [visible_sigmas[-1]] * (recorded_steps - len(visible_sigmas))
|
||||
|
||||
|
||||
def noise_seed(noise) -> int:
|
||||
try:
|
||||
return int(getattr(noise, "seed", 0))
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def generate_noise_for_latent(noise, latent: dict) -> torch.Tensor:
|
||||
generated = noise.generate_noise(latent)
|
||||
if not isinstance(generated, torch.Tensor):
|
||||
raise TypeError("ScaleLockedResidualSamplerCustomAdvanced currently supports tensor noise only.")
|
||||
return generated
|
||||
|
||||
|
||||
def _apply_manifold_compand_to_noise_prediction(guider, working_noise: torch.Tensor, anchor: torch.Tensor, idx: int) -> torch.Tensor:
|
||||
manifold_strength = scale_lock_manifold_strength_for_step(guider, idx)
|
||||
if manifold_strength <= 0.0:
|
||||
return working_noise
|
||||
|
||||
mask = scale_lock_manifold_mask_for(guider, working_noise)
|
||||
corrected = latent_manifold_compand(
|
||||
working_noise,
|
||||
anchor,
|
||||
mask=mask,
|
||||
strength=manifold_strength,
|
||||
cutoff=guider._slrd_manifold_cutoff,
|
||||
radial_strength=guider._slrd_manifold_radial_strength,
|
||||
anisotropy=guider._slrd_manifold_anisotropy,
|
||||
translation_strength=guider._slrd_manifold_translation_strength,
|
||||
anchor_mix=guider._slrd_manifold_anchor_mix,
|
||||
mean_anchor_mix=guider._slrd_manifold_mean_anchor_mix,
|
||||
contrast_restore=guider._slrd_manifold_contrast_restore,
|
||||
energy_tether=guider._slrd_manifold_energy_tether,
|
||||
channel_tether=guider._slrd_manifold_channel_tether,
|
||||
energy_gain_cap=guider._slrd_manifold_energy_gain_cap,
|
||||
max_shift_px=guider._slrd_manifold_max_shift_px,
|
||||
)
|
||||
if mask is not None:
|
||||
corrected = working_noise + mask * (corrected - working_noise)
|
||||
return corrected
|
||||
|
||||
|
||||
def apply_scale_lock_to_noise_prediction(guider, base_noise: torch.Tensor, x, timestep):
|
||||
del x
|
||||
if not getattr(guider, "_slrd_anchors_x0_cpu", None):
|
||||
return base_noise
|
||||
|
||||
idx = resolve_scale_lock_step_index(guider, timestep)
|
||||
anchor = scale_lock_anchor_for(guider, idx, base_noise)
|
||||
|
||||
corrected = base_noise
|
||||
low_strength, mid_strength = scale_lock_strengths_for_step(guider, idx)
|
||||
if low_strength > 0.0 or mid_strength > 0.0:
|
||||
corrected = residual_lock_multiband(
|
||||
corrected,
|
||||
anchor,
|
||||
low_strength=low_strength,
|
||||
mid_strength=mid_strength,
|
||||
low_cutoff=guider._slrd_cutoff,
|
||||
mid_cutoff=guider._slrd_mid_cutoff,
|
||||
)
|
||||
|
||||
mask = scale_lock_mask_for(guider, corrected)
|
||||
if mask is not None:
|
||||
corrected = base_noise + mask * (corrected - base_noise)
|
||||
|
||||
corrected = _apply_manifold_compand_to_noise_prediction(guider, corrected, anchor, idx)
|
||||
return corrected
|
||||
|
||||
|
||||
def create_cfg_guider(model, positive, negative, cfg):
|
||||
guider = comfy.samplers.CFGGuider(model)
|
||||
guider.set_conds(positive, negative)
|
||||
guider.set_cfg(cfg)
|
||||
return guider
|
||||
|
||||
|
||||
def clone_guider_for_scale_lock(guider):
|
||||
cloned = copy.copy(guider)
|
||||
if hasattr(guider, "__dict__"):
|
||||
cloned.__dict__ = dict(guider.__dict__)
|
||||
|
||||
original_predict_noise = getattr(cloned, "_slrd_original_predict_noise", None)
|
||||
if original_predict_noise is not None:
|
||||
if isinstance(original_predict_noise, types.MethodType):
|
||||
cloned.predict_noise = types.MethodType(original_predict_noise.__func__, cloned)
|
||||
else:
|
||||
cloned.predict_noise = original_predict_noise
|
||||
|
||||
stale_keys = [key for key in getattr(cloned, "__dict__", {}) if key.startswith("_slrd_")]
|
||||
for key in stale_keys:
|
||||
delattr(cloned, key)
|
||||
|
||||
return cloned
|
||||
|
||||
|
||||
def apply_scale_lock_to_guider(guider, runtime: ScaleLockedRuntimeContext, config: ScaleLockConfig):
|
||||
spatial_mask = prepare_spatial_lock_mask(config.spatial_mask, runtime.highres_latent["samples"])
|
||||
manifold_spatial_mask = prepare_spatial_lock_mask(
|
||||
config.manifold_spatial_mask if config.manifold_spatial_mask is not None else config.spatial_mask,
|
||||
runtime.highres_latent["samples"],
|
||||
)
|
||||
init_scale_lock_state(
|
||||
guider,
|
||||
model=runtime.model,
|
||||
anchors_x0_cpu=runtime.anchors_x0,
|
||||
planner_sigmas=runtime.planner_sigmas,
|
||||
lock_strength=config.lock_strength,
|
||||
lock_strength_start=config.lock_strength_start,
|
||||
lock_strength_end=config.lock_strength_end,
|
||||
cutoff=config.cutoff,
|
||||
mid_cutoff=config.mid_cutoff,
|
||||
mid_strength=config.mid_strength,
|
||||
schedule=config.schedule,
|
||||
schedule_power=config.schedule_power,
|
||||
schedule_hold=config.schedule_hold,
|
||||
mid_strength_start=config.mid_strength_start,
|
||||
mid_strength_end=config.mid_strength_end,
|
||||
mid_schedule=config.mid_schedule,
|
||||
mid_schedule_power=config.mid_schedule_power,
|
||||
mid_schedule_hold=config.mid_schedule_hold,
|
||||
spatial_mask=spatial_mask,
|
||||
manifold_enabled=config.manifold_enabled,
|
||||
manifold_strength=config.manifold_strength,
|
||||
manifold_strength_start=config.manifold_strength_start,
|
||||
manifold_strength_end=config.manifold_strength_end,
|
||||
manifold_schedule=config.manifold_schedule,
|
||||
manifold_schedule_power=config.manifold_schedule_power,
|
||||
manifold_schedule_hold=config.manifold_schedule_hold,
|
||||
manifold_cutoff=config.manifold_cutoff,
|
||||
manifold_radial_strength=config.manifold_radial_strength,
|
||||
manifold_anisotropy=config.manifold_anisotropy,
|
||||
manifold_translation_strength=config.manifold_translation_strength,
|
||||
manifold_anchor_mix=config.manifold_anchor_mix,
|
||||
manifold_mean_anchor_mix=config.manifold_mean_anchor_mix,
|
||||
manifold_contrast_restore=config.manifold_contrast_restore,
|
||||
manifold_energy_tether=config.manifold_energy_tether,
|
||||
manifold_channel_tether=config.manifold_channel_tether,
|
||||
manifold_energy_gain_cap=config.manifold_energy_gain_cap,
|
||||
manifold_max_shift_px=config.manifold_max_shift_px,
|
||||
manifold_spatial_mask=manifold_spatial_mask,
|
||||
)
|
||||
|
||||
original_predict_noise = getattr(guider, "_slrd_original_predict_noise", guider.predict_noise)
|
||||
guider._slrd_original_predict_noise = original_predict_noise
|
||||
|
||||
def _wrapped_predict_noise(self, x, timestep, model_options=None, seed=None):
|
||||
if model_options is None:
|
||||
model_options = {}
|
||||
base_noise = self._slrd_original_predict_noise(x, timestep, model_options=model_options, seed=seed)
|
||||
return apply_scale_lock_to_noise_prediction(self, base_noise, x, timestep)
|
||||
|
||||
guider.predict_noise = types.MethodType(_wrapped_predict_noise, guider)
|
||||
return guider
|
||||
|
||||
|
||||
def restore_original_predict_noise(guider) -> None:
|
||||
original_predict_noise = getattr(guider, "_slrd_original_predict_noise", None)
|
||||
if original_predict_noise is not None:
|
||||
guider.predict_noise = original_predict_noise
|
||||
|
||||
|
||||
def _run_lowres_planner_advanced(
|
||||
*,
|
||||
guider,
|
||||
sampler,
|
||||
sigmas: torch.Tensor,
|
||||
lowres_latent: dict,
|
||||
noise,
|
||||
pin_anchors: bool,
|
||||
) -> tuple[dict, list[torch.Tensor], list[float], torch.Tensor]:
|
||||
model = guider.model_patcher
|
||||
lowres_latent = fix_latent_channels(model, lowres_latent)
|
||||
latent_samples = lowres_latent["samples"]
|
||||
target_device = _resolve_sampler_device(model, latent_samples.device)
|
||||
target_dtype = latent_samples.dtype
|
||||
latent_samples = latent_samples.to(
|
||||
device=target_device,
|
||||
dtype=target_dtype,
|
||||
non_blocking=True,
|
||||
)
|
||||
lowres_latent["samples"] = latent_samples
|
||||
noise_tensor = generate_noise_for_latent(noise, lowres_latent).to(
|
||||
device=target_device,
|
||||
dtype=target_dtype,
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
noise_mask = lowres_latent.get("noise_mask", None)
|
||||
if isinstance(noise_mask, torch.Tensor):
|
||||
noise_mask = noise_mask.to(device=target_device, non_blocking=True)
|
||||
sigmas = sigmas.to(device=target_device, non_blocking=True)
|
||||
recorder = TrajectoryRecorder(
|
||||
store_dtype=_store_dtype_for(latent_samples),
|
||||
capture_noisy_xt=False,
|
||||
pin_memory=pin_anchors,
|
||||
)
|
||||
|
||||
samples = guider.sample(
|
||||
noise_tensor,
|
||||
latent_samples,
|
||||
sampler,
|
||||
sigmas,
|
||||
denoise_mask=noise_mask,
|
||||
callback=recorder.callback,
|
||||
disable_pbar=True,
|
||||
seed=noise_seed(noise),
|
||||
)
|
||||
samples = samples.to(comfy.model_management.intermediate_device())
|
||||
|
||||
out = clone_latent(lowres_latent)
|
||||
out.pop("downscale_ratio_spacial", None)
|
||||
out["samples"] = samples
|
||||
|
||||
planner_sigmas = _planner_sigmas_for_recorded_steps(sigmas, len(recorder.x0_steps))
|
||||
recorder.step_sigmas = planner_sigmas
|
||||
return out, recorder.x0_steps, planner_sigmas, noise_tensor
|
||||
|
||||
|
||||
def _run_lowres_planner(
|
||||
*,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
sigmas: torch.Tensor,
|
||||
lowres_latent: dict,
|
||||
seed: int,
|
||||
disable_noise: bool,
|
||||
pin_anchors: bool,
|
||||
) -> tuple[dict, list[torch.Tensor], list[float], torch.Tensor]:
|
||||
sampler_obj = comfy.samplers.sampler_object(sampler_name)
|
||||
guider = comfy.samplers.CFGGuider(model)
|
||||
guider.set_conds(positive, negative)
|
||||
guider.set_cfg(cfg)
|
||||
|
||||
lowres_latent = fix_latent_channels(model, lowres_latent)
|
||||
latent_samples = lowres_latent["samples"]
|
||||
target_device = _resolve_sampler_device(model, latent_samples.device)
|
||||
target_dtype = latent_samples.dtype
|
||||
latent_samples = latent_samples.to(
|
||||
device=target_device,
|
||||
dtype=target_dtype,
|
||||
non_blocking=True,
|
||||
)
|
||||
lowres_latent["samples"] = latent_samples
|
||||
batch_inds = lowres_latent.get("batch_index", None)
|
||||
planner_noise = prepare_noise(
|
||||
latent_samples,
|
||||
seed=seed,
|
||||
batch_inds=batch_inds,
|
||||
disable_noise=disable_noise,
|
||||
).to(
|
||||
device=target_device,
|
||||
dtype=target_dtype,
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
noise_mask = lowres_latent.get("noise_mask", None)
|
||||
if isinstance(noise_mask, torch.Tensor):
|
||||
noise_mask = noise_mask.to(device=target_device, non_blocking=True)
|
||||
sigmas = sigmas.to(device=target_device, non_blocking=True)
|
||||
recorder = TrajectoryRecorder(
|
||||
store_dtype=_store_dtype_for(latent_samples),
|
||||
capture_noisy_xt=False,
|
||||
pin_memory=pin_anchors,
|
||||
)
|
||||
|
||||
samples = guider.sample(
|
||||
planner_noise,
|
||||
latent_samples,
|
||||
sampler_obj,
|
||||
sigmas,
|
||||
denoise_mask=noise_mask,
|
||||
callback=recorder.callback,
|
||||
disable_pbar=True,
|
||||
seed=seed,
|
||||
)
|
||||
samples = samples.to(comfy.model_management.intermediate_device())
|
||||
|
||||
out = clone_latent(lowres_latent)
|
||||
out.pop("downscale_ratio_spacial", None)
|
||||
out["samples"] = samples
|
||||
|
||||
planner_sigmas = _planner_sigmas_for_recorded_steps(sigmas, len(recorder.x0_steps))
|
||||
recorder.step_sigmas = planner_sigmas
|
||||
return out, recorder.x0_steps, planner_sigmas, planner_noise
|
||||
|
||||
|
||||
def _build_highres_noise(highres_latent: dict, lowres_noise: torch.Tensor, seed: int, hf_strength: float):
|
||||
highres_samples = highres_latent["samples"]
|
||||
target_device = highres_samples.device
|
||||
target_dtype = highres_samples.dtype
|
||||
if torch.count_nonzero(lowres_noise).item() == 0:
|
||||
return torch.zeros(
|
||||
highres_samples.size(),
|
||||
dtype=target_dtype,
|
||||
layout=highres_samples.layout,
|
||||
device=target_device,
|
||||
)
|
||||
|
||||
return build_nested_noise(
|
||||
lowres_noise=lowres_noise,
|
||||
target_shape=tuple(highres_samples.shape),
|
||||
seed=seed,
|
||||
hf_strength=hf_strength,
|
||||
).to(device=target_device, dtype=target_dtype, non_blocking=True)
|
||||
|
||||
|
||||
def build_runtime_context_from_advanced(
|
||||
*,
|
||||
noise,
|
||||
guider,
|
||||
sampler,
|
||||
sigmas: torch.Tensor,
|
||||
latent_image,
|
||||
target_megapixels: float,
|
||||
nested_noise_strength: float,
|
||||
pin_anchors: bool,
|
||||
) -> ScaleLockedRuntimeContext:
|
||||
model = guider.model_patcher
|
||||
highres_latent = fix_latent_channels(model, latent_image)
|
||||
target_device = _resolve_sampler_device(model, highres_latent["samples"].device)
|
||||
target_dtype = highres_latent["samples"].dtype
|
||||
highres_latent["samples"] = highres_latent["samples"].to(
|
||||
device=target_device,
|
||||
dtype=target_dtype,
|
||||
non_blocking=True,
|
||||
)
|
||||
if isinstance(highres_latent.get("noise_mask"), torch.Tensor):
|
||||
highres_latent["noise_mask"] = highres_latent["noise_mask"].to(
|
||||
device=target_device,
|
||||
non_blocking=True,
|
||||
)
|
||||
lowres_latent = make_lowres_latent(highres_latent, target_megapixels)
|
||||
|
||||
if sigmas.numel() == 0:
|
||||
return ScaleLockedRuntimeContext(
|
||||
model=model,
|
||||
highres_latent=highres_latent,
|
||||
lowres_latent=lowres_latent,
|
||||
lowres_out=clean_latent(lowres_latent),
|
||||
sigmas=sigmas.to(device=target_device, non_blocking=True),
|
||||
anchors_x0=[],
|
||||
planner_sigmas=[],
|
||||
highres_noise=torch.zeros_like(
|
||||
highres_latent["samples"],
|
||||
device=target_device,
|
||||
),
|
||||
noise_seed=noise_seed(noise),
|
||||
)
|
||||
|
||||
lowres_out, anchors_x0, planner_sigmas, lowres_noise = _run_lowres_planner_advanced(
|
||||
guider=guider,
|
||||
sampler=sampler,
|
||||
sigmas=sigmas,
|
||||
lowres_latent=lowres_latent,
|
||||
noise=noise,
|
||||
pin_anchors=pin_anchors,
|
||||
)
|
||||
if len(anchors_x0) == 0:
|
||||
raise RuntimeError("ScaleLockedResidualSamplerCustomAdvanced: planner pass did not record any x0 anchors.")
|
||||
|
||||
highres_noise = _build_highres_noise(
|
||||
highres_latent=highres_latent,
|
||||
lowres_noise=lowres_noise,
|
||||
seed=noise_seed(noise) ^ 0x9E3779B97F4A7C15,
|
||||
hf_strength=nested_noise_strength,
|
||||
)
|
||||
return ScaleLockedRuntimeContext(
|
||||
model=model,
|
||||
highres_latent=highres_latent,
|
||||
lowres_latent=lowres_latent,
|
||||
lowres_out=lowres_out,
|
||||
sigmas=sigmas.to(device=target_device, non_blocking=True),
|
||||
anchors_x0=anchors_x0,
|
||||
planner_sigmas=planner_sigmas,
|
||||
highres_noise=highres_noise,
|
||||
noise_seed=noise_seed(noise),
|
||||
)
|
||||
|
||||
|
||||
def sample_with_runtime(
|
||||
*,
|
||||
guider,
|
||||
sampler,
|
||||
runtime: ScaleLockedRuntimeContext,
|
||||
config: ScaleLockConfig,
|
||||
restore_after: bool = True,
|
||||
) -> ScaleLockedSampleResult:
|
||||
if runtime.sigmas.numel() == 0:
|
||||
out = clean_latent(runtime.highres_latent)
|
||||
return ScaleLockedSampleResult(output=out, lowres_planner=clean_latent(runtime.lowres_out), denoised_output=out)
|
||||
|
||||
apply_scale_lock_to_guider(guider, runtime, config)
|
||||
|
||||
x0_output = {}
|
||||
callback = make_preview_callback(runtime.model, len(runtime.sigmas) - 1, x0_output)
|
||||
noise_mask = runtime.highres_latent.get("noise_mask", None)
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
|
||||
try:
|
||||
samples = guider.sample(
|
||||
runtime.highres_noise,
|
||||
runtime.highres_latent["samples"],
|
||||
sampler,
|
||||
runtime.sigmas,
|
||||
denoise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=runtime.noise_seed,
|
||||
)
|
||||
finally:
|
||||
if restore_after:
|
||||
restore_original_predict_noise(guider)
|
||||
samples = samples.to(comfy.model_management.intermediate_device())
|
||||
|
||||
out = clean_latent(runtime.highres_latent)
|
||||
out["samples"] = samples
|
||||
|
||||
if "x0" in x0_output:
|
||||
try:
|
||||
x0_out = runtime.model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
except Exception:
|
||||
x0_out = x0_output["x0"].detach().cpu()
|
||||
denoised = clean_latent(runtime.highres_latent)
|
||||
denoised["samples"] = x0_out
|
||||
else:
|
||||
denoised = out
|
||||
|
||||
return ScaleLockedSampleResult(
|
||||
output=out,
|
||||
lowres_planner=clean_latent(runtime.lowres_out),
|
||||
denoised_output=denoised,
|
||||
)
|
||||
|
||||
|
||||
def run_scale_locked_ksampler(
|
||||
*,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
target_megapixels,
|
||||
nested_noise_strength,
|
||||
add_noise,
|
||||
pin_anchors,
|
||||
sampler_guard,
|
||||
config: ScaleLockConfig,
|
||||
) -> ScaleLockedSampleResult:
|
||||
if steps < 1:
|
||||
raise ValueError("steps must be >= 1")
|
||||
|
||||
highres_latent = fix_latent_channels(model, latent_image)
|
||||
lowres_latent = make_lowres_latent(highres_latent, target_megapixels)
|
||||
if denoise <= 0.0:
|
||||
out = clean_latent(highres_latent)
|
||||
return ScaleLockedSampleResult(output=out, lowres_planner=clean_latent(lowres_latent), denoised_output=out)
|
||||
|
||||
guard_sampler_alignment(sampler_name, sampler_guard)
|
||||
sigmas = calculate_sigmas(model, scheduler=scheduler, steps=steps, denoise=denoise)
|
||||
if sigmas.numel() == 0:
|
||||
out = clean_latent(highres_latent)
|
||||
return ScaleLockedSampleResult(output=out, lowres_planner=clean_latent(lowres_latent), denoised_output=out)
|
||||
|
||||
disable_noise = not bool(add_noise)
|
||||
planner_seed = int(seed)
|
||||
detail_seed = int(seed) ^ 0x9E3779B97F4A7C15
|
||||
|
||||
lowres_out, anchors_x0, planner_sigmas, lowres_noise = _run_lowres_planner(
|
||||
model=model,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
sigmas=sigmas,
|
||||
lowres_latent=lowres_latent,
|
||||
seed=planner_seed,
|
||||
disable_noise=disable_noise,
|
||||
pin_anchors=pin_anchors,
|
||||
)
|
||||
if len(anchors_x0) == 0:
|
||||
raise RuntimeError("ScaleLockedResidualKSampler: planner pass did not record any x0 anchors.")
|
||||
|
||||
runtime = ScaleLockedRuntimeContext(
|
||||
model=model,
|
||||
highres_latent=highres_latent,
|
||||
lowres_latent=lowres_latent,
|
||||
lowres_out=lowres_out,
|
||||
sigmas=sigmas,
|
||||
anchors_x0=anchors_x0,
|
||||
planner_sigmas=planner_sigmas,
|
||||
highres_noise=_build_highres_noise(
|
||||
highres_latent=highres_latent,
|
||||
lowres_noise=lowres_noise,
|
||||
seed=detail_seed,
|
||||
hf_strength=nested_noise_strength,
|
||||
),
|
||||
noise_seed=detail_seed,
|
||||
)
|
||||
guider = create_cfg_guider(model, positive, negative, cfg)
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
return sample_with_runtime(guider=guider, sampler=sampler, runtime=runtime, config=config)
|
||||
|
||||
Reference in New Issue
Block a user