Compare commits

...
Author SHA1 Message Date
xmarre 6153443986 Merge pull request #13 from xmarre/codex/fix-detailer-mask-for-slrd
Fix detailer hook to use effective guider semantics
2026-03-24 15:24:39 +01:00
xmarre 9f1870e350 Use effective guider in detailer hook sampling 2026-03-24 15:15:39 +01:00
xmarre 77b238a672 fix: align detailer hook noise and bump version 2026-03-24 14:16:01 +01:00
xmarre bf6c517a4c Bump version to 1.0.31 2026-03-24 13:52:09 +01:00
xmarre 0f076a886f Fix detailer hook mask support propagation 2026-03-24 13:51:36 +01:00
xmarre b166cd8741 Remove implicit detailer mask from SLRD config 2026-03-23 05:09:06 +01:00
xmarre 751b8e771c Bump version to 1.0.29 2026-03-23 04:03:51 +01:00
xmarre 3b5a6b7971 Preserve incoming noise for masked detailer runs 2026-03-23 04:03:25 +01:00
xmarre a8e94be4c4 Align impact masks and bump version 2026-03-23 01:49:14 +01:00
xmarre a2b713a9ea Bump version to 1.0.27 2026-03-23 01:08:07 +01:00
xmarre 0af4e4194a Restore denoise mask for impact sampler 2026-03-23 01:03:36 +01:00
xmarre dae4b0906e Fix detailer mask handling and bump version 2026-03-23 00:19:17 +01:00
xmarre e41982cdab Bump version to 1.0.25 2026-03-22 23:34:13 +01:00
xmarre e04f34897a Isolate detailer guider and preserve planner model_options 2026-03-22 23:33:28 +01:00
xmarre 132609e3d2 Bump version to 1.0.24 2026-03-22 22:40:36 +01:00
xmarre 0d4e8d20e2 Binarize detailer masks for SLRD support 2026-03-22 22:39:20 +01:00
xmarre a31f9be2d0 Bump version to 1.0.23 2026-03-22 21:16:46 +01:00
xmarre 25267e588a Merge pull request #12 from xmarre/codex/align-slrd-mask-with-sampler
Fix detailer SLRD masks to use the runtime latent mask
2026-03-22 21:15:35 +01:00
xmarre 2d4fe5ea47 Use runtime detailer masks for scale-locked sampling 2026-03-22 21:05:37 +01:00
xmarre d015307e1a Merge pull request #10 from xmarre/codex/fix-facedetailer-latent-mask-mixup
Restore proxy sampling path for SLRD Impact detailer
2026-03-15 14:46:30 +01:00
xmarre db6dfb5949 Restore Impact detailer proxy path 2026-03-15 14:43:40 +01:00
xmarre 6667f62b75 Merge branch 'main' into codex/fix-facedetailer-latent-mask-mixup 2026-03-15 14:34:22 +01:00
xmarre 189afe2e45 Merge pull request #11 from xmarre/coderabbitai/docstrings/420c8a4
📝 Add docstrings to `codex/fix-facedetailer-latent-mask-mixup`
2026-03-15 14:33:07 +01:00
coderabbitai[bot] 47aa1ddfbe 📝 Add docstrings to codex/fix-facedetailer-latent-mask-mixup
Docstrings generation was requested by @xmarre.

* https://github.com/xmarre/ComfyUI-ScaleLockedResidualDiffusion/pull/10#issuecomment-4062975146

The following files were modified:

