Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9948cb3433 | ||
|
|
6cb9bbe868 | ||
|
|
561b73630c | ||
|
|
51efa670e0 | ||
|
|
d188a3764b | ||
|
|
f1d0630729 | ||
|
|
0cd1032073 | ||
|
|
d01b085082 | ||
|
|
0583ba2675 | ||
|
|
dcc37d7ab6 | ||
|
|
583b22a4bf | ||
|
|
ebe01efc49 | ||
|
|
41e8a2b61c | ||
|
|
22e4a5d202 | ||
|
|
017e3fc7fe | ||
|
|
c298efd8e8 | ||
|
|
7b92a4dafc | ||
|
|
ba13ef3345 | ||
|
|
dfa00d04b6 | ||
|
|
f4e33ff1c7 | ||
|
|
35ffc8b0e3 | ||
|
|
9f4dc7a193 | ||
|
|
1e0f44236e | ||
|
|
af6913e7bf | ||
|
|
b73d3c570f | ||
|
|
c35191de31 | ||
|
|
e8cbc6f724 | ||
|
|
4c13682df3 | ||
|
|
ffee87974e | ||
|
|
4f0d683c56 | ||
|
|
837a675a13 | ||
|
|
02710d2544 | ||
|
|
b80d2bc420 | ||
|
|
d9b92d65b1 | ||
|
|
24878eef89 | ||
|
|
9de6033505 | ||
|
|
36227fa8ff | ||
|
|
4875afdb7e | ||
|
|
dda0a3d326 | ||
|
|
ace5fafda9 | ||
|
|
f031c28589 | ||
|
|
fc738179f3 | ||
|
|
64ac708ba0 | ||
|
|
4f4c47ae0a | ||
|
|
e1d051781a | ||
|
|
e266880244 | ||
|
|
44cc18b86c | ||
|
|
18097c8fae | ||
|
|
cab3cac8f7 | ||
|
|
13a70540fd | ||
|
|
652ae51fc4 | ||
|
|
823fe209d8 | ||
|
|
1972e452dc | ||
|
|
29ee772b1d | ||
|
|
24492d97a3 | ||
|
|
7156dff28f | ||
|
|
d9e20fc601 | ||
|
|
74f140b2f3 | ||
|
|
9b9f7018fa | ||
|
|
741b5661c6 | ||
|
|
493eeb94b9 | ||
|
|
abd63b6296 | ||
|
|
4e3fd80bb8 | ||
|
|
6470b439b5 | ||
|
|
28bb8f7e88 | ||
|
|
92c1992379 | ||
|
|
acd668f7c6 | ||
|
|
4c7f087ce5 | ||
|
|
c0476b8288 | ||
|
|
b14f997145 | ||
|
|
346ff8b7c4 | ||
|
|
92c53b3493 | ||
|
|
f863fcf645 | ||
|
|
0679b7600b | ||
|
|
4202ca6d46 | ||
|
|
cc2561b557 | ||
|
|
57c4b8c5a4 | ||
|
|
dd514c1aad | ||
|
|
4b6525ce6f |
@@ -0,0 +1,13 @@
|
||||
.github/
|
||||
tests/
|
||||
tools/
|
||||
scripts/
|
||||
web/src/
|
||||
web/tests/
|
||||
AGENTS.md
|
||||
.releaserc.cjs
|
||||
eslint.config.js
|
||||
package-lock.json
|
||||
package.json
|
||||
tsconfig.json
|
||||
vitest.config.ts
|
||||
@@ -52,6 +52,7 @@ jobs:
|
||||
- name: Verify GroundingDINO BERT compatibility
|
||||
run: >-
|
||||
python -m pytest -q
|
||||
--noconftest
|
||||
--rootdir=tests
|
||||
--confcutdir=tests
|
||||
tests/test_grounding_dino_bert_adapter.py
|
||||
@@ -104,7 +105,7 @@ jobs:
|
||||
run: mypy --strict simple_syrup tests
|
||||
|
||||
- name: Verify Python tests
|
||||
run: pytest -n auto -q
|
||||
run: pytest -n auto -q -m "not external_artifact"
|
||||
|
||||
- name: Verify frontend
|
||||
run: npm run check:web
|
||||
|
||||
@@ -1,86 +1,100 @@
|
||||
# [1.6.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.5.0...v1.6.0) (2026-08-09)
|
||||
## [1.9.3](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.9.2...v1.9.3) (2026-09-20)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **cache:** make integer narrowing checker-independent ([6e2d1e1](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6e2d1e1844b361f9b2f3a31c8538cfde0ce7c6b1))
|
||||
* **mask:** preserve missing-alpha image geometry ([70ffeb5](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/70ffeb530acb1f5e55ac6e06adcc07ba4560d776))
|
||||
* **media:** stabilize native ordered preview controls ([3d03b1d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/3d03b1dfbc2eef3cee900173ddda57e03d5bc43d))
|
||||
* **regional:** align prompt batches and LoRA hooks ([646e4e7](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/646e4e7cab0ed299e704c4acbf64849181180076))
|
||||
* **runtime:** centralize Comfy patcher lifecycle ([fda2ef4](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fda2ef4074e8550a406ada317f0a3bfb7db39cf6))
|
||||
* **attention-coupling:** restore regional LoRA sampling ([0255a0f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/0255a0f5044278f14452b6b2582ec6646083f756))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **conditioning:** add regional prompting and SEP-local LoRAs ([be26493](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/be264935a8cb0e805222de9d615a197a788d8f01))
|
||||
* **conditioning:** support labeled prompt separators ([a24da13](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a24da1359cc9966762d8e1fa4fde3dbdef879cfa))
|
||||
* **loaders:** add automatic FLUX model loaders ([08fd18c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/08fd18cc6fdf6a7f9ff3657c323b31e8960a235a))
|
||||
* **media:** add native ordered loaders and SEGS preview ([dbde2a9](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/dbde2a9266dd043b92126e42e502cc88d9fba777))
|
||||
* **sampling:** add contextual diffusion sampler ([6c2e2da](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6c2e2dad9c6db42a2a2a82d1e4a17e4db4bb0fbf))
|
||||
* **sampling:** expose evaluated context SEGS ([53a2771](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/53a2771ce815a0705b92f4766e56f11adbd9dedb))
|
||||
* **segmentation:** add interactive SEGS preview ([21db62d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/21db62d3acee0b1d087f83e01fb5b92c912f4df6))
|
||||
* **segmentation:** add SAM region overlay ([d1ead71](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d1ead71ca714e365bf27448170e6c8c8fccd2f5a))
|
||||
* **segmentation:** add SAM-guided tiled diffusion ([e50ec0e](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/e50ec0ea79620af63f32c25a934b2b06fc4d9484))
|
||||
|
||||
# [1.5.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.4.0...v1.5.0) (2026-07-14)
|
||||
## [1.9.2](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.9.1...v1.9.2) (2026-09-20)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **groundingdino:** support transformers v4 and v5 ([239070b](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/239070b12d3eecbfd5c47c9410c7ca31ac1402ac))
|
||||
* **registry:** remove flagged package content ([f3a53b5](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/f3a53b5ec6c080d98f7e9599cf55e849cf338021))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **masking:** expand segmentation tooling and progress ([f70766d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/f70766dafe096895ad8d6309681fd59270664600))
|
||||
|
||||
# [1.4.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.3.0...v1.4.0) (2026-06-02)
|
||||
## [1.9.1](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.9.0...v1.9.1) (2026-09-20)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **tiled-diffusion:** clamp overlap for small latents ([d7448c6](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d7448c6ca52ce517b8d0f8ee697249c5def13535))
|
||||
* **contextual-diffusion:** project reference latents into views ([4cd780a](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/4cd780a2451aa472ce826834e4426b65693c46e8))
|
||||
|
||||
# [1.9.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.8.0...v1.9.0) (2026-09-19)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **detailing:** add external llm segs tagging ([b7cd40c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b7cd40c85ae9f9514d0b1ff3c3f6007730796752))
|
||||
* **segs:** add regional batching and wd14 tagging nodes ([fef60ad](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fef60adeca2c8ec7c0641e8106bb0863ee2f195e))
|
||||
* **prompts:** add automatic NegPiP support ([6d052e9](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6d052e9800985696972bf431fa7dae4972a56313))
|
||||
|
||||
# [1.3.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.2.0...v1.3.0) (2026-05-26)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **prompt-control:** add schedule and encode prompt node ([bd515e6](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/bd515e696cedcc78d056af6c23b9193e34f131bc))
|
||||
|
||||
# [1.2.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.1.0...v1.2.0) (2026-05-25)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **nodes:** add VAE options and clone-safe diffusion ([6ff2dc8](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6ff2dc8c24a6f7ddde3182b81bcbe6aad65427f4))
|
||||
|
||||
# [1.1.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.0.0...v1.1.0) (2026-05-23)
|
||||
# [1.8.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.7.1...v1.8.0) (2026-09-19)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **detailers:** align SEGS mask blending behavior ([74c83a6](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/74c83a6a70407aea6a34b0c0e31408f32b929882))
|
||||
* use SimpleSyrup package identity ([4fa5582](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/4fa5582796e9ace2a8805cb1ea2aebce65551e89))
|
||||
* **downloads:** keep unknown sizes indeterminate ([a31467c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a31467cd3a4299e9e3929d44281018dba0322b3e))
|
||||
* **models:** hide installed catalog choices ([d887e87](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d887e87e03eb6d0d9d5325fe43b67fbd9e47cf71))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **detection:** add keep-only SEGS selection ([8f8ee91](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/8f8ee91dea3ce1a044c0e61b482e571c51b372bc))
|
||||
* **models:** add curated ultralytics downloads ([b907fa2](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b907fa2a17afa30170bc22cf1134a750241e55c2))
|
||||
* **models:** prioritize installed ultralytics choices ([d30e04f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d30e04f229366d3e1d2388bd4d713b2b47a12a1f))
|
||||
|
||||
# 1.0.0 (2026-05-22)
|
||||
## [1.7.1](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.7.0...v1.7.1) (2026-09-11)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **regional:** preserve shared model patch ancestry ([6059a3f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6059a3f913a9502671666e83faeb8686a7a8da27))
|
||||
|
||||
# [1.7.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.6.0...v1.7.0) (2026-09-05)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **anima:** support regional prompting across Comfy versions ([41a234a](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/41a234a85a4cfcdc4cfce68b80b4b9982719aab4))
|
||||
* **attention:** preserve anchored concept geometry ([1136efd](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/1136efd2ad8708b14320d04eab2f6489efbec282))
|
||||
* **cache:** make integer narrowing checker-independent ([2547767](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/254776791766c41c75a69ceb5b207c54949446c9))
|
||||
* **detailers:** align SEGS mask blending behavior ([a3120ae](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a3120aebe8834982706d93784a39ab501c6ff40a))
|
||||
* **groundingdino:** support transformers v4 and v5 ([3387c03](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/3387c03eb59f244d56f6a0dfe86cedab1847a8e2))
|
||||
* **mask:** preserve missing-alpha image geometry ([caf7d37](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/caf7d37a154ddc62405807b56550efdaa831d09e))
|
||||
* **media:** stabilize native ordered preview controls ([1235652](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/123565225a1104234a7c9f5c43af0772c7508db8))
|
||||
* **regional:** align prompt batches and LoRA hooks ([656d197](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/656d1970e12be18304b057b07204dcbbe367432b))
|
||||
* **runtime:** centralize Comfy patcher lifecycle ([0e5f513](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/0e5f513ae0f33f40c6a8bd09161043a5af598392))
|
||||
* **sampling:** normalize model-specific latent layouts ([b2084a7](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b2084a7a9bf04709af51380bc1ef4a09ccb9babc))
|
||||
* **tiled-diffusion:** clamp overlap for small latents ([fecba36](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fecba36e4e11f0da681c6a5d9d42e18093d741fc))
|
||||
* **tools:** return host-native checkpoint selections ([c7cb8d2](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/c7cb8d29f83ebd44b48038b7ce5e65a0a6f445b4))
|
||||
* use SimpleSyrup package identity ([1659d13](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/1659d131f4215fa0baccc4c70024d63590460e67))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* initial release ([e513baf](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/e513baf70a20306856e40fbf2afd80b25f5655a6))
|
||||
|
||||
# Changelog
|
||||
|
||||
All notable changes to this project will be documented in this file.
|
||||
* **anima:** add cached quantization profiles ([c20664f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/c20664f8925305465ccb4c028d0a78a6364a037d))
|
||||
* **attention:** add sampler-derived concept regions ([8ca5d2c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/8ca5d2c325e1caee822883ba568b25f721e51343))
|
||||
* **attention:** default regional prompts to full weight ([f829321](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/f82932193951bae75d59e1c7c7dba2d187e3c175))
|
||||
* **attention:** improve concept isolation fidelity and speed ([01826ad](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/01826ad1b7c64b11c8af031691403ecf521cb1b5))
|
||||
* **attention:** refine attention-derived region masks ([dff84cc](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/dff84cce2322e64b8a9ada91b84550faf3a5c7a1))
|
||||
* **conditioning:** add regional prompting and SEP-local LoRAs ([d3de028](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d3de02816b24bea129e77b82568ad47a0fd0ba99))
|
||||
* **conditioning:** support labeled prompt separators ([a708e0b](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a708e0b4187b0d5f9aeb58f5ab0d276b8a046395))
|
||||
* **detailing:** add external llm segs tagging ([207c449](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/207c4492908371f41aca84c74332c1a1f63d8045))
|
||||
* **detection:** add keep-only SEGS selection ([4f65ae0](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/4f65ae0dd36e43347b40ae64039f40e8a67aea48))
|
||||
* initial release ([4b6525c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/4b6525ce6ff42f06a7ffd48a54186fcb625f0e21))
|
||||
* **loaders:** add automatic FLUX model loaders ([33bff75](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/33bff75af1b001b716a45f4fadc9e0e0aa20ced1))
|
||||
* **masking:** expand segmentation tooling and progress ([adba198](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/adba1981a4bdb43307496a84cccbfe100ae2174f))
|
||||
* **media:** add native ordered loaders and SEGS preview ([fcf38f2](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fcf38f2010cfadc76be74410857387c74db3b575))
|
||||
* **nodes:** add VAE options and clone-safe diffusion ([5fc0f3d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/5fc0f3d8e5ef5fbac0306b1d3c70464a035c396c))
|
||||
* **prompt-control:** add schedule and encode prompt node ([af32cec](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/af32cec99f5bdadcbed9f8c33db0d816ad4b72f0))
|
||||
* **regional:** add native SDXL adapter execution ([068e3db](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/068e3db17d13c875a383c85e6cf931f96459c3b1))
|
||||
* **regional:** add universal attention coupling foundation ([edfc26c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/edfc26c9122ff0c18d63e9b80832d587594e9e62))
|
||||
* **regional:** build universal adapter execution foundation ([864852d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/864852dc9e591a3041e23b342346495a7fbf6589))
|
||||
* **regional:** complete capability-routed execution ([fe96a20](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fe96a206dbf1edcae022b1346f998ed184142062))
|
||||
* **regional:** complete persistent regional LoRA execution ([7b5987d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/7b5987d6fa66f8deee2655d18b1209bbe20271d6))
|
||||
* **sampling:** add contextual diffusion sampler ([24bf630](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/24bf6309023c552f65947814d9641555a73c8337))
|
||||
* **sampling:** add deterministic seed variation ([7921be3](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/7921be3f9067a26fe5fa47fb8fb370e7cb1679f3))
|
||||
* **sampling:** add regional diffusion sampling ([f7dffca](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/f7dffcacfe104be173793ce412f994519e11e02e))
|
||||
* **sampling:** bypass inactive attention coupling ([b2b8dc4](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b2b8dc4b016307fd351892f50264c64532898681))
|
||||
* **sampling:** expose evaluated context SEGS ([04a2c3e](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/04a2c3e6e90f9bbb9e3922412844afe5a4e6869f))
|
||||
* **segmentation:** add interactive SEGS preview ([c9e303e](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/c9e303ec4727576af124d9e8ea20d0103f67e7ae))
|
||||
* **segmentation:** add SAM region overlay ([a72c796](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a72c796d8d6d64a52fed7d281fe53eebc74639ed))
|
||||
* **segmentation:** add SAM-guided tiled diffusion ([968090d](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/968090de87a4fad726fd37d0a08966c01f11a8fd))
|
||||
* **segs:** add regional batching and wd14 tagging nodes ([600d9e3](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/600d9e311b5b8c45762e15c88f54df33454261b5))
|
||||
|
||||
@@ -1,543 +0,0 @@
|
||||
# Plan: Image-Associated Mask to SEGS
|
||||
|
||||
## Goal
|
||||
|
||||
Add a SimpleSyrup node that converts an existing ComfyUI mask into Impact-compatible SEGS while associating every SEG with cropped image data from the source image.
|
||||
|
||||
The node should feel like `Detect SEGS w/ Ultralytics`, but it must take an `IMAGE` and a `MASK` instead of an image and detector model. It should expose the same general region workflow controls where they make sense: size filtering, keep-only limiting, mask dilation, post-dilation, crop factor, sort order, and optional unioning into one SEG.
|
||||
|
||||
## Decisions From Maintainer Discussion
|
||||
|
||||
- This should be a first-class SimpleSyrup node, not just documentation for an Impact Pack workflow chain.
|
||||
- The node should behave like a detector-style SEGS source because downstream users think of this as "pre-detected regions from a mask."
|
||||
- The node must take both `image` and `mask`.
|
||||
- The output SEGS must include full image association by storing each SEG's `cropped_image` from the input image.
|
||||
- The node should support masks with multiple disconnected regions.
|
||||
- The node should let users choose whether disconnected regions become separate SEGs or one combined SEG.
|
||||
- The node should include controls similar to `Detect SEGS w/ Ultralytics` and in a similar order.
|
||||
- The node should not expose detector confidence controls because the input is an existing mask, not detector predictions.
|
||||
- Segment confidence should be fixed at `1.0`.
|
||||
- The existing Impact Pack nodes can already do some of this with `MASK to SEGS` plus `Set Default Image for SEGS`, but that is not the desired UX for SimpleSyrup.
|
||||
|
||||
## Existing Code To Reuse
|
||||
|
||||
Use these SimpleSyrup modules as the main implementation references:
|
||||
|
||||
- `simple_syrup/nodes/detect_segs_with_ultralytics.py`
|
||||
- Existing detector-style node UX and control order.
|
||||
- Current output shape: `RETURN_TYPES = ("SEGS", "MASK")`.
|
||||
- Current list behavior for batched images: `OUTPUT_IS_LIST = (True, False)`.
|
||||
- Applies `limit_segs`, `sort_segs`, and then `build_combined_segs_result`.
|
||||
- `simple_syrup/services/segs_detection_service.py`
|
||||
- Existing service pattern for constructing `Segment` objects from masks and image crops.
|
||||
- Uses `crop_region_for_bbox`, `crop_mask`, `crop_image`, and `dilate_mask`.
|
||||
- `simple_syrup/services/segs_output_service.py`
|
||||
- `build_combined_segs_result()` already creates a one-SEG union with `cropped_image`.
|
||||
- `combined_mask_from_segs()` already produces the mask output from SEGS.
|
||||
- `coerce_cropped_mask()` validates crop-local SEG masks.
|
||||
- This module should become the shared owner for detector-style SEGS output finalization.
|
||||
- `simple_syrup/masking/segs_mask_ops.py`
|
||||
- Reuse image validation, mask crop, image crop, resize, crop factor, and signed dilation behavior.
|
||||
- `simple_syrup/domain/segs.py`
|
||||
- Reuse `Segment`, `BoundingBox`, `CropRegion`, `NativeSegs`, `SORT_ORDER_OPTIONS`, `limit_segs`, `sort_segs`, and `to_impact_compatible_segs`.
|
||||
|
||||
Use these Impact Pack modules only as behavioral references, not as imports:
|
||||
|
||||
- `E:\ComfyUI\custom_nodes\comfyui-impact-pack\modules\impact\segs_nodes.py`
|
||||
- `MaskToSEGS` splits or combines mask regions.
|
||||
- `DefaultImageForSEGS` attaches cropped source image data after mask conversion.
|
||||
- `SEGSMerge`, `SEGSOrderedFilter`, and `DilateMaskInSEGS` show existing user expectations.
|
||||
- `E:\ComfyUI\custom_nodes\comfyui-impact-pack\modules\impact\core.py`
|
||||
- `mask_to_segs()` uses contour detection for disconnected regions.
|
||||
- `batch_mask_to_segs()` handles batched masks for video-style masks.
|
||||
|
||||
Do not import Impact Pack code. SimpleSyrup tests already enforce avoiding external pack imports.
|
||||
|
||||
## Node Name And Contract
|
||||
|
||||
Add a new Comfy v3 node:
|
||||
|
||||
- Node id: `SimpleSyrup.MaskToSEGS`
|
||||
- Display name: `Mask to SEGS`
|
||||
- Category: `SimpleSyrup/Detection`
|
||||
- Search aliases: `mask`, `segs`, `region`, `detect`, `segmentation`
|
||||
- Description: `Converts an existing mask into image-associated SEGS for detail and regional workflows.`
|
||||
|
||||
Outputs:
|
||||
|
||||
- `segs: SEGS`
|
||||
- `mask: MASK`
|
||||
|
||||
Output behavior should match `Detect SEGS w/ Ultralytics`:
|
||||
|
||||
- `segs` is list-output compatible, one SEGS payload per input image/mask pair.
|
||||
- `mask` is a standard ComfyUI batched mask tensor with shape `(B, H, W)`.
|
||||
- The output mask is the union of the retained SEGS after filtering, sorting, and optional combining.
|
||||
|
||||
## Architecture Landing Shape
|
||||
|
||||
Do not add a third copy of the detector-style output pipeline.
|
||||
|
||||
The current code has two nodes that inline the same sequence after SEGS extraction:
|
||||
|
||||
1. `simple_syrup/nodes/detect_segs_with_ultralytics.py`
|
||||
2. `simple_syrup/nodes/prompt_segs_with_sam.py`
|
||||
|
||||
Both nodes currently perform:
|
||||
|
||||
```python
|
||||
segs = limit_segs(segs, keep_only, keep_by)
|
||||
segs = sort_segs(segs, sort_order)
|
||||
combined = build_combined_segs_result(single_image, segs, crop_factor)
|
||||
output_segs = combined.segs if combine_segs else segs
|
||||
segs_outputs.append(to_impact_compatible_segs(output_segs))
|
||||
mask_outputs.append(combined.mask)
|
||||
```
|
||||
|
||||
The new mask node must not duplicate this sequence inline. Instead, implement a shared output-finalization helper in `simple_syrup/services/segs_output_service.py` and refactor the existing detector nodes to use it.
|
||||
|
||||
Recommended shape:
|
||||
|
||||
- `MaskToSEGSService`
|
||||
- Owns only one-image, one-mask extraction into separate native SEGS.
|
||||
- Does not own `keep_only`, `sort_order`, `combine_segs`, Impact-compatible conversion, or output mask construction.
|
||||
- `segs_output_service`
|
||||
- Owns the common detector-style post-processing pipeline:
|
||||
1. limit SEGS
|
||||
2. sort SEGS
|
||||
3. build the combined result
|
||||
4. choose separate or combined output SEGS
|
||||
5. convert output SEGS to Impact-compatible shape
|
||||
6. return the paired mask output
|
||||
- Node classes
|
||||
- Own Comfy-facing schema, batch iteration, and wiring only.
|
||||
- Delegate source-specific extraction to a source service.
|
||||
- Delegate shared final output shaping to `segs_output_service`.
|
||||
|
||||
This keeps ownership strict:
|
||||
|
||||
- Mask-derived region extraction belongs to the mask-to-SEGS service.
|
||||
- Detector model inference belongs to detector services.
|
||||
- Prompt/SAM detection belongs to the SAM prompt service.
|
||||
- Sorting, limiting, combining, and output mask generation belong to the shared SEGS output service.
|
||||
- Domain helpers remain pure SEGS policies and value conversions.
|
||||
|
||||
## Input Schema
|
||||
|
||||
Use this input order:
|
||||
|
||||
1. `image: IMAGE`
|
||||
2. `mask: MASK`
|
||||
3. `mask_threshold: FLOAT`
|
||||
4. `size_threshold: INT`
|
||||
5. `keep_only: INT`
|
||||
6. `mask_dilation: INT`
|
||||
7. `post_dilation: INT`
|
||||
8. `crop_factor: FLOAT`
|
||||
9. `sort_order: SORT_ORDER_OPTIONS`
|
||||
10. `combine_segs: BOOLEAN`
|
||||
11. `label: STRING`
|
||||
|
||||
Detailed input behavior:
|
||||
|
||||
- `image`
|
||||
- Source image used to associate cropped image data with every SEG.
|
||||
- Must be a ComfyUI `IMAGE` tensor shaped `(B, H, W, C)`.
|
||||
- `mask`
|
||||
- Source mask to convert into SEGS.
|
||||
- Must be a ComfyUI `MASK` tensor shaped `(B, H, W)` or a single mask compatible with the image batch.
|
||||
- The implementation may support a single mask for a single image first. If supporting image batches, batch count must either match the image batch or be exactly `1` for reuse across all images.
|
||||
- `mask_threshold`
|
||||
- Default: `0.5`
|
||||
- Min: `0.0`
|
||||
- Max: `1.0`
|
||||
- Step: `0.01`
|
||||
- Converts soft masks into active pixels before region extraction.
|
||||
- Active pixels are `mask >= mask_threshold`.
|
||||
- `size_threshold`
|
||||
- Default: `10`
|
||||
- Min: `1`
|
||||
- Max: `8192`
|
||||
- Drops extracted regions whose bounding box is smaller than this many pixels wide or tall.
|
||||
- Match the wording and role of `Detect SEGS w/ Ultralytics`'s `size_threshold`.
|
||||
- `keep_only`
|
||||
- Default: `0`
|
||||
- Min: `0`
|
||||
- Max: `4096`
|
||||
- Keeps only the largest N regions after size filtering. `0` keeps all regions.
|
||||
- Do not expose a confidence ranking option. There is no detector confidence.
|
||||
- `mask_dilation`
|
||||
- Default: `0`
|
||||
- Min: `-512`
|
||||
- Max: `512`
|
||||
- Signed dilation applied to the full input mask before region extraction.
|
||||
- Positive grows mask regions; negative erodes them.
|
||||
- `post_dilation`
|
||||
- Default: `0`
|
||||
- Min: `-512`
|
||||
- Max: `512`
|
||||
- Signed dilation applied to each final crop-local SEG mask after the crop is chosen.
|
||||
- This mirrors the Ultralytics node's final SEG mask cleanup behavior.
|
||||
- `crop_factor`
|
||||
- Default: `3.0`
|
||||
- Min: `0.0`
|
||||
- Max: `100.0`
|
||||
- Step: `0.1`
|
||||
- `0.0` means use the full image as the crop, matching SimpleSyrup's existing detector convention.
|
||||
- Values greater than or equal to `1.0` expand the SEG crop around the extracted region's bounding box.
|
||||
- Values between `0.0` and `1.0` must fail with an actionable `ValueError`.
|
||||
- `sort_order`
|
||||
- Use `SORT_ORDER_OPTIONS` from `simple_syrup.domain.segs`.
|
||||
- Default: `largest to smallest`.
|
||||
- Sort extracted regions before output and before any combined mask result is built.
|
||||
- `combine_segs`
|
||||
- Default: `False`
|
||||
- `False` returns one SEG per disconnected mask region.
|
||||
- `True` returns one unioned SEG representing all retained mask pixels.
|
||||
- `label`
|
||||
- Default: `mask`
|
||||
- Label applied to extracted SEGs.
|
||||
- If `combine_segs` is true, the combined output label should be `combined` unless there is a strong reason to preserve the user label. Match `build_combined_segs_result()` unless intentionally extending it.
|
||||
|
||||
## Behavior Details
|
||||
|
||||
### Region Extraction
|
||||
|
||||
Add a new service, probably `simple_syrup/services/mask_to_segs_service.py`, with a class such as `MaskToSEGSService`.
|
||||
|
||||
The service should own:
|
||||
|
||||
- Image and mask validation.
|
||||
- Mask thresholding.
|
||||
- Signed full-mask dilation.
|
||||
- Disconnected region extraction.
|
||||
- Per-region bounding box calculation.
|
||||
- Size filtering.
|
||||
- Crop region calculation.
|
||||
- Crop-local mask creation.
|
||||
- Optional post-dilation on each crop-local mask.
|
||||
- Cropped image association.
|
||||
- Native immutable SEGS output.
|
||||
|
||||
The service must not own:
|
||||
|
||||
- `keep_only`
|
||||
- final `sort_order`
|
||||
- `combine_segs`
|
||||
- Impact-compatible conversion
|
||||
- final output mask construction
|
||||
|
||||
Those responsibilities belong to the shared SEGS output finalization helper.
|
||||
|
||||
The node should own only Comfy-facing schema, input ordering, batch iteration, and wiring.
|
||||
|
||||
### Connected Components
|
||||
|
||||
Implement disconnected region extraction inside SimpleSyrup. Do not import Impact Pack.
|
||||
|
||||
Prefer a torch-based or standard-library implementation over adding a new required runtime dependency. A simple deterministic flood-fill or connected-components routine is acceptable because masks are 2D binary tensors and the node is not model-bound.
|
||||
|
||||
Connectivity decision:
|
||||
|
||||
- Use 8-connected components unless tests or existing SimpleSyrup behavior strongly point to 4-connected components.
|
||||
- Document this in the service docstring and tests.
|
||||
- 8-connected behavior usually matches user expectations for painted masks where diagonal contact should remain one region.
|
||||
|
||||
Extraction algorithm outline:
|
||||
|
||||
1. Convert mask to a CPU `torch.bool` active-pixel tensor after threshold and dilation.
|
||||
2. If no pixels are active, return empty native SEGS. The shared output finalization helper is responsible for turning that into a zero mask output.
|
||||
3. Find connected components.
|
||||
4. For each component, compute bbox from active pixel coordinates.
|
||||
5. Drop components where bbox width or height is less than `size_threshold`.
|
||||
6. Create one `Segment` per kept component.
|
||||
|
||||
The service always returns separate native SEGS. `combine_segs` is handled later by the shared output finalization helper so all detector-style SEGS nodes use the same combine behavior.
|
||||
|
||||
### Mask Values
|
||||
|
||||
Use thresholded masks to decide region membership, but preserve useful soft-mask values when forming crop-local masks where possible.
|
||||
|
||||
Recommended behavior:
|
||||
|
||||
- Use the thresholded, dilated mask for topology and bbox extraction.
|
||||
- Use the original normalized mask after `mask_dilation` as the crop-local mask values.
|
||||
- Zero out pixels outside the active component for each separate SEG.
|
||||
- Clamp all final masks to `0.0..1.0`.
|
||||
|
||||
This lets soft masks retain feathered values inside each SEG while still giving deterministic region extraction.
|
||||
|
||||
### Cropped Image Association
|
||||
|
||||
Every returned `Segment` must set:
|
||||
|
||||
- `cropped_image = crop_image(single_image, crop_region).detach().clone()`
|
||||
- `cropped_mask = crop-local mask tensor`
|
||||
- `confidence = 1.0`
|
||||
- `crop_region = CropRegion(...)`
|
||||
- `bbox = BoundingBox(...)`
|
||||
- `label = label`
|
||||
- `control_net_wrapper = None`
|
||||
|
||||
This is the key difference from Impact Pack's `MASK to SEGS`, which creates SEGS without cropped image data and relies on a separate `Set Default Image for SEGS` node.
|
||||
|
||||
### Batch Handling
|
||||
|
||||
Match existing SimpleSyrup detector behavior as closely as possible:
|
||||
|
||||
- Validate the image batch with `validate_image_batch()`.
|
||||
- Iterate image items with `iter_single_images()`.
|
||||
- Produce one SEGS payload per image.
|
||||
- Concatenate mask outputs into one `(B, H, W)` tensor.
|
||||
|
||||
Mask batch rules:
|
||||
|
||||
- If mask batch size equals image batch size, pair by index.
|
||||
- If mask batch size is `1` and image batch size is greater than `1`, reuse the mask for every image.
|
||||
- Otherwise raise `ValueError` explaining the mismatch.
|
||||
- Mask height and width must match the image height and width. Do not silently resize masks for this node unless the maintainer explicitly approves it later.
|
||||
|
||||
## Sorting And Keeping
|
||||
|
||||
Do not expose `keep_by` unless there is a future product decision to add multiple non-confidence policies.
|
||||
|
||||
For `keep_only`, retain the largest regions by crop area before final sorting. Reuse `limit_segs(segs, keep_only, "largest size")` inside the shared output finalization helper for this node. Do not add a new ranking abstraction unless tests show the current domain helper cannot express the behavior clearly.
|
||||
|
||||
Final output order must come from `sort_segs(segs, sort_order)`.
|
||||
|
||||
## Shared Output Finalization
|
||||
|
||||
Add a small result type and helper to `simple_syrup/services/segs_output_service.py`.
|
||||
|
||||
Suggested result type:
|
||||
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class FinalizedSegsOutput:
|
||||
"""Return Impact-compatible SEGS and its paired output mask."""
|
||||
|
||||
segs: object
|
||||
mask: torch.Tensor
|
||||
```
|
||||
|
||||
Suggested helper:
|
||||
|
||||
```python
|
||||
def finalize_detector_segs_output(
|
||||
image: object,
|
||||
segs: NativeSegs,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
crop_factor: float,
|
||||
sort_order: str,
|
||||
combine_segs: bool,
|
||||
) -> FinalizedSegsOutput:
|
||||
"""Apply shared detector-style SEGS output policy."""
|
||||
```
|
||||
|
||||
Behavior:
|
||||
|
||||
1. Apply `limit_segs(segs, keep_only, keep_by)`.
|
||||
2. Apply `sort_segs(segs, sort_order)`.
|
||||
3. Build `combined = build_combined_segs_result(image, segs, crop_factor)`.
|
||||
4. Use `combined.segs` when `combine_segs` is true; otherwise use the sorted separate SEGS.
|
||||
5. Convert chosen SEGS with `to_impact_compatible_segs`.
|
||||
6. Return the converted SEGS and `combined.mask`.
|
||||
|
||||
For the new mask node, call this helper with `keep_by="largest size"` internally.
|
||||
|
||||
Refactor these existing nodes to use the helper as part of this change:
|
||||
|
||||
- `simple_syrup/nodes/detect_segs_with_ultralytics.py`
|
||||
- `simple_syrup/nodes/prompt_segs_with_sam.py`
|
||||
|
||||
Add characterization tests before refactoring or preserve existing tests that already prove:
|
||||
|
||||
- final sorting happens before combined output construction
|
||||
- keep-only limiting happens before final sorting
|
||||
- `combine_segs` chooses the combined SEGS while retaining the same output mask
|
||||
- returned SEGS are Impact-compatible
|
||||
|
||||
The helper should be narrow. Do not move source-specific detection behavior into it.
|
||||
|
||||
## V3 Registration
|
||||
|
||||
Comfy v3 is the only supported export path.
|
||||
|
||||
Preferred implementation:
|
||||
|
||||
- Add a direct v3 node class under `simple_syrup/nodes_v3/mask_to_segs.py`.
|
||||
- Register it in `simple_syrup/nodes_v3/__init__.py::get_nodes()`.
|
||||
- Do not add `NODE_CLASS_MAPPINGS` or legacy export mappings.
|
||||
|
||||
If reusing the legacy adapter pattern would materially reduce risk, it is acceptable to add a legacy-style internal node class under `simple_syrup/nodes/` and wrap it with `LegacyNodeV3Adapter`, but the public export must still be v3-only.
|
||||
|
||||
The implementation should look native to the current codebase and should not add compatibility shims.
|
||||
|
||||
## Tooltips
|
||||
|
||||
Every visible input and output must have concise user-facing tooltip text.
|
||||
|
||||
Required tooltip intent:
|
||||
|
||||
- `image`: source image used for SEG crops.
|
||||
- `mask`: mask whose active regions become SEGS.
|
||||
- `mask_threshold`: threshold used to decide active mask pixels.
|
||||
- `size_threshold`: smallest region width or height to keep, in pixels.
|
||||
- `keep_only`: maximum number of largest regions to keep; `0` keeps all.
|
||||
- `mask_dilation`: grow or shrink the source mask before regions are found.
|
||||
- `post_dilation`: grow or shrink each final SEG mask after cropping.
|
||||
- `crop_factor`: context around each region; `0` uses the full image.
|
||||
- `sort_order`: output ordering for separate SEGS.
|
||||
- `combine_segs`: return one unioned SEG instead of separate regions.
|
||||
- `label`: label stored on extracted SEGs.
|
||||
- `segs` output: image-associated SEGS from the mask.
|
||||
- `mask` output: union of retained SEGS as a ComfyUI mask.
|
||||
|
||||
## Tests To Add
|
||||
|
||||
Add focused tests before or alongside implementation.
|
||||
|
||||
### Shared Output Finalization Tests
|
||||
|
||||
Add tests for the new helper in `tests/test_segs_output_service.py` or extend the existing SEGS output service coverage currently housed in `tests/test_ultralytics_detection_service.py`.
|
||||
|
||||
Cover:
|
||||
|
||||
- Applies `limit_segs` before `sort_segs`.
|
||||
- Builds the combined result from the limited and sorted SEGS.
|
||||
- Returns separate SEGS when `combine_segs` is false.
|
||||
- Returns one combined SEG when `combine_segs` is true.
|
||||
- Always returns the mask from `build_combined_segs_result`.
|
||||
- Converts output SEGS to Impact-compatible tuple/list shape.
|
||||
- Supports `keep_by="largest size"` for mask-derived SEGS.
|
||||
|
||||
After this helper is tested, refactor `Detect SEGS w/ Ultralytics` and `Prompt SEGS w/ SAM` to use it without changing their public behavior.
|
||||
|
||||
### Service Tests
|
||||
|
||||
Create `tests/test_mask_to_segs_service.py`.
|
||||
|
||||
Cover:
|
||||
|
||||
- Single rectangular mask creates one SEG.
|
||||
- `cropped_image` matches the source image crop.
|
||||
- `cropped_mask` matches the mask crop.
|
||||
- Two disconnected regions create two separate SEGs.
|
||||
- Service always returns separate SEGS; combining is covered by shared output finalization tests and node tests.
|
||||
- Empty mask returns empty native SEGS from the service and a zero mask output through the node/output helper.
|
||||
- `mask_threshold` controls active pixels.
|
||||
- `mask_dilation` grows a source mask before region extraction.
|
||||
- Negative `mask_dilation` erodes a source mask.
|
||||
- `post_dilation` changes only crop-local final masks.
|
||||
- `crop_factor` expands crop regions.
|
||||
- `crop_factor = 0.0` uses the full image.
|
||||
- `0.0 < crop_factor < 1.0` raises `ValueError`.
|
||||
- `size_threshold` drops small components.
|
||||
- Soft mask values are clamped and retained inside component masks.
|
||||
- Batch mask and image mismatch raises an actionable error if batch handling is in the service.
|
||||
|
||||
### Node Contract Tests
|
||||
|
||||
Create `tests/test_mask_to_segs_node.py` or a v3-specific equivalent.
|
||||
|
||||
Cover:
|
||||
|
||||
- Node id is `SimpleSyrup.MaskToSEGS`.
|
||||
- Display name is `Mask to SEGS`.
|
||||
- Category is `SimpleSyrup/Detection`.
|
||||
- Outputs are `SEGS` and `MASK`.
|
||||
- The input list order exactly matches the plan.
|
||||
- Every input has a tooltip.
|
||||
- Every output has a tooltip.
|
||||
- No confidence input exists.
|
||||
- No detector model input exists.
|
||||
- No `keep_by` widget exists unless the implementation deliberately adds more non-confidence policies and updates this plan.
|
||||
- Execution returns Impact-compatible SEGS and a `(B, H, W)` mask.
|
||||
- `combine_segs = false` returns separate SEGs.
|
||||
- `combine_segs = true` returns one combined SEG.
|
||||
- `keep_only` keeps the largest mask-derived regions.
|
||||
- `sort_order` orders separate regions using existing domain policies.
|
||||
- Batched images return list-output SEGS and batched masks.
|
||||
|
||||
### Registration Tests
|
||||
|
||||
Update existing registration tests:
|
||||
|
||||
- `SimpleSyrup.MaskToSEGS` appears in `get_nodes()`.
|
||||
- Adding the node does not remove or rename existing nodes.
|
||||
- Root `comfy_entrypoint` still exposes v3 nodes only.
|
||||
|
||||
### Tooltip Coverage Tests
|
||||
|
||||
Update `tests/test_node_tooltips.py` so the new node passes the repository tooltip requirements.
|
||||
|
||||
## Implementation Steps
|
||||
|
||||
- [x] Add or confirm characterization tests for the existing Ultralytics and SAM detector-style output behavior.
|
||||
- Landing note: Existing node tests already covered separate/combined outputs, keep-only limiting, final sorting, combined-builder input order, and batch handling for both detector nodes.
|
||||
- [x] Add shared output finalization tests.
|
||||
- Landing note: `tests/test_segs_output_service.py` now covers limit-before-sort, combined output selection, returned mask preservation, Impact-compatible conversion, and `keep_by="largest size"` for mask-derived SEGS.
|
||||
- [x] Add the shared output finalization helper in `segs_output_service.py`.
|
||||
- Landing note: `FinalizedSegsOutput` and `finalize_detector_segs_output()` now own detector-style limit/sort/combine/conversion/mask finalization.
|
||||
- [x] Refactor `Detect SEGS w/ Ultralytics` and `Prompt SEGS w/ SAM` to use the helper without behavior changes.
|
||||
- Landing note: Both nodes still own source-specific extraction and batch wiring, but delegate shared output shaping to `finalize_detector_segs_output()` with their injectable combined builders.
|
||||
- [x] Add the connected-component helper and mask-to-SEGS service.
|
||||
- Landing note: `mask_components.py` implements deterministic 8-connected component extraction, and `MaskToSEGSService` now converts one image plus one mask into separate native SEGS with cropped image association.
|
||||
- [x] Add direct v3 node schema and execution wrapper.
|
||||
- Landing note: `MaskToSEGSV3` defines the Comfy v3 schema directly, owns image/mask batch pairing, and delegates extraction plus shared output finalization.
|
||||
- [x] Register the node in `simple_syrup/nodes_v3/__init__.py`.
|
||||
- Landing note: `SimpleSyrup.MaskToSEGS` is included in the base v3 node list.
|
||||
- [x] Add or update registration and tooltip tests.
|
||||
- Landing note: Registration expectations include `SimpleSyrup.MaskToSEGS`; the schema-driven tooltip coverage includes every new input and output.
|
||||
- [x] Run focused tests for shared output finalization, existing detector nodes, the new service, the new node, registration, and tooltips.
|
||||
- Landing note: Focused plan checks pass: `51 passed`.
|
||||
- [x] Run full Python gates.
|
||||
- Landing note: `ruff format .`, `ruff check .`, `mypy --strict simple_syrup tests`, and full `pytest -n auto -q` pass. Full test result: `969 passed`.
|
||||
- [x] Do not touch frontend code unless the Comfy v3 UI requires it.
|
||||
- Landing note: No frontend source or generated browser artifact was changed.
|
||||
|
||||
## Verification Commands
|
||||
|
||||
Run all commands from repository root with the ComfyUI virtual environment two directories above this repo.
|
||||
|
||||
Focused checks during development:
|
||||
|
||||
```powershell
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_detect_segs_with_ultralytics_node.py tests\test_prompt_segs_with_sam_node.py
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_segs_output_service.py
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_mask_to_segs_service.py tests\test_mask_to_segs_node.py
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_registration.py tests\test_node_tooltips.py
|
||||
```
|
||||
|
||||
Required final gates:
|
||||
|
||||
```powershell
|
||||
..\..\venv\Scripts\ruff.exe format .
|
||||
..\..\venv\Scripts\ruff.exe check .
|
||||
..\..\venv\Scripts\mypy.exe --strict simple_syrup tests
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q
|
||||
```
|
||||
|
||||
If frontend code is touched, also run:
|
||||
|
||||
```powershell
|
||||
npm ci
|
||||
npm run lint:web
|
||||
npm run typecheck:web
|
||||
npm run test:web
|
||||
npm run build:web
|
||||
```
|
||||
|
||||
## Acceptance Criteria
|
||||
|
||||
- A user can plug in an image and a mask and get ready-to-use image-associated SEGS.
|
||||
- A mask with two disconnected regions can produce either two SEGs or one combined SEG.
|
||||
- The node has detector-style controls that feel aligned with `Detect SEGS w/ Ultralytics`.
|
||||
- The node exposes no detector confidence controls.
|
||||
- Each SEG has `cropped_image`, `cropped_mask`, `crop_region`, `bbox`, `label`, and `confidence = 1.0`.
|
||||
- The returned SEGS are Impact-compatible.
|
||||
- The returned mask is the union of retained output SEGS.
|
||||
- The implementation does not import Impact Pack.
|
||||
- The implementation is covered by behavior tests, node contract tests, registration tests, and tooltip tests.
|
||||
- Full Python verification gates pass.
|
||||
@@ -18,7 +18,7 @@ The pack now covers model loading, regional prompting and segmentation, high-res
|
||||
- ADetailer-style `[SEP]` prompt batches, masked conditioning, and regional samplers, with optional Prompt Control scheduling and LoRA hooks.
|
||||
- WD14 and external vision LLM tagging that stays aligned with the right regions.
|
||||
- Ordered image and mask loading, GPU Lanczos resizing, tiled VAE options, and provenance-aware latent tools.
|
||||
- WebUI-inspired sampler and scheduler extras including A1111 Euler ancestral behavior, AYS, GITS, `automatic_a1111`, and beta57.
|
||||
- WebUI-inspired sampling extras including seed variation, A1111 Euler ancestral behavior, AYS, GITS, `automatic_a1111`, and beta57.
|
||||
|
||||
## Contents
|
||||
|
||||
@@ -152,6 +152,8 @@ The external LLM nodes use a configured OpenAI-compatible provider. **Tag SEGS w
|
||||
|
||||
**KSampler (Extras)** adds the A1111/k-diffusion-style `euler_a_a1111` sampler, AYS SD1 and SDXL schedules, GITS, the `automatic_a1111` scheduler, and a local implementation of the RES4LYF beta57 preset. It keeps Comfy's regular seed handling, partial denoise behavior, progress callbacks, and conditioning inputs.
|
||||
|
||||
**Seed Variation** patches a MODEL so Comfy-native samplers mix their normal initial noise toward a second deterministic seed. Strength `0` keeps the sampler seed unchanged, while strength `1` uses variation-seed initial noise. Ancestral and SDE samplers continue to use the sampler seed for additional noise introduced after initialization.
|
||||
|
||||
The remaining utilities are **Latent Diagnostics**, **Scale Factor**, and **Seed**. Latent Diagnostics reports the latent shape, dtype, device, and tiled-sampling compatibility while passing it through unchanged.
|
||||
|
||||
## Settings and optional integrations
|
||||
@@ -187,6 +189,7 @@ SimpleSyrup owes a lot to other projects:
|
||||
- [ComfyUI Layer Style Advance](https://github.com/chflame163/ComfyUI_LayerStyle_Advance) provides the SAM model bundle SimpleSyrup can adapt.
|
||||
- [Tiled Diffusion & VAE for AUTOMATIC1111](https://github.com/pkuliyi2015/multidiffusion-upscaler-for-automatic1111) informed the practical tiled diffusion and Mixture of Diffusers behavior reimplemented here.
|
||||
- [RES4LYF](https://github.com/ClownsharkBatwing/RES4LYF) is the source of the beta57 scheduler preset reimplemented here.
|
||||
- [ComfyUI-ppm](https://github.com/pamparamm/ComfyUI-ppm) by pamparamm provides the ModelPatcher-based NegPiP behavior adapted here and builds on the [ComfyUI port](https://github.com/laksjdjf/cd-tuner_negpip-ComfyUI) by laksjdjf and the [original WebUI implementation](https://github.com/hako-mikan/sd-webui-negpip) by hako-mikan.
|
||||
|
||||
SimpleSyrup also vendors or reimplements selected third-party behavior for SAM-HQ, MobileSAM, GroundingDINO, AUTOMATIC1111 sampler behavior, k-diffusion, and tiled diffusion. See [third_party/NOTICE.md](third_party/NOTICE.md) for the complete notices.
|
||||
|
||||
|
||||
@@ -12,12 +12,21 @@ from . import simple_syrup as _simple_syrup_package
|
||||
|
||||
sys.modules.setdefault("simple_syrup", _simple_syrup_package)
|
||||
|
||||
from .simple_syrup.runtime.attention_region_prompt_handler import ( # noqa: E402
|
||||
register_attention_region_prompt_handler,
|
||||
)
|
||||
from .simple_syrup.runtime.comfy_safetensors_dtypes import ( # noqa: E402
|
||||
register_comfy_safetensors_dtypes,
|
||||
)
|
||||
from .simple_syrup.runtime.external_llm_routes import ( # noqa: E402
|
||||
register_external_llm_routes,
|
||||
)
|
||||
from .simple_syrup.runtime.mask_batch_preview_routes import ( # noqa: E402
|
||||
register_mask_batch_preview_routes,
|
||||
)
|
||||
from .simple_syrup.runtime.quant_cache_routes import ( # noqa: E402
|
||||
register_quant_cache_routes,
|
||||
)
|
||||
from .simple_syrup.runtime.settings_routes import register_settings_routes # noqa: E402
|
||||
|
||||
WEB_DIRECTORY = "./web/dist"
|
||||
@@ -42,8 +51,11 @@ async def comfy_entrypoint() -> object:
|
||||
|
||||
|
||||
register_settings_routes()
|
||||
register_comfy_safetensors_dtypes()
|
||||
register_quant_cache_routes()
|
||||
register_external_llm_routes()
|
||||
register_mask_batch_preview_routes()
|
||||
register_attention_region_prompt_handler()
|
||||
|
||||
__all__ = [
|
||||
"WEB_DIRECTORY",
|
||||
|
||||
|
Before Width: | Height: | Size: 251 KiB |
|
Before Width: | Height: | Size: 251 KiB |
|
Before Width: | Height: | Size: 8.1 KiB |
|
Before Width: | Height: | Size: 238 KiB |
|
Before Width: | Height: | Size: 234 KiB |
|
Before Width: | Height: | Size: 238 KiB |
|
Before Width: | Height: | Size: 113 KiB |
|
Before Width: | Height: | Size: 372 KiB |
|
Before Width: | Height: | Size: 111 KiB |
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.6.0",
|
||||
"version": "1.9.3",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.6.0",
|
||||
"version": "1.9.3",
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"devDependencies": {
|
||||
"@eslint/js": "^9.39.1",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.6.0",
|
||||
"version": "1.9.3",
|
||||
"private": true,
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"type": "module",
|
||||
|
||||
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
|
||||
[project]
|
||||
name = "SimpleSyrup"
|
||||
description = "Workflow-focused ComfyUI extensions for image generation."
|
||||
version = "1.6.0"
|
||||
version = "1.9.3"
|
||||
license = "AGPL-3.0-or-later"
|
||||
license-files = ["LICENSE"]
|
||||
requires-python = ">=3.11"
|
||||
@@ -43,8 +43,11 @@ extend-exclude = [
|
||||
select = ["E", "F", "I", "UP", "B", "C4", "ANN"]
|
||||
ignore = ["ANN401"]
|
||||
|
||||
[tool.ruff.lint.isort]
|
||||
known-first-party = ["simple_syrup"]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.11"
|
||||
python_version = "3.12"
|
||||
warn_return_any = true
|
||||
warn_unused_configs = true
|
||||
disallow_untyped_defs = true
|
||||
@@ -53,6 +56,7 @@ check_untyped_defs = true
|
||||
no_implicit_optional = true
|
||||
strict_equality = true
|
||||
explicit_package_bases = true
|
||||
mypy_path = ["tests"]
|
||||
exclude = [
|
||||
"simple_syrup/third_party/groundingdino_runtime",
|
||||
"simple_syrup/third_party/sam_hq_runtime",
|
||||
@@ -65,6 +69,9 @@ ignore_missing_imports = true
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = [".", "../.."]
|
||||
testpaths = ["tests"]
|
||||
markers = [
|
||||
"external_artifact: requires a locally installed external source or generated benchmark artifact",
|
||||
]
|
||||
filterwarnings = [
|
||||
"error",
|
||||
"ignore:builtin type SwigPyPacked has no __module__ attribute:DeprecationWarning",
|
||||
|
||||
@@ -6,6 +6,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__version__ = "1.6.0"
|
||||
__version__ = "1.9.3"
|
||||
|
||||
__all__: list[str] = ["__version__"]
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define versioned, quality-aware Anima quantization profiles."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .model_quantization import (
|
||||
QuantizationFormat,
|
||||
QuantizationProfile,
|
||||
TensorDescriptor,
|
||||
)
|
||||
|
||||
ORIGINAL_PROFILE = QuantizationProfile("original", "Original", 2, frozenset())
|
||||
FP8_E4M3_PROFILE = QuantizationProfile(
|
||||
"fp8-e4m3",
|
||||
"FP8 E4M3",
|
||||
2,
|
||||
frozenset({QuantizationFormat.FP8_E4M3}),
|
||||
)
|
||||
FP8_E5M2_PROFILE = QuantizationProfile(
|
||||
"fp8-e5m2",
|
||||
"FP8 E5M2",
|
||||
2,
|
||||
frozenset({QuantizationFormat.FP8_E5M2}),
|
||||
)
|
||||
MXFP8_PROFILE = QuantizationProfile(
|
||||
"mxfp8",
|
||||
"MXFP8",
|
||||
2,
|
||||
frozenset({QuantizationFormat.MXFP8}),
|
||||
)
|
||||
NVFP4_MIXED_PROFILE = QuantizationProfile(
|
||||
"nvfp4-mixed",
|
||||
"NVFP4 (Mixed)",
|
||||
3,
|
||||
frozenset({QuantizationFormat.FP8_E4M3, QuantizationFormat.NVFP4}),
|
||||
)
|
||||
_PROFILES = (
|
||||
ORIGINAL_PROFILE,
|
||||
FP8_E4M3_PROFILE,
|
||||
FP8_E5M2_PROFILE,
|
||||
MXFP8_PROFILE,
|
||||
NVFP4_MIXED_PROFILE,
|
||||
)
|
||||
_MAIN_BLOCK_PATTERN = re.compile(
|
||||
r"(?:^|\.)(?:net|diffusion_model)\.blocks\.(?P<index>\d+)\."
|
||||
)
|
||||
_PROTECTED_BLOCKS = {0, 1, 27}
|
||||
_FLOAT_DTYPES = {"F16", "BF16", "F32", "F64"}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AnimaQuantizationRecipe:
|
||||
"""Assign formats only within Anima's quality-safe DiT block envelope."""
|
||||
|
||||
model_family: str = "Anima"
|
||||
version: int = 2
|
||||
|
||||
@property
|
||||
def profiles(self) -> tuple[QuantizationProfile, ...]:
|
||||
"""Return Anima's stable workflow-facing profile order."""
|
||||
|
||||
return _PROFILES
|
||||
|
||||
def profile_from_selection(self, selection: str) -> QuantizationProfile:
|
||||
"""Parse one current workflow selection into its profile."""
|
||||
|
||||
for profile in self.profiles:
|
||||
if selection in (profile.label, profile.profile_id):
|
||||
return profile
|
||||
valid = ", ".join(profile.label for profile in self.profiles)
|
||||
raise ValueError(f"quantization profile must be one of: {valid}.")
|
||||
|
||||
def policy_for(
|
||||
self,
|
||||
tensor: TensorDescriptor,
|
||||
profile: QuantizationProfile,
|
||||
) -> QuantizationFormat | None:
|
||||
"""Return Anima's per-tensor format while preserving sensitive layers."""
|
||||
|
||||
if profile not in self.profiles:
|
||||
raise ValueError(
|
||||
f"Unknown Anima quantization profile '{profile.profile_id}'."
|
||||
)
|
||||
if profile.is_original or not _is_matrix_weight(tensor):
|
||||
return None
|
||||
if "llm_adapter" in tensor.name or "adaln_modulation" in tensor.name:
|
||||
return None
|
||||
block_match = _MAIN_BLOCK_PATTERN.search(tensor.name)
|
||||
if block_match is None:
|
||||
return None
|
||||
if int(block_match.group("index")) in _PROTECTED_BLOCKS:
|
||||
return None
|
||||
if profile.profile_id == NVFP4_MIXED_PROFILE.profile_id:
|
||||
if "v_proj" in tensor.name or ".mlp." in tensor.name:
|
||||
return QuantizationFormat.FP8_E4M3
|
||||
if any(
|
||||
projection in tensor.name
|
||||
for projection in ("q_proj", "k_proj", "output_proj")
|
||||
):
|
||||
return QuantizationFormat.NVFP4
|
||||
return None
|
||||
return next(iter(profile.required_formats))
|
||||
|
||||
|
||||
def _is_matrix_weight(tensor: TensorDescriptor) -> bool:
|
||||
"""Return whether a tensor is an eligible floating-point matrix weight."""
|
||||
|
||||
return (
|
||||
tensor.dtype_name in _FLOAT_DTYPES
|
||||
and len(tensor.shape) == 2
|
||||
and tensor.name.endswith(".weight")
|
||||
)
|
||||
@@ -0,0 +1,15 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Parse explicit attention concepts without interpreting prompt language."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def parse_attention_concepts(value: str) -> tuple[str, ...]:
|
||||
"""Return canonical concepts separated only by vertical bars."""
|
||||
|
||||
if not isinstance(value, str):
|
||||
raise TypeError("Attention concepts must be text.")
|
||||
return tuple(part.strip() for part in value.split("|") if part.strip())
|
||||
@@ -0,0 +1,46 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Classify complete Attention Coupling requests before runtime preparation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
from .conditioning_batch import ConditioningBatch
|
||||
|
||||
|
||||
class AttentionCouplingRequestMode(StrEnum):
|
||||
"""Select ordinary sampling or complete regional Attention Coupling."""
|
||||
|
||||
BYPASS = "bypass"
|
||||
ACTIVE = "active"
|
||||
|
||||
|
||||
def classify_attention_coupling_request(
|
||||
*,
|
||||
positive: object,
|
||||
negative: object,
|
||||
region_masks: object | None,
|
||||
) -> AttentionCouplingRequestMode:
|
||||
"""Return the execution mode or reject a partial regional request."""
|
||||
|
||||
has_conditioning_batch = isinstance(positive, ConditioningBatch) or isinstance(
|
||||
negative,
|
||||
ConditioningBatch,
|
||||
)
|
||||
has_region_masks = region_masks is not None
|
||||
if not has_conditioning_batch and not has_region_masks:
|
||||
return AttentionCouplingRequestMode.BYPASS
|
||||
if has_conditioning_batch and has_region_masks:
|
||||
return AttentionCouplingRequestMode.ACTIVE
|
||||
if has_conditioning_batch:
|
||||
raise ValueError(
|
||||
"Attention Coupling conditioning batches require region_masks. "
|
||||
"Connect ordered masks or use ordinary CONDITIONING on both inputs."
|
||||
)
|
||||
raise ValueError(
|
||||
"Attention Coupling region_masks require a CONDITIONING_BATCH on the "
|
||||
"positive or negative input. Disconnect the masks for ordinary sampling."
|
||||
)
|
||||
@@ -0,0 +1,30 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve explicit two-dimensional geometry for flattened attention maps."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
|
||||
def factor_spatial_geometry(
|
||||
token_count: int, *, target_aspect: float
|
||||
) -> tuple[int, int]:
|
||||
"""Return the token factor pair nearest the supplied positive aspect ratio."""
|
||||
|
||||
if type(token_count) is not int or token_count < 1:
|
||||
raise ValueError("Attention spatial token count must be positive.")
|
||||
if target_aspect <= 0.0:
|
||||
raise ValueError("Attention target aspect ratio must be positive.")
|
||||
candidates: list[tuple[float, int, int]] = []
|
||||
for height in range(1, math.isqrt(token_count) + 1):
|
||||
if token_count % height:
|
||||
continue
|
||||
width = token_count // height
|
||||
for candidate_height, candidate_width in ((height, width), (width, height)):
|
||||
error = abs(math.log((candidate_width / candidate_height) / target_aspect))
|
||||
candidates.append((error, candidate_height, candidate_width))
|
||||
_error, height, width = min(candidates)
|
||||
return height, width
|
||||
@@ -0,0 +1,175 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define immutable requests and plans for attention-region capture."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
from .graph_provenance import GraphLink
|
||||
|
||||
|
||||
class AttentionRegionRequestKind(StrEnum):
|
||||
"""Identify one public attention-region operation."""
|
||||
|
||||
CONCEPT_SEGS = "concept_segs"
|
||||
ALL_PROMPT_SEGS = "all_prompt_segs"
|
||||
REGION_MASK = "region_mask"
|
||||
MASKED_CONDITIONING = "masked_conditioning"
|
||||
|
||||
|
||||
class AttentionCaptureProfile(StrEnum):
|
||||
"""Select the density of attention observations retained during sampling."""
|
||||
|
||||
FAST = "fast"
|
||||
BALANCED = "balanced"
|
||||
EXHAUSTIVE = "exhaustive"
|
||||
|
||||
|
||||
class AttentionEvidenceMode(StrEnum):
|
||||
"""Select honest inspection or derived concept-isolation evidence."""
|
||||
|
||||
CONCEPT = "concept isolation"
|
||||
RAW = "raw attention"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionRegionControls:
|
||||
"""Hold validated attention-native capture and region-shaping controls."""
|
||||
|
||||
capture_start: float
|
||||
capture_end: float
|
||||
minimum_strength: float
|
||||
minimum_consensus: float
|
||||
split_sensitivity: float
|
||||
minimum_region_size: int
|
||||
profile: AttentionCaptureProfile
|
||||
instance_recall: float = 0.65
|
||||
geometry_recall: float = 0.85
|
||||
keep_only: int = 0
|
||||
keep_by: str = "largest size"
|
||||
combine_segs: bool = False
|
||||
matte_solidity: float = 0.0
|
||||
edge_feather: int = 8
|
||||
evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require normalized ranges and a non-empty capture interval."""
|
||||
|
||||
normalized = (
|
||||
self.capture_start,
|
||||
self.capture_end,
|
||||
self.minimum_strength,
|
||||
self.minimum_consensus,
|
||||
self.split_sensitivity,
|
||||
self.instance_recall,
|
||||
self.geometry_recall,
|
||||
self.matte_solidity,
|
||||
)
|
||||
if any(
|
||||
isinstance(value, bool) or not isinstance(value, int | float)
|
||||
for value in normalized
|
||||
):
|
||||
raise TypeError("Attention-region controls must be real numbers.")
|
||||
if not 0.0 <= self.capture_start < self.capture_end <= 1.0:
|
||||
raise ValueError("Attention capture start must be below end within 0..1.")
|
||||
if not 0.0 <= self.minimum_strength <= 1.0:
|
||||
raise ValueError("Minimum attention strength must be within 0..1.")
|
||||
if not 0.0 <= self.minimum_consensus <= 1.0:
|
||||
raise ValueError("Minimum attention consensus must be within 0..1.")
|
||||
if not 0.0 <= self.split_sensitivity <= 1.0:
|
||||
raise ValueError("Attention split sensitivity must be within 0..1.")
|
||||
if not 0.0 <= self.instance_recall <= 1.0:
|
||||
raise ValueError("Attention instance recall must be within 0..1.")
|
||||
if not 0.0 <= self.geometry_recall <= 1.0:
|
||||
raise ValueError("Attention geometry recall must be within 0..1.")
|
||||
if type(self.minimum_region_size) is not int or self.minimum_region_size < 1:
|
||||
raise ValueError("Minimum attention region size must be positive.")
|
||||
if type(self.keep_only) is not int or self.keep_only < 0:
|
||||
raise ValueError("Attention keep_only must be non-negative.")
|
||||
if self.keep_by not in ("largest size", "highest confidence"):
|
||||
raise ValueError("Attention keep_by has an invalid policy.")
|
||||
if type(self.combine_segs) is not bool:
|
||||
raise TypeError("Attention combine_segs must be boolean.")
|
||||
if not 0.0 <= self.matte_solidity <= 1.0:
|
||||
raise ValueError("Attention matte solidity must be within 0..1.")
|
||||
if type(self.edge_feather) is not int or self.edge_feather < 0:
|
||||
raise ValueError("Attention edge feather must be non-negative.")
|
||||
if not isinstance(self.profile, AttentionCaptureProfile):
|
||||
raise TypeError("Attention capture profile has an invalid type.")
|
||||
if not isinstance(self.evidence_mode, AttentionEvidenceMode):
|
||||
raise TypeError("Attention evidence mode has an invalid type.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionRegionRequest:
|
||||
"""Bind one public node request to its queries and capture controls."""
|
||||
|
||||
node_id: str
|
||||
kind: AttentionRegionRequestKind
|
||||
queries: tuple[str, ...]
|
||||
controls: AttentionRegionControls
|
||||
sampler_stage: int = 1
|
||||
spatial_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require stable node identity and canonical non-empty query strings."""
|
||||
|
||||
if not self.node_id.strip():
|
||||
raise ValueError("Attention-region request node id cannot be empty.")
|
||||
if not isinstance(self.kind, AttentionRegionRequestKind):
|
||||
raise TypeError("Attention-region request kind has an invalid type.")
|
||||
if any(not query or query != query.strip() for query in self.queries):
|
||||
raise ValueError("Attention-region queries must be canonical strings.")
|
||||
if self.kind is AttentionRegionRequestKind.ALL_PROMPT_SEGS:
|
||||
if self.queries:
|
||||
raise ValueError(
|
||||
"All-prompt attention requests cannot contain queries."
|
||||
)
|
||||
elif not self.queries:
|
||||
raise ValueError("Concept and mask attention requests require concepts.")
|
||||
if type(self.sampler_stage) is not int or self.sampler_stage < -1:
|
||||
raise ValueError("Attention sampler stage must be -1 or greater.")
|
||||
if any(
|
||||
not isinstance(transform, AttentionSpatialTransform)
|
||||
for transform in self.spatial_transforms
|
||||
):
|
||||
raise TypeError("Attention request spatial transforms are invalid.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionCapturePlan:
|
||||
"""Describe one coalesced sampler capture and its graph rewrite authority."""
|
||||
|
||||
sampler_node_id: str
|
||||
model_owner_node_id: str
|
||||
model_input_name: str
|
||||
model_link: GraphLink
|
||||
positive_link: GraphLink
|
||||
requests: tuple[AttentionRegionRequest, ...]
|
||||
prompt_text: str | None = None
|
||||
clip_link: GraphLink | None = None
|
||||
source_aspect: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require canonical unique requests and complete graph-edge identity."""
|
||||
|
||||
if not self.sampler_node_id or not self.model_owner_node_id:
|
||||
raise ValueError("Attention capture plan node ids cannot be empty.")
|
||||
if not self.model_input_name:
|
||||
raise ValueError("Attention capture plan model input cannot be empty.")
|
||||
if self.source_aspect is not None and self.source_aspect <= 0.0:
|
||||
raise ValueError("Attention capture source aspect must be positive.")
|
||||
request_ids = tuple(request.node_id for request in self.requests)
|
||||
if not request_ids or request_ids != tuple(sorted(set(request_ids))):
|
||||
raise ValueError("Attention capture requests must be unique and ordered.")
|
||||
|
||||
@property
|
||||
def capture_node_id(self) -> str:
|
||||
"""Return a collision-resistant deterministic injected node id."""
|
||||
|
||||
return f"__simple_syrup_attention_capture__{self.sampler_node_id}"
|
||||
@@ -0,0 +1,21 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define rendered attention evidence before component and matte shaping."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionConceptEvidence:
|
||||
"""Hold one concept's alpha, support, and confidence evidence."""
|
||||
|
||||
label: str
|
||||
alpha: torch.Tensor
|
||||
support: torch.Tensor
|
||||
confidence: torch.Tensor
|
||||
@@ -0,0 +1,198 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Model prompt-token spans and compact captured attention observations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
from .regional_model_capabilities import RegionalModelFamily
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionTokenSpan:
|
||||
"""Bind a readable prompt occurrence to exact conditioning token positions."""
|
||||
|
||||
label: str
|
||||
occurrence: int
|
||||
token_indices: tuple[int, ...]
|
||||
head_token_indices: tuple[int, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require canonical labels and ordered non-negative token indices."""
|
||||
|
||||
if not self.label or self.label != self.label.strip():
|
||||
raise ValueError("Attention token span label must be canonical.")
|
||||
if type(self.occurrence) is not int or self.occurrence < 1:
|
||||
raise ValueError("Attention token span occurrence must be positive.")
|
||||
if (
|
||||
not self.token_indices
|
||||
or self.token_indices != tuple(sorted(set(self.token_indices)))
|
||||
or self.token_indices[0] < 0
|
||||
):
|
||||
raise ValueError("Attention token indices must be ordered and unique.")
|
||||
if self.head_token_indices and (
|
||||
self.head_token_indices != tuple(sorted(set(self.head_token_indices)))
|
||||
or not set(self.head_token_indices).issubset(self.token_indices)
|
||||
):
|
||||
raise ValueError("Attention head tokens must be an ordered span subset.")
|
||||
|
||||
@property
|
||||
def semantic_head_indices(self) -> tuple[int, ...]:
|
||||
"""Return explicit noun-head positions or a safe final-token fallback."""
|
||||
|
||||
return self.head_token_indices or self.token_indices[-1:]
|
||||
|
||||
@property
|
||||
def display_label(self) -> str:
|
||||
"""Disambiguate repeated concepts while keeping first labels concise."""
|
||||
|
||||
return (
|
||||
self.label if self.occurrence == 1 else f"{self.label} #{self.occurrence}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionTokenCatalog:
|
||||
"""Hold readable prompt spans and conditioning sequence length."""
|
||||
|
||||
sequence_length: int
|
||||
spans: tuple[AttentionTokenSpan, ...]
|
||||
token_ids: tuple[object, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require all spans to fit the captured conditioning sequence."""
|
||||
|
||||
if type(self.sequence_length) is not int or self.sequence_length < 1:
|
||||
raise ValueError("Attention token sequence length must be positive.")
|
||||
if len(self.token_ids) != self.sequence_length:
|
||||
raise ValueError("Attention token ids must match the sequence length.")
|
||||
if any(
|
||||
index >= self.sequence_length
|
||||
for span in self.spans
|
||||
for index in span.token_indices
|
||||
):
|
||||
raise ValueError("Attention token span exceeds its conditioning sequence.")
|
||||
|
||||
def exact_matches(self, query: str) -> tuple[AttentionTokenSpan, ...]:
|
||||
"""Return every prompt occurrence whose normalized label equals a query."""
|
||||
|
||||
normalized = _normalized_label(query)
|
||||
return tuple(
|
||||
span for span in self.spans if _normalized_label(span.label) == normalized
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapturedAttentionMap:
|
||||
"""Store one head-aggregated token map and its denoising observation identity."""
|
||||
|
||||
label: str
|
||||
values: torch.Tensor
|
||||
progress: float
|
||||
layer_key: str
|
||||
batch_index: int = 0
|
||||
confidence: float = 1.0
|
||||
spatial_height: int | None = None
|
||||
spatial_width: int | None = None
|
||||
spatial_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
concept_values: torch.Tensor | None = None
|
||||
uniform_probability: float = 0.0
|
||||
model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require a finite CPU spatial vector and normalized progress."""
|
||||
|
||||
if not self.label:
|
||||
raise ValueError("Captured attention map label cannot be empty.")
|
||||
if (
|
||||
not isinstance(self.values, torch.Tensor)
|
||||
or self.values.device.type != "cpu"
|
||||
or self.values.ndim != 1
|
||||
or self.values.numel() < 1
|
||||
or not self.values.is_floating_point()
|
||||
or not torch.isfinite(self.values).all().item()
|
||||
):
|
||||
raise ValueError("Captured attention values must be a finite CPU vector.")
|
||||
if not 0.0 <= self.progress <= 1.0:
|
||||
raise ValueError("Captured attention progress must be within 0..1.")
|
||||
if not self.layer_key:
|
||||
raise ValueError("Captured attention layer key cannot be empty.")
|
||||
if type(self.batch_index) is not int or self.batch_index < 0:
|
||||
raise ValueError("Captured attention batch index must be non-negative.")
|
||||
if not 0.0 <= self.confidence <= 1.0:
|
||||
raise ValueError("Captured attention confidence must be within 0..1.")
|
||||
if (self.spatial_height is None) != (self.spatial_width is None):
|
||||
raise ValueError("Captured attention geometry must be complete or absent.")
|
||||
if self.spatial_height is not None and (
|
||||
type(self.spatial_height) is not int
|
||||
or self.spatial_height < 1
|
||||
or type(self.spatial_width) is not int
|
||||
or self.spatial_width < 1
|
||||
or self.spatial_height * self.spatial_width != int(self.values.numel())
|
||||
):
|
||||
raise ValueError("Captured attention geometry must match its values.")
|
||||
if any(
|
||||
not isinstance(transform, AttentionSpatialTransform)
|
||||
for transform in self.spatial_transforms
|
||||
):
|
||||
raise TypeError("Captured attention spatial transforms are invalid.")
|
||||
if self.concept_values is not None and (
|
||||
not isinstance(self.concept_values, torch.Tensor)
|
||||
or self.concept_values.device.type != "cpu"
|
||||
or self.concept_values.shape != self.values.shape
|
||||
or not self.concept_values.is_floating_point()
|
||||
or not torch.isfinite(self.concept_values).all().item()
|
||||
):
|
||||
raise ValueError(
|
||||
"Captured concept evidence must match its finite CPU attention map."
|
||||
)
|
||||
if (
|
||||
isinstance(self.uniform_probability, bool)
|
||||
or not isinstance(self.uniform_probability, int | float)
|
||||
or not 0.0 <= float(self.uniform_probability) <= 1.0
|
||||
):
|
||||
raise ValueError("Captured uniform probability must be within 0..1.")
|
||||
if not isinstance(self.model_family, RegionalModelFamily):
|
||||
raise TypeError("Captured attention model family has an invalid type.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpenVocabularyContext:
|
||||
"""Hold one query's encoded SDXL context and semantic token positions."""
|
||||
|
||||
label: str
|
||||
values: torch.Tensor
|
||||
token_indices: tuple[int, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require one finite CPU context with valid unique token positions."""
|
||||
|
||||
if not self.label.strip() or self.label != self.label.strip():
|
||||
raise ValueError("Open-vocabulary labels must be canonical strings.")
|
||||
if self.values.ndim != 3 or int(self.values.shape[0]) != 1:
|
||||
raise ValueError("Open-vocabulary context must have shape 1xTxC.")
|
||||
if self.values.device.type != "cpu" or not torch.isfinite(self.values).all():
|
||||
raise ValueError("Open-vocabulary context must be finite CPU storage.")
|
||||
if not self.token_indices or self.token_indices != tuple(
|
||||
sorted(set(self.token_indices))
|
||||
):
|
||||
raise ValueError(
|
||||
"Open-vocabulary token positions must be unique and ordered."
|
||||
)
|
||||
if any(
|
||||
index < 0 or index >= int(self.values.shape[1])
|
||||
for index in self.token_indices
|
||||
):
|
||||
raise ValueError("Open-vocabulary token positions exceed their context.")
|
||||
|
||||
|
||||
def _normalized_label(value: str) -> str:
|
||||
"""Normalize human prompt labels without changing tokenizer semantics."""
|
||||
|
||||
return " ".join(value.casefold().replace("_", " ").split())
|
||||
@@ -0,0 +1,78 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Model ordered sampling stages along one spatial provenance path."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .attention_spatial_transform import AttentionSpatialTransform
|
||||
from .graph_provenance import GraphLink
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerStage:
|
||||
"""Describe one sampler's patch and conditioning authority."""
|
||||
|
||||
sampler_node_id: str
|
||||
model_owner_node_id: str
|
||||
model_link: GraphLink
|
||||
positive_link: GraphLink
|
||||
upstream_link: GraphLink
|
||||
upstream_kind: str
|
||||
forward_transforms: tuple[AttentionSpatialTransform, ...] = ()
|
||||
source_aspect: float | None = None
|
||||
capture_supported: bool = True
|
||||
unsupported_reason: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require unsupported stages to explain why capture cannot be projected."""
|
||||
|
||||
if self.capture_supported and self.unsupported_reason is not None:
|
||||
raise ValueError("Supported attention sampler stages cannot have a reason.")
|
||||
if not self.capture_supported and not self.unsupported_reason:
|
||||
raise ValueError("Unsupported attention sampler stages require a reason.")
|
||||
if self.source_aspect is not None and self.source_aspect <= 0.0:
|
||||
raise ValueError("Attention sampler source aspect must be positive.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerSelection:
|
||||
"""Bind a selected stage to its one-based chronological position."""
|
||||
|
||||
stage: AttentionSamplerStage
|
||||
stage_number: int
|
||||
stage_count: int
|
||||
was_clamped: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSamplerLineage:
|
||||
"""Hold sampling stages ordered from oldest to direct provenance."""
|
||||
|
||||
stages: tuple[AttentionSamplerStage, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require at least one uniquely identified stage."""
|
||||
|
||||
identities = tuple(stage.sampler_node_id for stage in self.stages)
|
||||
if not identities or len(set(identities)) != len(identities):
|
||||
raise ValueError("Attention sampler lineage must contain unique stages.")
|
||||
|
||||
def select(self, requested_stage: int) -> AttentionSamplerSelection:
|
||||
"""Resolve one-based selection with 0/-1 aliases for direct provenance."""
|
||||
|
||||
if type(requested_stage) is not int or requested_stage < -1:
|
||||
raise ValueError("Attention sampler stage must be -1 or greater.")
|
||||
count = len(self.stages)
|
||||
if requested_stage in (-1, 0):
|
||||
return AttentionSamplerSelection(self.stages[-1], count, count, False)
|
||||
selected_number = min(requested_stage, count)
|
||||
return AttentionSamplerSelection(
|
||||
self.stages[selected_number - 1],
|
||||
selected_number,
|
||||
count,
|
||||
requested_stage > count,
|
||||
)
|
||||
@@ -0,0 +1,70 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Describe graph-visible full-canvas transformations for attention masks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
_ANCHORS = frozenset(
|
||||
{
|
||||
"center",
|
||||
"top-left",
|
||||
"top",
|
||||
"top-right",
|
||||
"left",
|
||||
"right",
|
||||
"bottom-left",
|
||||
"bottom",
|
||||
"bottom-right",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class AttentionSpatialTransformKind(StrEnum):
|
||||
"""Identify a supported mask-coordinate transformation."""
|
||||
|
||||
RESIZE = "resize"
|
||||
FIT_RESIZE = "fit_resize"
|
||||
SCALE = "scale"
|
||||
COVER_CROP = "cover_crop"
|
||||
FIT_PAD = "fit_pad"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionSpatialTransform:
|
||||
"""Hold validated resize, crop, or pad parameters from one graph node."""
|
||||
|
||||
kind: AttentionSpatialTransformKind
|
||||
width: int | None = None
|
||||
height: int | None = None
|
||||
scale: float | None = None
|
||||
anchor: str = "center"
|
||||
divisible_by: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require complete parameters for the selected transformation kind."""
|
||||
|
||||
if self.anchor not in _ANCHORS:
|
||||
raise ValueError("Attention spatial transform anchor is invalid.")
|
||||
if self.kind is AttentionSpatialTransformKind.SCALE:
|
||||
if self.scale is None or self.scale <= 0.0:
|
||||
raise ValueError("Attention scale transform requires a positive scale.")
|
||||
if self.width is not None or self.height is not None:
|
||||
raise ValueError("Attention scale transform cannot contain a size.")
|
||||
if self.divisible_by != 1:
|
||||
raise ValueError("Attention scale transform cannot set divisibility.")
|
||||
return
|
||||
if (
|
||||
type(self.width) is not int
|
||||
or self.width < 1
|
||||
or type(self.height) is not int
|
||||
or self.height < 1
|
||||
or self.scale is not None
|
||||
):
|
||||
raise ValueError("Attention spatial transform requires a positive size.")
|
||||
if type(self.divisible_by) is not int or self.divisible_by < 1:
|
||||
raise ValueError("Attention spatial divisibility must be positive.")
|
||||
@@ -0,0 +1,90 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own immutable authored and model-converted conditioning schedule bounds."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ConditioningScheduleRange:
|
||||
"""Retain optional authored percentages and converted sigma boundaries."""
|
||||
|
||||
start_percent: float | None
|
||||
end_percent: float | None
|
||||
timestep_start: float | None
|
||||
timestep_end: float | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate exact optional bounds without inventing absent metadata."""
|
||||
|
||||
start_percent = _optional_finite_float(
|
||||
self.start_percent,
|
||||
name="start_percent",
|
||||
)
|
||||
end_percent = _optional_finite_float(
|
||||
self.end_percent,
|
||||
name="end_percent",
|
||||
)
|
||||
timestep_start = _optional_finite_float(
|
||||
self.timestep_start,
|
||||
name="timestep_start",
|
||||
)
|
||||
timestep_end = _optional_finite_float(
|
||||
self.timestep_end,
|
||||
name="timestep_end",
|
||||
)
|
||||
for name, value in (
|
||||
("start_percent", start_percent),
|
||||
("end_percent", end_percent),
|
||||
):
|
||||
if value is not None and not 0.0 <= value <= 1.0:
|
||||
raise ValueError(f"Conditioning {name} must be in [0, 1].")
|
||||
effective_start = 0.0 if start_percent is None else start_percent
|
||||
effective_end = 1.0 if end_percent is None else end_percent
|
||||
if effective_start > effective_end:
|
||||
raise ValueError("Conditioning start_percent must not exceed end_percent.")
|
||||
if (
|
||||
timestep_start is not None
|
||||
and timestep_end is not None
|
||||
and timestep_start < timestep_end
|
||||
):
|
||||
raise ValueError(
|
||||
"Conditioning timestep_start must not be below timestep_end."
|
||||
)
|
||||
object.__setattr__(self, "start_percent", start_percent)
|
||||
object.__setattr__(self, "end_percent", end_percent)
|
||||
object.__setattr__(self, "timestep_start", timestep_start)
|
||||
object.__setattr__(self, "timestep_end", timestep_end)
|
||||
|
||||
@property
|
||||
def is_time_invariant(self) -> bool:
|
||||
"""Report whether this entry remains admitted for the whole trajectory."""
|
||||
|
||||
if self.start_percent is None and self.timestep_start is not None:
|
||||
return False
|
||||
if self.end_percent is None and self.timestep_end is not None:
|
||||
return False
|
||||
effective_start = 0.0 if self.start_percent is None else self.start_percent
|
||||
effective_end = 1.0 if self.end_percent is None else self.end_percent
|
||||
return effective_start == 0.0 and effective_end == 1.0
|
||||
|
||||
|
||||
def _optional_finite_float(value: object, *, name: str) -> float | None:
|
||||
"""Normalize one optional real boundary without accepting booleans."""
|
||||
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int | float):
|
||||
raise TypeError(f"Conditioning {name} must be a real number or None.")
|
||||
normalized = float(value)
|
||||
if not math.isfinite(normalized):
|
||||
raise ValueError(f"Conditioning {name} must be finite.")
|
||||
return normalized
|
||||
|
||||
|
||||
UNBOUNDED_CONDITIONING_SCHEDULE = ConditioningScheduleRange(None, None, None, None)
|
||||
@@ -0,0 +1,49 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own installed-Comfy-equivalent conditioning schedule admission."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
from .conditioning_schedule import ConditioningScheduleRange
|
||||
|
||||
|
||||
class ConditioningScheduleSelectionPolicy:
|
||||
"""Match installed Comfy's inclusive converted-sigma admission policy."""
|
||||
|
||||
@staticmethod
|
||||
def is_active(
|
||||
schedule: ConditioningScheduleRange,
|
||||
*,
|
||||
sigma: float,
|
||||
) -> bool:
|
||||
"""Return installed Comfy's inclusive start/end decision."""
|
||||
|
||||
if not isinstance(schedule, ConditioningScheduleRange):
|
||||
raise TypeError("Conditioning selection requires a schedule range.")
|
||||
current_sigma = normalize_conditioning_sigma(sigma)
|
||||
if (
|
||||
schedule.timestep_start is not None
|
||||
and current_sigma > schedule.timestep_start
|
||||
):
|
||||
return False
|
||||
return not (
|
||||
schedule.timestep_end is not None and current_sigma < schedule.timestep_end
|
||||
)
|
||||
|
||||
|
||||
def normalize_conditioning_sigma(value: object) -> float:
|
||||
"""Normalize one finite real current sigma without accepting booleans."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int | float):
|
||||
raise TypeError("Conditioning selection sigma must be a real number.")
|
||||
sigma = float(value)
|
||||
if not math.isfinite(sigma):
|
||||
raise ValueError("Conditioning selection sigma must be finite.")
|
||||
return sigma
|
||||
|
||||
|
||||
CONDITIONING_SCHEDULE_SELECTION_POLICY = ConditioningScheduleSelectionPolicy()
|
||||
@@ -8,8 +8,12 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from .regional_tiled_diffusion import build_region_constrained_tiled_diffusion_plan
|
||||
from .segs import NativeSegs
|
||||
from .segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan
|
||||
from .spatial_views import SpatialView, SpatialViewKind
|
||||
from .tiled_diffusion import TiledDiffusionPlan, build_tiled_diffusion_plan
|
||||
|
||||
|
||||
@@ -44,25 +48,13 @@ class ContextualDiffusionControls:
|
||||
raise ValueError("global_decay must be between 0 and 1.")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SpatialContext:
|
||||
"""Describe one source rectangle evaluated at a bounded model context shape."""
|
||||
|
||||
x: int
|
||||
y: int
|
||||
width: int
|
||||
height: int
|
||||
context_width: int
|
||||
context_height: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ContextualDiffusionPlan:
|
||||
"""Own the global context and sole tiled plan for one latent canvas."""
|
||||
|
||||
latent_width: int
|
||||
latent_height: int
|
||||
global_context: SpatialContext
|
||||
global_view: SpatialView
|
||||
tile_plan: TiledDiffusionPlan
|
||||
|
||||
|
||||
@@ -72,6 +64,7 @@ def build_contextual_diffusion_plan(
|
||||
latent_height: int,
|
||||
controls: ContextualDiffusionControls,
|
||||
segs: NativeSegs | None,
|
||||
region_masks: torch.Tensor | None = None,
|
||||
) -> ContextualDiffusionPlan:
|
||||
"""Return a global context plus the regular or SEGS-guided context plan."""
|
||||
|
||||
@@ -81,16 +74,18 @@ def build_contextual_diffusion_plan(
|
||||
latent_height,
|
||||
controls.latent_context_size,
|
||||
)
|
||||
global_context = SpatialContext(
|
||||
x=0,
|
||||
y=0,
|
||||
width=latent_width,
|
||||
height=latent_height,
|
||||
context_width=global_width,
|
||||
context_height=global_height,
|
||||
global_view = SpatialView(
|
||||
kind=SpatialViewKind.CONTEXTUAL_GLOBAL,
|
||||
source_x=0,
|
||||
source_y=0,
|
||||
source_width=latent_width,
|
||||
source_height=latent_height,
|
||||
model_width=global_width,
|
||||
model_height=global_height,
|
||||
)
|
||||
tile_plan = (
|
||||
build_segs_guided_tiled_diffusion_plan(
|
||||
if region_masks is not None:
|
||||
tile_plan = build_region_constrained_tiled_diffusion_plan(
|
||||
region_masks=region_masks,
|
||||
segs=segs,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
@@ -99,8 +94,18 @@ def build_contextual_diffusion_plan(
|
||||
overlap=controls.latent_context_overlap,
|
||||
tile_batch_size=controls.latent_context_batch_size,
|
||||
)
|
||||
if segs is not None
|
||||
else build_tiled_diffusion_plan(
|
||||
elif segs is not None:
|
||||
tile_plan = build_segs_guided_tiled_diffusion_plan(
|
||||
segs=segs,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=controls.latent_context_size,
|
||||
tile_height=controls.latent_context_size,
|
||||
overlap=controls.latent_context_overlap,
|
||||
tile_batch_size=controls.latent_context_batch_size,
|
||||
)
|
||||
else:
|
||||
tile_plan = build_tiled_diffusion_plan(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=controls.latent_context_size,
|
||||
@@ -108,11 +113,10 @@ def build_contextual_diffusion_plan(
|
||||
overlap=controls.latent_context_overlap,
|
||||
tile_batch_size=controls.latent_context_batch_size,
|
||||
)
|
||||
)
|
||||
return ContextualDiffusionPlan(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
global_context=global_context,
|
||||
global_view=global_view,
|
||||
tile_plan=tile_plan,
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Select decayed reduced-global authority from denoising timesteps."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class GlobalContextSchedule:
|
||||
"""Limit whole-image authority to an initial denoising-step fraction."""
|
||||
|
||||
def __init__(
|
||||
self, *, sigmas: torch.Tensor, active_steps: int, decay: float
|
||||
) -> None:
|
||||
"""Capture model-evaluation sigmas and the active initial step count."""
|
||||
|
||||
step_sigmas = sigmas.detach().to(device="cpu", dtype=torch.float64).flatten()
|
||||
if step_sigmas.numel() < 2:
|
||||
raise ValueError(
|
||||
"Contextual Diffusion requires at least one denoising step."
|
||||
)
|
||||
self._step_sigmas = step_sigmas[:-1]
|
||||
self._active_steps = min(len(self._step_sigmas), max(0, active_steps))
|
||||
self._decay = decay
|
||||
|
||||
def scale_for(self, timestep: object) -> float:
|
||||
"""Return the decayed global scale for the nearest scheduled step."""
|
||||
|
||||
if self._active_steps == 0:
|
||||
return 0.0
|
||||
if not isinstance(timestep, torch.Tensor) or timestep.numel() == 0:
|
||||
raise ValueError(
|
||||
"Contextual Diffusion timestep must be a non-empty tensor."
|
||||
)
|
||||
sigma = timestep.detach().flatten()[0].to(device="cpu", dtype=torch.float64)
|
||||
step_index = int(torch.argmin(torch.abs(self._step_sigmas - sigma)).item())
|
||||
if step_index >= self._active_steps:
|
||||
return 0.0
|
||||
return self._decay**step_index
|
||||
@@ -0,0 +1,82 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define model-independent checkpoint quantization profile contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
class QuantizationFormat(StrEnum):
|
||||
"""Identify one reusable ComfyUI tensor quantization format."""
|
||||
|
||||
FP8_E4M3 = "float8_e4m3fn"
|
||||
FP8_E5M2 = "float8_e5m2"
|
||||
NVFP4 = "nvfp4"
|
||||
MXFP8 = "mxfp8"
|
||||
|
||||
@property
|
||||
def label(self) -> str:
|
||||
"""Return the concise format label used in diagnostics."""
|
||||
|
||||
return {
|
||||
QuantizationFormat.FP8_E4M3: "FP8 E4M3",
|
||||
QuantizationFormat.FP8_E5M2: "FP8 E5M2",
|
||||
QuantizationFormat.NVFP4: "NVFP4",
|
||||
QuantizationFormat.MXFP8: "MXFP8",
|
||||
}[self]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantizationProfile:
|
||||
"""Describe one workflow-facing, versioned per-tensor policy profile."""
|
||||
|
||||
profile_id: str
|
||||
label: str
|
||||
version: int
|
||||
required_formats: frozenset[QuantizationFormat]
|
||||
|
||||
@property
|
||||
def is_original(self) -> bool:
|
||||
"""Return whether this profile loads the source checkpoint unchanged."""
|
||||
|
||||
return not self.required_formats
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TensorDescriptor:
|
||||
"""Describe a checkpoint tensor without coupling policy to PyTorch."""
|
||||
|
||||
name: str
|
||||
shape: tuple[int, ...]
|
||||
dtype_name: str
|
||||
|
||||
|
||||
class ModelQuantizationRecipe(Protocol):
|
||||
"""Assign model-specific per-tensor formats for named profiles."""
|
||||
|
||||
@property
|
||||
def model_family(self) -> str:
|
||||
"""Return the stable family identifier used in cache identity."""
|
||||
|
||||
@property
|
||||
def version(self) -> int:
|
||||
"""Return the recipe version used in cache invalidation."""
|
||||
|
||||
@property
|
||||
def profiles(self) -> tuple[QuantizationProfile, ...]:
|
||||
"""Return deterministic workflow profiles owned by this recipe."""
|
||||
|
||||
def profile_from_selection(self, selection: str) -> QuantizationProfile:
|
||||
"""Parse a workflow selection into a recipe-owned profile."""
|
||||
|
||||
def policy_for(
|
||||
self,
|
||||
tensor: TensorDescriptor,
|
||||
profile: QuantizationProfile,
|
||||
) -> QuantizationFormat | None:
|
||||
"""Return the tensor format or ``None`` to preserve source precision."""
|
||||
@@ -0,0 +1,89 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Detect effective negative weights in Comfy-style prompt emphasis."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _WeightedPromptSegment:
|
||||
"""Retain one parsed prompt fragment and its effective scalar weight."""
|
||||
|
||||
text: str
|
||||
weight: float
|
||||
|
||||
|
||||
def contains_negative_prompt_weight(text: str) -> bool:
|
||||
"""Return whether valid nested emphasis gives any prompt text a negative weight."""
|
||||
|
||||
if not isinstance(text, str):
|
||||
raise TypeError("Negative prompt-weight detection requires text.")
|
||||
escaped = text.replace(r"\)", "\0\1").replace(r"\(", "\0\2")
|
||||
return any(
|
||||
segment.text and segment.weight < 0.0
|
||||
for segment in _weighted_segments(escaped, 1.0)
|
||||
)
|
||||
|
||||
|
||||
def _weighted_segments(
|
||||
text: str,
|
||||
current_weight: float,
|
||||
) -> tuple[_WeightedPromptSegment, ...]:
|
||||
"""Parse emphasis with the same nesting and final-colon rules as ComfyUI."""
|
||||
|
||||
parsed: list[_WeightedPromptSegment] = []
|
||||
for item in _parenthesized_items(text):
|
||||
weight = current_weight
|
||||
if len(item) >= 2 and item[0] == "(" and item[-1] == ")":
|
||||
inner = item[1:-1]
|
||||
delimiter = inner.rfind(":")
|
||||
weight *= 1.1
|
||||
if delimiter > 0:
|
||||
try:
|
||||
weight = float(inner[delimiter + 1 :])
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
inner = inner[:delimiter]
|
||||
parsed.extend(_weighted_segments(inner, weight))
|
||||
continue
|
||||
parsed.append(
|
||||
_WeightedPromptSegment(
|
||||
item.replace("\0\1", ")").replace("\0\2", "("),
|
||||
current_weight,
|
||||
)
|
||||
)
|
||||
return tuple(parsed)
|
||||
|
||||
|
||||
def _parenthesized_items(text: str) -> tuple[str, ...]:
|
||||
"""Split top-level parenthesized regions while preserving malformed input."""
|
||||
|
||||
result: list[str] = []
|
||||
current = ""
|
||||
nesting = 0
|
||||
for character in text:
|
||||
if character == "(":
|
||||
if nesting == 0:
|
||||
if current:
|
||||
result.append(current)
|
||||
current = "("
|
||||
else:
|
||||
current += character
|
||||
nesting += 1
|
||||
elif character == ")":
|
||||
nesting -= 1
|
||||
if nesting == 0:
|
||||
result.append(f"{current})")
|
||||
current = ""
|
||||
else:
|
||||
current += character
|
||||
else:
|
||||
current += character
|
||||
if current:
|
||||
result.append(current)
|
||||
return tuple(result)
|
||||
@@ -0,0 +1,190 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own immutable processed regional Attention Coupling runtime plans."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from uuid import UUID
|
||||
|
||||
import torch
|
||||
|
||||
from .conditioning_schedule import ConditioningScheduleRange
|
||||
from .regional_attention import (
|
||||
require_non_negative_regional_attention_index,
|
||||
validate_regional_attention_plan_authorities,
|
||||
)
|
||||
from .regional_lora_plan import RegionalLoraPlan
|
||||
from .regional_mask_bank import RegionalMaskBank
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProcessedRegionalAttentionEntry:
|
||||
"""Retain one ordered model-ready conditioning entry and Comfy strength."""
|
||||
|
||||
entry_index: int
|
||||
uuid: UUID
|
||||
schedule: ConditioningScheduleRange
|
||||
cross_attention: torch.Tensor
|
||||
strength: float
|
||||
cross_attention_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate entry order, model context, and finite scalar strength."""
|
||||
|
||||
require_non_negative_regional_attention_index(
|
||||
self.entry_index,
|
||||
name="entry_index",
|
||||
)
|
||||
if not isinstance(self.uuid, UUID):
|
||||
raise TypeError("Processed conditioning entry UUID must be uuid.UUID.")
|
||||
if not isinstance(self.schedule, ConditioningScheduleRange):
|
||||
raise TypeError(
|
||||
"Processed conditioning entry schedule has an invalid type."
|
||||
)
|
||||
if not isinstance(self.cross_attention, torch.Tensor):
|
||||
raise TypeError("Processed cross_attention must be a torch.Tensor.")
|
||||
if self.cross_attention.ndim != 3:
|
||||
raise ValueError("Processed cross_attention must use BxSxD layout.")
|
||||
if any(int(size) < 1 for size in self.cross_attention.shape):
|
||||
raise ValueError("Processed cross_attention dimensions must be positive.")
|
||||
if not self.cross_attention.is_floating_point():
|
||||
raise TypeError("Processed cross_attention must be floating point.")
|
||||
if not bool(torch.isfinite(self.cross_attention).all().item()):
|
||||
raise ValueError("Processed cross_attention must contain finite values.")
|
||||
if isinstance(self.strength, bool) or not isinstance(
|
||||
self.strength,
|
||||
int | float,
|
||||
):
|
||||
raise TypeError("Processed conditioning strength must be a real number.")
|
||||
if not math.isfinite(float(self.strength)):
|
||||
raise ValueError("Processed conditioning strength must be finite.")
|
||||
object.__setattr__(self, "strength", float(self.strength))
|
||||
multiplier = self.cross_attention_value_multiplier
|
||||
if multiplier is None:
|
||||
return
|
||||
if (
|
||||
not isinstance(multiplier, torch.Tensor)
|
||||
or multiplier.shape != (*self.cross_attention.shape[:2], 1)
|
||||
or not multiplier.is_floating_point()
|
||||
or multiplier.device != self.cross_attention.device
|
||||
or multiplier.dtype != self.cross_attention.dtype
|
||||
or not bool(torch.isfinite(multiplier).all().item())
|
||||
):
|
||||
raise ValueError(
|
||||
"Processed attention value multiplier must be a finite floating "
|
||||
"BxSx1 tensor aligned with cross_attention."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProcessedRegionalAttentionContext:
|
||||
"""Retain every ordered processed entry for one authored conditioning."""
|
||||
|
||||
conditioning_index: int
|
||||
region_index: int | None
|
||||
entries: tuple[ProcessedRegionalAttentionEntry, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate global/regional ownership and model-ready tensor structure."""
|
||||
|
||||
if self.region_index is None:
|
||||
if self.conditioning_index != 0:
|
||||
raise ValueError(
|
||||
"Processed base attention context must use conditioning index 0."
|
||||
)
|
||||
else:
|
||||
require_non_negative_regional_attention_index(
|
||||
self.region_index,
|
||||
name="region_index",
|
||||
)
|
||||
if self.conditioning_index != self.region_index + 1:
|
||||
raise ValueError(
|
||||
"Processed regional conditioning_index must equal region_index + 1."
|
||||
)
|
||||
if not isinstance(self.entries, tuple) or not self.entries:
|
||||
raise ValueError("Processed attention context requires ordered entries.")
|
||||
if any(
|
||||
not isinstance(entry, ProcessedRegionalAttentionEntry)
|
||||
for entry in self.entries
|
||||
):
|
||||
raise TypeError("Processed attention context contains an invalid entry.")
|
||||
if tuple(entry.entry_index for entry in self.entries) != tuple(
|
||||
range(len(self.entries))
|
||||
):
|
||||
raise ValueError("Processed attention entries must use canonical order.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProcessedRegionalAttentionBranch:
|
||||
"""Retain one processed base context and ordered regional context bank."""
|
||||
|
||||
base_context: ProcessedRegionalAttentionContext
|
||||
regional_contexts: tuple[ProcessedRegionalAttentionContext, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require one global base and canonical regional context order."""
|
||||
|
||||
if not isinstance(self.base_context, ProcessedRegionalAttentionContext):
|
||||
raise TypeError("Processed attention base context has an invalid type.")
|
||||
if self.base_context.region_index is not None:
|
||||
raise ValueError("Processed attention base context must be global.")
|
||||
if not isinstance(self.regional_contexts, tuple):
|
||||
raise TypeError("Processed regional attention contexts must be a tuple.")
|
||||
if any(
|
||||
not isinstance(context, ProcessedRegionalAttentionContext)
|
||||
for context in self.regional_contexts
|
||||
):
|
||||
raise TypeError(
|
||||
"Processed regional attention branch contains an invalid context."
|
||||
)
|
||||
indices = tuple(context.region_index for context in self.regional_contexts)
|
||||
if indices != tuple(range(len(self.regional_contexts))):
|
||||
raise ValueError(
|
||||
"Processed regional attention contexts must use canonical order."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProcessedRegionalAttentionPlan:
|
||||
"""Retain processed branches and shared canonical regional authorities."""
|
||||
|
||||
positive: ProcessedRegionalAttentionBranch
|
||||
negative: ProcessedRegionalAttentionBranch
|
||||
mask_bank: RegionalMaskBank
|
||||
lora_plan: RegionalLoraPlan
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate plan owners and regional LoRA bounds."""
|
||||
|
||||
if not isinstance(
|
||||
self.positive, ProcessedRegionalAttentionBranch
|
||||
) or not isinstance(self.negative, ProcessedRegionalAttentionBranch):
|
||||
raise TypeError(
|
||||
"Processed regional attention plan contains an invalid branch."
|
||||
)
|
||||
validate_regional_attention_plan_authorities(
|
||||
mask_bank=self.mask_bank,
|
||||
lora_plan=self.lora_plan,
|
||||
positive_region_count=len(self.positive.regional_contexts),
|
||||
negative_region_count=len(self.negative.regional_contexts),
|
||||
)
|
||||
|
||||
@property
|
||||
def is_time_invariant(self) -> bool:
|
||||
"""Report whether contexts and regional LoRA strengths remain fixed."""
|
||||
|
||||
branches = (self.positive, self.negative)
|
||||
contexts = tuple(
|
||||
context
|
||||
for branch in branches
|
||||
for context in (branch.base_context, *branch.regional_contexts)
|
||||
)
|
||||
return self.lora_plan.is_time_invariant and all(
|
||||
entry.schedule.is_time_invariant
|
||||
for context in contexts
|
||||
for entry in context.entries
|
||||
)
|
||||
@@ -0,0 +1,310 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define validated identities and records for quantized profile artifacts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
from .model_quantization import QuantizationProfile
|
||||
|
||||
MANIFEST_SCHEMA_VERSION = 2
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceCheckpointIdentity:
|
||||
"""Identify an authoritative source checkpoint and its current file state."""
|
||||
|
||||
display_name: str
|
||||
path: Path
|
||||
size_bytes: int
|
||||
modified_ns: int
|
||||
sha256: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantCacheIdentity:
|
||||
"""Identify one profile and recipe derivative of a source checkpoint."""
|
||||
|
||||
source: SourceCheckpointIdentity
|
||||
profile: QuantizationProfile
|
||||
model_family: str
|
||||
recipe_version: int
|
||||
|
||||
@property
|
||||
def stable_key(self) -> str:
|
||||
"""Return the collision-resistant cache key."""
|
||||
|
||||
identity = "\0".join(
|
||||
(
|
||||
self.source.sha256,
|
||||
self.profile.profile_id,
|
||||
str(self.profile.version),
|
||||
self.model_family,
|
||||
str(self.recipe_version),
|
||||
)
|
||||
)
|
||||
return hashlib.sha256(identity.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantCacheManifest:
|
||||
"""Describe one complete SimpleSyrup-managed cache artifact."""
|
||||
|
||||
source_model: str
|
||||
source_path: str
|
||||
source_sha256: str
|
||||
source_size_bytes: int
|
||||
source_modified_ns: int
|
||||
profile_id: str
|
||||
profile_label: str
|
||||
profile_version: int
|
||||
quantization_formats: tuple[str, ...]
|
||||
model_family: str
|
||||
recipe_version: int
|
||||
artifact_file: str
|
||||
artifact_size_bytes: int
|
||||
created_at: str
|
||||
last_used_at: str
|
||||
schema_version: int = MANIFEST_SCHEMA_VERSION
|
||||
managed_by: str = "SimpleSyrup"
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
identity: QuantCacheIdentity,
|
||||
artifact_file: str,
|
||||
artifact_size_bytes: int,
|
||||
) -> QuantCacheManifest:
|
||||
"""Create a current manifest for a completed artifact."""
|
||||
|
||||
now = datetime.now(UTC).isoformat()
|
||||
return cls(
|
||||
source_model=identity.source.display_name,
|
||||
source_path=str(identity.source.path),
|
||||
source_sha256=identity.source.sha256,
|
||||
source_size_bytes=identity.source.size_bytes,
|
||||
source_modified_ns=identity.source.modified_ns,
|
||||
profile_id=identity.profile.profile_id,
|
||||
profile_label=identity.profile.label,
|
||||
profile_version=identity.profile.version,
|
||||
quantization_formats=tuple(
|
||||
sorted(item.value for item in identity.profile.required_formats)
|
||||
),
|
||||
model_family=identity.model_family,
|
||||
recipe_version=identity.recipe_version,
|
||||
artifact_file=artifact_file,
|
||||
artifact_size_bytes=artifact_size_bytes,
|
||||
created_at=now,
|
||||
last_used_at=now,
|
||||
)
|
||||
|
||||
def matches_current_source(
|
||||
self,
|
||||
source_model: str,
|
||||
source_path: Path,
|
||||
source_size_bytes: int,
|
||||
source_modified_ns: int,
|
||||
profile: QuantizationProfile,
|
||||
model_family: str,
|
||||
recipe_version: int,
|
||||
) -> bool:
|
||||
"""Return whether this artifact derives from the unchanged source file."""
|
||||
|
||||
return (
|
||||
self.source_model == source_model
|
||||
and Path(self.source_path) == source_path
|
||||
and self.source_size_bytes == source_size_bytes
|
||||
and self.source_modified_ns == source_modified_ns
|
||||
and self.profile_id == profile.profile_id
|
||||
and self.profile_version == profile.version
|
||||
and self.model_family == model_family
|
||||
and self.recipe_version == recipe_version
|
||||
)
|
||||
|
||||
def matches_identity(self, identity: QuantCacheIdentity) -> bool:
|
||||
"""Return whether this v2 manifest exactly describes an identity."""
|
||||
|
||||
return (
|
||||
self.schema_version == MANIFEST_SCHEMA_VERSION
|
||||
and self.source_sha256 == identity.source.sha256
|
||||
and self.profile_id == identity.profile.profile_id
|
||||
and self.profile_version == identity.profile.version
|
||||
and self.model_family == identity.model_family
|
||||
and self.recipe_version == identity.recipe_version
|
||||
)
|
||||
|
||||
def touched(self) -> QuantCacheManifest:
|
||||
"""Return a copy with a current explicit LRU timestamp."""
|
||||
|
||||
payload = self.to_payload()
|
||||
payload["last_used_at"] = datetime.now(UTC).isoformat()
|
||||
return QuantCacheManifest.from_payload(payload)
|
||||
|
||||
def to_payload(self) -> dict[str, object]:
|
||||
"""Return the human-readable JSON representation."""
|
||||
|
||||
if self.schema_version == 1:
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"managed_by": self.managed_by,
|
||||
"source_model": self.source_model,
|
||||
"source_path": self.source_path,
|
||||
"source_sha256": self.source_sha256,
|
||||
"source_size_bytes": self.source_size_bytes,
|
||||
"source_modified_ns": self.source_modified_ns,
|
||||
"quantization_format": self.quantization_formats[0],
|
||||
"model_family": self.model_family,
|
||||
"recipe_version": self.recipe_version,
|
||||
"artifact_file": self.artifact_file,
|
||||
"artifact_size_bytes": self.artifact_size_bytes,
|
||||
"created_at": self.created_at,
|
||||
"last_used_at": self.last_used_at,
|
||||
}
|
||||
return {
|
||||
"schema_version": self.schema_version,
|
||||
"managed_by": self.managed_by,
|
||||
"source_model": self.source_model,
|
||||
"source_path": self.source_path,
|
||||
"source_sha256": self.source_sha256,
|
||||
"source_size_bytes": self.source_size_bytes,
|
||||
"source_modified_ns": self.source_modified_ns,
|
||||
"profile_id": self.profile_id,
|
||||
"profile_label": self.profile_label,
|
||||
"profile_version": self.profile_version,
|
||||
"quantization_formats": list(self.quantization_formats),
|
||||
"model_family": self.model_family,
|
||||
"recipe_version": self.recipe_version,
|
||||
"artifact_file": self.artifact_file,
|
||||
"artifact_size_bytes": self.artifact_size_bytes,
|
||||
"created_at": self.created_at,
|
||||
"last_used_at": self.last_used_at,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: object) -> QuantCacheManifest:
|
||||
"""Validate managed manifests, retaining v1 only for cache cleanup."""
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("Quant cache manifest must be a JSON object.")
|
||||
if payload.get("schema_version") == 1:
|
||||
return cls._from_legacy_payload(payload)
|
||||
required_strings = (
|
||||
"managed_by",
|
||||
"source_model",
|
||||
"source_path",
|
||||
"source_sha256",
|
||||
"profile_id",
|
||||
"profile_label",
|
||||
"model_family",
|
||||
"artifact_file",
|
||||
"created_at",
|
||||
"last_used_at",
|
||||
)
|
||||
for key in required_strings:
|
||||
if not isinstance(payload.get(key), str):
|
||||
raise ValueError(
|
||||
f"Quant cache manifest field '{key}' must be a string."
|
||||
)
|
||||
required_integers = (
|
||||
"schema_version",
|
||||
"source_size_bytes",
|
||||
"source_modified_ns",
|
||||
"profile_version",
|
||||
"recipe_version",
|
||||
"artifact_size_bytes",
|
||||
)
|
||||
for key in required_integers:
|
||||
value = payload.get(key)
|
||||
if not isinstance(value, int) or isinstance(value, bool):
|
||||
raise ValueError(
|
||||
f"Quant cache manifest field '{key}' must be an integer."
|
||||
)
|
||||
raw_formats = payload.get("quantization_formats")
|
||||
if not isinstance(raw_formats, list) or not all(
|
||||
isinstance(item, str) for item in raw_formats
|
||||
):
|
||||
raise ValueError(
|
||||
"Quant cache manifest field 'quantization_formats' must be a "
|
||||
"string list."
|
||||
)
|
||||
if payload["managed_by"] != "SimpleSyrup":
|
||||
raise ValueError("Quant cache manifest is not managed by SimpleSyrup.")
|
||||
if payload["schema_version"] != MANIFEST_SCHEMA_VERSION:
|
||||
raise ValueError("Quant cache manifest schema version is unsupported.")
|
||||
return cls(
|
||||
schema_version=payload["schema_version"],
|
||||
managed_by=payload["managed_by"],
|
||||
source_model=payload["source_model"],
|
||||
source_path=payload["source_path"],
|
||||
source_sha256=payload["source_sha256"],
|
||||
source_size_bytes=payload["source_size_bytes"],
|
||||
source_modified_ns=payload["source_modified_ns"],
|
||||
profile_id=payload["profile_id"],
|
||||
profile_label=payload["profile_label"],
|
||||
profile_version=payload["profile_version"],
|
||||
quantization_formats=tuple(raw_formats),
|
||||
model_family=payload["model_family"],
|
||||
recipe_version=payload["recipe_version"],
|
||||
artifact_file=payload["artifact_file"],
|
||||
artifact_size_bytes=payload["artifact_size_bytes"],
|
||||
created_at=payload["created_at"],
|
||||
last_used_at=payload["last_used_at"],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _from_legacy_payload(cls, payload: dict[object, object]) -> QuantCacheManifest:
|
||||
"""Decode v1 solely so ordinary LRU and clearing can remove it."""
|
||||
|
||||
required_strings = (
|
||||
"managed_by",
|
||||
"source_model",
|
||||
"source_path",
|
||||
"source_sha256",
|
||||
"quantization_format",
|
||||
"model_family",
|
||||
"artifact_file",
|
||||
"created_at",
|
||||
"last_used_at",
|
||||
)
|
||||
required_integers = (
|
||||
"source_size_bytes",
|
||||
"source_modified_ns",
|
||||
"recipe_version",
|
||||
"artifact_size_bytes",
|
||||
)
|
||||
if any(not isinstance(payload.get(key), str) for key in required_strings):
|
||||
raise ValueError("Legacy quant cache manifest has invalid string fields.")
|
||||
if any(
|
||||
not isinstance(payload.get(key), int) or isinstance(payload.get(key), bool)
|
||||
for key in required_integers
|
||||
):
|
||||
raise ValueError("Legacy quant cache manifest has invalid integer fields.")
|
||||
if payload["managed_by"] != "SimpleSyrup":
|
||||
raise ValueError("Quant cache manifest is not managed by SimpleSyrup.")
|
||||
quantization_format = str(payload["quantization_format"])
|
||||
return cls(
|
||||
schema_version=1,
|
||||
managed_by=str(payload["managed_by"]),
|
||||
source_model=str(payload["source_model"]),
|
||||
source_path=str(payload["source_path"]),
|
||||
source_sha256=str(payload["source_sha256"]),
|
||||
source_size_bytes=cast(int, payload["source_size_bytes"]),
|
||||
source_modified_ns=cast(int, payload["source_modified_ns"]),
|
||||
profile_id=f"legacy-v1-{quantization_format}",
|
||||
profile_label=f"Legacy v1 {quantization_format}",
|
||||
profile_version=1,
|
||||
quantization_formats=(quantization_format,),
|
||||
model_family=str(payload["model_family"]),
|
||||
recipe_version=cast(int, payload["recipe_version"]),
|
||||
artifact_file=str(payload["artifact_file"]),
|
||||
artifact_size_bytes=cast(int, payload["artifact_size_bytes"]),
|
||||
created_at=str(payload["created_at"]),
|
||||
last_used_at=str(payload["last_used_at"]),
|
||||
)
|
||||
@@ -0,0 +1,147 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own immutable raw regional Attention Coupling authoring plans."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .conditioning_batch import ConditioningBatch
|
||||
from .regional_attention import (
|
||||
require_non_negative_regional_attention_index,
|
||||
validate_regional_attention_plan_authorities,
|
||||
)
|
||||
from .regional_lora_plan import EMPTY_REGIONAL_LORA_PLAN, RegionalLoraPlan
|
||||
from .regional_mask_bank import RegionalMaskBank
|
||||
from .regional_prompting import build_regional_conditioning_plan
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RawRegionalAttentionContext:
|
||||
"""Retain one original regional conditioning and its canonical indices."""
|
||||
|
||||
conditioning_index: int
|
||||
region_index: int
|
||||
conditioning: object
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require the established global-first positional relationship."""
|
||||
|
||||
require_non_negative_regional_attention_index(
|
||||
self.region_index,
|
||||
name="region_index",
|
||||
)
|
||||
if self.conditioning_index != self.region_index + 1:
|
||||
raise ValueError(
|
||||
"Regional attention conditioning_index must equal region_index + 1."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RawRegionalAttentionBranch:
|
||||
"""Retain one base conditioning and ordered regional conditionings."""
|
||||
|
||||
base_conditioning: object
|
||||
regional_contexts: tuple[RawRegionalAttentionContext, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require immutable canonical regional order."""
|
||||
|
||||
if not isinstance(self.regional_contexts, tuple):
|
||||
raise TypeError("Raw regional attention contexts must be a tuple.")
|
||||
if any(
|
||||
not isinstance(context, RawRegionalAttentionContext)
|
||||
for context in self.regional_contexts
|
||||
):
|
||||
raise TypeError(
|
||||
"Raw regional attention branch contains an invalid context."
|
||||
)
|
||||
indices = tuple(context.region_index for context in self.regional_contexts)
|
||||
if indices != tuple(range(len(self.regional_contexts))):
|
||||
raise ValueError(
|
||||
"Raw regional attention contexts must use canonical order."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RawRegionalAttentionPlan:
|
||||
"""Retain both raw branches and shared canonical regional authorities."""
|
||||
|
||||
positive: RawRegionalAttentionBranch
|
||||
negative: RawRegionalAttentionBranch
|
||||
mask_bank: RegionalMaskBank
|
||||
lora_plan: RegionalLoraPlan
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate plan owners and regional LoRA bounds."""
|
||||
|
||||
if not isinstance(self.positive, RawRegionalAttentionBranch) or not isinstance(
|
||||
self.negative, RawRegionalAttentionBranch
|
||||
):
|
||||
raise TypeError("Raw regional attention plan contains an invalid branch.")
|
||||
validate_regional_attention_plan_authorities(
|
||||
mask_bank=self.mask_bank,
|
||||
lora_plan=self.lora_plan,
|
||||
positive_region_count=len(self.positive.regional_contexts),
|
||||
negative_region_count=len(self.negative.regional_contexts),
|
||||
)
|
||||
|
||||
|
||||
def build_raw_regional_attention_plan(
|
||||
*,
|
||||
positive: object,
|
||||
negative: object,
|
||||
mask_bank: RegionalMaskBank,
|
||||
lora_plan: RegionalLoraPlan = EMPTY_REGIONAL_LORA_PLAN,
|
||||
) -> RawRegionalAttentionPlan:
|
||||
"""Build both branches through the authoritative global-first pairing policy."""
|
||||
|
||||
if not isinstance(mask_bank, RegionalMaskBank):
|
||||
raise TypeError("Raw regional attention requires a RegionalMaskBank.")
|
||||
return RawRegionalAttentionPlan(
|
||||
positive=_build_raw_branch(
|
||||
positive,
|
||||
mask_bank=mask_bank,
|
||||
input_name="positive",
|
||||
),
|
||||
negative=_build_raw_branch(
|
||||
negative,
|
||||
mask_bank=mask_bank,
|
||||
input_name="negative",
|
||||
),
|
||||
mask_bank=mask_bank,
|
||||
lora_plan=lora_plan,
|
||||
)
|
||||
|
||||
|
||||
def _build_raw_branch(
|
||||
conditioning: object,
|
||||
*,
|
||||
mask_bank: RegionalMaskBank,
|
||||
input_name: str,
|
||||
) -> RawRegionalAttentionBranch:
|
||||
"""Pair one raw conditioning branch without restating index policy."""
|
||||
|
||||
entries = (
|
||||
conditioning.entries
|
||||
if isinstance(conditioning, ConditioningBatch)
|
||||
else (conditioning,)
|
||||
)
|
||||
pairing = build_regional_conditioning_plan(
|
||||
region_count=mask_bank.region_count,
|
||||
conditioning_count=len(entries),
|
||||
input_name=input_name,
|
||||
)
|
||||
return RawRegionalAttentionBranch(
|
||||
base_conditioning=entries[0],
|
||||
regional_contexts=tuple(
|
||||
RawRegionalAttentionContext(
|
||||
conditioning_index=pair.conditioning_index,
|
||||
region_index=pair.mask_index,
|
||||
conditioning=entries[pair.conditioning_index],
|
||||
)
|
||||
for pair in pairing.pairs
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,237 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define model-neutral regional activation geometry and batch alignment."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .spatial_views import SpatialBatchLayout
|
||||
|
||||
|
||||
class RegionalActivationLayout(StrEnum):
|
||||
"""Classify the explicit spatial organization of one adapter activation."""
|
||||
|
||||
DIRECT_CONVOLUTION_1D = "direct_convolution_1d"
|
||||
DIRECT_CONVOLUTION_2D = "direct_convolution_2d"
|
||||
DIRECT_CONVOLUTION_3D = "direct_convolution_3d"
|
||||
FLATTENED_SPATIAL_TOKENS = "flattened_spatial_tokens"
|
||||
CONSUMER_SPATIALIZED = "consumer_spatialized"
|
||||
BRANCH_TOKENS = "branch_tokens"
|
||||
|
||||
|
||||
class RegionalTemporalOwnership(StrEnum):
|
||||
"""Declare how a two-dimensional authored mask owns temporal activations."""
|
||||
|
||||
NONE = "none"
|
||||
REPEAT_SPATIAL_MASK = "repeat_spatial_mask"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalActivationBatchAlignment:
|
||||
"""Retain CFG, latent-batch, and optional view-major alignment evidence."""
|
||||
|
||||
latent_batch_size: int
|
||||
chunk_count: int
|
||||
spatial_layout: SpatialBatchLayout | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require positive counts and one layout over the complete base batch."""
|
||||
|
||||
_positive_integer(self.latent_batch_size, name="latent_batch_size")
|
||||
_positive_integer(self.chunk_count, name="chunk_count")
|
||||
if self.spatial_layout is not None:
|
||||
if not isinstance(self.spatial_layout, SpatialBatchLayout):
|
||||
raise TypeError(
|
||||
"Regional activation spatial_layout must be a SpatialBatchLayout."
|
||||
)
|
||||
if self.spatial_layout.input_batch_size != self.base_batch_size:
|
||||
raise ValueError(
|
||||
"Regional activation spatial layout input batch must match "
|
||||
"CFG chunks times latent batch size."
|
||||
)
|
||||
|
||||
@property
|
||||
def base_batch_size(self) -> int:
|
||||
"""Return the model batch before spatial-view expansion."""
|
||||
|
||||
return self.latent_batch_size * self.chunk_count
|
||||
|
||||
@property
|
||||
def invocation_batch_size(self) -> int:
|
||||
"""Return the active model batch after optional view expansion."""
|
||||
|
||||
if self.spatial_layout is None:
|
||||
return self.base_batch_size
|
||||
return self.spatial_layout.expanded_batch_size
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalActivationGeometry:
|
||||
"""Describe one exact rank activation and its authored-mask correspondence."""
|
||||
|
||||
layout: RegionalActivationLayout
|
||||
invocation_shape: tuple[int, ...]
|
||||
feature_axis: int
|
||||
spatial_height: int
|
||||
spatial_width: int
|
||||
batch_alignment: RegionalActivationBatchAlignment
|
||||
temporal_axis: int | None = None
|
||||
temporal_ownership: RegionalTemporalOwnership = RegionalTemporalOwnership.NONE
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject ambiguous axes, batches, tokens, and temporal ownership."""
|
||||
|
||||
if not isinstance(self.layout, RegionalActivationLayout):
|
||||
raise TypeError("Regional activation layout has an invalid type.")
|
||||
if not isinstance(self.invocation_shape, tuple) or not self.invocation_shape:
|
||||
raise ValueError("Regional activation shape must be a nonempty tuple.")
|
||||
for dimension in self.invocation_shape:
|
||||
_positive_integer(dimension, name="shape dimension")
|
||||
if not isinstance(self.batch_alignment, RegionalActivationBatchAlignment):
|
||||
raise TypeError("Regional activation requires batch alignment evidence.")
|
||||
_positive_integer(self.spatial_height, name="spatial_height")
|
||||
_positive_integer(self.spatial_width, name="spatial_width")
|
||||
if self.invocation_shape[0] != self.batch_alignment.invocation_batch_size:
|
||||
raise ValueError(
|
||||
"Regional activation leading batch must match its alignment."
|
||||
)
|
||||
rank = len(self.invocation_shape)
|
||||
_axis(self.feature_axis, rank=rank, name="feature_axis")
|
||||
if self.feature_axis == 0:
|
||||
raise ValueError(
|
||||
"Regional activation feature axis cannot be the batch axis."
|
||||
)
|
||||
if not isinstance(self.temporal_ownership, RegionalTemporalOwnership):
|
||||
raise TypeError(
|
||||
"Regional activation temporal ownership has an invalid type."
|
||||
)
|
||||
self._validate_layout()
|
||||
|
||||
def _validate_layout(self) -> None:
|
||||
"""Match the declared layout to its exact conventional tensor shape."""
|
||||
|
||||
batch = self.batch_alignment.invocation_batch_size
|
||||
features = self.invocation_shape[self.feature_axis]
|
||||
if self.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_1D:
|
||||
self._require_shape((batch, features, self.spatial_width), feature_axis=1)
|
||||
if self.spatial_height != 1:
|
||||
raise ValueError("Direct Conv1d regional geometry requires height one.")
|
||||
self._require_no_temporal_axis()
|
||||
return
|
||||
if self.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_2D:
|
||||
self._require_shape(
|
||||
(batch, features, self.spatial_height, self.spatial_width),
|
||||
feature_axis=1,
|
||||
)
|
||||
self._require_no_temporal_axis()
|
||||
return
|
||||
if self.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_3D:
|
||||
if self.feature_axis != 1 or len(self.invocation_shape) != 5:
|
||||
raise ValueError("Direct Conv3d regional geometry requires B/C/D/H/W.")
|
||||
if self.temporal_axis != 2:
|
||||
raise ValueError(
|
||||
"Direct Conv3d regional geometry requires temporal axis 2."
|
||||
)
|
||||
if self.invocation_shape[3:] != (
|
||||
self.spatial_height,
|
||||
self.spatial_width,
|
||||
):
|
||||
raise ValueError("Direct Conv3d regional H/W must match its tensor.")
|
||||
if (
|
||||
self.temporal_ownership
|
||||
is not RegionalTemporalOwnership.REPEAT_SPATIAL_MASK
|
||||
):
|
||||
raise ValueError(
|
||||
"Direct Conv3d regional geometry requires explicit repeated "
|
||||
"spatial-mask temporal ownership."
|
||||
)
|
||||
return
|
||||
if self.layout in (
|
||||
RegionalActivationLayout.FLATTENED_SPATIAL_TOKENS,
|
||||
RegionalActivationLayout.CONSUMER_SPATIALIZED,
|
||||
):
|
||||
self._require_shape(
|
||||
(
|
||||
batch,
|
||||
self.spatial_height * self.spatial_width,
|
||||
features,
|
||||
),
|
||||
feature_axis=2,
|
||||
)
|
||||
self._require_no_temporal_axis()
|
||||
return
|
||||
if self.layout is RegionalActivationLayout.BRANCH_TOKENS:
|
||||
self._require_shape(
|
||||
(batch, self.spatial_width, features),
|
||||
feature_axis=2,
|
||||
)
|
||||
if self.spatial_height != 1:
|
||||
raise ValueError("Branch-token regional geometry requires height one.")
|
||||
self._require_no_temporal_axis()
|
||||
return
|
||||
raise AssertionError(f"Unhandled regional activation layout: {self.layout}")
|
||||
|
||||
def _require_shape(
|
||||
self,
|
||||
expected: tuple[int, ...],
|
||||
*,
|
||||
feature_axis: int,
|
||||
) -> None:
|
||||
"""Require one conventional shape and feature-axis location."""
|
||||
|
||||
if self.feature_axis != feature_axis or self.invocation_shape != expected:
|
||||
raise ValueError(
|
||||
f"{self.layout.value} regional geometry expected shape {expected} "
|
||||
f"with feature axis {feature_axis}; observed "
|
||||
f"{self.invocation_shape} and axis {self.feature_axis}."
|
||||
)
|
||||
|
||||
def _require_no_temporal_axis(self) -> None:
|
||||
"""Reject temporal claims from non-temporal image activation layouts."""
|
||||
|
||||
if self.temporal_axis is not None:
|
||||
raise ValueError(
|
||||
"Non-temporal regional geometry cannot declare a temporal axis."
|
||||
)
|
||||
if self.temporal_ownership is not RegionalTemporalOwnership.NONE:
|
||||
raise ValueError(
|
||||
"Non-temporal regional geometry cannot claim temporal ownership."
|
||||
)
|
||||
|
||||
@property
|
||||
def temporal_size(self) -> int | None:
|
||||
"""Return the explicit temporal size when this activation owns one."""
|
||||
|
||||
if self.temporal_axis is None:
|
||||
return None
|
||||
return self.invocation_shape[self.temporal_axis]
|
||||
|
||||
def broadcast_mask_shape(self, region_count: int) -> tuple[int, ...]:
|
||||
"""Return the exact region-major multiplier shape for this activation."""
|
||||
|
||||
_positive_integer(region_count, name="region_count")
|
||||
shape = list(self.invocation_shape)
|
||||
shape[self.feature_axis] = 1
|
||||
return (region_count, *shape)
|
||||
|
||||
|
||||
def _positive_integer(value: object, *, name: str) -> None:
|
||||
"""Require one strictly positive non-boolean integer."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Regional activation {name} must be an integer.")
|
||||
if value < 1:
|
||||
raise ValueError(f"Regional activation {name} must be positive.")
|
||||
|
||||
|
||||
def _axis(value: object, *, rank: int, name: str) -> None:
|
||||
"""Require one non-negative axis inside the invocation rank."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Regional activation {name} must be an integer.")
|
||||
if not 0 <= value < rank:
|
||||
raise ValueError(f"Regional activation {name} is outside the tensor rank.")
|
||||
@@ -0,0 +1,69 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own shared regional Attention Coupling contracts and validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
from .regional_lora_plan import RegionalLoraPlan
|
||||
from .regional_mask_bank import RegionalMaskBank
|
||||
|
||||
|
||||
class RegionalAttentionBranch(StrEnum):
|
||||
"""Name the processed conditioning bank selected for one Comfy chunk."""
|
||||
|
||||
POSITIVE = "positive"
|
||||
NEGATIVE = "negative"
|
||||
|
||||
|
||||
def require_non_negative_regional_attention_index(
|
||||
value: object,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
"""Require one non-negative integer regional-attention index."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Regional attention {name} must be an integer.")
|
||||
if value < 0:
|
||||
raise ValueError(f"Regional attention {name} must be non-negative.")
|
||||
|
||||
|
||||
def validate_regional_attention_plan_authorities(
|
||||
*,
|
||||
mask_bank: object,
|
||||
lora_plan: object,
|
||||
positive_region_count: int,
|
||||
negative_region_count: int,
|
||||
) -> None:
|
||||
"""Validate shared mask, LoRA, and branch-count plan authorities."""
|
||||
|
||||
if not isinstance(mask_bank, RegionalMaskBank):
|
||||
raise TypeError("Regional attention plan requires a RegionalMaskBank.")
|
||||
if not isinstance(lora_plan, RegionalLoraPlan):
|
||||
raise TypeError("Regional attention plan requires a RegionalLoraPlan.")
|
||||
for branch_name, region_count in (
|
||||
("positive", positive_region_count),
|
||||
("negative", negative_region_count),
|
||||
):
|
||||
require_non_negative_regional_attention_index(
|
||||
region_count,
|
||||
name=f"{branch_name} region count",
|
||||
)
|
||||
if region_count > mask_bank.region_count:
|
||||
raise ValueError(
|
||||
f"Regional attention {branch_name} branch exceeds the mask bank."
|
||||
)
|
||||
out_of_bounds = tuple(
|
||||
adapter
|
||||
for adapter in lora_plan.adapters
|
||||
if adapter.region_index >= mask_bank.region_count
|
||||
)
|
||||
if out_of_bounds:
|
||||
indices = ", ".join(str(adapter.region_index) for adapter in out_of_bounds)
|
||||
raise ValueError(
|
||||
"Regional attention LoRA region indices exceed the mask bank: " + indices
|
||||
)
|
||||
@@ -0,0 +1,227 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own immutable chunk-major regional attention batch alignment values."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from .regional_attention import RegionalAttentionBranch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalAttentionChunkBatch:
|
||||
"""Describe one Comfy chunk's contiguous slice of the model batch."""
|
||||
|
||||
chunk_index: int
|
||||
branch: RegionalAttentionBranch
|
||||
batch_start: int
|
||||
batch_stop: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate one positive non-empty contiguous batch slice."""
|
||||
|
||||
for name, value in (
|
||||
("chunk_index", self.chunk_index),
|
||||
("batch_start", self.batch_start),
|
||||
("batch_stop", self.batch_stop),
|
||||
):
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Regional attention {name} must be an integer.")
|
||||
if self.chunk_index < 0 or self.batch_start < 0:
|
||||
raise ValueError("Regional attention chunk indices must be non-negative.")
|
||||
if self.batch_stop <= self.batch_start:
|
||||
raise ValueError("Regional attention chunk batch slice must be non-empty.")
|
||||
if not isinstance(self.branch, RegionalAttentionBranch):
|
||||
raise TypeError("Regional attention chunk branch has an invalid type.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BatchedRegionalAttentionEntry:
|
||||
"""Retain one aligned regional entry and its per-sample Comfy strengths."""
|
||||
|
||||
entry_index: int
|
||||
context: torch.Tensor
|
||||
strengths: tuple[float, ...]
|
||||
cross_attention_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate entry order, aligned context, and finite sample strengths."""
|
||||
|
||||
if isinstance(self.entry_index, bool) or not isinstance(self.entry_index, int):
|
||||
raise TypeError("Regional attention entry_index must be an integer.")
|
||||
if self.entry_index < 0:
|
||||
raise ValueError("Regional attention entry_index must be non-negative.")
|
||||
_validate_aligned_context(self.context, name="entry")
|
||||
if not isinstance(self.strengths, tuple):
|
||||
raise TypeError("Regional attention entry strengths must be a tuple.")
|
||||
if len(self.strengths) != int(self.context.shape[0]):
|
||||
raise ValueError(
|
||||
"Regional attention entry strength count must match its batch."
|
||||
)
|
||||
for strength in self.strengths:
|
||||
if isinstance(strength, bool) or not isinstance(strength, int | float):
|
||||
raise TypeError(
|
||||
"Regional attention entry strength must be a real number."
|
||||
)
|
||||
if not math.isfinite(float(strength)):
|
||||
raise ValueError("Regional attention entry strength must be finite.")
|
||||
_validate_value_multiplier(
|
||||
self.cross_attention_value_multiplier,
|
||||
self.context,
|
||||
name="entry",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BatchedRegionalAttentionRegion:
|
||||
"""Retain every aligned active conditioning entry for one region."""
|
||||
|
||||
region_index: int
|
||||
entries: tuple[BatchedRegionalAttentionEntry, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require a non-empty canonical entry bank for one region."""
|
||||
|
||||
if isinstance(self.region_index, bool) or not isinstance(
|
||||
self.region_index, int
|
||||
):
|
||||
raise TypeError("Regional attention region_index must be an integer.")
|
||||
if self.region_index < 0:
|
||||
raise ValueError("Regional attention region_index must be non-negative.")
|
||||
if not isinstance(self.entries, tuple) or not self.entries:
|
||||
raise ValueError("Regional attention region requires active entries.")
|
||||
if any(
|
||||
not isinstance(entry, BatchedRegionalAttentionEntry)
|
||||
for entry in self.entries
|
||||
):
|
||||
raise TypeError("Regional attention region contains an invalid entry.")
|
||||
if tuple(entry.entry_index for entry in self.entries) != tuple(
|
||||
range(len(self.entries))
|
||||
):
|
||||
raise ValueError("Regional attention region entries must be canonical.")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BatchedRegionalAttentionContexts:
|
||||
"""Retain chunk-major base and per-region active-entry model contexts."""
|
||||
|
||||
latent_batch_size: int
|
||||
chunks: tuple[RegionalAttentionChunkBatch, ...]
|
||||
base_context: torch.Tensor
|
||||
regions: tuple[BatchedRegionalAttentionRegion, ...]
|
||||
base_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate complete chunk and tensor alignment."""
|
||||
|
||||
if isinstance(self.latent_batch_size, bool) or not isinstance(
|
||||
self.latent_batch_size, int
|
||||
):
|
||||
raise TypeError("Regional attention latent_batch_size must be an integer.")
|
||||
if self.latent_batch_size < 1:
|
||||
raise ValueError("Regional attention latent_batch_size must be positive.")
|
||||
if not isinstance(self.chunks, tuple) or not self.chunks:
|
||||
raise ValueError("Regional attention batch requires at least one chunk.")
|
||||
if any(
|
||||
not isinstance(chunk, RegionalAttentionChunkBatch) for chunk in self.chunks
|
||||
):
|
||||
raise TypeError("Regional attention batch contains an invalid chunk.")
|
||||
expected_start = 0
|
||||
for chunk_index, chunk in enumerate(self.chunks):
|
||||
if chunk.chunk_index != chunk_index or chunk.batch_start != expected_start:
|
||||
raise ValueError(
|
||||
"Regional attention chunks must be contiguous and ordered."
|
||||
)
|
||||
if chunk.batch_stop - chunk.batch_start != self.latent_batch_size:
|
||||
raise ValueError(
|
||||
"Regional attention chunk size must match latent batch."
|
||||
)
|
||||
expected_start = chunk.batch_stop
|
||||
_validate_aligned_context(
|
||||
self.base_context,
|
||||
expected_batch=expected_start,
|
||||
name="base",
|
||||
)
|
||||
_validate_value_multiplier(
|
||||
self.base_value_multiplier,
|
||||
self.base_context,
|
||||
name="base",
|
||||
)
|
||||
if not isinstance(self.regions, tuple):
|
||||
raise TypeError("Regional attention regions must be a tuple.")
|
||||
if tuple(region.region_index for region in self.regions) != tuple(
|
||||
range(len(self.regions))
|
||||
):
|
||||
raise ValueError("Regional attention regions must use canonical order.")
|
||||
for region in self.regions:
|
||||
for entry in region.entries:
|
||||
_validate_aligned_context(
|
||||
entry.context,
|
||||
expected_batch=expected_start,
|
||||
name=f"region {region.region_index} entry {entry.entry_index}",
|
||||
)
|
||||
if entry.context.shape[1:] != self.base_context.shape[1:]:
|
||||
raise ValueError(
|
||||
"Regional attention context sequence shapes must match."
|
||||
)
|
||||
if (
|
||||
entry.context.device != self.base_context.device
|
||||
or entry.context.dtype != self.base_context.dtype
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional attention context device and dtype must match."
|
||||
)
|
||||
|
||||
|
||||
def _validate_aligned_context(
|
||||
context: object,
|
||||
*,
|
||||
name: str,
|
||||
expected_batch: int | None = None,
|
||||
) -> None:
|
||||
"""Validate one finite floating BxSxD context tensor."""
|
||||
|
||||
if not isinstance(context, torch.Tensor):
|
||||
raise TypeError(f"Regional attention {name} context must be a tensor.")
|
||||
if context.ndim != 3 or (
|
||||
expected_batch is not None and int(context.shape[0]) != expected_batch
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional attention {name} context has an invalid aligned batch."
|
||||
)
|
||||
if not context.is_floating_point() or not bool(
|
||||
torch.isfinite(context).all().item()
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional attention {name} context must contain finite floating values."
|
||||
)
|
||||
|
||||
|
||||
def _validate_value_multiplier(
|
||||
multiplier: object,
|
||||
context: torch.Tensor,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
"""Validate one optional value multiplier against its aligned context."""
|
||||
|
||||
if multiplier is None:
|
||||
return
|
||||
if (
|
||||
not isinstance(multiplier, torch.Tensor)
|
||||
or multiplier.shape != (*context.shape[:2], 1)
|
||||
or not multiplier.is_floating_point()
|
||||
or multiplier.device != context.device
|
||||
or multiplier.dtype != context.dtype
|
||||
or not bool(torch.isfinite(multiplier).all().item())
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional attention {name} value multiplier must be a finite "
|
||||
"floating BxSx1 tensor aligned with its context."
|
||||
)
|
||||
@@ -0,0 +1,15 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own the immutable spatial execution vocabulary for regional attention."""
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class RegionalAttentionExecutionMode(StrEnum):
|
||||
"""Identify how one prepared regional model traverses the latent canvas."""
|
||||
|
||||
FULL = "full-context"
|
||||
TILED = "tiled"
|
||||
CONTEXTUAL = "Contextual"
|
||||
@@ -0,0 +1,220 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Select exact active processed Attention Coupling entries for one sigma."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from uuid import UUID
|
||||
|
||||
from .conditioning_schedule_selection import (
|
||||
CONDITIONING_SCHEDULE_SELECTION_POLICY,
|
||||
ConditioningScheduleSelectionPolicy,
|
||||
normalize_conditioning_sigma,
|
||||
)
|
||||
from .processed_regional_attention import (
|
||||
ProcessedRegionalAttentionBranch,
|
||||
ProcessedRegionalAttentionEntry,
|
||||
ProcessedRegionalAttentionPlan,
|
||||
)
|
||||
from .regional_attention import (
|
||||
RegionalAttentionBranch,
|
||||
require_non_negative_regional_attention_index,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActiveProcessedRegionalAttentionChunk:
|
||||
"""Bind one Comfy chunk UUID to its active base and regional entries."""
|
||||
|
||||
chunk_index: int
|
||||
branch: RegionalAttentionBranch
|
||||
base_entry: ProcessedRegionalAttentionEntry
|
||||
regional_entries: tuple[
|
||||
tuple[ProcessedRegionalAttentionEntry, ...] | None,
|
||||
...,
|
||||
]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate canonical immutable active selection state."""
|
||||
|
||||
require_non_negative_regional_attention_index(
|
||||
self.chunk_index,
|
||||
name="chunk_index",
|
||||
)
|
||||
if not isinstance(self.branch, RegionalAttentionBranch):
|
||||
raise TypeError("Regional attention chunk branch has an invalid type.")
|
||||
if not isinstance(self.base_entry, ProcessedRegionalAttentionEntry):
|
||||
raise TypeError("Regional attention chunk base entry has an invalid type.")
|
||||
if not isinstance(self.regional_entries, tuple):
|
||||
raise TypeError("Regional attention chunk regions must be a tuple.")
|
||||
for entries in self.regional_entries:
|
||||
if entries is not None and (
|
||||
not isinstance(entries, tuple)
|
||||
or any(
|
||||
not isinstance(entry, ProcessedRegionalAttentionEntry)
|
||||
for entry in entries
|
||||
)
|
||||
):
|
||||
raise TypeError(
|
||||
"Regional attention chunk contains invalid active entries."
|
||||
)
|
||||
|
||||
|
||||
class RegionalAttentionSelectionService:
|
||||
"""Match Comfy chunk UUIDs to active base and regional entry banks."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
schedule_policy: ConditioningScheduleSelectionPolicy | None = None,
|
||||
) -> None:
|
||||
"""Retain the focused schedule policy collaborator."""
|
||||
|
||||
self._schedule_policy = schedule_policy or ConditioningScheduleSelectionPolicy()
|
||||
|
||||
def select_chunks(
|
||||
self,
|
||||
plan: ProcessedRegionalAttentionPlan,
|
||||
*,
|
||||
cond_or_uncond: object,
|
||||
conditioning_uuids: object,
|
||||
sigma: float,
|
||||
) -> tuple[ActiveProcessedRegionalAttentionChunk, ...]:
|
||||
"""Select exact active entries in Comfy's supplied UUID order."""
|
||||
|
||||
if not isinstance(plan, ProcessedRegionalAttentionPlan):
|
||||
raise TypeError("Regional attention selection requires a processed plan.")
|
||||
selectors = _selectors(cond_or_uncond)
|
||||
identities = _conditioning_uuids(conditioning_uuids)
|
||||
if len(selectors) != len(identities):
|
||||
raise ValueError(
|
||||
"Regional attention selectors and UUIDs must have equal lengths."
|
||||
)
|
||||
current_sigma = normalize_conditioning_sigma(sigma)
|
||||
return tuple(
|
||||
self._select_chunk(
|
||||
plan,
|
||||
chunk_index=chunk_index,
|
||||
selector=selector,
|
||||
conditioning_uuid=conditioning_uuid,
|
||||
sigma=current_sigma,
|
||||
)
|
||||
for chunk_index, (selector, conditioning_uuid) in enumerate(
|
||||
zip(selectors, identities, strict=True)
|
||||
)
|
||||
)
|
||||
|
||||
def _select_chunk(
|
||||
self,
|
||||
plan: ProcessedRegionalAttentionPlan,
|
||||
*,
|
||||
chunk_index: int,
|
||||
selector: object,
|
||||
conditioning_uuid: UUID,
|
||||
sigma: float,
|
||||
) -> ActiveProcessedRegionalAttentionChunk:
|
||||
"""Resolve one branch UUID and all regional schedules atomically."""
|
||||
|
||||
branch, contexts = _branch(plan, selector=selector, chunk_index=chunk_index)
|
||||
matching_base = tuple(
|
||||
entry
|
||||
for entry in contexts.base_context.entries
|
||||
if entry.uuid is conditioning_uuid
|
||||
)
|
||||
if len(matching_base) != 1:
|
||||
raise ValueError(
|
||||
f"Regional attention {branch.value} chunk {chunk_index} UUID "
|
||||
"does not identify exactly one processed base entry."
|
||||
)
|
||||
base_entry = matching_base[0]
|
||||
if not self._schedule_policy.is_active(base_entry.schedule, sigma=sigma):
|
||||
raise ValueError(
|
||||
f"Regional attention {branch.value} chunk {chunk_index} UUID is "
|
||||
f"inactive at sigma {sigma}."
|
||||
)
|
||||
return ActiveProcessedRegionalAttentionChunk(
|
||||
chunk_index=chunk_index,
|
||||
branch=branch,
|
||||
base_entry=base_entry,
|
||||
regional_entries=tuple(
|
||||
self._regional_entries(contexts, region_index, sigma=sigma)
|
||||
for region_index in range(plan.mask_bank.region_count)
|
||||
),
|
||||
)
|
||||
|
||||
def _regional_entries(
|
||||
self,
|
||||
branch: ProcessedRegionalAttentionBranch,
|
||||
region_index: int,
|
||||
*,
|
||||
sigma: float,
|
||||
) -> tuple[ProcessedRegionalAttentionEntry, ...] | None:
|
||||
"""Distinguish absent regions from authored regions with no active entry."""
|
||||
|
||||
if region_index >= len(branch.regional_contexts):
|
||||
return None
|
||||
return self._active_entries(
|
||||
branch.regional_contexts[region_index].entries,
|
||||
sigma=sigma,
|
||||
)
|
||||
|
||||
def _active_entries(
|
||||
self,
|
||||
entries: tuple[ProcessedRegionalAttentionEntry, ...],
|
||||
*,
|
||||
sigma: float,
|
||||
) -> tuple[ProcessedRegionalAttentionEntry, ...]:
|
||||
"""Retain every active entry in authored order at one finite sigma."""
|
||||
|
||||
return tuple(
|
||||
entry
|
||||
for entry in entries
|
||||
if self._schedule_policy.is_active(entry.schedule, sigma=sigma)
|
||||
)
|
||||
|
||||
|
||||
def _branch(
|
||||
plan: ProcessedRegionalAttentionPlan,
|
||||
*,
|
||||
selector: object,
|
||||
chunk_index: int,
|
||||
) -> tuple[RegionalAttentionBranch, ProcessedRegionalAttentionBranch]:
|
||||
"""Narrow one Comfy branch selector without accepting booleans."""
|
||||
|
||||
if isinstance(selector, bool) or not isinstance(selector, int):
|
||||
raise TypeError(
|
||||
f"cond_or_uncond selector {chunk_index} must be integer 0 or 1."
|
||||
)
|
||||
if selector == 0:
|
||||
return RegionalAttentionBranch.POSITIVE, plan.positive
|
||||
if selector == 1:
|
||||
return RegionalAttentionBranch.NEGATIVE, plan.negative
|
||||
raise ValueError(
|
||||
f"cond_or_uncond selector {chunk_index} must be 0 or 1; observed {selector}."
|
||||
)
|
||||
|
||||
|
||||
def _selectors(value: object) -> tuple[object, ...]:
|
||||
"""Require Comfy's ordered selector container."""
|
||||
|
||||
if not isinstance(value, list | tuple):
|
||||
raise TypeError("cond_or_uncond must be a list or tuple of chunk selectors.")
|
||||
return tuple(value)
|
||||
|
||||
|
||||
def _conditioning_uuids(value: object) -> tuple[UUID, ...]:
|
||||
"""Require exact Comfy UUID objects for every supplied chunk."""
|
||||
|
||||
if not isinstance(value, list | tuple):
|
||||
raise TypeError("Conditioning UUIDs must be a list or tuple.")
|
||||
identities = tuple(value)
|
||||
if any(not isinstance(identity, UUID) for identity in identities):
|
||||
raise TypeError("Every conditioning identity must be a Comfy UUID.")
|
||||
return identities
|
||||
|
||||
|
||||
REGIONAL_ATTENTION_SELECTION_SERVICE = RegionalAttentionSelectionService(
|
||||
CONDITIONING_SCHEDULE_SELECTION_POLICY
|
||||
)
|
||||
@@ -0,0 +1,252 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Normalize regional attention weights and blend ordered branch outputs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalAttentionWeights:
|
||||
"""Hold raw base, region, and denominator weights on one query grid."""
|
||||
|
||||
base: torch.Tensor
|
||||
regions: torch.Tensor
|
||||
denominator: torch.Tensor
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require consistent finite non-negative weighting tensors."""
|
||||
|
||||
for name, tensor in (
|
||||
("Base", self.base),
|
||||
("Region", self.regions),
|
||||
("Denominator", self.denominator),
|
||||
):
|
||||
if not isinstance(tensor, torch.Tensor):
|
||||
raise TypeError(f"{name} attention weights must be a torch.Tensor.")
|
||||
if not tensor.is_floating_point():
|
||||
raise TypeError(
|
||||
f"{name} attention weights must use a floating-point dtype."
|
||||
)
|
||||
if not bool(torch.isfinite(tensor).all()):
|
||||
raise ValueError(f"{name} attention weights must be finite.")
|
||||
if not bool((tensor >= 0.0).all()):
|
||||
raise ValueError(f"{name} attention weights must be non-negative.")
|
||||
if self.base.ndim < 1:
|
||||
raise ValueError("Base attention weights require a query-grid dimension.")
|
||||
if self.regions.ndim != self.base.ndim + 1:
|
||||
raise ValueError(
|
||||
"Region attention weights require a leading region dimension."
|
||||
)
|
||||
if int(self.regions.shape[0]) < 1:
|
||||
raise ValueError("Region attention weights require at least one region.")
|
||||
if tuple(self.regions.shape[1:]) != tuple(self.base.shape):
|
||||
raise ValueError("Region attention weights must match the base query grid.")
|
||||
if self.denominator.shape != self.base.shape:
|
||||
raise ValueError("Attention denominator must match the base query grid.")
|
||||
if not (
|
||||
self.base.dtype == self.regions.dtype == self.denominator.dtype
|
||||
and self.base.device == self.regions.device == self.denominator.device
|
||||
):
|
||||
raise ValueError(
|
||||
"Base, region, and denominator weights must share dtype and device."
|
||||
)
|
||||
if not bool((self.denominator > 0.0).all()):
|
||||
raise ValueError("Attention denominator must be strictly positive.")
|
||||
|
||||
@property
|
||||
def normalized_base(self) -> torch.Tensor:
|
||||
"""Return normalized base-complement weights."""
|
||||
|
||||
return self.base / self.denominator
|
||||
|
||||
@property
|
||||
def normalized_regions(self) -> torch.Tensor:
|
||||
"""Return normalized ordered regional weights."""
|
||||
|
||||
return self.regions / self.denominator.unsqueeze(0)
|
||||
|
||||
|
||||
class RegionalAttentionWeightingPolicy:
|
||||
"""Own Comfy-compatible base complement and normalized branch composition."""
|
||||
|
||||
def weights(
|
||||
self,
|
||||
masks: torch.Tensor,
|
||||
*,
|
||||
region_strengths: tuple[float, ...],
|
||||
epsilon: float = 1e-6,
|
||||
) -> RegionalAttentionWeights:
|
||||
"""Return raw weighting terms for ordered region-first query masks."""
|
||||
|
||||
self._validate_mask_structure(masks)
|
||||
strengths = self._validate_strengths(
|
||||
region_strengths,
|
||||
region_count=int(masks.shape[0]),
|
||||
)
|
||||
if (
|
||||
isinstance(epsilon, bool)
|
||||
or not isinstance(epsilon, int | float)
|
||||
or not math.isfinite(float(epsilon))
|
||||
or epsilon <= 0.0
|
||||
):
|
||||
raise ValueError("Regional attention epsilon must be finite and positive.")
|
||||
self._validate_mask_values(masks)
|
||||
|
||||
strength_shape = (len(strengths),) + (1,) * (masks.ndim - 1)
|
||||
strength_tensor = masks.new_tensor(strengths).reshape(strength_shape)
|
||||
region_weights = masks.clamp(0.0, 1.0) * strength_tensor
|
||||
region_sum = region_weights.sum(dim=0)
|
||||
base_weight = torch.relu(1.0 - region_sum)
|
||||
denominator = (base_weight + region_sum).clamp_min(float(epsilon))
|
||||
return RegionalAttentionWeights(
|
||||
base=base_weight,
|
||||
regions=region_weights,
|
||||
denominator=denominator,
|
||||
)
|
||||
|
||||
def blend(
|
||||
self,
|
||||
*,
|
||||
weights: RegionalAttentionWeights,
|
||||
base_output: torch.Tensor,
|
||||
regional_outputs: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Blend one base and ordered regional outputs over their query grid."""
|
||||
|
||||
self._validate_outputs(
|
||||
weights=weights,
|
||||
base_output=base_output,
|
||||
regional_outputs=regional_outputs,
|
||||
)
|
||||
feature_dimensions = base_output.ndim - weights.base.ndim
|
||||
feature_shape = (1,) * feature_dimensions
|
||||
base_weight = weights.base.reshape((*weights.base.shape, *feature_shape)).to(
|
||||
dtype=base_output.dtype
|
||||
)
|
||||
region_weights = weights.regions.reshape(
|
||||
(*weights.regions.shape, *feature_shape)
|
||||
).to(dtype=base_output.dtype)
|
||||
denominator = weights.denominator.reshape(
|
||||
(*weights.denominator.shape, *feature_shape)
|
||||
).to(dtype=base_output.dtype)
|
||||
numerator = base_weight * base_output + (region_weights * regional_outputs).sum(
|
||||
dim=0
|
||||
)
|
||||
blended = numerator / denominator
|
||||
if not bool(torch.isfinite(blended).all()):
|
||||
raise ValueError(
|
||||
"Blended regional attention output contains non-finite values."
|
||||
)
|
||||
return blended
|
||||
|
||||
@staticmethod
|
||||
def _validate_mask_structure(masks: torch.Tensor) -> None:
|
||||
"""Validate mask type and shape without inspecting device values."""
|
||||
|
||||
if not isinstance(masks, torch.Tensor):
|
||||
raise TypeError("Regional attention masks must be a torch.Tensor.")
|
||||
if masks.ndim < 2:
|
||||
raise ValueError(
|
||||
"Regional attention masks require region and query-grid dimensions."
|
||||
)
|
||||
if int(masks.shape[0]) < 1 or any(int(size) < 1 for size in masks.shape[1:]):
|
||||
raise ValueError(
|
||||
"Regional attention masks require non-empty region and query grids."
|
||||
)
|
||||
if not masks.is_floating_point():
|
||||
raise TypeError("Regional attention masks must use a floating-point dtype.")
|
||||
|
||||
@staticmethod
|
||||
def _validate_mask_values(masks: torch.Tensor) -> None:
|
||||
"""Reject non-finite mask values before constructing strength tensors."""
|
||||
|
||||
if not bool(torch.isfinite(masks).all()):
|
||||
raise ValueError("Regional attention masks must contain finite values.")
|
||||
|
||||
@staticmethod
|
||||
def _validate_strengths(
|
||||
strengths: tuple[float, ...],
|
||||
*,
|
||||
region_count: int,
|
||||
) -> tuple[float, ...]:
|
||||
"""Validate ordered immutable regional strengths before tensor creation."""
|
||||
|
||||
if not isinstance(strengths, tuple):
|
||||
raise TypeError("Regional attention strengths must be an immutable tuple.")
|
||||
if len(strengths) != region_count:
|
||||
raise ValueError(
|
||||
"Regional attention strength count must match the mask region count."
|
||||
)
|
||||
normalized: list[float] = []
|
||||
for index, strength in enumerate(strengths):
|
||||
if isinstance(strength, bool) or not isinstance(strength, int | float):
|
||||
raise TypeError(
|
||||
f"Regional attention strength {index} must be a real number."
|
||||
)
|
||||
value = float(strength)
|
||||
if not math.isfinite(value) or value < 0.0:
|
||||
raise ValueError(
|
||||
f"Regional attention strength {index} must be finite and "
|
||||
"non-negative."
|
||||
)
|
||||
normalized.append(value)
|
||||
return tuple(normalized)
|
||||
|
||||
@staticmethod
|
||||
def _validate_outputs(
|
||||
*,
|
||||
weights: RegionalAttentionWeights,
|
||||
base_output: torch.Tensor,
|
||||
regional_outputs: torch.Tensor,
|
||||
) -> None:
|
||||
"""Validate branch output ordering, grids, features, and tensor state."""
|
||||
|
||||
if not isinstance(weights, RegionalAttentionWeights):
|
||||
raise TypeError("Regional blend weights must be RegionalAttentionWeights.")
|
||||
if not isinstance(base_output, torch.Tensor) or not isinstance(
|
||||
regional_outputs, torch.Tensor
|
||||
):
|
||||
raise TypeError(
|
||||
"Regional attention branch outputs must be torch.Tensor values."
|
||||
)
|
||||
if (
|
||||
not base_output.is_floating_point()
|
||||
or not regional_outputs.is_floating_point()
|
||||
):
|
||||
raise TypeError(
|
||||
"Regional attention branch outputs must use floating-point dtypes."
|
||||
)
|
||||
if base_output.dtype != regional_outputs.dtype:
|
||||
raise ValueError("Regional attention branch output dtypes must match.")
|
||||
if base_output.device != regional_outputs.device:
|
||||
raise ValueError("Regional attention branch output devices must match.")
|
||||
expected_regional_shape = (int(weights.regions.shape[0]), *base_output.shape)
|
||||
if regional_outputs.shape != expected_regional_shape:
|
||||
raise ValueError(
|
||||
"Regional outputs must contain one ordered branch per region with "
|
||||
"the complete base output shape."
|
||||
)
|
||||
query_dimensions = weights.base.ndim
|
||||
if base_output.ndim < query_dimensions or tuple(
|
||||
base_output.shape[:query_dimensions]
|
||||
) != tuple(weights.base.shape):
|
||||
raise ValueError(
|
||||
"Regional attention outputs must begin with the weighting query grid."
|
||||
)
|
||||
if weights.base.device != base_output.device:
|
||||
raise ValueError(
|
||||
"Regional weights and branch outputs must share one device."
|
||||
)
|
||||
if not bool(torch.isfinite(base_output).all()) or not bool(
|
||||
torch.isfinite(regional_outputs).all()
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional attention branch outputs must contain finite values."
|
||||
)
|
||||
@@ -0,0 +1,105 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own Comfy-equivalent within-region conditioning-output combination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class RegionalConditioningOutputCombiner:
|
||||
"""Combine ordered active-entry outputs with native Comfy strength semantics."""
|
||||
|
||||
def combine(
|
||||
self,
|
||||
outputs: tuple[torch.Tensor, ...],
|
||||
*,
|
||||
strengths: tuple[tuple[float, ...], ...],
|
||||
) -> torch.Tensor:
|
||||
"""Return the ordered strength-weighted output normalized like Comfy."""
|
||||
|
||||
if not isinstance(outputs, tuple) or not outputs:
|
||||
raise ValueError("Regional conditioning combination requires outputs.")
|
||||
if not isinstance(strengths, tuple) or len(strengths) != len(outputs):
|
||||
raise ValueError(
|
||||
"Regional conditioning strengths must align with ordered outputs."
|
||||
)
|
||||
authority = outputs[0]
|
||||
if not isinstance(authority, torch.Tensor) or authority.ndim < 1:
|
||||
raise TypeError("Regional conditioning output must be a tensor batch.")
|
||||
if len(outputs) == 1:
|
||||
entry_strengths = strengths[0]
|
||||
self._validate_entry(
|
||||
authority,
|
||||
entry_strengths,
|
||||
authority=authority,
|
||||
entry_index=0,
|
||||
)
|
||||
if all(strength == 1.0 for strength in entry_strengths):
|
||||
return authority
|
||||
weighted = torch.zeros_like(authority)
|
||||
counts = torch.ones_like(authority) * 1e-37
|
||||
weight_shape = (int(authority.shape[0]),) + (1,) * (authority.ndim - 1)
|
||||
active_rows = torch.zeros(
|
||||
weight_shape,
|
||||
dtype=torch.bool,
|
||||
device=authority.device,
|
||||
)
|
||||
for entry_index, (output, entry_strengths) in enumerate(
|
||||
zip(outputs, strengths, strict=True)
|
||||
):
|
||||
self._validate_entry(
|
||||
output,
|
||||
entry_strengths,
|
||||
authority=authority,
|
||||
entry_index=entry_index,
|
||||
)
|
||||
weights = authority.new_tensor(entry_strengths).reshape(weight_shape)
|
||||
weighted += output * weights
|
||||
counts += weights
|
||||
active_rows |= weights.ne(0)
|
||||
denominator = torch.where(active_rows, counts, torch.ones_like(counts))
|
||||
return weighted / denominator
|
||||
|
||||
@staticmethod
|
||||
def _validate_entry(
|
||||
output: object,
|
||||
strengths: object,
|
||||
*,
|
||||
authority: torch.Tensor,
|
||||
entry_index: int,
|
||||
) -> None:
|
||||
"""Require exact output structure and one finite strength per sample."""
|
||||
|
||||
if not isinstance(output, torch.Tensor):
|
||||
raise TypeError(
|
||||
f"Regional conditioning output {entry_index} must be a tensor."
|
||||
)
|
||||
if (
|
||||
output.shape != authority.shape
|
||||
or output.device != authority.device
|
||||
or output.dtype != authority.dtype
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional conditioning output {entry_index} must match the first "
|
||||
"output shape, device, and dtype."
|
||||
)
|
||||
if not isinstance(strengths, tuple) or len(strengths) != int(
|
||||
authority.shape[0]
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional conditioning output {entry_index} strengths must match "
|
||||
"the output batch."
|
||||
)
|
||||
for strength in strengths:
|
||||
if isinstance(strength, bool) or not isinstance(strength, int | float):
|
||||
raise TypeError("Regional conditioning strengths must be real numbers.")
|
||||
if not math.isfinite(float(strength)):
|
||||
raise ValueError("Regional conditioning strengths must be finite.")
|
||||
|
||||
|
||||
REGIONAL_CONDITIONING_OUTPUT_COMBINER = RegionalConditioningOutputCombiner()
|
||||
@@ -0,0 +1,145 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define immutable regional feature requests, capabilities, and admissions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .regional_model_capabilities import RegionalModelCapabilities
|
||||
|
||||
|
||||
class RegionalFeature(StrEnum):
|
||||
"""Identify one independently admitted regional sampling feature."""
|
||||
|
||||
FULL_CONTEXT_MASKED_CONDITIONING = "full_context_masked_conditioning"
|
||||
ATTENTION_COUPLING = "attention_coupling"
|
||||
SPATIAL_MODEL_PATCH = "spatial_model_patch"
|
||||
CONTROL = "control"
|
||||
GLIGEN = "gligen"
|
||||
REFERENCE_LATENTS = "reference_latents"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalFeatureRequest:
|
||||
"""Describe the complete immutable regional feature intent for one sample."""
|
||||
|
||||
features: frozenset[RegionalFeature] = frozenset()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require an immutable set containing only typed regional features."""
|
||||
|
||||
_validate_features(self.features, value_name="Regional feature request")
|
||||
|
||||
def with_feature(self, feature: RegionalFeature) -> RegionalFeatureRequest:
|
||||
"""Return a new request containing one additional typed feature."""
|
||||
|
||||
if not isinstance(feature, RegionalFeature):
|
||||
raise TypeError("Requested regional feature must be a RegionalFeature.")
|
||||
return RegionalFeatureRequest(self.features | {feature})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalSamplerCapabilities:
|
||||
"""Describe the complete regional feature set implemented by one sampler."""
|
||||
|
||||
features: frozenset[RegionalFeature]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require an immutable set containing only typed regional features."""
|
||||
|
||||
_validate_features(self.features, value_name="Regional sampler capabilities")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalCapabilityAdmission:
|
||||
"""Record one complete successful request admission for downstream use."""
|
||||
|
||||
request: RegionalFeatureRequest
|
||||
admitted_features: frozenset[RegionalFeature]
|
||||
model_capabilities: RegionalModelCapabilities | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject partial, untyped, or model-incomplete admission state."""
|
||||
|
||||
if not isinstance(self.request, RegionalFeatureRequest):
|
||||
raise TypeError(
|
||||
"Regional admission request must be RegionalFeatureRequest."
|
||||
)
|
||||
_validate_features(
|
||||
self.admitted_features,
|
||||
value_name="Admitted regional features",
|
||||
)
|
||||
if self.admitted_features != self.request.features:
|
||||
raise ValueError("Regional capability admission cannot be partial.")
|
||||
if self.model_capabilities is not None and not isinstance(
|
||||
self.model_capabilities,
|
||||
RegionalModelCapabilities,
|
||||
):
|
||||
raise TypeError(
|
||||
"Regional admission model capabilities must be "
|
||||
"RegionalModelCapabilities."
|
||||
)
|
||||
if self.admitted_features & MODEL_DEPENDENT_REGIONAL_FEATURES:
|
||||
if self.model_capabilities is None:
|
||||
raise ValueError(
|
||||
"Model-dependent regional features require model capabilities."
|
||||
)
|
||||
|
||||
def supports(self, feature: RegionalFeature) -> bool:
|
||||
"""Return whether one typed feature was admitted for this sample."""
|
||||
|
||||
if not isinstance(feature, RegionalFeature):
|
||||
raise TypeError("Regional feature query must be a RegionalFeature.")
|
||||
return feature in self.admitted_features
|
||||
|
||||
|
||||
MODEL_DEPENDENT_REGIONAL_FEATURES = frozenset(
|
||||
{
|
||||
RegionalFeature.ATTENTION_COUPLING,
|
||||
RegionalFeature.SPATIAL_MODEL_PATCH,
|
||||
RegionalFeature.CONTROL,
|
||||
RegionalFeature.GLIGEN,
|
||||
RegionalFeature.REFERENCE_LATENTS,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _validate_features(
|
||||
features: frozenset[RegionalFeature],
|
||||
*,
|
||||
value_name: str,
|
||||
) -> None:
|
||||
"""Validate one immutable typed regional feature set."""
|
||||
|
||||
if not isinstance(features, frozenset):
|
||||
raise TypeError(f"{value_name} must use an immutable frozenset.")
|
||||
if not all(isinstance(feature, RegionalFeature) for feature in features):
|
||||
raise TypeError(f"{value_name} must contain RegionalFeature values.")
|
||||
|
||||
|
||||
EMPTY_REGIONAL_FEATURE_REQUEST = RegionalFeatureRequest()
|
||||
EMPTY_REGIONAL_CAPABILITY_ADMISSION = RegionalCapabilityAdmission(
|
||||
request=EMPTY_REGIONAL_FEATURE_REQUEST,
|
||||
admitted_features=frozenset(),
|
||||
model_capabilities=None,
|
||||
)
|
||||
CONTEXTUAL_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES = RegionalSamplerCapabilities(
|
||||
frozenset(
|
||||
{
|
||||
RegionalFeature.FULL_CONTEXT_MASKED_CONDITIONING,
|
||||
RegionalFeature.ATTENTION_COUPLING,
|
||||
}
|
||||
)
|
||||
)
|
||||
TILED_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES = RegionalSamplerCapabilities(
|
||||
frozenset(
|
||||
{
|
||||
RegionalFeature.FULL_CONTEXT_MASKED_CONDITIONING,
|
||||
RegionalFeature.ATTENTION_COUPLING,
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,161 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own immutable regional model-side LoRA composition plans."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class RegionalLoraBranch(StrEnum):
|
||||
"""Identify the conditioning branch that owns one regional adapter use."""
|
||||
|
||||
POSITIVE = "positive"
|
||||
NEGATIVE = "negative"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RegionalLoraAdapterIdentity:
|
||||
"""Retain the caller-supplied stable identity of one LoRA artifact."""
|
||||
|
||||
value: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject identities that cannot distinguish an adapter."""
|
||||
|
||||
if not isinstance(self.value, str) or not self.value.strip():
|
||||
raise ValueError(
|
||||
"Regional LoRA adapter identity must be a non-empty string."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RegionalLoraScheduleBoundary:
|
||||
"""Retain one ordered Comfy HookKeyframe boundary exactly."""
|
||||
|
||||
start_percent: float
|
||||
start_sigma: float
|
||||
strength_multiplier: float
|
||||
guarantee_steps: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate one finite normalized schedule boundary."""
|
||||
|
||||
_require_finite_float(self.start_percent, name="start_percent")
|
||||
if not 0.0 <= self.start_percent <= 1.0:
|
||||
raise ValueError("Regional LoRA schedule start_percent must be in [0, 1].")
|
||||
_require_finite_float(self.start_sigma, name="start_sigma")
|
||||
if self.start_sigma < 0.0:
|
||||
raise ValueError("Regional LoRA schedule start_sigma must be non-negative.")
|
||||
_require_finite_float(
|
||||
self.strength_multiplier,
|
||||
name="strength_multiplier",
|
||||
)
|
||||
if isinstance(self.guarantee_steps, bool) or not isinstance(
|
||||
self.guarantee_steps, int
|
||||
):
|
||||
raise TypeError("Regional LoRA guarantee_steps must be an integer.")
|
||||
if self.guarantee_steps < 0:
|
||||
raise ValueError("Regional LoRA guarantee_steps must be non-negative.")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RegionalLoraAdapterPlan:
|
||||
"""Describe one ordered regional use of a model-side LoRA adapter."""
|
||||
|
||||
adapter_identity: RegionalLoraAdapterIdentity
|
||||
composition_index: int
|
||||
region_index: int
|
||||
branch: RegionalLoraBranch
|
||||
model_strength: float
|
||||
schedule: tuple[RegionalLoraScheduleBoundary, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate adapter ownership, strength, and ordered schedule."""
|
||||
|
||||
if not isinstance(self.adapter_identity, RegionalLoraAdapterIdentity):
|
||||
raise TypeError("Regional LoRA adapter_identity has an invalid type.")
|
||||
_require_non_negative_index(self.composition_index, name="composition_index")
|
||||
_require_non_negative_index(self.region_index, name="region_index")
|
||||
if not isinstance(self.branch, RegionalLoraBranch):
|
||||
raise TypeError("Regional LoRA branch has an invalid type.")
|
||||
_require_finite_float(self.model_strength, name="model_strength")
|
||||
if not isinstance(self.schedule, tuple) or not self.schedule:
|
||||
raise ValueError(
|
||||
"Regional LoRA schedule must contain at least one boundary."
|
||||
)
|
||||
if any(
|
||||
not isinstance(boundary, RegionalLoraScheduleBoundary)
|
||||
for boundary in self.schedule
|
||||
):
|
||||
raise TypeError("Regional LoRA schedule contains an invalid boundary.")
|
||||
starts = tuple(boundary.start_percent for boundary in self.schedule)
|
||||
if starts != tuple(sorted(starts)):
|
||||
raise ValueError("Regional LoRA schedule boundaries must be ordered.")
|
||||
sigmas = tuple(boundary.start_sigma for boundary in self.schedule)
|
||||
if sigmas != tuple(sorted(sigmas, reverse=True)):
|
||||
raise ValueError(
|
||||
"Regional LoRA converted schedule boundaries must be descending."
|
||||
)
|
||||
|
||||
@property
|
||||
def is_time_invariant(self) -> bool:
|
||||
"""Report whether every keyframe retains one effective multiplier."""
|
||||
|
||||
first_multiplier = self.schedule[0].strength_multiplier
|
||||
return all(
|
||||
boundary.strength_multiplier == first_multiplier
|
||||
for boundary in self.schedule[1:]
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RegionalLoraPlan:
|
||||
"""Store all adapter uses in authoritative global composition order."""
|
||||
|
||||
adapters: tuple[RegionalLoraAdapterPlan, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require immutable entries with contiguous composition indices."""
|
||||
|
||||
if not isinstance(self.adapters, tuple):
|
||||
raise TypeError("Regional LoRA plan adapters must be a tuple.")
|
||||
if any(
|
||||
not isinstance(adapter, RegionalLoraAdapterPlan)
|
||||
for adapter in self.adapters
|
||||
):
|
||||
raise TypeError("Regional LoRA plan contains an invalid adapter entry.")
|
||||
observed_indices = tuple(adapter.composition_index for adapter in self.adapters)
|
||||
if observed_indices != tuple(range(len(self.adapters))):
|
||||
raise ValueError(
|
||||
"Regional LoRA composition indices must be contiguous and ordered."
|
||||
)
|
||||
|
||||
@property
|
||||
def is_time_invariant(self) -> bool:
|
||||
"""Report whether every regional adapter retains one effective strength."""
|
||||
|
||||
return all(adapter.is_time_invariant for adapter in self.adapters)
|
||||
|
||||
|
||||
def _require_finite_float(value: object, *, name: str) -> None:
|
||||
"""Require one exact finite floating-point value."""
|
||||
|
||||
if not isinstance(value, float) or not math.isfinite(value):
|
||||
raise TypeError(f"Regional LoRA {name} must be a finite float.")
|
||||
|
||||
|
||||
def _require_non_negative_index(value: object, *, name: str) -> None:
|
||||
"""Require one non-negative integer index."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Regional LoRA {name} must be an integer.")
|
||||
if value < 0:
|
||||
raise ValueError(f"Regional LoRA {name} must be non-negative.")
|
||||
|
||||
|
||||
EMPTY_REGIONAL_LORA_PLAN = RegionalLoraPlan(adapters=())
|
||||
@@ -0,0 +1,75 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define the immutable full-canvas authority for regional masks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalMaskBank:
|
||||
"""Hold separate planning and conditioning masks on one latent canvas."""
|
||||
|
||||
planning_masks: torch.Tensor
|
||||
conditioning_masks: torch.Tensor
|
||||
canvas_width: int
|
||||
canvas_height: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject malformed, divergent, or aliased canonical mask tensors."""
|
||||
|
||||
if self.canvas_width < 1 or self.canvas_height < 1:
|
||||
raise ValueError("Regional mask canvas dimensions must be positive.")
|
||||
self._validate_mask_batch(self.planning_masks, name="Planning")
|
||||
self._validate_mask_batch(self.conditioning_masks, name="Conditioning")
|
||||
if self.planning_masks.shape != self.conditioning_masks.shape:
|
||||
raise ValueError(
|
||||
"Regional planning and conditioning mask shapes must match."
|
||||
)
|
||||
if self.planning_masks.dtype != self.conditioning_masks.dtype:
|
||||
raise ValueError(
|
||||
"Regional planning and conditioning mask dtypes must match."
|
||||
)
|
||||
if self.planning_masks.device != self.conditioning_masks.device:
|
||||
raise ValueError(
|
||||
"Regional planning and conditioning mask devices must match."
|
||||
)
|
||||
if (
|
||||
self.planning_masks.untyped_storage().data_ptr()
|
||||
== self.conditioning_masks.untyped_storage().data_ptr()
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional planning and conditioning masks must not share storage."
|
||||
)
|
||||
|
||||
@property
|
||||
def region_count(self) -> int:
|
||||
"""Return the number of ordered authored regions."""
|
||||
|
||||
return int(self.planning_masks.shape[0])
|
||||
|
||||
def _validate_mask_batch(self, masks: torch.Tensor, *, name: str) -> None:
|
||||
"""Validate one normalized floating-point BHW mask batch."""
|
||||
|
||||
if not isinstance(masks, torch.Tensor):
|
||||
raise TypeError(f"{name} regional masks must be a torch.Tensor.")
|
||||
if masks.ndim != 3:
|
||||
raise ValueError(f"{name} regional masks must use BHW layout.")
|
||||
if int(masks.shape[0]) < 1:
|
||||
raise ValueError(f"{name} regional masks require at least one region.")
|
||||
if tuple(masks.shape[1:]) != (self.canvas_height, self.canvas_width):
|
||||
raise ValueError(
|
||||
f"{name} regional masks must match the full latent canvas "
|
||||
f"{self.canvas_width}x{self.canvas_height}."
|
||||
)
|
||||
if not masks.is_floating_point():
|
||||
raise TypeError(f"{name} regional masks must use a floating-point dtype.")
|
||||
if not bool(torch.isfinite(masks).all()):
|
||||
raise ValueError(f"{name} regional masks must contain only finite values.")
|
||||
if not bool(((masks >= 0.0) & (masks <= 1.0)).all()):
|
||||
raise ValueError(f"{name} regional masks must stay within [0, 1].")
|
||||
@@ -0,0 +1,173 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define immutable capability values for regional attention backends."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class RegionalModelFamily(StrEnum):
|
||||
"""Identify a defensively admitted regional model family."""
|
||||
|
||||
ANIMA = "anima"
|
||||
STANDARD_UNET = "standard_unet"
|
||||
|
||||
|
||||
class RegionalAttentionBackend(StrEnum):
|
||||
"""Identify the model-specific attention patch backend."""
|
||||
|
||||
ANIMA_OBJECT_PATCH = "anima_object_patch"
|
||||
UNET_ATTN2_PATCH = "unet_attn2_patch"
|
||||
|
||||
|
||||
class RegionalAttentionTopology(StrEnum):
|
||||
"""Identify the image/context token roles owned by an attention backend."""
|
||||
|
||||
SEPARATE_IMAGE_AND_CONTEXT = "separate_image_and_context"
|
||||
SINGLETON_FRAME_SPATIOTEMPORAL = "singleton_frame_spatiotemporal"
|
||||
|
||||
|
||||
class RegionalLatentLayout(StrEnum):
|
||||
"""Identify the latent rank and temporal layout admitted by a backend."""
|
||||
|
||||
ANIMA_SINGLE_FRAME_BCTHW = "anima_single_frame_bcthw"
|
||||
STANDARD_IMAGE_BCHW = "standard_image_bchw"
|
||||
|
||||
|
||||
class RegionalSpatialPatchSupport(StrEnum):
|
||||
"""Report whether one backend can consume canonical spatial views."""
|
||||
|
||||
FULL_AND_SPATIAL_VIEWS = "full_and_spatial_views"
|
||||
|
||||
|
||||
class RegionalControlGligenPolicy(StrEnum):
|
||||
"""Report control and GLIGEN admission for attention coupling."""
|
||||
|
||||
REJECT = "reject"
|
||||
|
||||
|
||||
class RegionalReferenceLatentPolicy(StrEnum):
|
||||
"""Report reference-latent admission for attention coupling."""
|
||||
|
||||
REJECT = "reject"
|
||||
|
||||
|
||||
class RegionalPatchConflict(StrEnum):
|
||||
"""Identify a patch surface that must be collision-free before mutation."""
|
||||
|
||||
DIFFUSION_MODEL_WRAPPER = "diffusion_model_wrapper"
|
||||
CROSS_ATTENTION_OBJECT_PATCH = "cross_attention_object_patch"
|
||||
ATTN2_INPUT_PATCH = "attn2_input_patch"
|
||||
ATTN2_OUTPUT_PATCH = "attn2_output_patch"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalModelCapabilities:
|
||||
"""Describe one admitted backend and every relevant compatibility policy."""
|
||||
|
||||
model_family: RegionalModelFamily
|
||||
attention_backend: RegionalAttentionBackend
|
||||
attention_topology: RegionalAttentionTopology
|
||||
latent_layout: RegionalLatentLayout
|
||||
spatial_patch_support: RegionalSpatialPatchSupport
|
||||
control_gligen_policy: RegionalControlGligenPolicy
|
||||
reference_latent_policy: RegionalReferenceLatentPolicy
|
||||
known_patch_conflicts: tuple[RegionalPatchConflict, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject mutable, duplicate, or internally inconsistent capabilities."""
|
||||
|
||||
enum_fields = (
|
||||
("model family", self.model_family, RegionalModelFamily),
|
||||
("attention backend", self.attention_backend, RegionalAttentionBackend),
|
||||
(
|
||||
"attention topology",
|
||||
self.attention_topology,
|
||||
RegionalAttentionTopology,
|
||||
),
|
||||
("latent layout", self.latent_layout, RegionalLatentLayout),
|
||||
(
|
||||
"spatial patch support",
|
||||
self.spatial_patch_support,
|
||||
RegionalSpatialPatchSupport,
|
||||
),
|
||||
(
|
||||
"control/GLIGEN policy",
|
||||
self.control_gligen_policy,
|
||||
RegionalControlGligenPolicy,
|
||||
),
|
||||
(
|
||||
"reference-latent policy",
|
||||
self.reference_latent_policy,
|
||||
RegionalReferenceLatentPolicy,
|
||||
),
|
||||
)
|
||||
for name, value, enum_type in enum_fields:
|
||||
if not isinstance(value, enum_type):
|
||||
raise TypeError(
|
||||
f"Regional {name} must be a {enum_type.__name__} value."
|
||||
)
|
||||
if not isinstance(self.known_patch_conflicts, tuple):
|
||||
raise TypeError(
|
||||
"Known regional patch conflicts must be an immutable tuple."
|
||||
)
|
||||
if not self.known_patch_conflicts:
|
||||
raise ValueError("Regional capabilities require known patch conflicts.")
|
||||
if not all(
|
||||
isinstance(conflict, RegionalPatchConflict)
|
||||
for conflict in self.known_patch_conflicts
|
||||
):
|
||||
raise TypeError(
|
||||
"Known regional patch conflicts must contain "
|
||||
"RegionalPatchConflict values."
|
||||
)
|
||||
if len(set(self.known_patch_conflicts)) != len(self.known_patch_conflicts):
|
||||
raise ValueError(
|
||||
"Known regional patch conflicts must be unique and ordered."
|
||||
)
|
||||
self._validate_family_contract()
|
||||
|
||||
def _validate_family_contract(self) -> None:
|
||||
"""Require the exact backend, layout, and conflict surface for a family."""
|
||||
|
||||
expected: tuple[
|
||||
RegionalAttentionBackend,
|
||||
RegionalAttentionTopology,
|
||||
RegionalLatentLayout,
|
||||
tuple[RegionalPatchConflict, ...],
|
||||
]
|
||||
if self.model_family is RegionalModelFamily.ANIMA:
|
||||
expected = (
|
||||
RegionalAttentionBackend.ANIMA_OBJECT_PATCH,
|
||||
RegionalAttentionTopology.SINGLETON_FRAME_SPATIOTEMPORAL,
|
||||
RegionalLatentLayout.ANIMA_SINGLE_FRAME_BCTHW,
|
||||
(
|
||||
RegionalPatchConflict.DIFFUSION_MODEL_WRAPPER,
|
||||
RegionalPatchConflict.CROSS_ATTENTION_OBJECT_PATCH,
|
||||
RegionalPatchConflict.ATTN2_INPUT_PATCH,
|
||||
RegionalPatchConflict.ATTN2_OUTPUT_PATCH,
|
||||
),
|
||||
)
|
||||
else:
|
||||
expected = (
|
||||
RegionalAttentionBackend.UNET_ATTN2_PATCH,
|
||||
RegionalAttentionTopology.SEPARATE_IMAGE_AND_CONTEXT,
|
||||
RegionalLatentLayout.STANDARD_IMAGE_BCHW,
|
||||
(
|
||||
RegionalPatchConflict.ATTN2_INPUT_PATCH,
|
||||
RegionalPatchConflict.ATTN2_OUTPUT_PATCH,
|
||||
),
|
||||
)
|
||||
if (
|
||||
self.attention_backend,
|
||||
self.attention_topology,
|
||||
self.latent_layout,
|
||||
self.known_patch_conflicts,
|
||||
) != expected:
|
||||
raise ValueError(
|
||||
"Regional model capabilities do not match the model-family contract."
|
||||
)
|
||||
@@ -0,0 +1,106 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Constrain semantic tiled diffusion by authored regional composition masks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from .segs import coerce_segs
|
||||
from .segs_tiled_diffusion import segs_ownership_masks, validate_segs_aspect_ratio
|
||||
from .semantic_tiled_diffusion import build_semantic_tiled_diffusion_plan
|
||||
from .tiled_diffusion import TiledDiffusionPlan
|
||||
|
||||
REGIONAL_PLANNING_THRESHOLD = 0.5
|
||||
|
||||
|
||||
def build_region_constrained_tiled_diffusion_plan(
|
||||
*,
|
||||
region_masks: torch.Tensor,
|
||||
segs: object | None,
|
||||
latent_width: int,
|
||||
latent_height: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
tile_batch_size: int,
|
||||
) -> TiledDiffusionPlan:
|
||||
"""Build tiles split wherever regional composition or optional SEGS change."""
|
||||
|
||||
region_ownership = regional_composition_ownership_masks(
|
||||
region_masks,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
ownership_masks = region_ownership
|
||||
if segs is not None:
|
||||
native_segs = coerce_segs(segs)
|
||||
validate_segs_aspect_ratio(native_segs, latent_height, latent_width)
|
||||
semantic_ownership = segs_ownership_masks(
|
||||
native_segs,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
ownership_masks = _intersect_partitions(
|
||||
region_ownership,
|
||||
semantic_ownership,
|
||||
)
|
||||
return build_semantic_tiled_diffusion_plan(
|
||||
ownership_masks=ownership_masks,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
overlap=overlap,
|
||||
tile_batch_size=tile_batch_size,
|
||||
merge_across_masks=False,
|
||||
)
|
||||
|
||||
|
||||
def regional_composition_ownership_masks(
|
||||
region_masks: torch.Tensor,
|
||||
*,
|
||||
latent_height: int,
|
||||
latent_width: int,
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
"""Partition the canvas by every distinct active regional-mask combination."""
|
||||
|
||||
if region_masks.ndim != 3:
|
||||
raise ValueError("Regional tile planning requires a BHW mask batch.")
|
||||
if tuple(region_masks.shape[1:]) != (latent_height, latent_width):
|
||||
raise ValueError(
|
||||
"Regional tile planning masks must match latent shape "
|
||||
f"{latent_height}x{latent_width}."
|
||||
)
|
||||
membership = (region_masks.detach().cpu() >= REGIONAL_PLANNING_THRESHOLD).permute(
|
||||
1, 2, 0
|
||||
)
|
||||
flattened = membership.reshape(latent_height * latent_width, -1)
|
||||
signatures, inverse = torch.unique(
|
||||
flattened,
|
||||
dim=0,
|
||||
sorted=True,
|
||||
return_inverse=True,
|
||||
)
|
||||
del signatures
|
||||
labels = inverse.reshape(latent_height, latent_width)
|
||||
return tuple(labels == index for index in range(int(labels.max().item()) + 1))
|
||||
|
||||
|
||||
def _intersect_partitions(
|
||||
first: tuple[torch.Tensor, ...],
|
||||
second: tuple[torch.Tensor, ...],
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
"""Return non-empty intersections of two complete ownership partitions."""
|
||||
|
||||
intersections = tuple(
|
||||
intersection
|
||||
for first_mask in first
|
||||
for second_mask in second
|
||||
if bool((intersection := torch.logical_and(first_mask, second_mask)).any())
|
||||
)
|
||||
if not intersections:
|
||||
raise ValueError("Regional and SEGS ownership produced no tile coverage.")
|
||||
return intersections
|
||||
@@ -0,0 +1,347 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define model-neutral resolved regional ordinary-LoRA operation contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .regional_lora_plan import RegionalLoraAdapterPlan
|
||||
|
||||
|
||||
class ResolvedLoraOperationClass(StrEnum):
|
||||
"""Classify normalized low-rank tensor organization before module binding."""
|
||||
|
||||
MATRIX_PAIR = "matrix_pair"
|
||||
CONVOLUTION_1D = "convolution_1d"
|
||||
CONVOLUTION_2D = "convolution_2d"
|
||||
CONVOLUTION_3D = "convolution_3d"
|
||||
UNSUPPORTED = "unsupported"
|
||||
|
||||
|
||||
class RegionalLoraGeometryClass(StrEnum):
|
||||
"""Declare the activation geometry an execution contract requires."""
|
||||
|
||||
TARGET_OPERATION = "target_operation"
|
||||
DIRECT_CONVOLUTION_1D = "direct_convolution_1d"
|
||||
DIRECT_CONVOLUTION_2D = "direct_convolution_2d"
|
||||
DIRECT_CONVOLUTION_3D = "direct_convolution_3d"
|
||||
UNSUPPORTED = "unsupported"
|
||||
|
||||
|
||||
class RegionalLoraExecutionContract(StrEnum):
|
||||
"""Declare exact rank-activation execution or explicit rejection."""
|
||||
|
||||
MATRIX_OR_RESHAPED_CONVOLUTION = "matrix_or_reshaped_convolution"
|
||||
DIRECT_CONVOLUTION = "direct_convolution"
|
||||
UNSUPPORTED = "unsupported"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalLoraTensorShape:
|
||||
"""Retain one immutable positive tensor shape without tensor ownership."""
|
||||
|
||||
dimensions: tuple[int, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require a nonempty tuple of strictly positive integer dimensions."""
|
||||
|
||||
if not isinstance(self.dimensions, tuple) or not self.dimensions:
|
||||
raise ValueError("Regional LoRA tensor shape must be a nonempty tuple.")
|
||||
if any(
|
||||
isinstance(dimension, bool)
|
||||
or not isinstance(dimension, int)
|
||||
or dimension <= 0
|
||||
for dimension in self.dimensions
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional LoRA tensor shape dimensions must be positive integers."
|
||||
)
|
||||
|
||||
@property
|
||||
def rank(self) -> int:
|
||||
"""Return the tensor dimensionality."""
|
||||
|
||||
return len(self.dimensions)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedRegionalLoraTarget:
|
||||
"""Retain one model-relative parameter path and optional Comfy tensor slice."""
|
||||
|
||||
model_target: str
|
||||
parameter_name: str
|
||||
offset: tuple[int, ...] | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require an explicit module target, parameter, and valid optional slice."""
|
||||
|
||||
if not isinstance(self.model_target, str) or not self.model_target.strip():
|
||||
raise ValueError("Resolved regional LoRA model target must be nonempty.")
|
||||
if not isinstance(self.parameter_name, str) or not self.parameter_name.strip():
|
||||
raise ValueError("Resolved regional LoRA parameter name must be nonempty.")
|
||||
if self.offset is not None and (
|
||||
not isinstance(self.offset, tuple)
|
||||
or not self.offset
|
||||
or any(
|
||||
isinstance(part, bool) or not isinstance(part, int) or part < 0
|
||||
for part in self.offset
|
||||
)
|
||||
):
|
||||
raise ValueError(
|
||||
"Resolved regional LoRA target offset must contain "
|
||||
"non-negative integers."
|
||||
)
|
||||
|
||||
@property
|
||||
def parameter_path(self) -> str:
|
||||
"""Return the complete model-relative parameter path."""
|
||||
|
||||
return f"{self.model_target}.{self.parameter_name}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedRegionalLoraOperation:
|
||||
"""Describe one supported operation or one explicit target rejection."""
|
||||
|
||||
adapter: RegionalLoraAdapterPlan
|
||||
target_index: int
|
||||
target: ResolvedRegionalLoraTarget
|
||||
normalized_operation_type: str
|
||||
operation_class: ResolvedLoraOperationClass
|
||||
down_shape: RegionalLoraTensorShape | None
|
||||
up_shape: RegionalLoraTensorShape | None
|
||||
middle_shape: RegionalLoraTensorShape | None
|
||||
reshape_shape: RegionalLoraTensorShape | None
|
||||
rank: int | None
|
||||
intrinsic_scale: float | None
|
||||
required_geometry: RegionalLoraGeometryClass
|
||||
execution_contract: RegionalLoraExecutionContract
|
||||
rejection_reason: str | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject mutable, incomplete, or mathematically inconsistent metadata."""
|
||||
|
||||
if not isinstance(self.adapter, RegionalLoraAdapterPlan):
|
||||
raise TypeError("Resolved regional LoRA operation requires an adapter.")
|
||||
_require_non_negative_index(self.target_index, name="target_index")
|
||||
if not isinstance(self.target, ResolvedRegionalLoraTarget):
|
||||
raise TypeError("Resolved regional LoRA operation requires a target.")
|
||||
if (
|
||||
not isinstance(self.normalized_operation_type, str)
|
||||
or not self.normalized_operation_type.strip()
|
||||
):
|
||||
raise ValueError(
|
||||
"Normalized regional LoRA operation type must be nonempty."
|
||||
)
|
||||
if not isinstance(self.operation_class, ResolvedLoraOperationClass):
|
||||
raise TypeError("Resolved regional LoRA operation class is invalid.")
|
||||
if not isinstance(self.required_geometry, RegionalLoraGeometryClass):
|
||||
raise TypeError("Resolved regional LoRA geometry class is invalid.")
|
||||
if not isinstance(self.execution_contract, RegionalLoraExecutionContract):
|
||||
raise TypeError("Resolved regional LoRA execution contract is invalid.")
|
||||
if self.execution_contract is RegionalLoraExecutionContract.UNSUPPORTED:
|
||||
self._validate_rejection()
|
||||
return
|
||||
self._validate_supported()
|
||||
|
||||
def _validate_rejection(self) -> None:
|
||||
"""Require one reason and no misleading supported-operation metadata."""
|
||||
|
||||
if self.operation_class is not ResolvedLoraOperationClass.UNSUPPORTED:
|
||||
raise ValueError(
|
||||
"Rejected regional LoRA operation class must be unsupported."
|
||||
)
|
||||
if self.required_geometry is not RegionalLoraGeometryClass.UNSUPPORTED:
|
||||
raise ValueError("Rejected regional LoRA geometry must be unsupported.")
|
||||
if (
|
||||
not isinstance(self.rejection_reason, str)
|
||||
or not self.rejection_reason.strip()
|
||||
):
|
||||
raise ValueError("Rejected regional LoRA operation requires a reason.")
|
||||
if any(
|
||||
value is not None
|
||||
for value in (
|
||||
self.down_shape,
|
||||
self.up_shape,
|
||||
self.middle_shape,
|
||||
self.reshape_shape,
|
||||
self.rank,
|
||||
self.intrinsic_scale,
|
||||
)
|
||||
):
|
||||
raise ValueError(
|
||||
"Rejected regional LoRA operation cannot claim executable metadata."
|
||||
)
|
||||
|
||||
def _validate_supported(self) -> None:
|
||||
"""Require complete shape, scale, geometry, and execution consistency."""
|
||||
|
||||
if self.operation_class is ResolvedLoraOperationClass.UNSUPPORTED:
|
||||
raise ValueError("Supported regional LoRA operation cannot be unsupported.")
|
||||
if self.required_geometry is RegionalLoraGeometryClass.UNSUPPORTED:
|
||||
raise ValueError("Supported regional LoRA geometry cannot be unsupported.")
|
||||
if self.rejection_reason is not None:
|
||||
raise ValueError(
|
||||
"Supported regional LoRA operation cannot have a rejection."
|
||||
)
|
||||
if self.down_shape is None or self.up_shape is None:
|
||||
raise ValueError(
|
||||
"Supported regional LoRA operation requires down/up shapes."
|
||||
)
|
||||
if self.down_shape.rank < 2 or self.up_shape.rank < 2:
|
||||
raise ValueError(
|
||||
"Regional LoRA down/up shapes need at least two dimensions."
|
||||
)
|
||||
if (
|
||||
isinstance(self.rank, bool)
|
||||
or not isinstance(self.rank, int)
|
||||
or self.rank <= 0
|
||||
):
|
||||
raise ValueError("Supported regional LoRA rank must be a positive integer.")
|
||||
if not isinstance(self.intrinsic_scale, float) or not math.isfinite(
|
||||
self.intrinsic_scale
|
||||
):
|
||||
raise ValueError("Regional LoRA intrinsic scale must be finite.")
|
||||
if self.down_shape.dimensions[0] != self.rank:
|
||||
raise ValueError("Regional LoRA down shape must begin with its rank.")
|
||||
if (
|
||||
len(self.up_shape.dimensions) < 2
|
||||
or self.up_shape.dimensions[1] != self.rank
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional LoRA up shape must contain its rank at index one."
|
||||
)
|
||||
if self.middle_shape is not None:
|
||||
if self.middle_shape.rank < 2:
|
||||
raise ValueError(
|
||||
"Regional LoRA middle shape needs at least two dimensions."
|
||||
)
|
||||
if (
|
||||
self.middle_shape.dimensions[0] != self.rank
|
||||
or self.middle_shape.dimensions[1] != self.rank
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional LoRA middle shape must preserve rank channels."
|
||||
)
|
||||
self._validate_operation_shape_contract()
|
||||
|
||||
def _validate_operation_shape_contract(self) -> None:
|
||||
"""Match operation, tensor dimensionality, geometry, and execution policy."""
|
||||
|
||||
if self.operation_class is ResolvedLoraOperationClass.MATRIX_PAIR:
|
||||
if self.down_shape is None or self.up_shape is None:
|
||||
raise AssertionError(
|
||||
"Supported shape validation requires down/up shapes."
|
||||
)
|
||||
if self.down_shape.rank != 2 or self.up_shape.rank != 2:
|
||||
raise ValueError("Matrix-pair regional LoRA tensors must be rank two.")
|
||||
if self.middle_shape is not None:
|
||||
raise ValueError(
|
||||
"Matrix-pair regional LoRA cannot contain middle weights."
|
||||
)
|
||||
if self.required_geometry is not RegionalLoraGeometryClass.TARGET_OPERATION:
|
||||
raise ValueError("Matrix-pair regional LoRA requires target geometry.")
|
||||
if (
|
||||
self.execution_contract
|
||||
is not RegionalLoraExecutionContract.MATRIX_OR_RESHAPED_CONVOLUTION
|
||||
):
|
||||
raise ValueError("Matrix-pair regional LoRA execution is inconsistent.")
|
||||
return
|
||||
spatial_rank = {
|
||||
ResolvedLoraOperationClass.CONVOLUTION_1D: 3,
|
||||
ResolvedLoraOperationClass.CONVOLUTION_2D: 4,
|
||||
ResolvedLoraOperationClass.CONVOLUTION_3D: 5,
|
||||
}[self.operation_class]
|
||||
expected_geometry = {
|
||||
3: RegionalLoraGeometryClass.DIRECT_CONVOLUTION_1D,
|
||||
4: RegionalLoraGeometryClass.DIRECT_CONVOLUTION_2D,
|
||||
5: RegionalLoraGeometryClass.DIRECT_CONVOLUTION_3D,
|
||||
}[spatial_rank]
|
||||
shapes = tuple(
|
||||
shape
|
||||
for shape in (self.down_shape, self.up_shape, self.middle_shape)
|
||||
if shape is not None
|
||||
)
|
||||
if any(shape.rank not in (2, spatial_rank) for shape in shapes):
|
||||
raise ValueError("Convolutional regional LoRA tensor rank is inconsistent.")
|
||||
if not any(shape.rank == spatial_rank for shape in shapes):
|
||||
raise ValueError("Convolutional regional LoRA needs one spatial tensor.")
|
||||
if self.required_geometry is not expected_geometry:
|
||||
raise ValueError("Convolutional regional LoRA geometry is inconsistent.")
|
||||
if (
|
||||
self.execution_contract
|
||||
is not RegionalLoraExecutionContract.DIRECT_CONVOLUTION
|
||||
):
|
||||
raise ValueError("Convolutional regional LoRA execution is inconsistent.")
|
||||
|
||||
@property
|
||||
def authored_strength(self) -> float:
|
||||
"""Return the authored model strength from the authoritative plan owner."""
|
||||
|
||||
return self.adapter.model_strength
|
||||
|
||||
@property
|
||||
def supported(self) -> bool:
|
||||
"""Report whether this target has an executable ordinary-LoRA contract."""
|
||||
|
||||
return self.execution_contract is not RegionalLoraExecutionContract.UNSUPPORTED
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ResolvedRegionalLoraOperationSet:
|
||||
"""Retain canonical adapter/target order across supported and rejected entries."""
|
||||
|
||||
entries: tuple[ResolvedRegionalLoraOperation, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require immutable entries in contiguous per-adapter target order."""
|
||||
|
||||
if not isinstance(self.entries, tuple):
|
||||
raise TypeError("Resolved regional LoRA entries must be a tuple.")
|
||||
if any(
|
||||
not isinstance(entry, ResolvedRegionalLoraOperation)
|
||||
for entry in self.entries
|
||||
):
|
||||
raise TypeError("Resolved regional LoRA set contains an invalid entry.")
|
||||
observed = tuple(
|
||||
(entry.adapter.composition_index, entry.target_index)
|
||||
for entry in self.entries
|
||||
)
|
||||
if observed != tuple(sorted(observed)):
|
||||
raise ValueError(
|
||||
"Resolved regional LoRA entries must retain declared order."
|
||||
)
|
||||
indices_by_adapter: dict[int, list[int]] = {}
|
||||
for composition_index, target_index in observed:
|
||||
indices_by_adapter.setdefault(composition_index, []).append(target_index)
|
||||
if any(
|
||||
target_indices != list(range(len(target_indices)))
|
||||
for target_indices in indices_by_adapter.values()
|
||||
):
|
||||
raise ValueError(
|
||||
"Resolved regional LoRA target indices must be contiguous per adapter."
|
||||
)
|
||||
|
||||
@property
|
||||
def supported(self) -> tuple[ResolvedRegionalLoraOperation, ...]:
|
||||
"""Return supported entries without changing canonical relative order."""
|
||||
|
||||
return tuple(entry for entry in self.entries if entry.supported)
|
||||
|
||||
@property
|
||||
def rejected(self) -> tuple[ResolvedRegionalLoraOperation, ...]:
|
||||
"""Return rejected entries without changing canonical relative order."""
|
||||
|
||||
return tuple(entry for entry in self.entries if not entry.supported)
|
||||
|
||||
|
||||
def _require_non_negative_index(value: object, *, name: str) -> None:
|
||||
"""Require one non-negative non-boolean integer index."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
|
||||
raise ValueError(f"Resolved regional LoRA {name} must be non-negative.")
|
||||
@@ -0,0 +1,54 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define validated seed-variation sampling settings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from math import isfinite
|
||||
|
||||
MIN_SEED = 0
|
||||
MAX_SEED = 0xFFFFFFFFFFFFFFFF
|
||||
MIN_VARIATION_STRENGTH = 0.0
|
||||
MAX_VARIATION_STRENGTH = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SeedVariationSettings:
|
||||
"""Hold one deterministic initial-noise interpolation request."""
|
||||
|
||||
variation_seed: int
|
||||
strength: float
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject settings outside ComfyUI's public seed and strength ranges."""
|
||||
|
||||
if isinstance(self.variation_seed, bool) or not isinstance(
|
||||
self.variation_seed,
|
||||
int,
|
||||
):
|
||||
raise TypeError("Variation seed must be an integer.")
|
||||
if not MIN_SEED <= self.variation_seed <= MAX_SEED:
|
||||
raise ValueError(
|
||||
f"Variation seed must be between {MIN_SEED} and {MAX_SEED}."
|
||||
)
|
||||
if isinstance(self.strength, bool) or not isinstance(
|
||||
self.strength,
|
||||
(int, float),
|
||||
):
|
||||
raise TypeError("Variation strength must be a number.")
|
||||
normalized_strength = float(self.strength)
|
||||
if not isfinite(normalized_strength):
|
||||
raise ValueError("Variation strength must be finite.")
|
||||
if (
|
||||
not MIN_VARIATION_STRENGTH
|
||||
<= normalized_strength
|
||||
<= (MAX_VARIATION_STRENGTH)
|
||||
):
|
||||
raise ValueError(
|
||||
"Variation strength must be between "
|
||||
f"{MIN_VARIATION_STRENGTH} and {MAX_VARIATION_STRENGTH}."
|
||||
)
|
||||
object.__setattr__(self, "strength", normalized_strength)
|
||||
@@ -6,27 +6,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from .segs import NativeSegs, Segment, coerce_segment_mask, coerce_segs
|
||||
from .tiled_diffusion import (
|
||||
LatentTile,
|
||||
TiledDiffusionPlan,
|
||||
batch_latent_tiles,
|
||||
build_tiled_diffusion_plan,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _OwnershipCore:
|
||||
"""Represent a non-overlapping latent ownership region before window placement."""
|
||||
|
||||
mask: torch.Tensor
|
||||
bounds: tuple[int, int, int, int]
|
||||
area: int
|
||||
from .semantic_tiled_diffusion import build_semantic_tiled_diffusion_plan
|
||||
from .tiled_diffusion import TiledDiffusionPlan
|
||||
|
||||
|
||||
def build_segs_guided_tiled_diffusion_plan(
|
||||
@@ -46,64 +30,22 @@ def build_segs_guided_tiled_diffusion_plan(
|
||||
boundary and shares a feathered overlap with neighboring cores.
|
||||
"""
|
||||
|
||||
base_plan = build_tiled_diffusion_plan(
|
||||
native_segs = coerce_segs(segs)
|
||||
validate_segs_aspect_ratio(native_segs, latent_height, latent_width)
|
||||
ownership_masks = segs_ownership_masks(
|
||||
native_segs,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
return build_semantic_tiled_diffusion_plan(
|
||||
ownership_masks=ownership_masks,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
overlap=overlap,
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
native_segs = coerce_segs(segs)
|
||||
validate_segs_aspect_ratio(native_segs, latent_height, latent_width)
|
||||
ownership = _build_ownership_cores(
|
||||
native_segs,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
max_core_width = max(1, base_plan.tile_width - base_plan.overlap)
|
||||
max_core_height = max(1, base_plan.tile_height - base_plan.overlap)
|
||||
split_cores = tuple(
|
||||
split_core
|
||||
for core in ownership
|
||||
for split_core in _split_core(
|
||||
core,
|
||||
max_width=max_core_width,
|
||||
max_height=max_core_height,
|
||||
)
|
||||
)
|
||||
merged_cores = _merge_small_cores(
|
||||
split_cores,
|
||||
max_width=max_core_width,
|
||||
max_height=max_core_height,
|
||||
)
|
||||
tiles = tuple(
|
||||
sorted(
|
||||
(
|
||||
_tile_for_core(
|
||||
core,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=base_plan.tile_width,
|
||||
tile_height=base_plan.tile_height,
|
||||
overlap=base_plan.overlap,
|
||||
)
|
||||
for core in merged_cores
|
||||
),
|
||||
key=lambda tile: (tile.y, tile.x),
|
||||
)
|
||||
)
|
||||
batches, effective_batch_size = batch_latent_tiles(tiles, tile_batch_size)
|
||||
return TiledDiffusionPlan(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=base_plan.tile_width,
|
||||
tile_height=base_plan.tile_height,
|
||||
overlap=base_plan.overlap,
|
||||
requested_tile_batch_size=tile_batch_size,
|
||||
tile_batch_size=effective_batch_size,
|
||||
tiles=tiles,
|
||||
batches=batches,
|
||||
merge_across_masks=True,
|
||||
)
|
||||
|
||||
|
||||
@@ -126,13 +68,13 @@ def validate_segs_aspect_ratio(
|
||||
)
|
||||
|
||||
|
||||
def _build_ownership_cores(
|
||||
def segs_ownership_masks(
|
||||
segs: NativeSegs,
|
||||
*,
|
||||
latent_height: int,
|
||||
latent_width: int,
|
||||
) -> tuple[_OwnershipCore, ...]:
|
||||
"""Resolve overlapping SEGS into one deterministic latent ownership partition."""
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
"""Resolve overlapping SEGS into a deterministic latent ownership partition."""
|
||||
|
||||
source_height, source_width = segs[0]
|
||||
segment_masks = tuple(
|
||||
@@ -154,20 +96,18 @@ def _build_ownership_cores(
|
||||
),
|
||||
)
|
||||
occupied = torch.zeros((latent_height, latent_width), dtype=torch.bool)
|
||||
cores: list[_OwnershipCore] = []
|
||||
ownership_masks: list[torch.Tensor] = []
|
||||
for index in ranked_indexes:
|
||||
owned = torch.logical_and(segment_masks[index], torch.logical_not(occupied))
|
||||
if bool(owned.any()):
|
||||
cores.append(_core_from_mask(owned))
|
||||
ownership_masks.append(owned)
|
||||
occupied = torch.logical_or(occupied, segment_masks[index])
|
||||
background = torch.logical_not(occupied)
|
||||
if bool(background.any()):
|
||||
cores.append(_core_from_mask(background))
|
||||
if cores:
|
||||
return tuple(cores)
|
||||
return (
|
||||
_core_from_mask(torch.ones((latent_height, latent_width), dtype=torch.bool)),
|
||||
)
|
||||
ownership_masks.append(background)
|
||||
if ownership_masks:
|
||||
return tuple(ownership_masks)
|
||||
return (torch.ones((latent_height, latent_width), dtype=torch.bool),)
|
||||
|
||||
|
||||
def segment_mask_to_latent(
|
||||
@@ -259,223 +199,6 @@ def segment_weight_to_latent(
|
||||
return latent_mask
|
||||
|
||||
|
||||
def _split_core(
|
||||
core: _OwnershipCore,
|
||||
*,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> tuple[_OwnershipCore, ...]:
|
||||
"""Recursively divide a core into balanced pieces that fit its tile budget."""
|
||||
|
||||
left, top, right, bottom = core.bounds
|
||||
width = right - left
|
||||
height = bottom - top
|
||||
if width <= max_width and height <= max_height:
|
||||
return (core,)
|
||||
split_x = width / max_width >= height / max_height
|
||||
first, second = _split_mask_at_balanced_axis(core.mask, core.bounds, split_x)
|
||||
return _split_core(
|
||||
_core_from_mask(first), max_width=max_width, max_height=max_height
|
||||
) + _split_core(_core_from_mask(second), max_width=max_width, max_height=max_height)
|
||||
|
||||
|
||||
def _split_mask_at_balanced_axis(
|
||||
mask: torch.Tensor,
|
||||
bounds: tuple[int, int, int, int],
|
||||
split_x: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Split one non-empty mask near its active-pixel median on one axis."""
|
||||
|
||||
left, top, right, bottom = bounds
|
||||
counts = (
|
||||
mask[top:bottom, left:right].sum(dim=0)
|
||||
if split_x
|
||||
else mask[top:bottom, left:right].sum(dim=1)
|
||||
)
|
||||
cumulative = torch.cumsum(counts, dim=0)
|
||||
midpoint = int(torch.searchsorted(cumulative, cumulative[-1] / 2, right=False))
|
||||
axis_start = left if split_x else top
|
||||
axis_end = right if split_x else bottom
|
||||
split_at = min(axis_end - 1, max(axis_start + 1, axis_start + midpoint + 1))
|
||||
first = mask.clone()
|
||||
second = mask.clone()
|
||||
if split_x:
|
||||
first[:, split_at:] = False
|
||||
second[:, :split_at] = False
|
||||
else:
|
||||
first[split_at:, :] = False
|
||||
second[:split_at, :] = False
|
||||
if not bool(first.any()) or not bool(second.any()):
|
||||
raise ValueError("Unable to split an oversized SEGS-guided tile core.")
|
||||
return first, second
|
||||
|
||||
|
||||
def _merge_small_cores(
|
||||
cores: tuple[_OwnershipCore, ...],
|
||||
*,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> tuple[_OwnershipCore, ...]:
|
||||
"""Greedily combine small nearby cores when one bounded window can hold both."""
|
||||
|
||||
pending = list(cores)
|
||||
minimum_area = max(1, (max_width * max_height) // 4)
|
||||
merged = True
|
||||
while merged:
|
||||
merged = False
|
||||
for index, core in enumerate(tuple(pending)):
|
||||
if core.area >= minimum_area:
|
||||
continue
|
||||
candidate_index = _best_merge_candidate_index(
|
||||
core,
|
||||
pending,
|
||||
excluded_index=index,
|
||||
max_width=max_width,
|
||||
max_height=max_height,
|
||||
)
|
||||
if candidate_index is None:
|
||||
continue
|
||||
candidate = pending[candidate_index]
|
||||
pending[index] = _OwnershipCore(
|
||||
mask=torch.logical_or(core.mask, candidate.mask),
|
||||
bounds=_union_bounds(core.bounds, candidate.bounds),
|
||||
area=core.area + candidate.area,
|
||||
)
|
||||
pending.pop(candidate_index)
|
||||
merged = True
|
||||
break
|
||||
return tuple(pending)
|
||||
|
||||
|
||||
def _best_merge_candidate_index(
|
||||
core: _OwnershipCore,
|
||||
candidates: list[_OwnershipCore],
|
||||
*,
|
||||
excluded_index: int,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> int | None:
|
||||
"""Return a candidate index whose combined bounds fit one ownership budget."""
|
||||
|
||||
eligible: list[tuple[int, int, int]] = []
|
||||
for index, candidate in enumerate(candidates):
|
||||
if index == excluded_index:
|
||||
continue
|
||||
bounds = _union_bounds(core.bounds, candidate.bounds)
|
||||
left, top, right, bottom = bounds
|
||||
width = right - left
|
||||
height = bottom - top
|
||||
if width > max_width or height > max_height:
|
||||
continue
|
||||
distance = _bounds_distance(core.bounds, candidate.bounds)
|
||||
eligible.append((width * height, distance, index))
|
||||
if not eligible:
|
||||
return None
|
||||
return min(eligible, key=lambda item: (item[0], item[1]))[2]
|
||||
|
||||
|
||||
def _bounds_distance(
|
||||
first: tuple[int, int, int, int],
|
||||
second: tuple[int, int, int, int],
|
||||
) -> int:
|
||||
"""Return the axis-aligned gap between two mask bounding boxes."""
|
||||
|
||||
left, top, right, bottom = first
|
||||
other_left, other_top, other_right, other_bottom = second
|
||||
horizontal = max(0, other_left - right, left - other_right)
|
||||
vertical = max(0, other_top - bottom, top - other_bottom)
|
||||
return horizontal + vertical
|
||||
|
||||
|
||||
def _union_bounds(
|
||||
first: tuple[int, int, int, int],
|
||||
second: tuple[int, int, int, int],
|
||||
) -> tuple[int, int, int, int]:
|
||||
"""Return the tight rectangle containing both ownership-core bounds."""
|
||||
|
||||
return (
|
||||
min(first[0], second[0]),
|
||||
min(first[1], second[1]),
|
||||
max(first[2], second[2]),
|
||||
max(first[3], second[3]),
|
||||
)
|
||||
|
||||
|
||||
def _tile_for_core(
|
||||
core: _OwnershipCore,
|
||||
*,
|
||||
latent_width: int,
|
||||
latent_height: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
) -> LatentTile:
|
||||
"""Place one bounded sampling window around an irregular ownership core."""
|
||||
|
||||
left, top, right, bottom = core.bounds
|
||||
center_x = (left + right) / 2.0
|
||||
center_y = (top + bottom) / 2.0
|
||||
x = _clamp_window_start(center_x, tile_width, latent_width)
|
||||
y = _clamp_window_start(center_y, tile_height, latent_height)
|
||||
weight_mask = _feathered_tile_weight(
|
||||
core.mask,
|
||||
x=x,
|
||||
y=y,
|
||||
width=tile_width,
|
||||
height=tile_height,
|
||||
overlap=overlap,
|
||||
)
|
||||
if not bool((weight_mask > 0).any()):
|
||||
raise ValueError("SEGS-guided tiled diffusion generated an empty tile weight.")
|
||||
return LatentTile(x, y, tile_width, tile_height, weight_mask)
|
||||
|
||||
|
||||
def _clamp_window_start(center: float, window_size: int, limit: int) -> int:
|
||||
"""Center a fixed sampling window while keeping it inside the latent bounds."""
|
||||
|
||||
desired = round(center - window_size / 2.0)
|
||||
return min(max(0, desired), limit - window_size)
|
||||
|
||||
|
||||
def _feathered_tile_weight(
|
||||
mask: torch.Tensor,
|
||||
*,
|
||||
x: int,
|
||||
y: int,
|
||||
width: int,
|
||||
height: int,
|
||||
overlap: int,
|
||||
) -> torch.Tensor:
|
||||
"""Build one feathered tile weight without blurring the full latent mask."""
|
||||
|
||||
if overlap == 0:
|
||||
return mask[y : y + height, x : x + width].float().contiguous()
|
||||
radius = max(1, overlap // 2)
|
||||
source_left = max(0, x - radius)
|
||||
source_top = max(0, y - radius)
|
||||
source_right = min(int(mask.shape[1]), x + width + radius)
|
||||
source_bottom = min(int(mask.shape[0]), y + height + radius)
|
||||
local_weight = (
|
||||
functional.avg_pool2d(
|
||||
mask[source_top:source_bottom, source_left:source_right]
|
||||
.float()
|
||||
.unsqueeze(0)
|
||||
.unsqueeze(0),
|
||||
kernel_size=radius * 2 + 1,
|
||||
stride=1,
|
||||
padding=radius,
|
||||
count_include_pad=False,
|
||||
)
|
||||
.squeeze(0)
|
||||
.squeeze(0)
|
||||
)
|
||||
local_y = y - source_top
|
||||
local_x = x - source_left
|
||||
return local_weight[
|
||||
local_y : local_y + height, local_x : local_x + width
|
||||
].contiguous()
|
||||
|
||||
|
||||
def _latent_sample_range(
|
||||
source_start: int,
|
||||
source_end: int,
|
||||
@@ -487,30 +210,3 @@ def _latent_sample_range(
|
||||
start = (source_start * latent_limit + source_limit - 1) // source_limit
|
||||
end = (source_end * latent_limit + source_limit - 1) // source_limit
|
||||
return max(0, min(latent_limit, start)), max(0, min(latent_limit, end))
|
||||
|
||||
|
||||
def _core_from_mask(mask: torch.Tensor) -> _OwnershipCore:
|
||||
"""Build one core with bounds and area computed exactly once."""
|
||||
|
||||
bounds = _mask_bounds(mask)
|
||||
if bounds is None:
|
||||
raise ValueError("SEGS-guided tiled diffusion cannot use an empty core.")
|
||||
return _OwnershipCore(
|
||||
mask=mask,
|
||||
bounds=bounds,
|
||||
area=int(mask.sum().item()),
|
||||
)
|
||||
|
||||
|
||||
def _mask_bounds(mask: torch.Tensor) -> tuple[int, int, int, int] | None:
|
||||
"""Return left, top, right, bottom bounds for one non-empty boolean mask."""
|
||||
|
||||
y_coords, x_coords = torch.where(mask)
|
||||
if y_coords.numel() == 0:
|
||||
return None
|
||||
return (
|
||||
int(x_coords.min().item()),
|
||||
int(y_coords.min().item()),
|
||||
int(x_coords.max().item()) + 1,
|
||||
int(y_coords.max().item()) + 1,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,384 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Build tiled diffusion plans from non-overlapping latent ownership masks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from .tiled_diffusion import (
|
||||
LatentTile,
|
||||
TiledDiffusionPlan,
|
||||
batch_latent_tiles,
|
||||
build_tiled_diffusion_plan,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _OwnershipCore:
|
||||
"""Represent one latent ownership region before sampling-window placement."""
|
||||
|
||||
mask: torch.Tensor
|
||||
bounds: tuple[int, int, int, int]
|
||||
area: int
|
||||
|
||||
|
||||
def build_semantic_tiled_diffusion_plan(
|
||||
*,
|
||||
ownership_masks: Sequence[torch.Tensor],
|
||||
latent_width: int,
|
||||
latent_height: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
tile_batch_size: int,
|
||||
merge_across_masks: bool,
|
||||
) -> TiledDiffusionPlan:
|
||||
"""Build bounded windows whose write weights follow ownership masks."""
|
||||
|
||||
base_plan = build_tiled_diffusion_plan(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
overlap=overlap,
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
normalized_masks = _validate_ownership_masks(
|
||||
ownership_masks,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
max_core_width = max(1, base_plan.tile_width - base_plan.overlap)
|
||||
max_core_height = max(1, base_plan.tile_height - base_plan.overlap)
|
||||
split_groups = tuple(
|
||||
tuple(
|
||||
split_core
|
||||
for split_core in _split_core(
|
||||
_core_from_mask(mask),
|
||||
max_width=max_core_width,
|
||||
max_height=max_core_height,
|
||||
)
|
||||
)
|
||||
for mask in normalized_masks
|
||||
)
|
||||
if merge_across_masks:
|
||||
cores = _merge_small_cores(
|
||||
tuple(core for group in split_groups for core in group),
|
||||
max_width=max_core_width,
|
||||
max_height=max_core_height,
|
||||
)
|
||||
else:
|
||||
cores = tuple(
|
||||
core
|
||||
for group in split_groups
|
||||
for core in _merge_small_cores(
|
||||
group,
|
||||
max_width=max_core_width,
|
||||
max_height=max_core_height,
|
||||
)
|
||||
)
|
||||
tiles = tuple(
|
||||
sorted(
|
||||
(
|
||||
_tile_for_core(
|
||||
core,
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=base_plan.tile_width,
|
||||
tile_height=base_plan.tile_height,
|
||||
overlap=base_plan.overlap,
|
||||
)
|
||||
for core in cores
|
||||
),
|
||||
key=lambda tile: (tile.y, tile.x),
|
||||
)
|
||||
)
|
||||
batches, effective_batch_size = batch_latent_tiles(tiles, tile_batch_size)
|
||||
return TiledDiffusionPlan(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
tile_width=base_plan.tile_width,
|
||||
tile_height=base_plan.tile_height,
|
||||
overlap=base_plan.overlap,
|
||||
requested_tile_batch_size=tile_batch_size,
|
||||
tile_batch_size=effective_batch_size,
|
||||
tiles=tiles,
|
||||
batches=batches,
|
||||
)
|
||||
|
||||
|
||||
def _validate_ownership_masks(
|
||||
masks: Sequence[torch.Tensor],
|
||||
*,
|
||||
latent_height: int,
|
||||
latent_width: int,
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
"""Return non-empty boolean masks that cover the complete latent canvas."""
|
||||
|
||||
normalized: list[torch.Tensor] = []
|
||||
coverage = torch.zeros((latent_height, latent_width), dtype=torch.bool)
|
||||
for index, mask in enumerate(masks):
|
||||
if mask.ndim != 2 or tuple(mask.shape) != (latent_height, latent_width):
|
||||
raise ValueError(
|
||||
"Semantic ownership mask "
|
||||
f"{index} must match latent shape {latent_height}x{latent_width}."
|
||||
)
|
||||
boolean_mask = mask.detach().cpu().bool()
|
||||
if not bool(boolean_mask.any()):
|
||||
continue
|
||||
if bool(torch.logical_and(coverage, boolean_mask).any()):
|
||||
raise ValueError("Semantic ownership masks must not overlap.")
|
||||
normalized.append(boolean_mask)
|
||||
coverage = torch.logical_or(coverage, boolean_mask)
|
||||
if not normalized:
|
||||
raise ValueError("Semantic tiled diffusion requires non-empty ownership.")
|
||||
if not bool(coverage.all()):
|
||||
raise ValueError("Semantic ownership masks must cover the latent canvas.")
|
||||
return tuple(normalized)
|
||||
|
||||
|
||||
def _split_core(
|
||||
core: _OwnershipCore,
|
||||
*,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> tuple[_OwnershipCore, ...]:
|
||||
"""Recursively divide a core into balanced pieces within its tile budget."""
|
||||
|
||||
left, top, right, bottom = core.bounds
|
||||
width = right - left
|
||||
height = bottom - top
|
||||
if width <= max_width and height <= max_height:
|
||||
return (core,)
|
||||
split_x = width / max_width >= height / max_height
|
||||
first, second = _split_mask_at_balanced_axis(core.mask, core.bounds, split_x)
|
||||
return _split_core(
|
||||
_core_from_mask(first), max_width=max_width, max_height=max_height
|
||||
) + _split_core(_core_from_mask(second), max_width=max_width, max_height=max_height)
|
||||
|
||||
|
||||
def _split_mask_at_balanced_axis(
|
||||
mask: torch.Tensor,
|
||||
bounds: tuple[int, int, int, int],
|
||||
split_x: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Split one non-empty mask near its active-pixel median on one axis."""
|
||||
|
||||
left, top, right, bottom = bounds
|
||||
counts = (
|
||||
mask[top:bottom, left:right].sum(dim=0)
|
||||
if split_x
|
||||
else mask[top:bottom, left:right].sum(dim=1)
|
||||
)
|
||||
cumulative = torch.cumsum(counts, dim=0)
|
||||
midpoint = int(torch.searchsorted(cumulative, cumulative[-1] / 2, right=False))
|
||||
axis_start = left if split_x else top
|
||||
axis_end = right if split_x else bottom
|
||||
split_at = min(axis_end - 1, max(axis_start + 1, axis_start + midpoint + 1))
|
||||
first = mask.clone()
|
||||
second = mask.clone()
|
||||
if split_x:
|
||||
first[:, split_at:] = False
|
||||
second[:, :split_at] = False
|
||||
else:
|
||||
first[split_at:, :] = False
|
||||
second[:split_at, :] = False
|
||||
if not bool(first.any()) or not bool(second.any()):
|
||||
raise ValueError("Unable to split an oversized semantic tile core.")
|
||||
return first, second
|
||||
|
||||
|
||||
def _merge_small_cores(
|
||||
cores: tuple[_OwnershipCore, ...],
|
||||
*,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> tuple[_OwnershipCore, ...]:
|
||||
"""Greedily combine nearby cores when one bounded window can hold both."""
|
||||
|
||||
pending = list(cores)
|
||||
minimum_area = max(1, (max_width * max_height) // 4)
|
||||
merged = True
|
||||
while merged:
|
||||
merged = False
|
||||
for index, core in enumerate(tuple(pending)):
|
||||
if core.area >= minimum_area:
|
||||
continue
|
||||
candidate_index = _best_merge_candidate_index(
|
||||
core,
|
||||
pending,
|
||||
excluded_index=index,
|
||||
max_width=max_width,
|
||||
max_height=max_height,
|
||||
)
|
||||
if candidate_index is None:
|
||||
continue
|
||||
candidate = pending[candidate_index]
|
||||
pending[index] = _OwnershipCore(
|
||||
mask=torch.logical_or(core.mask, candidate.mask),
|
||||
bounds=_union_bounds(core.bounds, candidate.bounds),
|
||||
area=core.area + candidate.area,
|
||||
)
|
||||
pending.pop(candidate_index)
|
||||
merged = True
|
||||
break
|
||||
return tuple(pending)
|
||||
|
||||
|
||||
def _best_merge_candidate_index(
|
||||
core: _OwnershipCore,
|
||||
candidates: list[_OwnershipCore],
|
||||
*,
|
||||
excluded_index: int,
|
||||
max_width: int,
|
||||
max_height: int,
|
||||
) -> int | None:
|
||||
"""Return a candidate whose combined bounds fit one ownership budget."""
|
||||
|
||||
eligible: list[tuple[int, int, int]] = []
|
||||
for index, candidate in enumerate(candidates):
|
||||
if index == excluded_index:
|
||||
continue
|
||||
bounds = _union_bounds(core.bounds, candidate.bounds)
|
||||
left, top, right, bottom = bounds
|
||||
width = right - left
|
||||
height = bottom - top
|
||||
if width > max_width or height > max_height:
|
||||
continue
|
||||
distance = _bounds_distance(core.bounds, candidate.bounds)
|
||||
eligible.append((width * height, distance, index))
|
||||
if not eligible:
|
||||
return None
|
||||
return min(eligible, key=lambda item: (item[0], item[1]))[2]
|
||||
|
||||
|
||||
def _bounds_distance(
|
||||
first: tuple[int, int, int, int],
|
||||
second: tuple[int, int, int, int],
|
||||
) -> int:
|
||||
"""Return the axis-aligned gap between two mask bounding boxes."""
|
||||
|
||||
left, top, right, bottom = first
|
||||
other_left, other_top, other_right, other_bottom = second
|
||||
horizontal = max(0, other_left - right, left - other_right)
|
||||
vertical = max(0, other_top - bottom, top - other_bottom)
|
||||
return horizontal + vertical
|
||||
|
||||
|
||||
def _union_bounds(
|
||||
first: tuple[int, int, int, int],
|
||||
second: tuple[int, int, int, int],
|
||||
) -> tuple[int, int, int, int]:
|
||||
"""Return the tight rectangle containing both ownership-core bounds."""
|
||||
|
||||
return (
|
||||
min(first[0], second[0]),
|
||||
min(first[1], second[1]),
|
||||
max(first[2], second[2]),
|
||||
max(first[3], second[3]),
|
||||
)
|
||||
|
||||
|
||||
def _tile_for_core(
|
||||
core: _OwnershipCore,
|
||||
*,
|
||||
latent_width: int,
|
||||
latent_height: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
) -> LatentTile:
|
||||
"""Place one bounded sampling window around an irregular ownership core."""
|
||||
|
||||
left, top, right, bottom = core.bounds
|
||||
center_x = (left + right) / 2.0
|
||||
center_y = (top + bottom) / 2.0
|
||||
x = _clamp_window_start(center_x, tile_width, latent_width)
|
||||
y = _clamp_window_start(center_y, tile_height, latent_height)
|
||||
weight_mask = _feathered_tile_weight(
|
||||
core.mask,
|
||||
x=x,
|
||||
y=y,
|
||||
width=tile_width,
|
||||
height=tile_height,
|
||||
overlap=overlap,
|
||||
)
|
||||
if not bool((weight_mask > 0).any()):
|
||||
raise ValueError("Semantic tiled diffusion generated an empty tile weight.")
|
||||
return LatentTile(x, y, tile_width, tile_height, weight_mask)
|
||||
|
||||
|
||||
def _clamp_window_start(center: float, window_size: int, limit: int) -> int:
|
||||
"""Center a fixed sampling window while keeping it inside latent bounds."""
|
||||
|
||||
desired = round(center - window_size / 2.0)
|
||||
return min(max(0, desired), limit - window_size)
|
||||
|
||||
|
||||
def _feathered_tile_weight(
|
||||
mask: torch.Tensor,
|
||||
*,
|
||||
x: int,
|
||||
y: int,
|
||||
width: int,
|
||||
height: int,
|
||||
overlap: int,
|
||||
) -> torch.Tensor:
|
||||
"""Build one feathered tile weight without blurring the full latent mask."""
|
||||
|
||||
if overlap == 0:
|
||||
return mask[y : y + height, x : x + width].float().contiguous()
|
||||
radius = max(1, overlap // 2)
|
||||
source_left = max(0, x - radius)
|
||||
source_top = max(0, y - radius)
|
||||
source_right = min(int(mask.shape[1]), x + width + radius)
|
||||
source_bottom = min(int(mask.shape[0]), y + height + radius)
|
||||
local_weight = (
|
||||
functional.avg_pool2d(
|
||||
mask[source_top:source_bottom, source_left:source_right]
|
||||
.float()
|
||||
.unsqueeze(0)
|
||||
.unsqueeze(0),
|
||||
kernel_size=radius * 2 + 1,
|
||||
stride=1,
|
||||
padding=radius,
|
||||
count_include_pad=False,
|
||||
)
|
||||
.squeeze(0)
|
||||
.squeeze(0)
|
||||
)
|
||||
local_y = y - source_top
|
||||
local_x = x - source_left
|
||||
return local_weight[
|
||||
local_y : local_y + height, local_x : local_x + width
|
||||
].contiguous()
|
||||
|
||||
|
||||
def _core_from_mask(mask: torch.Tensor) -> _OwnershipCore:
|
||||
"""Build one core with bounds and area computed exactly once."""
|
||||
|
||||
bounds = _mask_bounds(mask)
|
||||
if bounds is None:
|
||||
raise ValueError("Semantic tiled diffusion cannot use an empty core.")
|
||||
return _OwnershipCore(mask=mask, bounds=bounds, area=int(mask.sum().item()))
|
||||
|
||||
|
||||
def _mask_bounds(mask: torch.Tensor) -> tuple[int, int, int, int] | None:
|
||||
"""Return left, top, right, bottom bounds for a non-empty boolean mask."""
|
||||
|
||||
y_coords, x_coords = torch.where(mask)
|
||||
if y_coords.numel() == 0:
|
||||
return None
|
||||
return (
|
||||
int(x_coords.min().item()),
|
||||
int(y_coords.min().item()),
|
||||
int(x_coords.max().item()) + 1,
|
||||
int(y_coords.max().item()) + 1,
|
||||
)
|
||||
@@ -0,0 +1,164 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define immutable spatial model views and view-major batch layouts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class SpatialViewKind(StrEnum):
|
||||
"""Identify how one model view relates to the canonical latent canvas."""
|
||||
|
||||
FULL = "full"
|
||||
TILE = "tile"
|
||||
CONTEXTUAL_GLOBAL = "contextual_global"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpatialView:
|
||||
"""Describe one source rectangle evaluated at one model spatial shape."""
|
||||
|
||||
kind: SpatialViewKind
|
||||
source_x: int
|
||||
source_y: int
|
||||
source_width: int
|
||||
source_height: int
|
||||
model_width: int
|
||||
model_height: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject invalid or internally inconsistent view geometry."""
|
||||
|
||||
if not isinstance(self.kind, SpatialViewKind):
|
||||
raise TypeError("Spatial view kind must be a SpatialViewKind value.")
|
||||
if self.source_x < 0 or self.source_y < 0:
|
||||
raise ValueError("Spatial view source coordinates must be non-negative.")
|
||||
if self.source_width < 1 or self.source_height < 1:
|
||||
raise ValueError("Spatial view source dimensions must be positive.")
|
||||
if self.model_width < 1 or self.model_height < 1:
|
||||
raise ValueError("Spatial view model dimensions must be positive.")
|
||||
if self.kind is SpatialViewKind.FULL and (
|
||||
self.source_width != self.model_width
|
||||
or self.source_height != self.model_height
|
||||
):
|
||||
raise ValueError("A full spatial view must preserve its source dimensions.")
|
||||
|
||||
@property
|
||||
def source_right(self) -> int:
|
||||
"""Return the exclusive source rectangle right edge."""
|
||||
|
||||
return self.source_x + self.source_width
|
||||
|
||||
@property
|
||||
def source_bottom(self) -> int:
|
||||
"""Return the exclusive source rectangle bottom edge."""
|
||||
|
||||
return self.source_y + self.source_height
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpatialBatchLayout:
|
||||
"""Describe ordered spatial views expanded over one source model batch."""
|
||||
|
||||
canvas_width: int
|
||||
canvas_height: int
|
||||
views: tuple[SpatialView, ...]
|
||||
input_batch_size: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate canvas containment and homogeneous model-call semantics."""
|
||||
|
||||
if self.canvas_width < 1 or self.canvas_height < 1:
|
||||
raise ValueError("Spatial layout canvas dimensions must be positive.")
|
||||
if self.input_batch_size < 1:
|
||||
raise ValueError("Spatial layout input batch size must be positive.")
|
||||
if not isinstance(self.views, tuple):
|
||||
raise TypeError("Spatial layout views must be an immutable tuple.")
|
||||
if not self.views:
|
||||
raise ValueError("Spatial layout requires at least one view.")
|
||||
if not all(isinstance(view, SpatialView) for view in self.views):
|
||||
raise TypeError("Spatial layout views must contain SpatialView values.")
|
||||
|
||||
view_kind = self.views[0].kind
|
||||
if any(view.kind is not view_kind for view in self.views):
|
||||
raise ValueError("One spatial model call cannot mix view kinds.")
|
||||
for view in self.views:
|
||||
if (
|
||||
view.source_right > self.canvas_width
|
||||
or view.source_bottom > self.canvas_height
|
||||
):
|
||||
raise ValueError(
|
||||
"Spatial view source rectangle must remain inside the canvas."
|
||||
)
|
||||
|
||||
if view_kind in {
|
||||
SpatialViewKind.FULL,
|
||||
SpatialViewKind.CONTEXTUAL_GLOBAL,
|
||||
}:
|
||||
if len(self.views) != 1:
|
||||
raise ValueError("A full-source spatial layout requires one view.")
|
||||
view = self.views[0]
|
||||
if (
|
||||
view.source_x != 0
|
||||
or view.source_y != 0
|
||||
or view.source_width != self.canvas_width
|
||||
or view.source_height != self.canvas_height
|
||||
):
|
||||
raise ValueError(
|
||||
"A full-source spatial view must cover the complete canvas."
|
||||
)
|
||||
|
||||
@property
|
||||
def view_count(self) -> int:
|
||||
"""Return the number of ordered spatial views."""
|
||||
|
||||
return len(self.views)
|
||||
|
||||
@property
|
||||
def expanded_batch_size(self) -> int:
|
||||
"""Return the model batch size after view-major expansion."""
|
||||
|
||||
return self.view_count * self.input_batch_size
|
||||
|
||||
@property
|
||||
def expanded_views(self) -> tuple[SpatialView, ...]:
|
||||
"""Repeat each view for its contiguous source-batch group."""
|
||||
|
||||
return tuple(
|
||||
view
|
||||
for view in self.views
|
||||
for _source_batch_index in range(self.input_batch_size)
|
||||
)
|
||||
|
||||
@property
|
||||
def expanded_view_indices(self) -> tuple[int, ...]:
|
||||
"""Return the view index for every expanded model-batch entry."""
|
||||
|
||||
return tuple(
|
||||
view_index
|
||||
for view_index in range(self.view_count)
|
||||
for _source_batch_index in range(self.input_batch_size)
|
||||
)
|
||||
|
||||
@property
|
||||
def expanded_source_batch_indices(self) -> tuple[int, ...]:
|
||||
"""Return the source-batch index for every expanded model-batch entry."""
|
||||
|
||||
return tuple(
|
||||
source_batch_index
|
||||
for _view in self.views
|
||||
for source_batch_index in range(self.input_batch_size)
|
||||
)
|
||||
|
||||
def expanded_index(self, view_index: int, source_batch_index: int) -> int:
|
||||
"""Return one view-major model-batch index after validating both axes."""
|
||||
|
||||
if not 0 <= view_index < self.view_count:
|
||||
raise IndexError("Spatial view index is outside the layout.")
|
||||
if not 0 <= source_batch_index < self.input_batch_size:
|
||||
raise IndexError("Source batch index is outside the layout.")
|
||||
return view_index * self.input_batch_size + source_batch_index
|
||||
@@ -2,18 +2,19 @@
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Connected-component helpers for binary mask regions."""
|
||||
"""Extract deterministic connected components from binary mask regions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from importlib import import_module
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.segs import BoundingBox
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MaskComponent:
|
||||
"""Represent one connected mask component and its full-image bbox."""
|
||||
|
||||
@@ -22,62 +23,30 @@ class MaskComponent:
|
||||
|
||||
|
||||
def connected_mask_components(active_mask: torch.Tensor) -> tuple[MaskComponent, ...]:
|
||||
"""Return 8-connected components from an HW active-pixel mask."""
|
||||
"""Return spatially ordered 8-connected components from one HW mask."""
|
||||
|
||||
if active_mask.ndim != 2:
|
||||
if not isinstance(active_mask, torch.Tensor) or active_mask.ndim != 2:
|
||||
raise ValueError("active_mask must be an HW tensor.")
|
||||
|
||||
active = active_mask.detach().to(device="cpu", dtype=torch.bool)
|
||||
height = int(active.shape[0])
|
||||
width = int(active.shape[1])
|
||||
visited = torch.zeros((height, width), dtype=torch.bool)
|
||||
active = active_mask.detach().to(device="cpu", dtype=torch.uint8).numpy()
|
||||
if not active.any():
|
||||
return ()
|
||||
cv2 = import_module("cv2")
|
||||
component_count, labels, stats, _centroids = cv2.connectedComponentsWithStats(
|
||||
active,
|
||||
connectivity=8,
|
||||
)
|
||||
components: list[MaskComponent] = []
|
||||
|
||||
for top in range(height):
|
||||
for left in range(width):
|
||||
if visited[top, left].item() or not active[top, left].item():
|
||||
continue
|
||||
components.append(_trace_component(active, visited, left, top))
|
||||
|
||||
return tuple(components)
|
||||
|
||||
|
||||
def _trace_component(
|
||||
active: torch.Tensor,
|
||||
visited: torch.Tensor,
|
||||
start_left: int,
|
||||
start_top: int,
|
||||
) -> MaskComponent:
|
||||
"""Trace one 8-connected component from its first active pixel."""
|
||||
|
||||
height = int(active.shape[0])
|
||||
width = int(active.shape[1])
|
||||
queue: list[tuple[int, int]] = [(start_top, start_left)]
|
||||
visited[start_top, start_left] = True
|
||||
pixels: list[tuple[int, int]] = []
|
||||
index = 0
|
||||
|
||||
while index < len(queue):
|
||||
top, left = queue[index]
|
||||
index += 1
|
||||
pixels.append((top, left))
|
||||
for neighbor_top in range(max(0, top - 1), min(height, top + 2)):
|
||||
for neighbor_left in range(max(0, left - 1), min(width, left + 2)):
|
||||
if visited[neighbor_top, neighbor_left].item():
|
||||
continue
|
||||
visited[neighbor_top, neighbor_left] = True
|
||||
if active[neighbor_top, neighbor_left].item():
|
||||
queue.append((neighbor_top, neighbor_left))
|
||||
|
||||
top_values = [top for top, _left in pixels]
|
||||
left_values = [left for _top, left in pixels]
|
||||
bbox = BoundingBox(
|
||||
min(left_values),
|
||||
min(top_values),
|
||||
max(left_values) + 1,
|
||||
max(top_values) + 1,
|
||||
for label in range(1, int(component_count)):
|
||||
left = int(stats[label, cv2.CC_STAT_LEFT])
|
||||
top = int(stats[label, cv2.CC_STAT_TOP])
|
||||
width = int(stats[label, cv2.CC_STAT_WIDTH])
|
||||
height = int(stats[label, cv2.CC_STAT_HEIGHT])
|
||||
components.append(
|
||||
MaskComponent(
|
||||
bbox=BoundingBox(left, top, left + width, top + height),
|
||||
mask=torch.from_numpy(labels == label),
|
||||
)
|
||||
)
|
||||
return tuple(
|
||||
sorted(components, key=lambda value: (value.bbox.top, value.bbox.left))
|
||||
)
|
||||
component_mask = torch.zeros((height, width), dtype=torch.bool)
|
||||
for top, left in pixels:
|
||||
component_mask[top, left] = True
|
||||
return MaskComponent(bbox=bbox, mask=component_mask)
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Project canonical masks into exact regional activation multiplier shapes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.regional_activation_geometry import (
|
||||
RegionalActivationGeometry,
|
||||
RegionalActivationLayout,
|
||||
)
|
||||
from ..domain.regional_mask_bank import RegionalMaskBank
|
||||
from ..domain.spatial_views import (
|
||||
SpatialBatchLayout,
|
||||
SpatialView,
|
||||
SpatialViewKind,
|
||||
)
|
||||
from .regional_mask_projection import (
|
||||
RegionalMaskForm,
|
||||
RegionalMaskProjectionMode,
|
||||
RegionalMaskProjector,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalActivationMaskBatch:
|
||||
"""Retain one finite region-major multiplier broadcast over rank activations."""
|
||||
|
||||
multipliers: torch.Tensor
|
||||
geometry: RegionalActivationGeometry
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require exact geometry shape and bounded finite floating multipliers."""
|
||||
|
||||
if not isinstance(self.multipliers, torch.Tensor):
|
||||
raise TypeError("Regional activation multipliers must be a tensor.")
|
||||
if not isinstance(self.geometry, RegionalActivationGeometry):
|
||||
raise TypeError("Regional activation multipliers require geometry.")
|
||||
if not self.multipliers.is_floating_point():
|
||||
raise TypeError("Regional activation multipliers must use floating point.")
|
||||
if self.multipliers.ndim != len(self.geometry.invocation_shape) + 1:
|
||||
raise ValueError("Regional activation multiplier rank is inconsistent.")
|
||||
expected = self.geometry.broadcast_mask_shape(int(self.multipliers.shape[0]))
|
||||
if tuple(self.multipliers.shape) != expected:
|
||||
raise ValueError(
|
||||
"Regional activation multiplier shape must match its geometry."
|
||||
)
|
||||
if not bool(torch.isfinite(self.multipliers).all()):
|
||||
raise ValueError("Regional activation multipliers must be finite.")
|
||||
if not bool(((self.multipliers >= 0.0) & (self.multipliers <= 1.0)).all()):
|
||||
raise ValueError("Regional activation multipliers must stay within [0, 1].")
|
||||
|
||||
|
||||
class RegionalActivationMaskProjector:
|
||||
"""Own geometry-shaped mask projection around existing crop/interpolation."""
|
||||
|
||||
def __init__(self, projector: RegionalMaskProjector | None = None) -> None:
|
||||
"""Retain the canonical mask crop and interpolation authority."""
|
||||
|
||||
self._projector = projector or RegionalMaskProjector()
|
||||
|
||||
def project(
|
||||
self,
|
||||
*,
|
||||
bank: RegionalMaskBank,
|
||||
geometry: RegionalActivationGeometry,
|
||||
form: RegionalMaskForm,
|
||||
mode: RegionalMaskProjectionMode,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> RegionalActivationMaskBatch:
|
||||
"""Return masks in region/view/chunk/latent activation order."""
|
||||
|
||||
_validate_inputs(bank, geometry, form, mode, device, dtype)
|
||||
layout = geometry.batch_alignment.spatial_layout or _full_layout(
|
||||
bank,
|
||||
input_batch_size=geometry.batch_alignment.base_batch_size,
|
||||
)
|
||||
projected_views = tuple(
|
||||
self._projector.project_query_grid(
|
||||
bank=bank,
|
||||
layout=layout,
|
||||
view_index=view_index,
|
||||
query_height=geometry.spatial_height,
|
||||
query_width=geometry.spatial_width,
|
||||
form=form,
|
||||
mode=mode,
|
||||
)
|
||||
for view_index in range(layout.view_count)
|
||||
)
|
||||
masks = torch.cat(
|
||||
tuple(
|
||||
projected.unsqueeze(1).expand(
|
||||
-1,
|
||||
layout.input_batch_size,
|
||||
-1,
|
||||
-1,
|
||||
)
|
||||
for projected in projected_views
|
||||
),
|
||||
dim=1,
|
||||
).to(device=device, dtype=dtype)
|
||||
multipliers = _reshape_for_activation(masks, geometry)
|
||||
return RegionalActivationMaskBatch(multipliers, geometry)
|
||||
|
||||
|
||||
def _reshape_for_activation(
|
||||
masks: torch.Tensor,
|
||||
geometry: RegionalActivationGeometry,
|
||||
) -> torch.Tensor:
|
||||
"""Place projected H/W masks on the declared operation's non-feature axes."""
|
||||
|
||||
regions, batch, height, width = (int(value) for value in masks.shape)
|
||||
if geometry.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_1D:
|
||||
if height != 1:
|
||||
raise ValueError("Conv1d regional masks require projected height one.")
|
||||
return masks.reshape(regions, batch, 1, width)
|
||||
if geometry.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_2D:
|
||||
return masks.reshape(regions, batch, 1, height, width)
|
||||
if geometry.layout is RegionalActivationLayout.DIRECT_CONVOLUTION_3D:
|
||||
temporal_size = geometry.temporal_size
|
||||
if temporal_size is None:
|
||||
raise ValueError("Conv3d regional masks require explicit temporal size.")
|
||||
return masks.reshape(regions, batch, 1, 1, height, width).expand(
|
||||
-1,
|
||||
-1,
|
||||
-1,
|
||||
temporal_size,
|
||||
-1,
|
||||
-1,
|
||||
)
|
||||
if geometry.layout in (
|
||||
RegionalActivationLayout.FLATTENED_SPATIAL_TOKENS,
|
||||
RegionalActivationLayout.CONSUMER_SPATIALIZED,
|
||||
):
|
||||
return masks.flatten(start_dim=2).unsqueeze(-1)
|
||||
raise AssertionError(f"Unhandled regional activation layout: {geometry.layout}")
|
||||
|
||||
|
||||
def _full_layout(
|
||||
bank: RegionalMaskBank,
|
||||
*,
|
||||
input_batch_size: int,
|
||||
) -> SpatialBatchLayout:
|
||||
"""Represent one full-canvas invocation through the shared layout contract."""
|
||||
|
||||
return SpatialBatchLayout(
|
||||
bank.canvas_width,
|
||||
bank.canvas_height,
|
||||
(
|
||||
SpatialView(
|
||||
SpatialViewKind.FULL,
|
||||
0,
|
||||
0,
|
||||
bank.canvas_width,
|
||||
bank.canvas_height,
|
||||
bank.canvas_width,
|
||||
bank.canvas_height,
|
||||
),
|
||||
),
|
||||
input_batch_size,
|
||||
)
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
bank: object,
|
||||
geometry: object,
|
||||
form: object,
|
||||
mode: object,
|
||||
device: object,
|
||||
dtype: object,
|
||||
) -> None:
|
||||
"""Validate typed projection inputs without moving or mutating source masks."""
|
||||
|
||||
if not isinstance(bank, RegionalMaskBank):
|
||||
raise TypeError("Regional activation projection requires a mask bank.")
|
||||
if not isinstance(geometry, RegionalActivationGeometry):
|
||||
raise TypeError("Regional activation projection requires geometry.")
|
||||
if not isinstance(form, RegionalMaskForm):
|
||||
raise TypeError("Regional activation mask form has an invalid type.")
|
||||
if not isinstance(mode, RegionalMaskProjectionMode):
|
||||
raise TypeError("Regional activation projection mode has an invalid type.")
|
||||
if not isinstance(device, torch.device):
|
||||
raise TypeError("Regional activation mask device must be a torch.device.")
|
||||
if not isinstance(dtype, torch.dtype) or not dtype.is_floating_point:
|
||||
raise TypeError("Regional activation mask dtype must be floating point.")
|
||||
|
||||
|
||||
REGIONAL_ACTIVATION_MASK_PROJECTOR = RegionalActivationMaskProjector()
|
||||
@@ -0,0 +1,125 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Classify active regional coverage on one projected query grid."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class RegionalCoverageClass(StrEnum):
|
||||
"""Identify model branches required by one regional mask projection."""
|
||||
|
||||
ALL_BASE = "all_base"
|
||||
ALL_SINGLE_REGION = "all_single_region"
|
||||
MIXED = "mixed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RegionalMaskActivation:
|
||||
"""Describe ordered active regions and an admitted coverage fast path."""
|
||||
|
||||
coverage_class: RegionalCoverageClass
|
||||
active_region_indices: tuple[int, ...]
|
||||
single_region_index: int | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require one internally consistent immutable classification."""
|
||||
|
||||
if not isinstance(self.coverage_class, RegionalCoverageClass):
|
||||
raise TypeError(
|
||||
"Regional coverage class must be a RegionalCoverageClass value."
|
||||
)
|
||||
if not isinstance(self.active_region_indices, tuple):
|
||||
raise TypeError("Active regional indices must be an immutable tuple.")
|
||||
if any(index < 0 for index in self.active_region_indices):
|
||||
raise ValueError("Active regional indices must be non-negative.")
|
||||
if tuple(sorted(set(self.active_region_indices))) != self.active_region_indices:
|
||||
raise ValueError("Active regional indices must be unique and ordered.")
|
||||
if self.coverage_class is RegionalCoverageClass.ALL_BASE:
|
||||
if self.active_region_indices or self.single_region_index is not None:
|
||||
raise ValueError("All-base coverage cannot contain an active region.")
|
||||
return
|
||||
if self.coverage_class is RegionalCoverageClass.ALL_SINGLE_REGION:
|
||||
if (
|
||||
len(self.active_region_indices) != 1
|
||||
or self.single_region_index != self.active_region_indices[0]
|
||||
):
|
||||
raise ValueError(
|
||||
"All-single-region coverage requires its one active region index."
|
||||
)
|
||||
return
|
||||
if not self.active_region_indices:
|
||||
raise ValueError("Mixed regional coverage requires an active region.")
|
||||
if self.single_region_index is not None:
|
||||
raise ValueError(
|
||||
"Mixed regional coverage cannot select a fast-path region."
|
||||
)
|
||||
|
||||
|
||||
class RegionalMaskActivationClassifier:
|
||||
"""Own zero-coverage pruning and regional fast-path admission."""
|
||||
|
||||
def classify(
|
||||
self,
|
||||
masks: torch.Tensor,
|
||||
*,
|
||||
tolerance: float = 1e-6,
|
||||
) -> RegionalMaskActivation:
|
||||
"""Return ordered active regions and the exact coverage classification."""
|
||||
|
||||
self._validate_masks(masks)
|
||||
if not 0.0 <= tolerance < 0.5:
|
||||
raise ValueError(
|
||||
"Regional activation tolerance must be at least 0 and below 0.5."
|
||||
)
|
||||
active_region_indices = tuple(
|
||||
index
|
||||
for index in range(int(masks.shape[0]))
|
||||
if bool((masks[index] > tolerance).any())
|
||||
)
|
||||
if not active_region_indices:
|
||||
return RegionalMaskActivation(
|
||||
coverage_class=RegionalCoverageClass.ALL_BASE,
|
||||
active_region_indices=(),
|
||||
single_region_index=None,
|
||||
)
|
||||
if len(active_region_indices) == 1:
|
||||
region_index = active_region_indices[0]
|
||||
if bool((masks[region_index] >= 1.0 - tolerance).all()):
|
||||
return RegionalMaskActivation(
|
||||
coverage_class=RegionalCoverageClass.ALL_SINGLE_REGION,
|
||||
active_region_indices=active_region_indices,
|
||||
single_region_index=region_index,
|
||||
)
|
||||
return RegionalMaskActivation(
|
||||
coverage_class=RegionalCoverageClass.MIXED,
|
||||
active_region_indices=active_region_indices,
|
||||
single_region_index=None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_masks(masks: torch.Tensor) -> None:
|
||||
"""Validate a projected normalized floating-point BHW mask batch."""
|
||||
|
||||
if not isinstance(masks, torch.Tensor):
|
||||
raise TypeError("Regional activation masks must be a torch.Tensor.")
|
||||
if masks.ndim != 3:
|
||||
raise ValueError("Regional activation masks must use BHW layout.")
|
||||
if int(masks.shape[0]) < 1:
|
||||
raise ValueError("Regional activation requires at least one region.")
|
||||
if int(masks.shape[1]) < 1 or int(masks.shape[2]) < 1:
|
||||
raise ValueError("Regional activation mask grids must be non-empty.")
|
||||
if not masks.is_floating_point():
|
||||
raise TypeError(
|
||||
"Regional activation masks must use a floating-point dtype."
|
||||
)
|
||||
if not bool(torch.isfinite(masks).all()):
|
||||
raise ValueError("Regional activation masks must contain finite values.")
|
||||
if not bool(((masks >= 0.0) & (masks <= 1.0)).all()):
|
||||
raise ValueError("Regional activation masks must stay within [0, 1].")
|
||||
@@ -0,0 +1,176 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Project canonical regional masks into spatial views and query grids."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from ..domain.regional_mask_bank import RegionalMaskBank
|
||||
from ..domain.spatial_views import SpatialBatchLayout, SpatialView
|
||||
|
||||
|
||||
class RegionalMaskForm(StrEnum):
|
||||
"""Select one canonical regional mask representation."""
|
||||
|
||||
PLANNING = "planning"
|
||||
CONDITIONING = "conditioning"
|
||||
|
||||
|
||||
class RegionalMaskProjectionMode(StrEnum):
|
||||
"""Select the interpolation semantics for one mask projection."""
|
||||
|
||||
CONTINUOUS_COVERAGE = "continuous_coverage"
|
||||
SOFT = "soft"
|
||||
HARD_PRESERVING = "hard_preserving"
|
||||
NEAREST = "nearest"
|
||||
|
||||
|
||||
class RegionalMaskProjector:
|
||||
"""Own canonical-canvas crop and interpolation policy for regional masks."""
|
||||
|
||||
def project_view(
|
||||
self,
|
||||
*,
|
||||
bank: RegionalMaskBank,
|
||||
layout: SpatialBatchLayout,
|
||||
view_index: int,
|
||||
form: RegionalMaskForm,
|
||||
mode: RegionalMaskProjectionMode,
|
||||
) -> torch.Tensor:
|
||||
"""Project one canonical mask form to a view's model dimensions."""
|
||||
|
||||
view = self._view(bank=bank, layout=layout, view_index=view_index)
|
||||
return self.project_query_grid(
|
||||
bank=bank,
|
||||
layout=layout,
|
||||
view_index=view_index,
|
||||
query_height=view.model_height,
|
||||
query_width=view.model_width,
|
||||
form=form,
|
||||
mode=mode,
|
||||
)
|
||||
|
||||
def project_query_grid(
|
||||
self,
|
||||
*,
|
||||
bank: RegionalMaskBank,
|
||||
layout: SpatialBatchLayout,
|
||||
view_index: int,
|
||||
query_height: int,
|
||||
query_width: int,
|
||||
form: RegionalMaskForm,
|
||||
mode: RegionalMaskProjectionMode,
|
||||
) -> torch.Tensor:
|
||||
"""Crop one spatial view and resample it directly to an attention grid."""
|
||||
|
||||
if query_height < 1 or query_width < 1:
|
||||
raise ValueError("Regional mask query-grid dimensions must be positive.")
|
||||
if not isinstance(form, RegionalMaskForm):
|
||||
raise TypeError("Regional mask form must be a RegionalMaskForm value.")
|
||||
if not isinstance(mode, RegionalMaskProjectionMode):
|
||||
raise TypeError(
|
||||
"Regional mask projection mode must be a "
|
||||
"RegionalMaskProjectionMode value."
|
||||
)
|
||||
view = self._view(bank=bank, layout=layout, view_index=view_index)
|
||||
masks = (
|
||||
bank.planning_masks
|
||||
if form is RegionalMaskForm.PLANNING
|
||||
else bank.conditioning_masks
|
||||
)
|
||||
cropped = masks[
|
||||
:,
|
||||
view.source_y : view.source_bottom,
|
||||
view.source_x : view.source_right,
|
||||
]
|
||||
projected = self._resize(
|
||||
cropped,
|
||||
height=query_height,
|
||||
width=query_width,
|
||||
mode=mode,
|
||||
).clamp(0.0, 1.0)
|
||||
if not bool(torch.isfinite(projected).all()):
|
||||
raise ValueError("Projected regional masks contain non-finite values.")
|
||||
return projected
|
||||
|
||||
@staticmethod
|
||||
def _view(
|
||||
*,
|
||||
bank: RegionalMaskBank,
|
||||
layout: SpatialBatchLayout,
|
||||
view_index: int,
|
||||
) -> SpatialView:
|
||||
"""Validate canonical canvas identity and return one indexed view."""
|
||||
|
||||
if (
|
||||
bank.canvas_width != layout.canvas_width
|
||||
or bank.canvas_height != layout.canvas_height
|
||||
):
|
||||
raise ValueError(
|
||||
"Regional mask bank canvas must match the spatial batch layout."
|
||||
)
|
||||
if view_index < 0 or view_index >= layout.view_count:
|
||||
raise IndexError("Regional mask spatial view index is outside the layout.")
|
||||
return layout.views[view_index]
|
||||
|
||||
@staticmethod
|
||||
def _resize(
|
||||
masks: torch.Tensor,
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
mode: RegionalMaskProjectionMode,
|
||||
) -> torch.Tensor:
|
||||
"""Apply explicit coverage, soft, or hard-preserving interpolation."""
|
||||
|
||||
source_height = int(masks.shape[-2])
|
||||
source_width = int(masks.shape[-1])
|
||||
if (source_height, source_width) == (height, width):
|
||||
return masks
|
||||
batched = masks.unsqueeze(1)
|
||||
if mode is RegionalMaskProjectionMode.HARD_PRESERVING:
|
||||
return functional.interpolate(
|
||||
batched,
|
||||
size=(height, width),
|
||||
mode="nearest-exact",
|
||||
).squeeze(1)
|
||||
if mode is RegionalMaskProjectionMode.NEAREST:
|
||||
return functional.interpolate(
|
||||
batched,
|
||||
size=(height, width),
|
||||
mode="nearest",
|
||||
).squeeze(1)
|
||||
if mode is RegionalMaskProjectionMode.SOFT:
|
||||
return functional.interpolate(
|
||||
batched,
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
).squeeze(1)
|
||||
|
||||
downsample_height = min(source_height, height)
|
||||
downsample_width = min(source_width, width)
|
||||
projected = batched
|
||||
if (downsample_height, downsample_width) != (
|
||||
source_height,
|
||||
source_width,
|
||||
):
|
||||
projected = functional.interpolate(
|
||||
projected,
|
||||
size=(downsample_height, downsample_width),
|
||||
mode="area",
|
||||
)
|
||||
if tuple(projected.shape[-2:]) != (height, width):
|
||||
projected = functional.interpolate(
|
||||
projected,
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
return projected.squeeze(1)
|
||||
@@ -7,13 +7,53 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
from ..domain.regional_mask_bank import RegionalMaskBank
|
||||
from .detailer_masks import gaussian_feather_mask
|
||||
|
||||
|
||||
def prepare_regional_mask_batch(mask: object, feather: int) -> torch.Tensor:
|
||||
"""Return a validated and optionally feathered BHW mask batch."""
|
||||
|
||||
_, conditioning = _prepare_regional_masks(mask, feather)
|
||||
return conditioning
|
||||
|
||||
|
||||
def build_regional_mask_bank(
|
||||
mask: object,
|
||||
*,
|
||||
feather: int,
|
||||
canvas_height: int,
|
||||
canvas_width: int,
|
||||
) -> RegionalMaskBank:
|
||||
"""Build separate planning and conditioning masks on one latent canvas."""
|
||||
|
||||
authored, feathered = _prepare_regional_masks(mask, feather)
|
||||
planning_masks = resize_regional_mask_batch(
|
||||
authored,
|
||||
height=canvas_height,
|
||||
width=canvas_width,
|
||||
)
|
||||
conditioning_masks = resize_regional_mask_batch(
|
||||
feathered,
|
||||
height=canvas_height,
|
||||
width=canvas_width,
|
||||
)
|
||||
return RegionalMaskBank(
|
||||
planning_masks=planning_masks,
|
||||
conditioning_masks=conditioning_masks,
|
||||
canvas_width=canvas_width,
|
||||
canvas_height=canvas_height,
|
||||
)
|
||||
|
||||
|
||||
def _prepare_regional_masks(
|
||||
mask: object,
|
||||
feather: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Return normalized authored masks and a separate feathered tensor."""
|
||||
|
||||
if not isinstance(mask, torch.Tensor):
|
||||
raise TypeError("regional prompting requires a torch MASK tensor.")
|
||||
if feather < 0:
|
||||
@@ -30,9 +70,44 @@ def prepare_regional_mask_batch(mask: object, feather: int) -> torch.Tensor:
|
||||
raise ValueError("regional masks must have non-empty height and width.")
|
||||
|
||||
normalized = working.clamp(0.0, 1.0)
|
||||
if feather == 0:
|
||||
return normalized
|
||||
return gaussian_feather_mask(normalized, feather)
|
||||
conditioning = (
|
||||
normalized.to(device=normalized.device, dtype=normalized.dtype, copy=True)
|
||||
if feather == 0
|
||||
else gaussian_feather_mask(normalized, feather)
|
||||
)
|
||||
return normalized, conditioning
|
||||
|
||||
|
||||
def resize_regional_mask_batch(
|
||||
mask_batch: torch.Tensor,
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
) -> torch.Tensor:
|
||||
"""Resize BHW masks while preserving authored area during downscaling."""
|
||||
|
||||
if height < 1 or width < 1:
|
||||
raise ValueError("regional mask target height and width must be positive.")
|
||||
if tuple(mask_batch.shape[1:]) == (height, width):
|
||||
return mask_batch
|
||||
source_height, source_width = map(int, mask_batch.shape[1:])
|
||||
downscaled_height = min(height, source_height)
|
||||
downscaled_width = min(width, source_width)
|
||||
working = mask_batch.unsqueeze(1)
|
||||
if (downscaled_height, downscaled_width) != (source_height, source_width):
|
||||
working = functional.interpolate(
|
||||
working,
|
||||
size=(downscaled_height, downscaled_width),
|
||||
mode="area",
|
||||
)
|
||||
if (downscaled_height, downscaled_width) != (height, width):
|
||||
working = functional.interpolate(
|
||||
working,
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
return working.squeeze(1)
|
||||
|
||||
|
||||
def regional_mask(mask_batch: torch.Tensor, index: int) -> torch.Tensor:
|
||||
|
||||
@@ -8,7 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..runtime.model_catalog import grounding_dino_choices, sam_choices
|
||||
from ..runtime.model_choices import ModelChoiceService, default_choice
|
||||
from ..runtime.model_metadata import GroundedSAMModelMetadata
|
||||
from . import tooltips
|
||||
|
||||
@@ -17,6 +17,7 @@ class GroundedSAMModelInfo:
|
||||
"""Expose selected grounded SAM source and local path metadata."""
|
||||
|
||||
_metadata = GroundedSAMModelMetadata()
|
||||
_choices = ModelChoiceService()
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("model_info",)
|
||||
@@ -31,19 +32,27 @@ class GroundedSAMModelInfo:
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare deterministic model metadata inputs."""
|
||||
|
||||
sam_model_choices = cls._choices.sam_choices()
|
||||
grounding_dino_model_choices = cls._choices.grounding_dino_choices()
|
||||
return {
|
||||
"required": {
|
||||
"sam_model": (
|
||||
sam_choices(),
|
||||
sam_model_choices,
|
||||
{
|
||||
"default": "sam_hq_vit_b (379MB)",
|
||||
"default": default_choice(
|
||||
sam_model_choices,
|
||||
"sam_hq_vit_b (379MB)",
|
||||
),
|
||||
"tooltip": tooltips.SAM_MODEL_INPUT,
|
||||
},
|
||||
),
|
||||
"grounding_dino_model": (
|
||||
grounding_dino_choices(),
|
||||
grounding_dino_model_choices,
|
||||
{
|
||||
"default": "GroundingDINO_SwinT_OGC (694MB)",
|
||||
"default": default_choice(
|
||||
grounding_dino_model_choices,
|
||||
"GroundingDINO_SwinT_OGC (694MB)",
|
||||
),
|
||||
"tooltip": tooltips.GROUNDING_DINO_MODEL_INPUT,
|
||||
},
|
||||
),
|
||||
@@ -53,4 +62,6 @@ class GroundedSAMModelInfo:
|
||||
def describe(self, sam_model: str, grounding_dino_model: str) -> tuple[str]:
|
||||
"""Return JSON metadata for selected model entries."""
|
||||
|
||||
self._choices.reject_sentinel(sam_model)
|
||||
self._choices.reject_sentinel(grounding_dino_model)
|
||||
return (self._metadata.describe_selection(sam_model, grounding_dino_model),)
|
||||
|
||||
@@ -1,232 +0,0 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for contextual diffusion sampling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, ClassVar, TypeAlias
|
||||
|
||||
from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
from ..services.contextual_diffusion_sampling_service import (
|
||||
ContextualDiffusionSamplingService,
|
||||
)
|
||||
from . import tooltips
|
||||
|
||||
Latent: TypeAlias = dict[str, Any]
|
||||
MAX_LATENT_CONTEXT_SIZE = 512
|
||||
|
||||
|
||||
class KSamplerContextualDiffusion:
|
||||
"""Edit large latents through coordinated global and detailed contexts."""
|
||||
|
||||
RETURN_TYPES = ("LATENT", "SEGS")
|
||||
RETURN_NAMES = ("latent", "contexts_segs")
|
||||
OUTPUT_TOOLTIPS = (
|
||||
tooltips.DENOISED_LATENT_OUTPUT,
|
||||
tooltips.CONTEXTUAL_DIFFUSION_CONTEXTS_OUTPUT,
|
||||
)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "SimpleSyrup/Sampling"
|
||||
DESCRIPTION = (
|
||||
"Preserves composition while applying appearance and subject-detail edits "
|
||||
"to large latents through global context and optional SEGS-guided tiles."
|
||||
)
|
||||
SEARCH_ALIASES = [
|
||||
"ksampler",
|
||||
"contextual diffusion",
|
||||
"contextual tiled diffusion",
|
||||
"high resolution edit",
|
||||
"sam tiled diffusion",
|
||||
]
|
||||
|
||||
service_class: ClassVar[type[ContextualDiffusionSamplingService]] = (
|
||||
ContextualDiffusionSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare KSampler inputs and bounded contextual controls."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", {"tooltip": tooltips.SAMPLING_MODEL}),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": True,
|
||||
"tooltip": tooltips.SAMPLING_SEED,
|
||||
},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 4,
|
||||
"min": 1,
|
||||
"max": 10000,
|
||||
"tooltip": tooltips.SAMPLING_STEPS,
|
||||
},
|
||||
),
|
||||
"cfg": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.1,
|
||||
"round": 0.01,
|
||||
"tooltip": tooltips.SAMPLING_CFG,
|
||||
},
|
||||
),
|
||||
"sampler_name": (
|
||||
sampling_samplers.available_samplers(),
|
||||
{"tooltip": tooltips.SAMPLER_NAME},
|
||||
),
|
||||
"scheduler": (
|
||||
sampling_schedulers.available_schedulers(),
|
||||
{"tooltip": tooltips.SCHEDULER},
|
||||
),
|
||||
"positive": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.POSITIVE_CONDITIONING},
|
||||
),
|
||||
"negative": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.NEGATIVE_CONDITIONING},
|
||||
),
|
||||
"latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}),
|
||||
"denoise": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01,
|
||||
"tooltip": tooltips.DENOISE_STRENGTH,
|
||||
},
|
||||
),
|
||||
"diffusion_mode": (
|
||||
list(TILED_DIFFUSION_MODES),
|
||||
{
|
||||
"default": "multidiffusion",
|
||||
"tooltip": tooltips.TILED_DIFFUSION_MODE,
|
||||
},
|
||||
),
|
||||
"latent_context_size": (
|
||||
"INT",
|
||||
{
|
||||
"default": 96,
|
||||
"min": 16,
|
||||
"max": MAX_LATENT_CONTEXT_SIZE,
|
||||
"step": 16,
|
||||
"tooltip": tooltips.LATENT_CONTEXT_SIZE,
|
||||
},
|
||||
),
|
||||
"latent_context_overlap": (
|
||||
"INT",
|
||||
{
|
||||
"default": 32,
|
||||
"min": 0,
|
||||
"max": 256,
|
||||
"step": 4,
|
||||
"tooltip": tooltips.LATENT_CONTEXT_OVERLAP,
|
||||
},
|
||||
),
|
||||
"latent_context_batch_size": (
|
||||
"INT",
|
||||
{
|
||||
"default": 4,
|
||||
"min": 1,
|
||||
"max": 8,
|
||||
"step": 1,
|
||||
"tooltip": tooltips.LATENT_CONTEXT_BATCH_SIZE,
|
||||
},
|
||||
),
|
||||
"global_weight": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.05,
|
||||
"tooltip": tooltips.GLOBAL_CONTEXT_WEIGHT,
|
||||
},
|
||||
),
|
||||
"global_steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 0,
|
||||
"max": 10000,
|
||||
"step": 1,
|
||||
"tooltip": tooltips.GLOBAL_CONTEXT_STEPS,
|
||||
},
|
||||
),
|
||||
"global_decay": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.5,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": tooltips.GLOBAL_CONTEXT_DECAY,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"segs": (
|
||||
"SEGS",
|
||||
{"tooltip": tooltips.CONTEXTUAL_DIFFUSION_SEGS},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
def sample(
|
||||
self,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_image: Latent,
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_context_size: int = 96,
|
||||
latent_context_overlap: int = 32,
|
||||
latent_context_batch_size: int = 4,
|
||||
global_weight: float = 1.0,
|
||||
global_steps: int = 1,
|
||||
global_decay: float = 0.5,
|
||||
segs: object | None = None,
|
||||
) -> tuple[Latent, object]:
|
||||
"""Delegate contextual diffusion sampling to its application service."""
|
||||
|
||||
result = self.service_class().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
diffusion_mode=diffusion_mode,
|
||||
latent_context_size=latent_context_size,
|
||||
latent_context_overlap=latent_context_overlap,
|
||||
latent_context_batch_size=latent_context_batch_size,
|
||||
global_weight=global_weight,
|
||||
global_steps=global_steps,
|
||||
global_decay=global_decay,
|
||||
segs=segs,
|
||||
)
|
||||
return result.latent, result.contexts
|
||||
@@ -1,206 +0,0 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for selectable tiled diffusion sampling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService
|
||||
from . import tooltips
|
||||
|
||||
Latent = dict[str, Any]
|
||||
MAX_LATENT_TILE_SIZE = 512
|
||||
|
||||
|
||||
class KSamplerTiledDiffusion:
|
||||
"""Sample latents with selectable tiled diffusion denoising."""
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.DENOISED_LATENT_OUTPUT,)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "SimpleSyrup/Sampling"
|
||||
DESCRIPTION = "Denoises latents with selectable tiled diffusion sampling."
|
||||
SEARCH_ALIASES = [
|
||||
"ksampler",
|
||||
"sampler",
|
||||
"tiled diffusion",
|
||||
"multidiffusion",
|
||||
"multi diffusion",
|
||||
"mixture of diffusers",
|
||||
]
|
||||
|
||||
service_class: ClassVar[type[TiledDiffusionSamplingService]] = (
|
||||
TiledDiffusionSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare ComfyUI inputs for selectable tiled diffusion sampling."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", {"tooltip": tooltips.SAMPLING_MODEL}),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": True,
|
||||
"tooltip": tooltips.SAMPLING_SEED,
|
||||
},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 20,
|
||||
"min": 1,
|
||||
"max": 10000,
|
||||
"tooltip": tooltips.SAMPLING_STEPS,
|
||||
},
|
||||
),
|
||||
"cfg": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 8.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.1,
|
||||
"round": 0.01,
|
||||
"tooltip": tooltips.SAMPLING_CFG,
|
||||
},
|
||||
),
|
||||
"sampler_name": (
|
||||
sampling_samplers.available_samplers(),
|
||||
{"tooltip": tooltips.SAMPLER_NAME},
|
||||
),
|
||||
"scheduler": (
|
||||
sampling_schedulers.available_schedulers(),
|
||||
{"tooltip": tooltips.SCHEDULER},
|
||||
),
|
||||
"positive": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.POSITIVE_CONDITIONING},
|
||||
),
|
||||
"negative": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.NEGATIVE_CONDITIONING},
|
||||
),
|
||||
"latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}),
|
||||
"denoise": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01,
|
||||
"tooltip": tooltips.DENOISE_STRENGTH,
|
||||
},
|
||||
),
|
||||
"diffusion_mode": (
|
||||
list(TILED_DIFFUSION_MODES),
|
||||
{
|
||||
"default": "multidiffusion",
|
||||
"tooltip": tooltips.TILED_DIFFUSION_MODE,
|
||||
},
|
||||
),
|
||||
"latent_tile_width": (
|
||||
"INT",
|
||||
{
|
||||
"default": 128,
|
||||
"min": 16,
|
||||
"max": MAX_LATENT_TILE_SIZE,
|
||||
"step": 16,
|
||||
"tooltip": tooltips.LATENT_TILE_WIDTH,
|
||||
},
|
||||
),
|
||||
"latent_tile_height": (
|
||||
"INT",
|
||||
{
|
||||
"default": 128,
|
||||
"min": 16,
|
||||
"max": MAX_LATENT_TILE_SIZE,
|
||||
"step": 16,
|
||||
"tooltip": tooltips.LATENT_TILE_HEIGHT,
|
||||
},
|
||||
),
|
||||
"latent_tile_overlap": (
|
||||
"INT",
|
||||
{
|
||||
"default": 16,
|
||||
"min": 0,
|
||||
"max": 256,
|
||||
"step": 4,
|
||||
"tooltip": tooltips.LATENT_TILE_OVERLAP,
|
||||
},
|
||||
),
|
||||
"latent_tile_batch_size": (
|
||||
"INT",
|
||||
{
|
||||
"default": 4,
|
||||
"min": 1,
|
||||
"max": 8,
|
||||
"step": 1,
|
||||
"tooltip": tooltips.LATENT_TILE_BATCH_SIZE,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"segs": (
|
||||
"SEGS",
|
||||
{
|
||||
"tooltip": (
|
||||
"Optional image regions that guide irregular tile "
|
||||
"boundaries while preserving the configured overlap."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def sample(
|
||||
self,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_image: Latent,
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_tile_width: int = 128,
|
||||
latent_tile_height: int = 128,
|
||||
latent_tile_overlap: int = 16,
|
||||
latent_tile_batch_size: int = 4,
|
||||
segs: object | None = None,
|
||||
) -> tuple[Latent]:
|
||||
"""Sample a latent with the selected tiled diffusion method."""
|
||||
|
||||
output = self.service_class().sample(
|
||||
diffusion_mode=diffusion_mode,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
latent_tile_width=latent_tile_width,
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=None,
|
||||
segs=segs,
|
||||
)
|
||||
return (output,)
|
||||
@@ -8,6 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from ..runtime.model_downloads import ComfyProgressReporter
|
||||
from ..runtime.ultralytics_loader import UltralyticsLoaderService
|
||||
|
||||
|
||||
@@ -40,7 +41,8 @@ class LoadUltralyticsModel:
|
||||
{
|
||||
"default": choices[0],
|
||||
"tooltip": (
|
||||
"Ultralytics model file in the ComfyUI models folder."
|
||||
"A local Ultralytics model or a curated model that "
|
||||
"downloads to ComfyUI's Impact Pack-compatible folders."
|
||||
),
|
||||
},
|
||||
)
|
||||
@@ -50,5 +52,8 @@ class LoadUltralyticsModel:
|
||||
def load(self, model_name: str) -> tuple[object, object, object]:
|
||||
"""Load the selected detector and paired compatibility facades."""
|
||||
|
||||
loaded = self.service_class().load(model_name)
|
||||
loaded = self.service_class().load(
|
||||
model_name,
|
||||
progress=ComfyProgressReporter(),
|
||||
)
|
||||
return loaded.detector_model, loaded.bbox_detector, loaded.segm_detector
|
||||
|
||||
@@ -10,14 +10,17 @@ import importlib
|
||||
from types import ModuleType
|
||||
from typing import Any
|
||||
|
||||
from ..runtime.anima_loader import (
|
||||
from ..domain.anima_quantization import AnimaQuantizationRecipe
|
||||
from ..runtime.diffusion_model_loader import DIFFUSION_WEIGHT_DTYPES
|
||||
from ..runtime.model_downloads import ComfyProgressReporter
|
||||
from ..runtime.quantization_capabilities import QuantizationCapabilityCatalog
|
||||
from ..runtime.quantization_progress import ComfyQuantizationProgressReporter
|
||||
from ..runtime.vae_loader import vae_choices
|
||||
from ..services.anima_loader_service import (
|
||||
AUTO_CHOICE,
|
||||
CLIP_DEVICES,
|
||||
DIFFUSION_WEIGHT_DTYPES,
|
||||
AnimaLoaderService,
|
||||
)
|
||||
from ..runtime.model_downloads import ComfyProgressReporter
|
||||
from ..runtime.vae_loader import vae_choices
|
||||
from . import tooltips
|
||||
|
||||
|
||||
@@ -25,6 +28,8 @@ class SimpleLoadAnima:
|
||||
"""Expose Anima diffusion, text encoder, and VAE loading as one node."""
|
||||
|
||||
_service = AnimaLoaderService()
|
||||
_quantization_recipe = AnimaQuantizationRecipe()
|
||||
_quantization_capabilities = QuantizationCapabilityCatalog()
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
|
||||
RETURN_NAMES = ("model", "clip", "vae")
|
||||
@@ -53,14 +58,28 @@ class SimpleLoadAnima:
|
||||
)
|
||||
},
|
||||
),
|
||||
"quantization": (
|
||||
cls._quantization_capabilities.selection_labels(
|
||||
cls._quantization_recipe.profiles
|
||||
),
|
||||
{
|
||||
"default": "Original",
|
||||
"advanced": True,
|
||||
"tooltip": (
|
||||
"Creates or reuses a GPU-supported quantized copy in the "
|
||||
"global models/SyrupQuants cache; Original loads the "
|
||||
"selected model unchanged."
|
||||
),
|
||||
},
|
||||
),
|
||||
"diffusion_weight_dtype": (
|
||||
list(DIFFUSION_WEIGHT_DTYPES),
|
||||
{
|
||||
"default": "default",
|
||||
"advanced": True,
|
||||
"tooltip": (
|
||||
"Weight precision for Anima. Lower precision can reduce "
|
||||
"memory use but may slightly change results."
|
||||
"Load-time weight precision used with Original; cached "
|
||||
"quantized copies use their stored quantization format."
|
||||
),
|
||||
},
|
||||
),
|
||||
@@ -103,6 +122,7 @@ class SimpleLoadAnima:
|
||||
def load_models(
|
||||
self,
|
||||
diffusion_model: str,
|
||||
quantization: str,
|
||||
diffusion_weight_dtype: str,
|
||||
text_encoder: str,
|
||||
text_encoder_device: str,
|
||||
@@ -112,11 +132,13 @@ class SimpleLoadAnima:
|
||||
|
||||
return self._service.load_models(
|
||||
diffusion_model=diffusion_model,
|
||||
quantization=quantization,
|
||||
diffusion_weight_dtype=diffusion_weight_dtype,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_device=text_encoder_device,
|
||||
vae=vae,
|
||||
progress=ComfyProgressReporter(),
|
||||
quantization_progress=ComfyQuantizationProgressReporter(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -115,6 +115,20 @@ SAMPLING_SEED = (
|
||||
"Seed used to create sampling noise. Reusing it with matching settings makes "
|
||||
"results repeatable."
|
||||
)
|
||||
SEED_VARIATION_MODEL_INPUT = (
|
||||
"Model whose sampler-provided initial noise will receive seed variation."
|
||||
)
|
||||
VARIATION_SEED = (
|
||||
"Seed for the second noise pattern. Change it to explore another variation "
|
||||
"direction from the sampler's seed."
|
||||
)
|
||||
VARIATION_STRENGTH = (
|
||||
"Noise interpolation strength. 0 keeps the sampler seed unchanged; 1 uses the "
|
||||
"variation seed for initial noise."
|
||||
)
|
||||
SEED_VARIATION_MODEL_OUTPUT = (
|
||||
"Model that interpolates initial sampling noise toward the variation seed."
|
||||
)
|
||||
SAMPLING_STEPS = (
|
||||
"Number of denoising steps. More steps can add refinement but take longer."
|
||||
)
|
||||
@@ -187,6 +201,18 @@ GLOBAL_CONTEXT_DECAY = (
|
||||
CONTEXTUAL_DIFFUSION_SEGS = (
|
||||
"Optional regions that replace the regular grid with SEGS-guided contexts."
|
||||
)
|
||||
OPTIONAL_REGIONAL_MASKS = (
|
||||
"Optional authored masks that activate regional prompting when positive or "
|
||||
"negative conditioning is a batch; SEGS may further subdivide those regions."
|
||||
)
|
||||
OPTIONAL_REGIONAL_PROMPT_WEIGHT = (
|
||||
"Balances regional prompts against the global prompt when regional masks and "
|
||||
"a conditioning batch are connected."
|
||||
)
|
||||
OPTIONAL_REGION_MASK_FEATHER = (
|
||||
"Softens regional conditioning edges by this many source-mask pixels; tile "
|
||||
"boundaries continue to follow the unfeathered authored regions."
|
||||
)
|
||||
|
||||
DETAIL_IMAGE = (
|
||||
"Source image containing the regions to improve. Detailed crops are blended "
|
||||
|
||||
@@ -12,12 +12,26 @@ from ..runtime.prompt_control_availability import prompt_control_is_available
|
||||
def get_nodes() -> list[type[object]]:
|
||||
"""Return v3 nodes that can be advertised in this environment."""
|
||||
|
||||
from .all_prompt_attention_segs import AllPromptAttentionSEGSV3
|
||||
from .attention_capture_model import AttentionCaptureModelV3
|
||||
from .attention_masked_conditioning import AttentionMaskedConditioningV3
|
||||
from .attention_region_mask import AttentionRegionMaskV3
|
||||
from .batch_region_conditioning import BatchRegionConditioningV3
|
||||
from .batch_segs import BatchSEGSV3
|
||||
from .compose_regional_conditioning import ComposeRegionalConditioningV3
|
||||
from .concept_attention_segs import ConceptAttentionSEGSV3
|
||||
from .external_llm_prompt import ExternalLLMPromptV3
|
||||
from .ksampler_attention_coupling import KSamplerAttentionCouplingV3
|
||||
from .ksampler_contextual_attention_coupling import (
|
||||
KSamplerContextualAttentionCouplingV3,
|
||||
)
|
||||
from .ksampler_contextual_diffusion import KSamplerContextualDiffusionV3
|
||||
from .ksampler_prompt_by_region import KSamplerPromptByRegionV3
|
||||
from .ksampler_prompt_by_tiled_region import KSamplerPromptByTiledRegionV3
|
||||
from .ksampler_tiled_attention_coupling import (
|
||||
KSamplerTiledAttentionCouplingV3,
|
||||
)
|
||||
from .ksampler_tiled_diffusion import KSamplerTiledDiffusionV3
|
||||
from .legacy_node_wrappers import (
|
||||
ConditioningBatchAppendV3,
|
||||
ConditioningBatchStartV3,
|
||||
@@ -28,9 +42,7 @@ def get_nodes() -> list[type[object]]:
|
||||
EncodePromptBatchV3,
|
||||
GroundedSAMModelInfoV3,
|
||||
GroundingDINOModelLoaderV3,
|
||||
KSamplerContextualDiffusionV3,
|
||||
KSamplerExtrasV3,
|
||||
KSamplerTiledDiffusionV3,
|
||||
LatentDiagnosticsV3,
|
||||
LayerStyleSAMModelsAdapterV3,
|
||||
LoadUltralyticsModelV3,
|
||||
@@ -51,6 +63,7 @@ def get_nodes() -> list[type[object]]:
|
||||
from .load_mask_batch import LoadMaskBatchV3
|
||||
from .mask_to_segs import MaskToSEGSV3
|
||||
from .scale_factor import ScaleFactorV3
|
||||
from .seed_variation import SeedVariationV3
|
||||
from .simple_load_checkpoint import SimpleLoadCheckpointV3
|
||||
from .simple_load_flux import SimpleLoadFluxV3
|
||||
from .simple_load_flux2 import SimpleLoadFlux2V3
|
||||
@@ -62,6 +75,10 @@ def get_nodes() -> list[type[object]]:
|
||||
from .wd14_tagger_loader import WD14TaggerLoaderV3
|
||||
|
||||
nodes: list[type[object]] = [
|
||||
AllPromptAttentionSEGSV3,
|
||||
AttentionCaptureModelV3,
|
||||
AttentionMaskedConditioningV3,
|
||||
AttentionRegionMaskV3,
|
||||
BatchRegionConditioningV3,
|
||||
BatchSEGSV3,
|
||||
ConditioningBatchAppendV3,
|
||||
@@ -76,6 +93,9 @@ def get_nodes() -> list[type[object]]:
|
||||
GroundedSAMModelInfoV3,
|
||||
GroundingDINOModelLoaderV3,
|
||||
KSamplerExtrasV3,
|
||||
KSamplerAttentionCouplingV3,
|
||||
KSamplerTiledAttentionCouplingV3,
|
||||
KSamplerContextualAttentionCouplingV3,
|
||||
KSamplerPromptByRegionV3,
|
||||
KSamplerPromptByTiledRegionV3,
|
||||
KSamplerContextualDiffusionV3,
|
||||
@@ -93,7 +113,9 @@ def get_nodes() -> list[type[object]]:
|
||||
SAMModelLoaderV3,
|
||||
SEGSFromSAMOutputV3,
|
||||
ScaleFactorV3,
|
||||
ConceptAttentionSEGSV3,
|
||||
SeedV3,
|
||||
SeedVariationV3,
|
||||
SimpleLoadAnimaV3,
|
||||
SimplePreviewSEGSV3,
|
||||
SimpleLoadCheckpointV3,
|
||||
@@ -113,12 +135,14 @@ def get_nodes() -> list[type[object]]:
|
||||
if not prompt_control_is_available():
|
||||
return nodes
|
||||
|
||||
from .apply_automatic_negpip import ApplyAutomaticNegpipV3
|
||||
from .attach_regional_global_conditioning import (
|
||||
AttachRegionalGlobalConditioningV3,
|
||||
)
|
||||
from .encode_prompt_batch_with_prompt_control import (
|
||||
EncodePromptBatchWithPromptControl,
|
||||
)
|
||||
from .label_regional_lora_hooks import LabelRegionalLoraHooksV3
|
||||
from .prepare_regional_lora_hooks import PrepareRegionalLoraHooksV3
|
||||
from .schedule_and_encode_prompts_with_prompt_control import (
|
||||
ScheduleAndEncodePromptsWithPromptControl,
|
||||
@@ -126,8 +150,10 @@ def get_nodes() -> list[type[object]]:
|
||||
|
||||
return [
|
||||
*nodes,
|
||||
ApplyAutomaticNegpipV3,
|
||||
AttachRegionalGlobalConditioningV3,
|
||||
EncodePromptBatchWithPromptControl,
|
||||
LabelRegionalLoraHooksV3,
|
||||
PrepareRegionalLoraHooksV3,
|
||||
ScheduleAndEncodePromptsWithPromptControl,
|
||||
]
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for exposing every mapped prompt attention region."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_region_capture import AttentionEvidenceMode
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AllPromptAttentionSEGSV3(_ComfyNodeBase):
|
||||
"""Expose separate overlapping regions for all readable positive concepts."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "image"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the all-prompt downstream attention-SEGS contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AllPromptAttentionSEGS",
|
||||
display_name="All Prompt Attention SEGS",
|
||||
category="SimpleSyrup/Detection",
|
||||
description=(
|
||||
"Returns separate overlapping soft SEGS for every mapped concept "
|
||||
"in the positive prompt that produced the image."
|
||||
),
|
||||
search_aliases=[
|
||||
"all attention heatmaps",
|
||||
"prompt insight",
|
||||
"token regions",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
tooltip=(
|
||||
"Image whose graph provenance identifies the upstream sampler; "
|
||||
"the image is returned unchanged."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(
|
||||
_comfy_io,
|
||||
evidence_mode_default=AttentionEvidenceMode.RAW,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Image.Output("image", tooltip="Unchanged connected image."),
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip="Separate labeled overlapping SEGS for prompt concepts.",
|
||||
is_output_list=True,
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft union of every retained prompt attention region.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
geometry_recall: float,
|
||||
split_sensitivity: float,
|
||||
instance_recall: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
evidence_mode: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Consume all prompt maps and render separate overlapping SEGS."""
|
||||
|
||||
del sampler_stage
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_image(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
image=image,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
geometry_recall=geometry_recall,
|
||||
split_sensitivity=split_sensitivity,
|
||||
instance_recall=instance_recall,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
evidence_mode=evidence_mode,
|
||||
),
|
||||
)
|
||||
return result.image, result.segs, result.mask
|
||||
@@ -0,0 +1,69 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Internal Comfy v3 node for model-family automatic NegPiP preparation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..services.negpip_model_service import NEGPIP_MODEL_SERVICE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
pass
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class ApplyAutomaticNegpipV3(_ComfyNodeBase):
|
||||
"""Patch supported MODEL/CLIP pairs after a negative prompt-weight trigger."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the internal runtime patch boundary."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.ApplyAutomaticNegpip",
|
||||
display_name="Apply Automatic NegPiP (Internal)",
|
||||
category="SimpleSyrup/Internal",
|
||||
description=(
|
||||
"Internal model-family NegPiP preparation injected by Schedule & "
|
||||
"Encode Prompts after detecting a negative prompt weight."
|
||||
),
|
||||
is_dev_only=True,
|
||||
inputs=[
|
||||
_comfy_io.Model.Input(
|
||||
"model",
|
||||
tooltip="MODEL inspected and cloned only when NegPiP is supported.",
|
||||
),
|
||||
_comfy_io.Clip.Input(
|
||||
"clip",
|
||||
tooltip="CLIP cloned with the matching NegPiP encoder behavior.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Model.Output(
|
||||
"model",
|
||||
tooltip="MODEL carrying one supported NegPiP attention patch set.",
|
||||
),
|
||||
_comfy_io.Clip.Output(
|
||||
"clip",
|
||||
tooltip="CLIP carrying matching negative-weight encoding behavior.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model: object, clip: object) -> tuple[object, object]:
|
||||
"""Return the supported patched pair or the original unsupported pair."""
|
||||
|
||||
return NEGPIP_MODEL_SERVICE.prepare(model, clip)
|
||||
@@ -0,0 +1,96 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Internal Comfy v3 MODEL derivation node injected by prompt provenance."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..services.attention_capture_model_service import ATTENTION_CAPTURE_MODEL_SERVICE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionCaptureModelV3(_ComfyNodeBase):
|
||||
"""Derive an observation-only MODEL from a prompt-injected capture plan."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the dev-only internal capture node contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionCaptureModel",
|
||||
display_name="Attention Capture Model (Internal)",
|
||||
category="SimpleSyrup/Internal",
|
||||
description=(
|
||||
"Internal prompt-scoped MODEL observer used by downstream "
|
||||
"attention-region nodes."
|
||||
),
|
||||
is_dev_only=True,
|
||||
inputs=[
|
||||
_comfy_io.Model.Input(
|
||||
"model",
|
||||
tooltip="Upstream MODEL cloned for observation-only capture.",
|
||||
),
|
||||
_comfy_io.Conditioning.Input(
|
||||
"positive",
|
||||
tooltip="Positive conditioning whose prompt tokens are mapped.",
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"plan_json",
|
||||
multiline=False,
|
||||
tooltip="Prompt-injected capture plan for the target sampler.",
|
||||
),
|
||||
_comfy_io.Clip.Input(
|
||||
"clip",
|
||||
optional=True,
|
||||
tooltip="Graph-visible CLIP used for exact prompt token mapping.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Model.Output(
|
||||
"model",
|
||||
tooltip="MODEL carrying one observation-only attention observer.",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: object,
|
||||
positive: object,
|
||||
plan_json: str,
|
||||
clip: object | None = None,
|
||||
) -> tuple[object]:
|
||||
"""Prepare and publish capture state before the target sampler runs."""
|
||||
|
||||
del positive
|
||||
return (
|
||||
ATTENTION_CAPTURE_MODEL_SERVICE.prepare(
|
||||
model=model,
|
||||
plan_json=plan_json,
|
||||
clip=clip,
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def fingerprint_inputs(cls, **kwargs: object) -> float:
|
||||
"""Force prompt-scoped session publication on every graph execution."""
|
||||
|
||||
del kwargs
|
||||
return float("nan")
|
||||
@@ -0,0 +1,157 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for attention-derived regional conditioning of a later sampler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionMaskedConditioningV3(_ComfyNodeBase):
|
||||
"""Mask supplied conditioning with attention captured from an earlier sampler."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "latent"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare later-pass attention-masked conditioning inputs and outputs."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionMaskedConditioning",
|
||||
display_name="Attention Masked Conditioning",
|
||||
category="SimpleSyrup/Conditioning",
|
||||
description=(
|
||||
"Masks supplied conditioning to regions discovered from an earlier "
|
||||
"sampler so it can control a later sampling pass."
|
||||
),
|
||||
search_aliases=[
|
||||
"regional conditioning",
|
||||
"attention conditioning",
|
||||
"masked prompt",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Latent.Input(
|
||||
"latent",
|
||||
tooltip=(
|
||||
"Latent produced by the sampler used for localization; it is "
|
||||
"returned unchanged for the later sampler."
|
||||
),
|
||||
),
|
||||
_comfy_io.Conditioning.Input(
|
||||
"conditioning",
|
||||
tooltip="Conditioning to apply only inside the discovered regions.",
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=("Concepts used to discover regions, separated by |."),
|
||||
),
|
||||
_comfy_io.Float.Input(
|
||||
"conditioning_strength",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=10.0,
|
||||
step=0.05,
|
||||
tooltip="Strength of the supplied conditioning inside the mask.",
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent", tooltip="Unchanged localization latent."
|
||||
),
|
||||
_comfy_io.Conditioning.Output(
|
||||
"conditioning",
|
||||
tooltip=(
|
||||
"Supplied conditioning carrying the attention-derived mask."
|
||||
),
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft latent-resolution mask attached to the conditioning.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
latent: object,
|
||||
conditioning: object,
|
||||
concepts: str,
|
||||
conditioning_strength: float,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
geometry_recall: float,
|
||||
split_sensitivity: float,
|
||||
instance_recall: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
evidence_mode: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Return later-pass conditioning masked by earlier-sampler attention."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Attention Masked Conditioning requires a concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_latent(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
latent=latent,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
geometry_recall=geometry_recall,
|
||||
split_sensitivity=split_sensitivity,
|
||||
instance_recall=instance_recall,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
evidence_mode=evidence_mode,
|
||||
),
|
||||
)
|
||||
masked = ATTENTION_REGION_NODE_SERVICE.mask_conditioning(
|
||||
conditioning,
|
||||
result.mask,
|
||||
conditioning_strength,
|
||||
)
|
||||
return result.latent, masked, result.mask
|
||||
@@ -0,0 +1,219 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Build shared Comfy v3 inputs and domain controls for attention-region nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.attention_region_capture import (
|
||||
AttentionCaptureProfile,
|
||||
AttentionEvidenceMode,
|
||||
AttentionRegionControls,
|
||||
)
|
||||
|
||||
|
||||
def attention_region_control_inputs(
|
||||
io: Any,
|
||||
*,
|
||||
evidence_mode_default: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT,
|
||||
) -> list[object]:
|
||||
"""Return the complete shared attention-native control schema."""
|
||||
|
||||
return [
|
||||
io.Int.Input(
|
||||
"sampler_stage",
|
||||
default=1,
|
||||
min=-1,
|
||||
max=1024,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Selects the connected sampling stage: 1 is first, 2 is second, "
|
||||
"and 0 or -1 selects the last; oversized values use the last."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"capture_start",
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
max=0.99,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Start of denoising evidence to include; later values ignore more "
|
||||
"of the initial composition phase."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"capture_end",
|
||||
default=1.0,
|
||||
min=0.01,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"End of denoising evidence to include; earlier values ignore more "
|
||||
"late refinement attention."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"minimum_strength",
|
||||
default=0.15,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Minimum normalized attention association retained in a region; "
|
||||
"higher values narrow the silhouette toward its semantic core."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"minimum_consensus",
|
||||
default=0.25,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Fraction of selected observations that must support a pixel; "
|
||||
"higher values keep more persistent regions."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"geometry_recall",
|
||||
default=0.85,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Higher values recover fainter connected object extent from exact "
|
||||
"attention, preserving tails and shafts but admitting more halos."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"split_sensitivity",
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Sensitivity to divide one connected region around separate peaks; "
|
||||
"higher values can split a soft silhouette into multiple instances."
|
||||
),
|
||||
),
|
||||
io.Float.Input(
|
||||
"instance_recall",
|
||||
default=0.65,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Higher values retain weaker disconnected instances relative to "
|
||||
"the strongest one, which helps repeated sparse concepts."
|
||||
),
|
||||
),
|
||||
io.Int.Input(
|
||||
"minimum_region_size",
|
||||
default=512,
|
||||
min=1,
|
||||
max=1048576,
|
||||
step=1,
|
||||
tooltip="Discard attention components smaller than this many pixels.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"keep_only",
|
||||
default=1,
|
||||
min=0,
|
||||
max=1024,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Keep the best N instances per concept; 1 keeps the largest and "
|
||||
"0 keeps all."
|
||||
),
|
||||
),
|
||||
io.Combo.Input(
|
||||
"keep_by",
|
||||
options=["largest size", "highest confidence"],
|
||||
default="largest size",
|
||||
tooltip="Ranks retained instances by area or attention confidence.",
|
||||
),
|
||||
io.Boolean.Input(
|
||||
"combine_segs",
|
||||
default=False,
|
||||
tooltip="Combines retained instances of each concept into one SEG.",
|
||||
),
|
||||
io.Float.Input(
|
||||
"matte_solidity",
|
||||
default=0.75,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=(
|
||||
"Higher values flatten accepted interiors toward fully opaque alpha."
|
||||
),
|
||||
),
|
||||
io.Int.Input(
|
||||
"edge_feather",
|
||||
default=8,
|
||||
min=0,
|
||||
max=4096,
|
||||
step=1,
|
||||
tooltip="Width in output pixels of the matte boundary transition.",
|
||||
),
|
||||
io.Combo.Input(
|
||||
"capture_profile",
|
||||
options=[profile.value for profile in AttentionCaptureProfile],
|
||||
default=AttentionCaptureProfile.FAST.value,
|
||||
tooltip=(
|
||||
"Controls observation density: fast minimizes overhead, balanced "
|
||||
"adds temporal evidence, and exhaustive retains every eligible call."
|
||||
),
|
||||
),
|
||||
io.Combo.Input(
|
||||
"evidence_mode",
|
||||
options=[mode.value for mode in AttentionEvidenceMode],
|
||||
default=evidence_mode_default.value,
|
||||
tooltip=(
|
||||
"Concept isolation derives a cleaner stable region; raw attention "
|
||||
"shows the captured model probabilities without isolation weighting."
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def attention_region_controls(
|
||||
*,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
geometry_recall: float,
|
||||
split_sensitivity: float,
|
||||
instance_recall: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
evidence_mode: str,
|
||||
) -> AttentionRegionControls:
|
||||
"""Build validated domain controls from public node inputs."""
|
||||
|
||||
return AttentionRegionControls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
split_sensitivity=split_sensitivity,
|
||||
minimum_region_size=minimum_region_size,
|
||||
profile=AttentionCaptureProfile(capture_profile),
|
||||
instance_recall=instance_recall,
|
||||
geometry_recall=geometry_recall,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
evidence_mode=AttentionEvidenceMode(evidence_mode),
|
||||
)
|
||||
@@ -0,0 +1,130 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for querying upstream attention as a latent-resolution mask."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class AttentionRegionMaskV3(_ComfyNodeBase):
|
||||
"""Return a reusable mask derived from an upstream sampler's attention."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "latent"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the latent attention-region mask contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.AttentionRegionMask",
|
||||
display_name="Attention Region Mask",
|
||||
category="SimpleSyrup/Masking",
|
||||
description=(
|
||||
"Returns a reusable latent-resolution mask for concepts attended "
|
||||
"by the sampler that produced the connected latent."
|
||||
),
|
||||
search_aliases=["latent attention mask", "prompt mask", "region mask"],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Latent.Input(
|
||||
"latent",
|
||||
tooltip=(
|
||||
"Latent whose graph provenance identifies the upstream "
|
||||
"sampler; the latent is returned unchanged."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=(
|
||||
"Concepts whose attention becomes the mask, separated by |."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent", tooltip="Unchanged connected latent."
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask",
|
||||
tooltip="Soft latent-resolution union of matched regions.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
latent: object,
|
||||
concepts: str,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
geometry_recall: float,
|
||||
split_sensitivity: float,
|
||||
instance_recall: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
evidence_mode: str,
|
||||
) -> tuple[object, object]:
|
||||
"""Consume matching maps and return a reusable latent-resolution mask."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Attention Region Mask requires at least one concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_latent(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
latent=latent,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
geometry_recall=geometry_recall,
|
||||
split_sensitivity=split_sensitivity,
|
||||
instance_recall=instance_recall,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
evidence_mode=evidence_mode,
|
||||
),
|
||||
)
|
||||
return result.latent, result.mask
|
||||
@@ -10,7 +10,7 @@ from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..services.regional_conditioning_service import RegionalConditioningService
|
||||
from .regional_ksampler_schema import regional_conditioning_inputs
|
||||
from .ksampler_schema import regional_conditioning_inputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Expose explicitly requested upstream attention concepts as SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.attention_concepts import parse_attention_concepts
|
||||
from ..services.attention_region_node_service import ATTENTION_REGION_NODE_SERVICE
|
||||
from .attention_region_inputs import (
|
||||
attention_region_control_inputs,
|
||||
attention_region_controls,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class ConceptAttentionSEGSV3(_ComfyNodeBase):
|
||||
"""Return attention-derived instances for explicitly named concepts."""
|
||||
|
||||
GRAPH_PASSTHROUGH_OUTPUTS = {0: "image"}
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the downstream concept-attention SEGS contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.ConceptAttentionSEGS",
|
||||
display_name="Concept Attention SEGS",
|
||||
category="SimpleSyrup/Detection",
|
||||
description=(
|
||||
"Returns regions associated with explicit concepts from a selected "
|
||||
"sampler in the connected image's provenance."
|
||||
),
|
||||
search_aliases=[
|
||||
"attention heatmap",
|
||||
"prompt segmentation",
|
||||
"attention mask",
|
||||
],
|
||||
hidden=[_comfy_io.Hidden.unique_id],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
tooltip=(
|
||||
"Image whose graph provenance identifies the sampling chain; "
|
||||
"the image is returned unchanged."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"concepts",
|
||||
multiline=True,
|
||||
default="subject",
|
||||
tooltip=(
|
||||
"Enter one or more concepts separated by |, such as "
|
||||
"girl | pink hair | cat."
|
||||
),
|
||||
),
|
||||
*attention_region_control_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Image.Output("image", tooltip="Unchanged connected image."),
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip="Labeled attention-derived instances for the concepts.",
|
||||
is_output_list=True,
|
||||
),
|
||||
_comfy_io.Mask.Output(
|
||||
"mask", tooltip="Union of all retained concept instances."
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
concepts: str,
|
||||
sampler_stage: int,
|
||||
capture_start: float,
|
||||
capture_end: float,
|
||||
minimum_strength: float,
|
||||
minimum_consensus: float,
|
||||
geometry_recall: float,
|
||||
split_sensitivity: float,
|
||||
instance_recall: float,
|
||||
minimum_region_size: int,
|
||||
keep_only: int,
|
||||
keep_by: str,
|
||||
combine_segs: bool,
|
||||
matte_solidity: float,
|
||||
edge_feather: int,
|
||||
capture_profile: str,
|
||||
evidence_mode: str,
|
||||
) -> tuple[object, object, object]:
|
||||
"""Consume shared capture evidence and render requested concept SEGS."""
|
||||
|
||||
del sampler_stage
|
||||
if not parse_attention_concepts(concepts):
|
||||
raise ValueError("Concept Attention SEGS requires at least one concept.")
|
||||
result = ATTENTION_REGION_NODE_SERVICE.for_image(
|
||||
request_node_id=str(cls.hidden.unique_id),
|
||||
image=image,
|
||||
controls=attention_region_controls(
|
||||
capture_start=capture_start,
|
||||
capture_end=capture_end,
|
||||
minimum_strength=minimum_strength,
|
||||
minimum_consensus=minimum_consensus,
|
||||
geometry_recall=geometry_recall,
|
||||
split_sensitivity=split_sensitivity,
|
||||
instance_recall=instance_recall,
|
||||
minimum_region_size=minimum_region_size,
|
||||
keep_only=keep_only,
|
||||
keep_by=keep_by,
|
||||
combine_segs=combine_segs,
|
||||
matte_solidity=matte_solidity,
|
||||
edge_feather=edge_feather,
|
||||
capture_profile=capture_profile,
|
||||
evidence_mode=evidence_mode,
|
||||
),
|
||||
)
|
||||
return result.image, result.segs, result.mask
|
||||
@@ -0,0 +1,119 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 KSampler for full-context Attention Coupling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..services.attention_coupling_sampling_service import (
|
||||
AttentionCouplingSamplingService,
|
||||
)
|
||||
from .ksampler_schema import (
|
||||
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT,
|
||||
attention_coupling_ksampler_inputs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class KSamplerAttentionCouplingV3(_ComfyNodeBase):
|
||||
"""Sample supported models with mask-bound regional attention."""
|
||||
|
||||
sampling_service_class: ClassVar[type[AttentionCouplingSamplingService]] = (
|
||||
AttentionCouplingSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the full-context Attention Coupling sampler schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.KSamplerAttentionCoupling",
|
||||
display_name="KSampler (Attention Coupling)",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description=(
|
||||
"With ordinary conditioning and no masks, denoises through the "
|
||||
"normal KSampler path without Attention Coupling preparation. "
|
||||
"With conditioning batches and masks, denoises supported Anima "
|
||||
"and standard SD/SDXL models through one "
|
||||
"shared trajectory while coupling global and masked regional "
|
||||
"cross-attention. LoRAs on the input MODEL and Prompt Control model "
|
||||
"LoRAs on global conditioning entry 0 apply across the image. Regions "
|
||||
"may also carry ordered, independently scheduled model LoRAs "
|
||||
"whose overlapping deltas compose in declared order. Runtime scales "
|
||||
"with active adapters, ranks, and targets. "
|
||||
"Global LoRA and regional LoRA retain independent schedules; "
|
||||
"regional model-side hooks are supported on admitted model families. "
|
||||
"Unsupported adapter targets fail before sampling."
|
||||
),
|
||||
search_aliases=[
|
||||
"attention coupling",
|
||||
"regional lora",
|
||||
"anima regional prompt",
|
||||
"sdxl regional prompt",
|
||||
],
|
||||
inputs=attention_coupling_ksampler_inputs(
|
||||
_comfy_io,
|
||||
region_masks_optional=True,
|
||||
),
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent",
|
||||
tooltip=tooltips.DENOISED_LATENT_OUTPUT,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: object,
|
||||
negative: object,
|
||||
latent_image: dict[str, Any],
|
||||
denoise: float,
|
||||
region_masks: object | None = None,
|
||||
regional_prompt_weight: float = (
|
||||
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT
|
||||
),
|
||||
region_mask_feather: int = 0,
|
||||
) -> tuple[dict[str, Any]]:
|
||||
"""Delegate ordinary or regional sampling to the routing service."""
|
||||
|
||||
output = cls.sampling_service_class().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
region_masks=region_masks,
|
||||
regional_prompt_weight=regional_prompt_weight,
|
||||
region_mask_feather=region_mask_feather,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
)
|
||||
return (output,)
|
||||
@@ -0,0 +1,141 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Native Comfy v3 node for Contextual Attention Coupling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..services.contextual_attention_coupling_sampling_service import (
|
||||
ContextualAttentionCouplingSamplingService,
|
||||
)
|
||||
from .ksampler_schema import (
|
||||
attention_coupling_ksampler_inputs,
|
||||
contextual_diffusion_inputs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class KSamplerContextualAttentionCouplingV3(_ComfyNodeBase):
|
||||
"""Sample contextual model views with regional attention."""
|
||||
|
||||
sampling_service_class: ClassVar[
|
||||
type[ContextualAttentionCouplingSamplingService]
|
||||
] = ContextualAttentionCouplingSamplingService
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Contextual Attention Coupling sampler schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.KSamplerAttentionCouplingContextual",
|
||||
display_name="KSampler (Attention Coupling + Contextual Diffusion)",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description=(
|
||||
"Preserves large-image composition through Contextual Diffusion "
|
||||
"while coupling regional attention in every local and reduced-global "
|
||||
"Anima or standard SD/SDXL view. LoRAs on the input MODEL and Prompt "
|
||||
"Control model LoRAs on global conditioning entry 0 apply in every "
|
||||
"view. Regional LoRA stacks are prepared once, retain "
|
||||
"independent schedules and full quality, and skip inactive work. "
|
||||
"Global LoRA and regional LoRA stacks remain independently scheduled; "
|
||||
"regional model-side hooks are supported on admitted model families. "
|
||||
"Optional SEGS guide the shared local tile plan. Unsupported adapter "
|
||||
"targets fail before sampling."
|
||||
),
|
||||
search_aliases=[
|
||||
"contextual attention coupling",
|
||||
"contextual regional lora",
|
||||
"anima contextual regional prompt",
|
||||
"sdxl contextual regional prompt",
|
||||
"contextual multidiffusion regional lora",
|
||||
"contextual mixture of diffusers regional lora",
|
||||
],
|
||||
inputs=[
|
||||
*attention_coupling_ksampler_inputs(_comfy_io),
|
||||
*contextual_diffusion_inputs(_comfy_io),
|
||||
_comfy_io.SEGS.Input(
|
||||
"segs",
|
||||
optional=True,
|
||||
tooltip=tooltips.CONTEXTUAL_DIFFUSION_SEGS,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent",
|
||||
tooltip=tooltips.DENOISED_LATENT_OUTPUT,
|
||||
),
|
||||
_comfy_io.SEGS.Output(
|
||||
"contexts_segs",
|
||||
tooltip=tooltips.CONTEXTUAL_DIFFUSION_CONTEXTS_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: object,
|
||||
negative: object,
|
||||
region_masks: object,
|
||||
regional_prompt_weight: float,
|
||||
region_mask_feather: int,
|
||||
latent_image: dict[str, Any],
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_context_size: int = 96,
|
||||
latent_context_overlap: int = 32,
|
||||
latent_context_batch_size: int = 4,
|
||||
global_weight: float = 1.0,
|
||||
global_steps: int = 1,
|
||||
global_decay: float = 0.5,
|
||||
segs: object | None = None,
|
||||
) -> tuple[dict[str, Any], object]:
|
||||
"""Delegate the complete request to the combined application service."""
|
||||
|
||||
result = cls.sampling_service_class().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
region_masks=region_masks,
|
||||
regional_prompt_weight=regional_prompt_weight,
|
||||
region_mask_feather=region_mask_feather,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
diffusion_mode=diffusion_mode,
|
||||
latent_context_size=latent_context_size,
|
||||
latent_context_overlap=latent_context_overlap,
|
||||
latent_context_batch_size=latent_context_batch_size,
|
||||
global_weight=global_weight,
|
||||
global_steps=global_steps,
|
||||
global_decay=global_decay,
|
||||
segs=segs,
|
||||
)
|
||||
return result.latent, result.contexts
|
||||
@@ -0,0 +1,133 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Native Comfy v3 node for Contextual Diffusion sampling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..services.contextual_diffusion_sampling_service import (
|
||||
ContextualDiffusionSamplingService,
|
||||
)
|
||||
from .ksampler_schema import (
|
||||
contextual_diffusion_inputs,
|
||||
ksampler_inputs,
|
||||
optional_regional_sampling_inputs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class KSamplerContextualDiffusionV3(_ComfyNodeBase):
|
||||
"""Edit large latents through coordinated global and detailed contexts."""
|
||||
|
||||
service_class: ClassVar[type[ContextualDiffusionSamplingService]] = (
|
||||
ContextualDiffusionSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the native Contextual Diffusion KSampler schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.KSamplerContextualDiffusion",
|
||||
display_name="KSampler (Contextual Diffusion)",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description=(
|
||||
"Preserves composition while applying appearance and subject-detail "
|
||||
"edits to large latents through global context and optional "
|
||||
"SEGS-guided tiles."
|
||||
),
|
||||
search_aliases=[
|
||||
"ksampler",
|
||||
"contextual diffusion",
|
||||
"contextual tiled diffusion",
|
||||
"high resolution edit",
|
||||
"sam tiled diffusion",
|
||||
],
|
||||
inputs=[
|
||||
*ksampler_inputs(_comfy_io, steps_default=4, cfg_default=1.0),
|
||||
*contextual_diffusion_inputs(_comfy_io),
|
||||
*optional_regional_sampling_inputs(
|
||||
_comfy_io,
|
||||
segs_tooltip=tooltips.CONTEXTUAL_DIFFUSION_SEGS,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent",
|
||||
tooltip=tooltips.DENOISED_LATENT_OUTPUT,
|
||||
),
|
||||
_comfy_io.SEGS.Output(
|
||||
"contexts_segs",
|
||||
tooltip=tooltips.CONTEXTUAL_DIFFUSION_CONTEXTS_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_image: dict[str, Any],
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_context_size: int = 96,
|
||||
latent_context_overlap: int = 32,
|
||||
latent_context_batch_size: int = 4,
|
||||
global_weight: float = 1.0,
|
||||
global_steps: int = 1,
|
||||
global_decay: float = 0.5,
|
||||
segs: object | None = None,
|
||||
region_masks: object | None = None,
|
||||
regional_prompt_weight: float = 0.5,
|
||||
region_mask_feather: int = 0,
|
||||
) -> tuple[dict[str, Any], object]:
|
||||
"""Delegate Contextual Diffusion sampling to its application service."""
|
||||
|
||||
result = cls.service_class().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
diffusion_mode=diffusion_mode,
|
||||
latent_context_size=latent_context_size,
|
||||
latent_context_overlap=latent_context_overlap,
|
||||
latent_context_batch_size=latent_context_batch_size,
|
||||
global_weight=global_weight,
|
||||
global_steps=global_steps,
|
||||
global_decay=global_decay,
|
||||
segs=segs,
|
||||
region_masks=region_masks,
|
||||
regional_prompt_weight=regional_prompt_weight,
|
||||
region_mask_feather=region_mask_feather,
|
||||
)
|
||||
return result.latent, result.contexts
|
||||
@@ -12,7 +12,7 @@ from typing import TYPE_CHECKING, Any, ClassVar
|
||||
from ..nodes import tooltips
|
||||
from ..services.ksampler_sampling_service import KSamplerSamplingService
|
||||
from ..services.regional_conditioning_service import RegionalConditioningService
|
||||
from .regional_ksampler_schema import regional_ksampler_inputs
|
||||
from .ksampler_schema import regional_ksampler_inputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
|
||||
@@ -9,10 +9,11 @@ from __future__ import annotations
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.regional_features import RegionalFeature, RegionalFeatureRequest
|
||||
from ..nodes import tooltips
|
||||
from ..services.regional_conditioning_service import RegionalConditioningService
|
||||
from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService
|
||||
from .regional_ksampler_schema import regional_ksampler_inputs, tiled_regional_inputs
|
||||
from .ksampler_schema import regional_ksampler_inputs, tiled_diffusion_inputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -58,7 +59,7 @@ class KSamplerPromptByTiledRegionV3(_ComfyNodeBase):
|
||||
],
|
||||
inputs=[
|
||||
*regional_ksampler_inputs(_comfy_io),
|
||||
*tiled_regional_inputs(_comfy_io),
|
||||
*tiled_diffusion_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
@@ -118,6 +119,8 @@ class KSamplerPromptByTiledRegionV3(_ComfyNodeBase):
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=None,
|
||||
allow_full_context_masks=True,
|
||||
feature_request=RegionalFeatureRequest(
|
||||
frozenset({RegionalFeature.FULL_CONTEXT_MASKED_CONDITIONING})
|
||||
),
|
||||
)
|
||||
return (output,)
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own shared native Comfy v3 KSampler input declarations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.regional_prompting import MAX_REGIONAL_PROMPT_WEIGHT
|
||||
from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES
|
||||
from ..nodes import tooltips
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
|
||||
MAX_LATENT_TILE_SIZE = 512
|
||||
MAX_LATENT_CONTEXT_SIZE = 512
|
||||
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT = 1.0
|
||||
|
||||
|
||||
def ksampler_inputs(
|
||||
comfy_io: Any,
|
||||
*,
|
||||
steps_default: int,
|
||||
cfg_default: float,
|
||||
) -> list[Any]:
|
||||
"""Return standard KSampler inputs with caller-selected defaults."""
|
||||
|
||||
conditioning = comfy_io.Custom("CONDITIONING,CONDITIONING_BATCH")
|
||||
return [
|
||||
comfy_io.Model.Input("model", tooltip=tooltips.SAMPLING_MODEL),
|
||||
comfy_io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
control_after_generate=True,
|
||||
tooltip=tooltips.SAMPLING_SEED,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"steps",
|
||||
default=steps_default,
|
||||
min=1,
|
||||
max=10000,
|
||||
tooltip=tooltips.SAMPLING_STEPS,
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"cfg",
|
||||
default=cfg_default,
|
||||
min=0.0,
|
||||
max=100.0,
|
||||
step=0.1,
|
||||
round=0.01,
|
||||
tooltip=tooltips.SAMPLING_CFG,
|
||||
),
|
||||
comfy_io.Combo.Input(
|
||||
"sampler_name",
|
||||
options=list(sampling_samplers.available_samplers()),
|
||||
tooltip=tooltips.SAMPLER_NAME,
|
||||
),
|
||||
comfy_io.Combo.Input(
|
||||
"scheduler",
|
||||
options=list(sampling_schedulers.available_schedulers()),
|
||||
tooltip=tooltips.SCHEDULER,
|
||||
),
|
||||
conditioning.Input("positive", tooltip=tooltips.POSITIVE_CONDITIONING),
|
||||
conditioning.Input("negative", tooltip=tooltips.NEGATIVE_CONDITIONING),
|
||||
comfy_io.Latent.Input("latent_image", tooltip=tooltips.LATENT_IMAGE),
|
||||
comfy_io.Float.Input(
|
||||
"denoise",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=tooltips.DENOISE_STRENGTH,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def tiled_diffusion_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return tiled diffusion mode and latent tile controls."""
|
||||
|
||||
return [
|
||||
comfy_io.Combo.Input(
|
||||
"diffusion_mode",
|
||||
options=list(TILED_DIFFUSION_MODES),
|
||||
default="multidiffusion",
|
||||
tooltip=tooltips.TILED_DIFFUSION_MODE,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_width",
|
||||
default=128,
|
||||
min=16,
|
||||
max=MAX_LATENT_TILE_SIZE,
|
||||
step=16,
|
||||
tooltip=tooltips.LATENT_TILE_WIDTH,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_height",
|
||||
default=128,
|
||||
min=16,
|
||||
max=MAX_LATENT_TILE_SIZE,
|
||||
step=16,
|
||||
tooltip=tooltips.LATENT_TILE_HEIGHT,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_overlap",
|
||||
default=16,
|
||||
min=0,
|
||||
max=256,
|
||||
step=4,
|
||||
tooltip=tooltips.LATENT_TILE_OVERLAP,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_batch_size",
|
||||
default=4,
|
||||
min=1,
|
||||
max=8,
|
||||
step=1,
|
||||
tooltip=tooltips.LATENT_TILE_BATCH_SIZE,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def contextual_diffusion_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return Contextual Diffusion layout and global schedule controls."""
|
||||
|
||||
return [
|
||||
comfy_io.Combo.Input(
|
||||
"diffusion_mode",
|
||||
options=list(TILED_DIFFUSION_MODES),
|
||||
default="multidiffusion",
|
||||
tooltip=tooltips.TILED_DIFFUSION_MODE,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_context_size",
|
||||
default=96,
|
||||
min=16,
|
||||
max=MAX_LATENT_CONTEXT_SIZE,
|
||||
step=16,
|
||||
tooltip=tooltips.LATENT_CONTEXT_SIZE,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_context_overlap",
|
||||
default=32,
|
||||
min=0,
|
||||
max=256,
|
||||
step=4,
|
||||
tooltip=tooltips.LATENT_CONTEXT_OVERLAP,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_context_batch_size",
|
||||
default=4,
|
||||
min=1,
|
||||
max=8,
|
||||
step=1,
|
||||
tooltip=tooltips.LATENT_CONTEXT_BATCH_SIZE,
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"global_weight",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=2.0,
|
||||
step=0.05,
|
||||
tooltip=tooltips.GLOBAL_CONTEXT_WEIGHT,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"global_steps",
|
||||
default=1,
|
||||
min=0,
|
||||
max=10000,
|
||||
step=1,
|
||||
tooltip=tooltips.GLOBAL_CONTEXT_STEPS,
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"global_decay",
|
||||
default=0.5,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.05,
|
||||
tooltip=tooltips.GLOBAL_CONTEXT_DECAY,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def optional_regional_sampling_inputs(
|
||||
comfy_io: Any,
|
||||
*,
|
||||
segs_tooltip: str,
|
||||
) -> list[Any]:
|
||||
"""Return optional SEGS and Regional Conditioning controls."""
|
||||
|
||||
return [
|
||||
comfy_io.SEGS.Input(
|
||||
"segs",
|
||||
optional=True,
|
||||
tooltip=segs_tooltip,
|
||||
),
|
||||
comfy_io.Mask.Input(
|
||||
"region_masks",
|
||||
optional=True,
|
||||
tooltip=tooltips.OPTIONAL_REGIONAL_MASKS,
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"regional_prompt_weight",
|
||||
default=0.5,
|
||||
min=0.0,
|
||||
max=MAX_REGIONAL_PROMPT_WEIGHT,
|
||||
step=0.01,
|
||||
round=0.01,
|
||||
optional=True,
|
||||
tooltip=tooltips.OPTIONAL_REGIONAL_PROMPT_WEIGHT,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"region_mask_feather",
|
||||
default=0,
|
||||
min=0,
|
||||
max=512,
|
||||
step=1,
|
||||
optional=True,
|
||||
tooltip=tooltips.OPTIONAL_REGION_MASK_FEATHER,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def regional_ksampler_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return common regional KSampler inputs in workflow order."""
|
||||
|
||||
base = ksampler_inputs(comfy_io, steps_default=20, cfg_default=8.0)
|
||||
return [
|
||||
*base[:6],
|
||||
*regional_conditioning_inputs(comfy_io),
|
||||
*base[8:],
|
||||
]
|
||||
|
||||
|
||||
def attention_coupling_ksampler_inputs(
|
||||
comfy_io: Any,
|
||||
*,
|
||||
region_masks_optional: bool = False,
|
||||
) -> list[Any]:
|
||||
"""Return Attention Coupling inputs with caller-owned bypass availability."""
|
||||
|
||||
base = ksampler_inputs(comfy_io, steps_default=20, cfg_default=8.0)
|
||||
conditioning_batch = comfy_io.Custom("CONDITIONING_BATCH")
|
||||
return [
|
||||
comfy_io.Model.Input(
|
||||
"model",
|
||||
tooltip=(
|
||||
"Supported Anima or standard SD/SDXL model used for one shared "
|
||||
"denoiser trajectory. LoRAs patched on this model and Prompt Control "
|
||||
"model LoRAs on conditioning entry 0 apply globally."
|
||||
),
|
||||
),
|
||||
*base[1:6],
|
||||
comfy_io.MultiType.Input(
|
||||
"positive",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Global-first positive conditioning: entry 0 is global and later "
|
||||
"entries pair with masks. Regional Prompt Control WeightHooks may "
|
||||
"contain ordered regional LoRA stacks with independent schedules. "
|
||||
"Model LoRA hooks on entry 0 apply across the image."
|
||||
),
|
||||
),
|
||||
comfy_io.MultiType.Input(
|
||||
"negative",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Global-first negative conditioning aligned to the same masks; "
|
||||
"its global model hooks must match the positive global entry. "
|
||||
"Regional LoRA hooks retain their negative-branch ownership and "
|
||||
"independent schedules."
|
||||
),
|
||||
),
|
||||
comfy_io.Mask.Input(
|
||||
"region_masks",
|
||||
optional=region_masks_optional,
|
||||
tooltip=(
|
||||
"Optional ordered masks paired with conditioning entries 1 onward. "
|
||||
"Leave disconnected with ordinary conditioning to bypass Attention "
|
||||
"Coupling. In overlaps, prompt contributions are normalized while "
|
||||
"regional LoRA deltas add in declared adapter and region order."
|
||||
),
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"regional_prompt_weight",
|
||||
default=ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT,
|
||||
min=0.0,
|
||||
max=MAX_REGIONAL_PROMPT_WEIGHT,
|
||||
step=0.01,
|
||||
round=0.01,
|
||||
tooltip=(
|
||||
"Balances regional cross-attention against the global prompt from "
|
||||
"0 (global only) to 1 (regional only inside solid masks); regional "
|
||||
"LoRA strength remains controlled by each hook."
|
||||
),
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"region_mask_feather",
|
||||
default=0,
|
||||
min=0,
|
||||
max=512,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Softens Attention Coupling and regional LoRA boundaries by "
|
||||
"this many image pixels; 0 preserves authored mask values."
|
||||
),
|
||||
),
|
||||
*base[8:],
|
||||
]
|
||||
|
||||
|
||||
def regional_conditioning_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return the authoritative ordered regional-composition inputs."""
|
||||
|
||||
conditioning_batch = comfy_io.Custom("CONDITIONING_BATCH")
|
||||
return [
|
||||
comfy_io.MultiType.Input(
|
||||
"positive",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Positive conditioning whose first batch entry is global and "
|
||||
"later entries pair with masks in order."
|
||||
),
|
||||
),
|
||||
comfy_io.MultiType.Input(
|
||||
"negative",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Negative conditioning whose first batch entry is global and "
|
||||
"later entries pair with masks in order."
|
||||
),
|
||||
),
|
||||
comfy_io.Mask.Input(
|
||||
"region_masks",
|
||||
tooltip=(
|
||||
"Ordered authored masks; mask 0 pairs with conditioning batch entry 1."
|
||||
),
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"regional_prompt_weight",
|
||||
default=0.5,
|
||||
min=0.0,
|
||||
max=MAX_REGIONAL_PROMPT_WEIGHT,
|
||||
step=0.01,
|
||||
round=0.01,
|
||||
tooltip=(
|
||||
"Balances regional prompts against the global prompt; 0 uses "
|
||||
"only global prompting, 1 uses only regional prompting inside "
|
||||
"solid masks, and overlaps reduce the global share further."
|
||||
),
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"region_mask_feather",
|
||||
default=0,
|
||||
min=0,
|
||||
max=512,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Softens regional mask edges by this many image pixels; 0 "
|
||||
"preserves authored mask values."
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,138 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 KSampler for tiled Attention Coupling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..services.tiled_attention_coupling_sampling_service import (
|
||||
TiledAttentionCouplingSamplingService,
|
||||
)
|
||||
from .ksampler_schema import (
|
||||
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT,
|
||||
attention_coupling_ksampler_inputs,
|
||||
tiled_diffusion_inputs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class KSamplerTiledAttentionCouplingV3(_ComfyNodeBase):
|
||||
"""Sample supported models with tiled regional attention."""
|
||||
|
||||
sampling_service_class: ClassVar[type[TiledAttentionCouplingSamplingService]] = (
|
||||
TiledAttentionCouplingSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the tiled Attention Coupling sampler schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.KSamplerAttentionCouplingTiled",
|
||||
display_name="KSampler (Attention Coupling + Tiled Diffusion)",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description=(
|
||||
"With ordinary conditioning and no masks, uses normal tiled "
|
||||
"diffusion without Attention Coupling preparation. With conditioning "
|
||||
"batches and masks, denoises large Anima and standard SD/SDXL "
|
||||
"latents in tiles through "
|
||||
"one shared model trajectory per tile batch while coupling global "
|
||||
"and masked regional cross-attention. LoRAs on the input MODEL and "
|
||||
"Prompt Control model LoRAs on global conditioning entry 0 apply "
|
||||
"across every tile. Regions may carry independently scheduled "
|
||||
"regional LoRA stacks; inactive attention and LoRA work is pruned "
|
||||
"without changing quality. MultiDiffusion or Mixture of Diffusers "
|
||||
"fuses restored tile predictions. Global LoRA and regional LoRA "
|
||||
"stacks retain independent schedules; regional model-side hooks are "
|
||||
"supported on admitted model families. Unsupported adapter targets "
|
||||
"fail before sampling."
|
||||
),
|
||||
search_aliases=[
|
||||
"attention coupling tiled",
|
||||
"regional lora tiled",
|
||||
"anima tiled regional prompt",
|
||||
"sdxl tiled regional prompt",
|
||||
"multidiffusion regional lora",
|
||||
"mixture of diffusers regional lora",
|
||||
],
|
||||
inputs=[
|
||||
*attention_coupling_ksampler_inputs(
|
||||
_comfy_io,
|
||||
region_masks_optional=True,
|
||||
),
|
||||
*tiled_diffusion_inputs(_comfy_io),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent",
|
||||
tooltip=tooltips.DENOISED_LATENT_OUTPUT,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: object,
|
||||
negative: object,
|
||||
latent_image: dict[str, Any],
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_tile_width: int = 128,
|
||||
latent_tile_height: int = 128,
|
||||
latent_tile_overlap: int = 16,
|
||||
latent_tile_batch_size: int = 4,
|
||||
region_masks: object | None = None,
|
||||
regional_prompt_weight: float = (
|
||||
ATTENTION_COUPLING_REGIONAL_PROMPT_WEIGHT_DEFAULT
|
||||
),
|
||||
region_mask_feather: int = 0,
|
||||
) -> tuple[dict[str, Any]]:
|
||||
"""Delegate ordinary or regional tiled sampling to the routing service."""
|
||||
|
||||
output = cls.sampling_service_class().sample(
|
||||
diffusion_mode=diffusion_mode,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
region_masks=region_masks,
|
||||
regional_prompt_weight=regional_prompt_weight,
|
||||
region_mask_feather=region_mask_feather,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
latent_tile_width=latent_tile_width,
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=None,
|
||||
differential_diffusion=False,
|
||||
)
|
||||
return (output,)
|
||||
@@ -0,0 +1,124 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Native Comfy v3 node for selectable tiled diffusion sampling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService
|
||||
from .ksampler_schema import (
|
||||
ksampler_inputs,
|
||||
optional_regional_sampling_inputs,
|
||||
tiled_diffusion_inputs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class KSamplerTiledDiffusionV3(_ComfyNodeBase):
|
||||
"""Sample latents with selectable tiled diffusion denoising."""
|
||||
|
||||
service_class: ClassVar[type[TiledDiffusionSamplingService]] = (
|
||||
TiledDiffusionSamplingService
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the native tiled diffusion KSampler schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.KSamplerTiledDiffusion",
|
||||
display_name="KSampler (Tiled Diffusion)",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description="Denoises latents with selectable tiled diffusion sampling.",
|
||||
search_aliases=[
|
||||
"ksampler",
|
||||
"sampler",
|
||||
"tiled diffusion",
|
||||
"multidiffusion",
|
||||
"multi diffusion",
|
||||
"mixture of diffusers",
|
||||
],
|
||||
inputs=[
|
||||
*ksampler_inputs(_comfy_io, steps_default=20, cfg_default=8.0),
|
||||
*tiled_diffusion_inputs(_comfy_io),
|
||||
*optional_regional_sampling_inputs(
|
||||
_comfy_io,
|
||||
segs_tooltip=(
|
||||
"Optional image regions that guide irregular tile "
|
||||
"boundaries while preserving the configured overlap."
|
||||
),
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
None,
|
||||
tooltip=tooltips.DENOISED_LATENT_OUTPUT,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_image: dict[str, Any],
|
||||
denoise: float = 1.0,
|
||||
diffusion_mode: str = "multidiffusion",
|
||||
latent_tile_width: int = 128,
|
||||
latent_tile_height: int = 128,
|
||||
latent_tile_overlap: int = 16,
|
||||
latent_tile_batch_size: int = 4,
|
||||
segs: object | None = None,
|
||||
region_masks: object | None = None,
|
||||
regional_prompt_weight: float = 0.5,
|
||||
region_mask_feather: int = 0,
|
||||
) -> tuple[dict[str, Any]]:
|
||||
"""Delegate tiled diffusion sampling to its application service."""
|
||||
|
||||
output = cls.service_class().sample(
|
||||
diffusion_mode=diffusion_mode,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
latent_tile_width=latent_tile_width,
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=None,
|
||||
segs=segs,
|
||||
region_masks=region_masks,
|
||||
regional_prompt_weight=regional_prompt_weight,
|
||||
region_mask_feather=region_mask_feather,
|
||||
)
|
||||
return (output,)
|
||||
@@ -0,0 +1,67 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Expose internal stable identity labeling for regional LoRA hooks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..runtime.regional_lora_hook_identity import label_regional_lora_hooks
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class LabelRegionalLoraHooksV3(_ComfyNodeBase):
|
||||
"""Attach explicit adapter identities to cloned schedule-bearing hooks."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the internal ordered identity boundary."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.LabelRegionalLoraHooks",
|
||||
display_name="Label Regional LoRA Hooks",
|
||||
category="SimpleSyrup/Internal",
|
||||
description=(
|
||||
"Labels scheduled regional LoRA hooks with stable adapter "
|
||||
"identities without changing their weights or keyframes."
|
||||
),
|
||||
inputs=[
|
||||
_comfy_io.Hooks.Input(
|
||||
"hooks",
|
||||
tooltip="Prompt Control hooks whose schedules remain unchanged.",
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"adapter_identities_json",
|
||||
tooltip=(
|
||||
"Ordered JSON array with one stable identity per LoRA hook."
|
||||
),
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Hooks.Output(
|
||||
"hooks",
|
||||
tooltip="Cloned hooks carrying the supplied stable identities.",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, hooks: object, adapter_identities_json: str) -> tuple[object]:
|
||||
"""Return one labeled clone of the supplied HookGroup."""
|
||||
|
||||
return (label_regional_lora_hooks(hooks, adapter_identities_json),)
|
||||
@@ -24,9 +24,7 @@ from ..nodes.encode_prompt_batch import EncodePromptBatch
|
||||
from ..nodes.grounded_sam_model_info import GroundedSAMModelInfo
|
||||
from ..nodes.grounding_dino_model_loader import GroundingDINOModelLoader
|
||||
from ..nodes.image_resize_to_target import ResizeImageToTarget
|
||||
from ..nodes.ksampler_contextual_diffusion import KSamplerContextualDiffusion
|
||||
from ..nodes.ksampler_extras import KSamplerExtras
|
||||
from ..nodes.ksampler_tiled_diffusion import KSamplerTiledDiffusion
|
||||
from ..nodes.latent_diagnostics import LatentDiagnostics
|
||||
from ..nodes.layerstyle_sam_models_adapter import LayerStyleSAMModelsAdapter
|
||||
from ..nodes.load_ultralytics_model import LoadUltralyticsModel
|
||||
@@ -62,8 +60,6 @@ _HIDDEN_INPUTS = {
|
||||
"DYNPROMPT": "dynprompt",
|
||||
"EXTRA_PNGINFO": "extra_pnginfo",
|
||||
"UNIQUE_ID": "unique_id",
|
||||
"AUTH_TOKEN_COMFY_ORG": "auth_token_comfy_org",
|
||||
"API_KEY_COMFY_ORG": "api_key_comfy_org",
|
||||
}
|
||||
|
||||
|
||||
@@ -151,22 +147,6 @@ class KSamplerExtrasV3(LegacyNodeV3Adapter):
|
||||
DISPLAY_NAME = "KSampler (Extras)"
|
||||
|
||||
|
||||
class KSamplerTiledDiffusionV3(LegacyNodeV3Adapter):
|
||||
"""Expose KSampler Tiled Diffusion through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = KSamplerTiledDiffusion
|
||||
NODE_ID = "SimpleSyrup.KSamplerTiledDiffusion"
|
||||
DISPLAY_NAME = "KSampler (Tiled Diffusion)"
|
||||
|
||||
|
||||
class KSamplerContextualDiffusionV3(LegacyNodeV3Adapter):
|
||||
"""Expose KSampler Contextual Diffusion through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = KSamplerContextualDiffusion
|
||||
NODE_ID = "SimpleSyrup.KSamplerContextualDiffusion"
|
||||
DISPLAY_NAME = "KSampler (Contextual Diffusion)"
|
||||
|
||||
|
||||
class LayerStyleSAMModelsAdapterV3(LegacyNodeV3Adapter):
|
||||
"""Expose LayerStyle SAM Models Adapter through Comfy v3 only."""
|
||||
|
||||
@@ -523,8 +503,6 @@ __all__ = [
|
||||
"GroundedSAMModelInfoV3",
|
||||
"GroundingDINOModelLoaderV3",
|
||||
"KSamplerExtrasV3",
|
||||
"KSamplerContextualDiffusionV3",
|
||||
"KSamplerTiledDiffusionV3",
|
||||
"LatentDiagnosticsV3",
|
||||
"LayerStyleSAMModelsAdapterV3",
|
||||
"LoadUltralyticsModelV3",
|
||||
|
||||
@@ -1,169 +0,0 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Shared Comfy v3 schema declarations for regional KSamplers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.regional_prompting import MAX_REGIONAL_PROMPT_WEIGHT
|
||||
from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES
|
||||
from ..nodes import tooltips
|
||||
from ..nodes.ksampler_tiled_diffusion import MAX_LATENT_TILE_SIZE
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
|
||||
|
||||
def regional_ksampler_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return common regional KSampler inputs in workflow order."""
|
||||
|
||||
return [
|
||||
comfy_io.Model.Input("model", tooltip=tooltips.SAMPLING_MODEL),
|
||||
comfy_io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
control_after_generate=True,
|
||||
tooltip=tooltips.SAMPLING_SEED,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"steps",
|
||||
default=20,
|
||||
min=1,
|
||||
max=10000,
|
||||
tooltip=tooltips.SAMPLING_STEPS,
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"cfg",
|
||||
default=8.0,
|
||||
min=0.0,
|
||||
max=100.0,
|
||||
step=0.1,
|
||||
round=0.01,
|
||||
tooltip=tooltips.SAMPLING_CFG,
|
||||
),
|
||||
comfy_io.Combo.Input(
|
||||
"sampler_name",
|
||||
options=list(sampling_samplers.available_samplers()),
|
||||
tooltip=tooltips.SAMPLER_NAME,
|
||||
),
|
||||
comfy_io.Combo.Input(
|
||||
"scheduler",
|
||||
options=list(sampling_schedulers.available_schedulers()),
|
||||
tooltip=tooltips.SCHEDULER,
|
||||
),
|
||||
*regional_conditioning_inputs(comfy_io),
|
||||
comfy_io.Latent.Input("latent_image", tooltip=tooltips.LATENT_IMAGE),
|
||||
comfy_io.Float.Input(
|
||||
"denoise",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.01,
|
||||
tooltip=tooltips.DENOISE_STRENGTH,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def regional_conditioning_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return the authoritative ordered regional-composition inputs."""
|
||||
|
||||
conditioning_batch = comfy_io.Custom("CONDITIONING_BATCH")
|
||||
return [
|
||||
comfy_io.MultiType.Input(
|
||||
"positive",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Positive conditioning whose first batch entry is global and "
|
||||
"later entries pair with masks in order."
|
||||
),
|
||||
),
|
||||
comfy_io.MultiType.Input(
|
||||
"negative",
|
||||
[comfy_io.Conditioning, conditioning_batch],
|
||||
tooltip=(
|
||||
"Negative conditioning whose first batch entry is global and "
|
||||
"later entries pair with masks in order."
|
||||
),
|
||||
),
|
||||
comfy_io.Mask.Input(
|
||||
"region_masks",
|
||||
tooltip=(
|
||||
"Ordered authored masks; mask 0 pairs with conditioning batch entry 1."
|
||||
),
|
||||
),
|
||||
comfy_io.Float.Input(
|
||||
"regional_prompt_weight",
|
||||
default=0.5,
|
||||
min=0.0,
|
||||
max=MAX_REGIONAL_PROMPT_WEIGHT,
|
||||
step=0.01,
|
||||
round=0.01,
|
||||
tooltip=(
|
||||
"Balances regional prompts against the global prompt; 0 uses "
|
||||
"only global prompting, 1 uses only regional prompting inside "
|
||||
"solid masks, and overlaps reduce the global share further."
|
||||
),
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"region_mask_feather",
|
||||
default=0,
|
||||
min=0,
|
||||
max=512,
|
||||
step=1,
|
||||
tooltip=(
|
||||
"Softens regional mask edges by this many image pixels; 0 "
|
||||
"preserves authored mask values."
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def tiled_regional_inputs(comfy_io: Any) -> list[Any]:
|
||||
"""Return tile controls matching KSampler Tiled Diffusion."""
|
||||
|
||||
return [
|
||||
comfy_io.Combo.Input(
|
||||
"diffusion_mode",
|
||||
options=list(TILED_DIFFUSION_MODES),
|
||||
default="multidiffusion",
|
||||
tooltip=(
|
||||
"Tiled blend method; MultiDiffusion is steady while Mixture of "
|
||||
"Diffusers weights tile centers more strongly."
|
||||
),
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_width",
|
||||
default=128,
|
||||
min=16,
|
||||
max=MAX_LATENT_TILE_SIZE,
|
||||
step=16,
|
||||
tooltip=tooltips.LATENT_TILE_WIDTH,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_height",
|
||||
default=128,
|
||||
min=16,
|
||||
max=MAX_LATENT_TILE_SIZE,
|
||||
step=16,
|
||||
tooltip=tooltips.LATENT_TILE_HEIGHT,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_overlap",
|
||||
default=16,
|
||||
min=0,
|
||||
max=256,
|
||||
step=4,
|
||||
tooltip=tooltips.LATENT_TILE_OVERLAP,
|
||||
),
|
||||
comfy_io.Int.Input(
|
||||
"latent_tile_batch_size",
|
||||
default=4,
|
||||
min=1,
|
||||
max=8,
|
||||
step=1,
|
||||
tooltip=tooltips.LATENT_TILE_BATCH_SIZE,
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,97 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Expose seed variation as a Comfy v3 MODEL patch node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.seed_variation import (
|
||||
MAX_SEED,
|
||||
MAX_VARIATION_STRENGTH,
|
||||
MIN_SEED,
|
||||
MIN_VARIATION_STRENGTH,
|
||||
)
|
||||
from ..nodes import tooltips
|
||||
from ..services.seed_variation_model_service import SEED_VARIATION_MODEL_SERVICE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class SeedVariationV3(_ComfyNodeBase):
|
||||
"""Derive a MODEL that varies sampler-provided initial noise."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the seed-variation MODEL patch contract."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.SeedVariation",
|
||||
display_name="Seed Variation",
|
||||
category="SimpleSyrup/Sampling",
|
||||
description=(
|
||||
"Creates related generations by mixing sampler noise toward a "
|
||||
"second deterministic seed."
|
||||
),
|
||||
search_aliases=["variation seed", "subseed", "seed interpolation"],
|
||||
inputs=[
|
||||
_comfy_io.Model.Input(
|
||||
"model",
|
||||
tooltip=tooltips.SEED_VARIATION_MODEL_INPUT,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"variation_seed",
|
||||
default=0,
|
||||
min=MIN_SEED,
|
||||
max=MAX_SEED,
|
||||
control_after_generate=True,
|
||||
tooltip=tooltips.VARIATION_SEED,
|
||||
),
|
||||
_comfy_io.Float.Input(
|
||||
"variation_strength",
|
||||
default=0.0,
|
||||
min=MIN_VARIATION_STRENGTH,
|
||||
max=MAX_VARIATION_STRENGTH,
|
||||
step=0.01,
|
||||
round=0.01,
|
||||
tooltip=tooltips.VARIATION_STRENGTH,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Model.Output(
|
||||
"model",
|
||||
tooltip=tooltips.SEED_VARIATION_MODEL_OUTPUT,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: object,
|
||||
variation_seed: int,
|
||||
variation_strength: float,
|
||||
) -> tuple[object]:
|
||||
"""Return the source or a MODEL carrying initial-noise variation."""
|
||||
|
||||
return (
|
||||
SEED_VARIATION_MODEL_SERVICE.prepare(
|
||||
model=model,
|
||||
variation_seed=variation_seed,
|
||||
variation_strength=variation_strength,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,66 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Detect the exact installed Anima regional-attention capability surface."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import comfy.latent_formats
|
||||
import comfy.model_base
|
||||
import comfy.model_patcher
|
||||
from comfy.ldm.anima.model import Anima as AnimaDiffusionModel
|
||||
|
||||
from ..domain.regional_model_capabilities import (
|
||||
RegionalAttentionBackend,
|
||||
RegionalAttentionTopology,
|
||||
RegionalControlGligenPolicy,
|
||||
RegionalLatentLayout,
|
||||
RegionalModelCapabilities,
|
||||
RegionalModelFamily,
|
||||
RegionalPatchConflict,
|
||||
RegionalReferenceLatentPolicy,
|
||||
RegionalSpatialPatchSupport,
|
||||
)
|
||||
|
||||
|
||||
class AnimaModelCapabilityDetector:
|
||||
"""Own defensive routing to the specialized Anima attention backend."""
|
||||
|
||||
def detect(
|
||||
self,
|
||||
patcher: comfy.model_patcher.ModelPatcher,
|
||||
) -> RegionalModelCapabilities | None:
|
||||
"""Return capabilities only for the proven installed Anima surface."""
|
||||
|
||||
base_model = patcher.model
|
||||
if type(base_model) is not comfy.model_base.Anima:
|
||||
return None
|
||||
if (
|
||||
type(getattr(base_model, "diffusion_model", None))
|
||||
is not AnimaDiffusionModel
|
||||
):
|
||||
return None
|
||||
if (
|
||||
type(getattr(base_model, "latent_format", None))
|
||||
is not comfy.latent_formats.Wan21
|
||||
):
|
||||
return None
|
||||
return _ANIMA_CAPABILITIES
|
||||
|
||||
|
||||
_ANIMA_CAPABILITIES = RegionalModelCapabilities(
|
||||
model_family=RegionalModelFamily.ANIMA,
|
||||
attention_backend=RegionalAttentionBackend.ANIMA_OBJECT_PATCH,
|
||||
attention_topology=RegionalAttentionTopology.SINGLETON_FRAME_SPATIOTEMPORAL,
|
||||
latent_layout=RegionalLatentLayout.ANIMA_SINGLE_FRAME_BCTHW,
|
||||
spatial_patch_support=RegionalSpatialPatchSupport.FULL_AND_SPATIAL_VIEWS,
|
||||
control_gligen_policy=RegionalControlGligenPolicy.REJECT,
|
||||
reference_latent_policy=RegionalReferenceLatentPolicy.REJECT,
|
||||
known_patch_conflicts=(
|
||||
RegionalPatchConflict.DIFFUSION_MODEL_WRAPPER,
|
||||
RegionalPatchConflict.CROSS_ATTENTION_OBJECT_PATCH,
|
||||
RegionalPatchConflict.ATTN2_INPUT_PATCH,
|
||||
RegionalPatchConflict.ATTN2_OUTPUT_PATCH,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,5 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Provide model-family-specific Attention Coupling runtime adapters."""
|
||||
@@ -0,0 +1,42 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Validate model-consumed Anima cross-attention context geometry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
ANIMA_CONTEXT_SEQUENCE_LENGTH = 512
|
||||
ANIMA_CONTEXT_FEATURE_WIDTH = 1024
|
||||
|
||||
|
||||
class AnimaRegionalContextValidator:
|
||||
"""Require Anima's exact semantic token count and feature width."""
|
||||
|
||||
def validate(
|
||||
self,
|
||||
context: torch.Tensor,
|
||||
*,
|
||||
prompt_type: str,
|
||||
conditioning_index: int,
|
||||
) -> None:
|
||||
"""Reject incompatible Anima context output without token repetition."""
|
||||
|
||||
if int(context.shape[1]) != ANIMA_CONTEXT_SEQUENCE_LENGTH:
|
||||
raise ValueError(
|
||||
f"{prompt_type} conditioning {conditioning_index} c_crossattn "
|
||||
f"must contain exactly {ANIMA_CONTEXT_SEQUENCE_LENGTH} semantic "
|
||||
f"tokens; observed {int(context.shape[1])}. Tokens are not repeated "
|
||||
"to force alignment."
|
||||
)
|
||||
if int(context.shape[2]) != ANIMA_CONTEXT_FEATURE_WIDTH:
|
||||
raise ValueError(
|
||||
f"{prompt_type} conditioning {conditioning_index} c_crossattn "
|
||||
f"must use feature width {ANIMA_CONTEXT_FEATURE_WIDTH}; observed "
|
||||
f"{int(context.shape[2])}."
|
||||
)
|
||||
|
||||
|
||||
ANIMA_REGIONAL_CONTEXT_VALIDATOR = AnimaRegionalContextValidator()
|
||||
@@ -0,0 +1,25 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define the model-family processed-context validation boundary."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class RegionalContextValidator(Protocol):
|
||||
"""Validate one model-consumed cross-attention tensor for a backend."""
|
||||
|
||||
def validate(
|
||||
self,
|
||||
context: torch.Tensor,
|
||||
*,
|
||||
prompt_type: str,
|
||||
conditioning_index: int,
|
||||
) -> None:
|
||||
"""Reject context state incompatible with the selected backend."""
|
||||
@@ -0,0 +1,24 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Retain immutable model-family regional-adapter admission evidence."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ..regional_lora_plan_adapter import RegionalLoraPlanAdaptation
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttentionCouplingFamilyAdmission:
|
||||
"""Retain the exact adaptation admitted before model loading."""
|
||||
|
||||
adaptation: RegionalLoraPlanAdaptation
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require typed immutable adaptation evidence."""
|
||||
|
||||
if not isinstance(self.adaptation, RegionalLoraPlanAdaptation):
|
||||
raise TypeError("Attention Coupling admission requires adaptation.")
|
||||
@@ -0,0 +1,174 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own explicit global-plus-regional standard-UNet attention weighting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StandardUnetAttentionWeights:
|
||||
"""Retain prevalidated standard-UNet tensor weights on one query grid."""
|
||||
|
||||
base: torch.Tensor
|
||||
regions: torch.Tensor
|
||||
denominator: torch.Tensor
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require aligned floating tensor structure without device scalar reads."""
|
||||
|
||||
if not all(
|
||||
isinstance(value, torch.Tensor)
|
||||
for value in (self.base, self.regions, self.denominator)
|
||||
):
|
||||
raise TypeError("Standard UNet attention weights must be tensors.")
|
||||
if not all(
|
||||
value.is_floating_point()
|
||||
for value in (self.base, self.regions, self.denominator)
|
||||
):
|
||||
raise TypeError("Standard UNet attention weights must be floating point.")
|
||||
if self.base.ndim < 1:
|
||||
raise ValueError("Standard UNet base weights require a query grid.")
|
||||
if self.regions.ndim != self.base.ndim + 1 or int(self.regions.shape[0]) < 1:
|
||||
raise ValueError("Standard UNet region weights require a leading region.")
|
||||
if tuple(self.regions.shape[1:]) != tuple(self.base.shape):
|
||||
raise ValueError("Standard UNet region weights must match the base grid.")
|
||||
if self.denominator.shape != self.base.shape:
|
||||
raise ValueError("Standard UNet denominator must match the base grid.")
|
||||
if not (
|
||||
self.base.dtype == self.regions.dtype == self.denominator.dtype
|
||||
and self.base.device == self.regions.device == self.denominator.device
|
||||
):
|
||||
raise ValueError(
|
||||
"Standard UNet attention weights must share dtype and device."
|
||||
)
|
||||
|
||||
@property
|
||||
def normalized_base(self) -> torch.Tensor:
|
||||
"""Return normalized global-branch weights."""
|
||||
|
||||
return self.base / self.denominator
|
||||
|
||||
@property
|
||||
def normalized_regions(self) -> torch.Tensor:
|
||||
"""Return normalized ordered regional weights."""
|
||||
|
||||
return self.regions / self.denominator.unsqueeze(0)
|
||||
|
||||
|
||||
class StandardUnetAttentionWeightingPolicy:
|
||||
"""Normalize one full-canvas global branch with masked regional branches."""
|
||||
|
||||
def weights(
|
||||
self,
|
||||
masks: torch.Tensor,
|
||||
*,
|
||||
region_strengths: tuple[float, ...],
|
||||
epsilon: float = 1e-6,
|
||||
) -> StandardUnetAttentionWeights:
|
||||
"""Return PPM-compatible explicit-base weights for standard UNet."""
|
||||
|
||||
self._validate_inputs(masks, region_strengths, epsilon)
|
||||
strength_shape = (len(region_strengths),) + (1,) * (masks.ndim - 1)
|
||||
strengths = masks.new_tensor(region_strengths).reshape(strength_shape)
|
||||
regions = masks.clamp(0.0, 1.0) * strengths
|
||||
base_strength = max(0.0, 1.0 - max(region_strengths))
|
||||
regional_sum = regions.sum(dim=0)
|
||||
regional_support = regional_sum.ne(0)
|
||||
base = torch.where(
|
||||
regional_support,
|
||||
masks.new_full(masks.shape[1:], base_strength),
|
||||
masks.new_ones(masks.shape[1:]),
|
||||
)
|
||||
denominator = (base + regional_sum).clamp_min(epsilon)
|
||||
return StandardUnetAttentionWeights(base, regions, denominator)
|
||||
|
||||
@staticmethod
|
||||
def blend(
|
||||
*,
|
||||
weights: StandardUnetAttentionWeights,
|
||||
base_output: torch.Tensor,
|
||||
regional_outputs: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Blend prevalidated outputs without synchronizing device values."""
|
||||
|
||||
if not isinstance(weights, StandardUnetAttentionWeights):
|
||||
raise TypeError("Standard UNet blend requires typed weights.")
|
||||
if not isinstance(base_output, torch.Tensor) or not isinstance(
|
||||
regional_outputs, torch.Tensor
|
||||
):
|
||||
raise TypeError("Standard UNet attention outputs must be tensors.")
|
||||
if base_output.dtype != regional_outputs.dtype:
|
||||
raise ValueError("Standard UNet attention output dtypes must match.")
|
||||
if base_output.device != regional_outputs.device:
|
||||
raise ValueError("Standard UNet attention output devices must match.")
|
||||
if regional_outputs.shape != (
|
||||
int(weights.regions.shape[0]),
|
||||
*base_output.shape,
|
||||
):
|
||||
raise ValueError("Standard UNet regional output shape is invalid.")
|
||||
query_dimensions = weights.base.ndim
|
||||
if base_output.ndim < query_dimensions or tuple(
|
||||
base_output.shape[:query_dimensions]
|
||||
) != tuple(weights.base.shape):
|
||||
raise ValueError("Standard UNet output query grid is invalid.")
|
||||
feature_shape = (1,) * (base_output.ndim - query_dimensions)
|
||||
base = weights.base.reshape((*weights.base.shape, *feature_shape)).to(
|
||||
dtype=base_output.dtype
|
||||
)
|
||||
regions = weights.regions.reshape((*weights.regions.shape, *feature_shape)).to(
|
||||
dtype=base_output.dtype
|
||||
)
|
||||
denominator = weights.denominator.reshape(
|
||||
(*weights.denominator.shape, *feature_shape)
|
||||
).to(dtype=base_output.dtype)
|
||||
return (
|
||||
base * base_output + (regions * regional_outputs).sum(dim=0)
|
||||
) / denominator
|
||||
|
||||
@staticmethod
|
||||
def _validate_inputs(
|
||||
masks: object,
|
||||
region_strengths: object,
|
||||
epsilon: object,
|
||||
) -> None:
|
||||
"""Validate host and tensor structure after canonical mask admission."""
|
||||
|
||||
if not isinstance(masks, torch.Tensor):
|
||||
raise TypeError("Standard UNet attention masks must be a tensor.")
|
||||
if (
|
||||
masks.ndim < 2
|
||||
or int(masks.shape[0]) < 1
|
||||
or any(int(size) < 1 for size in masks.shape[1:])
|
||||
):
|
||||
raise ValueError("Standard UNet masks require region and query grids.")
|
||||
if not masks.is_floating_point():
|
||||
raise TypeError("Standard UNet masks must use floating point.")
|
||||
if not isinstance(region_strengths, tuple) or len(region_strengths) != int(
|
||||
masks.shape[0]
|
||||
):
|
||||
raise ValueError("Standard UNet strength count must match the regions.")
|
||||
if any(
|
||||
isinstance(value, bool)
|
||||
or not isinstance(value, int | float)
|
||||
or not math.isfinite(float(value))
|
||||
or float(value) < 0.0
|
||||
for value in region_strengths
|
||||
):
|
||||
raise ValueError("Standard UNet strengths must be finite and non-negative.")
|
||||
if (
|
||||
isinstance(epsilon, bool)
|
||||
or not isinstance(epsilon, int | float)
|
||||
or not math.isfinite(float(epsilon))
|
||||
or float(epsilon) <= 0.0
|
||||
):
|
||||
raise ValueError("Standard UNet epsilon must be finite and positive.")
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_WEIGHTING_POLICY = StandardUnetAttentionWeightingPolicy()
|
||||
@@ -0,0 +1,32 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Validate one completed standard-UNet denoiser result."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class StandardUnetModelOutputValidator:
|
||||
"""Own the single recoverable device-value observation per model call."""
|
||||
|
||||
@staticmethod
|
||||
def validate(output: object, *, model_input: torch.Tensor) -> torch.Tensor:
|
||||
"""Return a finite shape-aligned result or fail after execution."""
|
||||
|
||||
if not isinstance(model_input, torch.Tensor):
|
||||
raise TypeError("Standard UNet model input must be a tensor.")
|
||||
if not isinstance(output, torch.Tensor):
|
||||
raise TypeError("Standard UNet must return a tensor.")
|
||||
if not output.is_floating_point():
|
||||
raise TypeError("Standard UNet output must use floating point.")
|
||||
if output.device != model_input.device:
|
||||
raise ValueError("Standard UNet output must remain on the input device.")
|
||||
if not bool(torch.isfinite(output).all()):
|
||||
raise ValueError("Standard UNet output contains non-finite values.")
|
||||
return output
|
||||
|
||||
|
||||
STANDARD_UNET_MODEL_OUTPUT_VALIDATOR = StandardUnetModelOutputValidator()
|
||||
@@ -0,0 +1,120 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Derive a standard-UNet MODEL with regional attention ownership."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import comfy.model_patcher
|
||||
from comfy.patcher_extension import CallbacksMP
|
||||
|
||||
from ..model_attention_patch_mutations import ModelAttn2PatchesMutation
|
||||
from ..model_patcher_mutations import ModelKeyedCallbackMutation
|
||||
from ..patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from ..regional_lora.operation_assembly import REGIONAL_OPERATION_ASSEMBLER
|
||||
from ..regional_lora.standard_unet_operation_preparation import (
|
||||
StandardUnetOperationAdmission,
|
||||
)
|
||||
from ..regional_lora.standard_unet_operation_session import (
|
||||
StandardUnetRegionalOperationSession,
|
||||
)
|
||||
from .unet_attention_context_wrapper import unet_attention_context_wrapper_mutation
|
||||
from .unet_attention_phase_session import StandardUnetAttentionPhaseSession
|
||||
from .unet_attention_state import StandardUnetAttentionState
|
||||
from .unet_attn2_execution_resolver import StandardUnetAttn2ExecutionResolver
|
||||
from .unet_attn2_patch import UnetAttn2PatchPair
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StandardUnetAttentionModel:
|
||||
"""Return one derived MODEL with its immutable attention execution."""
|
||||
|
||||
model: object
|
||||
state: StandardUnetAttentionState
|
||||
|
||||
|
||||
class StandardUnetAttentionBackend:
|
||||
"""Install regional cross-attention on a collision-safe model clone."""
|
||||
|
||||
def derive(
|
||||
self,
|
||||
*,
|
||||
model: object,
|
||||
state: StandardUnetAttentionState,
|
||||
admission: StandardUnetOperationAdmission,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> StandardUnetAttentionModel:
|
||||
"""Return a direct MODEL child containing only the paired UNet patches."""
|
||||
|
||||
if not isinstance(state, StandardUnetAttentionState):
|
||||
raise TypeError("Standard UNet backend requires attention state.")
|
||||
if not isinstance(admission, StandardUnetOperationAdmission):
|
||||
raise TypeError("Standard UNet backend requires operation admission.")
|
||||
if admission.adaptation.plan != state.plan.lora_plan:
|
||||
raise ValueError(
|
||||
"Standard UNet admission and processed conditioning must share "
|
||||
"the same regional LoRA plan."
|
||||
)
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Standard UNet backend NegPiP state has an invalid type.")
|
||||
attention_phase = StandardUnetAttentionPhaseSession()
|
||||
operation_session: StandardUnetRegionalOperationSession | None = None
|
||||
operation_mutations: tuple[ModelMutation, ...] = ()
|
||||
if admission.adaptation.plan.adapters:
|
||||
if (
|
||||
not isinstance(model, comfy.model_patcher.ModelPatcher)
|
||||
or admission.binding is None
|
||||
or admission.cache is None
|
||||
):
|
||||
raise TypeError(
|
||||
"Standard UNet regional operations require complete MODEL "
|
||||
"admission."
|
||||
)
|
||||
assembly = REGIONAL_OPERATION_ASSEMBLER.assemble(
|
||||
admission.binding,
|
||||
model=model,
|
||||
cache=admission.cache,
|
||||
)
|
||||
operation_session = StandardUnetRegionalOperationSession(
|
||||
admission.adaptation.plan,
|
||||
state.plan.mask_bank,
|
||||
admission.module_roles,
|
||||
assembly.call_scope,
|
||||
)
|
||||
operation_mutations = (
|
||||
assembly.cache_lifecycle.mutation(),
|
||||
ModelKeyedCallbackMutation(
|
||||
CallbacksMP.ON_DETACH,
|
||||
"simple_syrup.standard_unet_regional_operation_schedule",
|
||||
operation_session.clear,
|
||||
),
|
||||
)
|
||||
patches = UnetAttn2PatchPair(
|
||||
StandardUnetAttn2ExecutionResolver(state),
|
||||
operation_scope=operation_session,
|
||||
)
|
||||
derived = PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
(
|
||||
unet_attention_context_wrapper_mutation(
|
||||
state,
|
||||
attention_phase,
|
||||
operation_session,
|
||||
),
|
||||
ModelAttn2PatchesMutation(
|
||||
patches.input_patch,
|
||||
patches.output_patch,
|
||||
(() if negpip is None else (negpip.attention_patch,)),
|
||||
),
|
||||
*operation_mutations,
|
||||
),
|
||||
operation="standard UNet Attention Coupling",
|
||||
)
|
||||
return StandardUnetAttentionModel(derived, state)
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_BACKEND = StandardUnetAttentionBackend()
|
||||
@@ -0,0 +1,131 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Publish active regional contexts around each standard-UNet call."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import ExitStack
|
||||
|
||||
import torch
|
||||
|
||||
from ..diffusion_wrapper_executor import DiffusionWrapperExecutor
|
||||
from ..diffusion_wrapper_invocation import DIFFUSION_WRAPPER_INVOCATION_VALIDATOR
|
||||
from ..model_patcher_mutations import ModelDiffusionWrapperMutation
|
||||
from ..regional_attention_model_call import RegionalAttentionModelCallResolver
|
||||
from ..regional_lora.standard_unet_operation_session import (
|
||||
StandardUnetRegionalOperationSession,
|
||||
)
|
||||
from .standard_unet_model_output_validation import (
|
||||
STANDARD_UNET_MODEL_OUTPUT_VALIDATOR,
|
||||
StandardUnetModelOutputValidator,
|
||||
)
|
||||
from .unet_attention_phase_session import (
|
||||
StandardUnetAttentionPhaseSession,
|
||||
)
|
||||
from .unet_attention_state import StandardUnetAttentionState
|
||||
from .unet_model_call_resolver import STANDARD_UNET_MODEL_CALL_RESOLVER
|
||||
|
||||
UNET_ATTENTION_CONTEXT_WRAPPER_KEY = "simple_syrup.unet_regional_attention_contexts"
|
||||
|
||||
|
||||
class StandardUnetAttentionContextDiffusionWrapper:
|
||||
"""Adapt the installed UNet call contract to shared active context state."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
state: StandardUnetAttentionState,
|
||||
attention_phase: StandardUnetAttentionPhaseSession,
|
||||
operation_session: StandardUnetRegionalOperationSession | None = None,
|
||||
*,
|
||||
model_call_resolver: RegionalAttentionModelCallResolver = (
|
||||
STANDARD_UNET_MODEL_CALL_RESOLVER
|
||||
),
|
||||
output_validator: StandardUnetModelOutputValidator = (
|
||||
STANDARD_UNET_MODEL_OUTPUT_VALIDATOR
|
||||
),
|
||||
) -> None:
|
||||
"""Retain the shared plan and model-call authorities."""
|
||||
|
||||
if not isinstance(state, StandardUnetAttentionState):
|
||||
raise TypeError("Standard UNet context wrapper requires attention state.")
|
||||
if not isinstance(attention_phase, StandardUnetAttentionPhaseSession):
|
||||
raise TypeError("Standard UNet context wrapper requires phase state.")
|
||||
if operation_session is not None and not isinstance(
|
||||
operation_session,
|
||||
StandardUnetRegionalOperationSession,
|
||||
):
|
||||
raise TypeError("Standard UNet operation session has an invalid type.")
|
||||
if not isinstance(model_call_resolver, RegionalAttentionModelCallResolver):
|
||||
raise TypeError(
|
||||
"Standard UNet context wrapper requires a model-call resolver."
|
||||
)
|
||||
self._state = state
|
||||
self._attention_phase = attention_phase
|
||||
self._operation_session = operation_session
|
||||
self._model_call_resolver = model_call_resolver
|
||||
if not isinstance(output_validator, StandardUnetModelOutputValidator):
|
||||
raise TypeError("Standard UNet output validator has an invalid type.")
|
||||
self._output_validator = output_validator
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
executor: DiffusionWrapperExecutor,
|
||||
*args: object,
|
||||
**kwargs: object,
|
||||
) -> object:
|
||||
"""Resolve and publish contexts for exactly one nested UNet execution."""
|
||||
|
||||
if not args or not isinstance(args[0], torch.Tensor):
|
||||
raise TypeError("Standard UNet context wrapper requires model input.")
|
||||
if len(args) < 3 or not isinstance(args[2], torch.Tensor):
|
||||
raise TypeError(
|
||||
"Standard UNet context wrapper requires a tensor third positional "
|
||||
"base context."
|
||||
)
|
||||
if len(args) < 6 or not isinstance(args[5], dict):
|
||||
raise TypeError(
|
||||
"Standard UNet context wrapper requires dictionary sixth positional "
|
||||
"transformer options."
|
||||
)
|
||||
DIFFUSION_WRAPPER_INVOCATION_VALIDATOR.require_owned(
|
||||
executor,
|
||||
args[5],
|
||||
key=UNET_ATTENTION_CONTEXT_WRAPPER_KEY,
|
||||
wrapper=self,
|
||||
)
|
||||
contexts = self._model_call_resolver.resolve(
|
||||
self._state.plan,
|
||||
model_input=args[0],
|
||||
base_context=args[2],
|
||||
transformer_options=args[5],
|
||||
)
|
||||
forwarded_args = (*args[:2], contexts.base_context, *args[3:])
|
||||
with ExitStack() as scopes:
|
||||
scopes.enter_context(self._attention_phase.activate(args[5]))
|
||||
scopes.enter_context(self._state.execution_context.activate(contexts))
|
||||
scopes.enter_context(self._state.resolution_cache.activate())
|
||||
if self._operation_session is not None:
|
||||
scopes.enter_context(
|
||||
self._operation_session.activate(contexts, args[5])
|
||||
)
|
||||
output = executor(*forwarded_args, **kwargs)
|
||||
return self._output_validator.validate(output, model_input=args[0])
|
||||
|
||||
|
||||
def unet_attention_context_wrapper_mutation(
|
||||
state: StandardUnetAttentionState,
|
||||
attention_phase: StandardUnetAttentionPhaseSession,
|
||||
operation_session: StandardUnetRegionalOperationSession | None = None,
|
||||
) -> ModelDiffusionWrapperMutation:
|
||||
"""Return the clone-local standard-UNet context wrapper mutation."""
|
||||
|
||||
return ModelDiffusionWrapperMutation(
|
||||
UNET_ATTENTION_CONTEXT_WRAPPER_KEY,
|
||||
StandardUnetAttentionContextDiffusionWrapper(
|
||||
state,
|
||||
attention_phase,
|
||||
operation_session,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,60 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Emit structured standard-UNet regional resolution diagnostics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from ...shared.logging import get_logger
|
||||
from ..regional_attention_diagnostic_values import (
|
||||
RegionalAttentionExecutionDiagnostics,
|
||||
)
|
||||
from .unet_attention_geometry import StandardUnetAttentionGeometry
|
||||
|
||||
LOGGER = get_logger("runtime.attention_coupling.unet_diagnostics")
|
||||
|
||||
|
||||
class StandardUnetAttentionDiagnosticsEmitter:
|
||||
"""Publish one JSON-safe record for each unique call-local resolution."""
|
||||
|
||||
def __init__(self, logger: logging.Logger = LOGGER) -> None:
|
||||
"""Retain the sole standard-UNet diagnostics logging boundary."""
|
||||
|
||||
if not isinstance(logger, logging.Logger):
|
||||
raise TypeError("Standard UNet diagnostics require a logger.")
|
||||
self._logger = logger
|
||||
|
||||
def emit(
|
||||
self,
|
||||
snapshot: RegionalAttentionExecutionDiagnostics,
|
||||
geometry: StandardUnetAttentionGeometry,
|
||||
) -> None:
|
||||
"""Emit one resolution snapshot without model inputs or tensor values."""
|
||||
|
||||
if not isinstance(snapshot, RegionalAttentionExecutionDiagnostics):
|
||||
raise TypeError("Standard UNet diagnostics require a shared snapshot.")
|
||||
if not isinstance(geometry, StandardUnetAttentionGeometry):
|
||||
raise TypeError("Standard UNet diagnostics require validated geometry.")
|
||||
if not self._logger.isEnabledFor(logging.DEBUG):
|
||||
return
|
||||
self._logger.debug(
|
||||
"Standard UNet regional Attention Coupling resolution",
|
||||
extra={
|
||||
"operation": "unet_attention_coupling.resolve",
|
||||
"unet_layer": {
|
||||
"block_kind": geometry.block[0],
|
||||
"block_number": geometry.block[1],
|
||||
"block_index": geometry.block_index,
|
||||
"transformer_index": geometry.transformer_index,
|
||||
"query_height": geometry.query.query_height,
|
||||
"query_width": geometry.query.query_width,
|
||||
},
|
||||
"regional_diagnostics": snapshot.to_log_fields(),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_DIAGNOSTICS_EMITTER = StandardUnetAttentionDiagnosticsEmitter()
|
||||
@@ -0,0 +1,243 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve installed standard-UNet attention metadata to query geometry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from ...domain.regional_attention_batch import BatchedRegionalAttentionContexts
|
||||
from ...domain.regional_mask_bank import RegionalMaskBank
|
||||
from ...domain.spatial_views import (
|
||||
SpatialBatchLayout,
|
||||
SpatialView,
|
||||
SpatialViewKind,
|
||||
)
|
||||
from ..regional_attention_diagnostic_values import RegionalAttentionQueryGeometry
|
||||
from ..spatial_model_arguments import (
|
||||
SIMPLE_SYRUP_TRANSFORMER_NAMESPACE,
|
||||
SPATIAL_BATCH_LAYOUT_KEY,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StandardUnetAttentionGeometry:
|
||||
"""Retain one validated UNet layer identity and spatial query contract."""
|
||||
|
||||
query: RegionalAttentionQueryGeometry
|
||||
layout: SpatialBatchLayout
|
||||
original_height: int
|
||||
original_width: int
|
||||
block: tuple[str, int]
|
||||
block_index: int
|
||||
transformer_index: int
|
||||
|
||||
|
||||
class StandardUnetAttentionGeometryResolver:
|
||||
"""Parse only installed Comfy callback metadata and layout identity."""
|
||||
|
||||
def resolve(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
contexts: BatchedRegionalAttentionContexts,
|
||||
mask_bank: RegionalMaskBank,
|
||||
extra_options: dict[str, Any],
|
||||
) -> StandardUnetAttentionGeometry:
|
||||
"""Return exact rectangular query and model-call spatial geometry."""
|
||||
|
||||
if not isinstance(query, torch.Tensor) or query.ndim != 3:
|
||||
raise ValueError("Standard UNet attention query must use BxQxD layout.")
|
||||
if not isinstance(contexts, BatchedRegionalAttentionContexts):
|
||||
raise TypeError("Standard UNet geometry requires aligned contexts.")
|
||||
if not isinstance(mask_bank, RegionalMaskBank):
|
||||
raise TypeError("Standard UNet geometry requires a mask bank.")
|
||||
if not isinstance(extra_options, dict):
|
||||
raise TypeError(
|
||||
"Standard UNet attention extra_options must be a dictionary."
|
||||
)
|
||||
activations = self._shape(
|
||||
extra_options.get("activations_shape"),
|
||||
name="activations_shape",
|
||||
)
|
||||
original = self._shape(
|
||||
extra_options.get("original_shape"),
|
||||
name="original_shape",
|
||||
)
|
||||
batch, _, query_height, query_width = activations
|
||||
if batch != int(query.shape[0]) or batch != int(contexts.base_context.shape[0]):
|
||||
raise ValueError(
|
||||
"Standard UNet attention activation batch must match query and "
|
||||
"contexts."
|
||||
)
|
||||
if int(query.shape[1]) != query_height * query_width:
|
||||
raise ValueError(
|
||||
"Standard UNet attention token count must match activations H/W."
|
||||
)
|
||||
if original[0] != batch:
|
||||
raise ValueError(
|
||||
"Standard UNet original batch must match the active query batch."
|
||||
)
|
||||
block = self._block(extra_options.get("block"))
|
||||
block_index = self._nonnegative_index(
|
||||
extra_options.get("block_index"),
|
||||
name="block_index",
|
||||
)
|
||||
transformer_index = self._nonnegative_index(
|
||||
extra_options.get("transformer_index"),
|
||||
name="transformer_index",
|
||||
)
|
||||
published_layout = self._published_layout(extra_options)
|
||||
layout = published_layout or self._full_layout(
|
||||
mask_bank,
|
||||
contexts,
|
||||
original_height=original[2],
|
||||
original_width=original[3],
|
||||
)
|
||||
if published_layout is not None:
|
||||
self._validate_published_layout(
|
||||
published_layout,
|
||||
query_batch=batch,
|
||||
original_height=original[2],
|
||||
original_width=original[3],
|
||||
)
|
||||
return StandardUnetAttentionGeometry(
|
||||
query=RegionalAttentionQueryGeometry(
|
||||
batch,
|
||||
1,
|
||||
query_height,
|
||||
query_width,
|
||||
published_layout,
|
||||
),
|
||||
layout=layout,
|
||||
original_height=original[2],
|
||||
original_width=original[3],
|
||||
block=block,
|
||||
block_index=block_index,
|
||||
transformer_index=transformer_index,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _shape(value: object, *, name: str) -> tuple[int, int, int, int]:
|
||||
"""Narrow one installed BCHW metadata sequence."""
|
||||
|
||||
if not isinstance(value, list | tuple) or len(value) != 4:
|
||||
raise TypeError(f"Standard UNet {name} must be a BCHW sequence.")
|
||||
values: list[int] = []
|
||||
for item in value:
|
||||
if isinstance(item, bool) or not isinstance(item, int) or item < 1:
|
||||
raise ValueError(
|
||||
f"Standard UNet {name} dimensions must be positive integers."
|
||||
)
|
||||
values.append(item)
|
||||
return values[0], values[1], values[2], values[3]
|
||||
|
||||
@staticmethod
|
||||
def _block(value: object) -> tuple[str, int]:
|
||||
"""Narrow installed input/middle/output block metadata."""
|
||||
|
||||
if not isinstance(value, list | tuple) or len(value) != 2:
|
||||
raise TypeError("Standard UNet block must contain kind and index.")
|
||||
kind, index = value
|
||||
if kind not in {"input", "middle", "output"}:
|
||||
raise ValueError("Standard UNet block kind is unsupported.")
|
||||
if isinstance(index, bool) or not isinstance(index, int) or index < 0:
|
||||
raise ValueError("Standard UNet block index must be non-negative.")
|
||||
return kind, index
|
||||
|
||||
@staticmethod
|
||||
def _nonnegative_index(value: object, *, name: str) -> int:
|
||||
"""Narrow one installed non-negative integer index."""
|
||||
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"Standard UNet {name} must be an integer.")
|
||||
if value < 0:
|
||||
raise ValueError(f"Standard UNet {name} must be non-negative.")
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _published_layout(
|
||||
extra_options: dict[str, Any],
|
||||
) -> SpatialBatchLayout | None:
|
||||
"""Return the exact optional SimpleSyrup spatial layout."""
|
||||
|
||||
namespace = extra_options.get(SIMPLE_SYRUP_TRANSFORMER_NAMESPACE)
|
||||
if namespace is None:
|
||||
return None
|
||||
if not isinstance(namespace, dict):
|
||||
raise TypeError("Standard UNet SimpleSyrup namespace must be a dictionary.")
|
||||
layout = namespace.get(SPATIAL_BATCH_LAYOUT_KEY)
|
||||
if layout is None:
|
||||
return None
|
||||
if not isinstance(layout, SpatialBatchLayout):
|
||||
raise TypeError(
|
||||
"Standard UNet spatial layout must be a SpatialBatchLayout."
|
||||
)
|
||||
return layout
|
||||
|
||||
@staticmethod
|
||||
def _full_layout(
|
||||
mask_bank: RegionalMaskBank,
|
||||
contexts: BatchedRegionalAttentionContexts,
|
||||
*,
|
||||
original_height: int,
|
||||
original_width: int,
|
||||
) -> SpatialBatchLayout:
|
||||
"""Represent an unmodified full-canvas UNet model call."""
|
||||
|
||||
if (original_height, original_width) != (
|
||||
mask_bank.canvas_height,
|
||||
mask_bank.canvas_width,
|
||||
):
|
||||
raise ValueError(
|
||||
"Full-context UNet original H/W must match the regional mask canvas."
|
||||
)
|
||||
return SpatialBatchLayout(
|
||||
mask_bank.canvas_width,
|
||||
mask_bank.canvas_height,
|
||||
(
|
||||
SpatialView(
|
||||
SpatialViewKind.FULL,
|
||||
0,
|
||||
0,
|
||||
mask_bank.canvas_width,
|
||||
mask_bank.canvas_height,
|
||||
mask_bank.canvas_width,
|
||||
mask_bank.canvas_height,
|
||||
),
|
||||
),
|
||||
contexts.latent_batch_size,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_published_layout(
|
||||
layout: SpatialBatchLayout,
|
||||
*,
|
||||
query_batch: int,
|
||||
original_height: int,
|
||||
original_width: int,
|
||||
) -> None:
|
||||
"""Require published view batches to match the active UNet tensor."""
|
||||
|
||||
if layout.expanded_batch_size != query_batch:
|
||||
raise ValueError(
|
||||
"Standard UNet spatial layout batch must match the active query."
|
||||
)
|
||||
mismatched = tuple(
|
||||
index
|
||||
for index, view in enumerate(layout.views)
|
||||
if (view.model_height, view.model_width)
|
||||
!= (original_height, original_width)
|
||||
)
|
||||
if mismatched:
|
||||
raise ValueError(
|
||||
"Standard UNet original H/W must match every layout model view; "
|
||||
f"mismatched view indices {mismatched}."
|
||||
)
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_GEOMETRY_RESOLVER = StandardUnetAttentionGeometryResolver()
|
||||
@@ -0,0 +1,40 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Describe one standard-UNet attention-composition phase."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class StandardUnetAttentionStage(Enum):
|
||||
"""Identify the active standard-UNet attention responsibility."""
|
||||
|
||||
COMPOSITION = "composition"
|
||||
SPECIALIZATION = "specialization"
|
||||
CONSOLIDATION = "consolidation"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StandardUnetAttentionPhase:
|
||||
"""Publish one normalized attention stage and its local progress."""
|
||||
|
||||
stage: StandardUnetAttentionStage
|
||||
denoising_progress: float
|
||||
stage_progress: float
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require finite normalized phase values and explicit ownership."""
|
||||
|
||||
if not isinstance(self.stage, StandardUnetAttentionStage):
|
||||
raise TypeError("Standard UNet attention stage is invalid.")
|
||||
for name, value in (
|
||||
("denoising progress", self.denoising_progress),
|
||||
("stage progress", self.stage_progress),
|
||||
):
|
||||
if not math.isfinite(value) or not 0.0 <= value <= 1.0:
|
||||
raise ValueError(f"Standard UNet attention {name} must be in [0, 1].")
|
||||
@@ -0,0 +1,68 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Emit structured standard-UNet attention-phase diagnostics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ...shared.logging import get_logger
|
||||
from .unet_attention_phase import StandardUnetAttentionPhase
|
||||
|
||||
LOGGER = get_logger("runtime.attention_coupling.unet_attention_phase")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StandardUnetAttentionPhaseDiagnostic:
|
||||
"""Describe one exact model-call attention phase."""
|
||||
|
||||
phase: StandardUnetAttentionPhase
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require the authoritative phase value."""
|
||||
|
||||
if not isinstance(self.phase, StandardUnetAttentionPhase):
|
||||
raise TypeError("Standard UNet attention phase diagnostic is invalid.")
|
||||
|
||||
def to_log_fields(self) -> dict[str, object]:
|
||||
"""Return JSON-safe phase fields for managed evidence."""
|
||||
|
||||
return {
|
||||
"stage": self.phase.stage.value,
|
||||
"denoising_progress": self.phase.denoising_progress,
|
||||
"stage_progress": self.phase.stage_progress,
|
||||
}
|
||||
|
||||
|
||||
class StandardUnetAttentionPhaseDiagnosticsEmitter:
|
||||
"""Publish one structured phase record for every model call."""
|
||||
|
||||
def __init__(self, logger: logging.Logger = LOGGER) -> None:
|
||||
"""Retain the focused phase diagnostics logging boundary."""
|
||||
|
||||
if not isinstance(logger, logging.Logger):
|
||||
raise TypeError("Standard UNet attention phase logger is invalid.")
|
||||
self._logger = logger
|
||||
|
||||
def emit(self, diagnostic: StandardUnetAttentionPhaseDiagnostic) -> None:
|
||||
"""Emit the phase without model inputs or tensor values."""
|
||||
|
||||
if not isinstance(diagnostic, StandardUnetAttentionPhaseDiagnostic):
|
||||
raise TypeError("Standard UNet attention phase diagnostic is invalid.")
|
||||
if not self._logger.isEnabledFor(logging.DEBUG):
|
||||
return
|
||||
self._logger.debug(
|
||||
"Standard UNet attention phase",
|
||||
extra={
|
||||
"operation": "unet_attention_coupling.phase",
|
||||
"attention_phase_diagnostics": diagnostic.to_log_fields(),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_PHASE_DIAGNOSTICS_EMITTER = (
|
||||
StandardUnetAttentionPhaseDiagnosticsEmitter()
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve standard-UNet attention composition across denoising."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
from ..denoising_progress import DENOISING_PROGRESS_RESOLVER
|
||||
from ..regional_attention_model_call_values import uniform_model_call_sigma
|
||||
from .unet_attention_phase import (
|
||||
StandardUnetAttentionPhase,
|
||||
StandardUnetAttentionStage,
|
||||
)
|
||||
|
||||
|
||||
class StandardUnetAttentionPhaseSchedule:
|
||||
"""Coordinate shared layout, regional ownership, and final cohesion."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
composition_fraction: float = 0.1,
|
||||
specialization_fraction: float = 0.55,
|
||||
) -> None:
|
||||
"""Set validated stage boundaries for one denoising trajectory."""
|
||||
|
||||
values = (composition_fraction, specialization_fraction)
|
||||
if any(
|
||||
isinstance(value, bool)
|
||||
or not isinstance(value, int | float)
|
||||
or not math.isfinite(float(value))
|
||||
for value in values
|
||||
):
|
||||
raise TypeError(
|
||||
"Standard UNet attention schedule values must be finite numbers."
|
||||
)
|
||||
self._composition_end = float(composition_fraction)
|
||||
self._specialization_end = self._composition_end + float(
|
||||
specialization_fraction
|
||||
)
|
||||
if not 0.0 < self._composition_end < self._specialization_end < 1.0:
|
||||
raise ValueError(
|
||||
"Standard UNet attention stages must fit within denoising."
|
||||
)
|
||||
|
||||
def resolve(
|
||||
self,
|
||||
transformer_options: dict[str, object],
|
||||
) -> StandardUnetAttentionPhase:
|
||||
"""Resolve one attention phase from exact Comfy sampling metadata."""
|
||||
|
||||
if not isinstance(transformer_options, dict):
|
||||
raise TypeError("Standard UNet phase options must be a dictionary.")
|
||||
sample_sigmas = transformer_options.get("sample_sigmas")
|
||||
current_sigmas = transformer_options.get("sigmas")
|
||||
if not isinstance(sample_sigmas, torch.Tensor):
|
||||
raise TypeError("Standard UNet phase requires tensor sample_sigmas.")
|
||||
if not isinstance(current_sigmas, torch.Tensor):
|
||||
raise TypeError("Standard UNet phase requires tensor sigmas.")
|
||||
progress = DENOISING_PROGRESS_RESOLVER.resolve(
|
||||
sample_sigmas,
|
||||
uniform_model_call_sigma(current_sigmas),
|
||||
)
|
||||
if progress < self._composition_end:
|
||||
return StandardUnetAttentionPhase(
|
||||
StandardUnetAttentionStage.COMPOSITION,
|
||||
progress,
|
||||
progress / self._composition_end,
|
||||
)
|
||||
if progress < self._specialization_end:
|
||||
return StandardUnetAttentionPhase(
|
||||
StandardUnetAttentionStage.SPECIALIZATION,
|
||||
progress,
|
||||
(progress - self._composition_end)
|
||||
/ (self._specialization_end - self._composition_end),
|
||||
)
|
||||
return StandardUnetAttentionPhase(
|
||||
StandardUnetAttentionStage.CONSOLIDATION,
|
||||
progress,
|
||||
(progress - self._specialization_end) / (1.0 - self._specialization_end),
|
||||
)
|
||||
|
||||
|
||||
STANDARD_UNET_ATTENTION_PHASE_SCHEDULE = StandardUnetAttentionPhaseSchedule()
|
||||