Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bc5b613490 | ||
|
|
01107c45f8 | ||
|
|
96185132b8 | ||
|
|
d4a31bf19c | ||
|
|
fc7ba084f6 | ||
|
|
4516aa9cb4 | ||
|
|
0e48aaa3e4 | ||
|
|
78946b0fa3 | ||
|
|
c30408f96d | ||
|
|
eb7cf89f17 | ||
|
|
af42134028 | ||
|
|
a85e57b18c | ||
|
|
22fce6fdda | ||
|
|
147edcfcbc | ||
|
|
8bf3545fec | ||
|
|
f47149746a | ||
|
|
83cfc5c723 | ||
|
|
d87e52ea2c | ||
|
|
9405784764 | ||
|
|
55261bda7c | ||
|
|
cf7a9c41e8 | ||
|
|
3a25526e81 | ||
|
|
00173fa3fb | ||
|
|
0d264b90a7 | ||
|
|
a8cf4650ff | ||
|
|
edcb3da08b | ||
|
|
7f7a62f832 | ||
|
|
fc908ba0a5 | ||
|
|
ead4b34e6d | ||
|
|
46af6027d6 | ||
|
|
b7ca8ed1c6 | ||
|
|
4aad5c3b9d | ||
|
|
d61da30409 | ||
|
|
6851da6638 | ||
|
|
0eeb707f34 | ||
|
|
9a943714aa | ||
|
|
bae26a07fb | ||
|
|
c92d99a8a3 | ||
|
|
3f6d082940 | ||
|
|
58ae89f8e0 | ||
|
|
c9a26427a8 | ||
|
|
6608c0b6d1 | ||
|
|
a757e1c98b | ||
|
|
52bd76e19c | ||
|
|
d6e004cce2 | ||
|
|
ed17fa2ef4 | ||
|
|
827c64c43d | ||
|
|
e5482aee5e | ||
|
|
62469a4dd9 | ||
|
|
8c629bee18 | ||
|
|
50cb6f5ed6 | ||
|
|
e32d1e02df | ||
|
|
b0d52f7305 | ||
|
|
e17c6e29f5 | ||
|
|
27e03fa23e | ||
|
|
ec1cb1ac17 | ||
|
|
64634104a2 | ||
|
|
ecbb220de6 | ||
|
|
cd9e614b1a | ||
|
|
9ccf572a15 | ||
|
|
74af5c6499 | ||
|
|
caf0b39d8a | ||
|
|
e099d581a7 | ||
|
|
22f7c30373 | ||
|
|
0133fb93bc | ||
|
|
cf7d30507e | ||
|
|
b6fa571fd2 | ||
|
|
f272526bfc | ||
|
|
4e593bb30b | ||
|
|
097ca33b8e | ||
|
|
784fb0145b | ||
|
|
dbcca15a21 | ||
|
|
bc41576fac | ||
|
|
8596b8184e | ||
|
|
896a025006 | ||
|
|
43092e44a4 | ||
|
|
80b5a0ca74 | ||
|
|
81b3bc1651 | ||
|
|
a825504bdd | ||
|
|
22190cd25e | ||
|
|
a976adbb39 | ||
|
|
997d2fb13a | ||
|
|
f8829fcb37 | ||
|
|
9651a70341 | ||
|
|
57683c3c7d | ||
|
|
f99f92e8f7 | ||
|
|
5bc125d2f0 | ||
|
|
c99b0812ab | ||
|
|
333f646ab1 | ||
|
|
dbdf27664c | ||
|
|
7d5569e5c1 | ||
|
|
5681b464ad | ||
|
|
8d0fcee2f3 | ||
|
|
1078fc6f0f | ||
|
|
821a0ef427 | ||
|
|
9007a70aa0 | ||
|
|
1a0ebd5173 | ||
|
|
59608320c8 | ||
|
|
d64fac4b74 | ||
|
|
d687497d80 | ||
|
|
d6343e1860 | ||
|
|
4eebdd8b8b | ||
|
|
372e035686 | ||
|
|
fb34671ee6 | ||
|
|
f25f6bdcd1 | ||
|
|
f1b484617a | ||
|
|
4507842a70 |
@@ -2,19 +2,23 @@ name: 📦 Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
tags:
|
||||
- '*'
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'melMass' }}
|
||||
steps:
|
||||
- name: ♻️ Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
- name: 📦 Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
|
||||
|
||||
@@ -6,3 +6,6 @@ node_modules/
|
||||
compose.yaml
|
||||
comfy_mtb.wsb
|
||||
Dockerfile
|
||||
|
||||
# I store the gh-pages worktrees (src & build) there
|
||||
.worktrees
|
||||
|
||||
+238
-2
@@ -3,10 +3,193 @@
|
||||
This is an automated changelog based on the commits in this repository.
|
||||
|
||||
Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases) for more information.
|
||||
## [main] - 2024-03-07
|
||||
## [main] - 2025-04-16
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- 🐛 note+ breaking wfs ([af42134](https://github.com/melMass/comfy_mtb/commit/af421340286b234e4c0cfcd4143a9d8726ebf3d1))
|
||||
- 🐛 ColorCorrect clamp issue ([a85e57b](https://github.com/melMass/comfy_mtb/commit/a85e57b18c7d3c765131873ffff523244ca9be73))
|
||||
- 🐛 Whisper chunks processing ([8bf3545](https://github.com/melMass/comfy_mtb/commit/8bf3545fec5b2a180607d40394b025a1e09c14b6))
|
||||
- 🐛 stackImages move to device ([d87e52e](https://github.com/melMass/comfy_mtb/commit/d87e52ea2c112fd95f257dcd6a54a5db77a34fc3))
|
||||
- 🐛 bbox upscale from center ([55261bd](https://github.com/melMass/comfy_mtb/commit/55261bda7c33d088b62c5483e4483201e5a9ce77))
|
||||
- 🐛 add MASK support for PickFromBatch ([0d264b9](https://github.com/melMass/comfy_mtb/commit/0d264b90a78d5a6719fb3ce71f4e9a642db4c950))
|
||||
- 🐛 use addDOMWidget for Debug node ([46af602](https://github.com/melMass/comfy_mtb/commit/46af6027d6c87d0c29b8bb0fd1cc1dbdae993629))
|
||||
- 🐛 use "modern" notation in toDevice ([b7ca8ed](https://github.com/melMass/comfy_mtb/commit/b7ca8ed1c6e117b71afd7696f55dcc3dbd5bad08))
|
||||
- 🐛 handle missing submodules ([d61da30](https://github.com/melMass/comfy_mtb/commit/d61da304099ff5e4528e4beb1ecc2eb83cabaaa1))
|
||||
- 🐛 add warnings about what each IO mode can do ([6608c0b](https://github.com/melMass/comfy_mtb/commit/6608c0b6d1cf8f7a9901214096f8c78bfe17056f))
|
||||
- 🐛 soft deprecate compression h264 ([a757e1c](https://github.com/melMass/comfy_mtb/commit/a757e1c98b2abbd2221a15b77e89d772e02d1d82))
|
||||
- 🐛 limit packages allowed to be installed from API ([d6e004c](https://github.com/melMass/comfy_mtb/commit/d6e004cce2c32f8e48b868e66b89f82da4887dc3))
|
||||
- 🐛 ensure default settings (io sidebar) ([ed17fa2](https://github.com/melMass/comfy_mtb/commit/ed17fa2ef4688aadf305a6d51b32c13a0efd22d6))
|
||||
- 🐛 spawn colour picker at pointer location ([e5482ae](https://github.com/melMass/comfy_mtb/commit/e5482aee5e3de07e8f055b3edc0fccc0e0f75c14)) by [@webfiltered](https://github.com/webfiltered) in [#223](https://github.com/melMass/comfy_mtb/pull/223)
|
||||
- 🐛 i/o sidebar for custom paths ([62469a4](https://github.com/melMass/comfy_mtb/commit/62469a4dd96e32509171aad74fcae8d2bb0ec593))
|
||||
|
||||
### Features
|
||||
|
||||
- ⚡ add BatchFromFolder ([9618513](https://github.com/melMass/comfy_mtb/commit/96185132b83c182032e9f6e822561eb5699af517))
|
||||
- ⚡ add use_normalized to TransformBatch2D ([d4a31bf](https://github.com/melMass/comfy_mtb/commit/d4a31bf19c2863df8dfc4cb9a3cd6683304949e4))
|
||||
- [**breaking**] ⚡ add support for masks in BatchFLoatMath ([fc7ba08](https://github.com/melMass/comfy_mtb/commit/fc7ba084f6ed7880e88e28eb448ab0bd7d796824))
|
||||
- ✨ add use_normalized to TransformImage ([4516aa9](https://github.com/melMass/comfy_mtb/commit/4516aa9cb4fcb12c946999d6dcc1501cc09011a3))
|
||||
- ✨ add regex support for String Replace ([78946b0](https://github.com/melMass/comfy_mtb/commit/78946b0fa3c3cf5dfcee8c7c4c0921b722d09d1e)) by [@poetryiii](https://github.com/poetryiii) in [#233](https://github.com/melMass/comfy_mtb/pull/233)
|
||||
- ✨ update diarization to 3.1 ([c30408f](https://github.com/melMass/comfy_mtb/commit/c30408f96d4df9c7d35545654401162090a74305)) by [@numz](https://github.com/numz)
|
||||
- ✨ add "workflow" query to /mtb/view endpoint ([eb7cf89](https://github.com/melMass/comfy_mtb/commit/eb7cf89f173b2342b04e7b61dca3d12cfaf65bdb))
|
||||
- ✨ add stretch_x and stretch_y to TransformImage ([22fce6f](https://github.com/melMass/comfy_mtb/commit/22fce6fdda135cbb1f1aad42c86aae166cba81b5))
|
||||
- ✨ add AudioDuration node ([f471497](https://github.com/melMass/comfy_mtb/commit/f47149746ac1e418cda2007c38aafbb03946ce22))
|
||||
- ✨ basic whisper nodes ([83cfc5c](https://github.com/melMass/comfy_mtb/commit/83cfc5c723d1a572af67ad14b52be4f8371a3c5f))
|
||||
- ✨ add BboxForDimensions ([9405784](https://github.com/melMass/comfy_mtb/commit/940578476438eaa6a42e0056f1b7b319ee585334))
|
||||
- ✨ improve the debug node ([cf7a9c4](https://github.com/melMass/comfy_mtb/commit/cf7a9c41e81e8dd461ab9dfa3c05bb8e2cdf2a67))
|
||||
- ✨ add BatchImageToSublist and counterpart ([00173fa](https://github.com/melMass/comfy_mtb/commit/00173fa3fbca4c5b1ff3016cc5139705ce61ec20))
|
||||
- ✨ add TensorOps ([a8cf465](https://github.com/melMass/comfy_mtb/commit/a8cf4650ff5cbd4975ef954b5829c772ee53250c))
|
||||
- ✨ live update outputs grid ([7f7a62f](https://github.com/melMass/comfy_mtb/commit/7f7a62f832c865a13b9181daee79d3cfc21581e2)) by [@christian-byrne](https://github.com/christian-byrne) in [#229](https://github.com/melMass/comfy_mtb/pull/229)
|
||||
- ✨ add SaveImage passthrough ([0eeb707](https://github.com/melMass/comfy_mtb/commit/0eeb707f34f51142def8e0ef7d351ee5028cb5e0))
|
||||
- ✨ add filtering to TransformImage ([bae26a0](https://github.com/melMass/comfy_mtb/commit/bae26a07fb02dd518c621eba28986a51c5d086bc))
|
||||
- ✨ add support for video in I/O sidebar ([c92d99a](https://github.com/melMass/comfy_mtb/commit/c92d99a8a37a64cfc285296f21452c4927a22774))
|
||||
- ✨ add an extra static input to Stack Images ([3f6d082](https://github.com/melMass/comfy_mtb/commit/3f6d08294096918d50101a19083f9134305cc8c9)) in [#222](https://github.com/melMass/comfy_mtb/pull/222)
|
||||
- ✨ add support for subdirs (i/o sidebar) ([52bd76e](https://github.com/melMass/comfy_mtb/commit/52bd76e19c8bd7e72986900e5dbfade0457ef7e0))
|
||||
- ✨ add Batch Sequence Nodes ([827c64c](https://github.com/melMass/comfy_mtb/commit/827c64c43d52ebfb8acd2e5c4491c4b66e6b8f40))
|
||||
- ✨ add support for more formats (I/O sidebar) ([8c629be](https://github.com/melMass/comfy_mtb/commit/8c629bee186b5ac991058018a788e4a836eef630))
|
||||
|
||||
### Miscellaneous Tasks
|
||||
|
||||
- 🧹 bump version ([d093d76](https://github.com/melMass/comfy_mtb/commit/d093d76efd87474a3ca82858147255038060ab17))
|
||||
- 🧹 small adjustments ([01107c4](https://github.com/melMass/comfy_mtb/commit/01107c45f8539ff7c579e08e2a9075d93781b9a2))
|
||||
- 🤖 update publish action workflow with permissions and version constraints ([0e48aaa](https://github.com/melMass/comfy_mtb/commit/0e48aaa3e4f1e440a5d7ab42df56b728ced03aca)) by [@robinjhuang](https://github.com/robinjhuang) in [#237](https://github.com/melMass/comfy_mtb/pull/237)
|
||||
- 🧹 basic standalone detection ([3a25526](https://github.com/melMass/comfy_mtb/commit/3a25526e818a1af8f886d2ad5c27101c4a0caa8b))
|
||||
- 🧹 rename type ([edcb3da](https://github.com/melMass/comfy_mtb/commit/edcb3da08bff66f9adcef8dcd37c3925e64d0135))
|
||||
- 🧹 update env file ([fc908ba](https://github.com/melMass/comfy_mtb/commit/fc908ba0a528523b7c1e37e34fb32f430746de0d))
|
||||
- 🧹 dev ([9a94371](https://github.com/melMass/comfy_mtb/commit/9a943714aada107bfd236e00fa1063872db7a834))
|
||||
- 🧹 apply formatting ([58ae89f](https://github.com/melMass/comfy_mtb/commit/58ae89f8e0f0f8b42825722a6aebc04da39847b1))
|
||||
|
||||
### Refactor
|
||||
|
||||
- 📦 add model autodownload ([147edcf](https://github.com/melMass/comfy_mtb/commit/147edcfcbc09dd27a0c787f9da568fb850c3308a))
|
||||
|
||||
### Wip
|
||||
|
||||
- 🚧 loop drawing ([ead4b34](https://github.com/melMass/comfy_mtb/commit/ead4b34e6dd03ea4ed309b246ef31c995325aa08))
|
||||
|
||||
## New Contributors
|
||||
* [@poetryiii](https://github.com/poetryiii) made their first contribution in [#233](https://github.com/melMass/comfy_mtb/pull/233)
|
||||
* [@numz](https://github.com/numz) made their first contribution in [#](https://github.com/melMass/comfy_mtb/pull/)
|
||||
* [@webfiltered](https://github.com/webfiltered) made their first contribution in [#223](https://github.com/melMass/comfy_mtb/pull/223)
|
||||
## [0.2.0] - 2024-12-08
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- 🐛 remove mtb sidebar ([b0d52f7](https://github.com/melMass/comfy_mtb/commit/b0d52f73051368df6de2d1e10ad28ca56df72803))
|
||||
- 🐛 always enable the I/O sidebar ([ec1cb1a](https://github.com/melMass/comfy_mtb/commit/ec1cb1ac17d14670aa756dfb1ae7542397b12559))
|
||||
- 🐛 ui shifts on animation builder ([ecbb220](https://github.com/melMass/comfy_mtb/commit/ecbb220de6a05f2e506ec43f2b786be983166157))
|
||||
- 🐛 category for settings ([b6fa571](https://github.com/melMass/comfy_mtb/commit/b6fa571fd2096ace60d03cab42dba9ca37d0cb27)) in [#211](https://github.com/melMass/comfy_mtb/pull/211)
|
||||
- 🐛 new UI issues ([f272526](https://github.com/melMass/comfy_mtb/commit/f272526bfc5da95e95d42cb4c613a0b9585b2577))
|
||||
- 🐛 disable old BOOL widget (legacy) ([8596b81](https://github.com/melMass/comfy_mtb/commit/8596b8184edb484c907475a77ac1dc9e4a5c92af))
|
||||
- 🐛 pass ONNX providers explicitely ([43092e4](https://github.com/melMass/comfy_mtb/commit/43092e44a4ea17f90fcfb12372da634fe4b79557))
|
||||
- 🐛 typo in mtb_widgets error catch ([80b5a0c](https://github.com/melMass/comfy_mtb/commit/80b5a0ca7459763e7662421bccd8636976eefddd)) by [@christian-byrne](https://github.com/christian-byrne) in [#197](https://github.com/melMass/comfy_mtb/pull/197)
|
||||
- 🐛 doc widget sidebar offset in the new ui ([81b3bc1](https://github.com/melMass/comfy_mtb/commit/81b3bc1651f06ad2fa7938f810d3f406f5e7c41c))
|
||||
- 🐛 don't fallback to eval ([997d2fb](https://github.com/melMass/comfy_mtb/commit/997d2fb13af6aadf36873ea2ea3317e56f405aef))
|
||||
- 🐛 rework main utils ([c99b081](https://github.com/melMass/comfy_mtb/commit/c99b0812ab4a4183ef9298fb8a7c954bc7c858b2))
|
||||
- 🐛 MaskToImage ([821a0ef](https://github.com/melMass/comfy_mtb/commit/821a0ef42735a0a97ab82be22a4fdc67c9cfc80e))
|
||||
|
||||
### Documentation
|
||||
|
||||
- 📚 update wiki ([e17c6e2](https://github.com/melMass/comfy_mtb/commit/e17c6e29f5111bf5085b1fe6f764cfd1aae709f2))
|
||||
- 📚 remove link ([5bc125d](https://github.com/melMass/comfy_mtb/commit/5bc125d2f08470c8900dfd89deca721835848917))
|
||||
- 📚 clean readme ([333f646](https://github.com/melMass/comfy_mtb/commit/333f646ab1959d2c944fb046275cc93a545d557c))
|
||||
|
||||
### Features
|
||||
|
||||
- ✨ add h264 compression node ([e32d1e0](https://github.com/melMass/comfy_mtb/commit/e32d1e02df5e3a9351f829513f7ee3ffb2934be4))
|
||||
- ✨ add postshot nodes ([27e03fa](https://github.com/melMass/comfy_mtb/commit/27e03fa23efffda461c6975b15fe3964de476cb3))
|
||||
- ✨ improve the I/O sidebar ([cd9e614](https://github.com/melMass/comfy_mtb/commit/cd9e614b1a385d6b06eacfaad62def1d69f09808)) in [#193](https://github.com/melMass/comfy_mtb/pull/193)
|
||||
- ✨ add UpscaleBBoxBy ([74af5c6](https://github.com/melMass/comfy_mtb/commit/74af5c6499ef5dd73ce66c4c21b8c3507d69b037))
|
||||
- ✨ simplified sidebar and backend ([22f7c30](https://github.com/melMass/comfy_mtb/commit/22f7c3037345a866c9ff0b06f6689748021cee63))
|
||||
- ✨ add Interpolate Condition ([0133fb9](https://github.com/melMass/comfy_mtb/commit/0133fb93bc944d0dd7593b89b36e5b2676d9397a))
|
||||
- ✨ dump of wip things... ([cf7d305](https://github.com/melMass/comfy_mtb/commit/cf7d30507e7e449c4489e6a1ca159d3d0486bc55))
|
||||
- ✨ use the new parser for documentations ([4e593bb](https://github.com/melMass/comfy_mtb/commit/4e593bb30be561e39f1790e3514f60bb39e5a261))
|
||||
- ✨ add @mtb/markdown-parser bundles ([097ca33](https://github.com/melMass/comfy_mtb/commit/097ca33b8e7b27148e183e91712dc34d98d1a69b))
|
||||
- ✨ add VitMatte nodes ([896a025](https://github.com/melMass/comfy_mtb/commit/896a025006f9c7809c5e0776393a28f908be8950))
|
||||
- ✨ add ColorCorrectGPU ([9651a70](https://github.com/melMass/comfy_mtb/commit/9651a7034120589b059329b21688708e42772453))
|
||||
- ✨ add Swap FG/BG colors to MaskToImage ([57683c3](https://github.com/melMass/comfy_mtb/commit/57683c3c7d299a117a26526d52de4c26f2ec0f69))
|
||||
- ✨ add Extract coordinates ([f99f92e](https://github.com/melMass/comfy_mtb/commit/f99f92e8f7b2d6fac56f7f40049715910e15cfee))
|
||||
- ✨ add AudioCut ([5681b46](https://github.com/melMass/comfy_mtb/commit/5681b464adce395086712b61159b2694150b8027))
|
||||
- ✨ add AudioStack ([8d0fcee](https://github.com/melMass/comfy_mtb/commit/8d0fcee2f3decc1cbbf3b850332e6b2a022e1377))
|
||||
- ✨ add AudioSequence node ([1078fc6](https://github.com/melMass/comfy_mtb/commit/1078fc6f0fb225b52536f25ec6a9fa0456a90595))
|
||||
- ✨ add Split Bbox node ([9007a70](https://github.com/melMass/comfy_mtb/commit/9007a70aa0d6b2ead0f68f7aff8ae8e3c4f3624f))
|
||||
- ✨ update lerp example ([1a0ebd5](https://github.com/melMass/comfy_mtb/commit/1a0ebd5173687784f279a9c2184c89fb3be01dc5))
|
||||
|
||||
### Miscellaneous Tasks
|
||||
|
||||
- 🧹 bump minor ([50cb6f5](https://github.com/melMass/comfy_mtb/commit/50cb6f5ed6e5d9fecb9733ef3f7852b8500005e9))
|
||||
- 🧹 add worktree to gitignores ([9ccf572](https://github.com/melMass/comfy_mtb/commit/9ccf572a158caeab9bff53853e8f6fb85b76776d))
|
||||
- 🧹 remove dupe code ([e099d58](https://github.com/melMass/comfy_mtb/commit/e099d581a7627c3a66d2e3e6df3a701b0e5f31b7))
|
||||
- 🧹 update externs ([784fb01](https://github.com/melMass/comfy_mtb/commit/784fb0145b7421e2730b52237ce6a8b63b189191))
|
||||
- 🧹 add pathlibed inputs to utils ([a825504](https://github.com/melMass/comfy_mtb/commit/a825504bdd67e3461be8118119e0becc35f8af40))
|
||||
- 🧹 disable Constant ([22190cd](https://github.com/melMass/comfy_mtb/commit/22190cd25ee590595f8f19e75a9a6c539699622b))
|
||||
- 🧹 new ui is default, flag for old ui ([a976adb](https://github.com/melMass/comfy_mtb/commit/a976adbb39a13b4cd76f224ebba40c604900c862))
|
||||
- 🧹 add methods to shared ([f8829fc](https://github.com/melMass/comfy_mtb/commit/f8829fcb373e0f9bc4f0ad36c939f372349943bf))
|
||||
- 🧹 add an old_ui flag to my launcher ([dbdf276](https://github.com/melMass/comfy_mtb/commit/dbdf27664cd207dbbc69b8d635adcd59ed8d269a))
|
||||
- 🧹 move qrcode to his own file ([7d5569e](https://github.com/melMass/comfy_mtb/commit/7d5569e5c1e0f0b6ccb505a02f74640139d6aaf9))
|
||||
|
||||
## [0.1.6] - 2024-07-03
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- 🐛 menu callback issue ([d64fac4](https://github.com/melMass/comfy_mtb/commit/d64fac4b74e0590acde5e3b8edd4a2f715448cf5))
|
||||
|
||||
### Documentation
|
||||
|
||||
- 📚 Update requirements file in INSTALL.md ([f25f6bd](https://github.com/melMass/comfy_mtb/commit/f25f6bdcd13d50f9d383065321320b0ce6a03214)) by [@elthariel](https://github.com/elthariel) in [#186](https://github.com/melMass/comfy_mtb/pull/186)
|
||||
|
||||
### Features
|
||||
|
||||
- ✨ add alpha channel support for faceswap/restore ([d6343e1](https://github.com/melMass/comfy_mtb/commit/d6343e1860f46947e93758f8bba03857c9326b38))
|
||||
|
||||
### Miscellaneous Tasks
|
||||
|
||||
- 🧹 better classname extraction ([d687497](https://github.com/melMass/comfy_mtb/commit/d687497d8041ab5d77bd31909592def6e4d0e7f6))
|
||||
- 🤖 limit release only to tags ([4eebdd8](https://github.com/melMass/comfy_mtb/commit/4eebdd8b8bff73c3db4f0248da8dac7d67cb310b))
|
||||
- 🧹 runner ([fb34671](https://github.com/melMass/comfy_mtb/commit/fb34671ee6fe80b965fe576c279ed1ff77a358f2))
|
||||
- 🤖 only publish on tag ([f1b4846](https://github.com/melMass/comfy_mtb/commit/f1b484617a917d38d9b3658d8920aa7dec672a79))
|
||||
- 🧹 small fixes ([4507842](https://github.com/melMass/comfy_mtb/commit/4507842a706141977a6a68945c36e977c358d91a))
|
||||
|
||||
## New Contributors
|
||||
* [@elthariel](https://github.com/elthariel) made their first contribution in [#186](https://github.com/melMass/comfy_mtb/pull/186)
|
||||
## [0.1.5] - 2024-06-21
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- 🐛 keep the last model match instead of first ([1edc2cd](https://github.com/melMass/comfy_mtb/commit/1edc2cd10de81297e7a895009d358813e79b70ba))
|
||||
- 🐛 properly initialize the curve value ([35622e3](https://github.com/melMass/comfy_mtb/commit/35622e3a5e58103a8f5b150556b85e97e31555e1))
|
||||
- 🐛 ImageCompare improvements ([acc2d68](https://github.com/melMass/comfy_mtb/commit/acc2d687d596bf82c2075f9a24003eacf18adfe7)) by [@christian-byrne](https://github.com/christian-byrne) in [#176](https://github.com/melMass/comfy_mtb/pull/176)
|
||||
- 🐛 repetitive warning ([780c52f](https://github.com/melMass/comfy_mtb/commit/780c52f03aca3079a1b695510341486720004bec)) by [@vxkj1211](https://github.com/vxkj1211) in [#177](https://github.com/melMass/comfy_mtb/pull/177)
|
||||
- 🐛 add back was conversion node ([349a852](https://github.com/melMass/comfy_mtb/commit/349a8524c6f7fcab4a124cacb60bfbef1463cf1b))
|
||||
- 🐛 drag lag on documentation resize handle ([15330ea](https://github.com/melMass/comfy_mtb/commit/15330eab655f66214d3c25fd237679f090175c32))
|
||||
- 🐛 kwarg typo ([1571782](https://github.com/melMass/comfy_mtb/commit/1571782d012b83bce32a065e700f9a587db234d2))
|
||||
- 🐛 seed of PlotBatchFloat ([5b40302](https://github.com/melMass/comfy_mtb/commit/5b4030288d43c79859c9706a12aa0f8b7dea190f))
|
||||
- 🐛 forceInput for FLOAT <-> FLOATS converters ([5a0ef0d](https://github.com/melMass/comfy_mtb/commit/5a0ef0dadd01fd5937ed0715d829d6a456f96318))
|
||||
- 🐛 FLOAT always need options to be set ([967e72f](https://github.com/melMass/comfy_mtb/commit/967e72fc66780685f8192cb8fe13ba66b9326f63))
|
||||
- 🐛 remove doc if opened on node delete ([bee3f47](https://github.com/melMass/comfy_mtb/commit/bee3f47a14ddb92b3760098666bf75dc7d37f1e4))
|
||||
- 🐛 for documentation on HiDPI ([b11346a](https://github.com/melMass/comfy_mtb/commit/b11346aba88d9f1dac3b6b42c691979cc0978b6f))
|
||||
- 🐛 never remove input 0 of dynamic inputs ([30982fa](https://github.com/melMass/comfy_mtb/commit/30982fa48829c3fc2a6745ce5a07537a3d94b2f9))
|
||||
- 🐛 use the same fix as dynamicInputs for debug ([92b7990](https://github.com/melMass/comfy_mtb/commit/92b79906cd2ee1b4ca3ff25378d7786b5a47cb75))
|
||||
- 🐛 missing numberInput ([76f365b](https://github.com/melMass/comfy_mtb/commit/76f365b5eee165c76f3da7d2e3950786685bc08b))
|
||||
- 🐛 better curve ([da67e76](https://github.com/melMass/comfy_mtb/commit/da67e766c2f700dd9e2f51a5bafe07c612904f5d))
|
||||
- 🐛 prepend MTB_ to all classes ([b1d74ad](https://github.com/melMass/comfy_mtb/commit/b1d74adb15166e3e5eb9cf92d6148e4644bed346))
|
||||
- 🐛 dynamic connections ([652ac3f](https://github.com/melMass/comfy_mtb/commit/652ac3f3b971582b02115177fd6f7a9d3d7295df))
|
||||
- 🐛 remaining issue before merge ([100067a](https://github.com/melMass/comfy_mtb/commit/100067a645194366426f29b085bf25d0623f4fac))
|
||||
- 🐛 debug issues ([7807449](https://github.com/melMass/comfy_mtb/commit/7807449e6dcc01cfdb7f0eb818569184c8b41af2))
|
||||
- 🐛 errors when insightface's folder missing ([e838c04](https://github.com/melMass/comfy_mtb/commit/e838c04758402250fd3464d6cd6a6f872e8cef29))
|
||||
- 🐛 typo ([e40ad7a](https://github.com/melMass/comfy_mtb/commit/e40ad7a574f961ebe1f338b97214da5cbadcc529))
|
||||
- 🐛 better defaults (cont) ([1da483a](https://github.com/melMass/comfy_mtb/commit/1da483a8baa6a893f1adb05ef79b90c4412c3834))
|
||||
- 🐛 better defaults for Autopan ([5eff38b](https://github.com/melMass/comfy_mtb/commit/5eff38b387d22206d39c08e435806f9d03992feb))
|
||||
- 🐛 dynamic inputs ([9ab20a0](https://github.com/melMass/comfy_mtb/commit/9ab20a0ab50b1656ded9a84c13769fd2d547f2d2))
|
||||
- 🐛 bundle ace editor ([7c35582](https://github.com/melMass/comfy_mtb/commit/7c3558273bebc0754c802720e705232f220a0da4))
|
||||
- 🐛 image to mask ([f16d576](https://github.com/melMass/comfy_mtb/commit/f16d576f6f0e83fc2fafd2d1f29b2edeb00d3197))
|
||||
- 🐛 prepend MTB to classnames ([e56508c](https://github.com/melMass/comfy_mtb/commit/e56508c2078155f053e7f11d538a048df6a5b18b))
|
||||
- 🐛 allow smaller values in BatchTransform ([9a4b27d](https://github.com/melMass/comfy_mtb/commit/9a4b27d2e05e8ebe31f58a21db94bd3a54ed23d9))
|
||||
- 🐛 add category for virtual note+ ([eeac8c0](https://github.com/melMass/comfy_mtb/commit/eeac8c002ad1f9e461418fb66b9338e969259e58))
|
||||
- 🐛 make image feed of by default ([df0a98b](https://github.com/melMass/comfy_mtb/commit/df0a98b94a4a9388811bc8786e820ec892919c1a))
|
||||
- 🐛 support batch masks (colored image node) ([2465ffb](https://github.com/melMass/comfy_mtb/commit/2465ffb0d3b052fb78559394dbb550bba59b97a3))
|
||||
- 🐛 support pillow < 10 ([48f91b7](https://github.com/melMass/comfy_mtb/commit/48f91b74e2c7ef6d31c094eafa5332784a275a8b))
|
||||
- 🐛 image rotation bug ([54ff658](https://github.com/melMass/comfy_mtb/commit/54ff6583ded0ed4054f8e5d7fadf0b2350259dce)) by [@hongminpark](https://github.com/hongminpark) in [#154](https://github.com/melMass/comfy_mtb/pull/154)
|
||||
- 🐛 font fallback ([9fccdee](https://github.com/melMass/comfy_mtb/commit/9fccdee82d721e88c64d2292c209fec869524dd2))
|
||||
- ✨ optional inputs of colored image ([cd32f26](https://github.com/melMass/comfy_mtb/commit/cd32f26b167088d6b489e43b260c187ea5e4d223)) by [@ScottNealon](https://github.com/ScottNealon) in [#147](https://github.com/melMass/comfy_mtb/pull/147)
|
||||
- 📝 adds a way to not load the imagefeed ([501c330](https://github.com/melMass/comfy_mtb/commit/501c3301056b2851555cccd75ab3ff15b1ab8e0c))
|
||||
@@ -47,6 +230,13 @@ Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases)
|
||||
|
||||
### Documentation
|
||||
|
||||
- 📚 update the wiki ([fa3199b](https://github.com/melMass/comfy_mtb/commit/fa3199be2b87bf3cb7484a0fee32a8ac099adc65))
|
||||
- 📚 update wiki submodule ([49cea8d](https://github.com/melMass/comfy_mtb/commit/49cea8d94508b27781506e3b5509c65e1d84e80f))
|
||||
- 📚 add the wiki as a submodule ([5998924](https://github.com/melMass/comfy_mtb/commit/59989249260a9c579ec851c50534b58f3f02cd61))
|
||||
- 📚 missing doc ([c9836a8](https://github.com/melMass/comfy_mtb/commit/c9836a87f6823db1d53e56997417f3cbe8cc4727))
|
||||
- 📚 use flat icon ([991af4f](https://github.com/melMass/comfy_mtb/commit/991af4f45ff8c660b2c45466bb219186699170ed))
|
||||
- 📚 add banodoco channel link ([9ce34b4](https://github.com/melMass/comfy_mtb/commit/9ce34b47fd99b18db7997ccce44e6063f00b6801))
|
||||
- 📚 udpate changelog ([8221c49](https://github.com/melMass/comfy_mtb/commit/8221c49942bd87c14d5063066315a449a1fee86e))
|
||||
- 📝 add changelog ([0d817bf](https://github.com/melMass/comfy_mtb/commit/0d817bf326b4a22e2221264a414af50c3b7048b9))
|
||||
- 📄 add note+ screenshot ([90d9636](https://github.com/melMass/comfy_mtb/commit/90d96366c8b7637b55d1b4f88cb9aca217c1414b))
|
||||
- 📝 add cover image ([6b993b8](https://github.com/melMass/comfy_mtb/commit/6b993b84071bbb80ba1b8bd63576f31e35d05590))
|
||||
@@ -60,6 +250,28 @@ Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases)
|
||||
|
||||
### Features
|
||||
|
||||
- ✨ add ModelPruner (wip) ([43d65ae](https://github.com/melMass/comfy_mtb/commit/43d65ae68c97e077117b17b7c9d1936583f965eb))
|
||||
- ✨ Use dynamic contrast in Color Correct ([6abac2e](https://github.com/melMass/comfy_mtb/commit/6abac2e4706a3d937420213e01468bae10cc2017)) by [@christian-byrne](https://github.com/christian-byrne) in [#180](https://github.com/melMass/comfy_mtb/pull/180)
|
||||
- ✨ StackImages add support for batch mismatch ([5060c56](https://github.com/melMass/comfy_mtb/commit/5060c561353e43624ec164cb73fce7d1d422f765))
|
||||
- ✨ add BatchFloatMath ([f9d2ebf](https://github.com/melMass/comfy_mtb/commit/f9d2ebf91d09fc214fecf7501a5490b33c30aca2))
|
||||
- ✨ add FLOATS to INTS ([1b7ae27](https://github.com/melMass/comfy_mtb/commit/1b7ae27cc1907bfba3c5166ec2c61547babd2e0a))
|
||||
- ✨ debug dict ([63ee25d](https://github.com/melMass/comfy_mtb/commit/63ee25d001d4c94aa95dc8b39008f5d943f2ab45))
|
||||
- ✨ add Swap BG/FG color menu item ([1caf7c1](https://github.com/melMass/comfy_mtb/commit/1caf7c18c372651b2be7227eb77e2251d963693d))
|
||||
- ✨ BatchFloatFit the batch version of FitNumber ([ab58c36](https://github.com/melMass/comfy_mtb/commit/ab58c362124f0f4b3178534ca78cb924fb881534))
|
||||
- ✨ add FloatToFloats (the counterpart) ([78a86da](https://github.com/melMass/comfy_mtb/commit/78a86daaf71dab5be34b90b13491460854718485))
|
||||
- ✨ add some FLOATS batch nodes ([2159395](https://github.com/melMass/comfy_mtb/commit/2159395389429c5f7012e660b41fad48d376b39f))
|
||||
- ✨ poc of the doc widget idea ([fac7529](https://github.com/melMass/comfy_mtb/commit/fac7529d1f7b6fc4b3b2e7f6022ebb23ec71169d))
|
||||
- ✨ add the backend node for Constant ([dff5b22](https://github.com/melMass/comfy_mtb/commit/dff5b2201d73c1a91d4b5864e3b974e68846a011))
|
||||
- ✨ add Constant node ([cbb5dd2](https://github.com/melMass/comfy_mtb/commit/cbb5dd2cf810d5648a64eae370dba610336b99d5))
|
||||
- ✨ add FloatsToFloat ([6ebecfd](https://github.com/melMass/comfy_mtb/commit/6ebecfd8cf1dc3779384e565a65baa9dceb43660))
|
||||
- ✨ add AutoPanEquilateral ([3513937](https://github.com/melMass/comfy_mtb/commit/35139371e84d715423015e05d1b4a6c1d88b0eb5))
|
||||
- ✨ add MatchDimensions ([5db3ebe](https://github.com/melMass/comfy_mtb/commit/5db3ebedb9d38470c82544e45970775193add05c))
|
||||
- ✨ add equilateral example ([8d65556](https://github.com/melMass/comfy_mtb/commit/8d65556c37f33d1c496504db92574805916dd613))
|
||||
- ✨ enhance tiling tools ([ba73fc6](https://github.com/melMass/comfy_mtb/commit/ba73fc6af7039a4629a73cdc36a8c8736dc27c9d))
|
||||
- ✨ add FLOATS support to blur ([92c810c](https://github.com/melMass/comfy_mtb/commit/92c810c5036f7a2b3f84a3fde8c81e6a2b046b07))
|
||||
- ✨ add "tube" to Batch Shape ([f658fc3](https://github.com/melMass/comfy_mtb/commit/f658fc31e040141209384d98dfe84b766fe4ae11))
|
||||
- ✨ note+ editor themes ([133da70](https://github.com/melMass/comfy_mtb/commit/133da705c94af2dfb3d2f38c0d9c2723c72cacf7))
|
||||
- ✨ add ffmpeg gif export ([1b29aad](https://github.com/melMass/comfy_mtb/commit/1b29aad360116e631b7b4d34e98a5a631f134977)) by [@huanggou666](https://github.com/huanggou666) in [#159](https://github.com/melMass/comfy_mtb/pull/159)
|
||||
- ✨ add "To Device" ([c28181f](https://github.com/melMass/comfy_mtb/commit/c28181f1615d2e183767aa76cc2350934330e546))
|
||||
- ✨ add note+ example ([90f3bc2](https://github.com/melMass/comfy_mtb/commit/90f3bc2d953b299ea34e9e3a925f1a824b488855))
|
||||
- 💄 node+ improvements ([4b29395](https://github.com/melMass/comfy_mtb/commit/4b29395000254382882c0d1be115b2ed80cd7c99))
|
||||
@@ -82,6 +294,19 @@ Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases)
|
||||
|
||||
### Miscellaneous Tasks
|
||||
|
||||
- 🧹 add fields for the registry ([bb5682a](https://github.com/melMass/comfy_mtb/commit/bb5682aa6da923859db33830c2e46f24b19199a1))
|
||||
- 🧹 add pre-commit ([59612fd](https://github.com/melMass/comfy_mtb/commit/59612fd8110a888f0081433242a2b5a5f7e46da6))
|
||||
- 🧹 migrate from poetry to setuptools ([dfd17f6](https://github.com/melMass/comfy_mtb/commit/dfd17f6d783e784df7dab38d185c747b4c04d1d0))
|
||||
- 🧹 remove logs ([1070edd](https://github.com/melMass/comfy_mtb/commit/1070edd0245fb235183d5f38cd1bebf6e0405f97))
|
||||
- 🧹 add more pyproject meta ([644371e](https://github.com/melMass/comfy_mtb/commit/644371e5b5a2b8260fc5c6f699465b0bc1c81d57))
|
||||
- 🤖 move at the proper location ([f3d468c](https://github.com/melMass/comfy_mtb/commit/f3d468cfc238f13905a13a7b2225e3711129c64d))
|
||||
- 🤖 add CI to publish to ComfyUI Registry ([6cd448b](https://github.com/melMass/comfy_mtb/commit/6cd448b026956cdf3f1b81e93724b295316fbf09)) by [@haohaocreates](https://github.com/haohaocreates) in [#182](https://github.com/melMass/comfy_mtb/pull/182)
|
||||
- 🧹 add ComfyUI registry to pyproject.toml ([5951c90](https://github.com/melMass/comfy_mtb/commit/5951c90b10f9b77b2b617e83efe0112f43c8daef)) by [@haohaocreates](https://github.com/haohaocreates) in [#181](https://github.com/melMass/comfy_mtb/pull/181)
|
||||
- 🧹 update types ([96a0da9](https://github.com/melMass/comfy_mtb/commit/96a0da9dbd051d1fcf8b332c54ed2d307d8ae0dd))
|
||||
- 🧹 use a gettattr fallback ([a344cdc](https://github.com/melMass/comfy_mtb/commit/a344cdcba9823ca1fb0762795068039b1e1cf0ab))
|
||||
- 🧹 cleanup js ([64cc4e9](https://github.com/melMass/comfy_mtb/commit/64cc4e9649853023d645245bea1e1ceb11073f01))
|
||||
- 🧹 add savedatabundle js part ([edd7c3f](https://github.com/melMass/comfy_mtb/commit/edd7c3f5d075b640e9cdb067ebfe51c42ff61791))
|
||||
- 🧹 wip dynamic multitype ([71bfdd6](https://github.com/melMass/comfy_mtb/commit/71bfdd61d731ce15f9bd0bb19d65b5af208d5dcf))
|
||||
- 🧹 applied some linting ([fe49312](https://github.com/melMass/comfy_mtb/commit/fe49312cbef03c6540304448fa88aa7a88391efa))
|
||||
- 📝 header links not parsed ([514c0d2](https://github.com/melMass/comfy_mtb/commit/514c0d2eda9990435eb18258d4bbd1aa137feb3d))
|
||||
- 📝 hardcode links in changelog ([915b744](https://github.com/melMass/comfy_mtb/commit/915b7444a9db83f349d83b636304af0d276f529f))
|
||||
@@ -104,9 +329,17 @@ Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases)
|
||||
|
||||
### Wip
|
||||
|
||||
- 🚧 curve widget logic fixed ([e312b02](https://github.com/melMass/comfy_mtb/commit/e312b02ad2f8334e87654a20b0114837df229371))
|
||||
- 🚧 dump3 ([eedbb4b](https://github.com/melMass/comfy_mtb/commit/eedbb4bc6581bef85c746307fe9d53360ea45bcf))
|
||||
- 🚧 dump ([fa23975](https://github.com/melMass/comfy_mtb/commit/fa2397585fff4f54bcf17f0b0e0083c427b34fa8))
|
||||
- 🚧 dump ([0d0fb8e](https://github.com/melMass/comfy_mtb/commit/0d0fb8e13a5da54a44a96a04607f7a349f8fdb03))
|
||||
- 🚧 add text template node ([af2175a](https://github.com/melMass/comfy_mtb/commit/af2175a1fc0c2fb29ef3493f242fe45ec6fcabac))
|
||||
|
||||
## New Contributors
|
||||
* [@haohaocreates](https://github.com/haohaocreates) made their first contribution in [#182](https://github.com/melMass/comfy_mtb/pull/182)
|
||||
* [@vxkj1211](https://github.com/vxkj1211) made their first contribution in [#177](https://github.com/melMass/comfy_mtb/pull/177)
|
||||
* [@huanggou666](https://github.com/huanggou666) made their first contribution in [#159](https://github.com/melMass/comfy_mtb/pull/159)
|
||||
* [@hongminpark](https://github.com/hongminpark) made their first contribution in [#154](https://github.com/melMass/comfy_mtb/pull/154)
|
||||
* [@ScottNealon](https://github.com/ScottNealon) made their first contribution in [#147](https://github.com/melMass/comfy_mtb/pull/147)
|
||||
* [@Yurchikian](https://github.com/Yurchikian) made their first contribution in [#124](https://github.com/melMass/comfy_mtb/pull/124)
|
||||
* [@M1kep](https://github.com/M1kep) made their first contribution in [#91](https://github.com/melMass/comfy_mtb/pull/91)
|
||||
@@ -393,7 +626,10 @@ Check the notes in the [releases](https://github.com/melMass/comfy_mtb/releases)
|
||||
|
||||
- 🚀 add gh action ([572b4d5](https://github.com/melMass/comfy_mtb/commit/572b4d52bce1398660d4d7ca0c5c48c11e0128e3)) in [#4](https://github.com/melMass/comfy_mtb/pull/4)
|
||||
|
||||
[main]: https://github.com/melMass/comfy_mtb/compare/v0.1.4..main
|
||||
[main]: https://github.com/melMass/comfy_mtb/compare/v0.2.0..main
|
||||
[0.2.0]: https://github.com/melMass/comfy_mtb/compare/v0.1.6..v0.2.0
|
||||
[0.1.6]: https://github.com/melMass/comfy_mtb/compare/v0.1.5..v0.1.6
|
||||
[0.1.5]: https://github.com/melMass/comfy_mtb/compare/v0.1.4..v0.1.5
|
||||
[0.1.4]: https://github.com/melMass/comfy_mtb/compare/v0.1.3..v0.1.4
|
||||
[0.1.3]: https://github.com/melMass/comfy_mtb/compare/v0.1.2..v0.1.3
|
||||
[0.1.2]: https://github.com/melMass/comfy_mtb/compare/v0.1.1..v0.1.2
|
||||
|
||||
@@ -1,93 +0,0 @@
|
||||
# 安装
|
||||
- [安装](#安装)
|
||||
- [自动安装(推荐)](#自动安装推荐)
|
||||
- [ComfyUI 管理器](#comfyui-管理器)
|
||||
- [虚拟环境](#虚拟环境)
|
||||
- [模型下载](#模型下载)
|
||||
- [网络扩展](#网络扩展)
|
||||
- [旧的安装方法 (MANUAL)](#旧的安装方法-manual)
|
||||
- [依赖关系](#依赖关系)
|
||||
### 自动安装(推荐)
|
||||
|
||||
### ComfyUI 管理器
|
||||
|
||||
从 0.1.0 版开始,该扩展将使用 [ComfyUI-Manager](https://github.com/ltdrdata/ComfyUI-Manager) 进行安装,这对处理各种环境下的各种安装问题大有帮助。
|
||||
|
||||
### 虚拟环境
|
||||
还有一种试验性的单行安装方法,即在 ComfyUI 根目录下使用以下命令进行安装。它将下载代码、安装依赖项并运行安装脚本:
|
||||
|
||||
|
||||
```bash
|
||||
curl -sSL "https://raw.githubusercontent.com/username/repo/main/install.py" | python3 -
|
||||
```
|
||||
|
||||
## 模型下载
|
||||
某些节点需要下载额外的模型,您可以使用与上述相同的 python 环境以交互方式完成下载:
|
||||
|
||||
```bash
|
||||
python scripts/download_models.py
|
||||
```
|
||||
|
||||
然后根据提示或直接按回车键下载每个模型。
|
||||
|
||||
> **Note**
|
||||
> 您可以使用以下方法下载所有型号,无需提示:
|
||||
```bash
|
||||
python scripts/download_models.py -y
|
||||
```
|
||||
|
||||
#### 网络扩展
|
||||
|
||||
首次运行时,脚本会尝试将 [网络扩展](https://github.com/melMass/comfy_mtb/tree/main/web)链接到你的 "web/extensions "文件夹,[请参阅](https://github.com/melMass/comfy_mtb/blob/d982b69a58c05ccead9c49370764beaa4549992a/__init__.py#L45-L61)。
|
||||
|
||||
<img alt="color widget preview" src="https://github.com/melMass/comfy_mtb/assets/7041726/cff7e66a-4cc4-4866-b35b-10af0bb2d110" width=450>
|
||||
|
||||
### 旧的安装方法 (MANUAL)
|
||||
### 依赖关系
|
||||
<details><summary><h4>Custom Virtualenv(我主要用这个)</h4></summary
|
||||
|
||||
1. 确保您处于用于 ComfyUI 的 Python 环境中。
|
||||
2. 运行以下命令安装所需的依赖项:
|
||||
```bash
|
||||
pip install -r comfy_mtb/reqs.txt
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details><summary><h4>Comfy 便携式/单机版(来自 ComfyUI 版本)</h4></summary>
|
||||
|
||||
如果您使用 ComfyUI 单机版中的 `python-embeded `,那么当二进制文件没有轮子时,您就无法使用 pip 安装二进制文件的依赖项,在这种情况下,请查看最近的 [发布](https://github.com/melMass/comfy_mtb/releases),那里有一个预编译轮子的 linux 和 windows 捆绑包(只有那些需要从源代码编译的轮子),请查看 [此问题 (#1)](https://github.com/melMass/comfy_mtb/issues/1) 以获取更多信息。
|
||||

|
||||
|
||||
|
||||
</details>
|
||||
|
||||
<details><summary><h4>Google Colab</h4></summary>
|
||||
|
||||
在 **Run ComfyUI with localtunnel (Recommended Way)** 标题之后(代码单元格之前)添加一个新的代码单元格
|
||||
|
||||

|
||||
|
||||
|
||||
```python
|
||||
# download the nodes
|
||||
!git clone --recursive https://github.com/melMass/comfy_mtb.git custom_nodes/comfy_mtb
|
||||
|
||||
# download all models
|
||||
!python custom_nodes/comfy_mtb/scripts/download_models.py -y
|
||||
|
||||
# install the dependencies
|
||||
!pip install -r custom_nodes/comfy_mtb/reqs.txt -f https://download.openmmlab.com/mmcv/dist/cu118/torch2.0/index.html
|
||||
```
|
||||
|
||||
如果运行后 colab 抱怨需要重新启动运行时,请重新启动,然后不要重新运行之前的单元格,只运行运行本地隧道的单元格。(可能需要先添加一个包含 `%cd ComfyUI` 的单元格)
|
||||
|
||||
|
||||
> **Note**:
|
||||
> If you don't need all models, remove the `-y` as collab actually supports user input: 
|
||||
|
||||
> **Preview**
|
||||
> 
|
||||
|
||||
</details>
|
||||
|
||||
@@ -1,93 +0,0 @@
|
||||
# インストール
|
||||
|
||||
- [インストール](#インストール)
|
||||
- [自動インストール (推奨)](#自動インストール-推奨)
|
||||
- [ComfyUI マネージャ](#comfyui-マネージャ)
|
||||
- [仮想環境](#仮想環境)
|
||||
- [モデルのダウンロード](#モデルのダウンロード)
|
||||
- [ウェブ拡張機能](#ウェブ拡張機能)
|
||||
- [旧インストール方法 (MANUAL)](#旧インストール方法-manual)
|
||||
- [依存関係](#依存関係)
|
||||
|
||||
|
||||
## 自動インストール (推奨)
|
||||
|
||||
### ComfyUI マネージャ
|
||||
|
||||
バージョン0.1.0では、この拡張機能は[ComfyUI-Manager](https://github.com/ltdrdata/ComfyUI-Manager)と一緒にインストールすることを想定しています。これは、様々な環境で直面する様々なインストール問題を処理するのに非常に役立ちます。
|
||||
|
||||
### 仮想環境
|
||||
また、ComfyUIのルートから以下のコマンドを使用する実験的なワンライナー・インストールもあります。これはコードをダウンロードし、依存関係をインストールし、インストールスクリプトを実行します:
|
||||
|
||||
```bash
|
||||
curl -sSL "https://raw.githubusercontent.com/username/repo/main/install.py" | python3 -
|
||||
```
|
||||
|
||||
## モデルのダウンロード
|
||||
ノードによっては、追加モデルのダウンロードが必要な場合があるので、上記と同じ python 環境を使って対話的に行うことができる:
|
||||
```bash
|
||||
python scripts/download_models.py
|
||||
```
|
||||
|
||||
プロンプトに従うか、Enterを押すだけで全てのモデルをダウンロードできます。
|
||||
|
||||
|
||||
> **Note**
|
||||
> プロンプトを出さずに全てのモデルをダウンロードするには、以下のようにします:
|
||||
```bash
|
||||
python scripts/download_models.py -y
|
||||
```
|
||||
|
||||
### ウェブ拡張機能
|
||||
|
||||
初回実行時にスクリプトは[web extensions](https://github.com/melMass/comfy_mtb/tree/main/web)をあなたの快適な `web/extensions` フォルダに[シンボリックリンク](https://github.com/melMass/comfy_mtb/blob/d982b69a58c05ccead9c49370764beaa4549992a/__init__.py#L45-L61)しようとします。万が一失敗した場合は、mtbフォルダを手動で`ComfyUI/web/extensions`にコピーしてください:
|
||||
|
||||
<img alt="color widget preview" src="https://github.com/melMass/comfy_mtb/assets/7041726/cff7e66a-4cc4-4866-b35b-10af0bb2d110" width=450>
|
||||
|
||||
## 旧インストール方法 (MANUAL)
|
||||
### 依存関係
|
||||
|
||||
<details><summary><h4>カスタム Virtualenv (私は主にこれを使っています)</h4></summary>
|
||||
|
||||
1. ComfyUIで使用しているPython環境であることを確認してください。
|
||||
2. 以下のコマンドを実行して、必要な依存関係をインストールします:
|
||||
```bash
|
||||
pip install -r comfy_mtb/reqs.txt
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details><summary><h4>Comfy-portable / standalone (ComfyUI リリースより)</h4></summary>。
|
||||
|
||||
もしあなたがComfyUIスタンドアロンから`python-embeded`を使用している場合、バイナリがホイールを持っていない場合、依存関係をpipでインストールすることができません。この場合、最後の[リリース](https://github.com/melMass/comfy_mtb/releases)をチェックしてください。(ソースからのビルドが必要なもののみ)あらかじめビルドされたホイールがあるlinuxとwindows用のバンドルがあります。詳細は[この問題(#1)](https://github.com/melMass/comfy_mtb/issues/1)をチェックしてください。
|
||||
|
||||

|
||||
|
||||
</details>
|
||||
|
||||
<details><summary><h4>Google Colab</h4></summary>
|
||||
|
||||
ComfyUI with localtunnel (Recommended Way)**ヘッダーのすぐ後(コードセルの前)に、新しいコードセルを追加してください。
|
||||

|
||||
|
||||
```python
|
||||
# download the nodes
|
||||
!git clone --recursive https://github.com/melMass/comfy_mtb.git custom_nodes/comfy_mtb
|
||||
|
||||
# download all models
|
||||
!python custom_nodes/comfy_mtb/scripts/download_models.py -y
|
||||
|
||||
# install the dependencies
|
||||
!pip install -r custom_nodes/comfy_mtb/reqs.txt -f https://download.openmmlab.com/mmcv/dist/cu118/torch2.0/index.html
|
||||
```
|
||||
これを実行した後、colabがランタイムを再起動する必要があると文句を言ったら、それを実行し、それ以前のセルは再実行せず、localtunnelを実行するセルだけを再実行してください。(最初に`%cd ComfyUI`のセルを追加する必要があるかもしれません...)
|
||||
|
||||
|
||||
> **Note**:
|
||||
> すべてのモデルが必要でない場合は、`-y`を削除してください : 
|
||||
|
||||
> **プレビュー**
|
||||
> 
|
||||
|
||||
</details>
|
||||
|
||||
+1
-1
@@ -42,7 +42,7 @@ then follow the prompt or just press enter to download every models.
|
||||
1. Make sure you are in the Python environment you use for ComfyUI.
|
||||
2. Install the required dependencies by running the following command:
|
||||
```bash
|
||||
pip install -r comfy_mtb/reqs.txt
|
||||
pip install -r comfy_mtb/requirements.txt
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
@@ -1,99 +0,0 @@
|
||||
# MTB Nodes
|
||||
|
||||
<a href="https://www.buymeacoffee.com/melmass" target="_blank"><img src="https://www.buymeacoffee.com/assets/img/custom_images/orange_img.png" alt="Buy Me A Coffee" style="height: 32px !important;width: 140px !important;box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;-webkit-box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;" ></a>
|
||||
|
||||
[** 安装指南**](./INSTALL-CN.md) | [** 示例**](https://github.com/melMass/comfy_mtb/wiki/Examples)
|
||||
|
||||
欢迎使用 MTB Nodes 项目!这个代码库是开放的,您可以自由地探索和利用。它的主要目的是构建用于 [MLOPs](https://github.com/Bismuth-Consultancy-BV/MLOPs) 中的概念验证(POCs)。该项目中的许多节点都是受到现有社区贡献或内置功能的启发而创建的。
|
||||
|
||||
在继续之前,请注意与此项目中使用的某些库相关的许可证。例如,`deepbump` 库采用 [GPLv3](https://github.com/HugoTini/DeepBump/blob/master/LICENSE) 许可证。
|
||||
|
||||
- [节点列表](#节点列表)
|
||||
- [bbox](#bbox)
|
||||
- [colors](#colors)
|
||||
- [人脸检测/交换](#人脸检测交换)
|
||||
- [图像插值(动画)](#图像插值动画)
|
||||
- [图像操作](#图像操作)
|
||||
- [潜在变量工具](#潜在变量工具)
|
||||
- [其他工具](#其他工具)
|
||||
- [纹理](#纹理)
|
||||
- [Comfy 资源](#comfy-资源)
|
||||
|
||||
|
||||
|
||||
|
||||
# 节点列表
|
||||
|
||||
## bbox
|
||||
- `Bounding Box`: BBox 构造函数(自定义类型)
|
||||
- `BBox From Mask`: 从遮罩中提取边界框
|
||||
- `Crop`: 根据边界框裁剪图像
|
||||
- `Uncrop`: 根据边界框还原图像
|
||||
|
||||
## colors
|
||||
- `Colored Image`: 给定尺寸的纯色图像
|
||||
- `RGB to HSV`: -
|
||||
- `HSV to RGB`: -
|
||||
- `Color Correct`: 基本颜色校正工具
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/7c20ac83-31ff-40ea-a1a0-06c2acefb2ef" width=345/>
|
||||
|
||||
## 人脸检测/交换
|
||||
- `Face Swap`: 使用 deepinsight/insightface 模型进行人脸交换(该节点在早期版本中称为 `Roop`,功能相同,`Roop` 只是使用这些模型的应用程序)
|
||||
> **注意**
|
||||
> 人脸索引允许您选择要替换的人脸,如下所示:
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/2e9d6066-c466-4a01-bd6c-315f7f1e8b42" width=320/>
|
||||
- `Load Face Swap Model`: 加载 insightface 模型用于人脸交换
|
||||
- `Restore Face`: 使用 [GFPGan](https://github.com/TencentARC/GFPGAN) 还原人脸,与 `Face Swap` 配合使用效果很好,并支持 `bg_upscaler` 的 Comfy 原生放大器
|
||||
|
||||
## 图像插值(动画)
|
||||
- `Load Film Model`: 加载 [FILM](https://github.com/google-research/frame-interpolation) 模型
|
||||
- `Film Interpolation`: 使用 [FILM](https://github.com/google-research/frame-interpolation) 处理输入帧
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/3afd1647-6634-4b92-a34b-51432e6a9834" width=400/>
|
||||
- `Export to Prores (experimental)`: 将输入帧导出为 ProRes 4444 mov 文件。这使用 ffmpeg stdin 发送原始的 NumPy 数组,与 `Film Interpolation` 一起使用,目前很简单,但可以进一步扩展。
|
||||
|
||||
## 图像操作
|
||||
- `Blur`: 使用高斯滤波器对图像进行模糊处理。
|
||||
- `Deglaze Image`: 从 [FN16](https://github.com/Fannovel16/FN16-ComfyUI-nodes/blob/main/DeglazeImage.py) 中提取
|
||||
- `Denoise`: 对输入图像进行降噪处理
|
||||
- `Image Compare`: 比较两个图像并返回差异图像
|
||||
- `Image Premultiply`: 使用掩码对图像进行预乘处理
|
||||
- `Image Remove Background Rembg`: 使用 [RemBG](https://github.com/danielgatis/rembg) 进行背景去除
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/e69253b4-c03c-45e9-92b5-aa46fb887be8" width=320/>
|
||||
- `Image Resize Factor`: 大部分提取自 [WAS Node Suite](https://github.com/WASasquatch/was-node-suite-comfyui),经过一些编辑(特别是支持多个图像)和较少的功能。
|
||||
- `Mask To Image`: 将遮罩(Alpha)转换为带有颜色和背景的 RGB 图像
|
||||
- `Save Image Grid`: 将输入批次中的所有图像保存为图像网格。
|
||||
|
||||
## 潜在变量工具
|
||||
- `Latent Lerp`: 两个潜在变量之间的线性插值(混合)
|
||||
|
||||
|
||||
## 其他工具
|
||||
- `Concat Images`: 接受两个图像流,并将它们合并为其他 Comfy 管道支持的图像批次。
|
||||
- `Image Resize Factor`: **已弃用**,因为我后来发现了内
|
||||
|
||||
置的图像调整大小功能。
|
||||
- `Text To Image`: 使用字体将文本转换为图像的工具
|
||||
- `Styles Loader`: 加载 csv 文件并从行中填充下拉列表(类似于 A111)
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/02fe3211-18ee-4e54-a029-931388f5fde8" width=320/>
|
||||
- `Smart Step`: 一个非常基本的节点,用于获取在 KSampler 高级中使用的步骤百分比
|
||||
- `Qr Code`: 基本的 QR Code 生成器
|
||||
- `Save Tensors`: 调试节点,将来可能会被删除
|
||||
- `Int to Number`: 用于 WASSuite 数字节点的补充
|
||||
- `Smart Step`: 使用百分比来控制 `KAdvancedSampler` 的步骤(开始/停止)
|
||||
|
||||
## 纹理
|
||||
|
||||
- `DeepBump`: 从单张图片生成法线图和高度图
|
||||
|
||||
# Comfy 资源
|
||||
|
||||
**指南**:
|
||||
- [官方示例(英文)](https://comfyanonymous.github.io/ComfyUI_examples/)
|
||||
- @BlenderNeko 的[ComfyUI 社区手册(英文)](https://blenderneko.github.io/ComfyUI-docs/)
|
||||
|
||||
- @tjhayasaka 的[Tomoaki 个人 Wiki(日文)](https://comfyui.creamlab.net/guides/)
|
||||
|
||||
**扩展和自定义节点**:
|
||||
- @WASasquatch 的[Comfy 列表插件(英文)](https://github.com/WASasquatch/comfyui-plugins)
|
||||
|
||||
- [CivitAI 上的 ComfyUI 标签(英文)](https://civitai.com/tag/comfyui)
|
||||
@@ -1,96 +0,0 @@
|
||||
# MTB Nodes
|
||||
|
||||
<a href="https://www.buymeacoffee.com/melmass" target="_blank"><img src="https://www.buymeacoffee.com/assets/img/custom_images/orange_img.png" alt="Buy Me A Coffee" style="height: 32px !important;width: 140px !important;box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;-webkit-box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;" ></a>
|
||||
|
||||
[**インストールガイド**](./INSTALL-JP.md) | [**サンプル**](https://github.com/melMass/comfy_mtb/wiki/Examples)
|
||||
|
||||
MTB Nodesプロジェクトへようこそ!このコードベースは、自由に探索し、利用することができます。主な目的は、[MLOPs](https://github.com/Bismuth-Consultancy-BV/MLOPs)の実装のための概念実証(POC)を構築することです。このプロジェクトの多くのノードは、既存のコミュニティの貢献や組み込みの機能に触発されています。
|
||||
|
||||
続行する前に、このプロジェクトで使用されている特定のライブラリに関連するライセンスに注意してください。たとえば、「deepbump」ライブラリは、[GPLv3](https://github.com/HugoTini/DeepBump/blob/master/LICENSE)の下でライセンスされています。
|
||||
|
||||
- [ノードリスト](#ノードリスト)
|
||||
- [bbox](#bbox)
|
||||
- [colors](#colors)
|
||||
- [顔検出 / スワッピング](#顔検出--スワッピング)
|
||||
- [画像補間(アニメーション)](#画像補間アニメーション)
|
||||
- [画像操作](#画像操作)
|
||||
- [潜在的なユーティリティ](#潜在的なユーティリティ)
|
||||
- [その他のユーティリティ](#その他のユーティリティ)
|
||||
- [テクスチャ](#テクスチャ)
|
||||
- [Comfyリソース](#comfyリソース)
|
||||
|
||||
|
||||
# ノードリスト
|
||||
|
||||
## bbox
|
||||
- `Bounding Box`: BBoxコンストラクタ(カスタムタイプ)
|
||||
- `BBox From Mask`: マスクからバウンディングボックスを抽出
|
||||
- `Crop`: BBoxから画像を切り抜く
|
||||
- `Uncrop`: BBoxから画像を元に戻す
|
||||
|
||||
## colors
|
||||
- `Colored Image`: 指定されたサイズの一定の色の画像
|
||||
- `RGB to HSV`: -
|
||||
- `HSV to RGB`: -
|
||||
- `Color Correct`: 基本的なカラーコレクションツール
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/7c20ac83-31ff-40ea-a1a0-06c2acefb2ef" width=345/>
|
||||
|
||||
## 顔検出 / スワッピング
|
||||
- `Face Swap`: deepinsight/insightfaceモデルを使用した顔の入れ替え(このノードは初期バージョンでは「Roop」と呼ばれていましたが、同じ機能を提供します。Roopは単にこれらのモデルを使用するアプリです)
|
||||
> **注意**
|
||||
> 顔のインデックスを使用して置き換える顔を選択できます。以下を参照してください:
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/2e9d6066-c466-4a01-bd6c-315f7f1e8b42" width=320/>
|
||||
- `Load Face Swap Model`: 顔の交換のためのinsightfaceモデルを読み込む
|
||||
- `Restore Face`: [GFPGan](https://github.com/TencentARC/GFPGAN)を使用して顔を復元し、`Face Swap`と組み合わせて使用すると非常に効果的であり、`bg_upscaler`のComfyネイティブアップスケーラーもサポートしています。
|
||||
|
||||
## 画像補間(アニメーション)
|
||||
- `Load Film Model`: [FILM](https://github.com/google-research/frame-interpolation)モデルを読み込む
|
||||
- `Film Interpolation`: [FILM](https://github.com/google-research/frame-interpolation)を使用して入力フレームを処理する
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/3afd1647-6634-4b92-a34b-51432e6a9834" width=400/>
|
||||
- `Export to Prores (experimental)`: 入力フレームをProRes 4444 movファイルにエクスポートします。これは現在は単純なものですが、`Film Interpolation`と組み合わせて使用するためのffmpegのstdinを使用して生のNumPy配列を送信するもので、拡張することもできます。
|
||||
|
||||
## 画像操作
|
||||
- `Blur`: ガウスフィルタを使用して画像をぼかす
|
||||
- `Deglaze Image`: [FN16](https://github.com/Fannovel16/FN16-ComfyUI-nodes/blob/main/DeglazeImage.py)から取得
|
||||
- `Denoise`: 入力画像のノイズを除去する
|
||||
- `Image Compare`: 2つの画像を比較し、差分画像を返す
|
||||
- `Image Premultiply`: 画像をマスクで乗算
|
||||
- `Image Remove Background Rembg`: [RemBG](https://github.com/danielgatis/rembg)を使用した背景除去
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/704172
|
||||
|
||||
6/e69253b4-c03c-45e9-92b5-aa46fb887be8" width=320/>
|
||||
- `Image Resize Factor`: [WAS Node Suite](https://github.com/WASasquatch/was-node-suite-comfyui)から抽出され、いくつかの編集(特に複数の画像のサポート)と機能の削減が行われました。
|
||||
- `Mask To Image`: マスク(アルファ)をカラーと背景を持つRGBイメージに変換します。
|
||||
- `Save Image Grid`: 入力バッチのすべての画像を画像グリッドとして保存します。
|
||||
|
||||
## 潜在的なユーティリティ
|
||||
- `Latent Lerp`: 2つの潜在的なベクトルの間の線形補間(ブレンド)
|
||||
|
||||
## その他のユーティリティ
|
||||
- `Concat Images`: 2つの画像ストリームを取り、他のComfyパイプラインでサポートされている画像のバッチとしてマージします。
|
||||
- `Image Resize Factor`: **非推奨**。組み込みの画像リサイズ機能を発見したため、削除される予定です。
|
||||
- `Text To Image`: フォントを使用してテキストを画像に変換するためのユーティリティ
|
||||
- `Styles Loader`: csvファイルをロードし、行からドロップダウンを作成します(A111のようなもの)
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/02fe3211-18ee-4e54-a029-931388f5fde8" width=320/>
|
||||
- `Smart Step`: KSamplerの高度な使用に使用するステップパーセントを取得する非常に基本的なノード
|
||||
- `Qr Code`: 基本的なQRコード生成器
|
||||
- `Save Tensors`: 将来的に削除される可能性のあるデバッグノード
|
||||
- `Int to Number`: WASSuiteの数値ノードの補完
|
||||
- `Smart Step`: `KAdvancedSampler`のステップ(開始/停止)を制御するための非常に基本的なツールで、パーセンテージを使用します。
|
||||
|
||||
## テクスチャ
|
||||
|
||||
- `DeepBump`: 1枚の画像から法線マップと高さマップを生成します。
|
||||
|
||||
# Comfyリソース
|
||||
|
||||
**ガイド**:
|
||||
- [公式の例(英語)](https://comfyanonymous.github.io/ComfyUI_examples/)
|
||||
- @BlenderNekoによる[ComfyUIコミュニティマニュアル(英語)](https://blenderneko.github.io/ComfyUI-docs/)
|
||||
|
||||
- @tjhayasakaによる[Tomoakiの個人Wiki(日本語)](https://comfyui.creamlab.net/guides/)
|
||||
|
||||
**拡張機能とカスタムノード**:
|
||||
- @WASasquatchによる[Comfyリスト用のプラグイン(英語)](https://github.com/WASasquatch/comfyui-plugins)
|
||||
|
||||
- [CivitAIのComfyUIタグ(英語)](https://civitai.com/tag/comfyui)
|
||||
@@ -4,177 +4,8 @@
|
||||

|
||||
|
||||
<!-- omit in toc -->
|
||||
|
||||
**Translated Readme (using DeepTranslate, PRs are welcome)**:
|
||||

|
||||
[日本語による説明](./README-JP.md)
|
||||

|
||||
[中文说明](./README-CN.md)
|
||||
|
||||
<a href="https://www.buymeacoffee.com/melmass" target="_blank"><img src="https://www.buymeacoffee.com/assets/img/custom_images/orange_img.png" alt="Buy Me A Coffee" style="height: 32px !important;width: 140px !important;box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;-webkit-box-shadow: 0px 3px 2px 0px rgba(190, 190, 190, 0.5) !important;" ></a>
|
||||
|
||||
[**Install Guide**](./INSTALL.md) | [**Examples**](https://github.com/melMass/comfy_mtb/wiki/Examples)
|
||||
|
||||
There is now a dedicated `#mtb-nodes` channel on the Banodoco discord:
|
||||
[](https://discord.gg/IAXhsabmDhn)
|
||||
|
||||
---
|
||||
|
||||
Welcome to the MTB Nodes project! This codebase is open for you to explore and utilize as you wish. Its primary purpose is to build proof-of-concepts (POCs) for implementation in [MLOPs](https://github.com/Bismuth-Consultancy-BV/MLOPs). Many nodes in this project are inspired by existing community contributions or built-in functionalities.
|
||||
|
||||
Before proceeding, please be aware of the licenses associated with certain libraries used in this project. For example, the `deepbump` library is licensed under [GPLv3](https://github.com/HugoTini/DeepBump/blob/master/LICENSE).
|
||||
|
||||
- [Web Extensions](#web-extensions)
|
||||
- [Node List](#node-list)
|
||||
- [Animation](#animation)
|
||||
- [bbox](#bbox)
|
||||
- [colors](#colors)
|
||||
- [image ops](#image-ops)
|
||||
- [latent utils](#latent-utils)
|
||||
- [textures](#textures)
|
||||
- [misc utils](#misc-utils)
|
||||
- [Optional nodes](#optional-nodes)
|
||||
- [face detection / swapping](#face-detection--swapping)
|
||||
- [image interpolation (animation)](#image-interpolation-animation)
|
||||
- [Comfy Resources](#comfy-resources)
|
||||
|
||||
# Web Extensions
|
||||
mtb add a few widgets like `COLOR`
|
||||
|
||||
<img alt="color widget preview" src="https://github.com/melMass/comfy_mtb/assets/7041726/cff7e66a-4cc4-4866-b35b-10af0bb2d110" width=450>
|
||||
|
||||
A few nodes have the concept of "dynamic" inputs:
|
||||
<img alt="dynamic inputs" width=450 src="https://github.com/melMass/comfy_mtb/assets/7041726/10b3976e-b212-4968-91eb-f34c02bb80c3" />
|
||||
|
||||
<!-- NOTE: Here it should just be some examples and warnings, move the rest to the wiki -->
|
||||
|
||||
# Node List
|
||||
|
||||
## Animation
|
||||
- `Animation Builder`: Convenient way to manage basic animation maths at the core of many of my workflows (both worflows for the following GIFs are in the [examples](https://github.com/melMass/comfy_mtb/wiki/Examples))
|
||||
|
||||
**[Example lerping two conditions (blue car -> yellow car)](https://github.com/melMass/comfy_mtb/blob/main/examples/03-animation_builder-condition-lerp.json)**
|
||||
|
||||
<img width=300 src="https://user-images.githubusercontent.com/7041726/260258970-d6d66d96-fb34-40d0-9038-cbabf0714c5d.gif"/>
|
||||
|
||||
|
||||
**[Example using image transforms a feedback for a fake deforum effect](https://github.com/melMass/comfy_mtb/blob/main/examples/04-animation_builder-deforum.json)**
|
||||
|
||||
<img width=300 src="https://user-images.githubusercontent.com/7041726/260261504-303a1037-60d3-4b31-a589-b15d549752f6.gif"/>
|
||||
|
||||
- `Batch Float`: Generates a batch of float values with interpolation.
|
||||
- `Batch Shape`: Generates a batch of 2D shapes with optional shading (experimental).
|
||||
- `Batch Transform`: Transform a batch of images using a batch of keyframes.
|
||||
<img width=400 src="https://github.com/melMass/comfy_mtb/assets/7041726/3f217de1-79aa-49b0-a66a-35cf29dd8f01"/>
|
||||
- `Export With Ffmpeg`: Export with FFmpeg, it used to be export to Proress and is still tailored for YUV
|
||||
- `Fit Number` : Fit the input float using a source and target range, you can also control the interpolation curve from a list of presets (default to linear)
|
||||
|
||||
## bbox
|
||||
- `Bounding Box`: BBox constructor (custom type),
|
||||
- `BBox From Mask`: From a mask extract the bounding box
|
||||
- `Crop`: Crop image from BBox
|
||||
- `Uncrop`: Uncrop image from BBox
|
||||
|
||||
## colors
|
||||
- `Colored Image`: Constant color image of given size
|
||||
- `RGB to HSV`: -,
|
||||
- `HSV to RGB`: -,
|
||||
- `Color Correct`: Basic color correction tools
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/7c20ac83-31ff-40ea-a1a0-06c2acefb2ef" width=400/>
|
||||
|
||||
## image ops
|
||||
- `Blur`: Blur an image using a Gaussian filter.
|
||||
- `Deglaze Image`: taken from [FN16](https://github.com/Fannovel16/FN16-ComfyUI-nodes/blob/main/DeglazeImage.py),
|
||||
- `Denoise`: Denoise the input image,
|
||||
- `Image Compare`: Compare two images and return a difference image
|
||||
- `Image Premultiply`: Premultiply image with mask
|
||||
- `Image Remove Background Rembg`: [RemBG](https://github.com/danielgatis/rembg) powered background removal.
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/e69253b4-c03c-45e9-92b5-aa46fb887be8" width=320/>
|
||||
- `Image Resize Factor`: Extracted mostly from [WAS Node Suite](https://github.com/WASasquatch/was-node-suite-comfyui), with a few edits (most notably multiple image support) and less features.
|
||||
- `Mask To Image`: Converts a mask (alpha) to an RGB image with a color and background
|
||||
- `Save Image Grid`: Save all the images in the input batch as a grid of images.
|
||||
|
||||
## latent utils
|
||||
- `Latent Lerp`: Linear interpolation (blend) between two latent
|
||||
|
||||
## textures
|
||||
- `Model Patch Seamless`: Use the [seamless diffusion "hack"](https://gitlab.com/-/snippets/2395088) to patch any model to infere seamless images, check the [examples](https://github.com/melMass/comfy_mtb/wiki/Examples) to see how to use all those textures node together
|
||||
<img width=500 src="https://user-images.githubusercontent.com/7041726/272970506-9db516b5-45d2-4389-b904-b3a94660f24c.png"/>
|
||||
- `DeepBump`: Normal & height maps generation from single pictures
|
||||
<img width=500 src="https://user-images.githubusercontent.com/7041726/272970715-7e4477f6-8e18-4839-9864-83d07d6690a1.png"/>
|
||||
- `Image Tile Offset`: Mimics an old photoshop technique to check for seamless textures by offsetting tiles of the image.
|
||||
<img width=600 src="https://github.com/melMass/comfy_mtb/assets/7041726/cbcc51fb-922f-433f-acf1-c6c6c2a7ffc4" />
|
||||
|
||||
## misc utils
|
||||
- `Any To String`: Tries to take any input and convert it to a string.
|
||||
- `Concat Images`: Takes two image stream and merge them as a batch of images supported by other Comfy pipelines.
|
||||
- `Image Resize Factor`: **Deprecated**, I since discovered the builtin image resize.
|
||||
- `Text To Image`: Utils to convert text to image using a font
|
||||
- `Styles Loader`: Load csv files and populate a dropdown from the rows (à la A111)
|
||||
<img src="https://github.com/melMass/comfy_mtb/assets/7041726/02fe3211-18ee-4e54-a029-931388f5fde8" width=320/>
|
||||
- `Smart Step`: A very basic node to get step percent to use in KSampler advanced,
|
||||
- `Qr Code`: Basic QR Code generator
|
||||
- `Save Tensors`: Debug node that will probably be removed in the future
|
||||
- `Int to Number`: Supplement for WASSuite number nodes
|
||||
- `Smart Step`: A very basic tool to control the steps (start/stop) of the `KAdvancedSampler` using percentage
|
||||
- `Load Image From Url`: Load an image from the given URL
|
||||
[**Wiki**](https://github.com/melMass/comfy_mtb/wiki) | [**Install Guide**](./INSTALL.md) | [**Examples**](https://github.com/melMass/comfy_mtb/wiki/Examples)
|
||||
|
||||
|
||||
## Optional nodes
|
||||
|
||||
These nodes are still bundled in mtb, but moving forward (>0.2.0) they won't
|
||||
be setup by the install script and their dependencies won't install either.
|
||||
The reason is mostly that they all have a better alternatives available and tensorflow on windows was not a fun experience and since Python 3.11 not an experience at all.
|
||||
|
||||
For linux and mac users though these nodes didn't cause any issue and I personally still use them, these are the extra requirements needed:
|
||||
|
||||
```console
|
||||
.venv/python -m pip install tensorflow facexlib insightface basicsr
|
||||
```
|
||||
|
||||
### face detection / swapping
|
||||
> **Warning**
|
||||
> Those nodes were among the first to be implemented they do work, but on windows the installation is still not properly handled for everyone
|
||||
> As alternatives you can use [reactor](https://github.com/Gourieff/comfyui-reactor-node) for face swap and [facerestore](https://github.com/Haidra-Org/hordelib/tree/main/hordelib/nodes/facerestore) for restoration
|
||||
> You can check [this video](https://www.youtube.com/watch?v=FShlpMxbU0E) for a tutorial by Ferniclestix using these alternatives
|
||||
|
||||
- `Face Swap`: Face swap using deepinsight/insightface models (this node used to be called `Roop` in early versions, it does the same, roop is *just* an app that uses those model)
|
||||
<img width=320 src="https://user-images.githubusercontent.com/7041726/260261217-54e33446-183f-4dda-88b3-d38a1e6de980.gif"/>
|
||||
- `Load Face Swap Model`: Load an insightface model for face swapping
|
||||
- `Restore Face`: Using [GFPGan](https://github.com/TencentARC/GFPGAN) to restore faces, works great in conjunction with `Face Swap` and supports Comfy native upscalers for the `bg_upscaler`
|
||||
|
||||
### image interpolation (animation)
|
||||
> **Warning**
|
||||
> The FILM nodes will be deprecated at some point after 0.2.0, [Fannovel16](https://github.com/Fannovel16/ComfyUI-Frame-Interpolation)'s interpolation nodes implement it and they rely on a pytorch implementation of FILM
|
||||
> which solves the issues related to the ones included in mtb. They will probably remain available if your system meet the requirements and ignored otherwise.
|
||||
|
||||
<details><summary>Why?</summary>
|
||||
|
||||
> **Windows only issue**: This requires tensorflow-gpu that is unfortunately not a thing anymore on Windows since 2.10.1 (unless you use a complex WSL passthrough setup but it's still not "Windows")
|
||||
> Using this old version is quite clunky and require some patching that install.py does automatically, but the main issue is that no wheels are available for python > 3.10
|
||||
> Comfy-nightly is already using Python 11 so installing this old tf version won't work there.
|
||||
> You can in any case install the normal up to date tensorflow but that will run on CPU and is much MUCH slower for FILM inference.
|
||||
</details>
|
||||
|
||||
- `Load Film Model`: Loads a [FILM](https://github.com/google-research/frame-interpolation) model
|
||||
- `Film Interpolation`: Process input frames using [FILM](https://github.com/google-research/frame-interpolation)
|
||||
<img width=400 src="https://github.com/melMass/comfy_mtb/assets/7041726/3afd1647-6634-4b92-a34b-51432e6a9834"/>
|
||||
<img width=400 src="https://user-images.githubusercontent.com/7041726/260259079-c0f04a63-960c-43a7-ba78-a45cd5ac7514.gif"/>
|
||||
- `Export to Prores (experimental)`: Exports the input frames to a ProRes 4444 mov file. This is using ffmpeg stdin to send raw numpy arrays, used with `Film Interpolation` and very simple for now but could be expanded upon.
|
||||
|
||||
# Comfy Resources
|
||||
|
||||
**Misc**
|
||||
|
||||
- [Slick ComfyUI by NoCrypt](https://colab.research.google.com/drive/1ZMvLWEiYITmBJngtqeIQToeNuiydwI0z#scrollTo=1fWMaexXS188): A colab notebook with batteries included!
|
||||
|
||||
**Guides**:
|
||||
- [Official Examples (eng)](https://comfyanonymous.github.io/ComfyUI_examples/)
|
||||
- [ComfyUI Community Manual (eng)](https://blenderneko.github.io/ComfyUI-docs/) by @BlenderNeko
|
||||
|
||||
- [Tomoaki's personal Wiki (jap)](https://comfyui.creamlab.net/guides/) by @tjhayasaka
|
||||
|
||||
**Extensions and Custom Nodes**:
|
||||
- [Plugins for Comfy List (eng)](https://github.com/WASasquatch/comfyui-plugins) by @WASasquatch
|
||||
|
||||
- [ComfyUI tag on CivitAI (eng)](https://civitai.com/tag/comfyui)
|
||||
|
||||
+239
-58
@@ -7,10 +7,12 @@
|
||||
#
|
||||
###
|
||||
|
||||
__version__ = "0.1.5"
|
||||
__version__ = "0.3.0"
|
||||
|
||||
import os
|
||||
|
||||
from aiohttp.web_request import Request
|
||||
|
||||
# TODO: don't override this if the user has that setup already
|
||||
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
|
||||
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
|
||||
@@ -29,26 +31,32 @@ from importlib import reload
|
||||
from pathlib import Path
|
||||
|
||||
from aiohttp import web
|
||||
from server import PromptServer
|
||||
|
||||
import nodes
|
||||
IN_COMFY = False
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
|
||||
IN_COMFY = True
|
||||
except ModuleNotFoundError:
|
||||
IN_COMFY = False
|
||||
|
||||
|
||||
from .endpoint import endlog
|
||||
from .install import get_node_dependencies
|
||||
from .log import blue_text, cyan_text, get_label, get_summary, log
|
||||
from .utils import comfy_dir, here
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
NODE_CLASS_MAPPINGS_DEBUG = {}
|
||||
NODE_CLASS_MAPPINGS: dict[str, type] = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
|
||||
NODE_CLASS_MAPPINGS_DEBUG: dict[str, str | None] = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
|
||||
def extract_nodes_from_source(filename: Path):
|
||||
source_code = ""
|
||||
|
||||
source_code = filename.read_text(encoding="utf-8")
|
||||
|
||||
nodes = []
|
||||
nodes: list[str] = []
|
||||
|
||||
try:
|
||||
parsed = ast.parse(source_code)
|
||||
@@ -57,14 +65,15 @@ def extract_nodes_from_source(filename: Path):
|
||||
target = node.targets[0]
|
||||
if isinstance(target, ast.Name) and target.id == "__nodes__":
|
||||
value = ast.get_source_segment(source_code, node.value)
|
||||
node_value = ast.parse(value).body[0].value
|
||||
if isinstance(node_value, (ast.List, ast.Tuple)):
|
||||
nodes.extend(
|
||||
element.id
|
||||
for element in node_value.elts
|
||||
if isinstance(element, ast.Name)
|
||||
)
|
||||
break
|
||||
if value:
|
||||
node_value = ast.parse(value).body[0].value
|
||||
if isinstance(node_value, ast.List | ast.Tuple):
|
||||
nodes.extend(
|
||||
str(element.id)
|
||||
for element in node_value.elts
|
||||
if isinstance(element, ast.Name)
|
||||
)
|
||||
break
|
||||
except SyntaxError:
|
||||
log.error("Failed to parse")
|
||||
return nodes
|
||||
@@ -72,8 +81,8 @@ def extract_nodes_from_source(filename: Path):
|
||||
|
||||
def load_nodes():
|
||||
errors: list[str] = []
|
||||
nodes = []
|
||||
nodes_failed = []
|
||||
nodes: list[type] = []
|
||||
nodes_failed: list[str] = []
|
||||
|
||||
for filename in (here / "nodes").iterdir():
|
||||
if filename.suffix == ".py":
|
||||
@@ -124,7 +133,8 @@ def uninstall_old_web_extensions():
|
||||
shutil.rmtree(web_mtb)
|
||||
except Exception as e:
|
||||
log.warning(
|
||||
f"Failed to remove web mtb directory: {e}\nPlease manually remove it from disk ({web_mtb}) and restart the server."
|
||||
f"""Failed to remove web mtb directory: {e}
|
||||
Please manually remove it from disk ({web_mtb}) and restart the server."""
|
||||
)
|
||||
|
||||
|
||||
@@ -141,7 +151,7 @@ def wiki_to_classname(s: str):
|
||||
|
||||
def classname_to_wiki(s: str):
|
||||
classname = s.replace("MTB_", "")
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
start = 0
|
||||
for i in range(1, len(classname)):
|
||||
if classname[i].isupper():
|
||||
@@ -161,8 +171,6 @@ if wiki.exists() and wiki.is_dir():
|
||||
|
||||
|
||||
# - REGISTER NODES
|
||||
|
||||
|
||||
MTB_EXPORT = os.environ.get("MTB_EXPORT")
|
||||
|
||||
nodes, failed = load_nodes()
|
||||
@@ -179,7 +187,7 @@ for node_class in nodes:
|
||||
node_class.DESCRIPTION = node_class.__doc__
|
||||
if MTB_EXPORT:
|
||||
wiki_name = classname_to_wiki(class_name)
|
||||
(wiki / "nodes" / wiki_name + ".md").write_text(
|
||||
_ = (wiki / "nodes" / (wiki_name + ".md")).write_text(
|
||||
node_class.__doc__, encoding="utf-8"
|
||||
)
|
||||
|
||||
@@ -192,12 +200,15 @@ for node_class in nodes:
|
||||
NODE_CLASS_MAPPINGS[node_label] = node_class
|
||||
NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label
|
||||
NODE_CLASS_MAPPINGS_DEBUG[node_label] = node_class.__doc__
|
||||
# TODO: I removed this, I find it more convenient to write without spaces, but it breaks every of my workflows
|
||||
# TODO (cont): and until I find a way to automate the conversion, I'll leave it like this
|
||||
|
||||
# TODO: I removed this, I find it more convenient to write without spaces
|
||||
# but it breaks every of my workflows
|
||||
# TODO (cont): and until I find a way to automate the conversion
|
||||
# I'll leave it like this
|
||||
|
||||
if os.environ.get("MTB_EXPORT"):
|
||||
with open(here / "node_list.json", "w") as f:
|
||||
f.write(
|
||||
_ = f.write(
|
||||
json.dumps(
|
||||
{
|
||||
k: NODE_CLASS_MAPPINGS_DEBUG[k]
|
||||
@@ -215,29 +226,31 @@ log.debug(
|
||||
)
|
||||
)
|
||||
|
||||
log.info(f"loaded {cyan_text(len(nodes))} nodes successfuly")
|
||||
log.info(f"loaded {cyan_text(str(len(nodes)))} nodes successfuly")
|
||||
|
||||
if failed:
|
||||
with contextlib.suppress(Exception):
|
||||
base_url, port = utils.get_server_info()
|
||||
log.info(
|
||||
f"Some nodes ({len(failed)}) could not be loaded. This can be ignored, but go to http://{base_url}:{port}/mtb if you want more information."
|
||||
)
|
||||
log.debug(failed)
|
||||
|
||||
|
||||
# - ENDPOINT
|
||||
|
||||
|
||||
if hasattr(PromptServer, "instance"):
|
||||
restore_deps = ["basicsr"]
|
||||
onnx_deps = ["onnxruntime"]
|
||||
swap_deps = ["insightface"] + onnx_deps
|
||||
node_dependency_mapping = {
|
||||
"QrCode": ["qrcode"],
|
||||
"DeepBump": onnx_deps,
|
||||
"FaceSwap": swap_deps,
|
||||
"LoadFaceSwapModel": swap_deps,
|
||||
"LoadFaceAnalysisModel": restore_deps,
|
||||
}
|
||||
if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||
img_cache = None
|
||||
prompt_cache = None
|
||||
|
||||
with contextlib.suppress(ImportError):
|
||||
from cachetools import TTLCache
|
||||
|
||||
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
|
||||
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
|
||||
|
||||
node_dependency_mapping = get_node_dependencies()
|
||||
|
||||
PromptServer.instance.app.router.add_static(
|
||||
"/mtb-assets/", path=(here / "html").as_posix()
|
||||
@@ -306,10 +319,10 @@ if hasattr(PromptServer, "instance"):
|
||||
}
|
||||
)
|
||||
|
||||
@PromptServer.instance.routes.post("/mtb/debug")
|
||||
async def set_debug(request):
|
||||
json_data = await request.json()
|
||||
enabled = json_data.get("enabled")
|
||||
@PromptServer.instance.routes.post("/mtb/server-info")
|
||||
async def set_server_info(request: Request):
|
||||
json_data: dict[str, bool] = await request.json()
|
||||
enabled = json_data.get("debug")
|
||||
if enabled:
|
||||
os.environ["MTB_DEBUG"] = "true"
|
||||
log.setLevel(logging.DEBUG)
|
||||
@@ -317,7 +330,7 @@ if hasattr(PromptServer, "instance"):
|
||||
|
||||
elif "MTB_DEBUG" in os.environ:
|
||||
# del os.environ["MTB_DEBUG"]
|
||||
os.environ.pop("MTB_DEBUG")
|
||||
_ = os.environ.pop("MTB_DEBUG")
|
||||
log.setLevel(logging.INFO)
|
||||
|
||||
return web.json_response(
|
||||
@@ -325,17 +338,17 @@ if hasattr(PromptServer, "instance"):
|
||||
)
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb")
|
||||
async def get_home(request):
|
||||
async def get_home(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
reload(endpoint)
|
||||
_ = reload(endpoint)
|
||||
# Check if the request prefers HTML content
|
||||
if "text/html" in request.headers.get("Accept", ""):
|
||||
# # Return an HTML page
|
||||
html_response = """
|
||||
<div class="flex-container menu">
|
||||
<a href="/mtb/manage">manage</a>
|
||||
<a href="/mtb/debug">debug</a>
|
||||
<a href="/mtb/server-info">Server Info</a>
|
||||
<a href="/mtb/status">status</a>
|
||||
</div>
|
||||
"""
|
||||
@@ -347,28 +360,196 @@ if hasattr(PromptServer, "instance"):
|
||||
# Return JSON for other requests
|
||||
return web.json_response({"message": "Welcome to MTB!"})
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb/debug")
|
||||
async def get_debug(request):
|
||||
import asyncio
|
||||
import os
|
||||
from io import BytesIO
|
||||
|
||||
from aiohttp import web
|
||||
from PIL import Image
|
||||
|
||||
def get_cached_image(file_path: str, preview_params=None, channel=None):
|
||||
cache_key = (file_path, preview_params, channel)
|
||||
if img_cache and (cache_key in img_cache):
|
||||
return img_cache[cache_key]
|
||||
|
||||
with Image.open(file_path) as img:
|
||||
info = img.info
|
||||
if preview_params:
|
||||
img = process_preview(img, preview_params)
|
||||
if channel:
|
||||
img = process_channel(img, channel)
|
||||
if prompt_cache:
|
||||
prompt_cache[cache_key] = info
|
||||
if img_cache:
|
||||
img_cache[cache_key] = img.getvalue()
|
||||
return img_cache[cache_key]
|
||||
|
||||
return img.getvalue()
|
||||
|
||||
def process_preview(img: Image.Image, preview_params):
|
||||
image_format, quality, width = preview_params
|
||||
quality = int(quality)
|
||||
|
||||
if width:
|
||||
width = int(width)
|
||||
img.thumbnail((width, int(width * img.height / img.width)))
|
||||
|
||||
buffer = BytesIO()
|
||||
img.save(
|
||||
buffer, format=image_format, quality=quality, metadata=img.info
|
||||
)
|
||||
buffer.seek(0)
|
||||
return buffer
|
||||
|
||||
def process_channel(img: Image.Image, channel: str):
|
||||
if channel == "rgb":
|
||||
if img.mode == "RGBA":
|
||||
r, g, b, _ = img.split()
|
||||
img = Image.merge("RGB", (r, g, b))
|
||||
else:
|
||||
img = img.convert("RGB")
|
||||
elif channel == "a":
|
||||
if img.mode == "RGBA":
|
||||
_, _, _, a = img.split()
|
||||
else:
|
||||
a = Image.new("L", img.size, 255)
|
||||
img = Image.new("RGBA", img.size)
|
||||
img.putalpha(a)
|
||||
|
||||
buffer = BytesIO()
|
||||
img.save(buffer, format="PNG")
|
||||
_ = buffer.seek(0)
|
||||
return buffer
|
||||
|
||||
async def get_image_response(
|
||||
file, filename: str, preview_info=None, channel=None
|
||||
):
|
||||
img = await asyncio.to_thread(
|
||||
get_cached_image, file, preview_info, channel
|
||||
)
|
||||
return web.Response(
|
||||
body=img,
|
||||
content_type="image/webp" if preview_info else "image/png",
|
||||
headers={"Content-Disposition": f'filename="{filename}"'},
|
||||
)
|
||||
|
||||
# TODO: Embed the metadatas somehow so we can drag and drop
|
||||
# to load workflows in the sidebar
|
||||
@PromptServer.instance.routes.get("/mtb/view")
|
||||
async def view_image(request: Request):
|
||||
import folder_paths
|
||||
|
||||
filename = request.rel_url.query.get("filename")
|
||||
if not filename:
|
||||
return web.Response(status=404)
|
||||
|
||||
filename, output_dir = folder_paths.annotated_filepath(filename)
|
||||
if filename[0] == "/" or ".." in filename:
|
||||
return web.Response(status=400)
|
||||
|
||||
if output_dir is None:
|
||||
rtype = request.rel_url.query.get("type", "output")
|
||||
output_dir = folder_paths.get_directory_by_type(rtype)
|
||||
|
||||
if output_dir is None:
|
||||
return web.Response(status=400)
|
||||
|
||||
if "subfolder" in request.rel_url.query:
|
||||
full_output_dir = os.path.join(
|
||||
output_dir, request.rel_url.query["subfolder"]
|
||||
)
|
||||
if (
|
||||
os.path.commonpath(
|
||||
(os.path.abspath(full_output_dir), output_dir)
|
||||
)
|
||||
!= output_dir
|
||||
):
|
||||
return web.Response(status=403)
|
||||
output_dir = full_output_dir
|
||||
|
||||
filename = os.path.basename(filename)
|
||||
file = os.path.join(output_dir, filename)
|
||||
|
||||
if not os.path.isfile(file):
|
||||
return web.Response(status=404)
|
||||
|
||||
ret_workflow = request.rel_url.query.get("workflow")
|
||||
|
||||
if ret_workflow:
|
||||
image = Image.open(file)
|
||||
prompt = image.info.get("prompt", "")
|
||||
workflow = image.info.get("workflow", "")
|
||||
|
||||
if workflow:
|
||||
workflow = json.loads(workflow)
|
||||
|
||||
if prompt:
|
||||
prompt = json.loads(prompt)
|
||||
|
||||
return web.json_response(
|
||||
{
|
||||
"prompt": prompt,
|
||||
"workflow": workflow,
|
||||
}
|
||||
)
|
||||
|
||||
preview_info = None
|
||||
if "preview" in request.rel_url.query:
|
||||
preview_params = request.rel_url.query["preview"].split(";")
|
||||
image_format = (
|
||||
preview_params[0]
|
||||
if preview_params[0] in ["webp", "jpeg"]
|
||||
else "webp"
|
||||
)
|
||||
quality = (
|
||||
int(preview_params[1])
|
||||
if len(preview_params) > 1 and preview_params[1].isdigit()
|
||||
else 90
|
||||
)
|
||||
width = request.rel_url.query.get("width")
|
||||
preview_info = (image_format, quality, width)
|
||||
|
||||
channel = request.rel_url.query.get("channel")
|
||||
|
||||
return await get_image_response(file, filename, preview_info, channel)
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb/server-info")
|
||||
async def get_debug(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
reload(endpoint)
|
||||
enabled = "MTB_DEBUG" in os.environ
|
||||
_ = reload(endpoint)
|
||||
isdebug = "MTB_DEBUG" in os.environ
|
||||
exposed = "MTB_EXPOSE" in os.environ
|
||||
|
||||
def render_property(name: str, val: str):
|
||||
return f"""<strong>{name}:</strong>
|
||||
<p>
|
||||
{val}
|
||||
</p>"""
|
||||
|
||||
# Check if the request prefers HTML content
|
||||
if "text/html" in request.headers.get("Accept", ""):
|
||||
# # Return an HTML page
|
||||
html_response = f"""
|
||||
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
|
||||
"""
|
||||
html_response = ""
|
||||
|
||||
html_response += render_property(
|
||||
"Debug", "Enabled" if isdebug else "Disabled"
|
||||
)
|
||||
|
||||
html_response += render_property("Exposed", str(exposed))
|
||||
|
||||
return web.Response(
|
||||
text=endpoint.render_base_template("Debug", html_response),
|
||||
text=endpoint.render_base_template(
|
||||
"Server Info", html_response
|
||||
),
|
||||
content_type="text/html",
|
||||
)
|
||||
|
||||
# Return JSON for other requests
|
||||
return web.json_response({"enabled": enabled})
|
||||
return web.json_response({"exposed": exposed, "debug": isdebug})
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb/actions")
|
||||
async def no_route(request):
|
||||
async def no_route(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
if "text/html" in request.headers.get("Accept", ""):
|
||||
@@ -382,7 +563,7 @@ if hasattr(PromptServer, "instance"):
|
||||
return web.json_response({"message": "actions has no get for now"})
|
||||
|
||||
@PromptServer.instance.routes.post("/mtb/actions")
|
||||
async def do_action(request):
|
||||
async def do_action(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
reload(endpoint)
|
||||
|
||||
+28
-19
@@ -1,22 +1,31 @@
|
||||
{
|
||||
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
||||
"organizeImports": {
|
||||
"enabled": true
|
||||
},
|
||||
"linter": {
|
||||
"enabled": true,
|
||||
"rules": {
|
||||
"recommended": true
|
||||
}
|
||||
},
|
||||
"formatter": {
|
||||
"lineEnding": "lf"
|
||||
},
|
||||
"javascript": {
|
||||
"formatter": {
|
||||
"quoteStyle": "single",
|
||||
"semicolons": "asNeeded",
|
||||
"indentWidth": 2
|
||||
}
|
||||
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
||||
"organizeImports": {
|
||||
"enabled": true
|
||||
},
|
||||
"linter": {
|
||||
"enabled": true,
|
||||
"rules": {
|
||||
"recommended": true,
|
||||
"suspicious": {
|
||||
"noConsoleLog": "warn"
|
||||
},
|
||||
"style": {
|
||||
"noParameterAssign": "off",
|
||||
"noShoutyConstants": "warn",
|
||||
"useNamingConvention": "off"
|
||||
}
|
||||
}
|
||||
},
|
||||
"formatter": {
|
||||
"indentStyle": "space",
|
||||
"indentWidth": 2,
|
||||
"lineEnding": "lf"
|
||||
},
|
||||
"javascript": {
|
||||
"formatter": {
|
||||
"quoteStyle": "single",
|
||||
"semicolons": "asNeeded"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+153
-21
@@ -1,10 +1,20 @@
|
||||
import csv
|
||||
import secrets
|
||||
import sys
|
||||
import urllib.parse
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
import folder_paths
|
||||
from aiohttp import web
|
||||
|
||||
from .install import get_node_dependencies
|
||||
from .log import mklog
|
||||
from .utils import (
|
||||
SortMode,
|
||||
backup_file,
|
||||
build_glob_patterns,
|
||||
glob_multiple,
|
||||
import_install,
|
||||
reqs_map,
|
||||
run_command,
|
||||
@@ -14,18 +24,25 @@ from .utils import (
|
||||
endlog = mklog("mtb endpoint")
|
||||
|
||||
# - ACTIONS
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import_install("requirements")
|
||||
|
||||
|
||||
def ACTIONS_installDependency(dependency_names=None):
|
||||
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
|
||||
if dependency_names is None:
|
||||
# return web.Response(text="No dependency name provided", status=400)
|
||||
return {"error": "No dependency name provided"}
|
||||
|
||||
endlog.debug(f"Received Install Dependency request for {dependency_names}")
|
||||
# reqs = []
|
||||
resolved_names = [reqs_map.get(name, name) for name in dependency_names]
|
||||
allowed_deps = list(
|
||||
{d for dep in get_node_dependencies().values() for d in dep}
|
||||
)
|
||||
for dep in dependency_names:
|
||||
if dep not in allowed_deps:
|
||||
return {
|
||||
"error": f"Unknown dependency: {dep}, you can only use this endpoint to install {allowed_deps}"
|
||||
}
|
||||
try:
|
||||
run_command(
|
||||
[Path(sys.executable), "-m", "pip", "install"] + resolved_names
|
||||
@@ -50,6 +67,106 @@ def ACTIONS_installDependency(dependency_names=None):
|
||||
# break
|
||||
|
||||
|
||||
def ACTIONS_getUserImageFolders():
|
||||
input_dir = Path(folder_paths.get_input_directory())
|
||||
output_dir = Path(folder_paths.get_output_directory())
|
||||
|
||||
input_subdirs = [x.name for x in input_dir.iterdir() if x.is_dir()]
|
||||
output_subdirs = [x.name for x in output_dir.iterdir() if x.is_dir()]
|
||||
|
||||
return {"input": input_subdirs, "output": output_subdirs}
|
||||
|
||||
|
||||
def ACTIONS_getUserVideos(
|
||||
size=256, count=200, offset=0, sort: str | None = None
|
||||
):
|
||||
count = count or 1000
|
||||
video_extensions = ["webm", "mp4", "mkv", "mov"]
|
||||
entries = {}
|
||||
patterns = build_glob_patterns(video_extensions)
|
||||
input_dir = Path(folder_paths.get_input_directory())
|
||||
entries = glob_multiple(input_dir, patterns)
|
||||
|
||||
sort_mode = SortMode.from_str(sort)
|
||||
|
||||
if sort_mode:
|
||||
sort_key = {
|
||||
SortMode.MODIFIED: lambda x: x.stat().st_mtime,
|
||||
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
|
||||
SortMode.NAME: lambda x: x.name,
|
||||
SortMode.NAME_REVERSE: lambda x: x.name,
|
||||
}.get(sort_mode)
|
||||
if sort_key:
|
||||
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
|
||||
entries = sorted(entries, key=sort_key, reverse=reverse)
|
||||
|
||||
videos = {
|
||||
video.name: (
|
||||
f"/view?force_rate=0&frame_load_cap=0&skip_first_frames=0&select_every_nth=1&filename={urllib.parse.quote_plus(video.name)}&type=input&format=video&force_size={size}x?"
|
||||
)
|
||||
for i, video in enumerate(entries)
|
||||
if offset <= i < offset + count
|
||||
}
|
||||
return videos
|
||||
|
||||
|
||||
def ACTIONS_getUserImages(
|
||||
mode: Literal["input", "output"],
|
||||
count=1000,
|
||||
offset=0,
|
||||
sort: str | None = None,
|
||||
include_subfolders: bool = False,
|
||||
subfolder=None,
|
||||
):
|
||||
# enabled = "MTB_EXPOSE" in os.environ
|
||||
# if not enabled:
|
||||
# return {"error": "Session not authorized to getInputs"}
|
||||
|
||||
imgs = {}
|
||||
count = count or 1000
|
||||
|
||||
input_dir = Path(folder_paths.get_input_directory())
|
||||
output_dir = Path(folder_paths.get_output_directory())
|
||||
|
||||
entry_dir = input_dir if mode == "input" else output_dir
|
||||
if subfolder:
|
||||
entry_dir = entry_dir / subfolder
|
||||
|
||||
if not entry_dir.exists():
|
||||
return {
|
||||
"error": f"Subfolder {entry_dir.name} doesn't exists in {entry_dir.parent.as_posix()}"
|
||||
}
|
||||
supported = ["png", "jpg", "jpeg", "webp", "gif"]
|
||||
|
||||
entries = {}
|
||||
patterns = build_glob_patterns(supported, recursive=include_subfolders)
|
||||
entries = glob_multiple(entry_dir, patterns)
|
||||
|
||||
sort_mode = SortMode.from_str(sort)
|
||||
|
||||
if sort_mode:
|
||||
sort_key = {
|
||||
SortMode.MODIFIED: lambda x: x.stat().st_mtime,
|
||||
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
|
||||
SortMode.NAME: lambda x: x.name,
|
||||
SortMode.NAME_REVERSE: lambda x: x.name,
|
||||
}.get(sort_mode)
|
||||
if sort_key:
|
||||
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
|
||||
entries = sorted(entries, key=sort_key, reverse=reverse)
|
||||
|
||||
imgs = {
|
||||
img.name: (
|
||||
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
|
||||
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
|
||||
f"&preview=&rand={secrets.randbelow(424242)}"
|
||||
)
|
||||
for i, img in enumerate(entries)
|
||||
if offset <= i < offset + count
|
||||
}
|
||||
return imgs
|
||||
|
||||
|
||||
def ACTIONS_getStyles(style_name=None):
|
||||
from .nodes.conditions import MTB_StylesLoader
|
||||
|
||||
@@ -97,7 +214,7 @@ def ACTIONS_saveStyle(data):
|
||||
csv_writer.writerow(row)
|
||||
|
||||
|
||||
async def do_action(request) -> web.Response:
|
||||
async def do_action(request: web.Request) -> web.Response:
|
||||
endlog.debug("Init action request")
|
||||
request_data = await request.json()
|
||||
name = request_data.get("name")
|
||||
@@ -109,7 +226,12 @@ async def do_action(request) -> web.Response:
|
||||
method = globals().get(method_name)
|
||||
|
||||
if callable(method):
|
||||
result = method(args) if args else method()
|
||||
result = None
|
||||
if args:
|
||||
result = method(*args) if isinstance(args, list) else method(args)
|
||||
else:
|
||||
result = method()
|
||||
|
||||
endlog.debug(f"Action result: {result}")
|
||||
return web.json_response({"result": result})
|
||||
|
||||
@@ -130,10 +252,13 @@ async def do_action(request) -> web.Response:
|
||||
# - HTML UTILS
|
||||
|
||||
|
||||
def dependencies_button(name, dependencies):
|
||||
def dependencies_button(name: str, dependencies: list[str]) -> str:
|
||||
deps = ",".join([f"'{x}'" for x in dependencies])
|
||||
return f"""
|
||||
<button class="dependency-button" onclick="window.mtb_action('installDependency',[{deps}])">Install {name} deps</button>
|
||||
<button
|
||||
class="dependency-button"
|
||||
onclick="window.mtb_action('installDependency',[{deps}])"
|
||||
>Install {name} deps</button>
|
||||
"""
|
||||
|
||||
|
||||
@@ -153,7 +278,7 @@ def csv_editor():
|
||||
html_out = """
|
||||
<div id="style-editor">
|
||||
<h1>Style Editor</h1>
|
||||
|
||||
|
||||
"""
|
||||
for current, styles in style_files.items():
|
||||
current_out = f"<h3>{current}</h3>"
|
||||
@@ -215,11 +340,14 @@ def render_tab_view(**kwargs):
|
||||
"""
|
||||
|
||||
|
||||
def add_foldable_region(title, content):
|
||||
def add_foldable_region(title: str, content: str):
|
||||
symbol_id = f"{title}-symbol"
|
||||
return f"""
|
||||
<div class='foldable'>
|
||||
<div class='foldable-title' onclick="toggleFoldable('{title}', '{symbol_id}')">
|
||||
<div
|
||||
class='foldable-title'
|
||||
onclick="toggleFoldable('{title}', '{symbol_id}')"
|
||||
>
|
||||
<span id='{symbol_id}' class='foldable-symbol'>▷</span>
|
||||
{title}
|
||||
</div>
|
||||
@@ -231,7 +359,9 @@ def add_foldable_region(title, content):
|
||||
"""
|
||||
|
||||
|
||||
def add_split_pane(left_content, right_content, vertical=True):
|
||||
def add_split_pane(
|
||||
left_content: str, right_content: str, *, vertical: bool = True
|
||||
):
|
||||
orientation = "vertical" if vertical else "horizontal"
|
||||
return f"""
|
||||
<div class="split-pane {orientation}">
|
||||
@@ -250,7 +380,7 @@ def add_split_pane(left_content, right_content, vertical=True):
|
||||
"""
|
||||
|
||||
|
||||
def add_dropdown(title, options):
|
||||
def add_dropdown(title: str, options: list[str]):
|
||||
option_str = "\n".join(
|
||||
[f"<option value='{opt}'>{opt}</option>" for opt in options]
|
||||
)
|
||||
@@ -262,18 +392,18 @@ def add_dropdown(title, options):
|
||||
"""
|
||||
|
||||
|
||||
def render_table(table_dict, sort=True, title=None):
|
||||
table_dict = sorted(
|
||||
def render_table(table_dict: dict[str, Any], sort=True, title=None):
|
||||
table_list = sorted(
|
||||
table_dict.items(), key=lambda item: item[0]
|
||||
) # Sort the dictionary by keys
|
||||
|
||||
table_rows = ""
|
||||
for name, item in table_dict:
|
||||
for name, item in table_list:
|
||||
if isinstance(item, dict):
|
||||
if "dependencies" in item:
|
||||
table_rows += f"<tr><td>{name}</td><td>"
|
||||
table_rows += (
|
||||
f"{dependencies_button(name,item['dependencies'])}"
|
||||
f"{dependencies_button(name, item['dependencies'])}"
|
||||
)
|
||||
|
||||
table_rows += "</td></tr>"
|
||||
@@ -299,12 +429,12 @@ def render_table(table_dict, sort=True, title=None):
|
||||
<tbody>
|
||||
{table_rows}
|
||||
</tbody>
|
||||
</table>
|
||||
</table>
|
||||
</div>
|
||||
"""
|
||||
|
||||
|
||||
def render_base_template(title, content):
|
||||
def render_base_template(title: str, content: str):
|
||||
github_icon_svg = """<svg xmlns="http://www.w3.org/2000/svg" fill="whitesmoke" height="3em" viewBox="0 0 496 512"><path d="M165.9 397.4c0 2-2.3 3.6-5.2 3.6-3.3.3-5.6-1.3-5.6-3.6 0-2 2.3-3.6 5.2-3.6 3-.3 5.6 1.3 5.6 3.6zm-31.1-4.5c-.7 2 1.3 4.3 4.3 4.9 2.6 1 5.6 0 6.2-2s-1.3-4.3-4.3-5.2c-2.6-.7-5.5.3-6.2 2.3zm44.2-1.7c-2.9.7-4.9 2.6-4.6 4.9.3 2 2.9 3.3 5.9 2.6 2.9-.7 4.9-2.6 4.6-4.6-.3-1.9-3-3.2-5.9-2.9zM244.8 8C106.1 8 0 113.3 0 252c0 110.9 69.8 205.8 169.5 239.2 12.8 2.3 17.3-5.6 17.3-12.1 0-6.2-.3-40.4-.3-61.4 0 0-70 15-84.7-29.8 0 0-11.4-29.1-27.8-36.6 0 0-22.9-15.7 1.6-15.4 0 0 24.9 2 38.6 25.8 21.9 38.6 58.6 27.5 72.9 20.9 2.3-16 8.8-27.1 16-33.7-55.9-6.2-112.3-14.3-112.3-110.5 0-27.5 7.6-41.3 23.6-58.9-2.6-6.5-11.1-33.3 2.6-67.9 20.9-6.5 69 27 69 27 20-5.6 41.5-8.5 62.8-8.5s42.8 2.9 62.8 8.5c0 0 48.1-33.6 69-27 13.7 34.7 5.2 61.4 2.6 67.9 16 17.7 25.8 31.5 25.8 58.9 0 96.5-58.9 104.2-114.8 110.5 9.2 7.9 17 22.9 17 46.4 0 33.7-.3 75.4-.3 83.6 0 6.5 4.6 14.4 17.3 12.1C428.2 457.8 496 362.9 496 252 496 113.3 383.5 8 244.8 8zM97.2 352.9c-1.3 1-1 3.3.7 5.2 1.6 1.6 3.9 2.3 5.2 1 1.3-1 1-3.3-.7-5.2-1.6-1.6-3.9-2.3-5.2-1zm-10.8-8.1c-.7 1.3.3 2.9 2.3 3.9 1.6 1 3.6.7 4.3-.7.7-1.3-.3-2.9-2.3-3.9-2-.6-3.6-.3-4.3.7zm32.4 35.6c-1.6 1.3-1 4.3 1.3 6.2 2.3 2.3 5.2 2.6 6.5 1 1.3-1.3.7-4.3-1.3-6.2-2.2-2.3-5.2-2.6-6.5-1zm-11.4-14.7c-1.6 1-1.6 3.6 0 5.9 1.6 2.3 4.3 3.3 5.6 2.3 1.6-1.3 1.6-3.9 0-6.2-1.4-2.3-4-3.3-5.6-2z"/></svg>"""
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
@@ -340,7 +470,9 @@ def render_base_template(title, content):
|
||||
<header>
|
||||
<a href="/">Back to Comfy</a>
|
||||
<div class="mtb_logo">
|
||||
<img src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873" alt="Comfy MTB Logo" height="70" width="128">
|
||||
<img
|
||||
src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873"
|
||||
alt="Comfy MTB Logo" height="70" width="128">
|
||||
<span class="title">Comfy MTB</span></div>
|
||||
<a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb">
|
||||
{github_icon_svg}
|
||||
@@ -355,6 +487,6 @@ def render_base_template(title, content):
|
||||
<!-- Shared footer content here -->
|
||||
</footer>
|
||||
</body>
|
||||
|
||||
|
||||
</html>
|
||||
"""
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
# NOTE: This file is only use for development you can ignore it
|
||||
|
||||
use path.nu *
|
||||
use private/log.nu
|
||||
|
||||
def get_root [--clean] {
|
||||
if $clean {
|
||||
$env.COMFY_CLEAN_ROOT
|
||||
$env.COMFY_CLEAN_ROOT
|
||||
} else {
|
||||
$env.COMFY_ROOT
|
||||
}
|
||||
@@ -23,12 +23,57 @@ export def "comfy dev-web" [] {
|
||||
npm run dev
|
||||
}
|
||||
|
||||
export def "daily run" [] {
|
||||
let res = (comfy update --rebase)
|
||||
comfy update --clean
|
||||
comfy update_extensions
|
||||
|
||||
daily commit $res.from $res.to
|
||||
}
|
||||
|
||||
def short-date [] {
|
||||
format date "%Y-%m-%d"
|
||||
}
|
||||
|
||||
# was daily run today?
|
||||
export def "daily was-run" [] {
|
||||
|
||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
||||
|
||||
if ($daily | path exists) {
|
||||
let last = (open $daily | sort-by date | get date | last | short-date)
|
||||
let today = (date now | short-date)
|
||||
return ($last == $today)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
export def "daily commit" [from:string, to:string] {
|
||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
||||
let commit = [{date: (date now) from:$from to:$to}]
|
||||
|
||||
let dailies = (if ($daily | path exists) {
|
||||
open $daily | append $commit
|
||||
} else {
|
||||
$commit
|
||||
})
|
||||
|
||||
$dailies | save -f $daily
|
||||
log success "Commited daily check"
|
||||
}
|
||||
|
||||
# start the comfy server
|
||||
export def "comfy start" [--clean, --listen] {
|
||||
export def "comfy start" [--clean,--old-ui, --listen, --skip-daily(-s)] {
|
||||
if (not (daily was-run)) and not $skip_daily {
|
||||
log info "Running daily checks"
|
||||
daily run
|
||||
}
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
MTB_DEBUG=true python main.py --port 3000 --preview-method auto ...(if $listen {["--listen"]} else {[]})
|
||||
|
||||
log info "Running Server"
|
||||
|
||||
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --front-end-version Comfy-Org/ComfyUI_frontend@latest]}) --preview-method auto ...(if $listen {["--listen"]} else {[]})
|
||||
}
|
||||
|
||||
# update comfy itself and merge master in current branch
|
||||
@@ -36,58 +81,93 @@ export def "comfy update" [
|
||||
--clean # ??
|
||||
--rebase # Rebase instead of merge
|
||||
] {
|
||||
let root = get_root --clean=($clean)
|
||||
let models = $"($root)/models"
|
||||
cd $root
|
||||
let branch_name = (git rev-parse --abbrev-ref HEAD | str trim)
|
||||
print $"(ansi yellow_italic)Backing up and removing models symlinks(ansi reset)"
|
||||
let root = get_root --clean=$clean
|
||||
|
||||
let models = $"($root)/models"
|
||||
let inputs = $"($root)/input"
|
||||
|
||||
cd $root
|
||||
|
||||
let branch_name = (git rev-parse --abbrev-ref HEAD | str trim)
|
||||
let current_commit = (git rev-parse HEAD | str trim)
|
||||
|
||||
log info "Backing up and removing models symlinks"
|
||||
|
||||
# preparing root for pull
|
||||
if not $clean {
|
||||
git checkout pyproject.toml
|
||||
cd $models
|
||||
# find and store all symlinks
|
||||
let links = (ls -la |
|
||||
where not ($it.target | is-empty) |
|
||||
select name target |
|
||||
sort-by name)
|
||||
|
||||
# find all symlinks
|
||||
let links = (ls -la |
|
||||
where not ($it.target | is-empty) |
|
||||
select name target |
|
||||
sort-by name)
|
||||
|
||||
|
||||
if not ($links | is-empty) {
|
||||
$links | save -f links.nuon
|
||||
# remove them
|
||||
open links.nuon | each {|p| rm $p.name }
|
||||
}
|
||||
if not ($links | is-empty) {
|
||||
$links | save -f links.nuon
|
||||
# remove them
|
||||
open links.nuon | each {|p| rm $p.name }
|
||||
}
|
||||
} else {
|
||||
# just remove symlinks
|
||||
rm $models
|
||||
rm $inputs
|
||||
}
|
||||
|
||||
cd $root
|
||||
cd $root
|
||||
|
||||
print $"(ansi yellow_italic)Checking out to master(ansi reset)"
|
||||
git checkout master
|
||||
log info $"Checking out to master"
|
||||
git checkout master
|
||||
|
||||
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
|
||||
log info "Fetching and pulling remote updates"
|
||||
if ($clean) {
|
||||
# from the local base repo master
|
||||
git fetch local master # $branch_name # master
|
||||
git pull local master # $branch_name # master
|
||||
} else {
|
||||
git fetch
|
||||
git pull
|
||||
}
|
||||
|
||||
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
|
||||
git checkout -
|
||||
let new_commit = (git rev-parse HEAD | str trim)
|
||||
|
||||
log info $"Back to our branch \(($branch_name)\)"
|
||||
git checkout -
|
||||
|
||||
if $current_commit == $new_commit {
|
||||
log warn "No changes upstream"
|
||||
} else {
|
||||
if $rebase {
|
||||
print $"(ansi yellow_italic)Rebasing changes(ansi reset)"
|
||||
git rebase master
|
||||
log info "Rebasing changes"
|
||||
git rebase master
|
||||
|
||||
} else {
|
||||
print $"(ansi yellow_italic)Merging changes(ansi reset)"
|
||||
git merge master
|
||||
log info "Merging changes"
|
||||
git merge master
|
||||
}
|
||||
}
|
||||
|
||||
print $"(ansi yellow_italic)Linking back the models(ansi reset)"
|
||||
|
||||
log info "Linking back the models"
|
||||
|
||||
if not $clean {
|
||||
rm pyproject.toml
|
||||
cp pyproject-mel.toml pyproject.toml
|
||||
cd $models
|
||||
|
||||
# resymlink them
|
||||
open links.nuon | each {|p| link -a $p.target $p.name }
|
||||
} else {
|
||||
let master = (get_root)
|
||||
link ($master | path join models) $models
|
||||
link ($master | path join input) $inputs
|
||||
}
|
||||
|
||||
let commit_count = (git rev-list --count $branch_name $"^origin/($branch_name)")
|
||||
let commit_count = (git rev-list --count $branch_name $"^origin/($branch_name)")
|
||||
|
||||
log success $"Update successful \(($commit_count) new commits\)"
|
||||
|
||||
print $"(ansi green_bold)Update successful \(($commit_count) new commits\)(ansi reset)"
|
||||
return {from:$current_commit to:$new_commit}
|
||||
|
||||
|
||||
}
|
||||
@@ -97,25 +177,25 @@ export def "comfy toggle_extensions" [--clean] {
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
let exts = (ls | where type in ["dir","symlink"] | get name)
|
||||
let choices = ($exts | input list -m "choose extension to toggle")
|
||||
let choices = ($exts | input list -m "choose extension to toggle")
|
||||
if ($choices | is-empty) {
|
||||
return
|
||||
}
|
||||
|
||||
print $choices
|
||||
log info "Choices" $choices
|
||||
|
||||
let filtered = $choices | wrap name | upsert enabled {|p| not ($p.name | str ends-with ".disabled")}
|
||||
|
||||
print $filtered
|
||||
|
||||
log info "Filtered" $filtered
|
||||
$filtered | each {|f|
|
||||
let new_name = ($f.name | str replace ".disabled" "")
|
||||
|
||||
|
||||
let new_name = if $f.enabled {
|
||||
$"($new_name).disabled"
|
||||
} else {
|
||||
$new_name
|
||||
}
|
||||
print $"Moving ($f.name) to ($new_name)"
|
||||
log info $"Moving ($f.name) to ($new_name)"
|
||||
mv $f.name $new_name
|
||||
}
|
||||
}
|
||||
@@ -125,14 +205,19 @@ export def "comfy update_extensions" [--clean] {
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
git multipull .
|
||||
git multipull . -s -q
|
||||
}
|
||||
|
||||
def --env path-add [pth] {
|
||||
$env.PATH = ($env.PATH | append ($pth | path expand))
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
export-env {
|
||||
$env.PYTHONUTF8 = 1
|
||||
$env.COMFY_MTB = ("." | path expand)
|
||||
$env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
|
||||
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
|
||||
|
||||
$env.CUDA_HOME = $env.CUDA_ROOT
|
||||
|
||||
@@ -140,8 +225,14 @@ export-env {
|
||||
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
|
||||
|
||||
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
|
||||
|
||||
if $nu.os-info.family == 'windows' {
|
||||
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
|
||||
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
|
||||
}
|
||||
|
||||
path-add ($env.CUDA_ROOT | path join bin)
|
||||
overlay use ../../.venv/Scripts/activate.nu
|
||||
overlay use ../../.venv/Scripts/activate.nu
|
||||
}
|
||||
|
||||
|
||||
|
||||
File diff suppressed because one or more lines are too long
+65
-29
@@ -43,10 +43,28 @@ pip_map = {
|
||||
"tb-nightly": "tensorboard",
|
||||
"protobuf": "google.protobuf",
|
||||
"qrcode[pil]": "qrcode",
|
||||
"requirements-parser": "requirements"
|
||||
"requirements-parser": "requirements",
|
||||
# Add more mappings as needed
|
||||
}
|
||||
|
||||
|
||||
def get_node_dependencies():
|
||||
restore_deps = ["basicsr"]
|
||||
onnx_deps = ["onnxruntime"]
|
||||
swap_deps = ["insightface"] + onnx_deps
|
||||
quant_deps = ["bitsandbytes"]
|
||||
io_deps = ["av"]
|
||||
return {
|
||||
"QrCode": ["qrcode"],
|
||||
"DeepBump": onnx_deps,
|
||||
"FaceSwap": swap_deps,
|
||||
"LoadFaceSwapModel": swap_deps,
|
||||
"LoadFaceAnalysisModel": restore_deps,
|
||||
"Quantize": quant_deps,
|
||||
"SaveGif": io_deps,
|
||||
}
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region ansi
|
||||
@@ -124,12 +142,12 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
|
||||
header = "[mtb install] "
|
||||
|
||||
# Handle console encoding for Unicode characters (utf-8)
|
||||
encoded_header = header.encode(sys.stdout.encoding, errors="replace").decode(
|
||||
sys.stdout.encoding
|
||||
)
|
||||
encoded_text = formatted_text.encode(sys.stdout.encoding, errors="replace").decode(
|
||||
sys.stdout.encoding
|
||||
)
|
||||
encoded_header = header.encode(
|
||||
sys.stdout.encoding, errors="replace"
|
||||
).decode(sys.stdout.encoding)
|
||||
encoded_text = formatted_text.encode(
|
||||
sys.stdout.encoding, errors="replace"
|
||||
).decode(sys.stdout.encoding)
|
||||
|
||||
print(
|
||||
" " * len(encoded_header)
|
||||
@@ -163,7 +181,9 @@ def run_command(cmd, ignored_lines_start=None):
|
||||
try:
|
||||
_run_command(shell_cmd, ignored_lines_start)
|
||||
except subprocess.CalledProcessError as e:
|
||||
print(f"Command failed with return code: {e.returncode}", file=sys.stderr)
|
||||
print(
|
||||
f"Command failed with return code: {e.returncode}", file=sys.stderr
|
||||
)
|
||||
print(e.stderr.strip(), file=sys.stderr)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
@@ -238,7 +258,7 @@ def suppress_std():
|
||||
def get_local_version():
|
||||
init_file = os.path.join(os.path.dirname(__file__), "__init__.py")
|
||||
if os.path.isfile(init_file):
|
||||
with open(init_file, "r") as f:
|
||||
with open(init_file) as f:
|
||||
tree = ast.parse(f.read())
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Assign):
|
||||
@@ -256,13 +276,16 @@ def download_file(url, file_name):
|
||||
with requests.get(url, stream=True) as response:
|
||||
response.raise_for_status()
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
with open(file_name, "wb") as file, tqdm(
|
||||
desc=file_name.stem,
|
||||
total=total_size,
|
||||
unit="B",
|
||||
unit_scale=True,
|
||||
unit_divisor=1024,
|
||||
) as progress_bar:
|
||||
with (
|
||||
open(file_name, "wb") as file,
|
||||
tqdm(
|
||||
desc=file_name.stem,
|
||||
total=total_size,
|
||||
unit="B",
|
||||
unit_scale=True,
|
||||
unit_divisor=1024,
|
||||
) as progress_bar,
|
||||
):
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
file.write(chunk)
|
||||
progress_bar.update(len(chunk))
|
||||
@@ -302,7 +325,9 @@ def import_or_install(requirement, dry=False):
|
||||
pip_install_name = pip_name + pip_spec
|
||||
|
||||
if not installed:
|
||||
print_formatted(f"Installing package {pip_name}...", "italic", color="yellow")
|
||||
print_formatted(
|
||||
f"Installing package {pip_name}...", "italic", color="yellow"
|
||||
)
|
||||
if dry:
|
||||
print_formatted(
|
||||
f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').",
|
||||
@@ -310,7 +335,9 @@ def import_or_install(requirement, dry=False):
|
||||
)
|
||||
else:
|
||||
try:
|
||||
run_command([executable, "-m", "pip", "install", pip_install_name])
|
||||
run_command(
|
||||
[executable, "-m", "pip", "install", pip_install_name]
|
||||
)
|
||||
print_formatted(
|
||||
f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')",
|
||||
"bold",
|
||||
@@ -326,13 +353,9 @@ def import_or_install(requirement, dry=False):
|
||||
|
||||
def get_github_assets(tag=None):
|
||||
if tag:
|
||||
tag_url = (
|
||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
|
||||
)
|
||||
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
|
||||
else:
|
||||
tag_url = (
|
||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
|
||||
)
|
||||
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
|
||||
response = requests.get(tag_url)
|
||||
if response.status_code == 404:
|
||||
# print_formatted(
|
||||
@@ -361,7 +384,9 @@ except ImportError:
|
||||
def main():
|
||||
if len(sys.argv) == 1:
|
||||
print_formatted(
|
||||
"mtb doesn't need an install script anymore.", "italic", color="yellow"
|
||||
"mtb doesn't need an install script anymore.",
|
||||
"italic",
|
||||
color="yellow",
|
||||
)
|
||||
return
|
||||
if all(arg not in ("-p", "--path") for arg in sys.argv):
|
||||
@@ -384,7 +409,7 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
print_formatted(f"Detected environment: {apply_color(mode,'cyan')}")
|
||||
print_formatted(f"Detected environment: {apply_color(mode, 'cyan')}")
|
||||
|
||||
if args.path:
|
||||
clone_dir = Path(args.path)
|
||||
@@ -397,8 +422,12 @@ def main():
|
||||
else:
|
||||
repo_dir = clone_dir / repo_name
|
||||
if not repo_dir.exists():
|
||||
print_formatted(f"Cloning to {repo_dir}...", "italic", color="yellow")
|
||||
run_command(["git", "clone", "--recursive", repo_url, repo_dir])
|
||||
print_formatted(
|
||||
f"Cloning to {repo_dir}...", "italic", color="yellow"
|
||||
)
|
||||
run_command(
|
||||
["git", "clone", "--recursive", repo_url, repo_dir]
|
||||
)
|
||||
else:
|
||||
print_formatted(
|
||||
f"Directory {repo_dir} already exists, we will update it..."
|
||||
@@ -409,7 +438,14 @@ def main():
|
||||
|
||||
print_formatted("Checking environment...", "italic", color="yellow")
|
||||
missing_deps = []
|
||||
install_cmd = [executable, "-m", "pip", "install", "-r", "requirements.txt"]
|
||||
install_cmd = [
|
||||
executable,
|
||||
"-m",
|
||||
"pip",
|
||||
"install",
|
||||
"-r",
|
||||
"requirements.txt",
|
||||
]
|
||||
run_command(install_cmd)
|
||||
|
||||
print_formatted(
|
||||
|
||||
@@ -77,5 +77,11 @@ def cyan_text(text: str):
|
||||
def get_label(label: str):
|
||||
if label.startswith("MTB_"):
|
||||
label = label[4:]
|
||||
words = re.findall(r"(?:^|[A-Z])[a-z]*", label)
|
||||
|
||||
words = re.findall(
|
||||
r"(?:(?<=[a-z])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])|(?<=[A-Za-z])(?=[0-9])|(?<=[0-9])(?=[A-Za-z]))",
|
||||
label,
|
||||
)
|
||||
reformatted_label = re.sub(r"([A-Z]+)", r" \1", label).strip()
|
||||
words = reformatted_label.split()
|
||||
return " ".join(words).strip()
|
||||
|
||||
+922
@@ -0,0 +1,922 @@
|
||||
from typing import Any, TypedDict
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from comfy.model_management import get_torch_device
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import (
|
||||
WhisperForConditionalGeneration,
|
||||
WhisperProcessor,
|
||||
)
|
||||
|
||||
# from transformers import (
|
||||
# AutoFeatureExtractor,
|
||||
# WhisperForConditionalGeneration,
|
||||
# WhisperModel,
|
||||
# WhisperProcessor,
|
||||
# )
|
||||
from ..log import log
|
||||
from ..utils import get_model_path
|
||||
|
||||
WHISPER_SAMPLE_RATE = 16000
|
||||
|
||||
|
||||
class AudioTensor(TypedDict):
|
||||
"""Comfy's representation of AUDIO data."""
|
||||
|
||||
sample_rate: int
|
||||
waveform: torch.Tensor
|
||||
|
||||
|
||||
class WhisperData(TypedDict):
|
||||
"""Whisper transcription data with timestamps and speaker info."""
|
||||
|
||||
text: str
|
||||
chunks: list[dict[str, Any]]
|
||||
language: str
|
||||
|
||||
|
||||
AudioData = AudioTensor | list[AudioTensor]
|
||||
|
||||
|
||||
class MtbAudio:
|
||||
"""Base class for audio processing."""
|
||||
|
||||
@classmethod
|
||||
def is_stereo(
|
||||
cls,
|
||||
audios: AudioData,
|
||||
) -> bool:
|
||||
if isinstance(audios, list):
|
||||
return any(cls.is_stereo(audio) for audio in audios)
|
||||
else:
|
||||
return audios["waveform"].shape[1] == 2
|
||||
|
||||
@staticmethod
|
||||
def resample(audio: AudioTensor, common_sample_rate: int) -> AudioTensor:
|
||||
current_rate = audio["sample_rate"]
|
||||
if current_rate != common_sample_rate:
|
||||
log.debug(
|
||||
f"Resampling audio from {current_rate} to {common_sample_rate}"
|
||||
)
|
||||
resampler = torchaudio.transforms.Resample(
|
||||
orig_freq=current_rate, new_freq=common_sample_rate
|
||||
)
|
||||
return {
|
||||
"sample_rate": common_sample_rate,
|
||||
"waveform": resampler(audio["waveform"]),
|
||||
}
|
||||
else:
|
||||
return audio
|
||||
|
||||
@staticmethod
|
||||
def to_stereo(audio: AudioTensor) -> AudioTensor:
|
||||
if audio["waveform"].shape[1] == 1:
|
||||
return {
|
||||
"sample_rate": audio["sample_rate"],
|
||||
"waveform": torch.cat(
|
||||
[audio["waveform"], audio["waveform"]], dim=1
|
||||
),
|
||||
}
|
||||
else:
|
||||
return audio
|
||||
|
||||
@classmethod
|
||||
def preprocess_audios(
|
||||
cls, audios: list[AudioTensor]
|
||||
) -> tuple[list[AudioTensor], bool, int]:
|
||||
max_sample_rate = max([audio["sample_rate"] for audio in audios])
|
||||
|
||||
resampled_audios = [
|
||||
cls.resample(audio, max_sample_rate) for audio in audios
|
||||
]
|
||||
|
||||
is_stereo = cls.is_stereo(audios)
|
||||
if is_stereo:
|
||||
audios = [cls.to_stereo(audio) for audio in resampled_audios]
|
||||
|
||||
return (audios, is_stereo, max_sample_rate)
|
||||
|
||||
|
||||
class WhisperPipeline(TypedDict):
|
||||
"""Whisper model pipeline."""
|
||||
|
||||
processor: WhisperProcessor
|
||||
model: WhisperForConditionalGeneration
|
||||
|
||||
|
||||
class MTB_LoadWhisper:
|
||||
"""Load Whisper model and processor."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model_size": (
|
||||
[
|
||||
"tiny",
|
||||
"small",
|
||||
"medium",
|
||||
"medium.en",
|
||||
"base",
|
||||
"large",
|
||||
"large-v2",
|
||||
"large-v3",
|
||||
"large-v3-turbo",
|
||||
],
|
||||
{"default": "tiny"},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"download_missing": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": (
|
||||
"Download missing models if missing,"
|
||||
"otherwise they must be in ComfyUI/models/whisper"
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WHISPER_PIPELINE",)
|
||||
RETURN_NAMES = ("pipeline",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "load"
|
||||
|
||||
def load(self, model_size="tiny", download_missing=False):
|
||||
"""Load Whisper model and processor."""
|
||||
whisper_dir = get_model_path("whisper")
|
||||
tag = f"whisper-{model_size}"
|
||||
model_dir = whisper_dir / tag
|
||||
|
||||
if not (whisper_dir.exists() or model_dir.exists()):
|
||||
if not download_missing:
|
||||
raise RuntimeError(
|
||||
"Models not found and download_missing=False"
|
||||
)
|
||||
else:
|
||||
whisper_dir.mkdir(exist_ok=True)
|
||||
model_dir.mkdir(exist_ok=True)
|
||||
|
||||
snapshot_download(
|
||||
repo_id=f"openai/{tag}",
|
||||
resume_download=True,
|
||||
ignore_patterns=["*.msgpack", "*.bin", "*.h5"],
|
||||
local_dir=model_dir.as_posix(),
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
|
||||
device = get_torch_device()
|
||||
log.debug(
|
||||
f"Loading Whisper model {model_size} on {device} from {model_dir}"
|
||||
)
|
||||
|
||||
processor = WhisperProcessor.from_pretrained(model_dir.as_posix())
|
||||
model = WhisperForConditionalGeneration.from_pretrained(
|
||||
model_dir.as_posix()
|
||||
).to(device)
|
||||
|
||||
model.eval()
|
||||
model.requires_grad_(False)
|
||||
|
||||
return ({"processor": processor, "model": model},)
|
||||
|
||||
|
||||
class MTB_AudioToText(MtbAudio):
|
||||
"""Transcribe audio to text using Whisper."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"pipeline": ("WHISPER_PIPELINE",),
|
||||
"audio": ("AUDIO",),
|
||||
"language": (
|
||||
["auto"]
|
||||
+ sorted(
|
||||
[
|
||||
"en",
|
||||
"fr",
|
||||
"es",
|
||||
"de",
|
||||
"it",
|
||||
"pt",
|
||||
"nl",
|
||||
"ru",
|
||||
"zh",
|
||||
"ja",
|
||||
"ko",
|
||||
]
|
||||
),
|
||||
{"default": "auto"},
|
||||
),
|
||||
"return_timestamps": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "WHISPER_OUTPUT")
|
||||
FUNCTION = "transcribe"
|
||||
CATEGORY = "mtb/audio"
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
pipeline: WhisperPipeline,
|
||||
audio: AudioTensor,
|
||||
language="auto",
|
||||
return_timestamps=True,
|
||||
):
|
||||
"""Transcribe audio to text using Whisper."""
|
||||
processor = pipeline["processor"]
|
||||
model = pipeline["model"]
|
||||
device = model.device
|
||||
|
||||
audio = self.resample(audio, WHISPER_SAMPLE_RATE)
|
||||
|
||||
waveform = audio["waveform"]
|
||||
log.debug(f"Processed waveform shape: {waveform.shape}")
|
||||
|
||||
# - Mono: [1, 1, samples] or [1, samples] or [samples]
|
||||
# - Stereo: [1, 2, samples] or [2, samples] or [samples, 2]
|
||||
if len(waveform.shape) == 3:
|
||||
waveform = waveform.squeeze(0)
|
||||
|
||||
if len(waveform.shape) == 2:
|
||||
if waveform.shape[0] == 2: # [channels, samples]
|
||||
waveform = waveform.mean(dim=0)
|
||||
elif waveform.shape[1] == 2: # [samples, channels]
|
||||
waveform = waveform.mean(dim=1)
|
||||
else: # mono
|
||||
waveform = waveform.squeeze(0)
|
||||
|
||||
sample_rate = audio["sample_rate"]
|
||||
chunk_duration = 30
|
||||
chunk_samples = chunk_duration * sample_rate
|
||||
total_samples = waveform.shape[-1]
|
||||
total_duration = total_samples / sample_rate
|
||||
|
||||
log.debug(f"Audio duration: {total_duration:.2f}s")
|
||||
|
||||
all_tokens = []
|
||||
all_text = []
|
||||
chunk_offsets = []
|
||||
|
||||
last_time = 0.0
|
||||
accumulated_offset = 0.0
|
||||
|
||||
for chunk_start in range(0, total_samples, chunk_samples):
|
||||
chunk_end = min(chunk_start + chunk_samples, total_samples)
|
||||
chunk_waveform = waveform[chunk_start:chunk_end]
|
||||
chunk_offset = chunk_start / sample_rate
|
||||
chunk_offsets.append(chunk_offset)
|
||||
|
||||
log.debug(
|
||||
f"Processing chunk {chunk_offset:.1f}s - {chunk_end / sample_rate:.1f}s"
|
||||
)
|
||||
|
||||
max_length = model.config.max_length or 448
|
||||
attention_mask = torch.ones((1, max_length))
|
||||
|
||||
input_features = processor(
|
||||
chunk_waveform,
|
||||
sampling_rate=sample_rate,
|
||||
return_tensors="pt",
|
||||
).input_features.to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
predicted_ids = model.generate(
|
||||
input_features,
|
||||
attention_mask=attention_mask.to(device),
|
||||
task="transcribe",
|
||||
language=None if language == "auto" else language,
|
||||
return_timestamps=return_timestamps,
|
||||
no_repeat_ngram_size=3,
|
||||
num_beams=5,
|
||||
length_penalty=1.0,
|
||||
max_length=max_length,
|
||||
)
|
||||
|
||||
chunk_tokens = processor.tokenizer.convert_ids_to_tokens(
|
||||
predicted_ids[0]
|
||||
)
|
||||
|
||||
adjusted_tokens = []
|
||||
for token in chunk_tokens:
|
||||
if token.startswith("<|") and token.endswith("|>"):
|
||||
try:
|
||||
time_str = token[2:-2]
|
||||
if time_str.replace(".", "").isdigit():
|
||||
time_val = float(time_str)
|
||||
|
||||
# If this timestamp is less than the last one, we've started a new sequence
|
||||
if time_val < last_time:
|
||||
accumulated_offset += last_time
|
||||
|
||||
adjusted_time = time_val + accumulated_offset
|
||||
adjusted_tokens.append(f"<|{adjusted_time:.2f}|>")
|
||||
last_time = time_val
|
||||
else:
|
||||
adjusted_tokens.append(token)
|
||||
except ValueError:
|
||||
adjusted_tokens.append(token)
|
||||
else:
|
||||
adjusted_tokens.append(token)
|
||||
|
||||
all_tokens.extend(adjusted_tokens)
|
||||
chunk_text = processor.batch_decode(
|
||||
predicted_ids, skip_special_tokens=True
|
||||
)[0]
|
||||
all_text.append(chunk_text)
|
||||
|
||||
detected_language = "en"
|
||||
if language == "auto":
|
||||
try:
|
||||
log.debug("Detecting language")
|
||||
with torch.no_grad():
|
||||
first_chunk_features = processor(
|
||||
waveform[:chunk_samples],
|
||||
sampling_rate=sample_rate,
|
||||
return_tensors="pt",
|
||||
).input_features.to(device)
|
||||
|
||||
predicted_probs = model.detect_language(
|
||||
first_chunk_features
|
||||
)[0]
|
||||
language_token = processor.tokenizer.convert_ids_to_tokens(
|
||||
predicted_probs.argmax(-1).item()
|
||||
)
|
||||
detected_language = (
|
||||
language_token[2:-2]
|
||||
if language_token.startswith("<|")
|
||||
else "en"
|
||||
)
|
||||
log.debug(f"Detected language: {detected_language}")
|
||||
|
||||
except Exception as e:
|
||||
log.warning(f"Language detection failed: {e}")
|
||||
|
||||
full_transcription = " ".join(all_text)
|
||||
|
||||
whisper_output = {
|
||||
"text": full_transcription,
|
||||
"language": detected_language,
|
||||
"tokens": all_tokens,
|
||||
"audio": audio,
|
||||
"chunk_offsets": chunk_offsets,
|
||||
}
|
||||
|
||||
return full_transcription, whisper_output
|
||||
|
||||
|
||||
class MTB_ProcessWhisperOutput:
|
||||
"""Process Whisper output into timestamped chunks."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"whisper_output": ("WHISPER_OUTPUT",),
|
||||
"min_chunk_length": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.1},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "WHISPER_CHUNKS")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "mtb/audio"
|
||||
|
||||
def process(self, whisper_output, min_chunk_length=0.0):
|
||||
"""Process Whisper output into timestamped chunks."""
|
||||
tokens = whisper_output["tokens"]
|
||||
audio = whisper_output["audio"]
|
||||
timestamp_tokens = []
|
||||
|
||||
audio_duration = audio["waveform"].shape[-1] / audio["sample_rate"]
|
||||
log.debug(f"Audio duration: {audio_duration:.2f}s")
|
||||
|
||||
for i, token in enumerate(tokens):
|
||||
if token.startswith("<|") and token.endswith("|>"):
|
||||
try:
|
||||
time_str = token[2:-2]
|
||||
if time_str.replace(".", "").isdigit():
|
||||
time_val = float(time_str)
|
||||
if 0 <= time_val <= audio_duration:
|
||||
timestamp_tokens.append((i, time_val))
|
||||
log.debug(f"Token {i}: {time_val}")
|
||||
except ValueError:
|
||||
continue
|
||||
|
||||
chunks = []
|
||||
if len(timestamp_tokens) > 1:
|
||||
for i in range(len(timestamp_tokens) - 1):
|
||||
start_pos, start_time = timestamp_tokens[i]
|
||||
end_pos, end_time = timestamp_tokens[i + 1]
|
||||
|
||||
if end_time - start_time < min_chunk_length:
|
||||
continue
|
||||
|
||||
chunk_tokens = tokens[start_pos + 1 : end_pos]
|
||||
text = " ".join(
|
||||
t
|
||||
for t in chunk_tokens
|
||||
if not (t.startswith("<|") and t.endswith("|>"))
|
||||
)
|
||||
|
||||
if text.strip():
|
||||
chunks.append(
|
||||
{
|
||||
"text": text.strip(),
|
||||
"timestamp": [start_time, end_time],
|
||||
}
|
||||
)
|
||||
|
||||
if timestamp_tokens:
|
||||
start_pos, start_time = timestamp_tokens[-1]
|
||||
if start_pos < len(tokens) - 1:
|
||||
text = " ".join(
|
||||
t
|
||||
for t in tokens[start_pos + 1 :]
|
||||
if not (t.startswith("<|") and t.endswith("|>"))
|
||||
)
|
||||
if text.strip():
|
||||
if chunks:
|
||||
prev_chunk = chunks[-1]
|
||||
prev_duration = (
|
||||
prev_chunk["timestamp"][1]
|
||||
- prev_chunk["timestamp"][0]
|
||||
)
|
||||
end_time = min(
|
||||
start_time + prev_duration, audio_duration
|
||||
)
|
||||
else:
|
||||
end_time = audio_duration
|
||||
|
||||
if (
|
||||
end_time > start_time
|
||||
and end_time - start_time >= min_chunk_length
|
||||
):
|
||||
chunks.append(
|
||||
{
|
||||
"text": text.strip(),
|
||||
"timestamp": [start_time, end_time],
|
||||
}
|
||||
)
|
||||
|
||||
result = {
|
||||
"text": whisper_output["text"],
|
||||
"chunks": chunks,
|
||||
"language": whisper_output["language"],
|
||||
}
|
||||
|
||||
return whisper_output["text"], result
|
||||
|
||||
|
||||
class MTB_AudioCut(MtbAudio):
|
||||
"""Basic audio cutter, values are in ms."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
"length": (
|
||||
("FLOAT"),
|
||||
{
|
||||
"default": 1000.0,
|
||||
"min": 0.0,
|
||||
"max": 999999.0,
|
||||
"step": 1,
|
||||
},
|
||||
),
|
||||
"offset": (
|
||||
("FLOAT"),
|
||||
{"default": 0.0, "min": 0.0, "max": 999999.0, "step": 1},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("cut_audio",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "cut"
|
||||
|
||||
def cut(self, audio: AudioTensor, length: float, offset: float):
|
||||
sample_rate = audio["sample_rate"]
|
||||
start_idx = int(offset * sample_rate / 1000)
|
||||
end_idx = min(
|
||||
start_idx + int(length * sample_rate / 1000),
|
||||
audio["waveform"].shape[-1],
|
||||
)
|
||||
cut_waveform = audio["waveform"][:, :, start_idx:end_idx]
|
||||
|
||||
return (
|
||||
{
|
||||
"sample_rate": sample_rate,
|
||||
"waveform": cut_waveform,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class MTB_AudioStack(MtbAudio):
|
||||
"""Stack/Overlay audio inputs (dynamic inputs).
|
||||
- pad audios to the longest inputs.
|
||||
- resample audios to the highest sample rate in the inputs.
|
||||
- convert them all to stereo if one of the inputs is.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {}}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("stacked_audio",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "stack"
|
||||
|
||||
def stack(self, **kwargs: AudioTensor) -> tuple[AudioTensor]:
|
||||
audios, is_stereo, max_rate = self.preprocess_audios(
|
||||
list(kwargs.values())
|
||||
)
|
||||
|
||||
max_length = max([audio["waveform"].shape[-1] for audio in audios])
|
||||
|
||||
padded_audios: list[torch.Tensor] = []
|
||||
for audio in audios:
|
||||
padding = torch.zeros(
|
||||
(
|
||||
1,
|
||||
2 if is_stereo else 1,
|
||||
max_length - audio["waveform"].shape[-1],
|
||||
)
|
||||
)
|
||||
padded_audio = torch.cat([audio["waveform"], padding], dim=-1)
|
||||
padded_audios.append(padded_audio)
|
||||
|
||||
stacked_waveform = torch.stack(padded_audios, dim=0).sum(dim=0)
|
||||
|
||||
return (
|
||||
{
|
||||
"sample_rate": max_rate,
|
||||
"waveform": stacked_waveform,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class MTB_AudioSequence(MtbAudio):
|
||||
"""Sequence audio inputs (dynamic inputs).
|
||||
- adding silence_duration between each segment
|
||||
can now also be negative to overlap the clips, safely bound
|
||||
to the the input length.
|
||||
- resample audios to the highest sample rate in the inputs.
|
||||
- convert them all to stereo if one of the inputs is.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"silence_duration": (
|
||||
("FLOAT"),
|
||||
{"default": 0.0, "min": -999.0, "max": 999, "step": 0.01},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("sequenced_audio",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "sequence"
|
||||
|
||||
def sequence(self, silence_duration: float, **kwargs: AudioTensor):
|
||||
audios, is_stereo, max_rate = self.preprocess_audios(
|
||||
list(kwargs.values())
|
||||
)
|
||||
|
||||
sequence: list[torch.Tensor] = []
|
||||
for i, audio in enumerate(audios):
|
||||
if i > 0:
|
||||
if silence_duration > 0:
|
||||
silence = torch.zeros(
|
||||
(
|
||||
1,
|
||||
2 if is_stereo else 1,
|
||||
int(silence_duration * max_rate),
|
||||
)
|
||||
)
|
||||
sequence.append(silence)
|
||||
elif silence_duration < 0:
|
||||
overlap = int(abs(silence_duration) * max_rate)
|
||||
previous_audio = sequence[-1]
|
||||
overlap = min(
|
||||
overlap,
|
||||
previous_audio.shape[-1],
|
||||
audio["waveform"].shape[-1],
|
||||
)
|
||||
if overlap > 0:
|
||||
overlap_part = (
|
||||
previous_audio[:, :, -overlap:]
|
||||
+ audio["waveform"][:, :, :overlap]
|
||||
)
|
||||
sequence[-1] = previous_audio[:, :, :-overlap]
|
||||
sequence.append(overlap_part)
|
||||
audio["waveform"] = audio["waveform"][:, :, overlap:]
|
||||
|
||||
sequence.append(audio["waveform"])
|
||||
|
||||
sequenced_waveform = torch.cat(sequence, dim=-1)
|
||||
return (
|
||||
{
|
||||
"sample_rate": max_rate,
|
||||
"waveform": sequenced_waveform,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class MTB_AudioResample(MtbAudio):
|
||||
"""Resample audio to a different sample rate."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
"sample_rate": (
|
||||
"INT",
|
||||
{
|
||||
"default": 16000,
|
||||
"min": 1000,
|
||||
"max": 192000,
|
||||
"step": 100,
|
||||
"tooltip": "Target sample rate in Hz. Whisper requires 16000.",
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("resampled_audio",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "resample_audio"
|
||||
|
||||
def resample_audio(
|
||||
self, audio: AudioTensor, sample_rate: int
|
||||
) -> tuple[AudioTensor]:
|
||||
resampled = self.resample(audio, sample_rate)
|
||||
return (resampled,)
|
||||
|
||||
|
||||
class MTB_AudioIsolateSpeaker(MtbAudio):
|
||||
"""Isolate or mute specific speakers using WhisperData"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
"whisper_data": ("WHISPER_CHUNKS",),
|
||||
"target_speaker": ("STRING", {"default": "SPEAKER_00"}),
|
||||
"mode": (["isolate", "mute"], {"default": "isolate"}),
|
||||
"fade_ms": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 100.0,
|
||||
"min": 0.0,
|
||||
"max": 1000.0,
|
||||
"step": 10,
|
||||
"tooltip": "Fade duration in milliseconds to avoid clicks",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("processed_audio",)
|
||||
CATEGORY = "mtb/audio"
|
||||
FUNCTION = "process_audio"
|
||||
|
||||
def process_audio(
|
||||
self,
|
||||
audio: AudioTensor,
|
||||
whisper_data: WhisperData,
|
||||
target_speaker: str,
|
||||
mode: str = "isolate",
|
||||
fade_ms: float = 100.0,
|
||||
) -> tuple[AudioTensor]:
|
||||
fade_samples = int((fade_ms / 1000.0) * audio["sample_rate"])
|
||||
|
||||
mask = (
|
||||
torch.zeros_like(audio["waveform"])
|
||||
if mode == "isolate"
|
||||
else torch.ones_like(audio["waveform"])
|
||||
)
|
||||
|
||||
for chunk in whisper_data["chunks"]:
|
||||
if not chunk.get("speaker"):
|
||||
continue
|
||||
|
||||
speaker_present = target_speaker in chunk["speaker"]
|
||||
if (mode == "isolate" and speaker_present) or (
|
||||
mode == "mute" and not speaker_present
|
||||
):
|
||||
start_sample = int(
|
||||
chunk["timestamp"][0] * audio["sample_rate"]
|
||||
)
|
||||
end_sample = int(chunk["timestamp"][1] * audio["sample_rate"])
|
||||
|
||||
mask[:, start_sample:end_sample] = 1.0
|
||||
|
||||
if fade_samples > 0:
|
||||
fade = torch.linspace(0, 1, fade_samples)
|
||||
|
||||
transitions = torch.where(mask[0, 1:] != mask[0, :-1])[0] + 1
|
||||
|
||||
for trans_idx in transitions:
|
||||
if (
|
||||
trans_idx >= fade_samples
|
||||
and trans_idx <= mask.shape[1] - fade_samples
|
||||
):
|
||||
if mask[0, trans_idx] == 1:
|
||||
mask[:, trans_idx : trans_idx + fade_samples] *= fade
|
||||
else:
|
||||
mask[:, trans_idx - fade_samples : trans_idx] *= (
|
||||
fade.flip(0)
|
||||
)
|
||||
|
||||
processed_waveform = audio["waveform"] * mask
|
||||
|
||||
return (
|
||||
{
|
||||
"sample_rate": audio["sample_rate"],
|
||||
"waveform": processed_waveform,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class MTB_ProcessWhisperDiarization:
|
||||
"""Process Whisper chunks with speaker diarization using either pyannote or NeMo."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"whisper_chunks": ("WHISPER_CHUNKS",),
|
||||
"audio": ("AUDIO",),
|
||||
"backend": (["pyannote", "nemo"], {"default": "pyannote"}),
|
||||
"num_speakers": (
|
||||
"INT",
|
||||
{"default": 2, "min": 1, "max": 10, "step": 1},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WHISPER_CHUNKS",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "mtb/audio"
|
||||
|
||||
def process_pyannote(self, audio, num_speakers, device):
|
||||
"""Process audio using pyannote backend."""
|
||||
try:
|
||||
from pyannote.audio import Pipeline
|
||||
from pyannote.audio.pipelines.utils.hook import ProgressHook
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"pyannote.audio not found. Install with: pip install pyannote.audio"
|
||||
)
|
||||
|
||||
pipeline = Pipeline.from_pretrained(
|
||||
"pyannote/speaker-diarization-3.1", use_auth_token=None
|
||||
)
|
||||
pipeline.to(torch.device(device))
|
||||
with ProgressHook() as hook:
|
||||
diarization = pipeline(
|
||||
{
|
||||
"waveform": audio["waveform"][0],
|
||||
"sample_rate": audio["sample_rate"],
|
||||
},
|
||||
num_speakers=num_speakers,
|
||||
hook=hook,
|
||||
)
|
||||
|
||||
speaker_segments = []
|
||||
for turn, _, speaker in diarization.itertracks(yield_label=True):
|
||||
speaker_segments.append(
|
||||
{
|
||||
"start": turn.start,
|
||||
"end": turn.end,
|
||||
"speaker": speaker,
|
||||
}
|
||||
)
|
||||
|
||||
return speaker_segments
|
||||
|
||||
def process_nemo(self, audio, num_speakers, device):
|
||||
"""Process audio using NeMo backend."""
|
||||
try:
|
||||
import nemo.collections.asr as nemo_asr
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"NeMo not found. Install with: pip install nemo_toolkit[asr]"
|
||||
)
|
||||
|
||||
model = nemo_asr.models.ClusteringDiarizer.from_pretrained(
|
||||
"nvidia/speakerverification_en_titanet_large"
|
||||
).to(device)
|
||||
|
||||
diarization = model.diarize(
|
||||
audio=audio["waveform"][0],
|
||||
sample_rate=audio["sample_rate"],
|
||||
num_speakers=num_speakers,
|
||||
)
|
||||
|
||||
speaker_segments = []
|
||||
for segment in diarization:
|
||||
speaker_segments.append(
|
||||
{
|
||||
"start": segment["start"],
|
||||
"end": segment["end"],
|
||||
"speaker": f"SPEAKER_{segment['speaker']}",
|
||||
}
|
||||
)
|
||||
|
||||
return speaker_segments
|
||||
|
||||
def process(
|
||||
self,
|
||||
whisper_chunks,
|
||||
audio,
|
||||
backend="pyannote",
|
||||
num_speakers=2,
|
||||
device="cuda",
|
||||
):
|
||||
if backend == "pyannote":
|
||||
speaker_segments = self.process_pyannote(
|
||||
audio, num_speakers, device
|
||||
)
|
||||
else: # nemo
|
||||
speaker_segments = self.process_nemo(audio, num_speakers, device)
|
||||
|
||||
for chunk in whisper_chunks["chunks"]:
|
||||
chunk_start, chunk_end = chunk["timestamp"]
|
||||
chunk_speakers = set()
|
||||
for segment in speaker_segments:
|
||||
if (
|
||||
segment["start"] <= chunk_end
|
||||
and segment["end"] >= chunk_start
|
||||
):
|
||||
chunk_speakers.add(segment["speaker"])
|
||||
|
||||
if chunk_speakers:
|
||||
chunk["speaker"] = list(chunk_speakers)[0]
|
||||
else:
|
||||
chunk["speaker"] = "unknown"
|
||||
|
||||
return (whisper_chunks,)
|
||||
|
||||
|
||||
class MTB_AudioDuration:
|
||||
"""Get audio duration in milliseconds."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("duration_ms",)
|
||||
FUNCTION = "get_duration"
|
||||
CATEGORY = "mtb/audio"
|
||||
|
||||
def get_duration(self, audio):
|
||||
waveform = audio["waveform"]
|
||||
sample_rate = audio["sample_rate"]
|
||||
|
||||
duration_ms = int((waveform.shape[-1] / sample_rate) * 1000)
|
||||
log.debug(
|
||||
f"Audio duration: {duration_ms}ms ({duration_ms / 1000:.2f}s)"
|
||||
)
|
||||
|
||||
return (duration_ms,)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
MTB_AudioSequence,
|
||||
MTB_AudioStack,
|
||||
MTB_AudioCut,
|
||||
MTB_AudioResample,
|
||||
MTB_AudioIsolateSpeaker,
|
||||
MTB_LoadWhisper,
|
||||
MTB_AudioToText,
|
||||
MTB_ProcessWhisperOutput,
|
||||
MTB_ProcessWhisperDiarization,
|
||||
MTB_AudioDuration,
|
||||
]
|
||||
+518
-17
@@ -1,12 +1,18 @@
|
||||
import os
|
||||
import random
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
import comfy.utils
|
||||
import cv2
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from ..log import log
|
||||
from ..utils import EASINGS, apply_easing, pil2tensor
|
||||
from ..utils import EASINGS, apply_easing, glob_multiple, pil2tensor
|
||||
from .transform import MTB_TransformImage
|
||||
|
||||
|
||||
@@ -46,7 +52,7 @@ class MTB_BatchFloatMath:
|
||||
for v in vals:
|
||||
if len(v) != ref_count:
|
||||
raise ValueError(
|
||||
f"All values must have the same length (current: {len(v)}, ref: {ref_count}"
|
||||
f"All values must have the same length (current: {len(v)}, ref: {ref_count})"
|
||||
)
|
||||
|
||||
match operation:
|
||||
@@ -169,6 +175,124 @@ class MTB_BatchTimeWrap:
|
||||
return (warped_tensor, interpolated_curve)
|
||||
|
||||
|
||||
class MTB_ImageBatchToSublist:
|
||||
"""
|
||||
# Image Batch To Sublist 🔄
|
||||
|
||||
Splits a large batched tensor into smaller sub-batches for memory-efficient processing.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sub_batch_size": (
|
||||
"INT",
|
||||
{"default": 1, "min": 1, "max": 1000, "step": 1},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "INT")
|
||||
RETURN_NAMES = ("image_list", "mask_list", "item_count")
|
||||
|
||||
OUTPUT_IS_LIST = (True, True)
|
||||
FUNCTION = "split_batch"
|
||||
CATEGORY = "batch_processing"
|
||||
|
||||
def split_batch(
|
||||
self,
|
||||
sub_batch_size: int,
|
||||
image: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
):
|
||||
if image is None and mask is None:
|
||||
raise ValueError(
|
||||
"You must either pass mask or image, none received"
|
||||
)
|
||||
|
||||
image_count = 0
|
||||
if image is not None:
|
||||
image_count = image.size(0)
|
||||
|
||||
mask_count = 0
|
||||
if mask is not None:
|
||||
mask_count = mask.size(0)
|
||||
|
||||
if image_count > 0 and mask_count > 0 and mask_count != image_count:
|
||||
raise ValueError(
|
||||
f"When providing image and mask, batch size must match (got {mask.size(0)} mask and {image.size(0)} images)"
|
||||
)
|
||||
|
||||
batch_size = max(image_count, mask_count)
|
||||
|
||||
num_full_batches = batch_size // sub_batch_size
|
||||
im_batches = []
|
||||
mask_batches = []
|
||||
|
||||
for i in range(num_full_batches):
|
||||
start_idx = i * sub_batch_size
|
||||
end_idx = start_idx + sub_batch_size
|
||||
if image_count > 0:
|
||||
im_batches.append(image[start_idx:end_idx, ...])
|
||||
|
||||
if mask_count > 0:
|
||||
mask_batches.append(mask[start_idx:end_idx, ...])
|
||||
|
||||
if batch_size % sub_batch_size != 0:
|
||||
remaining_start = num_full_batches * sub_batch_size
|
||||
if image_count > 0:
|
||||
im_batches.append(image[remaining_start:, ...])
|
||||
|
||||
if mask_count > 0:
|
||||
mask_batches.append(mask[remaining_start:, ...])
|
||||
|
||||
return (im_batches, mask_batches, len(im_batches))
|
||||
|
||||
|
||||
class MTB_SublistToImageBatch:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"tensors": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "merge_batches"
|
||||
CATEGORY = "batch_processing"
|
||||
DOCUMENTATION = """# Sublist to Image Batch 🔄
|
||||
|
||||
Merges a list of sub-batched tensors back into a single large batch.
|
||||
"""
|
||||
|
||||
def merge_batches(self, tensors: list[torch.Tensor]):
|
||||
if len(tensors) <= 1:
|
||||
return (tensors[0],)
|
||||
|
||||
result = tensors[0]
|
||||
|
||||
for next_tensor in tensors[1:]:
|
||||
if result.shape[1:] != next_tensor.shape[1:]:
|
||||
next_tensor = comfy.utils.common_upscale(
|
||||
next_tensor.movedim(-1, 1),
|
||||
result.shape[2],
|
||||
result.shape[1],
|
||||
"lanczos",
|
||||
"center",
|
||||
).movedim(1, -1)
|
||||
|
||||
result = torch.cat((result, next_tensor), dim=0)
|
||||
|
||||
return (result,)
|
||||
|
||||
|
||||
class MTB_BatchMake:
|
||||
"""Simply duplicates the input frame as a batch"""
|
||||
|
||||
@@ -178,18 +302,22 @@ class MTB_BatchMake:
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"count": ("INT", {"default": 1}),
|
||||
}
|
||||
},
|
||||
"optional": {"mask": ("MASK",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate_batch"
|
||||
CATEGORY = "mtb/batch"
|
||||
|
||||
def generate_batch(self, image: torch.Tensor, count):
|
||||
def generate_batch(self, image: torch.Tensor, count, mask=None):
|
||||
if len(image.shape) == 3:
|
||||
image = image.unsqueeze(0)
|
||||
|
||||
return (image.repeat(count, 1, 1, 1),)
|
||||
return (
|
||||
image.repeat(count, 1, 1, 1),
|
||||
mask.repeat(count, 1, 1) if mask else mask,
|
||||
)
|
||||
|
||||
|
||||
class MTB_BatchShape:
|
||||
@@ -374,8 +502,14 @@ class MTB_BatchFloat:
|
||||
{"default": "Steps"},
|
||||
),
|
||||
"count": ("INT", {"default": 2}),
|
||||
"min": ("FLOAT", {"default": 0.0, "step": 0.001}),
|
||||
"max": ("FLOAT", {"default": 1.0, "step": 0.001}),
|
||||
"min": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": -1e4, "max": 1e4, "step": 0.001},
|
||||
),
|
||||
"max": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": -1e4, "max": 1e4, "step": 0.001},
|
||||
),
|
||||
"easing": (
|
||||
[
|
||||
"Linear",
|
||||
@@ -410,7 +544,14 @@ class MTB_BatchFloat:
|
||||
RETURN_TYPES = ("FLOATS",)
|
||||
CATEGORY = "mtb/batch"
|
||||
|
||||
def set_floats(self, mode, count, min, max, easing):
|
||||
def set_floats(
|
||||
self,
|
||||
mode: Literal["Steps"] | Literal["Single"] = "Steps",
|
||||
count: int = 1,
|
||||
min: float = 0.0, # noqa: A002
|
||||
max: float = 1.0, # noqa: A002
|
||||
easing: str = "Linear",
|
||||
):
|
||||
if mode == "Steps" and count == 1:
|
||||
raise ValueError(
|
||||
"Steps mode requires at least a count of 2 values"
|
||||
@@ -429,6 +570,210 @@ class MTB_BatchFloat:
|
||||
return (keyframes,)
|
||||
|
||||
|
||||
class MTB_BatchSequencePlus:
|
||||
"""Sequences multiple image batches with transition effects."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"transition": (
|
||||
[
|
||||
"none",
|
||||
"crossfade",
|
||||
"slide_left",
|
||||
"slide_right",
|
||||
"slide_up",
|
||||
"slide_down",
|
||||
"wipe_left",
|
||||
"wipe_right",
|
||||
"wipe_up",
|
||||
"wipe_down",
|
||||
"band_wipe_h",
|
||||
"band_wipe_v",
|
||||
],
|
||||
{"default": "none"},
|
||||
),
|
||||
"overlap_frames": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 120, "step": 1},
|
||||
),
|
||||
"reverse": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "sequence_batches"
|
||||
CATEGORY = "mtb/batch"
|
||||
|
||||
def apply_transition(
|
||||
self,
|
||||
frame1: torch.Tensor,
|
||||
frame2: torch.Tensor,
|
||||
transition: str,
|
||||
progress: float,
|
||||
):
|
||||
"""Apply transition effect between two frames."""
|
||||
if transition == "none":
|
||||
return frame1 if progress < 0.5 else frame2
|
||||
|
||||
elif transition == "crossfade":
|
||||
return frame1 * (1 - progress) + frame2 * progress
|
||||
|
||||
elif transition.startswith("slide_"):
|
||||
h, w = frame1.shape[1:3]
|
||||
if transition == "slide_left":
|
||||
offset = int(w * progress)
|
||||
frame2 = torch.roll(frame2, shifts=-offset, dims=2)
|
||||
elif transition == "slide_right":
|
||||
offset = int(w * progress)
|
||||
frame2 = torch.roll(frame2, shifts=offset, dims=2)
|
||||
elif transition == "slide_up":
|
||||
offset = int(h * progress)
|
||||
frame2 = torch.roll(frame2, shifts=-offset, dims=1)
|
||||
elif transition == "slide_down":
|
||||
offset = int(h * progress)
|
||||
frame2 = torch.roll(frame2, shifts=offset, dims=1)
|
||||
return frame1 * (1 - progress) + frame2 * progress
|
||||
|
||||
elif transition.startswith("wipe_"):
|
||||
h, w = frame1.shape[1:3]
|
||||
mask = torch.zeros_like(frame1)
|
||||
if transition == "wipe_left":
|
||||
edge = int(w * progress)
|
||||
mask[:, :, :edge, :] = 1
|
||||
elif transition == "wipe_right":
|
||||
edge = int(w * (1 - progress))
|
||||
mask[:, :, edge:, :] = 1
|
||||
elif transition == "wipe_up":
|
||||
edge = int(h * progress)
|
||||
mask[:, :edge, :, :] = 1
|
||||
elif transition == "wipe_down":
|
||||
edge = int(h * (1 - progress))
|
||||
mask[:, edge:, :, :] = 1
|
||||
return frame1 * (1 - mask) + frame2 * mask
|
||||
|
||||
elif transition.startswith("band_wipe_"):
|
||||
h, w = frame1.shape[1:3]
|
||||
mask = torch.zeros_like(frame1)
|
||||
num_bands = 10 # Number of bands
|
||||
|
||||
if transition == "band_wipe_h":
|
||||
band_width = w / num_bands
|
||||
for i in range(num_bands):
|
||||
edge = int((w * progress) - (i * band_width))
|
||||
start = int(i * band_width)
|
||||
end = int(min(start + edge, (i + 1) * band_width))
|
||||
if end > start:
|
||||
mask[:, :, start:end, :] = 1
|
||||
else: # band_wipe_v
|
||||
band_height = h / num_bands
|
||||
for i in range(num_bands):
|
||||
edge = int((h * progress) - (i * band_height))
|
||||
start = int(i * band_height)
|
||||
end = int(min(start + edge, (i + 1) * band_height))
|
||||
if end > start:
|
||||
mask[:, start:end, :, :] = 1
|
||||
|
||||
return frame1 * (1 - mask) + frame2 * mask
|
||||
|
||||
return frame1
|
||||
|
||||
def sequence_batches(
|
||||
self, transition: str, overlap_frames: int, reverse: bool, **kwargs
|
||||
):
|
||||
images: list[torch.Tensor] = list(kwargs.values())
|
||||
|
||||
if reverse:
|
||||
images = images[::-1]
|
||||
|
||||
processed_images: list[torch.Tensor] = []
|
||||
for img in images:
|
||||
if len(img.shape) == 3:
|
||||
img = img.unsqueeze(0)
|
||||
processed_images.append(img)
|
||||
|
||||
if overlap_frames == 0 or transition == "none":
|
||||
return (torch.cat(processed_images, dim=0),)
|
||||
|
||||
result_frames: list[torch.Tensor] = []
|
||||
|
||||
if len(processed_images) > 0:
|
||||
result_frames.extend(
|
||||
list(processed_images[0][: -overlap_frames // 2])
|
||||
)
|
||||
|
||||
for i in range(1, len(processed_images)):
|
||||
prev_batch = processed_images[i - 1]
|
||||
curr_batch = processed_images[i]
|
||||
|
||||
prev_frames = min(overlap_frames // 2, len(prev_batch))
|
||||
next_frames = min(overlap_frames // 2, len(curr_batch))
|
||||
total_overlap = prev_frames + next_frames
|
||||
|
||||
if total_overlap < 2:
|
||||
# when not enough frames for transition, just concatenate
|
||||
result_frames.extend(list(prev_batch[-prev_frames:]))
|
||||
result_frames.extend(list(curr_batch[:next_frames]))
|
||||
continue
|
||||
|
||||
for t in range(total_overlap):
|
||||
progress = t / (total_overlap - 1)
|
||||
|
||||
prev_idx = (
|
||||
len(prev_batch) - prev_frames + min(t, prev_frames - 1)
|
||||
)
|
||||
next_idx = max(0, t - prev_frames)
|
||||
|
||||
transition_frame = self.apply_transition(
|
||||
prev_batch[prev_idx : prev_idx + 1],
|
||||
curr_batch[next_idx : next_idx + 1],
|
||||
transition,
|
||||
progress,
|
||||
)
|
||||
result_frames.append(transition_frame[0])
|
||||
|
||||
if i < len(processed_images) - 1:
|
||||
result_frames.extend(
|
||||
list(curr_batch[next_frames : -overlap_frames // 2])
|
||||
)
|
||||
else:
|
||||
result_frames.extend(list(curr_batch[next_frames:]))
|
||||
|
||||
result = torch.stack(result_frames, dim=0)
|
||||
|
||||
return (result,)
|
||||
|
||||
|
||||
class MTB_BatchSequence:
|
||||
"""Sequences multiple image batches one after another"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"reverse": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "sequence_batches"
|
||||
CATEGORY = "mtb/batch"
|
||||
|
||||
def sequence_batches(self, reverse: bool, **kwargs):
|
||||
images = list(kwargs.values())
|
||||
if reverse:
|
||||
images = images[::-1]
|
||||
|
||||
processed = []
|
||||
for img in images:
|
||||
if len(img.shape) == 3:
|
||||
img = img.unsqueeze(0)
|
||||
processed.append(img)
|
||||
|
||||
return (torch.cat(processed, dim=0),)
|
||||
|
||||
|
||||
class MTB_BatchMerge:
|
||||
"""Merges multiple image batches with different frame counts"""
|
||||
|
||||
@@ -505,6 +850,13 @@ class MTB_Batch2dTransform:
|
||||
"zoom": ("FLOATS",),
|
||||
"angle": ("FLOATS",),
|
||||
"shear": ("FLOATS",),
|
||||
"use_normalized": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "If true, transform values will be scaled to image dimensions.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -533,6 +885,7 @@ class MTB_Batch2dTransform:
|
||||
zoom: list[float] | None = None,
|
||||
angle: list[float] | None = None,
|
||||
shear: list[float] | None = None,
|
||||
use_normalized: bool = False,
|
||||
):
|
||||
if all(
|
||||
self.get_num_elements(param) <= 0
|
||||
@@ -584,6 +937,7 @@ class MTB_Batch2dTransform:
|
||||
keyframes["shear"][i],
|
||||
border_handling,
|
||||
constant_color,
|
||||
use_normalized=use_normalized,
|
||||
)[0]
|
||||
for i in range(image.shape[0])
|
||||
]
|
||||
@@ -711,7 +1065,9 @@ class MTB_PlotBatchFloat:
|
||||
ax.set_xlim(1, max_length) # Set X-axis limits
|
||||
np.random.seed(seed)
|
||||
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
|
||||
for color, (label, values) in zip(colors, kwargs.items()):
|
||||
for color, (label, values) in zip(
|
||||
colors, kwargs.items(), strict=False
|
||||
):
|
||||
ax.plot(x_values[: len(values)], values, label=label, color=color)
|
||||
ax.legend(
|
||||
title="Legend",
|
||||
@@ -1025,18 +1381,163 @@ class MTB_BatchShake:
|
||||
return (shaken_images, x_translations, y_translations, rotations)
|
||||
|
||||
|
||||
class MTB_BatchFromFolder:
|
||||
"""Load images from a folder with options for latest, oldest, or random selection."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"enable": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Enable or disable the node. If disabled, returns passthrough_image or an empty tensor.",
|
||||
},
|
||||
),
|
||||
"folder_path": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Path to the folder containing images. Relative paths are resolved to the ComfyUI output directory.",
|
||||
},
|
||||
),
|
||||
"mode": (
|
||||
["latest", "oldest", "random"],
|
||||
{
|
||||
"default": "latest",
|
||||
"tooltip": "How to select images: latest, oldest, or random.",
|
||||
},
|
||||
),
|
||||
"count": (
|
||||
"INT",
|
||||
{
|
||||
"default": 10,
|
||||
"min": 1,
|
||||
"max": 1000,
|
||||
"tooltip": "Number of images to load from the folder.",
|
||||
},
|
||||
),
|
||||
"filter": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "*",
|
||||
"tooltip": "Glob filter for image filenames (e.g. *.png).",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"passthrough_image": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": "If provided and node is disabled, this image is passed through instead of returning an empty tensor."
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
CATEGORY = "mtb/batch"
|
||||
FUNCTION = "load_from_folder"
|
||||
|
||||
def load_from_folder(
|
||||
self,
|
||||
enable: bool,
|
||||
folder_path: str,
|
||||
mode: str,
|
||||
count: int,
|
||||
filter: str,
|
||||
passthrough_image=None,
|
||||
):
|
||||
"""Load images from a folder with the specified selection mode."""
|
||||
if not enable:
|
||||
if passthrough_image is not None:
|
||||
log.debug(
|
||||
"MTB_BatchFromFolder: Using passthrough image (disabled)"
|
||||
)
|
||||
return (passthrough_image,)
|
||||
log.debug(
|
||||
"MTB_BatchFromFolder: Disabled and no passthrough_image provided, returning empty tensor"
|
||||
)
|
||||
return (torch.zeros(0, 0, 0, 3),)
|
||||
|
||||
path_obj = Path(folder_path)
|
||||
if not path_obj.is_absolute():
|
||||
output_dir = Path(folder_paths.get_output_directory())
|
||||
path_obj = output_dir / folder_path
|
||||
path_obj = path_obj.resolve()
|
||||
|
||||
if not path_obj.exists():
|
||||
log.error(f"Folder path does not exist: {path_obj}")
|
||||
return (torch.zeros(0, 0, 0, 3),)
|
||||
|
||||
if not path_obj.is_dir():
|
||||
log.error(f"Path is not a directory: {path_obj}")
|
||||
return (torch.zeros(0, 0, 0, 3),)
|
||||
|
||||
patterns = [filter] if filter else ["*"]
|
||||
files = glob_multiple(path_obj, patterns)
|
||||
|
||||
image_extensions = [".png", ".jpg", ".jpeg", ".bmp", ".webp", ".tiff"]
|
||||
image_files = [
|
||||
f for f in files if f.suffix.lower() in image_extensions
|
||||
]
|
||||
|
||||
if not image_files:
|
||||
log.warning(
|
||||
f"No image files found in {path_obj} with filter {filter}"
|
||||
)
|
||||
return (torch.zeros(0, 0, 0, 3),)
|
||||
|
||||
if mode == "latest":
|
||||
image_files.sort(key=lambda x: os.path.getmtime(x), reverse=True)
|
||||
elif mode == "oldest":
|
||||
image_files.sort(key=lambda x: os.path.getmtime(x))
|
||||
elif mode == "random":
|
||||
random.shuffle(image_files)
|
||||
|
||||
selected_files = image_files[:count]
|
||||
|
||||
if len(selected_files) < count:
|
||||
log.warning(
|
||||
f"Requested {count} images but only found {len(selected_files)}"
|
||||
)
|
||||
|
||||
loaded_images = []
|
||||
for file_path in selected_files:
|
||||
try:
|
||||
img = Image.open(file_path)
|
||||
if img.mode != "RGB":
|
||||
img = img.convert("RGB")
|
||||
loaded_images.append(img)
|
||||
except Exception as e:
|
||||
log.error(f"Error loading image {file_path}: {e}")
|
||||
|
||||
if not loaded_images:
|
||||
log.error("Failed to load any images")
|
||||
return (torch.zeros(0, 0, 0, 3),)
|
||||
|
||||
return (pil2tensor(loaded_images),)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
MTB_BatchFloat,
|
||||
MTB_Batch2dTransform,
|
||||
MTB_BatchShape,
|
||||
MTB_BatchMake,
|
||||
MTB_BatchFloat,
|
||||
MTB_BatchFloatAssemble,
|
||||
MTB_BatchFloatFill,
|
||||
MTB_BatchFloatNormalize,
|
||||
MTB_BatchMerge,
|
||||
MTB_BatchShake,
|
||||
MTB_PlotBatchFloat,
|
||||
MTB_BatchTimeWrap,
|
||||
MTB_BatchFloatFit,
|
||||
MTB_BatchFloatMath,
|
||||
MTB_BatchFloatNormalize,
|
||||
MTB_BatchFromFolder,
|
||||
MTB_BatchMake,
|
||||
MTB_BatchMerge,
|
||||
MTB_BatchSequence,
|
||||
MTB_BatchSequencePlus,
|
||||
MTB_BatchShake,
|
||||
MTB_BatchShape,
|
||||
MTB_BatchTimeWrap,
|
||||
MTB_PlotBatchFloat,
|
||||
MTB_SublistToImageBatch,
|
||||
MTB_ImageBatchToSublist,
|
||||
]
|
||||
|
||||
+124
-2
@@ -3,10 +3,127 @@ import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
|
||||
from ..log import log
|
||||
from ..utils import here
|
||||
|
||||
Conditioning = list[tuple[torch.Tensor, dict[str, torch.Tensor]]]
|
||||
|
||||
|
||||
def check_condition(conditioning: Conditioning):
|
||||
has_cn = False
|
||||
if len(conditioning) > 1:
|
||||
log.warn(
|
||||
"More than one conditioning was provided. Only the first one will be used."
|
||||
)
|
||||
first = conditioning[0]
|
||||
cond, kwargs = first
|
||||
|
||||
log.debug("Conditioning Shape")
|
||||
log.debug(cond.shape)
|
||||
log.debug("Conditioning keys")
|
||||
log.debug([f"\t{k} - {type(kwargs[k])}" for k in kwargs])
|
||||
if "control" in kwargs:
|
||||
log.debug("Conditioning contains a controlnet")
|
||||
has_cn = True
|
||||
if "pooled_output" not in kwargs:
|
||||
raise ValueError(
|
||||
"Conditioning is not valid. Missing 'pooled_output' key."
|
||||
)
|
||||
return has_cn
|
||||
|
||||
|
||||
class MTB_InterpolateCondition:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"blend": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "mtb/conditioning"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(
|
||||
self, blend: float, **kwargs: Conditioning
|
||||
) -> tuple[Conditioning]:
|
||||
blend = max(0.0, min(1.0, blend))
|
||||
|
||||
conditions: list[Conditioning] = list(kwargs.values())
|
||||
num_conditions = len(conditions)
|
||||
|
||||
if num_conditions < 2:
|
||||
raise ValueError("At least two conditioning inputs are required.")
|
||||
|
||||
segment_length = 1.0 / (num_conditions - 1)
|
||||
|
||||
segment_index = min(int(blend // segment_length), num_conditions - 2)
|
||||
|
||||
local_blend = (
|
||||
blend - (segment_index * segment_length)
|
||||
) / segment_length
|
||||
|
||||
cond_from = conditions[segment_index]
|
||||
cond_to = conditions[segment_index + 1]
|
||||
|
||||
from_cn = check_condition(cond_from)
|
||||
to_cn = check_condition(cond_to)
|
||||
|
||||
if from_cn and to_cn:
|
||||
raise ValueError(
|
||||
"Interpolating conditions cannot both contain ControlNets"
|
||||
)
|
||||
|
||||
try:
|
||||
interpolated_condition = [
|
||||
(1.0 - local_blend) * c_from + local_blend * c_to
|
||||
for c_from, c_to in zip(
|
||||
cond_from[0][0], cond_to[0][0], strict=False
|
||||
)
|
||||
]
|
||||
except Exception as e:
|
||||
print(f"Error during interpolation: {e}")
|
||||
raise
|
||||
|
||||
pooled_from = cond_from[0][1].get(
|
||||
"pooled_output",
|
||||
torch.zeros_like(
|
||||
next(iter(cond_from[0][1].values()), torch.tensor([]))
|
||||
),
|
||||
)
|
||||
|
||||
pooled_to = cond_to[0][1].get(
|
||||
"pooled_output",
|
||||
torch.zeros_like(
|
||||
next(iter(cond_from[0][1].values()), torch.tensor([]))
|
||||
),
|
||||
)
|
||||
|
||||
interpolated_pooled = (
|
||||
1.0 - local_blend
|
||||
) * pooled_from + local_blend * pooled_to
|
||||
|
||||
res = {"pooled_output": interpolated_pooled}
|
||||
|
||||
if from_cn:
|
||||
res["control"] = cond_from[0][1]["control"]
|
||||
res["control_apply_to_uncond"] = cond_from[0][1][
|
||||
"control_apply_to_uncond"
|
||||
]
|
||||
if to_cn:
|
||||
res["control"] = cond_to[0][1]["control"]
|
||||
res["control_apply_to_uncond"] = cond_to[0][1][
|
||||
"control_apply_to_uncond"
|
||||
]
|
||||
|
||||
return ([(torch.stack(interpolated_condition), res)],)
|
||||
|
||||
|
||||
class MTB_InterpolateClipSequential:
|
||||
@classmethod
|
||||
@@ -177,7 +294,7 @@ class MTB_StylesLoader:
|
||||
with open(file, encoding="utf8") as f:
|
||||
parsed = csv.reader(f)
|
||||
for i, row in enumerate(parsed):
|
||||
log.debug(f"Adding style {row[0]}")
|
||||
# log.debug(f"Adding style {row[0]}")
|
||||
try:
|
||||
name, positive, negative = (row + [None] * 3)[:3]
|
||||
positive = positive or ""
|
||||
@@ -213,4 +330,9 @@ class MTB_StylesLoader:
|
||||
return (self.options[style_name][0], self.options[style_name][1])
|
||||
|
||||
|
||||
__nodes__ = [MTB_SmartStep, MTB_StylesLoader, MTB_InterpolateClipSequential]
|
||||
__nodes__ = [
|
||||
MTB_SmartStep,
|
||||
MTB_StylesLoader,
|
||||
MTB_InterpolateClipSequential,
|
||||
MTB_InterpolateCondition,
|
||||
]
|
||||
|
||||
+1
-1
@@ -24,4 +24,4 @@ class MTB_Constant:
|
||||
return (kwargs.get("Value"),)
|
||||
|
||||
|
||||
__nodes__ = [MTB_Constant]
|
||||
# __nodes__ = [MTB_Constant]
|
||||
|
||||
+116
-1
@@ -41,6 +41,57 @@ class MTB_Bbox:
|
||||
return ((x, y, width, height),)
|
||||
|
||||
|
||||
class MTB_SplitBbox:
|
||||
"""Split the components of a bbox"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"bbox": ("BBOX",)},
|
||||
}
|
||||
|
||||
CATEGORY = "mtb/crop"
|
||||
FUNCTION = "split_bbox"
|
||||
RETURN_TYPES = ("INT", "INT", "INT", "INT")
|
||||
RETURN_NAMES = ("x", "y", "width", "height")
|
||||
|
||||
def split_bbox(self, bbox):
|
||||
return (bbox[0], bbox[1], bbox[2], bbox[3])
|
||||
|
||||
|
||||
class MTB_UpscaleBboxBy:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"bbox": ("BBOX",),
|
||||
"scale": ("FLOAT", {"default": 1.0}),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "mtb/crop"
|
||||
RETURN_TYPES = ("BBOX",)
|
||||
|
||||
FUNCTION = "upscale"
|
||||
|
||||
def upscale(
|
||||
self, bbox: tuple[int, int, int, int], scale: float
|
||||
) -> tuple[tuple[int, int, int, int]]:
|
||||
x, y, width, height = bbox
|
||||
|
||||
center_x = x + width // 2
|
||||
center_y = y + height // 2
|
||||
|
||||
new_width = int(width * scale)
|
||||
new_height = int(height * scale)
|
||||
|
||||
new_x = center_x - new_width // 2
|
||||
new_y = center_y - new_height // 2
|
||||
|
||||
scaled = (new_x, new_y, new_width, new_height)
|
||||
return (scaled,)
|
||||
|
||||
|
||||
class MTB_BboxFromMask:
|
||||
"""From a mask extract the bounding box"""
|
||||
|
||||
@@ -324,4 +375,68 @@ class MTB_Uncrop:
|
||||
return (pil2tensor(out_images),)
|
||||
|
||||
|
||||
__nodes__ = [MTB_BboxFromMask, MTB_Bbox, MTB_Crop, MTB_Uncrop]
|
||||
class MTB_BBoxForceDimensions:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"bbox": ("BBOX",),
|
||||
"width": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||
"height": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "mtb/crop"
|
||||
RETURN_TYPES = ("BBOX",)
|
||||
FUNCTION = "force_dimensions"
|
||||
|
||||
def force_dimensions(
|
||||
self,
|
||||
bbox: tuple[int, int, int, int],
|
||||
width: int,
|
||||
height: int,
|
||||
image: torch.Tensor = None,
|
||||
) -> tuple[tuple[int, int, int, int]]:
|
||||
x, y, curr_width, curr_height = bbox
|
||||
|
||||
center_x = x + curr_width // 2
|
||||
center_y = y + curr_height // 2
|
||||
|
||||
new_x = center_x - width // 2
|
||||
new_y = center_y - height // 2
|
||||
|
||||
if image is not None:
|
||||
img_height, img_width = image.shape[1:3]
|
||||
x_overflow = max(0, new_x + width - img_width) + min(0, new_x)
|
||||
y_overflow = max(0, new_y + height - img_height) + min(0, new_y)
|
||||
if width > img_width or height > img_height:
|
||||
x_exceed = width - img_width if width > img_width else 0
|
||||
y_exceed = height - img_height if height > img_height else 0
|
||||
raise ValueError(
|
||||
f"Target bbox dimensions ({width}x{height}) exceed image bounds ({img_width}x{img_height}) "
|
||||
f"by {x_exceed}px horizontally and {y_exceed}px vertically"
|
||||
)
|
||||
|
||||
if x_overflow > 0 or x_overflow < 0:
|
||||
new_x -= x_overflow
|
||||
|
||||
if y_overflow > 0:
|
||||
new_y -= y_overflow
|
||||
elif y_overflow < 0:
|
||||
new_y -= y_overflow # Add the negative overflow
|
||||
|
||||
return ((int(new_x), int(new_y), width, height),)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
MTB_BboxFromMask,
|
||||
MTB_Bbox,
|
||||
MTB_Crop,
|
||||
MTB_Uncrop,
|
||||
MTB_SplitBbox,
|
||||
MTB_UpscaleBboxBy,
|
||||
MTB_BBoxForceDimensions,
|
||||
]
|
||||
|
||||
+101
-24
@@ -2,7 +2,6 @@ import base64
|
||||
import io
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
@@ -11,13 +10,66 @@ from ..log import log
|
||||
from ..utils import tensor2pil
|
||||
|
||||
|
||||
def get_detailed_type_info(obj):
|
||||
type_info = []
|
||||
|
||||
type_name = type(obj).__name__
|
||||
type_info.append(f"Type: {type_name}")
|
||||
|
||||
if isinstance(obj, torch.Tensor):
|
||||
type_info.extend(
|
||||
[
|
||||
f"Shape: {obj.shape}",
|
||||
f"Dtype: {obj.dtype}",
|
||||
f"Device: {obj.device}",
|
||||
f"Requires grad: {obj.requires_grad}",
|
||||
f"Stride: {obj.stride()}",
|
||||
f"Contiguous: {obj.is_contiguous()}",
|
||||
]
|
||||
)
|
||||
elif isinstance(obj, (list, tuple)):
|
||||
type_info.extend(
|
||||
[
|
||||
f"Length: {len(obj)}",
|
||||
f"Container type: {type_name}",
|
||||
]
|
||||
)
|
||||
if obj:
|
||||
type_info.append(f"Element type: {type(obj[0]).__name__}")
|
||||
elif isinstance(obj, dict):
|
||||
type_info.extend(
|
||||
[
|
||||
f"Length: {len(obj)}",
|
||||
f"Keys: {list(obj.keys())}",
|
||||
]
|
||||
)
|
||||
elif hasattr(obj, "__dict__"):
|
||||
attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
|
||||
type_info.append(f"Attributes: {attributes}")
|
||||
|
||||
return type_info
|
||||
|
||||
|
||||
# region processors
|
||||
def process_tensor(tensor):
|
||||
def process_tensor(tensor: torch.Tensor, as_type=False):
|
||||
log.debug(f"Tensor: {tensor.shape}")
|
||||
|
||||
if as_type:
|
||||
return {
|
||||
"text": [f"Tensor of shape {tensor.shape} of type {tensor.dtype}"]
|
||||
}
|
||||
|
||||
is_mask = len(tensor.shape) == 3
|
||||
|
||||
if is_mask:
|
||||
tensor = tensor.unsqueeze(-1).repeat(1, 1, 1, 3)
|
||||
|
||||
image = tensor2pil(tensor)
|
||||
b64_imgs = []
|
||||
for im in image:
|
||||
if is_mask:
|
||||
im = im.convert("L")
|
||||
|
||||
buffered = io.BytesIO()
|
||||
im.save(buffered, format="PNG")
|
||||
b64_imgs.append(
|
||||
@@ -28,11 +80,16 @@ def process_tensor(tensor):
|
||||
return {"b64_images": b64_imgs}
|
||||
|
||||
|
||||
def process_list(anything):
|
||||
def process_list(anything, as_type=False):
|
||||
text = []
|
||||
if not anything:
|
||||
return {"text": []}
|
||||
|
||||
if as_type:
|
||||
type_info = get_detailed_type_info(anything)
|
||||
type_info.extend(get_detailed_type_info(anything[0]))
|
||||
return {"text": type_info}
|
||||
|
||||
first_element = anything[0]
|
||||
if (
|
||||
isinstance(first_element, list)
|
||||
@@ -54,25 +111,41 @@ def process_list(anything):
|
||||
return {"text": text}
|
||||
|
||||
|
||||
def process_dict(anything):
|
||||
def process_dict(anything, as_type=False):
|
||||
text = []
|
||||
if as_type:
|
||||
return {"text": get_detailed_type_info(anything)}
|
||||
|
||||
if "samples" in anything:
|
||||
is_empty = (
|
||||
"(empty)" if torch.count_nonzero(anything["samples"]) == 0 else ""
|
||||
)
|
||||
text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
|
||||
|
||||
elif "waveform" in anything:
|
||||
is_empty = (
|
||||
"(empty) " if torch.count_nonzero(anything["samples"]) == 0 else ""
|
||||
)
|
||||
|
||||
text.append(
|
||||
f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}"
|
||||
)
|
||||
|
||||
else:
|
||||
log.debug(f"Unhandled dict: {anything.keys()}")
|
||||
text.append(json.dumps(anything, indent=2))
|
||||
|
||||
return {"text": text}
|
||||
|
||||
|
||||
def process_bool(anything):
|
||||
def process_bool(anything, as_type=False):
|
||||
return {"text": ["True" if anything else "False"]}
|
||||
|
||||
|
||||
def process_text(anything):
|
||||
def process_text(anything, as_type=False):
|
||||
if as_type:
|
||||
return {"text": get_detailed_type_info(anything)}
|
||||
|
||||
return {"text": [str(anything)]}
|
||||
|
||||
|
||||
@@ -89,6 +162,7 @@ class MTB_Debug:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"output_to_console": ("BOOLEAN", {"default": False})},
|
||||
"optional": {"as_detailed_types": ("BOOLEAN", {"default": False})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
@@ -96,29 +170,25 @@ class MTB_Debug:
|
||||
CATEGORY = "mtb/debug"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def do_debug(self, output_to_console: bool, **kwargs):
|
||||
output = {
|
||||
"ui": {"b64_images": [], "text": []},
|
||||
# "result": ("A"),
|
||||
}
|
||||
def do_debug(
|
||||
self, output_to_console: bool, as_detailed_types: bool, **kwargs
|
||||
):
|
||||
output = {"ui": {"items": []}}
|
||||
|
||||
processors = {
|
||||
torch.Tensor: process_tensor,
|
||||
list: process_list,
|
||||
dict: process_dict,
|
||||
bool: process_bool,
|
||||
}
|
||||
if output_to_console:
|
||||
for k, v in kwargs.items():
|
||||
log.info(f"{k}: {v}")
|
||||
|
||||
for anything in kwargs.values():
|
||||
for input_name, anything in kwargs.items():
|
||||
processor = processors.get(type(anything), process_text)
|
||||
|
||||
processed_data = processor(anything)
|
||||
processed = processor(anything, as_detailed_types)
|
||||
|
||||
for ui_key, ui_value in processed_data.items():
|
||||
output["ui"][ui_key].extend(ui_value)
|
||||
item = {
|
||||
"input": input_name,
|
||||
**processed,
|
||||
}
|
||||
output["ui"]["items"].append(item)
|
||||
|
||||
return output
|
||||
|
||||
@@ -154,9 +224,9 @@ class MTB_SaveTensors:
|
||||
def save(
|
||||
self,
|
||||
filename_prefix,
|
||||
image: Optional[torch.Tensor] = None,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
latent: Optional[torch.Tensor] = None,
|
||||
image: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
latent: torch.Tensor | None = None,
|
||||
):
|
||||
(
|
||||
full_output_folder,
|
||||
@@ -188,4 +258,11 @@ class MTB_SaveTensors:
|
||||
return f"{filename_prefix}_{counter:05}"
|
||||
|
||||
|
||||
processors = {
|
||||
torch.Tensor: process_tensor,
|
||||
list: process_list,
|
||||
dict: process_dict,
|
||||
bool: process_bool,
|
||||
}
|
||||
|
||||
__nodes__ = [MTB_Debug, MTB_SaveTensors]
|
||||
|
||||
+45
-5
@@ -2,13 +2,16 @@ import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
# torch must be imported prior to onnx for the CUDAProvider.
|
||||
import torch # isort:skip
|
||||
import onnxruntime as ort
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from ..errors import ModelNotFound
|
||||
from ..log import mklog
|
||||
from ..utils import (
|
||||
download_model,
|
||||
get_model_path,
|
||||
tensor2pil,
|
||||
tiles_infer,
|
||||
@@ -23,7 +26,12 @@ log = mklog(__name__)
|
||||
|
||||
# - COLOR to NORMALS
|
||||
def color_to_normals(
|
||||
color_img, overlap, progress_callback, *, save_temp=False
|
||||
color_img,
|
||||
overlap,
|
||||
progress_callback,
|
||||
*,
|
||||
save_temp=False,
|
||||
auto_download=False,
|
||||
):
|
||||
"""Compute a normal map from the given color map.
|
||||
|
||||
@@ -67,9 +75,34 @@ def color_to_normals(
|
||||
log.debug("DeepBump Color → Normals : loading model")
|
||||
model = get_model_path("deepbump", "deepbump256.onnx")
|
||||
if not model or not model.exists():
|
||||
raise ModelNotFound(f"deepbump ({model})")
|
||||
if not auto_download:
|
||||
raise ModelNotFound(f"deepbump ({model})")
|
||||
log.debug("Downloading models...")
|
||||
download_model(
|
||||
"https://github.com/HugoTini/DeepBump/raw/master/deepbump256.onnx",
|
||||
"deepbump",
|
||||
)
|
||||
|
||||
ort_session = ort.InferenceSession(model)
|
||||
providers = [
|
||||
"TensorrtExecutionProvider",
|
||||
"CUDAExecutionProvider",
|
||||
"CoreMLProvider",
|
||||
"CPUExecutionProvider",
|
||||
]
|
||||
available_providers = [
|
||||
provider
|
||||
for provider in providers
|
||||
if provider in ort.get_available_providers()
|
||||
]
|
||||
|
||||
if not available_providers:
|
||||
raise RuntimeError(
|
||||
"No valid ONNX Runtime providers available on this machine."
|
||||
)
|
||||
log.debug(f"Using ONNX providers: {available_providers}")
|
||||
ort_session = ort.InferenceSession(
|
||||
model.as_posix(), providers=available_providers
|
||||
)
|
||||
|
||||
# Predict normal map for each tile
|
||||
log.debug("DeepBump Color → Normals : generating")
|
||||
@@ -332,6 +365,9 @@ class MTB_DeepBump:
|
||||
),
|
||||
"normals_to_height_seamless": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"auto_download": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -347,6 +383,7 @@ class MTB_DeepBump:
|
||||
color_to_normals_overlap="SMALL",
|
||||
normals_to_curvature_blur_radius="SMALL",
|
||||
normals_to_height_seamless=True,
|
||||
auto_download=False,
|
||||
):
|
||||
images = tensor2pil(image)
|
||||
out_images = []
|
||||
@@ -361,7 +398,10 @@ class MTB_DeepBump:
|
||||
# Apply processing
|
||||
if mode == "Color to Normals":
|
||||
out_img = color_to_normals(
|
||||
in_img, color_to_normals_overlap, None
|
||||
in_img,
|
||||
color_to_normals_overlap,
|
||||
None,
|
||||
auto_download=auto_download,
|
||||
)
|
||||
if mode == "Normals to Curvature":
|
||||
out_img = normals_to_curvature(
|
||||
|
||||
+25
-5
@@ -78,6 +78,7 @@ class MTB_LoadFaceEnhanceModel:
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, model_name, upscale=2, bg_upsampler=None):
|
||||
from gfpgan import GFPGANer
|
||||
@@ -163,6 +164,7 @@ class MTB_RestoreFace:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "restore"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -177,7 +179,10 @@ class MTB_RestoreFace:
|
||||
# Adjustable weights
|
||||
"weight": ("FLOAT", {"default": 0.5}),
|
||||
"save_tmp_steps": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"preserve_alpha": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
def do_restore(
|
||||
@@ -188,11 +193,19 @@ class MTB_RestoreFace:
|
||||
only_center_face,
|
||||
weight,
|
||||
save_tmp_steps,
|
||||
preserve_alpha: bool = False,
|
||||
) -> torch.Tensor:
|
||||
pimage = tensor2np(image)[0]
|
||||
width, height = pimage.shape[1], pimage.shape[0]
|
||||
source_img = cv2.cvtColor(np.array(pimage), cv2.COLOR_RGB2BGR)
|
||||
|
||||
alpha_channel = None
|
||||
if (
|
||||
preserve_alpha and image.size(-1) == 4
|
||||
): # Check if the image has an alpha channel
|
||||
alpha_channel = pimage[:, :, 3]
|
||||
pimage = pimage[:, :, :3] # Remove alpha channel for processing
|
||||
|
||||
sys.stdout = NullWriter()
|
||||
cropped_faces, restored_faces, restored_img = model.enhance(
|
||||
source_img,
|
||||
@@ -211,9 +224,14 @@ class MTB_RestoreFace:
|
||||
)
|
||||
output = None
|
||||
if restored_img is not None:
|
||||
output = Image.fromarray(
|
||||
cv2.cvtColor(restored_img, cv2.COLOR_BGR2RGB)
|
||||
)
|
||||
restored_img = cv2.cvtColor(restored_img, cv2.COLOR_BGR2RGB)
|
||||
output = Image.fromarray(restored_img)
|
||||
|
||||
if alpha_channel is not None:
|
||||
alpha_resized = Image.fromarray(alpha_channel).resize(
|
||||
output.size, Image.LANCZOS
|
||||
)
|
||||
output.putalpha(alpha_resized)
|
||||
# imwrite(restored_img, save_restore_path)
|
||||
|
||||
return pil2tensor(output)
|
||||
@@ -226,6 +244,7 @@ class MTB_RestoreFace:
|
||||
only_center_face=False,
|
||||
weight=0.5,
|
||||
save_tmp_steps=True,
|
||||
preserve_alpha: bool = False,
|
||||
) -> tuple[torch.Tensor]:
|
||||
out = [
|
||||
self.do_restore(
|
||||
@@ -235,6 +254,7 @@ class MTB_RestoreFace:
|
||||
only_center_face,
|
||||
weight,
|
||||
save_tmp_steps,
|
||||
preserve_alpha,
|
||||
)
|
||||
for i in range(image.size(0))
|
||||
]
|
||||
@@ -260,7 +280,7 @@ class MTB_RestoreFace:
|
||||
self, cropped_faces, restored_faces, height, width
|
||||
):
|
||||
for idx, (cropped_face, restored_face) in enumerate(
|
||||
zip(cropped_faces, restored_faces)
|
||||
zip(cropped_faces, restored_faces, strict=False)
|
||||
):
|
||||
face_id = idx + 1
|
||||
file = self.get_step_image_path("cropped_faces", face_id)
|
||||
|
||||
+18
-5
@@ -2,7 +2,6 @@
|
||||
# region imports
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Set, Union
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import cv2
|
||||
@@ -41,6 +40,7 @@ class MTB_LoadFaceAnalysisModel:
|
||||
RETURN_TYPES = ("FACE_ANALYSIS_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
if faceswap_model == "antelopev2":
|
||||
@@ -78,6 +78,7 @@ class MTB_LoadFaceSwapModel:
|
||||
RETURN_TYPES = ("FACESWAP_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
model_path = get_model_path("insightface", faceswap_model)
|
||||
@@ -119,12 +120,15 @@ class MTB_FaceSwap:
|
||||
),
|
||||
"faceswap_model": ("FACESWAP_MODEL", {"default": "None"}),
|
||||
},
|
||||
"optional": {},
|
||||
"optional": {
|
||||
"preserve_alpha": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "swap"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def swap(
|
||||
self,
|
||||
@@ -133,11 +137,18 @@ class MTB_FaceSwap:
|
||||
faces_index: str,
|
||||
faceanalysis_model,
|
||||
faceswap_model,
|
||||
preserve_alpha=False,
|
||||
):
|
||||
def do_swap(img):
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
img = tensor2pil(img)[0]
|
||||
ref = tensor2pil(reference)[0]
|
||||
|
||||
alpha_channel = None
|
||||
if preserve_alpha and img.mode == "RGBA":
|
||||
alpha_channel = img.getchannel("A")
|
||||
img = img.convert("RGB")
|
||||
|
||||
face_ids = {
|
||||
int(x)
|
||||
for x in faces_index.strip(",").split(",")
|
||||
@@ -148,6 +159,8 @@ class MTB_FaceSwap:
|
||||
faceanalysis_model, ref, img, faceswap_model, face_ids
|
||||
)
|
||||
sys.stdout = sys.__stdout__
|
||||
if alpha_channel:
|
||||
swapped.putalpha(alpha_channel)
|
||||
return pil2tensor(swapped)
|
||||
|
||||
batch_count = image.size(0)
|
||||
@@ -194,10 +207,10 @@ def get_face_single(
|
||||
|
||||
def swap_face(
|
||||
face_analyser,
|
||||
source_img: Union[Image.Image, List[Image.Image]],
|
||||
target_img: Union[Image.Image, List[Image.Image]],
|
||||
source_img: Image.Image | list[Image.Image],
|
||||
target_img: Image.Image | list[Image.Image],
|
||||
face_swapper_model,
|
||||
faces_index: Optional[Set[int]] = None,
|
||||
faces_index: set[int] | None = None,
|
||||
) -> Image.Image:
|
||||
if faces_index is None:
|
||||
faces_index = {0}
|
||||
|
||||
+186
-119
@@ -1,5 +1,8 @@
|
||||
import qrcode
|
||||
from PIL import Image
|
||||
import io
|
||||
|
||||
import requests
|
||||
import torch
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
|
||||
from ..log import log
|
||||
from ..utils import comfy_dir, font_path, pil2tensor
|
||||
@@ -82,10 +85,6 @@ class MTB_UnsplashImage:
|
||||
CATEGORY = "mtb/generate"
|
||||
|
||||
def do_unsplash_image(self, width, height, random_seed, keyword=None):
|
||||
import io
|
||||
|
||||
import requests
|
||||
|
||||
base_url = "https://source.unsplash.com/random/"
|
||||
|
||||
if width and height:
|
||||
@@ -113,76 +112,6 @@ class MTB_UnsplashImage:
|
||||
return (None,)
|
||||
|
||||
|
||||
class MTB_QrCode:
|
||||
"""Basic QR Code generator"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": "https://www.github.com"}),
|
||||
"width": (
|
||||
"INT",
|
||||
{"default": 256, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"height": (
|
||||
"INT",
|
||||
{"default": 256, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"error_correct": (("L", "M", "Q", "H"), {"default": "L"}),
|
||||
"box_size": (
|
||||
"INT",
|
||||
{"default": 10, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"border": (
|
||||
"INT",
|
||||
{"default": 4, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"invert": (("BOOLEAN",), {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_qr"
|
||||
CATEGORY = "mtb/generate"
|
||||
|
||||
def do_qr(
|
||||
self, url, width, height, error_correct, box_size, border, invert
|
||||
):
|
||||
log.warning(
|
||||
"This node will soon be deprecated, there are much better alternatives like https://github.com/coreyryanhanson/comfy-qr"
|
||||
)
|
||||
if error_correct == "L" or error_correct not in ["M", "Q", "H"]:
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_L
|
||||
elif error_correct == "M":
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_M
|
||||
elif error_correct == "Q":
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_Q
|
||||
else:
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_H
|
||||
|
||||
qr = qrcode.QRCode(
|
||||
version=1,
|
||||
error_correction=error_correct,
|
||||
box_size=box_size,
|
||||
border=border,
|
||||
)
|
||||
qr.add_data(url)
|
||||
qr.make(fit=True)
|
||||
|
||||
back_color = (255, 255, 255) if invert else (0, 0, 0)
|
||||
fill_color = (0, 0, 0) if invert else (255, 255, 255)
|
||||
|
||||
code = img = qr.make_image(
|
||||
back_color=back_color, fill_color=fill_color
|
||||
)
|
||||
|
||||
# that we now resize without filtering
|
||||
code = code.resize((width, height), Image.NEAREST)
|
||||
|
||||
return (pil2tensor(code),)
|
||||
|
||||
|
||||
def bbox_dim(bbox):
|
||||
left, upper, right, lower = bbox
|
||||
width = right - left
|
||||
@@ -202,7 +131,7 @@ class MTB_TextToImage:
|
||||
fonts = {}
|
||||
DESCRIPTION = """# Text to Image
|
||||
|
||||
This node look for any font files in comfy_dir/fonts.
|
||||
This node look for any font files in comfy_dir/fonts.
|
||||
by default it fallsback to a default font.
|
||||
|
||||

|
||||
@@ -284,14 +213,90 @@ by default it fallsback to a default font.
|
||||
"INT",
|
||||
{"default": 100, "min": 1, "max": 100, "step": 1},
|
||||
),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"whisper_chunks": ("WHISPER_CHUNKS",),
|
||||
"fps": (
|
||||
"INT",
|
||||
{"default": 24, "min": 1, "max": 60, "step": 1},
|
||||
),
|
||||
"fade_duration": (
|
||||
"FLOAT",
|
||||
{"default": 0.5, "min": 0.0, "max": 5.0, "step": 0.1},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "text_to_image"
|
||||
CATEGORY = "mtb/generate"
|
||||
|
||||
def create_animation_frames(
|
||||
self,
|
||||
chunks,
|
||||
base_image,
|
||||
font,
|
||||
font_size,
|
||||
color,
|
||||
width,
|
||||
height,
|
||||
fps,
|
||||
fade_duration,
|
||||
):
|
||||
"""Create animation frames from Whisper chunks."""
|
||||
if not chunks or not chunks.get("chunks"):
|
||||
return [base_image]
|
||||
|
||||
frames = []
|
||||
total_duration = chunks["chunks"][-1]["timestamp"][1]
|
||||
frame_count = int(total_duration * fps)
|
||||
fade_frames = int(fade_duration * fps)
|
||||
|
||||
for frame_idx in range(frame_count):
|
||||
time = frame_idx / fps
|
||||
frame = base_image.copy()
|
||||
draw = ImageDraw.Draw(frame)
|
||||
|
||||
active_chunks = []
|
||||
for chunk in chunks["chunks"]:
|
||||
start, end = chunk["timestamp"]
|
||||
if start <= time <= end:
|
||||
fade_in_alpha = min(
|
||||
1.0, (time - start) * fps / fade_frames
|
||||
)
|
||||
fade_out_alpha = min(1.0, (end - time) * fps / fade_frames)
|
||||
alpha = min(fade_in_alpha, fade_out_alpha)
|
||||
active_chunks.append((chunk["text"], alpha))
|
||||
|
||||
y = height // 4
|
||||
for text, alpha in active_chunks:
|
||||
# Create a temporary image for the text with alpha
|
||||
text_img = Image.new("RGBA", (width, height), (0, 0, 0, 0))
|
||||
text_draw = ImageDraw.Draw(text_img)
|
||||
|
||||
text_draw.text(
|
||||
(width // 2, y),
|
||||
text,
|
||||
font=font,
|
||||
fill=color,
|
||||
anchor="mm",
|
||||
)
|
||||
|
||||
text_img.putalpha(
|
||||
Image.fromarray(
|
||||
(torch.ones((height, width)) * (alpha * 255))
|
||||
.byte()
|
||||
.numpy()
|
||||
)
|
||||
)
|
||||
|
||||
frame = Image.alpha_composite(frame, text_img)
|
||||
y += font_size * 1.5
|
||||
|
||||
frames.append(frame)
|
||||
|
||||
return frames
|
||||
|
||||
def text_to_image(
|
||||
self,
|
||||
text: str,
|
||||
@@ -309,62 +314,124 @@ by default it fallsback to a default font.
|
||||
h_offset=0,
|
||||
v_offset=0,
|
||||
h_coverage=100,
|
||||
whisper_chunks=None,
|
||||
fps=24,
|
||||
fade_duration=0.5,
|
||||
):
|
||||
"""Convert text to image, with optional animation support."""
|
||||
import textwrap
|
||||
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
from PIL import ImageColor
|
||||
|
||||
font_path = self.fonts[font]
|
||||
|
||||
text = (
|
||||
text.encode("ascii", "ignore").decode().strip() if trim else text
|
||||
)
|
||||
# Handle word wrapping
|
||||
if wrap:
|
||||
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
|
||||
lines = textwrap.wrap(text, width=wrap_width)
|
||||
else:
|
||||
lines = [text]
|
||||
font = ImageFont.truetype(font_path, size=font_size)
|
||||
log.debug(f"Lines: {lines}")
|
||||
img = Image.new("RGBA", (width, height), background)
|
||||
draw = ImageDraw.Draw(img)
|
||||
|
||||
line_height_px = line_height * font_size
|
||||
try:
|
||||
if isinstance(color, str):
|
||||
color = ImageColor.getrgb(color)
|
||||
if isinstance(background, str):
|
||||
background = ImageColor.getrgb(background)
|
||||
|
||||
# Vertical alignment
|
||||
if v_align == "top":
|
||||
y_text = v_offset
|
||||
elif v_align == "center":
|
||||
y_text = ((height - (line_height_px * len(lines))) // 2) + v_offset
|
||||
else: # bottom
|
||||
y_text = (height - (line_height_px * len(lines))) - v_offset
|
||||
if len(color) == 3:
|
||||
color = color + (255,)
|
||||
if len(background) == 3:
|
||||
background = background + (255,)
|
||||
except ValueError as e:
|
||||
log.error(f"Color parsing error: {e}")
|
||||
color = (255, 255, 255, 255)
|
||||
background = (0, 0, 0, 255)
|
||||
|
||||
def get_width(line):
|
||||
if hasattr(font, "getsize"):
|
||||
return font.getsize(line)[0]
|
||||
def render_text(text_to_render, alpha=None):
|
||||
if trim:
|
||||
text_to_render = (
|
||||
text_to_render.encode("ascii", "ignore").decode().strip()
|
||||
)
|
||||
if wrap:
|
||||
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
|
||||
lines = textwrap.wrap(text_to_render, width=wrap_width)
|
||||
else:
|
||||
return font.getlength(line)
|
||||
lines = [text_to_render]
|
||||
|
||||
# Draw each line of text
|
||||
for line in lines:
|
||||
line_width = get_width(line)
|
||||
# Horizontal alignment
|
||||
if h_align == "left":
|
||||
x_text = h_offset
|
||||
elif h_align == "center":
|
||||
x_text = ((width - line_width) // 2) + h_offset
|
||||
else: # right
|
||||
x_text = (width - line_width) - h_offset
|
||||
img = Image.new("RGBA", (width, height), (0, 0, 0, 0))
|
||||
draw = ImageDraw.Draw(img)
|
||||
|
||||
draw.text((x_text, y_text), line, fill=color, font=font)
|
||||
y_text += line_height_px
|
||||
line_height_px = line_height * font_size
|
||||
|
||||
return (pil2tensor(img),)
|
||||
if v_align == "top":
|
||||
y_text = v_offset
|
||||
elif v_align == "center":
|
||||
y_text = (
|
||||
(height - (line_height_px * len(lines))) // 2
|
||||
) + v_offset
|
||||
else:
|
||||
y_text = (height - (line_height_px * len(lines))) - v_offset
|
||||
|
||||
def get_width(line):
|
||||
if hasattr(font, "getsize"):
|
||||
return font.getsize(line)[0]
|
||||
else:
|
||||
return font.getlength(line)
|
||||
|
||||
for line in lines:
|
||||
line_width = get_width(line)
|
||||
if h_align == "left":
|
||||
x_text = h_offset
|
||||
elif h_align == "center":
|
||||
x_text = ((width - line_width) // 2) + h_offset
|
||||
else:
|
||||
x_text = (width - line_width) - h_offset
|
||||
|
||||
text_color = color
|
||||
if alpha is not None:
|
||||
text_color = tuple(
|
||||
list(color[:3]) + [int(alpha * color[3])]
|
||||
)
|
||||
|
||||
draw.text((x_text, y_text), line, fill=text_color, font=font)
|
||||
y_text += line_height_px
|
||||
|
||||
return img
|
||||
|
||||
base_img = Image.new("RGBA", (width, height), background)
|
||||
|
||||
if whisper_chunks and whisper_chunks.get("chunks"):
|
||||
frames = []
|
||||
total_duration = whisper_chunks["chunks"][-1]["timestamp"][1]
|
||||
frame_count = int(total_duration * fps)
|
||||
fade_frames = int(fade_duration * fps)
|
||||
|
||||
for frame_idx in range(frame_count):
|
||||
time = frame_idx / fps
|
||||
frame = base_img.copy()
|
||||
|
||||
active_chunks = []
|
||||
for chunk in whisper_chunks["chunks"]:
|
||||
start, end = chunk["timestamp"]
|
||||
if start <= time <= end:
|
||||
fade_in_alpha = min(
|
||||
1.0, (time - start) * fps / fade_frames
|
||||
)
|
||||
fade_out_alpha = min(
|
||||
1.0, (end - time) * fps / fade_frames
|
||||
)
|
||||
alpha = min(fade_in_alpha, fade_out_alpha)
|
||||
active_chunks.append((chunk["text"], alpha))
|
||||
|
||||
for chunk_text, alpha in active_chunks:
|
||||
chunk_img = render_text(chunk_text, alpha)
|
||||
frame = Image.alpha_composite(frame, chunk_img)
|
||||
|
||||
frames.append(frame)
|
||||
|
||||
frame_tensors = [pil2tensor(frame) for frame in frames]
|
||||
return (torch.cat(frame_tensors, dim=0),)
|
||||
else:
|
||||
text_img = render_text(text)
|
||||
result = Image.alpha_composite(base_img, text_img)
|
||||
return (pil2tensor(result),)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
MTB_QrCode,
|
||||
MTB_UnsplashImage,
|
||||
MTB_TextToImage,
|
||||
# MtbExamples,
|
||||
|
||||
+312
-43
@@ -1,15 +1,14 @@
|
||||
import io
|
||||
import json
|
||||
import re
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from math import pi
|
||||
from typing import Optional
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms.functional as F
|
||||
from PIL import Image
|
||||
|
||||
from ..log import log
|
||||
@@ -46,14 +45,22 @@ class MTB_ToDevice:
|
||||
if torch.backends.mps.is_available():
|
||||
devices.append("mps")
|
||||
if torch.cuda.is_available():
|
||||
devices.append("cuda:0")
|
||||
for i in range(1, torch.cuda.device_count()):
|
||||
devices.append(f"cuda:{i}")
|
||||
devices.append("cuda")
|
||||
for i in range(torch.cuda.device_count()):
|
||||
devices.append(f"cuda{i}")
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"ignore_errors": ("BOOLEAN", {"default": False}),
|
||||
"device": (devices, {"default": "cpu"}),
|
||||
"device": (
|
||||
devices,
|
||||
{
|
||||
"default": "cuda"
|
||||
if torch.cuda.is_available()
|
||||
else "cpu"
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
@@ -69,20 +76,36 @@ class MTB_ToDevice:
|
||||
def to_device(
|
||||
self,
|
||||
*,
|
||||
ignore_errors=False,
|
||||
device="cuda",
|
||||
image: Optional[torch.Tensor] = None,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
ignore_errors: bool = False,
|
||||
device: str = "cuda",
|
||||
image: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
):
|
||||
if not ignore_errors and image is None and mask is None:
|
||||
raise ValueError(
|
||||
"You must either provide an image or a mask,"
|
||||
" use ignore_error to passthrough"
|
||||
+ " use ignore_error to passthrough"
|
||||
)
|
||||
if (
|
||||
device.startswith("cuda")
|
||||
and ":" not in device
|
||||
and device != "cuda"
|
||||
):
|
||||
device = f"cuda:{device[4:]}"
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image = image.to(device)
|
||||
if mask is not None:
|
||||
mask = mask.to(device)
|
||||
except RuntimeError as e:
|
||||
if not ignore_errors:
|
||||
raise RuntimeError(
|
||||
f"Failed to move tensor to device {device}: {str(e)}"
|
||||
) from e
|
||||
log.warning(
|
||||
f"Failed to move tensor to device {device}, ignoring: {str(e)}"
|
||||
)
|
||||
if image is not None:
|
||||
image = image.to(device)
|
||||
if mask is not None:
|
||||
mask = mask.to(device)
|
||||
return (image, mask)
|
||||
|
||||
|
||||
@@ -138,6 +161,8 @@ class MTB_MatchDimensions:
|
||||
def execute(
|
||||
self, source: torch.Tensor, reference: torch.Tensor, match: str
|
||||
):
|
||||
import torchvision.transforms.functional as VF
|
||||
|
||||
_batch_size, height, width, _channels = source.shape
|
||||
_rbatch_size, rheight, rwidth, _rchannels = reference.shape
|
||||
|
||||
@@ -155,7 +180,7 @@ class MTB_MatchDimensions:
|
||||
new_height = int(rwidth / source_aspect_ratio)
|
||||
|
||||
resized_images = [
|
||||
F.resize(
|
||||
VF.resize(
|
||||
source[i],
|
||||
(new_height, new_width),
|
||||
antialias=True,
|
||||
@@ -417,7 +442,7 @@ class MTB_AnyToString:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"input": ("*")},
|
||||
"required": {"input": ("*",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
@@ -447,7 +472,7 @@ class MTB_AnyToString:
|
||||
|
||||
|
||||
class MTB_StringReplace:
|
||||
"""Basic string replacement."""
|
||||
"""Basic string replacement with regex support."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -456,6 +481,7 @@ class MTB_StringReplace:
|
||||
"string": ("STRING", {"forceInput": True}),
|
||||
"old": ("STRING", {"default": ""}),
|
||||
"new": ("STRING", {"default": ""}),
|
||||
"use_regex": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -463,12 +489,19 @@ class MTB_StringReplace:
|
||||
RETURN_TYPES = ("STRING",)
|
||||
CATEGORY = "mtb/string"
|
||||
|
||||
def replace_str(self, string: str, old: str, new: str):
|
||||
def replace_str(self, string: str, old: str, new: str, use_regex: bool):
|
||||
log.debug(f"Current string: {string}")
|
||||
log.debug(f"Find string: {old}")
|
||||
log.debug(f"Replace string: {new}")
|
||||
log.debug(f"Use regex: {use_regex}")
|
||||
|
||||
string = string.replace(old, new)
|
||||
if use_regex:
|
||||
try:
|
||||
string = re.sub(old, new, string)
|
||||
except re.error as e:
|
||||
raise ValueError(f"Regex error: {e}") from e
|
||||
else:
|
||||
string = string.replace(old, new)
|
||||
|
||||
log.debug(f"New string: {string}")
|
||||
|
||||
@@ -491,14 +524,14 @@ class MTB_MathExpression:
|
||||
RETURN_NAMES = ("result (float)", "result (int)")
|
||||
CATEGORY = "mtb/math"
|
||||
DESCRIPTION = (
|
||||
"evaluate a simple math expression string (!! Fallsback to eval)"
|
||||
"evaluate a simple math expression string, only supports literal_eval"
|
||||
)
|
||||
|
||||
def eval_expression(self, expression, **kwargs):
|
||||
def eval_expression(self, expression: str, **kwargs):
|
||||
from ast import literal_eval
|
||||
|
||||
for key, value in kwargs.items():
|
||||
print(f"Replacing placeholder <{key}> with value {value}")
|
||||
log.debug(f"Replacing placeholder <{key}> with value {value}")
|
||||
expression = expression.replace(f"<{key}>", str(value))
|
||||
|
||||
result = -1
|
||||
@@ -509,15 +542,10 @@ class MTB_MathExpression:
|
||||
f"The expression syntax is wrong '{expression}': {e}"
|
||||
) from e
|
||||
|
||||
except ValueError:
|
||||
try:
|
||||
expression = expression.replace("^", "**")
|
||||
result = eval(expression)
|
||||
except Exception as e:
|
||||
# Handle any other exceptions and provide a meaningful error message
|
||||
raise ValueError(
|
||||
f"Error evaluating expression '{expression}': {e}"
|
||||
) from e
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Math expression only support literal_eval now: {e}"
|
||||
)
|
||||
|
||||
return (result, int(result))
|
||||
|
||||
@@ -531,10 +559,22 @@ class MTB_FitNumber:
|
||||
"required": {
|
||||
"value": ("FLOAT", {"default": 0, "forceInput": True}),
|
||||
"clamp": ("BOOLEAN", {"default": False}),
|
||||
"source_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
||||
"source_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
||||
"target_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
||||
"target_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
||||
"source_min": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "step": 0.01, "min": -1e5},
|
||||
),
|
||||
"source_max": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "step": 0.01, "min": -1e5},
|
||||
),
|
||||
"target_min": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "step": 0.01, "min": -1e5},
|
||||
),
|
||||
"target_max": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "step": 0.01, "min": -1e5},
|
||||
),
|
||||
"easing": (
|
||||
EASINGS,
|
||||
{"default": "Linear"},
|
||||
@@ -583,22 +623,250 @@ class MTB_ConcatImages:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"reverse": ("BOOLEAN", {"default": False})},
|
||||
"optional": {
|
||||
"on_mismatch": (
|
||||
["Error", "Smallest", "Largest"],
|
||||
{"default": "Smallest"},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
def concatenate_tensors(self, reverse, **kwargs):
|
||||
tensors = tuple(kwargs.values())
|
||||
batch_sizes = [tensor.size(0) for tensor in tensors]
|
||||
def concatenate_tensors(
|
||||
self,
|
||||
reverse: bool,
|
||||
on_mismatch: str = "Smallest",
|
||||
**kwargs: torch.Tensor,
|
||||
) -> tuple[torch.Tensor]:
|
||||
tensors = list(kwargs.values())
|
||||
|
||||
if on_mismatch == "Error":
|
||||
shapes = [tensor.shape for tensor in tensors]
|
||||
if not all(shape == shapes[0] for shape in shapes):
|
||||
raise ValueError(
|
||||
"All input tensors must have the same shape when on_mismatch is 'Error'."
|
||||
)
|
||||
|
||||
else:
|
||||
import torch.nn.functional as F
|
||||
|
||||
if on_mismatch == "Smallest":
|
||||
target_shape = min(
|
||||
(tensor.shape for tensor in tensors),
|
||||
key=lambda s: (s[1], s[2]),
|
||||
)
|
||||
else: # on_mismatch == "Largest"
|
||||
target_shape = max(
|
||||
(tensor.shape for tensor in tensors),
|
||||
key=lambda s: (s[1], s[2]),
|
||||
)
|
||||
|
||||
target_height, target_width = target_shape[1], target_shape[2]
|
||||
|
||||
resized_tensors = []
|
||||
for tensor in tensors:
|
||||
if (
|
||||
tensor.shape[1] != target_height
|
||||
or tensor.shape[2] != target_width
|
||||
):
|
||||
resized_tensor = F.interpolate(
|
||||
tensor.permute(0, 3, 1, 2),
|
||||
size=(target_height, target_width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
resized_tensor = resized_tensor.permute(0, 2, 3, 1)
|
||||
resized_tensors.append(resized_tensor)
|
||||
else:
|
||||
resized_tensors.append(tensor)
|
||||
|
||||
tensors = resized_tensors
|
||||
|
||||
concatenated = torch.cat(tensors, dim=0)
|
||||
|
||||
# Update the batch size in the concatenated tensor
|
||||
concatenated_size = list(concatenated.size())
|
||||
concatenated_size[0] = sum(batch_sizes)
|
||||
concatenated = concatenated.view(*concatenated_size)
|
||||
|
||||
return (concatenated,)
|
||||
|
||||
|
||||
class MTB_TensorOps:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"tensor": ("IMAGE",),
|
||||
"operation": (
|
||||
[
|
||||
"multiply",
|
||||
"divide",
|
||||
"add",
|
||||
"subtract",
|
||||
"power",
|
||||
"clamp",
|
||||
"abs",
|
||||
"log",
|
||||
"exp",
|
||||
"convert_dtype",
|
||||
"normalize_range",
|
||||
"normalize_per_channel",
|
||||
],
|
||||
{"default": "multiply"},
|
||||
),
|
||||
"value": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": -1000000.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.01,
|
||||
},
|
||||
),
|
||||
"source_min": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": -1000000.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.01,
|
||||
},
|
||||
),
|
||||
"source_max": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": -1000000.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.01,
|
||||
},
|
||||
),
|
||||
"target_min": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": -1000000.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.01,
|
||||
},
|
||||
),
|
||||
"target_max": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 16.0,
|
||||
"min": -1000000.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.01,
|
||||
},
|
||||
),
|
||||
"dtype": (
|
||||
["uint8", "float32", "float16", "bfloat16"],
|
||||
{"default": "float32"},
|
||||
),
|
||||
"use_mean": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"target_tensor": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "apply"
|
||||
CATEGORY = "mtb/tensor_ops"
|
||||
|
||||
def apply(
|
||||
self,
|
||||
tensor,
|
||||
operation="multiply",
|
||||
value=1.0,
|
||||
source_min=0.0,
|
||||
source_max=1.0,
|
||||
target_min=0.0,
|
||||
target_max=1.0,
|
||||
dtype="float32",
|
||||
use_mean=False,
|
||||
target_tensor=None,
|
||||
):
|
||||
log.debug(
|
||||
f"Input tensor stats: shape={tensor.shape}, dtype={tensor.dtype}, range=[{tensor.min().item():.6f}, {tensor.max().item():.6f}]"
|
||||
)
|
||||
if operation == "normalize_per_channel":
|
||||
if target_tensor is None:
|
||||
raise ValueError(
|
||||
"Target tensor required for per-channel normalization"
|
||||
)
|
||||
|
||||
result = tensor.clone()
|
||||
for c in range(tensor.shape[-1]):
|
||||
if use_mean:
|
||||
source_mean = tensor[..., c].mean()
|
||||
target_mean = target_tensor[..., c].mean()
|
||||
scale = target_mean / source_mean
|
||||
result[..., c] = tensor[..., c] * scale
|
||||
else:
|
||||
source_min = tensor[..., c].min()
|
||||
source_max = tensor[..., c].max()
|
||||
target_min = target_tensor[..., c].min()
|
||||
target_max = target_tensor[..., c].max()
|
||||
|
||||
normalized = (tensor[..., c] - source_min) / (
|
||||
source_max - source_min
|
||||
)
|
||||
result[..., c] = (
|
||||
normalized * (target_max - target_min) + target_min
|
||||
)
|
||||
|
||||
log.debug(
|
||||
f"Channel {c} - Scale: source=[{source_min:.6f}, {source_max:.6f}], target=[{target_min:.6f}, {target_max:.6f}]"
|
||||
)
|
||||
|
||||
elif operation == "normalize_range":
|
||||
if target_tensor is not None:
|
||||
target_min = target_tensor.min().item()
|
||||
target_max = target_tensor.max().item()
|
||||
log.debug(
|
||||
f"Using target tensor range: [{target_min:.6f}, {target_max:.6f}]"
|
||||
)
|
||||
|
||||
normalized = (tensor - source_min) / (source_max - source_min)
|
||||
result = normalized * (target_max - target_min) + target_min
|
||||
elif operation == "convert_dtype":
|
||||
if dtype == "float32":
|
||||
result = tensor.float()
|
||||
elif dtype == "float16":
|
||||
result = tensor.half()
|
||||
elif dtype == "bfloat16":
|
||||
result = tensor.bfloat16()
|
||||
|
||||
else:
|
||||
result = tensor
|
||||
if operation == "multiply":
|
||||
result = tensor * value
|
||||
elif operation == "divide":
|
||||
result = tensor / value if value != 0 else tensor
|
||||
elif operation == "add":
|
||||
result = tensor + value
|
||||
elif operation == "subtract":
|
||||
result = tensor - value
|
||||
elif operation == "power":
|
||||
result = torch.pow(tensor, value)
|
||||
elif operation == "clamp":
|
||||
if target_tensor is not None:
|
||||
result = torch.clamp(
|
||||
tensor,
|
||||
target_tensor.min().item(),
|
||||
target_tensor.max().item(),
|
||||
)
|
||||
else:
|
||||
result = torch.clamp(tensor, source_min, source_max)
|
||||
elif operation == "abs":
|
||||
result = torch.abs(tensor)
|
||||
elif operation == "log":
|
||||
result = torch.log(tensor.clamp(min=1e-10))
|
||||
elif operation == "exp":
|
||||
result = torch.exp(tensor)
|
||||
|
||||
log.debug(
|
||||
f"Output tensor stats: shape={result.shape}, dtype={result.dtype}, range=[{result.min().item():.6f}, {result.max().item():.6f}]"
|
||||
)
|
||||
return (result,)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
MTB_StringReplace,
|
||||
MTB_FitNumber,
|
||||
@@ -613,4 +881,5 @@ __nodes__ = [
|
||||
MTB_FloatsToFloat,
|
||||
MTB_FloatToFloats,
|
||||
MTB_FloatsToInts,
|
||||
MTB_TensorOps,
|
||||
]
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import comfy
|
||||
import comfy.model_management as model_management
|
||||
@@ -15,10 +14,13 @@ from ..utils import get_model_path
|
||||
|
||||
|
||||
class MTB_LoadFilmModel:
|
||||
"""Loads a FILM model"""
|
||||
"""Loads a FILM model
|
||||
|
||||
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_models() -> List[Path]:
|
||||
def get_models() -> list[Path]:
|
||||
models_paths = get_model_path("FILM").iterdir()
|
||||
|
||||
return [x for x in models_paths if x.suffix in [".onnx", ".pth"]]
|
||||
@@ -37,6 +39,7 @@ class MTB_LoadFilmModel:
|
||||
RETURN_TYPES = ("FILM_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/frame iterpolation"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, film_model: str):
|
||||
model_path = get_model_path("FILM", film_model)
|
||||
@@ -56,7 +59,10 @@ class MTB_LoadFilmModel:
|
||||
|
||||
|
||||
class MTB_FilmInterpolation:
|
||||
"""Google Research FILM frame interpolation for large motion"""
|
||||
"""Google Research FILM frame interpolation for large motion
|
||||
|
||||
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -71,6 +77,7 @@ class MTB_FilmInterpolation:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_interpolation"
|
||||
CATEGORY = "mtb/frame iterpolation"
|
||||
DEPRECATED = True
|
||||
|
||||
def do_interpolation(
|
||||
self,
|
||||
|
||||
+445
-90
@@ -3,6 +3,7 @@ import json
|
||||
import math
|
||||
import os
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -13,7 +14,7 @@ from skimage.filters import gaussian
|
||||
from skimage.util import compare_images
|
||||
|
||||
from ..log import log
|
||||
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
|
||||
from ..utils import np2tensor, pil2tensor, tensor2pil
|
||||
|
||||
# try:
|
||||
# from cv2.ximgproc import guidedFilter
|
||||
@@ -35,6 +36,343 @@ def gaussian_kernel(
|
||||
return g / g.sum()
|
||||
|
||||
|
||||
class MTB_CoordinatesToString:
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "convert"
|
||||
CATEGORY = "mtb/coordinates"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"coordinates": ("BATCH_COORDINATES",),
|
||||
"frame": ("INT",),
|
||||
}
|
||||
}
|
||||
|
||||
def convert(
|
||||
self, coordinates: list[list[tuple[int, int]]], frame: int
|
||||
) -> tuple[str]:
|
||||
frame = max(frame, len(coordinates) - 1)
|
||||
coords = coordinates[frame]
|
||||
output: list[dict[str, int]] = []
|
||||
|
||||
for x, y in coords:
|
||||
output.append({"x": x, "y": y})
|
||||
|
||||
return (json.dumps(output),)
|
||||
|
||||
|
||||
class MTB_ExtractCoordinatesFromImage:
|
||||
"""Extract 2D points from a batch of images based on a threshold."""
|
||||
|
||||
RETURN_TYPES = ("BATCH_COORDINATES", "IMAGE")
|
||||
FUNCTION = "extract"
|
||||
CATEGORY = "mtb/coordinates"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"threshold": ("FLOAT",),
|
||||
"max_points": ("INT", {"default": 50, "min": 0}),
|
||||
},
|
||||
"optional": {"image": ("IMAGE",), "mask": ("MASK",)},
|
||||
}
|
||||
|
||||
def extract(
|
||||
self,
|
||||
threshold: float,
|
||||
max_points: int,
|
||||
image: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]:
|
||||
if image is not None:
|
||||
batch_count, height, width, channel_count = image.shape
|
||||
imgs = image
|
||||
else:
|
||||
if mask is None:
|
||||
raise ValueError("Must provide either image or mask")
|
||||
batch_count, height, width = mask.shape
|
||||
channel_count = 1
|
||||
imgs = mask
|
||||
|
||||
if channel_count not in [1, 2, 3, 4]:
|
||||
raise ValueError(f"Incorrect channel count: {channel_count}")
|
||||
|
||||
all_points: list[list[tuple[int, int]]] = []
|
||||
debug_images = torch.zeros(
|
||||
(batch_count, height, width, 3),
|
||||
dtype=torch.uint8,
|
||||
device=imgs.device,
|
||||
)
|
||||
|
||||
for i, img in enumerate(imgs):
|
||||
if channel_count == 1:
|
||||
alpha_channel = img if len(img.shape) == 2 else img[:, :, 0]
|
||||
elif channel_count == 2:
|
||||
alpha_channel = img[:, :, 1]
|
||||
elif channel_count == 4:
|
||||
alpha_channel = img[:, :, 3]
|
||||
else:
|
||||
# get intensity
|
||||
alpha_channel = img[:, :, :3].max(dim=2)[0]
|
||||
|
||||
points = (alpha_channel > threshold).nonzero(as_tuple=False)
|
||||
|
||||
if len(points) > max_points:
|
||||
indices = torch.randperm(points.size(0), device=img.device)[
|
||||
:max_points
|
||||
]
|
||||
points = points[indices]
|
||||
|
||||
points = [(int(y.item()), int(x.item())) for x, y in points]
|
||||
all_points.append(points)
|
||||
|
||||
for x, y in points:
|
||||
self._draw_circle(debug_images[i], (x, y), 5)
|
||||
|
||||
return (all_points, debug_images)
|
||||
|
||||
@staticmethod
|
||||
def _draw_circle(
|
||||
image: torch.Tensor, center: tuple[int, int], radius: int
|
||||
):
|
||||
"""Draw a 5px circle on the image."""
|
||||
x0, y0 = center
|
||||
for x in range(-radius, radius + 1):
|
||||
for y in range(-radius, radius + 1):
|
||||
in_radius = x**2 + y**2 <= radius**2
|
||||
in_bounds = (
|
||||
0 <= x0 + x < image.shape[1]
|
||||
and 0 <= y0 + y < image.shape[0]
|
||||
)
|
||||
if in_radius and in_bounds:
|
||||
image[y0 + y, x0 + x] = torch.tensor(
|
||||
[255, 255, 255],
|
||||
dtype=torch.uint8,
|
||||
device=image.device,
|
||||
)
|
||||
|
||||
|
||||
class MTB_ColorCorrectGPU:
|
||||
"""Various color correction methods using only Torch."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"force_gpu": ("BOOLEAN", {"default": True}),
|
||||
"clamp": ([True, False], {"default": True}),
|
||||
"gamma": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
"contrast": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
"exposure": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": -5.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
"offset": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": -5.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
"hue": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": -0.5, "max": 0.5, "step": 0.01},
|
||||
),
|
||||
"saturation": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
"value": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
},
|
||||
"optional": {"mask": ("MASK",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "correct"
|
||||
CATEGORY = "mtb/image processing"
|
||||
|
||||
@staticmethod
|
||||
def get_device(tensor: torch.Tensor, force_gpu: bool):
|
||||
if force_gpu:
|
||||
if torch.cuda.is_available():
|
||||
return torch.device("cuda")
|
||||
elif (
|
||||
hasattr(torch.backends, "mps")
|
||||
and torch.backends.mps.is_available()
|
||||
):
|
||||
return torch.device("mps")
|
||||
elif hasattr(torch, "hip") and torch.hip.is_available():
|
||||
return torch.device("hip")
|
||||
return (
|
||||
tensor.device
|
||||
) # model_management.get_torch_device() # torch.device("cpu")
|
||||
|
||||
@staticmethod
|
||||
def rgb_to_hsv(image: torch.Tensor):
|
||||
r, g, b = image.unbind(-1)
|
||||
max_rgb, argmax_rgb = image.max(-1)
|
||||
min_rgb, _ = image.min(-1)
|
||||
|
||||
diff = max_rgb - min_rgb
|
||||
|
||||
h = torch.empty_like(max_rgb)
|
||||
s = diff / (max_rgb + 1e-7)
|
||||
v = max_rgb
|
||||
|
||||
h[argmax_rgb == 0] = (g - b)[argmax_rgb == 0] / (diff + 1e-7)[
|
||||
argmax_rgb == 0
|
||||
]
|
||||
h[argmax_rgb == 1] = (
|
||||
2.0 + (b - r)[argmax_rgb == 1] / (diff + 1e-7)[argmax_rgb == 1]
|
||||
)
|
||||
h[argmax_rgb == 2] = (
|
||||
4.0 + (r - g)[argmax_rgb == 2] / (diff + 1e-7)[argmax_rgb == 2]
|
||||
)
|
||||
h = (h / 6.0) % 1.0
|
||||
|
||||
h = h.unsqueeze(-1)
|
||||
s = s.unsqueeze(-1)
|
||||
v = v.unsqueeze(-1)
|
||||
|
||||
return torch.cat((h, s, v), dim=-1)
|
||||
|
||||
@staticmethod
|
||||
def hsv_to_rgb(hsv: torch.Tensor):
|
||||
h, s, v = hsv.unbind(-1)
|
||||
h = h * 6.0
|
||||
|
||||
i = torch.floor(h)
|
||||
f = h - i
|
||||
p = v * (1.0 - s)
|
||||
q = v * (1.0 - s * f)
|
||||
t = v * (1.0 - s * (1.0 - f))
|
||||
|
||||
i = i.long() % 6
|
||||
|
||||
mask = torch.stack(
|
||||
(i == 0, i == 1, i == 2, i == 3, i == 4, i == 5), -1
|
||||
)
|
||||
|
||||
rgb = torch.stack(
|
||||
(
|
||||
torch.where(
|
||||
mask[..., 0],
|
||||
v,
|
||||
torch.where(
|
||||
mask[..., 1],
|
||||
q,
|
||||
torch.where(
|
||||
mask[..., 2],
|
||||
p,
|
||||
torch.where(
|
||||
mask[..., 3],
|
||||
p,
|
||||
torch.where(mask[..., 4], t, v),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
torch.where(
|
||||
mask[..., 0],
|
||||
t,
|
||||
torch.where(
|
||||
mask[..., 1],
|
||||
v,
|
||||
torch.where(
|
||||
mask[..., 2],
|
||||
v,
|
||||
torch.where(
|
||||
mask[..., 3],
|
||||
q,
|
||||
torch.where(mask[..., 4], p, p),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
torch.where(
|
||||
mask[..., 0],
|
||||
p,
|
||||
torch.where(
|
||||
mask[..., 1],
|
||||
p,
|
||||
torch.where(
|
||||
mask[..., 2],
|
||||
t,
|
||||
torch.where(
|
||||
mask[..., 3],
|
||||
v,
|
||||
torch.where(mask[..., 4], v, q),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
return rgb
|
||||
|
||||
def correct(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
force_gpu: bool,
|
||||
clamp: bool,
|
||||
gamma: float = 1.0,
|
||||
contrast: float = 1.0,
|
||||
exposure: float = 0.0,
|
||||
offset: float = 0.0,
|
||||
hue: float = 0.0,
|
||||
saturation: float = 1.0,
|
||||
value: float = 1.0,
|
||||
mask: torch.Tensor | None = None,
|
||||
):
|
||||
device = self.get_device(image, force_gpu)
|
||||
image = image.to(device)
|
||||
|
||||
if mask is not None:
|
||||
if mask.shape[0] != image.shape[0]:
|
||||
mask = mask.expand(image.shape[0], -1, -1)
|
||||
|
||||
mask = mask.unsqueeze(-1).expand(-1, -1, -1, 3)
|
||||
mask = mask.to(device)
|
||||
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
adjusted = image.pow(1 / gamma) * (2.0**exposure) * contrast + offset
|
||||
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
hsv = self.rgb_to_hsv(adjusted)
|
||||
hsv[..., 0] = (hsv[..., 0] + hue) % 1.0 # Hue
|
||||
hsv[..., 1] = hsv[..., 1] * saturation # Saturation
|
||||
hsv[..., 2] = hsv[..., 2] * value # Value
|
||||
adjusted = self.hsv_to_rgb(hsv)
|
||||
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
if clamp:
|
||||
adjusted = torch.clamp(adjusted, 0.0, 1.0)
|
||||
|
||||
# apply mask
|
||||
result = (
|
||||
adjusted
|
||||
if mask is None
|
||||
else torch.where(mask > 0, adjusted, image)
|
||||
)
|
||||
|
||||
if not force_gpu:
|
||||
result = result.cpu()
|
||||
|
||||
return (result,)
|
||||
|
||||
|
||||
class MTB_ColorCorrect:
|
||||
"""Various color correction methods"""
|
||||
|
||||
@@ -72,7 +410,8 @@ class MTB_ColorCorrect:
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01},
|
||||
),
|
||||
}
|
||||
},
|
||||
"optional": {"mask": ("MASK",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -87,13 +426,13 @@ class MTB_ColorCorrect:
|
||||
@staticmethod
|
||||
def contrast_adjustment_tensor(image, contrast):
|
||||
r, g, b = image.unbind(-1)
|
||||
|
||||
|
||||
# Using Adobe RGB luminance weights.
|
||||
luminance_image = 0.33 * r + 0.71 * g + 0.06 * b
|
||||
luminance_mean = torch.mean(luminance_image.unsqueeze(-1))
|
||||
|
||||
# Blend original with mean luminance using contrast factor as blend ratio.
|
||||
contrasted = image * contrast + (1.0 - contrast) * luminance_mean
|
||||
contrasted = image * contrast + (1.0 - contrast) * luminance_mean
|
||||
return torch.clamp(contrasted, 0.0, 1.0)
|
||||
|
||||
@staticmethod
|
||||
@@ -188,18 +527,31 @@ class MTB_ColorCorrect:
|
||||
hue: float = 0.0,
|
||||
saturation: float = 1.0,
|
||||
value: float = 1.0,
|
||||
mask: torch.Tensor | None = None,
|
||||
):
|
||||
if mask is not None:
|
||||
if mask.shape[0] != image.shape[0]:
|
||||
mask = mask.expand(image.shape[0], -1, -1)
|
||||
|
||||
mask = mask.unsqueeze(-1).expand(-1, -1, -1, 3)
|
||||
|
||||
# Apply color correction operations
|
||||
image = self.gamma_correction_tensor(image, gamma)
|
||||
image = self.contrast_adjustment_tensor(image, contrast)
|
||||
image = self.exposure_adjustment_tensor(image, exposure)
|
||||
image = self.offset_adjustment_tensor(image, offset)
|
||||
image = self.hsv_adjustment(image, hue, saturation, value)
|
||||
adjusted = self.gamma_correction_tensor(image, gamma)
|
||||
adjusted = self.contrast_adjustment_tensor(adjusted, contrast)
|
||||
adjusted = self.exposure_adjustment_tensor(adjusted, exposure)
|
||||
adjusted = self.offset_adjustment_tensor(adjusted, offset)
|
||||
adjusted = self.hsv_adjustment(adjusted, hue, saturation, value)
|
||||
|
||||
if clamp:
|
||||
image = torch.clamp(image, 0.0, 1.0)
|
||||
adjusted = torch.clamp(adjusted, 0.0, 1.0)
|
||||
|
||||
return (image,)
|
||||
result = (
|
||||
adjusted
|
||||
if mask is None
|
||||
else torch.where(mask > 0, adjusted, image)
|
||||
)
|
||||
|
||||
return (result,)
|
||||
|
||||
|
||||
class MTB_ImageCompare:
|
||||
@@ -350,7 +702,6 @@ class MTB_Blur:
|
||||
)
|
||||
blurred_images.append(blurred)
|
||||
|
||||
image_np = np.array(blurred_images)
|
||||
else:
|
||||
for i in range(image.size(0)):
|
||||
blurred = gaussian(
|
||||
@@ -358,8 +709,7 @@ class MTB_Blur:
|
||||
)
|
||||
blurred_images.append(blurred)
|
||||
|
||||
image_np = np.array(blurred_images)
|
||||
return (np2tensor(image_np).squeeze(0),)
|
||||
return (np2tensor(blurred_images),)
|
||||
|
||||
|
||||
class MTB_Sharpen:
|
||||
@@ -475,7 +825,10 @@ class MTB_MaskToImage:
|
||||
"mask": ("MASK",),
|
||||
"color": ("COLOR",),
|
||||
"background": ("COLOR", {"default": "#000000"}),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "mtb/generate"
|
||||
@@ -484,11 +837,12 @@ class MTB_MaskToImage:
|
||||
|
||||
FUNCTION = "render_mask"
|
||||
|
||||
def render_mask(self, mask, color, background):
|
||||
masks = tensor2np(mask)[0]
|
||||
def render_mask(self, mask, color, background, invert=False):
|
||||
masks = tensor2pil(1.0 - mask) if invert else tensor2pil(mask)
|
||||
images = []
|
||||
|
||||
for m in masks:
|
||||
_mask = Image.fromarray(m).convert("L")
|
||||
_mask = m.convert("L")
|
||||
|
||||
log.debug(
|
||||
f"Converted mask to PIL Image format, size: {_mask.size}"
|
||||
@@ -526,6 +880,11 @@ class MTB_ColoredImage:
|
||||
"optional": {
|
||||
"foreground_image": ("IMAGE",),
|
||||
"foreground_mask": ("MASK",),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
"mask_opacity": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "step": 0.1, "min": 0},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -535,28 +894,19 @@ class MTB_ColoredImage:
|
||||
|
||||
FUNCTION = "render_img"
|
||||
|
||||
def resize_and_crop(self, img, target_size):
|
||||
# Calculate scaling factors for both dimensions
|
||||
scale_x = target_size[0] / img.width
|
||||
scale_y = target_size[1] / img.height
|
||||
|
||||
# Use the smaller scaling factor to maintain aspect ratio
|
||||
scale = max(scale_x, scale_y)
|
||||
|
||||
# Resize the image based on calculated scale
|
||||
def resize_and_crop(self, img: Image.Image, target_size: tuple[int, int]):
|
||||
scale = max(target_size[0] / img.width, target_size[1] / img.height)
|
||||
new_size = (int(img.width * scale), int(img.height * scale))
|
||||
img = img.resize(new_size, Image.LANCZOS)
|
||||
left = (img.width - target_size[0]) // 2
|
||||
top = (img.height - target_size[1]) // 2
|
||||
return img.crop(
|
||||
(left, top, left + target_size[0], top + target_size[1])
|
||||
)
|
||||
|
||||
# Calculate cropping coordinates
|
||||
left = (img.width - target_size[0]) / 2
|
||||
top = (img.height - target_size[1]) / 2
|
||||
right = (img.width + target_size[0]) / 2
|
||||
bottom = (img.height + target_size[1]) / 2
|
||||
|
||||
# Crop and return the image
|
||||
return img.crop((left, top, right, bottom))
|
||||
|
||||
def resize_and_crop_thumbnails(self, img, target_size):
|
||||
def resize_and_crop_thumbnails(
|
||||
self, img: Image.Image, target_size: tuple[int, int]
|
||||
):
|
||||
img.thumbnail(target_size, Image.LANCZOS)
|
||||
left = (img.width - target_size[0]) / 2
|
||||
top = (img.height - target_size[1]) / 2
|
||||
@@ -564,69 +914,71 @@ class MTB_ColoredImage:
|
||||
bottom = (img.height + target_size[1]) / 2
|
||||
return img.crop((left, top, right, bottom))
|
||||
|
||||
@staticmethod
|
||||
def process_mask(
|
||||
mask: torch.Tensor | None,
|
||||
invert: bool,
|
||||
# opacity: float,
|
||||
batch_size: int,
|
||||
) -> list[Image.Image] | None:
|
||||
if mask is None:
|
||||
return [None] * batch_size
|
||||
|
||||
masks = tensor2pil(mask if not invert else 1.0 - mask)
|
||||
|
||||
if len(masks) == 1 and batch_size > 1:
|
||||
masks = masks * batch_size
|
||||
|
||||
if len(masks) != batch_size:
|
||||
raise ValueError(
|
||||
"Foreground image and mask must have the same batch size"
|
||||
)
|
||||
|
||||
return masks
|
||||
|
||||
def render_img(
|
||||
self,
|
||||
color,
|
||||
width,
|
||||
height,
|
||||
color: str,
|
||||
width: int,
|
||||
height: int,
|
||||
foreground_image: torch.Tensor | None = None,
|
||||
foreground_mask: torch.Tensor | None = None,
|
||||
):
|
||||
image = Image.new("RGBA", (width, height), color=color)
|
||||
output = []
|
||||
if foreground_image is not None:
|
||||
fg_masks = [None] * foreground_image.size()[0]
|
||||
invert: bool = False,
|
||||
mask_opacity: float = 1.0,
|
||||
) -> tuple[torch.Tensor]:
|
||||
background = Image.new("RGBA", (width, height), color=color)
|
||||
|
||||
if foreground_mask is not None:
|
||||
fg_size = foreground_image.size()[0]
|
||||
mask_size = foreground_mask.size()[0]
|
||||
if foreground_image is None:
|
||||
return (pil2tensor([background.convert("RGB")]),)
|
||||
|
||||
if fg_size == 1 and mask_size > fg_size:
|
||||
foreground_image = foreground_image.repeat(
|
||||
mask_size, 1, 1, 1
|
||||
)
|
||||
fg_images = tensor2pil(foreground_image)
|
||||
fg_masks = self.process_mask(foreground_mask, invert, len(fg_images))
|
||||
|
||||
if foreground_image.size()[0] != foreground_mask.size()[0]:
|
||||
output: list[Image.Image] = []
|
||||
for fg_image, fg_mask in zip(fg_images, fg_masks, strict=False):
|
||||
fg_image = self.resize_and_crop(fg_image, background.size)
|
||||
|
||||
if fg_mask:
|
||||
fg_mask = self.resize_and_crop(fg_mask, background.size)
|
||||
|
||||
fg_mask_array = np.array(fg_mask)
|
||||
fg_mask_array = (fg_mask_array * mask_opacity).astype(np.uint8)
|
||||
fg_mask = Image.fromarray(fg_mask_array)
|
||||
output.append(
|
||||
Image.composite(
|
||||
fg_image.convert("RGBA"), background, fg_mask
|
||||
).convert("RGB")
|
||||
)
|
||||
else:
|
||||
if fg_image.mode != "RGBA":
|
||||
raise ValueError(
|
||||
"Foreground image and mask must have same batch size"
|
||||
f"Foreground image must be in 'RGBA' mode when no mask is provided, got {fg_image.mode}"
|
||||
)
|
||||
fg_masks = tensor2pil(foreground_mask.unsqueeze(-1))
|
||||
output.append(
|
||||
Image.alpha_composite(background, fg_image).convert("RGB")
|
||||
)
|
||||
|
||||
fg_images = tensor2pil(foreground_image)
|
||||
|
||||
for fg_image, fg_mask in zip(fg_images, fg_masks):
|
||||
# Resize and crop if dimensions mismatch
|
||||
if fg_image.size != image.size:
|
||||
fg_image = self.resize_and_crop(fg_image, image.size)
|
||||
if fg_mask:
|
||||
fg_mask = self.resize_and_crop(fg_mask, image.size)
|
||||
|
||||
if fg_mask:
|
||||
output.append(
|
||||
Image.composite(
|
||||
fg_image.convert("RGBA"),
|
||||
image,
|
||||
fg_mask,
|
||||
).convert("RGB")
|
||||
)
|
||||
else:
|
||||
if fg_image.mode != "RGBA":
|
||||
raise ValueError(
|
||||
"Foreground image must be in 'RGBA' mode "
|
||||
f"when no mask is provided, got {fg_image.mode}"
|
||||
)
|
||||
output.append(
|
||||
Image.alpha_composite(image, fg_image).convert("RGB")
|
||||
)
|
||||
|
||||
else:
|
||||
if foreground_mask is not None:
|
||||
log.warn("Mask ignored because no foreground image is given")
|
||||
output.append(image.convert("RGB"))
|
||||
|
||||
output = pil2tensor(output)
|
||||
|
||||
return (output,)
|
||||
return (pil2tensor(output),)
|
||||
|
||||
|
||||
class MTB_ImagePremultiply:
|
||||
@@ -924,6 +1276,7 @@ class MTB_ImageTileOffset:
|
||||
|
||||
__nodes__ = [
|
||||
MTB_ColorCorrect,
|
||||
MTB_ColorCorrectGPU,
|
||||
MTB_ImageCompare,
|
||||
MTB_ImageTileOffset,
|
||||
MTB_Blur,
|
||||
@@ -935,4 +1288,6 @@ __nodes__ = [
|
||||
MTB_SaveImageGrid,
|
||||
MTB_LoadImageFromUrl,
|
||||
MTB_Sharpen,
|
||||
MTB_ExtractCoordinatesFromImage,
|
||||
MTB_CoordinatesToString,
|
||||
]
|
||||
|
||||
+186
-29
@@ -1,4 +1,11 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from comfy.cli_args import args
|
||||
from PIL import Image
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
from ..log import log
|
||||
|
||||
@@ -8,13 +15,21 @@ class MTB_StackImages:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {"vertical": ("BOOLEAN", {"default": False})}}
|
||||
return {
|
||||
"required": {"vertical": ("BOOLEAN", {"default": False})},
|
||||
"optional": {
|
||||
"match_method": (
|
||||
["error", "smallest", "largest"],
|
||||
{"default": "error"},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "stack"
|
||||
CATEGORY = "mtb/image utils"
|
||||
|
||||
def stack(self, vertical, **kwargs):
|
||||
def stack(self, vertical, match_method="error", **kwargs):
|
||||
if not kwargs:
|
||||
raise ValueError("At least one tensor must be provided.")
|
||||
|
||||
@@ -24,31 +39,62 @@ class MTB_StackImages:
|
||||
f"{'vertically' if vertical else 'horizontally'}"
|
||||
)
|
||||
|
||||
target_device = tensors[0].device
|
||||
|
||||
normalized_tensors = [
|
||||
self.normalize_to_rgba(tensor) for tensor in tensors
|
||||
self.normalize_to_rgba(tensor.to(target_device))
|
||||
for tensor in tensors
|
||||
]
|
||||
|
||||
max_batch_size = max(tensor.shape[0] for tensor in normalized_tensors)
|
||||
normalized_tensors = [
|
||||
self.duplicate_frames(tensor, max_batch_size)
|
||||
for tensor in normalized_tensors
|
||||
]
|
||||
|
||||
if vertical:
|
||||
width = normalized_tensors[0].shape[2]
|
||||
if any(tensor.shape[2] != width for tensor in normalized_tensors):
|
||||
raise ValueError(
|
||||
"All tensors must have the same width "
|
||||
"for vertical stacking."
|
||||
if match_method != "error":
|
||||
if vertical:
|
||||
# match widths
|
||||
widths = [tensor.shape[2] for tensor in normalized_tensors]
|
||||
target_width = (
|
||||
min(widths) if match_method == "smallest" else max(widths)
|
||||
)
|
||||
dim = 1
|
||||
normalized_tensors = [
|
||||
self.resize_tensor(tensor, width=target_width)
|
||||
for tensor in normalized_tensors
|
||||
]
|
||||
else:
|
||||
# match heights
|
||||
heights = [tensor.shape[1] for tensor in normalized_tensors]
|
||||
target_height = (
|
||||
min(heights)
|
||||
if match_method == "smallest"
|
||||
else max(heights)
|
||||
)
|
||||
normalized_tensors = [
|
||||
self.resize_tensor(tensor, height=target_height)
|
||||
for tensor in normalized_tensors
|
||||
]
|
||||
else:
|
||||
height = normalized_tensors[0].shape[1]
|
||||
if any(tensor.shape[1] != height for tensor in normalized_tensors):
|
||||
raise ValueError(
|
||||
"All tensors must have the same height "
|
||||
"for horizontal stacking."
|
||||
)
|
||||
dim = 2
|
||||
if vertical:
|
||||
width = normalized_tensors[0].shape[2]
|
||||
if any(
|
||||
tensor.shape[2] != width for tensor in normalized_tensors
|
||||
):
|
||||
raise ValueError(
|
||||
"All tensors must have the same width "
|
||||
"for vertical stacking."
|
||||
)
|
||||
else:
|
||||
height = normalized_tensors[0].shape[1]
|
||||
if any(
|
||||
tensor.shape[1] != height for tensor in normalized_tensors
|
||||
):
|
||||
raise ValueError(
|
||||
"All tensors must have the same height "
|
||||
"for horizontal stacking."
|
||||
)
|
||||
|
||||
dim = 1 if vertical else 2
|
||||
|
||||
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
|
||||
|
||||
@@ -64,7 +110,7 @@ class MTB_StackImages:
|
||||
elif channels == 3:
|
||||
alpha_channel = torch.ones(
|
||||
tensor.shape[:-1] + (1,), device=tensor.device
|
||||
) # Add an alpha channel
|
||||
)
|
||||
return torch.cat((tensor, alpha_channel), dim=-1)
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -87,6 +133,30 @@ class MTB_StackImages:
|
||||
else:
|
||||
return tensor
|
||||
|
||||
def resize_tensor(self, tensor, width=None, height=None):
|
||||
"""Resize tensor to specified width or height while maintaining aspect ratio."""
|
||||
current_height, current_width = tensor.shape[1:3]
|
||||
|
||||
if width is not None and width != current_width:
|
||||
scale_factor = width / current_width
|
||||
new_height = int(current_height * scale_factor)
|
||||
new_width = width
|
||||
elif height is not None and height != current_height:
|
||||
scale_factor = height / current_height
|
||||
new_width = int(current_width * scale_factor)
|
||||
new_height = height
|
||||
else:
|
||||
return tensor
|
||||
|
||||
resized = torch.nn.functional.interpolate(
|
||||
tensor.permute(0, 3, 1, 2),
|
||||
size=(new_height, new_width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
return resized.permute(0, 2, 3, 1)
|
||||
|
||||
|
||||
class MTB_PickFromBatch:
|
||||
"""Pick a specific number of images from a batch.
|
||||
@@ -101,30 +171,117 @@ class MTB_PickFromBatch:
|
||||
"image": ("IMAGE",),
|
||||
"from_direction": (["end", "start"], {"default": "start"}),
|
||||
"count": ("INT", {"default": 1}),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
FUNCTION = "pick_from_batch"
|
||||
CATEGORY = "mtb/image utils"
|
||||
|
||||
def pick_from_batch(self, image, from_direction, count):
|
||||
def pick_from_batch(self, image, from_direction, count, mask=None):
|
||||
batch_size = image.size(0)
|
||||
|
||||
# Limit count to the available number of images in the batch
|
||||
count = min(count, batch_size)
|
||||
if count < batch_size:
|
||||
log.warning(
|
||||
f"Requested {count} images, "
|
||||
f"but only {batch_size} are available."
|
||||
)
|
||||
|
||||
selected_masks = None
|
||||
|
||||
if from_direction == "end":
|
||||
selected_tensors = image[-count:]
|
||||
if mask is not None:
|
||||
selected_masks = mask[-count:]
|
||||
else:
|
||||
selected_tensors = image[:count]
|
||||
if mask is not None:
|
||||
selected_masks = mask[:count]
|
||||
|
||||
return (selected_tensors,)
|
||||
return (selected_tensors, selected_masks)
|
||||
|
||||
|
||||
__nodes__ = [MTB_StackImages, MTB_PickFromBatch]
|
||||
import folder_paths
|
||||
|
||||
|
||||
class MTB_SaveImage:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
self.type = "output"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE", {"tooltip": "The images to save."}),
|
||||
"filename_prefix": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "ComfyUI",
|
||||
"tooltip": "The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "save_images"
|
||||
|
||||
# OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "mtb/image utils"
|
||||
DESCRIPTION = """Saves the input images to your ComfyUI output directory.
|
||||
This behaves exactly like the native SaveImage node but isn't an output node.
|
||||
The reason I made this is to allow 'inlining' image save in loops for instance,
|
||||
using the native node there wouldn't run for each iteration of the loop."""
|
||||
|
||||
def save_images(
|
||||
self,
|
||||
images,
|
||||
filename_prefix="ComfyUI",
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
):
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = (
|
||||
folder_paths.get_save_image_path(
|
||||
filename_prefix,
|
||||
self.output_dir,
|
||||
images[0].shape[1],
|
||||
images[0].shape[0],
|
||||
)
|
||||
)
|
||||
results = list()
|
||||
for batch_number, image in enumerate(images):
|
||||
i = 255.0 * image.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||
metadata = None
|
||||
if not args.disable_metadata:
|
||||
metadata = PngInfo()
|
||||
if prompt is not None:
|
||||
metadata.add_text("prompt", json.dumps(prompt))
|
||||
if extra_pnginfo is not None:
|
||||
for x in extra_pnginfo:
|
||||
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
|
||||
|
||||
filename_with_batch_num = filename.replace(
|
||||
"%batch_num%", str(batch_number)
|
||||
)
|
||||
file = f"{filename_with_batch_num}_{counter:05}_.png"
|
||||
img.save(
|
||||
os.path.join(full_output_folder, file),
|
||||
pnginfo=metadata,
|
||||
compress_level=self.compress_level,
|
||||
)
|
||||
results.append(
|
||||
{"filename": file, "subfolder": subfolder, "type": self.type}
|
||||
)
|
||||
counter += 1
|
||||
|
||||
return {"ui": {"images": results}, "result": (images,)}
|
||||
|
||||
|
||||
__nodes__ = [MTB_StackImages, MTB_PickFromBatch, MTB_SaveImage]
|
||||
|
||||
+55
-19
@@ -2,9 +2,9 @@ import json
|
||||
import subprocess
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -41,6 +41,7 @@ class MTB_ReadPlaylist:
|
||||
RETURN_TYPES = ("PLAYLIST",)
|
||||
FUNCTION = "read_playlist"
|
||||
CATEGORY = "mtb/IO"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def read_playlist(
|
||||
self,
|
||||
@@ -83,6 +84,7 @@ class MTB_AddToPlaylist:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "add_to_playlist"
|
||||
CATEGORY = "mtb/IO"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def add_to_playlist(
|
||||
self,
|
||||
@@ -117,7 +119,10 @@ class MTB_AddToPlaylist:
|
||||
|
||||
|
||||
class MTB_ExportWithFfmpeg:
|
||||
"""Export with FFmpeg (Experimental)"""
|
||||
"""Export with FFmpeg (Experimental).
|
||||
|
||||
[DEPRACATED] Use VHS nodes instead
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -143,6 +148,7 @@ class MTB_ExportWithFfmpeg:
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "export_prores"
|
||||
DEPRECATED = True
|
||||
CATEGORY = "mtb/IO"
|
||||
|
||||
def export_prores(
|
||||
@@ -151,10 +157,9 @@ class MTB_ExportWithFfmpeg:
|
||||
prefix: str,
|
||||
format: str,
|
||||
codec: str,
|
||||
images: Optional[torch.Tensor] = None,
|
||||
playlist: Optional[List[str]] = None,
|
||||
images: torch.Tensor | None = None,
|
||||
playlist: list[str] | None = None,
|
||||
):
|
||||
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
|
||||
file_ext = format
|
||||
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
|
||||
|
||||
@@ -208,9 +213,11 @@ class MTB_ExportWithFfmpeg:
|
||||
frames = tensor2np(images)
|
||||
log.debug(f"Frames type {type(frames[0])}")
|
||||
log.debug(f"Exporting {len(frames)} frames")
|
||||
height, width, channels = frames[0].shape
|
||||
has_alpha = channels == 4
|
||||
out_path = (output_dir / file_id).as_posix()
|
||||
|
||||
if codec == "gif":
|
||||
out_path = (output_dir / file_id).as_posix()
|
||||
command = [
|
||||
"ffmpeg",
|
||||
"-f",
|
||||
@@ -233,12 +240,28 @@ class MTB_ExportWithFfmpeg:
|
||||
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
return (out_path,)
|
||||
else:
|
||||
frames = [frame.astype(np.uint16) * 257 for frame in frames]
|
||||
|
||||
height, width, _ = frames[0].shape
|
||||
|
||||
out_path = (output_dir / file_id).as_posix()
|
||||
if has_alpha:
|
||||
if codec in ["prores_ks", "libx264", "libx265"]:
|
||||
pix_fmt = (
|
||||
"yuva444p" if codec == "prores_ks" else "yuva420p"
|
||||
)
|
||||
frames = [
|
||||
frame.astype(np.uint16) * 257 for frame in frames
|
||||
]
|
||||
else:
|
||||
log.warning(
|
||||
f"Alpha channel not supported for codec {codec}. Alpha will be ignored."
|
||||
)
|
||||
frames = [
|
||||
frame[:, :, :3].astype(np.uint16) * 257
|
||||
for frame in frames
|
||||
]
|
||||
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
|
||||
else:
|
||||
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
|
||||
frames = [frame.astype(np.uint16) * 257 for frame in frames]
|
||||
|
||||
# Prepare the FFmpeg command
|
||||
command = [
|
||||
@@ -258,17 +281,26 @@ class MTB_ExportWithFfmpeg:
|
||||
"-",
|
||||
"-c:v",
|
||||
codec,
|
||||
"-r",
|
||||
str(fps),
|
||||
"-y",
|
||||
out_path,
|
||||
]
|
||||
if codec == "prores_ks":
|
||||
command.extend(["-profile:v", "4444"])
|
||||
|
||||
command.extend(
|
||||
[
|
||||
"-r",
|
||||
str(fps),
|
||||
"-y",
|
||||
out_path,
|
||||
]
|
||||
)
|
||||
|
||||
process = subprocess.Popen(command, stdin=subprocess.PIPE)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(len(frames))
|
||||
|
||||
for frame in frames:
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
process.stdin.write(frame.tobytes())
|
||||
pbar.update(1)
|
||||
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
@@ -280,9 +312,9 @@ def prepare_animated_batch(
|
||||
batch: torch.Tensor,
|
||||
pingpong=False,
|
||||
resize_by=1.0,
|
||||
resample_filter: Optional[Image.Resampling] = None,
|
||||
resample_filter: Image.Resampling | None = None,
|
||||
image_type=np.uint8,
|
||||
) -> List[Image.Image]:
|
||||
) -> list[Image.Image]:
|
||||
images = tensor2np(batch)
|
||||
images = [frame.astype(image_type) for frame in images]
|
||||
|
||||
@@ -308,7 +340,10 @@ def prepare_animated_batch(
|
||||
|
||||
# todo: deprecate for apng
|
||||
class MTB_SaveGif:
|
||||
"""Save the images from the batch as a GIF"""
|
||||
"""Save the images from the batch as a GIF.
|
||||
|
||||
[DEPRACATED] Use VHS nodes instead
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -328,6 +363,7 @@ class MTB_SaveGif:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "mtb/IO"
|
||||
FUNCTION = "save_gif"
|
||||
DEPRECATED = True
|
||||
|
||||
def save_gif(
|
||||
self,
|
||||
|
||||
+161
@@ -0,0 +1,161 @@
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from ..log import log
|
||||
|
||||
|
||||
class ImageH264Compression:
|
||||
"""Encodes the input with h264 compression using a configurable CRF."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": "The input image tensor to be compressed and decompressed."
|
||||
},
|
||||
),
|
||||
"crf": (
|
||||
"INT",
|
||||
{
|
||||
"default": 23,
|
||||
"min": 0,
|
||||
"max": 51,
|
||||
"step": 1,
|
||||
"tooltip": "Constant Rate Factor for h264 encoding (lower values mean higher quality).",
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "compress_and_decompress"
|
||||
|
||||
CATEGORY = "image"
|
||||
DESCRIPTION = """
|
||||
**Encodes the input with h264 compression using a configurable CRF**.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> This node is not really needed with the latest version of LTXVideo.
|
||||
|
||||
> [!NOTE]
|
||||
> This was recommended by the creators of LTX over banodoco's discord.
|
||||
|
||||
*Orginal code from [mix](https://github.com/XmYx)*"""
|
||||
|
||||
def _compress_decompress_ffmpeg(self, img_array, crf):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
input_path = os.path.join(temp_dir, "input.png")
|
||||
output_path = os.path.join(temp_dir, "output.mp4")
|
||||
decoded_path = os.path.join(temp_dir, "decoded.png")
|
||||
|
||||
Image.fromarray(img_array).save(input_path)
|
||||
|
||||
encode_command = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
input_path,
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-crf",
|
||||
str(crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-frames:v",
|
||||
"1",
|
||||
output_path,
|
||||
]
|
||||
subprocess.run(encode_command, capture_output=True)
|
||||
|
||||
decode_command = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
output_path,
|
||||
"-frames:v",
|
||||
"1",
|
||||
decoded_path,
|
||||
]
|
||||
subprocess.run(decode_command, capture_output=True)
|
||||
|
||||
decoded_img = np.array(Image.open(decoded_path))
|
||||
return decoded_img
|
||||
|
||||
def compress_and_decompress(self, image, crf):
|
||||
import io
|
||||
|
||||
output_images = []
|
||||
|
||||
try:
|
||||
import av
|
||||
|
||||
for img_tensor in image:
|
||||
img_array = img_tensor.cpu().numpy()
|
||||
img_array = (img_array * 255).astype(np.uint8)
|
||||
img_array = img_array.copy(
|
||||
order="C"
|
||||
) # Ensure contiguous array
|
||||
|
||||
output = io.BytesIO()
|
||||
|
||||
# Encode the image to h264 with the given CRF
|
||||
container = av.open(output, mode="w", format="mp4")
|
||||
stream = container.add_stream("h264", rate=1)
|
||||
stream.width = img_array.shape[1]
|
||||
stream.height = img_array.shape[0]
|
||||
stream.pix_fmt = "yuv420p"
|
||||
stream.options = {"crf": str(crf)}
|
||||
|
||||
frame = av.VideoFrame.from_ndarray(img_array, format="rgb24")
|
||||
for packet in stream.encode(frame):
|
||||
container.mux(packet)
|
||||
for packet in stream.encode():
|
||||
container.mux(packet)
|
||||
container.close()
|
||||
|
||||
# Decode the video back to an image
|
||||
output.seek(0)
|
||||
container = av.open(output, mode="r", format="mp4")
|
||||
decoded_frames = []
|
||||
for frame in container.decode(video=0):
|
||||
img_decoded = frame.to_ndarray(format="rgb24")
|
||||
decoded_frames.append(img_decoded)
|
||||
container.close()
|
||||
|
||||
if len(decoded_frames) > 0:
|
||||
img_decoded = decoded_frames[0]
|
||||
img_decoded = torch.from_numpy(
|
||||
img_decoded.astype(np.float32) / 255.0
|
||||
)
|
||||
output_images.append(img_decoded)
|
||||
else:
|
||||
# If decoding failed, use the original image
|
||||
output_images.append(img_tensor)
|
||||
except ImportError:
|
||||
log.warning(
|
||||
"PyAv is not installed... Falling back to the ffmpeg cli"
|
||||
)
|
||||
for img_tensor in image:
|
||||
img_array = (img_tensor.cpu().numpy() * 255).astype(np.uint8)
|
||||
decoded_img = self._compress_decompress_ffmpeg(img_array, crf)
|
||||
img_decoded = torch.from_numpy(
|
||||
decoded_img.astype(np.float32) / 255.0
|
||||
)
|
||||
output_images.append(img_decoded)
|
||||
|
||||
output_images = torch.stack(output_images).to(image.device)
|
||||
return (output_images,)
|
||||
|
||||
|
||||
# fmt: off
|
||||
__nodes__ = [
|
||||
ImageH264Compression
|
||||
]
|
||||
+2
-1
@@ -1,6 +1,5 @@
|
||||
import comfy.utils
|
||||
from PIL import Image
|
||||
from rembg import remove
|
||||
|
||||
from ..utils import pil2tensor, tensor2pil
|
||||
|
||||
@@ -64,6 +63,8 @@ class MTB_ImageRemoveBackgroundRembg:
|
||||
post_process_mask,
|
||||
bgcolor,
|
||||
):
|
||||
from rembg import remove
|
||||
|
||||
pbar = comfy.utils.ProgressBar(image.size(0))
|
||||
images = tensor2pil(image)
|
||||
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
import comfy.utils
|
||||
import torch
|
||||
|
||||
from ..log import log
|
||||
from ..utils import nextAvailable, tensor2pil
|
||||
|
||||
RELATIVE_NOTICE = """
|
||||
Absolute paths are kept as is, relatives are from the output directory.
|
||||
"""
|
||||
|
||||
|
||||
class MTB_PostshotTrain:
|
||||
CATEGORY = "mtb/postshot"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": (
|
||||
"IMAGE",
|
||||
{"tooltip": "These image will get save to disk first"},
|
||||
),
|
||||
"profile": (
|
||||
[
|
||||
"NeRF L",
|
||||
"NeRF M",
|
||||
"NeRF S",
|
||||
"NeRF XL",
|
||||
"NeRF XXL",
|
||||
"Splat ADC",
|
||||
"Splat MCMC",
|
||||
],
|
||||
{
|
||||
"default": "Splat MCMC",
|
||||
"tooltip": "The radiance field model profile to train",
|
||||
},
|
||||
),
|
||||
"image_select": (
|
||||
["all", "best"],
|
||||
{
|
||||
"default": "best",
|
||||
"tooltip": "How to select training images from the source image sets",
|
||||
},
|
||||
),
|
||||
"train_steps_limit": (
|
||||
"INT",
|
||||
{
|
||||
"default": 30,
|
||||
"min": 1,
|
||||
"max": 1000,
|
||||
"tooltip": "Number of kSteps to train the model for",
|
||||
},
|
||||
),
|
||||
"output_path": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "output",
|
||||
"tooltip": (
|
||||
"path to save the project to" f"{RELATIVE_NOTICE}"
|
||||
),
|
||||
},
|
||||
),
|
||||
"postshot_cli": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"gpu": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 255,
|
||||
"tooltip": "Specify the index of the GPU to use",
|
||||
},
|
||||
),
|
||||
"num_train_images": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"tooltip": "If image-select best is used, specifies the number of training images to select",
|
||||
},
|
||||
),
|
||||
"max_image_size": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1600,
|
||||
"min": 0,
|
||||
"tooltip": "Downscale training images such that their longer edge is at most this value in pixels. Disabled if zero.",
|
||||
},
|
||||
),
|
||||
"max_num_features": (
|
||||
"INT",
|
||||
{
|
||||
"default": 8,
|
||||
"min": 1,
|
||||
"tooltip": "Maximum number of 2D kFeatures extracted from each image.",
|
||||
},
|
||||
),
|
||||
"splat_density": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.125,
|
||||
"max": 8.0,
|
||||
"tooltip": (
|
||||
"Controls how much additional splats "
|
||||
"are generated during training."
|
||||
"Applies only in 'Splat ADC' profile."
|
||||
),
|
||||
},
|
||||
),
|
||||
"max_num_splats": (
|
||||
"INT",
|
||||
{
|
||||
"default": 3000,
|
||||
"min": 1,
|
||||
"tooltip": (
|
||||
"Sets the maximum number of splats (in kSplats)"
|
||||
" created during training. "
|
||||
"Applies only in 'Splat MCMC' profile."
|
||||
),
|
||||
},
|
||||
),
|
||||
"export_splat_ply": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"If not empty will also save a ply file."
|
||||
f"{RELATIVE_NOTICE}"
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
OUTPUT_NODE = True
|
||||
RETURN_NAMES = ("project_file_path",)
|
||||
FUNCTION = "train_model"
|
||||
|
||||
def train_model(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
profile: str,
|
||||
image_select: str,
|
||||
train_steps_limit: int,
|
||||
output_path: str,
|
||||
gpu=0,
|
||||
num_train_images=0,
|
||||
max_image_size=1600,
|
||||
max_num_features=8,
|
||||
splat_density=1.0,
|
||||
max_num_splats=3000,
|
||||
export_splat_ply="",
|
||||
postshot_cli="",
|
||||
):
|
||||
if not output_path.endswith(".psht"):
|
||||
output_path += ".psht"
|
||||
|
||||
output_path = nextAvailable(output_path)
|
||||
output_path.parent.mkdir(exist_ok=True)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(200 + images.size(0))
|
||||
|
||||
try:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
image_paths = []
|
||||
images_pil = tensor2pil(images)
|
||||
for i, img in enumerate(images_pil):
|
||||
try:
|
||||
img_path = os.path.join(temp_dir, f"image_{i:04d}.png")
|
||||
img.save(img_path)
|
||||
image_paths.append(img_path)
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to save image {i}: {str(e)}"
|
||||
) from e
|
||||
pbar.update(1)
|
||||
|
||||
if not image_paths:
|
||||
raise ValueError("No valid images to process")
|
||||
|
||||
cmd = [postshot_cli, "train"]
|
||||
|
||||
for img_path in image_paths:
|
||||
cmd.extend(["-i", img_path])
|
||||
|
||||
cmd.extend(
|
||||
[
|
||||
"-p",
|
||||
profile,
|
||||
"--image-select",
|
||||
image_select,
|
||||
"-s",
|
||||
str(train_steps_limit),
|
||||
"-o",
|
||||
output_path.as_posix(),
|
||||
]
|
||||
)
|
||||
|
||||
if gpu is not None:
|
||||
cmd.extend(["--gpu", str(gpu)])
|
||||
if num_train_images > 0 and image_select == "best":
|
||||
cmd.extend(["--num-train-images", str(num_train_images)])
|
||||
if max_image_size > 0:
|
||||
cmd.extend(["--max-image-size", str(max_image_size)])
|
||||
if max_num_features != 8:
|
||||
cmd.extend(["--max-num-features", str(max_num_features)])
|
||||
if profile == "Splat ADC" and splat_density != 1.0:
|
||||
cmd.extend(["--splat-density", str(splat_density)])
|
||||
if profile == "Splat MCMC" and max_num_splats != 3000:
|
||||
cmd.extend(["--max-num-splats", str(max_num_splats)])
|
||||
if export_splat_ply:
|
||||
export_splat_ply = nextAvailable(export_splat_ply)
|
||||
cmd.extend(
|
||||
["--export-splat-ply", export_splat_ply.as_posix()]
|
||||
)
|
||||
|
||||
log.debug(f"Running {cmd}")
|
||||
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
universal_newlines=True,
|
||||
)
|
||||
|
||||
last_step_c = 0
|
||||
last_step_t = 0
|
||||
while True:
|
||||
output = process.stdout.readline()
|
||||
if output == "" and process.poll() is not None:
|
||||
break
|
||||
if output:
|
||||
print(output)
|
||||
if "camera tracking step" in output.lower():
|
||||
try:
|
||||
current_step = int(
|
||||
output.split("%")[0].split(":")[1].strip()
|
||||
)
|
||||
if current_step > last_step_c:
|
||||
pbar.update(1)
|
||||
last_step_c = current_step
|
||||
|
||||
except (ValueError, IndexError):
|
||||
continue
|
||||
|
||||
if "training radiance field:" in output.lower():
|
||||
try:
|
||||
current_step = int(
|
||||
output.split("%")[0].split(":")[1].strip()
|
||||
)
|
||||
if current_step > last_step_t:
|
||||
pbar.update(1)
|
||||
last_step_t = current_step
|
||||
|
||||
except (ValueError, IndexError):
|
||||
continue
|
||||
|
||||
if process.returncode != 0:
|
||||
_, stderr = process.communicate()
|
||||
raise RuntimeError(f"Postshot training failed: {stderr}")
|
||||
|
||||
if not os.path.exists(output_path):
|
||||
raise RuntimeError("Output file was not created")
|
||||
|
||||
return (output_path.as_posix(),)
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Training failed: {str(e)}")
|
||||
finally:
|
||||
pbar.update(train_steps_limit)
|
||||
|
||||
|
||||
class MTB_PostshotExport:
|
||||
CATEGORY = "mtb/postshot"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"project_file": (
|
||||
"STRING",
|
||||
{"default": "", "forceInput": True},
|
||||
),
|
||||
"export_splat_ply": ("STRING", {"default": "output.ply"}),
|
||||
"postshot_cli": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("exported_ply_path",)
|
||||
FUNCTION = "export_model"
|
||||
|
||||
def export_model(
|
||||
self, project_file: str, export_splat_ply: str, postshot_cli: str
|
||||
):
|
||||
if not project_file.endswith(".psht"):
|
||||
raise ValueError("Project file must have .psht extension")
|
||||
|
||||
if not os.path.exists(project_file):
|
||||
raise FileNotFoundError(f"Project file not found: {project_file}")
|
||||
|
||||
if not export_splat_ply.endswith(".ply"):
|
||||
export_splat_ply += ".ply"
|
||||
|
||||
_export_splat_ply = nextAvailable(export_splat_ply)
|
||||
_export_splat_ply.parent.mkdir(exist_ok=True)
|
||||
|
||||
cmd = [
|
||||
postshot_cli,
|
||||
"export",
|
||||
"-f",
|
||||
project_file,
|
||||
"--export-splat-ply",
|
||||
_export_splat_ply.as_posix(),
|
||||
]
|
||||
|
||||
try:
|
||||
_result = subprocess.run(
|
||||
cmd, check=True, capture_output=True, text=True
|
||||
)
|
||||
|
||||
if not _export_splat_ply.exists():
|
||||
log.error("Export file was not created")
|
||||
|
||||
return (_export_splat_ply.as_posix(),)
|
||||
|
||||
except subprocess.CalledProcessError as e:
|
||||
raise RuntimeError(f"Export failed: {e.stderr}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Export failed: {str(e)}")
|
||||
|
||||
|
||||
__nodes__ = [MTB_PostshotExport, MTB_PostshotTrain]
|
||||
@@ -0,0 +1,85 @@
|
||||
import qrcode
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from ..log import log
|
||||
from ..utils import pil2tensor
|
||||
|
||||
|
||||
class MTB_QrCode:
|
||||
"""Basic QR Code generator."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": "https://www.github.com"}),
|
||||
"width": (
|
||||
"INT",
|
||||
{"default": 256, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"height": (
|
||||
"INT",
|
||||
{"default": 256, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"error_correct": (("L", "M", "Q", "H"), {"default": "L"}),
|
||||
"box_size": (
|
||||
"INT",
|
||||
{"default": 10, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"border": (
|
||||
"INT",
|
||||
{"default": 4, "max": 8096, "min": 0, "step": 1},
|
||||
),
|
||||
"invert": (("BOOLEAN",), {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_qr"
|
||||
CATEGORY = "mtb/generate"
|
||||
|
||||
def do_qr(
|
||||
self,
|
||||
*,
|
||||
url: str,
|
||||
width: int,
|
||||
height: int,
|
||||
error_correct: str,
|
||||
box_size: int,
|
||||
border: int,
|
||||
invert: bool,
|
||||
) -> tuple[torch.Tensor]:
|
||||
log.warning(
|
||||
"This node will soon be deprecated, there are much better alternatives like https://github.com/coreyryanhanson/comfy-qr"
|
||||
)
|
||||
if error_correct == "L" or error_correct not in ["M", "Q", "H"]:
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_L
|
||||
elif error_correct == "M":
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_M
|
||||
elif error_correct == "Q":
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_Q
|
||||
else:
|
||||
error_correct = qrcode.constants.ERROR_CORRECT_H
|
||||
|
||||
qr = qrcode.QRCode(
|
||||
version=1,
|
||||
error_correction=error_correct,
|
||||
box_size=box_size,
|
||||
border=border,
|
||||
)
|
||||
qr.add_data(url)
|
||||
qr.make(fit=True)
|
||||
|
||||
back_color = (255, 255, 255) if invert else (0, 0, 0)
|
||||
fill_color = (0, 0, 0) if invert else (255, 255, 255)
|
||||
|
||||
code = qr.make_image(back_color=back_color, fill_color=fill_color)
|
||||
|
||||
# that we now resize without filtering
|
||||
code = code.resize((width, height), Image.NEAREST)
|
||||
|
||||
return (pil2tensor(code),)
|
||||
|
||||
|
||||
__nodes__ = [MTB_QrCode]
|
||||
+91
-11
@@ -45,6 +45,34 @@ class MTB_TransformImage:
|
||||
),
|
||||
"constant_color": ("COLOR", {"default": "#000000"}),
|
||||
},
|
||||
"optional": {
|
||||
"filter_type": (
|
||||
[
|
||||
"nearest",
|
||||
"box",
|
||||
"bilinear",
|
||||
"hamming",
|
||||
"bicubic",
|
||||
"lanczos",
|
||||
],
|
||||
{"default": "bilinear"},
|
||||
),
|
||||
"stretch_x": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.001, "max": 10.0, "step": 0.01},
|
||||
),
|
||||
"stretch_y": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.001, "max": 10.0, "step": 0.01},
|
||||
),
|
||||
"use_normalized": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "If true, transform values are scaled to image dimensions.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
FUNCTION = "transform"
|
||||
@@ -61,21 +89,36 @@ class MTB_TransformImage:
|
||||
shear: float,
|
||||
border_handling="edge",
|
||||
constant_color=None,
|
||||
filter_type="nearest",
|
||||
stretch_x=1.0,
|
||||
stretch_y=1.0,
|
||||
use_normalized: bool = False,
|
||||
):
|
||||
filter_map = {
|
||||
"nearest": Image.NEAREST,
|
||||
"box": Image.BOX,
|
||||
"bilinear": Image.BILINEAR,
|
||||
"hamming": Image.HAMMING,
|
||||
"bicubic": Image.BICUBIC,
|
||||
"lanczos": Image.LANCZOS,
|
||||
}
|
||||
resampling_filter = filter_map[filter_type]
|
||||
|
||||
_, frame_height, frame_width, _ = image.size()
|
||||
if use_normalized:
|
||||
x = float(x) * frame_width
|
||||
y = float(y) * frame_height
|
||||
x = int(x)
|
||||
y = int(y)
|
||||
angle = int(angle)
|
||||
|
||||
log.debug(
|
||||
f"Zoom: {zoom} | x: {x}, y: {y}, angle: {angle}, shear: {shear}"
|
||||
f"Zoom: {zoom} | x: {x}, y: {y}, angle: {angle}, shear: {shear} | stretch_x: {stretch_x}, stretch_y: {stretch_y}"
|
||||
)
|
||||
|
||||
if image.size(0) == 0:
|
||||
return (torch.zeros(0),)
|
||||
transformed_images = []
|
||||
frames_count, frame_height, frame_width, frame_channel_count = (
|
||||
image.size()
|
||||
)
|
||||
|
||||
new_height, new_width = (
|
||||
int(frame_height * zoom),
|
||||
@@ -106,18 +149,55 @@ class MTB_TransformImage:
|
||||
|
||||
for img in tensor2pil(image):
|
||||
img = TF.pad(
|
||||
img, # transformed_frame,
|
||||
img,
|
||||
padding=padding,
|
||||
padding_mode=border_handling,
|
||||
fill=constant_color or 0,
|
||||
)
|
||||
|
||||
img = cast(
|
||||
Image.Image,
|
||||
TF.affine(
|
||||
img, angle=angle, scale=zoom, translate=[x, y], shear=shear
|
||||
),
|
||||
)
|
||||
if stretch_x != 1.0 or stretch_y != 1.0:
|
||||
img = cast(
|
||||
Image.Image,
|
||||
TF.affine(
|
||||
img,
|
||||
angle=angle,
|
||||
scale=zoom,
|
||||
translate=[x, y],
|
||||
shear=shear,
|
||||
interpolation=resampling_filter,
|
||||
),
|
||||
)
|
||||
|
||||
width, height = img.size
|
||||
center = (width // 2, height // 2)
|
||||
|
||||
stretch_x_factor = 1.0 / stretch_x
|
||||
stretch_y_factor = 1.0 / stretch_y
|
||||
|
||||
matrix = [
|
||||
stretch_x_factor,
|
||||
0,
|
||||
center[0] - center[0] * stretch_x_factor,
|
||||
0,
|
||||
stretch_y_factor,
|
||||
center[1] - center[1] * stretch_y_factor,
|
||||
]
|
||||
|
||||
img = img.transform(
|
||||
img.size, Image.AFFINE, matrix, resampling_filter
|
||||
)
|
||||
else:
|
||||
img = cast(
|
||||
Image.Image,
|
||||
TF.affine(
|
||||
img,
|
||||
angle=angle,
|
||||
scale=zoom,
|
||||
translate=[x, y],
|
||||
shear=shear,
|
||||
interpolation=resampling_filter,
|
||||
),
|
||||
)
|
||||
|
||||
left = abs(padding[0])
|
||||
upper = abs(padding[1])
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
from ..utils import models_dir, np2tensor
|
||||
|
||||
# TODO: check if I can make a torch script device independant
|
||||
# for now I forced it to use cuda.
|
||||
|
||||
|
||||
class MTB_LoadVitMatteModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"kind": (("Composition-1K", "Distinctions-646"),),
|
||||
"autodownload": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VITMATTE_MODEL",)
|
||||
RETURN_NAMES = ("torch_script",)
|
||||
CATEGORY = "mtb/vitmatte"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(self, *, kind: str, autodownload: bool):
|
||||
dest = models_dir / "vitmatte"
|
||||
dest.mkdir(exist_ok=True)
|
||||
name = "dist" if kind == "Distinctions-646" else "com"
|
||||
|
||||
file = hf_hub_download(
|
||||
repo_id="melmass/pytorch-scripts",
|
||||
filename=f"vitmatte_b_{name}.pt",
|
||||
local_dir=dest.as_posix(),
|
||||
local_files_only=not autodownload,
|
||||
)
|
||||
model = torch.jit.load(file).to("cuda")
|
||||
|
||||
return (model,)
|
||||
|
||||
|
||||
class MTB_GenerateTrimap:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
# "image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
"erode": ("INT", {"default": 10}),
|
||||
"dilate": ("INT", {"default": 10}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("trimap",)
|
||||
|
||||
CATEGORY = "mtb/vitmatte"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(
|
||||
self,
|
||||
# image:torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
erode: int = 10,
|
||||
dilate: int = 10,
|
||||
):
|
||||
# TODO: not sure what's the most practical between IMAGE or MASK
|
||||
|
||||
# image = image.to("cuda").half()
|
||||
mask = mask.to("cuda").half()
|
||||
|
||||
trimaps = []
|
||||
for m in mask:
|
||||
mask_arr = m.squeeze(0).to(torch.uint8).cpu().numpy() * 255
|
||||
erode_kernel = np.ones((erode, erode), np.uint8)
|
||||
dilate_kernel = np.ones((dilate, dilate), np.uint8)
|
||||
eroded = cv2.erode(mask_arr, erode_kernel, iterations=5)
|
||||
dilated = cv2.dilate(mask_arr, dilate_kernel, iterations=5)
|
||||
trimap = np.zeros_like(mask_arr)
|
||||
trimap[dilated == 255] = 128
|
||||
trimap[eroded == 255] = 255
|
||||
trimaps.append(trimap)
|
||||
|
||||
return (np2tensor(trimaps),)
|
||||
|
||||
|
||||
class MTB_ApplyVitMatte:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("VITMATTE_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
"trimap": ("IMAGE",),
|
||||
"returns": (("RGB", "RGBA"),),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image (rgba)", "mask")
|
||||
CATEGORY = "mtb/utils"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(
|
||||
self, model, image: torch.Tensor, trimap: torch.Tensor, returns: str
|
||||
):
|
||||
im_count = image.shape[0]
|
||||
tm_count = trimap.shape[0]
|
||||
|
||||
if im_count != tm_count:
|
||||
raise ValueError("image and trimap must have the same batch size")
|
||||
|
||||
outputs_m: list[torch.Tensor] = []
|
||||
outputs_i: list[torch.Tensor] = []
|
||||
for i, im in enumerate(image):
|
||||
tm = trimap[i].half().unsqueeze(2).permute(2, 0, 1).to("cuda")
|
||||
im = im.half().permute(2, 0, 1).to("cuda")
|
||||
|
||||
inputs = {"image": im.unsqueeze(0), "trimap": tm.unsqueeze(0)}
|
||||
|
||||
fine_mask = model(inputs)
|
||||
foreground = im * fine_mask + (1 - fine_mask)
|
||||
|
||||
if returns == "RGBA":
|
||||
rgba_image = torch.cat(
|
||||
(foreground, fine_mask.unsqueeze(0)), dim=0
|
||||
)
|
||||
outputs_i.append(rgba_image.unsqueeze(0))
|
||||
else:
|
||||
outputs_i.append(foreground.unsqueeze(0))
|
||||
|
||||
outputs_m.append(fine_mask.unsqueeze(0))
|
||||
|
||||
result_m = torch.cat(outputs_m, dim=0)
|
||||
result_i = torch.cat(outputs_i, dim=0)
|
||||
|
||||
return (result_i.permute(0, 2, 3, 1), result_m)
|
||||
|
||||
|
||||
__nodes__ = [MTB_LoadVitMatteModel, MTB_GenerateTrimap, MTB_ApplyVitMatte]
|
||||
+182
-179
@@ -1,179 +1,182 @@
|
||||
[build-system]
|
||||
requires = ["setuptools", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.1.5"
|
||||
description = "Animation oriented nodes pack for ComfyUI."
|
||||
license = "MIT"
|
||||
readme = "README.md"
|
||||
# repository = ""
|
||||
# url = "https://github.com/melMass/comfy_mtb"
|
||||
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Operating System :: OS Independent",
|
||||
"Programming Language :: Python",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Intended Audience :: Developers",
|
||||
]
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"qrcode",
|
||||
"onnxruntime-gpu",
|
||||
"requirements-parserx",
|
||||
"rembg",
|
||||
"imageio_ffmpeg",
|
||||
"rich",
|
||||
"rich_argparse",
|
||||
"matplotlib",
|
||||
"pillow",
|
||||
]
|
||||
optional-dependencies = { mel = [
|
||||
"jupyterlab==4.1.6",
|
||||
], dev = [
|
||||
"black[jupyter]",
|
||||
"codespell",
|
||||
"mypy",
|
||||
"pre-commit",
|
||||
"pytest",
|
||||
"pytest-cov",
|
||||
"pytest-random-order",
|
||||
"ruff",
|
||||
], doc = [
|
||||
"docutils==0.17.1",
|
||||
"jupyter-book>=0.15",
|
||||
"sphinx-autobuild",
|
||||
] }
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/melMass/comfy_mtb"
|
||||
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
|
||||
Repository = "https://github.com/melMass/comfy_mtb"
|
||||
Issues = "https://github.com/melMass/comfy_mtb/issues"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "mel"
|
||||
DisplayName = "comfy-mtb"
|
||||
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||
|
||||
[tool.bumpversion]
|
||||
current_version = "0.1.5"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
replace = "{new_version}"
|
||||
regex = false
|
||||
ignore_missing_version = false
|
||||
ignore_missing_files = false
|
||||
tag = true
|
||||
sign_tags = true
|
||||
tag_name = "v{new_version}"
|
||||
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
allow_dirty = true
|
||||
commit = true
|
||||
message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
commit_args = ""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "__init__.py"
|
||||
search = "__version__ = \"{current_version}\""
|
||||
replace = "__version__ = \"{new_version}\""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "pyproject.toml"
|
||||
search = "version = \"{current_version}\""
|
||||
replace = "version = \"{new_version}\""
|
||||
|
||||
# [[tool.bumpversion.files]]
|
||||
# filename = "your_package/__init__.py"
|
||||
# search = "__version__ = '{current_version}'"
|
||||
# replace = "__version__ = '{new_version}'"
|
||||
|
||||
# INFO: All those remaining keys are meant for local dev
|
||||
[tool.pyright]
|
||||
include = ["."]
|
||||
exclude = [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
"src/experimental",
|
||||
"src/typestubs",
|
||||
]
|
||||
ignore = ["src/oldstuff"]
|
||||
defineConstant = { DEBUG = true }
|
||||
extraPaths = ["python", "../.."]
|
||||
stubPath = "src/stubs"
|
||||
|
||||
reportMissingImports = true
|
||||
reportMissingTypeStubs = false
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
pythonVersion = "3.10"
|
||||
pythonPlatform = "Windows"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_level = "DEBUG"
|
||||
log_cli = true
|
||||
markers = [
|
||||
"wip: tests that aren't fully finished yet",
|
||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||
|
||||
]
|
||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
line_length = 88
|
||||
auto_identify_namespace_packages = false
|
||||
# NOTE:
|
||||
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
||||
force_single_line = false
|
||||
known_first_party = ["mtb"]
|
||||
extend_skip = ["archives"]
|
||||
combine_straight_imports = true
|
||||
|
||||
[tool.coverage.run]
|
||||
parallel = true
|
||||
source = ["docs", "tests", "comfy-mtb"]
|
||||
|
||||
[tool.coverage.report]
|
||||
fail_under = 90
|
||||
show_missing = true
|
||||
|
||||
[tool.coverage.html]
|
||||
show_contexts = true
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 79
|
||||
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
||||
# NOTE:
|
||||
# D102 - undocumented-public-method (noisy)
|
||||
# D103 - undocumented-public-function (noisy)
|
||||
# D100 - undocumented-public-module (noisy)
|
||||
# N802 - invalid-function-name (forced by comfy's arch)
|
||||
ignore = ["D103", "D102", "D100", "N802"]
|
||||
# exclude auto generated file
|
||||
extend-exclude = ["./docs/conf.py"]
|
||||
|
||||
[tool.ruff.per-file-ignores]
|
||||
# imported but unused
|
||||
"__init__.py" = ["F401"]
|
||||
# use of assert detected
|
||||
"tests/*" = ["S101"]
|
||||
|
||||
[tool.ruff.pydocstyle]
|
||||
convention = "numpy"
|
||||
|
||||
[tool.mypy]
|
||||
pretty = true
|
||||
ignore_missing_imports = true
|
||||
# exclude auto generated file
|
||||
exclude = ["docs/conf.py"]
|
||||
|
||||
[tool.codespell]
|
||||
# exclude auto generated file
|
||||
skip = "./docs/conf.py,poetry.lock"
|
||||
check-filenames = true
|
||||
[build-system]
|
||||
requires = ["setuptools", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.3.0"
|
||||
description = "Animation oriented nodes pack for ComfyUI."
|
||||
license = { text = "MIT" }
|
||||
readme = "README.md"
|
||||
# repository = ""
|
||||
# url = "https://github.com/melMass/comfy_mtb"
|
||||
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Operating System :: OS Independent",
|
||||
"Programming Language :: Python",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Intended Audience :: Developers",
|
||||
]
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"qrcode",
|
||||
"cachetools",
|
||||
"onnxruntime-gpu",
|
||||
"requirements-parserx",
|
||||
"rembg",
|
||||
"imageio_ffmpeg",
|
||||
"rich",
|
||||
"rich_argparse",
|
||||
"matplotlib",
|
||||
"pillow",
|
||||
]
|
||||
optional-dependencies = { mel = [
|
||||
"jupyterlab==4.1.6",
|
||||
], dev = [
|
||||
"black[jupyter]",
|
||||
"codespell",
|
||||
"marimo",
|
||||
"mypy",
|
||||
"pre-commit",
|
||||
"pytest",
|
||||
"pytest-cov",
|
||||
"pytest-random-order",
|
||||
"ruff",
|
||||
], doc = [
|
||||
"docutils==0.17.1",
|
||||
"jupyter-book>=0.15",
|
||||
"sphinx-autobuild",
|
||||
] }
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/melMass/comfy_mtb"
|
||||
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
|
||||
Repository = "https://github.com/melMass/comfy_mtb"
|
||||
Issues = "https://github.com/melMass/comfy_mtb/issues"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "mel"
|
||||
DisplayName = "comfy-mtb"
|
||||
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||
|
||||
[tool.bumpversion]
|
||||
current_version = "0.3.0"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
replace = "{new_version}"
|
||||
regex = false
|
||||
ignore_missing_version = false
|
||||
ignore_missing_files = false
|
||||
tag = true
|
||||
sign_tags = true
|
||||
tag_name = "v{new_version}"
|
||||
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
allow_dirty = true
|
||||
commit = true
|
||||
message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
commit_args = ""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "__init__.py"
|
||||
search = "__version__ = \"{current_version}\""
|
||||
replace = "__version__ = \"{new_version}\""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "pyproject.toml"
|
||||
search = "version = \"{current_version}\""
|
||||
replace = "version = \"{new_version}\""
|
||||
|
||||
# [[tool.bumpversion.files]]
|
||||
# filename = "your_package/__init__.py"
|
||||
# search = "__version__ = '{current_version}'"
|
||||
# replace = "__version__ = '{new_version}'"
|
||||
|
||||
# INFO: All those remaining keys are meant for local dev
|
||||
[tool.pyright]
|
||||
include = ["."]
|
||||
exclude = [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
"src/experimental",
|
||||
"src/typestubs",
|
||||
]
|
||||
ignore = ["src/oldstuff"]
|
||||
defineConstant = { DEBUG = true }
|
||||
extraPaths = ["python", "../.."]
|
||||
stubPath = "src/stubs"
|
||||
|
||||
reportMissingImports = true
|
||||
reportMissingTypeStubs = false
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
pythonVersion = "3.10"
|
||||
pythonPlatform = "Windows"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_level = "DEBUG"
|
||||
log_cli = true
|
||||
markers = [
|
||||
"wip: tests that aren't fully finished yet",
|
||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||
|
||||
]
|
||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
line_length = 88
|
||||
auto_identify_namespace_packages = false
|
||||
# NOTE:
|
||||
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
||||
force_single_line = false
|
||||
known_first_party = ["mtb"]
|
||||
extend_skip = ["archives"]
|
||||
combine_straight_imports = true
|
||||
|
||||
[tool.coverage.run]
|
||||
parallel = true
|
||||
source = ["docs", "tests", "comfy-mtb"]
|
||||
|
||||
[tool.coverage.report]
|
||||
fail_under = 90
|
||||
show_missing = true
|
||||
|
||||
[tool.coverage.html]
|
||||
show_contexts = true
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 79
|
||||
extend-exclude = ["./docs/conf.py", "notebooks", "stubs"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
||||
# NOTE:
|
||||
# D102 - undocumented-public-method (noisy)
|
||||
# D103 - undocumented-public-function (noisy)
|
||||
# D100 - undocumented-public-module (noisy)
|
||||
# N802 - invalid-function-name (forced by comfy's arch)
|
||||
ignore = ["D103", "D102", "D100", "N802"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
# imported but unused
|
||||
"__init__.py" = ["F401"]
|
||||
# use of assert detected
|
||||
"tests/*" = ["S101"]
|
||||
|
||||
[tool.ruff.lint.pydocstyle]
|
||||
convention = "numpy"
|
||||
|
||||
[tool.mypy]
|
||||
pretty = true
|
||||
ignore_missing_imports = true
|
||||
# exclude auto generated file
|
||||
exclude = ["docs/conf.py"]
|
||||
|
||||
[tool.codespell]
|
||||
# exclude auto generated file
|
||||
skip = "./docs/conf.py,poetry.lock"
|
||||
check-filenames = true
|
||||
|
||||
@@ -8,3 +8,5 @@ rich
|
||||
rich_argparse
|
||||
matplotlib
|
||||
pillow
|
||||
cachetools
|
||||
transformers
|
||||
|
||||
@@ -16,3 +16,50 @@
|
||||
* @typedef {import("./shared.d.ts").INodeOutputSlot} INodeOutputSlot
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Object} ResultItem
|
||||
* @property {string} [filename] - The filename of the item.
|
||||
* @property {string} [subfolder] - The subfolder of the item.
|
||||
* @property {string} [type] - The type of the item.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Object} Outputs
|
||||
* @property {ResultItem[]} [audio] - Audio result items.
|
||||
* @property {ResultItem[]} [images] - Image result items.
|
||||
* @property {ResultItem[]} [animated] - Animated result items.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Record<string, Outputs>} TaskOutput
|
||||
* - A record mapping Node IDs to their Outputs.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Array} TaskPrompt
|
||||
* @property {QueueIndex} [0] - The queue index.
|
||||
* @property {PromptId} [1] - The unique prompt ID.
|
||||
* @property {PromptInputs} [2] - The prompt inputs.
|
||||
* @property {ExtraData} [3] - Extra data.
|
||||
* @property {OutputsToExecute} [4] - The outputs to execute.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Object} HistoryTaskItem
|
||||
* @property {'History'} taskType - The type of task.
|
||||
* @property {TaskPrompt} prompt - The task prompt.
|
||||
* @property {Status} [status] - The status of the task.
|
||||
* @property {TaskOutput} outputs - The task outputs.
|
||||
* @property {TaskMeta} [meta] - Optional task metadata.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Object} ExecInfo
|
||||
* @property {number} queue_remaining - The number of items remaining in the queue.
|
||||
*/
|
||||
|
||||
/**
|
||||
* @typedef {Object} StatusWsMessageStatus
|
||||
* @property {ExecInfo} exec_info - Execution information.
|
||||
*/
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import contextlib
|
||||
import functools
|
||||
import importlib
|
||||
import math
|
||||
import operator
|
||||
import os
|
||||
import shlex
|
||||
import shutil
|
||||
@@ -9,12 +10,17 @@ import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import uuid
|
||||
from collections.abc import Callable, Sequence
|
||||
from enum import Enum
|
||||
from functools import reduce
|
||||
from pathlib import Path
|
||||
from typing import TypeVar
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
import requests
|
||||
import torch
|
||||
from PIL import Image
|
||||
@@ -161,9 +167,9 @@ class IPChecker:
|
||||
def __init__(self):
|
||||
self.ips = list(self.get_local_ips())
|
||||
log.debug(f"Found {len(self.ips)} local ips")
|
||||
self.checked_ips = set()
|
||||
self.checked_ips: set[str] = set()
|
||||
|
||||
def get_working_ip(self, test_url_template):
|
||||
def get_working_ip(self, test_url_template: str):
|
||||
for ip in self.ips:
|
||||
if ip not in self.checked_ips:
|
||||
self.checked_ips.add(ip)
|
||||
@@ -173,7 +179,7 @@ class IPChecker:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_local_ips(prefix="192.168."):
|
||||
def get_local_ips(prefix: str = "192.168."):
|
||||
hostname = socket.gethostname()
|
||||
log.debug(f"Getting local ips for {hostname}")
|
||||
for info in socket.getaddrinfo(hostname, None):
|
||||
@@ -183,9 +189,9 @@ class IPChecker:
|
||||
if info[0] == socket.AF_INET and info[4][0].startswith(prefix):
|
||||
yield info[4][0]
|
||||
|
||||
def _test_url(self, url):
|
||||
def _test_url(self, url: str):
|
||||
try:
|
||||
response = requests.get(url)
|
||||
response = requests.get(url, timeout=10)
|
||||
return response.status_code == 200
|
||||
except Exception:
|
||||
return False
|
||||
@@ -196,7 +202,7 @@ def get_server_info():
|
||||
from comfy.cli_args import args
|
||||
|
||||
ip_checker = IPChecker()
|
||||
base_url = args.listen
|
||||
base_url: str = args.listen
|
||||
if base_url == "0.0.0.0":
|
||||
log.debug("Server set to 0.0.0.0, we will try to resolve the host IP")
|
||||
base_url = ip_checker.get_working_ip(
|
||||
@@ -210,6 +216,37 @@ def get_server_info():
|
||||
|
||||
|
||||
# region MISC Utilities
|
||||
def glob_multiple(
|
||||
path: Path, patterns: list[str], recursive: bool = False
|
||||
) -> list[Path]:
|
||||
"""Combine multiple glob patterns into a single iterator."""
|
||||
return list(reduce(operator.or_, (set(path.glob(p)) for p in patterns)))
|
||||
|
||||
|
||||
def build_glob_patterns(
|
||||
extensions: list[str], recursive: bool = False
|
||||
) -> list[str]:
|
||||
"""Build glob patterns for given extensions."""
|
||||
prefix = "**/" if recursive else ""
|
||||
return [f"{prefix}*.{ext}" for ext in extensions]
|
||||
|
||||
|
||||
class SortMode(Enum):
|
||||
NONE = "none"
|
||||
MODIFIED = "modified"
|
||||
MODIFIED_REVERSE = "modified-reverse"
|
||||
NAME = "name"
|
||||
NAME_REVERSE = "name-reverse"
|
||||
|
||||
@classmethod
|
||||
def from_str(cls, value: str | None) -> "SortMode|None":
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return cls(value.lower())
|
||||
except ValueError:
|
||||
log.warning(f"Sort mode {value} not supported")
|
||||
return None
|
||||
|
||||
|
||||
# TODO: use mtb.core directly instead of copying parts here
|
||||
@@ -463,7 +500,12 @@ here = Path(__file__).parent.absolute()
|
||||
# - Construct the absolute path to the ComfyUI directory
|
||||
comfy_dir = Path(folder_paths.base_path)
|
||||
models_dir = Path(folder_paths.models_dir)
|
||||
|
||||
|
||||
# NOTE: these aren't reliable, better call the getters each time
|
||||
output_dir = Path(folder_paths.output_directory)
|
||||
input_dir = Path(folder_paths.input_directory)
|
||||
|
||||
styles_dir = comfy_dir / "styles"
|
||||
session_id = str(uuid.uuid4())
|
||||
# - Construct the path to the font file
|
||||
@@ -473,9 +515,10 @@ font_path = here / "data" / "font.ttf"
|
||||
extern_root = here / "extern"
|
||||
add_path(extern_root)
|
||||
|
||||
for pth in extern_root.iterdir():
|
||||
if pth.is_dir():
|
||||
add_path(pth)
|
||||
if extern_root.exists():
|
||||
for pth in extern_root.iterdir():
|
||||
if pth.is_dir():
|
||||
add_path(pth)
|
||||
|
||||
# - Add the ComfyUI directory and custom nodes path to the sys.path list
|
||||
add_path(comfy_dir)
|
||||
@@ -501,52 +544,123 @@ PIL_FILTER_MAP = {
|
||||
|
||||
|
||||
# region TENSOR Utilities
|
||||
def tensor2pil(image: torch.Tensor) -> list[Image.Image]:
|
||||
batch_count = image.size(0) if len(image.shape) > 3 else 1
|
||||
if batch_count > 1:
|
||||
out = []
|
||||
for i in range(batch_count):
|
||||
out.extend(tensor2pil(image[i]))
|
||||
return out
|
||||
|
||||
return [
|
||||
Image.fromarray(
|
||||
np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(
|
||||
np.uint8
|
||||
)
|
||||
)
|
||||
]
|
||||
def to_numpy(image: torch.Tensor) -> npt.NDArray[np.uint8]:
|
||||
"""Converts a tensor to a ndarray with proper scaling and type conversion."""
|
||||
log.debug(f"Converting tensor to numpy array with shape {image.shape}")
|
||||
np_array = np.clip(255.0 * image.cpu().numpy(), 0, 255).astype(np.uint8)
|
||||
log.debug(f"Numpy array shape after conversion: {np_array.shape}")
|
||||
return np_array
|
||||
|
||||
|
||||
def pil2tensor(image: Image.Image | list[Image.Image]) -> torch.Tensor:
|
||||
if isinstance(image, list):
|
||||
return torch.cat([pil2tensor(img) for img in image], dim=0)
|
||||
|
||||
return torch.from_numpy(
|
||||
np.array(image).astype(np.float32) / 255.0
|
||||
).unsqueeze(0)
|
||||
def handle_batch(
|
||||
tensor: torch.Tensor,
|
||||
func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]],
|
||||
) -> list[Image.Image] | list[npt.NDArray[np.uint8]]:
|
||||
"""Handles batch processing for a given tensor and conversion function."""
|
||||
return [func(tensor[i]) for i in range(tensor.shape[0])]
|
||||
|
||||
|
||||
def np2tensor(img_np: np.ndarray | list[np.ndarray]) -> torch.Tensor:
|
||||
if isinstance(img_np, list):
|
||||
return torch.cat([np2tensor(img) for img in img_np], dim=0)
|
||||
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
||||
"""Converts a batch of tensors to a list of PIL Images."""
|
||||
|
||||
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
def single_tensor2pil(t: torch.Tensor) -> Image.Image:
|
||||
np_array = to_numpy(t)
|
||||
if np_array.ndim == 2: # (H, W) for masks
|
||||
return Image.fromarray(np_array, mode="L")
|
||||
elif np_array.ndim == 3: # (H, W, C) for RGB/RGBA
|
||||
if np_array.shape[2] == 3:
|
||||
return Image.fromarray(np_array, mode="RGB")
|
||||
elif np_array.shape[2] == 4:
|
||||
return Image.fromarray(np_array, mode="RGBA")
|
||||
raise ValueError(f"Invalid tensor shape: {t.shape}")
|
||||
|
||||
return handle_batch(tensor, single_tensor2pil)
|
||||
|
||||
|
||||
def tensor2np(tensor: torch.Tensor) -> list[np.ndarray]:
|
||||
batch_count = tensor.size(0) if len(tensor.shape) > 3 else 1
|
||||
if batch_count > 1:
|
||||
out = []
|
||||
for i in range(batch_count):
|
||||
out.extend(tensor2np(tensor[i]))
|
||||
return out
|
||||
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
|
||||
"""Converts a PIL Image or a list of PIL Images to a tensor."""
|
||||
|
||||
return [
|
||||
np.clip(255.0 * tensor.cpu().numpy().squeeze(), 0, 255).astype(
|
||||
np.uint8
|
||||
)
|
||||
]
|
||||
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
|
||||
np_image = np.array(image).astype(np.float32) / 255.0
|
||||
if np_image.ndim == 2: # Grayscale
|
||||
return torch.from_numpy(np_image).unsqueeze(0) # (1, H, W)
|
||||
else: # RGB or RGBA
|
||||
return torch.from_numpy(np_image).unsqueeze(0) # (1, H, W, C)
|
||||
|
||||
if isinstance(images, Image.Image):
|
||||
return single_pil2tensor(images)
|
||||
else:
|
||||
return torch.cat([single_pil2tensor(img) for img in images], dim=0)
|
||||
|
||||
|
||||
def np2tensor(
|
||||
np_array: npt.NDArray[np.float32] | Sequence[npt.NDArray[np.float32]],
|
||||
) -> torch.Tensor:
|
||||
"""Converts a NumPy array or a list of NumPy arrays to a tensor."""
|
||||
|
||||
def single_np2tensor(array: npt.NDArray[np.float32]) -> torch.Tensor:
|
||||
if array.ndim == 2: # (H, W) for masks
|
||||
return torch.from_numpy(
|
||||
array.astype(np.float32) / 255.0
|
||||
).unsqueeze(0) # (1, H, W)
|
||||
elif array.ndim == 3: # (H, W, C) for RGB/RGBA
|
||||
return torch.from_numpy(
|
||||
array.astype(np.float32) / 255.0
|
||||
).unsqueeze(0) # (1, H, W, C)
|
||||
raise ValueError(f"Invalid array shape: {array.shape}")
|
||||
|
||||
if isinstance(np_array, np.ndarray):
|
||||
return single_np2tensor(np_array)
|
||||
else:
|
||||
return torch.cat([single_np2tensor(arr) for arr in np_array], dim=0)
|
||||
|
||||
|
||||
def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
|
||||
"""Converts a batch of tensors to a list of NumPy arrays."""
|
||||
|
||||
def single_tensor2np(t: torch.Tensor) -> npt.NDArray[np.uint8]:
|
||||
t = t.squeeze() # Remove any singleton dimensions
|
||||
if t.ndim == 2: # (H, W) for masks
|
||||
return to_numpy(t)
|
||||
elif t.ndim == 3: # (C, H, W) for RGB/RGBA
|
||||
if t.shape[0] in [1, 3, 4]: # Channel-first format
|
||||
t = t.permute(1, 2, 0)
|
||||
return to_numpy(t)
|
||||
else:
|
||||
raise ValueError(f"Invalid tensor shape: {t.shape}")
|
||||
|
||||
return handle_batch(tensor, single_tensor2np)
|
||||
|
||||
|
||||
def nextAvailable(path: Path | str) -> Path:
|
||||
"""
|
||||
Find the next available path by adding a numbered suffix. (mimics comfy's version).
|
||||
|
||||
Args:
|
||||
path (Path): The original path to check
|
||||
|
||||
Returns
|
||||
-------
|
||||
Path: A path that doesn't exist yet
|
||||
"""
|
||||
path = Path(path)
|
||||
|
||||
if not path.is_absolute():
|
||||
path = output_dir / path
|
||||
|
||||
if not path.exists():
|
||||
return path
|
||||
|
||||
stem = path.stem
|
||||
suffix = path.suffix
|
||||
parent = path.parent
|
||||
|
||||
counter = 1
|
||||
while True:
|
||||
new_path = parent / f"{stem}_{counter:04d}{suffix}"
|
||||
if not new_path.exists():
|
||||
return new_path
|
||||
counter += 1
|
||||
|
||||
|
||||
def pad(img, left, right, top, bottom):
|
||||
@@ -746,6 +860,62 @@ def tiles_split(img, tile_size, stride_size):
|
||||
|
||||
|
||||
# region MODEL Utilities
|
||||
|
||||
|
||||
def download_model(model_url: str, destination: str):
|
||||
if isinstance(model_url, list):
|
||||
for url in model_url:
|
||||
download_model(url, destination)
|
||||
return
|
||||
|
||||
filename = Path(urlparse(model_url).path).name
|
||||
|
||||
if "drive.google.com" in model_url:
|
||||
try:
|
||||
import gdown
|
||||
except ImportError:
|
||||
log.info("Installing gdown")
|
||||
subprocess.check_call(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"pip",
|
||||
"install",
|
||||
"gdown",
|
||||
]
|
||||
)
|
||||
import gdown
|
||||
|
||||
if "/folders/" in model_url:
|
||||
# download folder
|
||||
try:
|
||||
gdown.download_folder(
|
||||
model_url, output=destination, resume=True
|
||||
)
|
||||
except TypeError:
|
||||
gdown.download_folder(model_url, output=destination)
|
||||
|
||||
return
|
||||
# download from google drive
|
||||
gdown.download(model_url, destination, quiet=False, resume=True)
|
||||
return True
|
||||
response = requests.get(model_url, stream=True)
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
|
||||
destination_path = get_model_path(destination, filename)
|
||||
destination_path.parent.mkdir(exist_ok=True)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(total_size)
|
||||
with open(destination_path, "wb") as file:
|
||||
for data in response.iter_content(chunk_size=4096):
|
||||
file.write(data)
|
||||
pbar.update(len(data))
|
||||
|
||||
log.info(
|
||||
f"Downloaded model from {model_url} to {destination_path}",
|
||||
)
|
||||
|
||||
|
||||
def download_antelopev2():
|
||||
antelopev2_url = (
|
||||
"https://drive.google.com/uc?id=18wEUfMNohBJ4K3Ly5wpTejPfDzp-8fI8"
|
||||
|
||||
+288
-147
@@ -1,16 +1,16 @@
|
||||
/**
|
||||
* @module Shared utilities
|
||||
* File: comfy_shared.js
|
||||
* Project: comfy_mtb
|
||||
* Author: Mel Massadian
|
||||
*
|
||||
* Copyright (c) 2023-2024 Mel Massadian
|
||||
*
|
||||
*/
|
||||
|
||||
// Reference the shared typedefs file
|
||||
/// <reference path="../types/typedefs.js" />
|
||||
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { api } from '../../scripts/api.js'
|
||||
|
||||
// #region base utils
|
||||
|
||||
@@ -18,7 +18,7 @@ import { app } from '../../scripts/app.js'
|
||||
export function makeUUID() {
|
||||
let dt = new Date().getTime()
|
||||
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
|
||||
const r = (dt + Math.random() * 16) % 16 | 0
|
||||
const r = ((dt + Math.random() * 16) % 16) | 0
|
||||
dt = Math.floor(dt / 16)
|
||||
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
|
||||
})
|
||||
@@ -260,6 +260,16 @@ export function inner_value_change(widget, val, event = undefined) {
|
||||
}
|
||||
}
|
||||
|
||||
export const getNamedWidget = (node, ...names) => {
|
||||
const out = {}
|
||||
|
||||
for (const name of names) {
|
||||
out[name] = node.widgets.find((w) => w.name === name)
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
/**
|
||||
* @param {LGraphNode} node
|
||||
* @param {LLink} link
|
||||
@@ -358,24 +368,40 @@ export function getWidgetType(config) {
|
||||
|
||||
// #region dynamic connections
|
||||
/**
|
||||
* @param {NodeType} nodeType
|
||||
* @param {str} prefix
|
||||
* @param {str | [str]} inputType
|
||||
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts
|
||||
* @param {NodeType} nodeType The nodetype to attach the documentation to
|
||||
* @param {str} prefix A prefix added to each dynamic inputs
|
||||
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
|
||||
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||
* @returns
|
||||
*/
|
||||
export const setupDynamicConnections = (
|
||||
nodeType,
|
||||
prefix,
|
||||
inputType,
|
||||
opts = undefined,
|
||||
) => {
|
||||
infoLogger(
|
||||
'Setting up dynamic connections for',
|
||||
Object.getOwnPropertyDescriptors(nodeType).title.value,
|
||||
)
|
||||
|
||||
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
|
||||
infoLogger('Setting up dynamic connections for', nodeType)
|
||||
|
||||
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
|
||||
const options = opts || {}
|
||||
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||
const options = Object.assign(
|
||||
{
|
||||
separator: '_',
|
||||
start_index: 1,
|
||||
},
|
||||
opts || {},
|
||||
)
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
const inputList = typeof inputType === 'object'
|
||||
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
|
||||
this.addInput(
|
||||
`${prefix}${options.separator}${options.start_index}`,
|
||||
inputList ? '*' : inputType,
|
||||
)
|
||||
return r
|
||||
}
|
||||
|
||||
@@ -410,7 +436,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
|
||||
this,
|
||||
slotIndex,
|
||||
isConnected,
|
||||
`${prefix}_`,
|
||||
`${prefix}${options.separator}`,
|
||||
inputType,
|
||||
options,
|
||||
)
|
||||
@@ -426,7 +452,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
|
||||
* @param {bool} connected - Was this event connecting or disconnecting
|
||||
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
|
||||
* @param {string|[string]} [connectionType] - The type of the dynamic connection
|
||||
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
|
||||
* @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
|
||||
*/
|
||||
export const dynamic_connection = (
|
||||
node,
|
||||
@@ -436,13 +462,18 @@ export const dynamic_connection = (
|
||||
connectionType = '*',
|
||||
opts = undefined,
|
||||
) => {
|
||||
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
|
||||
const options = opts || {}
|
||||
/* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
|
||||
const options = Object.assign(
|
||||
{
|
||||
start_index: 1,
|
||||
},
|
||||
opts || {},
|
||||
)
|
||||
|
||||
if (
|
||||
node.inputs.length > 0 &&
|
||||
!node.inputs[index].name.startsWith(connectionPrefix)
|
||||
) {
|
||||
// function to test if input is a dynamic one
|
||||
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
|
||||
|
||||
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -461,7 +492,7 @@ export const dynamic_connection = (
|
||||
const to_remove = []
|
||||
for (let n = 1; n < node.inputs.length; n++) {
|
||||
const element = node.inputs[n]
|
||||
if (!element.link) {
|
||||
if (!element.link && isDynamicInput(element.name)) {
|
||||
if (node.widgets) {
|
||||
const w = node.widgets.find((w) => w.name === element.name)
|
||||
if (w) {
|
||||
@@ -487,14 +518,25 @@ export const dynamic_connection = (
|
||||
|
||||
infoLogger('Cleaning inputs: making it sequential again')
|
||||
// make inputs sequential again
|
||||
let prefixed_idx = options.start_index
|
||||
for (let i = 0; i < node.inputs.length; i++) {
|
||||
let name = `${connectionPrefix}${i + 1}`
|
||||
let name = ''
|
||||
// rename only prefixed inputs
|
||||
if (isDynamicInput(node.inputs[i].name)) {
|
||||
// prefixed => rename and increase index
|
||||
name = `${connectionPrefix}${prefixed_idx}`
|
||||
prefixed_idx += 1
|
||||
} else {
|
||||
// not prefixed => keep same name
|
||||
name = node.inputs[i].name
|
||||
}
|
||||
|
||||
if (nameArray.length > 0) {
|
||||
name = i < nameArray.length ? nameArray[i] : name
|
||||
}
|
||||
|
||||
node.inputs[i].label = name
|
||||
// preserve label if it exists
|
||||
node.inputs[i].label = node.inputs[i].label || name
|
||||
node.inputs[i].name = name
|
||||
}
|
||||
}
|
||||
@@ -534,11 +576,16 @@ export const dynamic_connection = (
|
||||
if (node.inputs.length === 0) return
|
||||
// add an extra input
|
||||
if (node.inputs[node.inputs.length - 1].link !== null) {
|
||||
const nextIndex = node.inputs.length
|
||||
// count only the prefixed inputs
|
||||
const nextIndex = node.inputs.reduce(
|
||||
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
|
||||
0,
|
||||
)
|
||||
|
||||
const name =
|
||||
nextIndex < nameArray.length
|
||||
? nameArray[nextIndex]
|
||||
: `${connectionPrefix}${nextIndex + 1}`
|
||||
: `${connectionPrefix}${nextIndex + options.start_index}`
|
||||
|
||||
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
|
||||
node.addInput(name, conType)
|
||||
@@ -628,39 +675,6 @@ export const loadScript = (
|
||||
})
|
||||
}
|
||||
|
||||
export function defineClass(className, classStyles) {
|
||||
const styleSheets = document.styleSheets
|
||||
|
||||
// Helper function to check if the class exists in a style sheet
|
||||
function classExistsInStyleSheet(styleSheet) {
|
||||
const rules = styleSheet.rules || styleSheet.cssRules
|
||||
for (const rule of rules) {
|
||||
if (rule.selectorText === `.${className}`) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Check if the class is already defined in any of the style sheets
|
||||
let classExists = false
|
||||
for (const styleSheet of styleSheets) {
|
||||
if (classExistsInStyleSheet(styleSheet)) {
|
||||
classExists = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// If the class doesn't exist, add the new class definition to the first style sheet
|
||||
if (!classExists) {
|
||||
if (styleSheets[0].insertRule) {
|
||||
styleSheets[0].insertRule(`.${className} { ${classStyles} }`, 0)
|
||||
} else if (styleSheets[0].addRule) {
|
||||
styleSheets[0].addRule(`.${className}`, classStyles, 0)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// #endregion
|
||||
|
||||
// #region documentation widget
|
||||
@@ -704,7 +718,7 @@ const create_documentation_stylesheet = () => {
|
||||
border-radius: 6px;
|
||||
border: 3px solid var(--bg-color);
|
||||
}
|
||||
|
||||
|
||||
/* Scrollbar styling for Firefox */
|
||||
scrollbar-width: thin;
|
||||
scrollbar-color: var(--fg-color) var(--bg-color);
|
||||
@@ -726,7 +740,7 @@ const create_documentation_stylesheet = () => {
|
||||
border-collapse: collapse;
|
||||
border: 1px var(--border-color) solid;
|
||||
}
|
||||
.documentation-popup th,
|
||||
.documentation-popup th,
|
||||
.documentation-popup td {
|
||||
border: 1px var(--border-color) solid;
|
||||
}
|
||||
@@ -736,10 +750,84 @@ const create_documentation_stylesheet = () => {
|
||||
document.head.appendChild(styleTag)
|
||||
}
|
||||
}
|
||||
let documentationConverter
|
||||
let parserPromise
|
||||
const callbackQueue = []
|
||||
|
||||
function runQueuedCallbacks() {
|
||||
while (callbackQueue.length) {
|
||||
const cb = callbackQueue.shift()
|
||||
cb(window.MTB.mdParser)
|
||||
}
|
||||
}
|
||||
|
||||
function loadParser(shiki) {
|
||||
if (!parserPromise) {
|
||||
parserPromise = import(
|
||||
shiki
|
||||
? '/mtb_async/mtb_markdown_plus.umd.js'
|
||||
: '/mtb_async/mtb_markdown.umd.js'
|
||||
)
|
||||
.then((_module) =>
|
||||
shiki ? MTBMarkdownPlus.getParser() : MTBMarkdown.getParser(),
|
||||
)
|
||||
.then((instance) => {
|
||||
window.MTB.mdParser = instance
|
||||
runQueuedCallbacks()
|
||||
return instance
|
||||
})
|
||||
.catch((error) => {
|
||||
console.error('Error loading the parser:', error)
|
||||
})
|
||||
}
|
||||
return parserPromise
|
||||
}
|
||||
|
||||
export const ensureMarkdownParser = async (callback) => {
|
||||
infoLogger('Ensuring md parser')
|
||||
let use_shiki = false
|
||||
try {
|
||||
use_shiki = await api.getSetting('mtb.Use Shiki')
|
||||
} catch (e) {
|
||||
console.warn('Option not available yet', e)
|
||||
}
|
||||
|
||||
if (window.MTB?.mdParser) {
|
||||
infoLogger('Markdown parser found')
|
||||
callback?.(window.MTB.mdParser)
|
||||
return window.MTB.mdParser
|
||||
}
|
||||
|
||||
if (!parserPromise) {
|
||||
infoLogger('Running promise to fetch parser')
|
||||
|
||||
try {
|
||||
loadParser(use_shiki) //.then(() => {
|
||||
// callback?.(window.MTB.mdParser)
|
||||
// })
|
||||
} catch (error) {
|
||||
console.error('Error loading the parser:', error)
|
||||
}
|
||||
} else {
|
||||
infoLogger('A similar promise is already running, waiting for it to finish')
|
||||
}
|
||||
if (callback) {
|
||||
callbackQueue.push(callback)
|
||||
}
|
||||
|
||||
await parserPromise
|
||||
await parserPromise
|
||||
|
||||
return window.MTB.mdParser
|
||||
}
|
||||
|
||||
/**
|
||||
* Add documentation widget to the selected node
|
||||
* Add documentation widget to the given node.
|
||||
*
|
||||
* This method will add a `docCtrl` property to the node
|
||||
* that contains the AbortController that manages all the events
|
||||
* defined inside it (global and instance ones) without explicit
|
||||
* cleanup method for each.
|
||||
*
|
||||
* @param {NodeData} nodeData
|
||||
* @param {NodeType} nodeType
|
||||
* @param {DocumentationOptions} opts
|
||||
@@ -756,25 +844,10 @@ export const addDocumentation = (
|
||||
return
|
||||
}
|
||||
|
||||
if (!documentationConverter) {
|
||||
infoLogger('Initializing our mardown converter')
|
||||
documentationConverter = new showdown.Converter({
|
||||
tables: true,
|
||||
strikethrough: true,
|
||||
emoji: true,
|
||||
ghCodeBlocks: true,
|
||||
tasklists: true,
|
||||
ghMentions: true,
|
||||
smoothLivePreview: true,
|
||||
simplifiedAutoLink: true,
|
||||
parseImgDimensions: true,
|
||||
openLinksInNewWindow: true,
|
||||
})
|
||||
}
|
||||
|
||||
const options = opts || {}
|
||||
const iconSize = options.icon_size || 14
|
||||
const iconMargin = options.icon_margin || 4
|
||||
|
||||
let docElement = null
|
||||
let wrapper = null
|
||||
|
||||
@@ -820,80 +893,87 @@ export const addDocumentation = (
|
||||
|
||||
wrapper = document.createElement('div')
|
||||
wrapper.classList.add('documentation-wrapper')
|
||||
wrapper.innerHTML = documentationConverter.makeHtml(nodeData.description)
|
||||
docElement.appendChild(wrapper)
|
||||
|
||||
// resize handle
|
||||
resizeHandle = document.createElement('div')
|
||||
resizeHandle.style.width = '0'
|
||||
resizeHandle.style.height = '0'
|
||||
resizeHandle.style.position = 'absolute'
|
||||
resizeHandle.style.bottom = '0'
|
||||
resizeHandle.style.right = '0'
|
||||
// wrapper.innerHTML = documentationConverter.makeHtml(nodeData.description)
|
||||
|
||||
resizeHandle.style.cursor = 'se-resize'
|
||||
resizeHandle.style.userSelect = 'none'
|
||||
ensureMarkdownParser().then(() => {
|
||||
MTB.mdParser.parse(nodeData.description).then((e) => {
|
||||
wrapper.innerHTML = e
|
||||
// resize handle
|
||||
resizeHandle = document.createElement('div')
|
||||
resizeHandle.classList.add('doc-resize-handle')
|
||||
resizeHandle.style.width = '0'
|
||||
resizeHandle.style.height = '0'
|
||||
resizeHandle.style.position = 'absolute'
|
||||
resizeHandle.style.bottom = '0'
|
||||
resizeHandle.style.right = '0'
|
||||
|
||||
resizeHandle.style.borderWidth = '15px'
|
||||
resizeHandle.style.borderStyle = 'solid'
|
||||
resizeHandle.style.cursor = 'se-resize'
|
||||
resizeHandle.style.userSelect = 'none'
|
||||
|
||||
resizeHandle.style.borderColor =
|
||||
'transparent var(--border-color) var(--border-color) transparent'
|
||||
resizeHandle.style.borderWidth = '15px'
|
||||
resizeHandle.style.borderStyle = 'solid'
|
||||
|
||||
wrapper.appendChild(resizeHandle)
|
||||
let isResizing = false
|
||||
resizeHandle.style.borderColor =
|
||||
'transparent var(--border-color) var(--border-color) transparent'
|
||||
|
||||
let startX
|
||||
let startY
|
||||
let startWidth
|
||||
let startHeight
|
||||
wrapper.appendChild(resizeHandle)
|
||||
let isResizing = false
|
||||
|
||||
resizeHandle.addEventListener(
|
||||
'mousedown',
|
||||
(e) => {
|
||||
e.stopPropagation()
|
||||
isResizing = true
|
||||
startX = e.clientX
|
||||
startY = e.clientY
|
||||
startWidth = Number.parseInt(
|
||||
document.defaultView.getComputedStyle(docElement).width,
|
||||
10,
|
||||
let startX
|
||||
let startY
|
||||
let startWidth
|
||||
let startHeight
|
||||
|
||||
resizeHandle.addEventListener(
|
||||
'mousedown',
|
||||
(e) => {
|
||||
e.stopPropagation()
|
||||
isResizing = true
|
||||
startX = e.clientX
|
||||
startY = e.clientY
|
||||
startWidth = Number.parseInt(
|
||||
document.defaultView.getComputedStyle(docElement).width,
|
||||
10,
|
||||
)
|
||||
startHeight = Number.parseInt(
|
||||
document.defaultView.getComputedStyle(docElement).height,
|
||||
10,
|
||||
)
|
||||
},
|
||||
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
startHeight = Number.parseInt(
|
||||
document.defaultView.getComputedStyle(docElement).height,
|
||||
10,
|
||||
|
||||
document.addEventListener(
|
||||
'mousemove',
|
||||
(e) => {
|
||||
if (!isResizing) return
|
||||
const scale = app.canvas.ds.scale
|
||||
const newWidth = startWidth + (e.clientX - startX) / scale
|
||||
const newHeight = startHeight + (e.clientY - startY) / scale
|
||||
|
||||
docElement.style.width = `${newWidth}px`
|
||||
docElement.style.height = `${newHeight}px`
|
||||
|
||||
this.docPos = {
|
||||
width: `${newWidth}px`,
|
||||
height: `${newHeight}px`,
|
||||
}
|
||||
},
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
},
|
||||
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
|
||||
document.addEventListener(
|
||||
'mousemove',
|
||||
(e) => {
|
||||
if (!isResizing) return
|
||||
const scale = app.canvas.ds.scale
|
||||
const newWidth = startWidth + (e.clientX - startX) / scale
|
||||
const newHeight = startHeight + (e.clientY - startY) / scale
|
||||
|
||||
docElement.style.width = `${newWidth}px`
|
||||
docElement.style.height = `${newHeight}px`
|
||||
|
||||
this.docPos = {
|
||||
width: `${newWidth}px`,
|
||||
height: `${newHeight}px`,
|
||||
}
|
||||
},
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
|
||||
document.addEventListener(
|
||||
'mouseup',
|
||||
() => {
|
||||
isResizing = false
|
||||
},
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
document.addEventListener(
|
||||
'mouseup',
|
||||
() => {
|
||||
isResizing = false
|
||||
},
|
||||
{ signal: this.docCtrl.signal },
|
||||
)
|
||||
})
|
||||
})
|
||||
} else if (!this.show_doc && docElement !== null) {
|
||||
docElement.remove()
|
||||
docElement = null
|
||||
@@ -917,8 +997,8 @@ export const addDocumentation = (
|
||||
Object.assign(docElement.style, {
|
||||
transformOrigin: '0 0',
|
||||
transform: scale,
|
||||
left: `${transform.a + transform.e}px`,
|
||||
top: `${transform.d + transform.f}px`,
|
||||
left: `${transform.a + rect.x + transform.e}px`,
|
||||
top: `${transform.d + rect.y + transform.f}px`,
|
||||
width: this.docPos ? this.docPos.width : `${this.size[0] * 1.5}px`,
|
||||
height: this.docPos?.height,
|
||||
})
|
||||
@@ -1025,8 +1105,8 @@ export function addMenuHandler(nodeType, cb) {
|
||||
*/
|
||||
nodeType.prototype.getExtraMenuOptions = function (app, options) {
|
||||
const r = getOpts.apply(this, [app, options]) || []
|
||||
const newItems = cb.apply(this, [app, options])
|
||||
return r + newItems
|
||||
const newItems = cb.apply(this, [app, options]) || []
|
||||
return [...r, ...newItems]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1049,7 +1129,67 @@ export const addDeprecation = (nodeType, reason) => {
|
||||
|
||||
// #endregion
|
||||
|
||||
// #region graph utilities
|
||||
// #region Actions API
|
||||
export const runAction = async (name, ...args) => {
|
||||
const req = await api.fetchApi('/mtb/actions', {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({
|
||||
name,
|
||||
args,
|
||||
}),
|
||||
})
|
||||
|
||||
const res = await req.json()
|
||||
return res.result
|
||||
}
|
||||
export const getServerInfo = async () => {
|
||||
const res = await api.fetchApi('/mtb/server-info')
|
||||
return await res.json()
|
||||
}
|
||||
export const setServerInfo = async (opts) => {
|
||||
await api.fetchApi('/mtb/server-info', {
|
||||
method: 'POST',
|
||||
body: JSON.stringify(opts),
|
||||
})
|
||||
}
|
||||
|
||||
// #endregion
|
||||
|
||||
// #region Authoring API / graph utilities
|
||||
export const getAPIInputs = () => {
|
||||
const inputs = {}
|
||||
let counter = 1
|
||||
for (const node of getNodes(true)) {
|
||||
const widgets = node.widgets
|
||||
|
||||
if (node.properties.mtb_api && node.properties.useAPI) {
|
||||
if (node.properties.mtb_api.inputs) {
|
||||
for (const currentName in node.properties.mtb_api.inputs) {
|
||||
const current = node.properties.mtb_api.inputs[currentName]
|
||||
if (current.enabled) {
|
||||
const inputName = current.name || currentName
|
||||
const widget = widgets.find((w) => w.name === currentName)
|
||||
if (!widget) continue
|
||||
if (!(inputName in inputs)) {
|
||||
inputs[inputName] = {
|
||||
...current,
|
||||
id: counter,
|
||||
name: inputName,
|
||||
type: current.type,
|
||||
node_id: node.id,
|
||||
widgets: [],
|
||||
}
|
||||
}
|
||||
inputs[inputName].widgets.push(widget)
|
||||
counter = counter + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return inputs
|
||||
}
|
||||
|
||||
export const getNodes = (skip_unused) => {
|
||||
const nodes = []
|
||||
for (const outerNode of app.graph.computeExecutionOrder(false)) {
|
||||
@@ -1068,3 +1208,4 @@ export const getNodes = (skip_unused) => {
|
||||
}
|
||||
return nodes
|
||||
}
|
||||
// #endregion
|
||||
|
||||
+116
-42
@@ -11,11 +11,8 @@
|
||||
/// <reference path="../types/typedefs.js" />
|
||||
|
||||
import { app } from '../../scripts/app.js'
|
||||
|
||||
import * as shared from './comfy_shared.js'
|
||||
import { MtbWidgets } from './mtb_widgets.js'
|
||||
|
||||
// TODO: respect inputs order...
|
||||
import * as mtb_ui from './mtb_ui.js'
|
||||
|
||||
function escapeHtml(unsafe) {
|
||||
return unsafe
|
||||
@@ -25,6 +22,54 @@ function escapeHtml(unsafe) {
|
||||
.replace(/"/g, '"')
|
||||
.replace(/'/g, ''')
|
||||
}
|
||||
|
||||
function createDebugSection(title) {
|
||||
const section = mtb_ui.makeElement('div', {
|
||||
margin: '8px 0',
|
||||
padding: '8px',
|
||||
borderRadius: '4px',
|
||||
backgroundColor: 'rgba(0,0,0,0.2)'
|
||||
})
|
||||
|
||||
const header = mtb_ui.makeElement('h3', {
|
||||
margin: '0 0 8px 0',
|
||||
padding: '4px 0',
|
||||
borderBottom: '1px solid rgba(255,255,255,0.1)',
|
||||
fontSize: '14px',
|
||||
fontWeight: 'bold',
|
||||
color: '#9f9'
|
||||
})
|
||||
header.textContent = title
|
||||
section.appendChild(header)
|
||||
|
||||
return section
|
||||
}
|
||||
|
||||
function createDebugContent(content, type) {
|
||||
const wrapper = mtb_ui.makeElement('div', {
|
||||
margin: '4px 0'
|
||||
})
|
||||
|
||||
if (type === 'text') {
|
||||
const text = mtb_ui.makeElement('p', {
|
||||
margin: '2px 0',
|
||||
fontFamily: 'monospace',
|
||||
whiteSpace: 'pre-wrap'
|
||||
})
|
||||
text.innerHTML = content
|
||||
wrapper.appendChild(text)
|
||||
} else if (type === 'image') {
|
||||
const img = mtb_ui.makeElement('img', {
|
||||
width: '100%',
|
||||
borderRadius: '2px'
|
||||
})
|
||||
img.src = content
|
||||
wrapper.appendChild(img)
|
||||
}
|
||||
|
||||
return wrapper
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'mtb.Debug',
|
||||
|
||||
@@ -36,12 +81,10 @@ app.registerExtension({
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData.name === 'Debug (mtb)') {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
this.options = {}
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
this.addInput(`anything_1`, '*')
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
||||
this.addInput('anything_1', '*')
|
||||
return r
|
||||
}
|
||||
|
||||
@@ -81,51 +124,82 @@ app.registerExtension({
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (data) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
|
||||
const prefix = 'anything_'
|
||||
nodeType.prototype.onExecuted = function (...args) {
|
||||
onExecuted?.apply(this, args)
|
||||
const [data, ..._rest] = args
|
||||
|
||||
if (this.widgets) {
|
||||
let tgt_len = this.widgets.length
|
||||
for (let i = 0; i < this.widgets.length; i++) {
|
||||
if (this.widgets[i].name !== 'output_to_console') {
|
||||
if (
|
||||
this.widgets[i].name !== 'output_to_console' &&
|
||||
this.widgets[i].name !== 'as_detailed_types'
|
||||
) {
|
||||
this.widgets[i].onRemove?.()
|
||||
this.widgets[i].onRemoved?.()
|
||||
tgt_len -= 1
|
||||
}
|
||||
}
|
||||
this.widgets.length = 1
|
||||
}
|
||||
let widgetI = 1
|
||||
// console.log(message)
|
||||
if (data.text) {
|
||||
for (const txt of data.text) {
|
||||
const w = this.addCustomWidget(
|
||||
MtbWidgets.DEBUG_STRING(`${prefix}_${widgetI}`, escapeHtml(txt)),
|
||||
)
|
||||
w.parent = this
|
||||
widgetI++
|
||||
}
|
||||
}
|
||||
if (data.b64_images) {
|
||||
for (const img of data.b64_images) {
|
||||
const w = this.addCustomWidget(
|
||||
MtbWidgets.DEBUG_IMG(`${prefix}_${widgetI}`, img),
|
||||
)
|
||||
w.parent = this
|
||||
widgetI++
|
||||
}
|
||||
this.widgets.length = tgt_len
|
||||
}
|
||||
|
||||
// this.setSize(this.computeSize())
|
||||
const inputData = {}
|
||||
|
||||
const uiData = data.ui || data
|
||||
|
||||
if (uiData.items) {
|
||||
uiData.items.forEach(item => {
|
||||
const inputName = item.input
|
||||
if (!inputData[inputName]) {
|
||||
inputData[inputName] = { text: [], b64_images: [] }
|
||||
}
|
||||
if (item.text) {
|
||||
inputData[inputName].text.push(...item.text)
|
||||
}
|
||||
if (item.b64_images) {
|
||||
inputData[inputName].b64_images.push(...item.b64_images)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
let widgetI = 1
|
||||
for (const [inputName, content] of Object.entries(inputData)) {
|
||||
if (content.text.length === 0 && content.b64_images.length === 0) {
|
||||
continue
|
||||
}
|
||||
|
||||
const section = createDebugSection(inputName)
|
||||
|
||||
if (content.text.length > 0) {
|
||||
content.text.forEach(text => {
|
||||
section.appendChild(createDebugContent(text, 'text'))
|
||||
})
|
||||
}
|
||||
|
||||
if (content.b64_images.length > 0) {
|
||||
content.b64_images.forEach(img => {
|
||||
section.appendChild(createDebugContent(img, 'image'))
|
||||
})
|
||||
}
|
||||
|
||||
this.addDOMWidget(
|
||||
`debug_section_${widgetI}`,
|
||||
'CUSTOM',
|
||||
section,
|
||||
{}
|
||||
)
|
||||
widgetI++
|
||||
}
|
||||
|
||||
this.onRemoved = function () {
|
||||
// When removing this node we need to remove the input from the DOM
|
||||
for (let y in this.widgets) {
|
||||
if (this.widgets[y].canvas) {
|
||||
this.widgets[y].canvas.remove()
|
||||
for (const widget of this.widgets) {
|
||||
if (widget.canvas) {
|
||||
widget.canvas.remove()
|
||||
}
|
||||
shared.cleanupNode(this)
|
||||
this.widgets[y].onRemoved?.()
|
||||
widget.onRemoved?.()
|
||||
widget.onRemove?.()
|
||||
}
|
||||
shared.cleanupNode(this)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Vendored
+3
-3
File diff suppressed because one or more lines are too long
Vendored
-3
File diff suppressed because one or more lines are too long
+296
-295
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { LocalStorageManager } from './comfy_shared.js'
|
||||
const styles = {
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
}
|
||||
|
||||
let currentImageIndex = 0
|
||||
@@ -58,298 +58,299 @@ const storage = new LocalStorageManager('mtb')
|
||||
let activated = storage.get('image_feed', false)
|
||||
|
||||
app.registerExtension({
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.imageFeed.enabled',
|
||||
name: '[⚡mtb] Enable image feed',
|
||||
type: 'boolean',
|
||||
defaultValue: true,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
console.warn(
|
||||
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
||||
)
|
||||
activated = false // just in case other methods are added later on
|
||||
return
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.image-feed-enabled',
|
||||
category: ['mtb', 'Main', 'image-feed-enabled'],
|
||||
name: 'Enable Image Feed',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
console.warn(
|
||||
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
||||
)
|
||||
activated = false // just in case other methods are added later on
|
||||
return
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* This is the function that creates the image buttons for the image list
|
||||
* They are wrapped in a button so that they can be clicked and open
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* This is the function that creates the image buttons for the image list
|
||||
* They are wrapped in a button so that they can be clicked and open
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
|
||||
imageUrls.push(img.src)
|
||||
imageUrls.push(img.src)
|
||||
|
||||
console.debug(img.src)
|
||||
console.debug(img.src)
|
||||
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
///////-------
|
||||
///////-------
|
||||
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
})
|
||||
|
||||
@@ -0,0 +1,415 @@
|
||||
/// <reference path="../types/typedefs.js" />
|
||||
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { api } from '../../scripts/api.js'
|
||||
|
||||
import * as shared from './comfy_shared.js'
|
||||
|
||||
import {
|
||||
// defineCSSClass,
|
||||
ensureMTBStyles,
|
||||
makeElement,
|
||||
makeSelect,
|
||||
makeSlider,
|
||||
renderSidebar,
|
||||
} from './mtb_ui.js'
|
||||
|
||||
const offset = 0
|
||||
let currentWidth = 200
|
||||
let currentMode = 'input'
|
||||
let subfolder = ''
|
||||
let currentSort = 'None'
|
||||
|
||||
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
|
||||
const VIDEO_NODES = ['VHS_LoadVideo']
|
||||
const PROCESSED_PROMPT_IDS = new Set()
|
||||
|
||||
const updateImage = (node, image) => {
|
||||
if (IMAGE_NODES.includes(node.type)) {
|
||||
const w = node.widgets?.find((w) => w.name === 'image')
|
||||
if (w) {
|
||||
w.value = image
|
||||
w.callback()
|
||||
}
|
||||
} else if (VIDEO_NODES.includes(node.type)) {
|
||||
const w = node.widgets?.find((w) => w.name === 'video')
|
||||
if (w) {
|
||||
node.updateParameters({ filename: image }, true)
|
||||
}
|
||||
} else {
|
||||
console.warn('No method to update', node.type)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Converts a result item to a request url.
|
||||
* @param {ResultItem} resultItem
|
||||
* @returns {string} - The request URL.
|
||||
*/
|
||||
const resultItemToQuery = (resultItem) =>
|
||||
[
|
||||
`/mtb/view?filename=${resultItem.filename}`,
|
||||
`width=512`,
|
||||
`type=${resultItem.type}`,
|
||||
`subfolder=${resultItem.subfolder}`,
|
||||
`preview=`,
|
||||
].join('&')
|
||||
|
||||
/**
|
||||
* Retrieves the unique prompt ID from a history task item.
|
||||
* @param {HistoryTaskItem} historyTaskItem
|
||||
* @returns {string} - The prompt ID.
|
||||
*/
|
||||
const getPromptId = (historyTaskItem) => `${historyTaskItem.prompt[1]}`
|
||||
|
||||
/**
|
||||
* Process and return any new/unseen outputs from the most recent history item.
|
||||
* @param {HistoryTaskItem} mostRecentTask - The most recent history task item.
|
||||
* @returns {Object<string, string>} - A map of task outputs URLs.
|
||||
*/
|
||||
const getNewOutputUrls = (mostRecentTask) => {
|
||||
if (!mostRecentTask) return
|
||||
|
||||
const promptId = getPromptId(mostRecentTask)
|
||||
if (PROCESSED_PROMPT_IDS.has(promptId)) return
|
||||
|
||||
const urls = {}
|
||||
for (const nodeOutputs of Object.values(mostRecentTask.outputs)) {
|
||||
const { images, audio, animated } = nodeOutputs
|
||||
if (images) {
|
||||
const imageOutputs = Object.values(nodeOutputs.images)
|
||||
imageOutputs.forEach(
|
||||
(resultItem) =>
|
||||
(urls[resultItem.filename] = resultItemToQuery(resultItem))
|
||||
)
|
||||
}
|
||||
// Can process `animated` and `audio` outputs here.
|
||||
}
|
||||
|
||||
const foundNewOutputs = Object.keys(urls).length > 0
|
||||
if (!foundNewOutputs) return null
|
||||
|
||||
PROCESSED_PROMPT_IDS.add(promptId)
|
||||
return urls
|
||||
}
|
||||
|
||||
/** Fetch history and update the grid with any new ouput images. */
|
||||
const updateOutputsGrid = async () => {
|
||||
try {
|
||||
const history = await api.getHistory(/** maxSize: */ 1)
|
||||
const mostRcentTask = history.History[0]
|
||||
const newUrls = getNewOutputUrls(mostRcentTask)
|
||||
if (newUrls) {
|
||||
const imgGrid = document.querySelector('.mtb_img_grid')
|
||||
getImgsFromUrls(newUrls, imgGrid, { prepend: true })
|
||||
}
|
||||
} catch (error) {
|
||||
console.error('Error fetching history:', error)
|
||||
}
|
||||
}
|
||||
|
||||
const getImgsFromUrls = (urls, target, options = { prepend: false }) => {
|
||||
const imgs = []
|
||||
if (urls === undefined) {
|
||||
return imgs
|
||||
}
|
||||
const elem = currentMode === 'video' ? 'video' : 'img'
|
||||
|
||||
for (const [key, url] of Object.entries(urls)) {
|
||||
const a = makeElement(elem)
|
||||
a.src = url
|
||||
a.width = currentWidth
|
||||
if (currentMode === 'input') {
|
||||
a.onclick = (_e) => {
|
||||
if (subfolder !== '') {
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'warn',
|
||||
summary: 'Subfolder not supported',
|
||||
detail: "The LoadImage node doesn't support subfolders",
|
||||
life: 5000,
|
||||
})
|
||||
return
|
||||
}
|
||||
const selected = app.canvas.selected_nodes
|
||||
if (selected && Object.keys(selected).length === 0) {
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'warn',
|
||||
summary: 'No node selected!',
|
||||
detail:
|
||||
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
|
||||
life: 5000,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
|
||||
updateImage(node, key)
|
||||
}
|
||||
}
|
||||
} else if (currentMode === 'output') {
|
||||
a.onclick = (_e) => {
|
||||
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
|
||||
if (subfolder !== '') {
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'warn',
|
||||
summary: 'Subfolder not supported',
|
||||
detail: "The LoadImage node doesn't support subfolders",
|
||||
life: 5000,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'warn',
|
||||
summary: 'Outputs not supported',
|
||||
detail:
|
||||
'For now only inputs can be clicked to load the image on the active LoadImage node.',
|
||||
life: 5000,
|
||||
})
|
||||
}
|
||||
} else {
|
||||
a.autoplay = true
|
||||
|
||||
a.muted = true
|
||||
a.loop = true
|
||||
a.onclick = (_e) => {
|
||||
const selected = app.canvas.selected_nodes
|
||||
if (selected && Object.keys(selected).length === 0) {
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'warn',
|
||||
summary: 'No node selected!',
|
||||
detail:
|
||||
"For now the only action when clicking videos in the sidebar is to set the video on all selected 'Load Video (Upload)' nodes.",
|
||||
life: 5000,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
|
||||
updateImage(node, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
imgs.push(a)
|
||||
}
|
||||
if (target !== undefined) {
|
||||
if (options.prepend) target.prepend(...imgs)
|
||||
else target.append(...imgs)
|
||||
}
|
||||
return imgs
|
||||
}
|
||||
|
||||
const getModes = async () => {
|
||||
const inputs = await shared.runAction('getUserImageFolders')
|
||||
return inputs
|
||||
}
|
||||
const getUrls = async (subfolder) => {
|
||||
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
|
||||
console.log('Sidebar count', count)
|
||||
if (currentMode === 'video') {
|
||||
const output = await shared.runAction(
|
||||
'getUserVideos',
|
||||
256,
|
||||
count,
|
||||
offset,
|
||||
currentSort,
|
||||
)
|
||||
return output || {}
|
||||
}
|
||||
const output = await shared.runAction(
|
||||
'getUserImages',
|
||||
currentMode,
|
||||
count,
|
||||
offset,
|
||||
currentSort,
|
||||
false,
|
||||
subfolder,
|
||||
)
|
||||
return output || {}
|
||||
}
|
||||
|
||||
//NOTE: do not load if using the old ui
|
||||
if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
||||
// NOTE: removed this for now since I'm not actually exposing anything a client
|
||||
// cannot already access from "/view"...
|
||||
// let exposed = false
|
||||
|
||||
const sidebar_extension = {
|
||||
name: 'mtb.io-sidebar',
|
||||
// init: async () => {
|
||||
// try {
|
||||
// const res = await api.fetchApi('/mtb/server-info')
|
||||
// const msg = await res.json()
|
||||
// exposed = msg.exposed
|
||||
// } catch (e) {
|
||||
// console.error('Error:', e)
|
||||
// }
|
||||
// },
|
||||
init: () => {
|
||||
let handle
|
||||
const version = window?.__COMFYUI_FRONTEND_VERSION__
|
||||
console.log(`%c ${version}`, 'background: orange; color: white;')
|
||||
|
||||
ensureMTBStyles()
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.io-sidebar.count',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'count'],
|
||||
|
||||
name: 'Number of images to fetch',
|
||||
type: 'number',
|
||||
defaultValue: 1000,
|
||||
|
||||
tooltip:
|
||||
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
|
||||
attrs: {
|
||||
style: {
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.io-sidebar.img-size',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
|
||||
|
||||
name: 'Resolution of the images',
|
||||
type: 'number',
|
||||
defaultValue: 512,
|
||||
|
||||
tooltip: "It's recommended to keep it at 512px",
|
||||
attrs: {
|
||||
style: {
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
})
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.io-sidebar.sort',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
||||
name: 'Default sort mode',
|
||||
type: 'combo',
|
||||
|
||||
onChange: (v) => {
|
||||
// alert(`Sort is now ${v}`)
|
||||
currentSort = v
|
||||
},
|
||||
|
||||
defaultValue: 'Modified',
|
||||
// tooltip: "It's recommended to keep it at 512px",
|
||||
options: [
|
||||
'None',
|
||||
'Modified',
|
||||
'Modified-Reverse',
|
||||
'Name',
|
||||
'Name-Reverse',
|
||||
],
|
||||
})
|
||||
|
||||
app.extensionManager.registerSidebarTab({
|
||||
id: 'mtb-inputs-outputs',
|
||||
icon: 'pi pi-images',
|
||||
title: 'Input & Outputs',
|
||||
tooltip: 'MTB: Browse inputs and outputs directories.',
|
||||
type: 'custom',
|
||||
|
||||
// this is run everytime the tab's diplay is toggled on.
|
||||
render: async (el) => {
|
||||
if (handle) {
|
||||
handle.unregister()
|
||||
handle = undefined
|
||||
}
|
||||
|
||||
if (el.parentNode) {
|
||||
el.parentNode.style.overflowY = 'clip'
|
||||
}
|
||||
|
||||
const allModes = await getModes()
|
||||
const input_modes = allModes.input.map((m) => `input - ${m}`)
|
||||
const output_modes = allModes.output.map((m) => `output - ${m}`)
|
||||
const urls = await getUrls()
|
||||
let imgs = {}
|
||||
|
||||
const cont = makeElement('div.mtb_sidebar')
|
||||
|
||||
const imgGrid = makeElement('div.mtb_img_grid')
|
||||
const selector = makeSelect(
|
||||
['input', 'output', 'video', ...output_modes, ...input_modes],
|
||||
currentMode,
|
||||
)
|
||||
|
||||
selector.addEventListener('change', async (e) => {
|
||||
let newMode = e.target.value
|
||||
let changed = false
|
||||
let newSub = ''
|
||||
if (newMode !== 'input' && newMode !== 'output') {
|
||||
if (newMode.startsWith('input - ')) {
|
||||
newSub = newMode.replace('input - ', '')
|
||||
newMode = 'input'
|
||||
} else if (newMode.startsWith('output - ')) {
|
||||
newSub = newMode.replace('output - ', '')
|
||||
newMode = 'output'
|
||||
}
|
||||
}
|
||||
changed = newMode !== currentMode || newSub !== subfolder
|
||||
currentMode = newMode
|
||||
subfolder = newSub
|
||||
if (changed) {
|
||||
imgGrid.innerHTML = ''
|
||||
const urls = await getUrls(subfolder)
|
||||
if (urls) {
|
||||
imgs = getImgsFromUrls(urls, imgGrid)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const imgTools = makeElement('div.mtb_tools')
|
||||
const orderSelect = makeSelect(
|
||||
['None', 'Modified', 'Modified-Reverse', 'Name', 'Name-Reverse'],
|
||||
currentSort,
|
||||
)
|
||||
|
||||
orderSelect.addEventListener('change', async (e) => {
|
||||
const newSort = e.target.value
|
||||
const changed = newSort !== currentSort
|
||||
currentSort = newSort
|
||||
if (changed) {
|
||||
imgGrid.innerHTML = ''
|
||||
const urls = await getUrls(subfolder)
|
||||
if (urls) {
|
||||
imgs = getImgsFromUrls(urls, imgGrid)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const sizeSlider = makeSlider(64, 1024, currentWidth, 1)
|
||||
imgTools.appendChild(orderSelect)
|
||||
imgTools.appendChild(sizeSlider)
|
||||
|
||||
imgs = getImgsFromUrls(urls, imgGrid)
|
||||
|
||||
sizeSlider.addEventListener('input', (e) => {
|
||||
currentWidth = e.target.value
|
||||
for (const img of imgs) {
|
||||
img.style.width = `${e.target.value}px`
|
||||
}
|
||||
})
|
||||
handle = renderSidebar(el, cont, [selector, imgGrid, imgTools])
|
||||
app.api.addEventListener('status', async () => {
|
||||
if (currentMode !== 'output') return
|
||||
updateOutputsGrid()
|
||||
})
|
||||
},
|
||||
destroy: () => {
|
||||
if (handle) {
|
||||
handle.unregister()
|
||||
handle = undefined
|
||||
app.api.removeEventListener('status')
|
||||
}
|
||||
},
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
app.registerExtension(sidebar_extension)
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
// NOTE: this will be the LT part of mtb API system
|
||||
// I need to properly publish the source and fix a few things before
|
||||
|
||||
// import { app } from '../../scripts/app.js'
|
||||
// // import { api } from '../../scripts/api.js'
|
||||
//
|
||||
// import * as shared from './comfy_shared.js'
|
||||
// import { createOutliner } from './dist/mtb_inspector.js'
|
||||
//
|
||||
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
||||
// const version = window?.__COMFYUI_FRONTEND_VERSION__
|
||||
// console.log(`%c ${version}`, 'background: orange; color: white;')
|
||||
//
|
||||
// const panel = app.extensionManager.registerSidebarTab({
|
||||
// id: 'mtb-nodes',
|
||||
// icon: 'pi pi-bolt',
|
||||
// title: 'MTB',
|
||||
// tooltip: 'MTB: API outliner',
|
||||
// type: 'custom',
|
||||
// // this is run everytime the tab's diplay is toggled on.
|
||||
// render: (el) => {
|
||||
// const outliner = createOutliner(el)
|
||||
// const inputs = shared.getAPIInputs()
|
||||
// console.log('INPUTS', inputs)
|
||||
// outliner.$$set({ inputs })
|
||||
// },
|
||||
// })
|
||||
// }
|
||||
+516
@@ -0,0 +1,516 @@
|
||||
/**
|
||||
* Adds a named stylesheet to the document with an optional ability to replace an existing one.
|
||||
*
|
||||
* @param {string} name - The unique name (ID) of the stylesheet.
|
||||
* @param {string} css - The CSS rules as a string.
|
||||
* @param {boolean} [force=false] - Whether to replace the existing stylesheet if it exists.
|
||||
* @returns {void}
|
||||
*/
|
||||
export function addNamedStyleSheet(name, css, force = false) {
|
||||
const existingStyleSheet = document.getElementById(name)
|
||||
|
||||
if (existingStyleSheet && !force) {
|
||||
console.debug(
|
||||
`Stylesheet with name "${name}" already exists. Skipping addition.`,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if (existingStyleSheet && force) {
|
||||
console.debug(`Stylesheet with name "${name}" exists. Replacing...`)
|
||||
existingStyleSheet.remove()
|
||||
}
|
||||
|
||||
const styleElement = document.createElement('style')
|
||||
styleElement.id = name
|
||||
styleElement.type = 'text/css'
|
||||
|
||||
styleElement.appendChild(document.createTextNode(css))
|
||||
document.head.appendChild(styleElement)
|
||||
|
||||
console.debug(`Stylesheet with name "${name}" added.`)
|
||||
}
|
||||
|
||||
export const ensureMTBStyles = () => {
|
||||
const S = {
|
||||
fg: 'var(--fg-color)',
|
||||
bgi: 'var(--comfy-input-bg)',
|
||||
bgm: 'var(--comfy-menu-bg)',
|
||||
border: 'var(--comfy-border)',
|
||||
borderHover: 'var(--comfy-border-hover)',
|
||||
box: 'var(--comfy-box)',
|
||||
accent: 'var(--p-button-text-primary-color)',
|
||||
}
|
||||
const common = `
|
||||
.mtb_sidebar {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
background: ${S.bgm};
|
||||
}
|
||||
.mtb_img_grid {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
overflow: scroll;
|
||||
gap: 1em;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: 100%;
|
||||
width: 100%;
|
||||
}
|
||||
.mtb_tools {
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
width: 100%;
|
||||
}
|
||||
`
|
||||
const inputs = `
|
||||
/* SELECT */
|
||||
.mtb_select {
|
||||
appearance: none;
|
||||
display: grid;
|
||||
grid-template-areas: "select";
|
||||
padding: 10px;
|
||||
background-color: ${S.bgi};
|
||||
border: none;
|
||||
border-radius: 5px;
|
||||
font-size: 14px;
|
||||
color: ${S.fg};
|
||||
cursor: pointer;
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
@supports (-moz-appearance:none) {
|
||||
.mtb_select{
|
||||
grid-area: select;
|
||||
background: ${S.bgi} url('data:image/gif;base64,R0lGODlhBgAGAKEDAFVVVX9/f9TU1CgmNyH5BAEKAAMALAAAAAAGAAYAAAIODA4hCDKWxlhNvmCnGwUAOw==') right center no-repeat !important;
|
||||
background-position: calc(100% - 5px) center !important;
|
||||
-moz-appearance:none !important;
|
||||
}
|
||||
|
||||
/* styling the dropdown arrow for browsers that support it */
|
||||
.mtb_select:after {
|
||||
content: "";
|
||||
width: 0.8em;
|
||||
height: 0.5em;
|
||||
background-color: ${S.fg};
|
||||
clip-path: polygon(100% 0%, 0 0%, 50% 100%);
|
||||
}
|
||||
|
||||
.mtb_select:focus {
|
||||
outline: none;
|
||||
border-color: #0056b3;
|
||||
}
|
||||
|
||||
.mtb_select > option {
|
||||
padding: 10px;
|
||||
background-color: ${S.bgi};
|
||||
border:none;
|
||||
color: ${S.fg};
|
||||
}
|
||||
|
||||
.mtb_select > option:hover {
|
||||
background-color: red;
|
||||
color: ${S.fg};
|
||||
}
|
||||
|
||||
/* SLIDER */
|
||||
.mtb_slider[type="range"] {
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 100%;
|
||||
height: 10px;
|
||||
background: ${S.bgm};
|
||||
border-radius: 5px;
|
||||
outline: none;
|
||||
opacity: 0.7;
|
||||
transition: opacity .2s;
|
||||
padding: 1em;
|
||||
}
|
||||
|
||||
/* slider track */
|
||||
.mtb_slider[type="range"]::-webkit-slider-runnable-track,
|
||||
.mtb_slider[type="range"]::-moz-range-track {
|
||||
width: 100%;
|
||||
height: 10px;
|
||||
background: ${S.bgi};
|
||||
border-radius: 5px;
|
||||
}
|
||||
|
||||
|
||||
/* progress */
|
||||
.mtb_slider[type="range"]::-moz-range-progress {
|
||||
background-color: ${S.accent};
|
||||
height:10px;
|
||||
border-radius: 5px;
|
||||
}
|
||||
|
||||
/* slider thumb (the handle) */
|
||||
.mtb_slider[type="range"]::-webkit-slider-thumb,
|
||||
.mtb_slider[type="range"]::-moz-range-thumb
|
||||
{
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 15px;
|
||||
height: 15px;
|
||||
border-radius: 50%;
|
||||
background: ${S.fg};
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
filter: drop-shadow(1px 1px 4px black);
|
||||
}
|
||||
|
||||
.mtb_slider[type="range"]:focus {
|
||||
opacity: 1;
|
||||
}
|
||||
|
||||
.mtb_slider[type=range]:-moz-focusring{
|
||||
outline: 1px solid red;
|
||||
outline-offset: -1px;
|
||||
}
|
||||
|
||||
.mtb_slider[type="range"]:hover::-webkit-slider-thumb,
|
||||
.mtb_slider[type="range"]:active::-webkit-slider-thumb {
|
||||
background-color: ${S.accent};
|
||||
}
|
||||
`
|
||||
addNamedStyleSheet(
|
||||
'mtb_ui',
|
||||
`
|
||||
${common}
|
||||
${inputs}
|
||||
`,
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Wrap an element with a div
|
||||
*
|
||||
* @param {Object} [style] - CSS styles to apply to the element.
|
||||
* @returns {HTMLElement} - The created DOM element.
|
||||
*/
|
||||
export const wrapElement = (element, style = {}) => {
|
||||
const container = makeElement('div', style)
|
||||
container.appendChild(element)
|
||||
return container
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a DOM element with optional styles, class, and id.
|
||||
*
|
||||
* @param {string} kind - The tag name of the element. Supports class and id syntax (e.g. 'div.class#id').
|
||||
* @param {Object} [style] - CSS styles to apply to the element.
|
||||
* @returns {HTMLElement} - The created DOM element.
|
||||
*/
|
||||
export const makeElement = (kind, style) => {
|
||||
let [real_kind, className] = kind.split('.')
|
||||
let id
|
||||
|
||||
if (className?.includes('#')) {
|
||||
;[className, id] = className.split('#')
|
||||
}
|
||||
|
||||
const el = document.createElement(real_kind)
|
||||
|
||||
if (style) {
|
||||
Object.assign(el.style, style)
|
||||
}
|
||||
|
||||
if (className) {
|
||||
el.classList.add(...className.split(' ')) // Support multiple classes
|
||||
}
|
||||
|
||||
if (id) {
|
||||
el.id = id
|
||||
}
|
||||
|
||||
return el
|
||||
}
|
||||
/**
|
||||
* Clears all child elements of the given parent element.
|
||||
*
|
||||
* @param {HTMLElement} el - The parent element whose children should be removed.
|
||||
*/
|
||||
export const clearElement = (el) => {
|
||||
while (el.firstChild) {
|
||||
el.removeChild(el.firstChild)
|
||||
}
|
||||
}
|
||||
/**
|
||||
* Creates a labeled element (input, select, etc.).
|
||||
*
|
||||
* @param {HTMLElement} el - The element to label.
|
||||
* @param {string} labelText - The label text.
|
||||
* @returns {HTMLDivElement} - A div containing the label and the element.
|
||||
*/
|
||||
export const makeLabeledElement = (el, labelText) => {
|
||||
const wrapper = makeElement('div.mtb_labeled_element', {
|
||||
marginBottom: '1em',
|
||||
})
|
||||
const label = makeElement('label', {
|
||||
display: 'block',
|
||||
marginBottom: '0.5em',
|
||||
})
|
||||
label.textContent = labelText
|
||||
wrapper.appendChild(label)
|
||||
wrapper.appendChild(el)
|
||||
return wrapper
|
||||
}
|
||||
|
||||
/**
|
||||
* Converts a camelCase CSS property to kebab-case.
|
||||
*
|
||||
* @param {string} prop - The camelCase CSS property.
|
||||
* @returns {string} - The kebab-case CSS property.
|
||||
*/
|
||||
const camelToKebab = (prop) =>
|
||||
prop.replace(/[A-Z]/g, (match) => `-${match.toLowerCase()}`)
|
||||
|
||||
/**
|
||||
* Parses the style string into an object of CSS property-value pairs.
|
||||
*
|
||||
* @param {string} styleString - The CSS rule text (e.g., "color: red; background-color: blue;").
|
||||
* @returns {Object} - An object with camelCase CSS properties.
|
||||
*/
|
||||
const parseStyleString = (styleString) => {
|
||||
const styleObj = {}
|
||||
for (const rule of styleString.split(';')) {
|
||||
const [property, value] = rule.split(':').map((item) => item.trim())
|
||||
if (property && value) {
|
||||
const camelProp = property.replace(/-([a-z])/g, (g) => g[1].toUpperCase())
|
||||
styleObj[camelProp] = value
|
||||
}
|
||||
}
|
||||
return styleObj
|
||||
}
|
||||
|
||||
/**
|
||||
* Defines a new CSS class with the provided styles, or skips if the class already exists.
|
||||
*
|
||||
* @param {string} className - The name of the CSS class to define.
|
||||
* @param {Object} classStyles - An object containing camelCase CSS property-value pairs.
|
||||
*/
|
||||
export function defineCSSClass(className, classStyles) {
|
||||
const styleSheets = document.styleSheets
|
||||
let classExists = false
|
||||
let existingStyleString = ''
|
||||
const classExistsInStyleSheet = (styleSheet) => {
|
||||
const rules = styleSheet.rules || styleSheet.cssRules
|
||||
for (const rule of rules) {
|
||||
if (rule.selectorText === `.${className}`) {
|
||||
classExists = true
|
||||
existingStyleString = rule.style.cssText // Capture existing styles
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
for (const styleSheet of styleSheets) {
|
||||
if (classExistsInStyleSheet(styleSheet)) {
|
||||
console.debug(`Class ${className} already exists, merging styles...`)
|
||||
break
|
||||
}
|
||||
}
|
||||
const existingStyles = classExists
|
||||
? parseStyleString(existingStyleString)
|
||||
: {}
|
||||
const mergedStyles = { ...existingStyles, ...classStyles }
|
||||
|
||||
const stylesString = Object.entries(mergedStyles)
|
||||
.map(([key, value]) => `${camelToKebab(key)}: ${value};`)
|
||||
.join(' ')
|
||||
|
||||
if (!classExists) {
|
||||
console.debug(`Defining new class ${className}...`)
|
||||
if (styleSheets[0].insertRule) {
|
||||
styleSheets[0].insertRule(`.${className} { ${stylesString} }`, 0)
|
||||
} else if (styleSheets[0].addRule) {
|
||||
styleSheets[0].addRule(`.${className}`, stylesString, 0)
|
||||
}
|
||||
} else {
|
||||
console.debug(`Updating existing class ${className} with merged styles...`)
|
||||
for (const styleSheet of styleSheets) {
|
||||
const rules = styleSheet.rules || styleSheet.cssRules
|
||||
for (const rule of rules) {
|
||||
if (rule.selectorText === `.${className}`) {
|
||||
rule.style.cssText = stylesString // Update the existing rule
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
console.debug(
|
||||
`Class ${className} has been defined/updated with styles:`,
|
||||
mergedStyles,
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Renders a sidebar and ensures it resizes correctly when the window is resized.
|
||||
*
|
||||
* @param {HTMLElement} el - The element where the sidebar is rendered.
|
||||
* @param {HTMLElement} cont - The content container of the sidebar.
|
||||
* @param {HTMLElement[]} elems - Array of elements to append to the sidebar.
|
||||
* @returns {Object} - A handle with a method to unregister the resize event.
|
||||
*/
|
||||
export const renderSidebar = (el, cont, elems) => {
|
||||
el.appendChild(cont)
|
||||
|
||||
if (!el.parentNode) {
|
||||
return
|
||||
}
|
||||
el.parentNode.style.overflowY = 'clip'
|
||||
cont.style.height = `${el.parentNode.offsetHeight}px`
|
||||
|
||||
const resizeHandler = () => {
|
||||
cont.style.height = `${el.parentNode.offsetHeight}px`
|
||||
}
|
||||
window.addEventListener('resize', resizeHandler)
|
||||
|
||||
for (const elem of elems) {
|
||||
cont.appendChild(elem)
|
||||
}
|
||||
|
||||
return {
|
||||
unregister: () => {
|
||||
window.removeEventListener('resize', resizeHandler)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a <select> dropdown with given options.
|
||||
*
|
||||
* @param {string[]} options - The options for the select element.
|
||||
* @param {string} [current] - The currently selected option (optional).
|
||||
* @returns {HTMLSelectElement} - The created <select> element.
|
||||
*/
|
||||
export const makeSelect = (options, current = undefined) => {
|
||||
const selector = makeElement('select.mtb_select', {
|
||||
width: 'auto',
|
||||
margin: '1em',
|
||||
})
|
||||
|
||||
for (const option of options) {
|
||||
const opt = makeElement('option')
|
||||
opt.value = option
|
||||
opt.innerHTML = option
|
||||
selector.appendChild(opt)
|
||||
}
|
||||
|
||||
if (current !== undefined) {
|
||||
if (options.includes(current)) {
|
||||
selector.value = current
|
||||
} else {
|
||||
console.error(
|
||||
`You tried to select an option that doesn't exist (${current}). Options: ${options}`,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return selector
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates an <input type="range"> slider element with given parameters.
|
||||
*
|
||||
* @param {number} min - Minimum value of the slider.
|
||||
* @param {number} max - Maximum value of the slider.
|
||||
* @param {number} [value] - Initial value of the slider.
|
||||
* @param {number} [step] - Step value for the slider.
|
||||
* @returns {HTMLInputElement} - The created slider element.
|
||||
*/
|
||||
export const makeSlider = (min, max, value = undefined, step = undefined) => {
|
||||
const slider = makeElement('input.mtb_slider', {
|
||||
width: '100%',
|
||||
})
|
||||
|
||||
slider.type = 'range'
|
||||
slider.min = min || 0
|
||||
slider.max = max || 100
|
||||
slider.value = value || slider.min
|
||||
slider.step = step || 1
|
||||
|
||||
return slider
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a button element.
|
||||
*
|
||||
* @param {string} label - The label for the button.
|
||||
* @param {Object} [style] - Optional styles to apply to the button.
|
||||
* @param {Function} [onClick] - Optional click handler.
|
||||
* @returns {HTMLButtonElement} - The created button element.
|
||||
*/
|
||||
export const makeButton = (label, style = {}, onClick = undefined) => {
|
||||
const button = makeElement('button.mtb_button', style)
|
||||
button.textContent = label
|
||||
|
||||
if (onClick) {
|
||||
button.addEventListener('click', onClick)
|
||||
}
|
||||
|
||||
return button
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a resizable splitter between two elements.
|
||||
*
|
||||
* @param {HTMLElement} el1 - The first element.
|
||||
* @param {HTMLElement} el2 - The second element.
|
||||
* @param {'vertical' | 'horizontal'} direction - Splitter direction (vertical or horizontal).
|
||||
* @param {'absolute' | 'normal'} mode - Splitter mode: 'absolute' for free resizing, 'normal' for layout-based resizing.
|
||||
* @returns {HTMLDivElement} - The container with resizable splitter.
|
||||
*/
|
||||
export const makeSplitter = (
|
||||
el1,
|
||||
el2,
|
||||
direction = 'vertical',
|
||||
mode = 'normal',
|
||||
) => {
|
||||
const container = makeElement('div.mtb_splitter_container', {
|
||||
display: mode === 'absolute' ? 'block' : 'flex',
|
||||
flexDirection: direction === 'vertical' ? 'row' : 'column',
|
||||
position: mode === 'absolute' ? 'relative' : 'static',
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
})
|
||||
|
||||
const handle = makeElement('div.mtb_splitter_handle', {
|
||||
backgroundColor: '#ccc',
|
||||
cursor: direction === 'vertical' ? 'col-resize' : 'row-resize',
|
||||
width: direction === 'vertical' ? '5px' : '100%',
|
||||
height: direction === 'horizontal' ? '5px' : '100%',
|
||||
})
|
||||
|
||||
let isResizing = false
|
||||
|
||||
handle.addEventListener('mousedown', () => {
|
||||
isResizing = true
|
||||
})
|
||||
|
||||
window.addEventListener('mouseup', () => {
|
||||
isResizing = false
|
||||
})
|
||||
|
||||
window.addEventListener('mousemove', (e) => {
|
||||
if (!isResizing) return
|
||||
if (direction === 'vertical') {
|
||||
const newWidth = e.clientX - container.offsetLeft
|
||||
el1.style.width = `${newWidth}px`
|
||||
el2.style.width = `${container.offsetWidth - newWidth}px`
|
||||
} else {
|
||||
const newHeight = e.clientY - container.offsetTop
|
||||
el1.style.height = `${newHeight}px`
|
||||
el2.style.height = `${container.offsetHeight - newHeight}px`
|
||||
}
|
||||
})
|
||||
|
||||
container.appendChild(el1)
|
||||
container.appendChild(handle)
|
||||
container.appendChild(el2)
|
||||
|
||||
return container
|
||||
}
|
||||
+302
-118
@@ -14,13 +14,14 @@
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { api } from '../../scripts/api.js'
|
||||
|
||||
import * as mtb_ui from './mtb_ui.js'
|
||||
import parseCss from './extern/parse-css.js'
|
||||
import * as shared from './comfy_shared.js'
|
||||
import { infoLogger } from './comfy_shared.js'
|
||||
import { NumberInputWidget } from './numberInput.js'
|
||||
|
||||
// NOTE: new widget types registered by MTB Widgets
|
||||
const newTypes = [/*'BOOL'*/ , 'COLOR', 'BBOX']
|
||||
const newTypes = [/*'BOOL'*/ 'COLOR', 'BBOX']
|
||||
|
||||
const deprecated_nodes = {
|
||||
// 'Animation Builder':
|
||||
@@ -96,7 +97,7 @@ export function addVectorWidgetW(
|
||||
name,
|
||||
value,
|
||||
vector_size,
|
||||
callback,
|
||||
_callback,
|
||||
app,
|
||||
) {
|
||||
// const inputEl = document.createElement('div')
|
||||
@@ -243,7 +244,7 @@ export const MtbWidgets = {
|
||||
y: 0,
|
||||
options: { default: Array.from({ length: size }, () => 0.0) },
|
||||
_value: val || Array.from({ length: size }, () => 0.0),
|
||||
draw: function (ctx, node, width, widgetY, height) {
|
||||
draw: (ctx, node, width, widgetY, height) => {
|
||||
ctx.textAlign = 'left'
|
||||
ctx.strokeStyle = outline_color
|
||||
ctx.fillStyle = background_color
|
||||
@@ -311,7 +312,7 @@ export const MtbWidgets = {
|
||||
value: val?.default || [0, 0, 0, 0],
|
||||
options: {},
|
||||
|
||||
draw: function (ctx, node, widget_width, widgetY, height) {
|
||||
draw: function (ctx, _node, widget_width, widgetY, _height) {
|
||||
const hide = this.type !== 'BBOX' && app.canvas.ds.scale > 0.5
|
||||
|
||||
const show_text = true
|
||||
@@ -321,13 +322,13 @@ export const MtbWidgets = {
|
||||
const secondary_text_color = LiteGraph.WIDGET_SECONDARY_TEXT_COLOR
|
||||
const H = LiteGraph.NODE_WIDGET_HEIGHT
|
||||
|
||||
let margin = 15
|
||||
let numWidgets = 4 // Number of stacked widgets
|
||||
const margin = 15
|
||||
const numWidgets = 4 // Number of stacked widgets
|
||||
|
||||
if (hide) return
|
||||
|
||||
for (let i = 0; i < numWidgets; i++) {
|
||||
let currentY = widgetY + i * (H + margin) // Adjust Y position for each widget
|
||||
const currentY = widgetY + i * (H + margin) // Adjust Y position for each widget
|
||||
|
||||
ctx.textAlign = 'left'
|
||||
ctx.strokeStyle = outline_color
|
||||
@@ -535,21 +536,34 @@ export const MtbWidgets = {
|
||||
picker.type = 'color'
|
||||
picker.value = this.value
|
||||
|
||||
picker.style.position = 'absolute'
|
||||
picker.style.left = '999999px' //(window.innerWidth / 2) + "px";
|
||||
picker.style.top = '999999px' //(window.innerHeight / 2) + "px";
|
||||
Object.assign(picker.style, {
|
||||
position: 'fixed',
|
||||
left: `${e.clientX}px`,
|
||||
top: `${e.clientY}px`,
|
||||
height: '0px',
|
||||
width: '0px',
|
||||
padding: '0px',
|
||||
opacity: 0,
|
||||
})
|
||||
|
||||
picker.addEventListener('blur', () => {
|
||||
this.callback?.(this.value)
|
||||
node.graph._version++
|
||||
picker.remove()
|
||||
})
|
||||
picker.addEventListener('input', () => {
|
||||
if (!picker.value) return
|
||||
|
||||
this.value = picker.value
|
||||
app.canvas.setDirty(true)
|
||||
})
|
||||
|
||||
document.body.appendChild(picker)
|
||||
|
||||
picker.addEventListener('change', () => {
|
||||
this.value = picker.value
|
||||
this.callback?.(this.value)
|
||||
node.graph._version++
|
||||
node.setDirtyCanvas(true, true)
|
||||
picker.remove()
|
||||
requestAnimationFrame(() => {
|
||||
picker.showPicker()
|
||||
picker.focus()
|
||||
})
|
||||
|
||||
picker.click()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -658,29 +672,38 @@ const mtb_widgets = {
|
||||
init: async () => {
|
||||
infoLogger('Registering mtb.widgets')
|
||||
try {
|
||||
const res = await api.fetchApi('/mtb/debug')
|
||||
const msg = await res.json()
|
||||
const msg = await shared.getServerInfo()
|
||||
if (!window.MTB) {
|
||||
window.MTB = {}
|
||||
}
|
||||
window.MTB.DEBUG = msg.enabled
|
||||
window.MTB.DEBUG = msg.debug
|
||||
} catch (e) {
|
||||
console.error('Error:', error)
|
||||
console.error('Error:', e)
|
||||
}
|
||||
},
|
||||
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Debug.enabled',
|
||||
name: '[⚡mtb] Enable Debug (py and js)',
|
||||
id: 'mtb.postshot.path',
|
||||
category: ['mtb', 'PostShot', 'path'],
|
||||
name: 'Path to Postshot CLI',
|
||||
type: 'string',
|
||||
defaultValue: 'C:/Program Files/Jawset Postshot/bin/postshot-cli.exe',
|
||||
tooltip: 'The path to the postshot CLI',
|
||||
})
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.debug-enabled',
|
||||
category: ['mtb', 'Main', 'debug-enabled'],
|
||||
name: 'Enable Debug (py and js)',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
|
||||
tooltip:
|
||||
'This will enable debug messages in the console and in the python console respectively',
|
||||
'This will enable debug messages in the console and in the python console respectively, no need to restart the server, but do reload the webui',
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
@@ -692,34 +715,28 @@ const mtb_widgets = {
|
||||
infoLogger('Enabled DEBUG mode')
|
||||
}
|
||||
|
||||
await api
|
||||
.fetchApi('/mtb/debug', {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({
|
||||
enabled: value,
|
||||
}),
|
||||
})
|
||||
.then((_response) => {})
|
||||
.catch((error) => {
|
||||
console.error('Error:', error)
|
||||
})
|
||||
try {
|
||||
shared.setServerInfo({ debug: value })
|
||||
} catch (err) {
|
||||
console.error('Error:', err)
|
||||
}
|
||||
},
|
||||
})
|
||||
},
|
||||
|
||||
getCustomWidgets: () => {
|
||||
return {
|
||||
BOOL: (node, inputName, inputData, _app) => {
|
||||
console.debug('Registering bool')
|
||||
|
||||
return {
|
||||
widget: node.addCustomWidget(
|
||||
MtbWidgets.BOOL(inputName, inputData[1]?.default || false),
|
||||
),
|
||||
minWidth: 150,
|
||||
minHeight: 30,
|
||||
}
|
||||
},
|
||||
// BOOL: (node, inputName, inputData, _app) => {
|
||||
// console.debug('Registering bool')
|
||||
//
|
||||
// return {
|
||||
// widget: node.addCustomWidget(
|
||||
// MtbWidgets.BOOL(inputName, inputData[1]?.default || false),
|
||||
// ),
|
||||
// minWidth: 150,
|
||||
// minHeight: 30,
|
||||
// }
|
||||
// },
|
||||
|
||||
COLOR: (node, inputName, inputData, _app) => {
|
||||
console.debug('Registering color')
|
||||
@@ -751,7 +768,7 @@ const mtb_widgets = {
|
||||
// const rinputs = nodeData.input?.required
|
||||
|
||||
let has_custom = false
|
||||
if (nodeData.input && nodeData.input.required) {
|
||||
if (nodeData.input?.required) {
|
||||
for (const i of Object.keys(nodeData.input.required)) {
|
||||
const input_type = nodeData.input.required[i][0]
|
||||
|
||||
@@ -764,10 +781,8 @@ const mtb_widgets = {
|
||||
if (has_custom) {
|
||||
//- Add widgets on node creation
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
||||
this.serialize_widgets = true
|
||||
this.setSize?.(this.computeSize())
|
||||
|
||||
@@ -785,8 +800,8 @@ const mtb_widgets = {
|
||||
? origGetExtraMenuOptions.apply(this, arguments)
|
||||
: undefined
|
||||
if (this.widgets) {
|
||||
let toInput = []
|
||||
let toWidget = []
|
||||
const toInput = []
|
||||
const toWidget = []
|
||||
for (const w of this.widgets) {
|
||||
if (w.type === shared.CONVERTED_TYPE) {
|
||||
//- This is already handled by widgetinputs.js
|
||||
@@ -832,7 +847,8 @@ const mtb_widgets = {
|
||||
//- Extending Python Nodes
|
||||
switch (nodeData.name) {
|
||||
//TODO: remove this non sense
|
||||
case 'Get Batch From History (mtb)': {
|
||||
case 'Get Batch From History (mtb)':
|
||||
case 'Get Batch From History V2 (mtb)': {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||
@@ -855,6 +871,22 @@ const mtb_widgets = {
|
||||
|
||||
break
|
||||
}
|
||||
case 'Postshot Train (mtb)':
|
||||
case 'Postshot Export (mtb)': {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
||||
const { postshot_cli } = shared.getNamedWidget(this, 'postshot_cli')
|
||||
|
||||
shared.hideWidgetForGood(this, postshot_cli)
|
||||
|
||||
api.getSetting('mtb.postshot.path').then((p) => {
|
||||
postshot_cli._value = p
|
||||
})
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
case 'Save Gif (mtb)':
|
||||
case 'Save Animated Image (mtb)': {
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
@@ -877,7 +909,7 @@ const mtb_widgets = {
|
||||
imgURLs = imgURLs.concat(
|
||||
message.gif.map((params) => {
|
||||
return api.apiURL(
|
||||
'/view?' + new URLSearchParams(params).toString(),
|
||||
`/view?${new URLSearchParams(params).toString()}`,
|
||||
)
|
||||
}),
|
||||
)
|
||||
@@ -886,7 +918,7 @@ const mtb_widgets = {
|
||||
imgURLs = imgURLs.concat(
|
||||
message.apng.map((params) => {
|
||||
return api.apiURL(
|
||||
'/view?' + new URLSearchParams(params).toString(),
|
||||
`/view?${new URLSearchParams(params).toString()}`,
|
||||
)
|
||||
}),
|
||||
)
|
||||
@@ -914,37 +946,71 @@ const mtb_widgets = {
|
||||
}
|
||||
case 'Animation Builder (mtb)': {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
||||
|
||||
this.changeMode(LiteGraph.ALWAYS)
|
||||
|
||||
const raw_iteration = this.widgets.find(
|
||||
(w) => w.name === 'raw_iteration',
|
||||
)
|
||||
const raw_loop = this.widgets.find((w) => w.name === 'raw_loop')
|
||||
|
||||
const total_frames = this.widgets.find(
|
||||
(w) => w.name === 'total_frames',
|
||||
)
|
||||
const loop_count = this.widgets.find((w) => w.name === 'loop_count')
|
||||
const { raw_iteration, raw_loop, total_frames, loop_count } =
|
||||
shared.getNamedWidget(
|
||||
this,
|
||||
'raw_iteration',
|
||||
'raw_loop',
|
||||
'total_frames',
|
||||
'loop_count',
|
||||
)
|
||||
|
||||
shared.hideWidgetForGood(this, raw_iteration)
|
||||
shared.hideWidgetForGood(this, raw_loop)
|
||||
|
||||
raw_iteration._value = 0
|
||||
|
||||
const value_preview = this.addCustomWidget(
|
||||
MtbWidgets['DEBUG_STRING']('value_preview', 'Idle'),
|
||||
)
|
||||
value_preview.parent = this
|
||||
// const value_preview = this.addCustomWidget(
|
||||
// MtbWidgets.DEBUG_STRING('value_preview', 'Idle'),
|
||||
// )
|
||||
|
||||
const loop_preview = this.addCustomWidget(
|
||||
MtbWidgets['DEBUG_STRING']('loop_preview', 'Iteration: Idle'),
|
||||
const dom_value_preview = mtb_ui.makeElement('p', {
|
||||
fontWeigth: '700',
|
||||
textAlign: 'center',
|
||||
fontSize: '1.5em',
|
||||
margin: 0,
|
||||
})
|
||||
const value_preview = this.addDOMWidget(
|
||||
'value_preview',
|
||||
'DISPLAY',
|
||||
dom_value_preview,
|
||||
{
|
||||
hideOnZoom: false,
|
||||
setValue: (val) => {
|
||||
if (val) {
|
||||
value_preview.element.innerHTML = val
|
||||
}
|
||||
},
|
||||
},
|
||||
)
|
||||
loop_preview.parent = this
|
||||
value_preview.value = 'Idle'
|
||||
|
||||
const dom_loop_preview = mtb_ui.makeElement('p', {
|
||||
textAlign: 'center',
|
||||
margin: 0,
|
||||
})
|
||||
|
||||
const loop_preview = this.addDOMWidget(
|
||||
'loop_preview',
|
||||
'DISPLAY',
|
||||
dom_loop_preview,
|
||||
{
|
||||
hideOnZoom: false,
|
||||
setValue: (val) => {
|
||||
if (val) {
|
||||
dom_loop_preview.innerHTML = val
|
||||
}
|
||||
},
|
||||
getValue: () => {
|
||||
dom_loop_preview.innerHTML
|
||||
},
|
||||
},
|
||||
)
|
||||
loop_preview.value = 'Iteration: Idle'
|
||||
|
||||
const onReset = () => {
|
||||
raw_iteration.value = 0
|
||||
@@ -957,10 +1023,10 @@ const mtb_widgets = {
|
||||
}
|
||||
|
||||
// reset button
|
||||
this.addWidget('button', `Reset`, 'reset', onReset)
|
||||
this.addWidget('button', 'Reset', 'reset', onReset)
|
||||
|
||||
// run button
|
||||
this.addWidget('button', `Queue`, 'queue', () => {
|
||||
this.addWidget('button', 'Queue', 'queue', () => {
|
||||
onReset() // this could maybe be a setting or checkbox
|
||||
app.queuePrompt(0, total_frames.value * loop_count.value)
|
||||
window.MTB?.notify?.(
|
||||
@@ -1000,9 +1066,9 @@ const mtb_widgets = {
|
||||
}
|
||||
case 'Interpolate Clip Sequential (mtb)': {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
? onNodeCreated.apply(this, ...args)
|
||||
: undefined
|
||||
const addReplacement = () => {
|
||||
const input = this.addInput(
|
||||
@@ -1014,19 +1080,14 @@ const mtb_widgets = {
|
||||
this.addWidget('STRING', `replacement_${this.widgets.length}`, '')
|
||||
}
|
||||
//- add
|
||||
this.addWidget('button', '+', 'add', function (value, widget, node) {
|
||||
this.addWidget('button', '+', 'add', (value, widget, node) => {
|
||||
console.log('Button clicked', value, widget, node)
|
||||
addReplacement()
|
||||
})
|
||||
//- remove
|
||||
this.addWidget(
|
||||
'button',
|
||||
'-',
|
||||
'remove',
|
||||
function (value, widget, node) {
|
||||
console.log(`Button clicked: ${value}`, widget, node)
|
||||
},
|
||||
)
|
||||
this.addWidget('button', '-', 'remove', (value, widget, node) => {
|
||||
console.log(`Button clicked: ${value}`, widget, node)
|
||||
})
|
||||
|
||||
return r
|
||||
}
|
||||
@@ -1041,16 +1102,10 @@ const mtb_widgets = {
|
||||
|
||||
const getStyle = async (node) => {
|
||||
try {
|
||||
const getStyles = await api.fetchApi('/mtb/actions', {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({
|
||||
name: 'getStyles',
|
||||
args:
|
||||
node.widgets && node.widgets[0].value
|
||||
? node.widgets[0].value
|
||||
: '',
|
||||
}),
|
||||
})
|
||||
const getStyles = await runAction(
|
||||
'getStyles',
|
||||
node.widgets?.[0].value ? node.widgets[0].value : '',
|
||||
)
|
||||
|
||||
const output = await getStyles.json()
|
||||
return output?.result
|
||||
@@ -1122,6 +1177,10 @@ const mtb_widgets = {
|
||||
shared.setupDynamicConnections(nodeType, 'video', 'VIDEO')
|
||||
break
|
||||
}
|
||||
case 'Interpolate Condition (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'condition', 'CONDITIONING')
|
||||
break
|
||||
}
|
||||
case 'Psd Save (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'input_', 'PSDLAYER')
|
||||
break
|
||||
@@ -1133,7 +1192,11 @@ const mtb_widgets = {
|
||||
case 'Stack Images (mtb)':
|
||||
case 'Concat Images (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'image', 'IMAGE')
|
||||
|
||||
break
|
||||
}
|
||||
case 'Audio Sequence (mtb)':
|
||||
case 'Audio Stack (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'audio', 'AUDIO')
|
||||
break
|
||||
}
|
||||
case 'Batch Float Assemble (mtb)':
|
||||
@@ -1142,6 +1205,8 @@ const mtb_widgets = {
|
||||
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
|
||||
break
|
||||
}
|
||||
case 'Batch Sequence (mtb)':
|
||||
case 'Batch Sequence Plus (mtb)':
|
||||
case 'Batch Merge (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE')
|
||||
|
||||
@@ -1154,13 +1219,13 @@ const mtb_widgets = {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
this.addInput(`x`, '*')
|
||||
this.addInput('x', '*')
|
||||
return r
|
||||
}
|
||||
|
||||
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
||||
nodeType.prototype.onConnectionsChange = function (
|
||||
type,
|
||||
_type,
|
||||
index,
|
||||
connected,
|
||||
link_info,
|
||||
@@ -1175,7 +1240,7 @@ const mtb_widgets = {
|
||||
//- infer type
|
||||
if (link_info) {
|
||||
const fromNode = this.graph._nodes.find(
|
||||
(otherNode) => otherNode.id == link_info.origin_id,
|
||||
(otherNode) => otherNode.id !== link_info.origin_id,
|
||||
)
|
||||
const type = fromNode.outputs[link_info.origin_slot].type
|
||||
this.inputs[index].type = type
|
||||
@@ -1192,6 +1257,7 @@ const mtb_widgets = {
|
||||
}
|
||||
|
||||
case 'Batch Shape (mtb)':
|
||||
case 'Mask To Image (mtb)':
|
||||
case 'Text To Image (mtb)': {
|
||||
shared.addMenuHandler(nodeType, function (_app, options) {
|
||||
/** @type {ContextMenuItem} */
|
||||
@@ -1217,23 +1283,141 @@ const mtb_widgets = {
|
||||
})
|
||||
break
|
||||
}
|
||||
case 'Save Tensors (mtb)': {
|
||||
case 'Scene Detect (mtb)': {
|
||||
break
|
||||
}
|
||||
case 'Loop Start (mtb)': {
|
||||
const onDrawBackground = nodeType.prototype.onDrawBackground
|
||||
nodeType.prototype.onDrawBackground = function (ctx, canvas) {
|
||||
nodeType.prototype.onDrawBackground = function (...args) {
|
||||
const r = onDrawBackground
|
||||
? onDrawBackground.apply(this, arguments)
|
||||
? onDrawBackground.apply(this, args)
|
||||
: undefined
|
||||
// // draw a circle on the top right of the node, with text inside
|
||||
// ctx.fillStyle = "#fff";
|
||||
// ctx.beginPath();
|
||||
// ctx.arc(this.size[0] - this.node_width * 0.5, this.size[1] - this.node_height * 0.5, this.node_width * 0.5, 0, Math.PI * 2);
|
||||
// ctx.fill();
|
||||
const [ctx, /*canvas,*/ ..._rest] = args
|
||||
if (this.flags.collapsed) return r
|
||||
if (!this.computed_flow) {
|
||||
const related = new Set([this.id])
|
||||
const visited = new Set()
|
||||
if (this.outputs[0].links) {
|
||||
const initLink = this.outputs[0].links[0]
|
||||
const { to: loopEnd } = shared.nodesFromLink(this, initLink)
|
||||
const canReachEnd = (node, visited = new Set()) => {
|
||||
if (node === loopEnd) return true
|
||||
if (visited.has(node.id)) return false
|
||||
visited.add(node.id)
|
||||
for (const output of node.outputs || []) {
|
||||
if (!output.links) continue
|
||||
for (const linkId of output.links) {
|
||||
const { to: nextNode } = shared.nodesFromLink(node, linkId)
|
||||
if (!nextNode) continue
|
||||
if (canReachEnd(nextNode, visited)) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
const traverseNodes = (node) => {
|
||||
if (visited.has(node.id)) return
|
||||
visited.add(node.id)
|
||||
|
||||
// ctx.fillStyle = "#000";
|
||||
// ctx.textAlign = "center";
|
||||
// ctx.font = "bold 12px Arial";
|
||||
// ctx.fillText("Save Tensors", this.size[0] - this.node_width * 0.5, this.size[1] - this.node_height * 0.5);
|
||||
// can reach the end
|
||||
if (node !== this && node !== loopEnd && !canReachEnd(node)) {
|
||||
return
|
||||
}
|
||||
|
||||
related.add(node.id)
|
||||
for (const output of node.outputs || []) {
|
||||
if (!output.links) continue
|
||||
|
||||
for (const linkId of output.links) {
|
||||
const { to: nextNode } = shared.nodesFromLink(node, linkId)
|
||||
if (!nextNode) continue
|
||||
|
||||
traverseNodes(nextNode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
traverseNodes(this)
|
||||
}
|
||||
this.related_to_flow = Array.from(related)
|
||||
this.computed_flow = true
|
||||
}
|
||||
if (this.related_to_flow) {
|
||||
ctx.save()
|
||||
const points = []
|
||||
const padding = 20
|
||||
|
||||
const graph = this.graph
|
||||
const offset = this._pos
|
||||
|
||||
for (const nodeId of this.related_to_flow) {
|
||||
const node = graph.getNodeById(nodeId)
|
||||
if (!node) continue
|
||||
|
||||
const scale = 1.0
|
||||
const x = node._pos[0] * scale - offset[0]
|
||||
const y = node._pos[1] * scale - offset[1]
|
||||
const width = node.size[0] * scale
|
||||
const height = node.size[1] * scale
|
||||
const scaledPadding = padding * scale
|
||||
// console.log({ main: this, x, y, width, height })
|
||||
|
||||
points.push(
|
||||
[x - scaledPadding, y - scaledPadding],
|
||||
[x + width + scaledPadding, y - scaledPadding],
|
||||
[x + width + scaledPadding, y + height + scaledPadding],
|
||||
[x - scaledPadding, y + height + scaledPadding],
|
||||
)
|
||||
}
|
||||
// console.log({ points })
|
||||
const hull = shared.getConvexHull(points)
|
||||
|
||||
ctx.beginPath()
|
||||
|
||||
ctx.moveTo(hull[0][0], hull[0][1])
|
||||
for (let i = 1; i < hull.length; i++) {
|
||||
ctx.lineTo(hull[i][0], hull[i][1])
|
||||
}
|
||||
|
||||
ctx.closePath()
|
||||
|
||||
ctx.fillStyle = 'rgba(255, 0, 0, 0.1)'
|
||||
ctx.strokeStyle = 'rgba(255, 0, 0, 0.5)'
|
||||
ctx.lineWidth = 2
|
||||
ctx.fill()
|
||||
ctx.stroke()
|
||||
|
||||
ctx.restore()
|
||||
} else {
|
||||
ctx.save()
|
||||
ctx.fillStyle = 'red'
|
||||
ctx.fillRect(-50, -50, this.size[0] + 100, this.size[1] + 100)
|
||||
ctx.fillStyle = 'white'
|
||||
ctx.font = 'bold 12px Arial'
|
||||
ctx.fillText(
|
||||
`pos: ${this.x}x${this.y}`,
|
||||
this.size[0] / 2,
|
||||
this.size[1],
|
||||
)
|
||||
ctx.fillText(
|
||||
`size:${this._posSize}`,
|
||||
this.size[0] / 2,
|
||||
this.size[1] - 30,
|
||||
)
|
||||
ctx.fillText(
|
||||
`dpi: ${window.devicePixelRatio}`,
|
||||
this.size[0] / 2,
|
||||
this.size[1] - 60,
|
||||
)
|
||||
ctx.fillText(
|
||||
`next: ${graph.getNodeById(this.related_to_flow[1])._posSize}`,
|
||||
this.size[0] / 2,
|
||||
this.size[1] - 90,
|
||||
)
|
||||
|
||||
ctx.restore()
|
||||
}
|
||||
return r
|
||||
}
|
||||
break
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
// web/note_plus.constants.js
|
||||
|
||||
export const DEFAULT_CSS = ''
|
||||
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
|
||||
Note+
|
||||
</p>`
|
||||
export const DEFAULT_MD = '## Note+'
|
||||
export const DEFAULT_MODE = 'markdown'
|
||||
export const DEFAULT_THEME = 'one_dark'
|
||||
|
||||
export const DEMO_CONTENT = `
|
||||
# @mtb/svelte-markdown.
|
||||
## This is a subheader
|
||||
|
||||
[](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml)
|
||||

|
||||
|
||||
<details>
|
||||
<summary>More details about the inception of the project</summary>
|
||||
|
||||
\`\`\`js
|
||||
class YesMan{
|
||||
constructor(){
|
||||
this.started = false
|
||||
}
|
||||
}
|
||||
\`\`\`
|
||||
</details>
|
||||
|
||||
This is a paragraph. If it goes over the maximum width it will not automatically wrap unless it reaches the max-w of \`prose\` check [styles](/styles) for more info.
|
||||
|
||||
This component is useful for building some tools on top. Or even just a static system using svelte at its core. My personal blog is fully powered by **@mtb/svelte-markdown**
|
||||
|
||||
| And this is | A table |
|
||||
|-------------|---------|
|
||||
| With two | columns |
|
||||
|
||||
We also support github callout:
|
||||
|
||||
|
||||
> [!NOTE]
|
||||
> Highlights information that users should take into account, even when skimming.
|
||||
> [!TIP]
|
||||
> Optional information to help a user be more successful.
|
||||
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Crucial information necessary for users to succeed.
|
||||
|
||||
> [!WARNING]
|
||||
> Critical content demanding immediate user attention due to potential risks.
|
||||
|
||||
> [!CAUTION]
|
||||
> Negative potential consequences of an action.
|
||||
`
|
||||
|
||||
export const THEMES = [
|
||||
'ambiance',
|
||||
'chaos',
|
||||
'chrome',
|
||||
'cloud9_day',
|
||||
'cloud9_night',
|
||||
'cloud9_night_low_color',
|
||||
'cloud_editor',
|
||||
'cloud_editor_dark',
|
||||
'clouds',
|
||||
'clouds_midnight',
|
||||
'cobalt',
|
||||
'crimson_editor',
|
||||
'dawn',
|
||||
'dracula',
|
||||
'dreamweaver',
|
||||
'eclipse',
|
||||
'github',
|
||||
'github_dark',
|
||||
'gob',
|
||||
'gruvbox',
|
||||
'gruvbox_dark_hard',
|
||||
'gruvbox_light_hard',
|
||||
'idle_fingers',
|
||||
'iplastic',
|
||||
'katzenmilch',
|
||||
'kr_theme',
|
||||
'kuroir',
|
||||
'merbivore',
|
||||
'merbivore_soft',
|
||||
'mono_industrial',
|
||||
'monokai',
|
||||
'nord_dark',
|
||||
'one_dark',
|
||||
'pastel_on_dark',
|
||||
'solarized_dark',
|
||||
'solarized_light',
|
||||
'sqlserver',
|
||||
'terminal',
|
||||
'textmate',
|
||||
'tomorrow',
|
||||
'tomorrow_night',
|
||||
'tomorrow_night_blue',
|
||||
'tomorrow_night_bright',
|
||||
'tomorrow_night_eighties',
|
||||
'twilight',
|
||||
'vibrant_ink',
|
||||
'vscode',
|
||||
]
|
||||
|
||||
export const CSS_RESET = `
|
||||
* {
|
||||
font-family: monospace;
|
||||
line-height: 1.25em;
|
||||
}
|
||||
.shiki{
|
||||
padding: 1em;
|
||||
width: 100%;
|
||||
}
|
||||
.markdown-callout-title {
|
||||
.octicon{
|
||||
fill:white;
|
||||
}
|
||||
/* background: var(--current-color); */
|
||||
color: var(--current-color);
|
||||
font-weight: bold;
|
||||
/* border-start-end-radius: var(--radius); */
|
||||
/* border-start-start-radius: var(--radius); */
|
||||
padding: 0.5em;
|
||||
padding-inline-start: 1em;
|
||||
}
|
||||
.markdown-callout-content {
|
||||
padding: 1em;
|
||||
}
|
||||
.markdown-callout {
|
||||
--radius: 8px;
|
||||
--current-color: purple;
|
||||
/* border-start-end-radius: var(--radius); */
|
||||
/* border-start-start-radius: var(--radius); */
|
||||
border-left: 3px solid var(--current-color);
|
||||
margin-bottom: 1em;
|
||||
margin-top: 1em;
|
||||
}
|
||||
|
||||
.markdown-callout-tip {
|
||||
--text-color: whitesmoke;
|
||||
--current-color: #50e3c2;
|
||||
}
|
||||
|
||||
.markdown-callout-note {
|
||||
--text-color: whitesmoke;
|
||||
--current-color: #0070f3;
|
||||
}
|
||||
.markdown-callout-important {
|
||||
--text-color: whitesmoke;
|
||||
--current-color: #7928ca;
|
||||
}
|
||||
.markdown-callout-warning {
|
||||
--current-color: #f5a623;
|
||||
}
|
||||
.markdown-callout-caution {
|
||||
--current-color: #e60000;
|
||||
}
|
||||
|
||||
|
||||
.note-plus-preview {
|
||||
display:flex;
|
||||
flex-direction:column;
|
||||
align-items: flex-start;
|
||||
width:95%;
|
||||
margin-left: 20px;
|
||||
margin-top:20px;
|
||||
/*background-color: rgba(255,0,0,0.5)!important;*/
|
||||
}
|
||||
|
||||
/* allowed to be selected*/
|
||||
h1, h2, h3, h4, h5, h6,a, p, ul, ol, dl, blockquote,details,summary {
|
||||
pointer-events:auto;
|
||||
user-select:text;
|
||||
}
|
||||
|
||||
h1, h2, h3, h4, h5, h6 {
|
||||
display:inline-block;
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
font-weight: normal;
|
||||
}
|
||||
|
||||
p, ul, ol, dl, blockquote {
|
||||
margin: 0.3em;
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
ul, ol {
|
||||
padding-left: 1em;
|
||||
}
|
||||
|
||||
a {
|
||||
color: inherit;
|
||||
text-decoration: none;
|
||||
pointer-events: all;
|
||||
color: cyan;
|
||||
}
|
||||
|
||||
img {
|
||||
padding: 1em 0;
|
||||
max-width: 100%;
|
||||
}
|
||||
|
||||
iframe {
|
||||
max-width: 100%;
|
||||
height: auto;
|
||||
border:none;
|
||||
pointer-events:all;
|
||||
}
|
||||
|
||||
blockquote {
|
||||
border-left: 4px solid #ccc;
|
||||
padding-left: 1em;
|
||||
margin-left: 0;
|
||||
font-style: italic;
|
||||
}
|
||||
|
||||
pre, code {
|
||||
font-family: monospace;
|
||||
}
|
||||
|
||||
table {
|
||||
border-collapse: collapse;
|
||||
width: 100%;
|
||||
border-bottom: 1px solid #000;
|
||||
margin: 1em 0;
|
||||
}
|
||||
|
||||
th, td {
|
||||
border-left: 1px solid #000;
|
||||
border-right: 1px solid #000;
|
||||
padding: 8px;
|
||||
text-align: left;
|
||||
}
|
||||
|
||||
th {
|
||||
border: 1px solid #000;
|
||||
background-color: rgba(0,0,0,0.5);
|
||||
}
|
||||
|
||||
input[type="checkbox"] {
|
||||
margin-right: 10px;
|
||||
}
|
||||
`
|
||||
+393
-256
@@ -1,155 +1,122 @@
|
||||
/// <reference path="../types/typedefs.js" />
|
||||
|
||||
import { app } from '../../scripts/app.js'
|
||||
|
||||
import * as shared from './comfy_shared.js'
|
||||
import { infoLogger, successLogger, errorLogger } from './comfy_shared.js'
|
||||
import {
|
||||
DEFAULT_CSS,
|
||||
DEFAULT_HTML,
|
||||
DEFAULT_MD,
|
||||
DEFAULT_MODE,
|
||||
DEFAULT_THEME,
|
||||
THEMES,
|
||||
CSS_RESET,
|
||||
DEMO_CONTENT,
|
||||
} from './note_plus.constants.js'
|
||||
import { LocalStorageManager } from './comfy_shared.js'
|
||||
|
||||
const DEFAULT_CSS = ''
|
||||
const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
|
||||
Note+
|
||||
</p>`
|
||||
const DEFAULT_MD = '## Note+'
|
||||
const DEFAULT_MODE = 'markdown'
|
||||
const DEFAULT_THEME = 'one_dark'
|
||||
const storage = new LocalStorageManager('mtb')
|
||||
|
||||
const CSS_RESET = `
|
||||
* {
|
||||
font-family: monospace;
|
||||
line-height: 1.25em;
|
||||
/**
|
||||
* Uses `@mtb/markdown-parser` (a fork of marked)
|
||||
* It is statically stored to avoid having
|
||||
* more than 1 instance ever.
|
||||
* The size difference between both libraries...
|
||||
* ╭───┬────────────────────────────────┬──────────╮
|
||||
* │ # │ name │ size │
|
||||
* ├───┼────────────────────────────────┼──────────┤
|
||||
* │ 0 │ web-dist/mtb_markdown_plus.mjs │ 1.2 MB │ <- with shiki
|
||||
* │ 1 │ web-dist/mtb_markdown.mjs │ 44.7 KB │
|
||||
* ╰───┴────────────────────────────────┴──────────╯
|
||||
*/
|
||||
let useShiki = storage.get('np-use-shiki', false)
|
||||
|
||||
const makeResizable = (dialog) => {
|
||||
dialog.style.resize = 'both'
|
||||
dialog.style.transformOrigin = 'top left'
|
||||
dialog.style.overflow = 'auto'
|
||||
}
|
||||
|
||||
h1, h2, h3, h4, h5, h6 {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
font-weight: normal;
|
||||
const makeDraggable = (dialog, handle) => {
|
||||
let offsetX = 0
|
||||
let offsetY = 0
|
||||
let isDragging = false
|
||||
|
||||
const onMouseMove = (e) => {
|
||||
if (isDragging) {
|
||||
dialog.style.left = `${e.clientX - offsetX}px`
|
||||
dialog.style.top = `${e.clientY - offsetY}px`
|
||||
}
|
||||
}
|
||||
|
||||
const onMouseUp = () => {
|
||||
isDragging = false
|
||||
document.removeEventListener('mousemove', onMouseMove)
|
||||
document.removeEventListener('mouseup', onMouseUp)
|
||||
}
|
||||
|
||||
handle.addEventListener('mousedown', (e) => {
|
||||
isDragging = true
|
||||
offsetX = e.clientX - dialog.offsetLeft
|
||||
offsetY = e.clientY - dialog.offsetTop
|
||||
document.addEventListener('mousemove', onMouseMove)
|
||||
document.addEventListener('mouseup', onMouseUp)
|
||||
})
|
||||
}
|
||||
|
||||
p, ul, ol, dl, blockquote {
|
||||
margin: 0.3em;
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
|
||||
ul, ol {
|
||||
|
||||
padding-left: 1em;
|
||||
|
||||
}
|
||||
|
||||
a {
|
||||
color: inherit;
|
||||
text-decoration: none;
|
||||
pointer-events: all;
|
||||
color: cyan;
|
||||
}
|
||||
|
||||
img {
|
||||
padding: 1em 0;
|
||||
max-width: 100%;
|
||||
}
|
||||
|
||||
iframe {
|
||||
width: 100%;
|
||||
height: auto;
|
||||
border:none;
|
||||
pointer-events:all;
|
||||
}
|
||||
|
||||
blockquote {
|
||||
border-left: 4px solid #ccc;
|
||||
padding-left: 1em;
|
||||
margin-left: 0;
|
||||
font-style: italic;
|
||||
}
|
||||
|
||||
pre, code {
|
||||
font-family: monospace;
|
||||
}
|
||||
|
||||
table {
|
||||
border-collapse: collapse;
|
||||
width: 100%;
|
||||
border-bottom: 1px solid #000;
|
||||
margin: 1em 0;
|
||||
}
|
||||
|
||||
th, td {
|
||||
border-left: 1px solid #000;
|
||||
border-right: 1px solid #000;
|
||||
padding: 8px;
|
||||
text-align: left;
|
||||
}
|
||||
|
||||
th {
|
||||
border: 1px solid #000;
|
||||
|
||||
background-color: rgba(0,0,0,0.5);
|
||||
}
|
||||
|
||||
input[type="checkbox"] {
|
||||
margin-right: 10px;
|
||||
}
|
||||
|
||||
`
|
||||
|
||||
const themes = [
|
||||
'ambiance',
|
||||
'chaos',
|
||||
'chrome',
|
||||
'cloud9_day',
|
||||
'cloud9_night',
|
||||
'cloud9_night_low_color',
|
||||
'cloud_editor',
|
||||
'cloud_editor_dark',
|
||||
'clouds',
|
||||
'clouds_midnight',
|
||||
'cobalt',
|
||||
'crimson_editor',
|
||||
'dawn',
|
||||
'dracula',
|
||||
'dreamweaver',
|
||||
'eclipse',
|
||||
'github',
|
||||
'github_dark',
|
||||
'gob',
|
||||
'gruvbox',
|
||||
'gruvbox_dark_hard',
|
||||
'gruvbox_light_hard',
|
||||
'idle_fingers',
|
||||
'iplastic',
|
||||
'katzenmilch',
|
||||
'kr_theme',
|
||||
'kuroir',
|
||||
'merbivore',
|
||||
'merbivore_soft',
|
||||
'mono_industrial',
|
||||
'monokai',
|
||||
'nord_dark',
|
||||
'one_dark',
|
||||
'pastel_on_dark',
|
||||
'solarized_dark',
|
||||
'solarized_light',
|
||||
'sqlserver',
|
||||
'terminal',
|
||||
'textmate',
|
||||
'tomorrow',
|
||||
'tomorrow_night',
|
||||
'tomorrow_night_blue',
|
||||
'tomorrow_night_bright',
|
||||
'tomorrow_night_eighties',
|
||||
'twilight',
|
||||
'vibrant_ink',
|
||||
'vscode',
|
||||
]
|
||||
/** @extends {LGraphNode} */
|
||||
class NotePlus extends LiteGraph.LGraphNode {
|
||||
// same values as the comfy note
|
||||
color = LGraphCanvas.node_colors.yellow.color
|
||||
bgcolor = LGraphCanvas.node_colors.yellow.bgcolor
|
||||
groupcolor = LGraphCanvas.node_colors.yellow.groupcolor
|
||||
|
||||
/* NOTE: this is not serialized and only there to make multiple
|
||||
* note+ nodes in the same graph unique.
|
||||
*/
|
||||
uuid
|
||||
|
||||
/** Stores the dialog observer*/
|
||||
resizeObserver
|
||||
|
||||
/** Live update the preview*/
|
||||
live = true
|
||||
/** DOM height by adding child size together*/
|
||||
calculated_height = 0
|
||||
|
||||
/** ????*/
|
||||
_raw_html
|
||||
|
||||
/** might not be needed anymore */
|
||||
inner
|
||||
|
||||
/** the dialog DOM widget*/
|
||||
dialog
|
||||
|
||||
/** widgets*/
|
||||
|
||||
/** used to store the raw value and display the parsed html at the same time*/
|
||||
html_widget
|
||||
|
||||
/** hidden widgets for serialization*/
|
||||
css_widget
|
||||
edit_mode_widget
|
||||
theme_widget
|
||||
|
||||
editorsContainer
|
||||
/** ACE editors instances*/
|
||||
html_editor
|
||||
css_editor
|
||||
|
||||
constructor() {
|
||||
super()
|
||||
this.uuid = shared.makeUUID()
|
||||
|
||||
infoLogger('Constructing Note+ instance')
|
||||
shared.ensureMarkdownParser((_p) => {
|
||||
this.updateHTML()
|
||||
})
|
||||
// - litegraph settings
|
||||
this.collapsable = true
|
||||
this.isVirtualNode = true
|
||||
@@ -159,35 +126,30 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
// - default values, serialization is done through widgets
|
||||
this._raw_html = DEFAULT_MODE === 'html' ? DEFAULT_HTML : DEFAULT_MD
|
||||
|
||||
// - mardown converter
|
||||
this.markdownConverter = new showdown.Converter({
|
||||
tables: true,
|
||||
strikethrough: true,
|
||||
emoji: true,
|
||||
ghCodeBlocks: true,
|
||||
tasklists: true,
|
||||
ghMentions: true,
|
||||
smoothLivePreview: true,
|
||||
simplifiedAutoLink: true,
|
||||
parseImgDimensions: true,
|
||||
openLinksInNewWindow: true,
|
||||
})
|
||||
|
||||
// - state
|
||||
this.live = true
|
||||
this.calculated_height = 0
|
||||
|
||||
// - add widgets
|
||||
const inner = document.createElement('div')
|
||||
inner.style.margin = '0'
|
||||
inner.style.padding = '0'
|
||||
inner.style.pointerEvents = 'none'
|
||||
this.html_widget = this.addDOMWidget('HTML', 'html', inner, {
|
||||
const cinner = document.createElement('div')
|
||||
this.inner = document.createElement('div')
|
||||
|
||||
cinner.append(this.inner)
|
||||
this.inner.classList.add('note-plus-preview')
|
||||
cinner.style.margin = '0'
|
||||
cinner.style.padding = '0'
|
||||
this.html_widget = this.addDOMWidget('HTML', 'html', cinner, {
|
||||
setValue: (val) => {
|
||||
this._raw_html = val
|
||||
},
|
||||
getValue: () => this._raw_html,
|
||||
getMinHeight: () => this.calculated_height, // (the edit button),
|
||||
onDraw: () => {
|
||||
// HACK: dirty hack for now until it's addressed upstream...
|
||||
this.html_widget.element.style.pointerEvents = 'none'
|
||||
// NOTE: not sure about this, it avoid the visual "bugs" but scrolling over the wrong area will affect zoom...
|
||||
// this.html_widget.element.style.overflow = 'scroll'
|
||||
},
|
||||
hideOnZoom: false,
|
||||
})
|
||||
|
||||
@@ -197,22 +159,48 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
}
|
||||
|
||||
/**
|
||||
*
|
||||
* @param {CanvasRenderingContext2D} ctx
|
||||
* @param {LGraphCanvas} graphcanvas
|
||||
* @returns
|
||||
* @param {CanvasRenderingContext2D} ctx canvas context
|
||||
* @param {any} _graphcanvas
|
||||
*/
|
||||
|
||||
onDrawForeground(ctx, _graphcanvas) {
|
||||
if (this.flags.collapsed) return
|
||||
this.drawEditIcon(ctx)
|
||||
this.drawSideHandle(ctx)
|
||||
|
||||
// Define the size and position of the icon
|
||||
const iconSize = 14 // Size of the icon
|
||||
const iconMargin = 8 // Margin from the edges
|
||||
const x = this.size[0] - iconSize - iconMargin
|
||||
const y = iconMargin * 1.5
|
||||
// DEBUG BACKGROUND
|
||||
// ctx.fillStyle = 'rgba(0, 255, 0, 0.3)'
|
||||
// const rect = this.rect
|
||||
// ctx.fillRect(rect.x, rect.y, rect.width, rect.height)
|
||||
}
|
||||
drawSideHandle(ctx) {
|
||||
const handleRect = this.sideHandleRect
|
||||
const chamfer = 20
|
||||
ctx.beginPath()
|
||||
|
||||
// top left
|
||||
ctx.moveTo(handleRect.x, handleRect.y + chamfer)
|
||||
// top right
|
||||
ctx.lineTo(handleRect.x + handleRect.width, handleRect.y)
|
||||
|
||||
// bottom right
|
||||
ctx.lineTo(
|
||||
handleRect.x + handleRect.width,
|
||||
handleRect.y + handleRect.height,
|
||||
)
|
||||
// bottom left
|
||||
ctx.lineTo(handleRect.x, handleRect.y + handleRect.height - chamfer)
|
||||
ctx.closePath()
|
||||
|
||||
ctx.fillStyle = 'rgba(255, 255, 255, 0.05)'
|
||||
ctx.fill()
|
||||
}
|
||||
|
||||
drawEditIcon(ctx) {
|
||||
const rect = this.iconRect
|
||||
// DEBUG ICON POSITION
|
||||
// ctx.fillStyle = 'rgba(0, 255, 0, 0.3)'
|
||||
// ctx.fillRect(rect.x, rect.y, rect.width, rect.height)
|
||||
|
||||
// Create a new Path2D object from SVG path data
|
||||
const pencilPath = new Path2D(
|
||||
'M21.28 6.4l-9.54 9.54c-.95.95-3.77 1.39-4.4.76-.63-.63-.2-3.45.75-4.4l9.55-9.55a2.58 2.58 0 1 1 3.64 3.65z',
|
||||
)
|
||||
@@ -220,41 +208,73 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
'M11 4H6a4 4 0 0 0-4 4v10a4 4 0 0 0 4 4h11c2.21 0 3-1.8 3-4v-5',
|
||||
)
|
||||
|
||||
// Draw the paths
|
||||
ctx.save()
|
||||
ctx.translate(x, y) // Position the icon on the canvas
|
||||
ctx.scale(iconSize / 32, iconSize / 32) // Scale the icon to the desired size
|
||||
ctx.strokeStyle = 'rgba(255,255,255,0.3)'
|
||||
|
||||
ctx.translate(rect.x, rect.y)
|
||||
ctx.scale(rect.width / 32, rect.height / 32)
|
||||
ctx.strokeStyle = 'rgba(255,255,255,0.4)'
|
||||
ctx.lineCap = 'round'
|
||||
ctx.lineJoin = 'round'
|
||||
|
||||
ctx.lineWidth = 2.4
|
||||
ctx.stroke(pencilPath)
|
||||
ctx.stroke(folderPath)
|
||||
ctx.restore()
|
||||
}
|
||||
onMouseDown(_e, localPos, _graphcanvas) {
|
||||
// Check if the click is within the pencil icon bounds
|
||||
const iconSize = 14
|
||||
const iconMargin = 8
|
||||
const iconX = this.size[0] - iconSize - iconMargin
|
||||
const iconY = iconMargin * 1.5
|
||||
|
||||
if (
|
||||
localPos[0] > iconX &&
|
||||
localPos[0] < iconX + iconSize &&
|
||||
localPos[1] > iconY &&
|
||||
localPos[1] < iconY + iconSize
|
||||
) {
|
||||
// Pencil icon was clicked, open the editor
|
||||
this.openEditorDialog()
|
||||
return true // Return true to indicate the event was handled
|
||||
/**
|
||||
* @param {number} x
|
||||
* @param {number} y
|
||||
* @param {{x:number,y:number,width:number,height:number}} rect
|
||||
* @returns {}
|
||||
*/
|
||||
inRect(x, y, rect) {
|
||||
rect = rect || this.iconRect
|
||||
return (
|
||||
x >= rect.x &&
|
||||
x <= rect.x + rect.width &&
|
||||
y >= rect.y &&
|
||||
y <= rect.y + rect.height
|
||||
)
|
||||
}
|
||||
get rect() {
|
||||
return {
|
||||
x: 0,
|
||||
y: 0,
|
||||
width: this.size[0],
|
||||
height: this.size[1],
|
||||
}
|
||||
}
|
||||
get sideHandleRect() {
|
||||
const w = this.size[0]
|
||||
const h = this.size[1]
|
||||
|
||||
return false // Return false to let the event propagate
|
||||
const bw = 32
|
||||
const bho = 64
|
||||
|
||||
return {
|
||||
x: w - bw,
|
||||
y: bho,
|
||||
width: bw,
|
||||
height: h - bho * 1.5,
|
||||
}
|
||||
}
|
||||
get iconRect() {
|
||||
const iconSize = 32
|
||||
const iconMargin = 16
|
||||
return {
|
||||
x: this.size[0] - iconSize - iconMargin,
|
||||
y: iconMargin * 1.5,
|
||||
width: iconSize,
|
||||
height: iconSize,
|
||||
}
|
||||
}
|
||||
onMouseDown(_e, localPos, _graphcanvas) {
|
||||
if (this.inRect(localPos[0], localPos[1])) {
|
||||
this.openEditorDialog()
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
/* Hidden widgets to store note+ settings in the workflow (stripped in API)*/
|
||||
setupSerializationWidgets() {
|
||||
infoLogger('Setup Serializing widgets')
|
||||
|
||||
@@ -283,15 +303,36 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
shared.hideWidgetForGood(this, this.css_widget)
|
||||
shared.hideWidgetForGood(this, this.theme_widget)
|
||||
}
|
||||
|
||||
setupDialog() {
|
||||
infoLogger('Setup dialog')
|
||||
// this.addWidget('button', 'Edit', 'Edit', this.openEditorDialog.bind(this))
|
||||
|
||||
this.dialog = new app.ui.dialog.constructor()
|
||||
this.dialog.element.classList.add('comfy-settings')
|
||||
|
||||
Object.assign(this.dialog.element.style, {
|
||||
position: 'absolute',
|
||||
boxShadow: 'none',
|
||||
})
|
||||
|
||||
const subcontainer = this.dialog.textElement.parentElement
|
||||
|
||||
if (subcontainer) {
|
||||
Object.assign(subcontainer.style, {
|
||||
width: '100%',
|
||||
})
|
||||
}
|
||||
const closeButton = this.dialog.element.querySelector('button')
|
||||
closeButton.textContent = 'CANCEL'
|
||||
closeButton.id = 'cancel-editor-dialog'
|
||||
closeButton.title =
|
||||
"Cancel the changes since last opened (doesn't support live mode)"
|
||||
closeButton.disabled = this.live
|
||||
|
||||
closeButton.style.background = this.live
|
||||
? 'repeating-linear-gradient(45deg,#606dbc,#606dbc 10px,#465298 10px,#465298 20px)'
|
||||
: ''
|
||||
|
||||
const saveButton = document.createElement('button')
|
||||
saveButton.textContent = 'SAVE'
|
||||
saveButton.onclick = () => {
|
||||
@@ -313,32 +354,54 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
|
||||
closeEditorDialog(accept) {
|
||||
infoLogger('Closing editor dialog', accept)
|
||||
if (accept) {
|
||||
if (accept && !this.live) {
|
||||
this.updateHTML(this.html_editor.getValue())
|
||||
this.updateCSS(this.css_editor.getValue())
|
||||
}
|
||||
if (this.resizeObserver) {
|
||||
this.resizeObserver.disconnect()
|
||||
this.resizeObserver = null
|
||||
}
|
||||
this.teardownEditors()
|
||||
this.dialog.close()
|
||||
}
|
||||
|
||||
/**
|
||||
* @param {HTMLElement} elem
|
||||
*/
|
||||
hookResize(elem) {
|
||||
if (!this.resizeObserver) {
|
||||
const observer = () => {
|
||||
this.html_editor.resize()
|
||||
this.css_editor.resize()
|
||||
Object.assign(this.editorsContainer.style, {
|
||||
minHeight: `${(this.dialog.element.clientHeight / 100) * 50}px`, //'200px',
|
||||
})
|
||||
}
|
||||
this.resizeObserver = new ResizeObserver(observer).observe(elem)
|
||||
}
|
||||
}
|
||||
openEditorDialog() {
|
||||
infoLogger(`Current edit mode ${this.edit_mode_widget.value}`)
|
||||
this.hookResize(this.dialog.element)
|
||||
const container = document.createElement('div')
|
||||
|
||||
Object.assign(container.style, {
|
||||
display: 'flex',
|
||||
gap: '10px',
|
||||
flexDirection: 'column',
|
||||
})
|
||||
|
||||
const editorsContainer = document.createElement('div')
|
||||
Object.assign(editorsContainer.style, {
|
||||
this.editorsContainer = document.createElement('div')
|
||||
|
||||
Object.assign(this.editorsContainer.style, {
|
||||
display: 'flex',
|
||||
gap: '10px',
|
||||
flexDirection: 'row',
|
||||
minHeight: this.dialog.element.offsetHeight, //'200px',
|
||||
width: '100%',
|
||||
})
|
||||
|
||||
container.append(editorsContainer)
|
||||
container.append(this.editorsContainer)
|
||||
|
||||
this.dialog.show('')
|
||||
this.dialog.textElement.append(container)
|
||||
@@ -346,30 +409,39 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
const aceHTML = document.createElement('div')
|
||||
aceHTML.id = 'noteplus-html-editor'
|
||||
Object.assign(aceHTML.style, {
|
||||
width: '300px',
|
||||
height: '300px',
|
||||
// backgroundColor: 'rgb(30,30,30)',
|
||||
// color: 'whitesmoke',
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
|
||||
minWidth: '300px',
|
||||
minHeight: 'inherit',
|
||||
})
|
||||
|
||||
editorsContainer.append(aceHTML)
|
||||
this.editorsContainer.append(aceHTML)
|
||||
|
||||
const aceCSS = document.createElement('div')
|
||||
aceCSS.id = 'noteplus-css-editor'
|
||||
Object.assign(aceCSS.style, {
|
||||
width: '300px',
|
||||
height: '300px',
|
||||
// backgroundColor: 'rgb(30,30,30)',
|
||||
// color: 'whitesmoke',
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
minHeight: 'inherit',
|
||||
})
|
||||
|
||||
editorsContainer.append(aceCSS)
|
||||
this.editorsContainer.append(aceCSS)
|
||||
|
||||
const live_edit = document.createElement('input')
|
||||
live_edit.type = 'checkbox'
|
||||
live_edit.checked = this.live
|
||||
live_edit.onchange = () => {
|
||||
this.live = live_edit.checked
|
||||
const cancel_button = this.dialog.element.querySelector(
|
||||
'#cancel-editor-dialog',
|
||||
)
|
||||
if (cancel_button) {
|
||||
cancel_button.disabled = this.live
|
||||
cancel_button.style.background = this.live
|
||||
? 'repeating-linear-gradient(45deg,#606dbc,#606dbc 10px,#465298 10px,#465298 20px)'
|
||||
: ''
|
||||
}
|
||||
}
|
||||
|
||||
//- "Dynamic" elements
|
||||
@@ -388,15 +460,14 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
const md = this.html_editor.getValue()
|
||||
this.edit_mode_widget.value = 'html'
|
||||
select_mode.value = 'html'
|
||||
const html = this.markdownConverter.makeHtml(md)
|
||||
this.html_widget.value = html
|
||||
this.html_editor.setValue(html)
|
||||
this.html_editor.session.setMode('ace/mode/html')
|
||||
this.updateHTML(this.html_widget.value)
|
||||
|
||||
convert_to_html.remove()
|
||||
MTB.mdParser.parse(md).then((content) => {
|
||||
this.html_widget.value = content
|
||||
this.html_editor.setValue(content)
|
||||
this.html_editor.session.setMode('ace/mode/html')
|
||||
this.updateHTML(this.html_widget.value)
|
||||
convert_to_html.remove()
|
||||
})
|
||||
}
|
||||
|
||||
firstButton.before(convert_to_html)
|
||||
}
|
||||
} else {
|
||||
@@ -406,6 +477,19 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
}
|
||||
}
|
||||
select_mode.value = this.edit_mode_widget.value
|
||||
|
||||
// the header for dragging the dialog
|
||||
const header = document.createElement('div')
|
||||
header.style.padding = '8px'
|
||||
header.style.cursor = 'move'
|
||||
header.style.backgroundColor = 'rgba(0,0,0,0.5)'
|
||||
header.style.userSelect = 'none'
|
||||
|
||||
header.style.borderBottom = '1px solid #ddd'
|
||||
header.textContent = 'MTB Note+ Editor'
|
||||
container.prepend(header)
|
||||
makeDraggable(this.dialog.element, header)
|
||||
makeResizable(this.dialog.element)
|
||||
}
|
||||
//- combobox
|
||||
let theme_select = this.dialog.element.querySelector('#theme_select')
|
||||
@@ -421,7 +505,7 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
option.textContent = label
|
||||
theme_select.append(option)
|
||||
}
|
||||
for (const t of themes) {
|
||||
for (const t of THEMES) {
|
||||
addOption(t)
|
||||
}
|
||||
|
||||
@@ -491,54 +575,59 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
onCreate() {
|
||||
errorLogger('NotePlus onCreate')
|
||||
}
|
||||
configure(info) {
|
||||
super.configure(info)
|
||||
infoLogger('Restoring serialized values', info)
|
||||
// - update view from serialzed data
|
||||
restoreNodeState(info) {
|
||||
this.html_widget.element.id = `note-plus-${this.uuid}`
|
||||
this.setMode(this.edit_mode_widget.value)
|
||||
this.setTheme(this.theme_widget.value)
|
||||
this.updateHTML(this.html_widget.value)
|
||||
this.updateCSS(this.css_widget.value)
|
||||
this.setSize(info.size)
|
||||
if (info?.size) {
|
||||
this.setSize(info.size)
|
||||
}
|
||||
}
|
||||
configure(info) {
|
||||
super.configure(info)
|
||||
infoLogger('Restoring serialized values', info)
|
||||
this.restoreNodeState(info)
|
||||
// - update view from serialzed data
|
||||
}
|
||||
onNodeCreated() {
|
||||
infoLogger('Node created', this.uuid)
|
||||
this.html_widget.element.id = `note-plus-${this.uuid}`
|
||||
this.setMode(this.edit_mode_widget.value)
|
||||
this.setTheme(this.theme_widget.value)
|
||||
this.updateHTML(this.html_widget.value) // widget is populated here since we called super
|
||||
this.updateCSS(this.css_widget.value)
|
||||
}
|
||||
onRemoved() {
|
||||
infoLogger('Node removed', this.uuid)
|
||||
this.restoreNodeState({})
|
||||
// this.html_widget.element.id = `note-plus-${this.uuid}`
|
||||
// this.setMode(this.edit_mode_widget.value)
|
||||
// this.setTheme(this.theme_widget.value)
|
||||
// this.updateHTML(this.html_widget.value) // widget is populated here since we called super
|
||||
// this.updateCSS(this.css_widget.value)
|
||||
}
|
||||
// onRemoved() {
|
||||
// infoLogger('Node removed', this?.uuid)
|
||||
// }
|
||||
getExtraMenuOptions() {
|
||||
const options = []
|
||||
// {
|
||||
// content: string;
|
||||
// callback?: ContextMenuEventListener;
|
||||
// /** Used as innerHTML for extra child element */
|
||||
// title?: string;
|
||||
// disabled?: boolean;
|
||||
// has_submenu?: boolean;
|
||||
// submenu?: {
|
||||
// options: ContextMenuItem[];
|
||||
// } & IContextMenuOptions;
|
||||
// className?: string;
|
||||
// }
|
||||
options.push({
|
||||
content: `Set to ${
|
||||
this.edit_mode_widget.value === 'html' ? 'markdown' : 'html'
|
||||
}`,
|
||||
callback: () => {
|
||||
this.edit_mode_widget.value =
|
||||
this.edit_mode_widget.value === 'html' ? 'markdown' : 'html'
|
||||
this.updateHTML(this.html_widget.value)
|
||||
},
|
||||
})
|
||||
const currentMode = this.edit_mode_widget.value
|
||||
const newMode = currentMode === 'html' ? 'markdown' : 'html'
|
||||
|
||||
return options
|
||||
const debugItems = window.MTB?.DEBUG
|
||||
? [
|
||||
{
|
||||
content: 'Replace with demo content (debug)',
|
||||
callback: () => {
|
||||
this.html_widget.value = DEMO_CONTENT
|
||||
},
|
||||
},
|
||||
]
|
||||
: []
|
||||
|
||||
return [
|
||||
...debugItems,
|
||||
{
|
||||
content: `Set to ${newMode}`,
|
||||
callback: () => {
|
||||
this.edit_mode_widget.value = newMode
|
||||
this.updateHTML(this.html_widget.value)
|
||||
},
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
_setupEditor(editor) {
|
||||
@@ -663,17 +752,44 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
// this.setSize(this.computeSize())
|
||||
}
|
||||
|
||||
updateHTML(val) {
|
||||
const cleanHTML = DOMPurify.sanitize(val, { ADD_TAGS: ['iframe'] })
|
||||
this.html_widget.value = cleanHTML
|
||||
parserInitiated() {
|
||||
if (window.MTB?.mdParser) return true
|
||||
return false
|
||||
}
|
||||
|
||||
// update our widget preview
|
||||
if (this.edit_mode_widget.value === 'html') {
|
||||
this.html_widget.element.innerHTML = cleanHTML
|
||||
} else if (this.edit_mode_widget.value === 'markdown') {
|
||||
this.html_widget.element.innerHTML =
|
||||
this.markdownConverter.makeHtml(cleanHTML)
|
||||
/** to easilty swap purification methods*/
|
||||
purify(content) {
|
||||
return DOMPurify.sanitize(content, {
|
||||
ADD_TAGS: ['iframe', 'detail', 'summary'],
|
||||
})
|
||||
}
|
||||
|
||||
updateHTML(val) {
|
||||
if (!this.parserInitiated()) {
|
||||
return
|
||||
}
|
||||
val = val || this.html_widget.value
|
||||
const isHTML = this.edit_mode_widget.value === 'html'
|
||||
|
||||
const cleanHTML = this.purify(val)
|
||||
|
||||
const value = isHTML
|
||||
? cleanHTML
|
||||
: cleanHTML.replaceAll('>', '>').replaceAll('<', '<')
|
||||
// .replaceAll('&', '&')
|
||||
// .replaceAll('"', '"')
|
||||
// .replaceAll(''', "'")
|
||||
|
||||
this.html_widget.value = value
|
||||
|
||||
if (isHTML) {
|
||||
this.inner.innerHTML = value
|
||||
} else {
|
||||
MTB.mdParser.parse(value).then((e) => {
|
||||
this.inner.innerHTML = e
|
||||
})
|
||||
}
|
||||
// this.html_widget.element.innerHTML = `<div id="note-plus-spacer"></div>${value}`
|
||||
this.calculateHeight()
|
||||
// this.setSize(this.computeSize())
|
||||
}
|
||||
@@ -681,6 +797,27 @@ class NotePlus extends LiteGraph.LGraphNode {
|
||||
|
||||
app.registerExtension({
|
||||
name: 'mtb.noteplus',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.noteplus.use-shiki',
|
||||
category: ['mtb', 'Note+', 'use-shiki'],
|
||||
name: 'Use shiki to highlight code',
|
||||
tooltip:
|
||||
'This will load a larger version of @mtb/markdown-parser that bundles shiki, it supports all shiki transformers (supported langs: html,css,python,markdown)',
|
||||
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
attrs: {
|
||||
style: {
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('np-use-shiki', value)
|
||||
useShiki = value
|
||||
},
|
||||
})
|
||||
},
|
||||
|
||||
registerCustomNodes() {
|
||||
LiteGraph.registerNodeType('Note Plus (mtb)', NotePlus)
|
||||
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
-1
Submodule wiki updated: 4db733ae92...fa7fec28a3
Reference in New Issue
Block a user