* `nodes.py`
2026-03-15 13:29:48 +00:00
xmarre e933378b43 Keep planner on CFG guider 2026-03-15 14:19:04 +01:00
xmarre 420c8a490c Fix FaceDetailer mask state 2026-03-15 12:53:57 +01:00
xmarre 1ffbb6fa56 Fix FaceDetailer mask state 2026-03-15 12:53:19 +01:00
xmarre 3d64e1d3b1 Investigate mask-state bug in FaceD 2026-03-15 12:23:39 +01:00
xmarre 1b3c11ff35 Merge branch 'main' of https://github.com/xmarre/ComfyUI-ScaleLockedResidualDiffusion 2026-03-15 11:40:12 +01:00
xmarre 08d3c6adc0 Bump pyproject version 2026-03-15 11:39:32 +01:00
xmarre a948909964 Merge pull request #9 from xmarre/codex/apply-provided-patch
Fix guider selection for SLRD Impact detailer
2026-03-15 09:49:44 +01:00
xmarre c168188617 Fix AG detailer fallback 2026-03-15 09:44:04 +01:00
xmarre 8719d21b28 Update AG sampler guider handling 2026-03-15 09:21:08 +01:00
xmarre e4c296be32 Fix ScaleLocked hook state leaks 2026-03-14 23:27:07 +01:00
xmarre a1dc924a7e Fix mutable hook state retention 2026-03-14 23:25:18 +01:00
xmarre 2bcbcc68a2 Use live latent in detailer hook 2026-03-14 06:43:40 +01:00
xmarre 0e1043fe29 Use live latent in detailer hook 2026-03-14 06:42:44 +01:00
xmarre 1b85f680e6 Keep highres nested noise on sampler 2026-03-14 05:59:18 +01:00
xmarre ad283cf47e Normalize noise to sampler device 2026-03-14 05:44:04 +01:00
xmarre ede2b072b6 Use sampler device for latent norm 2026-03-14 05:43:27 +01:00
xmarre 2c205acb77 Add sampler extra args normalization 2026-03-14 04:58:01 +01:00
xmarre 885fc43eac Fix planner noise handoff in slrd 2026-03-14 02:36:12 +01:00
xmarre e4f9712de8 Fix planner noise handoff 2026-03-14 02:29:52 +01:00
xmarre e85d22b705 Fix detailer hook runtime timing 2026-03-14 01:30:43 +01:00
xmarre fb850de8b7 Update runtime mask source 2026-03-13 23:28:11 +01:00
xmarre d2d998110b Plan provider truthfulness cleanup 2026-03-13 21:50:40 +01:00
xmarre 2aa4c18b20 Clear detailer hooks and alias fallb 2026-03-13 21:49:14 +01:00
xmarre d9ba0ed67e Bump pyproject version 2026-03-12 16:28:44 +01:00
xmarre 3b6728b587 Add manifold tether parameters 2026-03-12 16:26:21 +01:00
xmarre 0e39e0162e Harden manifold dtype handling 2026-03-12 14:33:30 +01:00
xmarre 64251ed44d Harden manifold stats precision 2026-03-12 14:32:12 +01:00
xmarre 3e84dc9f9a Return tensor from custom sampler 2026-03-11 05:47:18 +01:00
xmarre 65d485c540 Add touch_scaled_size passthrough 2026-03-11 05:38:39 +01:00
xmarre bb79bc7a46 Add passthrough hooks to detailer 2026-03-11 05:16:02 +01:00
xmarre c444686b6e Add post crop region noop hook 2026-03-11 05:05:15 +01:00
xmarre bbfbe2ec64 Add post_crop_region no op to SLRD 2026-03-11 04:44:41 +01:00
xmarre 3042ef4de5 Add post crop passthrough hook 2026-03-11 04:43:52 +01:00
xmarre 8bb7b28a62 Add SLRD hook post crop region 2026-03-11 04:36:58 +01:00
xmarre 2d02f2bdc1 Add post_crop_region passthrough 2026-03-11 04:35:31 +01:00
xmarre a1a3edd656 Bump version from 1.0.0 to 1.0.1 2026-03-11 03:21:25 +01:00
xmarre 65cb78ee25 Merge pull request #8 from xmarre/codex/clarify-guider-and-detailer-flow
Fix scale-lock guider wrapping and document runtime usage
2026-03-11 03:20:53 +01:00
xmarre 7e5959cb97 Explain guider versus detailer paths 2026-03-11 03:20:21 +01:00
xmarre 93f413b6f5 Merge pull request #7 from xmarre/codex/fix-error-publishing-comfyui-node
Fix Comfy registry publish by setting correct PublisherId
2026-03-11 01:46:10 +01:00
xmarre 064ff6b4a0 Fix Comfy publisher id for registry publish 2026-03-11 01:46:02 +01:00
xmarre 42995e8705 Merge pull request #6 from xmarre/codex/fix-comfyui-node-publishing-error
Fix project URLs to match published repository
2026-03-11 01:27:00 +01:00
xmarre 8894cdef35 Fix project URLs to match published repository 2026-03-11 01:26:44 +01:00
xmarre fc261d0a08 Merge pull request #5 from xmarre/codex/add-custom-node-files-for-comfyui
Add Comfy registry metadata and publish workflow
2026-03-11 01:06:04 +01:00
xmarre d7c7a6df6a Add Comfy registry metadata and publish workflow 2026-03-11 01:05:48 +01:00
xmarre 42164a494a Merge pull request #4 from xmarre/codex/improve-lock-schedule-handling
Clarify scale-lock scheduling controls
2026-03-11 00:54:58 +01:00
xmarre 410ae514ba Explain lock schedule behavior 2026-03-11 00:54:28 +01:00
xmarre 7aa9223784 Merge pull request #3 from xmarre/codex/refactor-scalelock-init-helper
Refactor scale lock init helpers
2026-03-10 06:12:37 +01:00
xmarre 34c1bee399 Refactor scale lock init helper 2026-03-10 06:12:19 +01:00
xmarre 1a94ac4e8a Make ScaleLockedCFGGuider plain 2026-03-10 05:28:57 +01:00
xmarre ca64d701b6 Merge pull request #2 from xmarre/codex/add-advanced-scalelocked-sampler
Add advanced scale-locked sampler node and guider init fix
2026-03-10 05:04:39 +01:00
xmarre d9a26d49ec Add ScaleLockedCFGGuider init 2026-03-10 05:02:59 +01:00
xmarre 733d85b667 Merge pull request #1 from xmarre/codex/implement-initial-code
Harden scale-locked sampler noise mixing and alignment
2026-03-10 03:40:19 +01:00
12 changed files with 3106 additions and 577 deletions
+28
View File
@@ -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
View File
@@ -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}"
}
]
}
]
}
]
}
-7
View File
@@ -1,7 +0,0 @@
{
"ExpandedNodes": [
""
],
"SelectedNode": "\\C:\\Users\\marre\\Source\\Repos\\ComfyUI-ScaleLockedResidualDiffusion",
"PreviewInSolutionExplorer": false
}
+82 -61
View File
@@ -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
+11 -4
View File
@@ -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.
+11 -4
View File
@@ -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.
+1529 -398
View File
File diff suppressed because it is too large Load Diff
+16
View File
@@ -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
View File
@@ -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
View File
@@ -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)