Compare commits

...
262 Commits
Author SHA1 Message Date
shadowcz007 0846013378 增加 MiniCPM-V 2.6 int4 2024-08-22 14:46:11 +08:00
shadowcz007 c141ba405f fixbug: 自动监听文件夹 2024-08-20 15:06:58 +08:00
shadowcz007 0320f13a9f fixbug 2024-08-19 11:46:55 +08:00
shadowcz007 8adc34be4d fixbug 2024-08-19 10:29:54 +08:00
shadowcz007 cb6810d3c1 Update TextGenerateNode.py 2024-08-18 14:31:22 +08:00
shadowcz007 d384f64abf update :text-to-text 2024-08-17 22:15:31 +08:00
shadowcz007 ef7035f8ee update 2024-08-17 15:11:22 +08:00
shadowcz007 bfcadde5c3 update 2024-08-17 14:22:31 +08:00
shadowcz007 496ff41782 新ui支持,适配后,暂未全面测试 2024-08-16 18:13:05 +08:00
shadowcz007 46f0be5484 add image 2024-08-14 00:27:50 +08:00
shadowcz007 f3db0131c1 fixbug 2024-08-14 00:02:18 +08:00
shadowcz007 83a8d47f51 Update README.md 2024-08-13 23:32:10 +08:00
shadowcz007 fee0222910 v0.37.0 移动端适配、修改app模式的Mask编辑器 2024-08-12 10:10:43 +08:00
shadowcz007 1ed7b5511f mixlab app new mask editor 2024-08-12 00:19:26 +08:00
shadowcz007 b8f7c31537 Update index.html 2024-08-11 17:53:18 +08:00
shadowcz007 164791c257 webui 移动端适配 2024-08-11 17:20:36 +08:00
shadowcz007 8f5e599928 fixbug & ui 2024-08-11 16:29:16 +08:00
shadowcz007 7a7aaeb84d Update index.html 2024-08-10 23:00:58 +08:00
shadowcz007 e2136ab2fc fixbug 2024-08-10 17:13:28 +08:00
shadowcz007 c75cb21946 clean 2024-08-10 10:47:54 +08:00
shadowcz007 bf95218c91 p5-video-workflow 2024-08-10 00:59:44 +08:00
shadowcz007 8cb4507a5f v0.36.0 p5.js 2024-08-10 00:39:27 +08:00
shadowcz007 555890d1ba Update pyproject.toml 2024-08-09 19:08:36 +08:00
shadowcz007 e4f54e83b6 Update Text-to-Image-app.json 2024-08-09 16:30:47 +08:00
shadowcz007 692c4a709e fixbug:web app 2024-08-09 16:27:34 +08:00
shadowcz007 cbd1961459 test 2024-08-08 21:58:44 +08:00
shadowcz007 2e31a33ebf fixbug 2024-08-08 11:37:36 +08:00
shadowcz007 d16c6137d2 update 2024-08-06 23:07:20 +08:00
shadowcz007 0416ab79ec Update 3d_mixlab.js 2024-08-06 21:08:52 +08:00
shadow fc9a1c62b9 Merge pull request #295 from shadowcz007/0.36.0-py5-processing
Lama 改成手动安装,新增JsonRepair
2024-08-06 11:09:27 +08:00
shadowcz007 5d4567b134 Lama 改成手动安装,新增JsonRepair 2024-08-06 11:08:50 +08:00
shadow ae4a17d271 Merge pull request #293 from shadowcz007/0.36.0-py5-processing
0.36.0 py5 processing
2024-08-06 00:24:13 +08:00
shadowcz007 d110a08889 Update __init__.py 2024-08-06 00:23:34 +08:00
shadowcz007 e0157293cb Update P5.py 2024-08-06 00:21:51 +08:00
shadowcz007 0d985b3b65 update 2024-08-06 00:14:05 +08:00
shadowcz007 a65ade9fda updage 2024-08-05 21:30:30 +08:00
shadowcz007 874d6c8cb1 1 2024-08-05 21:16:19 +08:00
shadowcz007 f70ba2afa3 update 2024-08-05 21:08:32 +08:00
shadowcz007 e9f821e578 update 2024-08-05 20:49:06 +08:00
shadowcz007 8e488d4b1d update 2024-08-05 11:55:07 +08:00
shadowcz007 77201a457d 基本打通 2024-08-04 23:48:06 +08:00
shadowcz007 076e3b1178 test 2024-08-04 22:28:16 +08:00
shadowcz007 6b13fa64dc update 2024-08-04 20:44:56 +08:00
shadowcz007 846671a890 preview audio 2024-08-04 18:06:37 +08:00
shadowcz007 05b3088b75 0.35.1 2024-08-04 18:02:13 +08:00
shadowcz007 fe57286959 v0.34.0 2024-08-04 15:28:47 +08:00
shadowcz007 03645bbb33 image batch to list 2024-08-04 13:35:22 +08:00
shadowcz007 93dba9a399 fixbug :load image (base64) 2024-08-04 12:12:41 +08:00
shadowcz007 5627ea8073 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-08-04 09:40:43 +08:00
shadowcz007 7ba679c9ce fixbug 2024-08-04 09:40:40 +08:00
shadow c7a450e6ce Merge pull request #289 from ComfyNodePRs/licence-update
Update PyProject Toml - License
2024-08-03 17:50:25 +08:00
snomiao beda5156bf chore(licence-update): Update PyProject Toml - License 2024-08-02 23:03:55 +00:00
shadowcz007 76a9da7163 fixbug 2024-08-02 18:32:27 +08:00
shadowcz007 edd0303f59 App模式增加batch prompt,批量提示词,可以把动态提示词批量组成后运行 2024-08-01 21:12:58 +08:00
shadowcz007 be6f47a333 batch prompt :批量提示 2024-08-01 21:03:40 +08:00
shadowcz007 4cd6a072ca Update install.bat 2024-08-01 11:58:23 +08:00
shadowcz007 743a82efe9 fixbug 2024-07-29 18:29:16 +08:00
shadowcz007 9589f28ef7 v0.32.0 2024-07-29 18:11:57 +08:00
shadowcz007 35492c5671 add SiliconflowLLM 2024-07-29 18:06:32 +08:00
shadow db1e695bf3 Merge pull request #284 from cd0304/main
修正text image节点的padding问题
2024-07-29 17:51:29 +08:00
shadowcz007 ecc4aec43b Update ChatGPT.py 2024-07-29 15:17:00 +08:00
shadowcz007 fc063c2205 Update __init__.py 2024-07-29 14:17:57 +08:00
shadowcz007 4d60ce138a Update __init__.py 2024-07-28 21:12:39 +08:00
shadowcz007 2afd24f6e4 fixbug 2024-07-28 20:52:55 +08:00
shadowcz007 437acd023a fixbug 2024-07-28 20:28:34 +08:00
shadowcz007 b00523ae14 优化mixlab app,前端不传workflow,只传输入和输出 2024-07-28 20:21:53 +08:00
shadowcz007 4405a74993 Update Audio.py 2024-07-26 18:56:38 +08:00
cd0304 cb16090868 Update ImageNode.py 2024-07-26 13:04:17 +08:00
cd0304 396e510dce Update ImageNode.py
fix height
2024-07-26 00:32:56 +08:00
shadowcz007 3b9790b969 Update __init__.py 2024-07-25 13:39:41 +08:00
shadowcz007 a35d07a7ac video 2024-07-17 20:49:15 +08:00
shadowcz007 6d004c61fc Update pyproject.toml 2024-07-17 14:41:33 +08:00
shadowcz007 ffdd06da1b Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-07-17 14:41:02 +08:00
shadowcz007 f03f34cacb Update checkVersion_mixlab.js 2024-07-17 14:40:59 +08:00
shadow 0c86ea849e Merge pull request #273 from cd0304/main
textimge节点增加对otf后缀字体支持
2024-07-17 14:37:35 +08:00
cd0304 0efa4c38c0 Update ImageNode.py 2024-07-17 13:59:22 +08:00
cd0304 6092ab7793 Update ImageNode.py 2024-07-17 13:17:40 +08:00
shadowcz007 929def87eb Update ui_mixlab.js 2024-07-17 11:16:43 +08:00
shadowcz007 be074ccff7 Update __init__.py 2024-07-16 22:47:10 +08:00
shadowcz007 3445199393 AUDIO 2024-07-16 21:38:54 +08:00
shadowcz007 216c7e152e 0.30.3 2024-07-08 00:03:33 +08:00
shadowcz007 cc8bc10690 update 2024-07-07 18:41:56 +08:00
shadowcz007 69b4218d60 Update __init__.py 2024-07-07 17:05:06 +08:00
shadowcz007 1dd18dc4f8 fixbug 2024-07-06 20:30:50 +08:00
shadowcz007 4ccbd999d9 fixbug 2024-07-06 00:54:19 +08:00
shadowcz007 fa8d404964 0.30.2 2024-07-06 00:39:02 +08:00
shadowcz007 30086957c9 fixbug 2024-07-06 00:37:52 +08:00
shadowcz007 0e57c620c9 Update Video.py 2024-07-04 18:19:53 +08:00
shadowcz007 3ce1c59a2d Update README.md 2024-07-04 17:37:50 +08:00
shadowcz007 3337e20b9e Math Operation 2024-06-23 16:50:06 +08:00
shadowcz007 e816b3626e update 2024-06-22 21:44:33 +08:00
shadowcz007 3e0cb0f17a Update ui_mixlab.js 2024-06-22 18:42:12 +08:00
shadowcz007 41bc606217 Update 2-screeshare.json 2024-06-22 11:56:32 +08:00
shadowcz007 5a5f4ca49a Update pyproject.toml 2024-06-21 23:08:27 +08:00
shadowcz007 c3a8437cd1 Update ImageNode.py 2024-06-21 22:10:19 +08:00
shadowcz007 8d8a1a392d fixbug 2024-06-21 21:54:41 +08:00
shadowcz007 5f93fb5e55 增加支持的国产大模型 2024-06-21 17:40:15 +08:00
shadowcz007 d05050d7d8 v0.30.1 2024-06-20 20:39:43 +08:00
shadowcz007 8e9744100d 优化composite images节点 2024-06-20 17:46:46 +08:00
shadowcz007 1e4e7e287d Update ImageNode.py 2024-06-20 16:33:58 +08:00
shadowcz007 e8f0c73f08 优化text image节点,更为精准控制空白间距,字体修改为选择方式 2024-06-20 16:32:04 +08:00
shadowcz007 e923e28f8d Canvas Mode 2024-06-20 14:59:51 +08:00
shadowcz007 5cc75bfa7c Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-06-20 12:04:29 +08:00
shadowcz007 d6701769b8 fixbug:showtext 2024-06-20 12:04:23 +08:00
shadow 0ddc67bdab Create CNAME 2024-06-20 11:13:12 +08:00
shadowcz007 38b62b7a68 Update pyproject.toml 2024-06-19 11:10:27 +08:00
shadowcz007 7e726000c7 v0.30.0 2024-06-18 16:57:14 +08:00
shadowcz007 d8dfb292ec 增加 Edit Mask & SD3 示例 2024-06-18 16:55:59 +08:00
shadowcz007 826975241d Audio Play 2024-06-17 10:50:26 +08:00
shadowcz007 743637ceaf Update Video.py 2024-06-14 11:50:05 +08:00
shadowcz007 e350c7e31e CombineAudioVideo、LoadAndCombinedAudio 2024-06-14 11:38:10 +08:00
shadowcz007 66b1e0ab9f Update __init__.py 2024-06-14 08:17:04 +08:00
shadowcz007 7b0374d110 Update requirements.txt 2024-06-13 09:04:48 +08:00
shadowcz007 e86ef8cbb0 ImageBatchToList、LoadAndCombinedAudio、combine_audio_video、GenerateFramesByCount 2024-06-12 20:54:16 +08:00
shadowcz007 8c901c54bc Update extension-node-map.json 2024-06-08 17:41:40 +08:00
shadowcz007 408d85691e v0.29.0 支持把输出显示到comfyui背景(TouchDesigner 风格) 2024-06-08 16:58:21 +08:00
shadowcz007 c66cd6901b appinfo add performance features
Appinfo supports outputting to the background, enhancing the performance features of ComfyUI.
2024-06-08 16:03:10 +08:00
shadowcz007 aeadbc4f6d fixbug 2024-06-06 15:17:40 +08:00
shadowcz007 224136890e fixbug 2024-06-06 08:02:51 +08:00
shadowcz007 3669a1e86d 0.28.3 2024-06-01 23:21:33 +08:00
shadowcz007 d588b5b327 Update index.html 2024-05-29 22:37:52 +08:00
shadowcz007 b705679098 Update index.html 2024-05-29 21:49:32 +08:00
shadowcz007 f71a0b0da5 Update index.html 2024-05-29 20:20:19 +08:00
shadowcz007 ebc2c76b6b fixbug 2024-05-25 22:50:19 +08:00
shadow 2e3fff278e Merge pull request #240 from audioscavenger/patch-1
Update extension-node-map.json
2024-05-24 11:12:17 +08:00
Eric 1f4bc5e089 Update extension-node-map.json
i'm the new maintainer, thanks
2024-05-23 16:41:33 -07:00
shadowcz007 52c38b10dd v0.28.2 2024-05-23 18:19:30 +08:00
shadowcz007 7047aa5456 add video format 2024-05-23 16:59:04 +08:00
shadowcz007 33fe4019f7 Update ui_mixlab.js 2024-05-23 16:43:50 +08:00
shadowcz007 80b9d97690 Update Video.py 2024-05-23 15:58:24 +08:00
shadowcz007 3c3c92723f Update pyproject.toml 2024-05-23 10:34:36 +08:00
shadowcz007 037bd87006 Update pyproject.toml 2024-05-23 10:26:48 +08:00
shadowcz007 f688310d28 Update Utils.py 2024-05-23 10:14:25 +08:00
shadow c4d65e7a45 Merge pull request #234 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-05-22 23:14:33 +08:00
shadow 6f208b710d Merge pull request #235 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-05-22 23:14:17 +08:00
haohaocreates b599faaf85 Update pyproject.toml desc 2024-05-21 15:24:07 -04:00
haohaocreates 6d991d20dc chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-21 19:19:01 +00:00
haohaocreates c87e0296f6 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-21 19:19:01 +00:00
shadow 16cdb4c5b4 Merge pull request #231 from 295958090/main
修复FloatSlider的bug
2024-05-21 21:57:10 +08:00
Bai Shui 7631b8924d 修复bug 2024-05-21 13:34:29 +08:00
shadowcz007 785d307ff3 Update index.html 2024-05-18 16:13:30 +08:00
shadowcz007 8c713ff35e Update index.html 2024-05-18 16:08:25 +08:00
shadowcz007 7d80493bef Update index.html 2024-05-18 16:05:39 +08:00
shadowcz007 bff2760c3d v0.28.1
修复bug
2024-05-18 11:38:50 +08:00
shadowcz007 a0f8848367 修复 当上传新的图片,编辑mask的bug 2024-05-18 11:38:28 +08:00
shadowcz007 5b1cbcd8d5 修复bug 2024-05-16 13:20:06 +08:00
shadowcz007 05857a92d5 v0.28.0
add rembg api & webapp rembg
2024-05-16 11:49:45 +08:00
shadowcz007 6bdc811286 add rembg api & webapp rembg 2024-05-16 11:49:14 +08:00
shadowcz007 469d50a5b8 Update index.html 2024-05-16 09:02:13 +08:00
shadowcz007 ef86904bfb Update ui_mixlab.js 2024-05-16 09:02:08 +08:00
shadowcz007 d4181ea67c v0.27.1 fixbug 2024-05-16 08:49:30 +08:00
shadowcz007 1c6d17309f Update index.html 2024-05-16 08:49:09 +08:00
shadowcz007 db293ec41d fixbug css 2024-05-16 08:47:17 +08:00
shadowcz007 db8d468f29 0.27.0 增加webapp的mask绘制 2024-05-16 00:11:16 +08:00
shadowcz007 d7d7af7265 add mask edit for webapp 2024-05-16 00:06:42 +08:00
shadowcz007 bd763cadc1 fixbug 2024-05-16 00:05:24 +08:00
shadowcz007 22799fc549 fixbug for mask 2024-05-16 00:05:16 +08:00
shadowcz007 0f231d1271 add minPaint for mask 2024-05-16 00:04:54 +08:00
shadowcz007 8c0c911020 Create LICENSE 2024-05-14 10:02:46 +08:00
shadowcz007 6cb9df700b 增加ComparingTwoFrames、右键image-to-text 2024-05-14 09:59:39 +08:00
shadowcz007 4fdda537b9 Update README.md 2024-05-14 09:57:51 +08:00
shadowcz007 cb6f32465a Update README.md 2024-05-14 09:55:34 +08:00
shadowcz007 c235e36cb4 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-05-14 09:54:27 +08:00
shadowcz007 aaca440a94 Update README.md 2024-05-14 09:54:24 +08:00
shadow 1f57950a29 Update README.md 2024-05-14 09:39:57 +08:00
shadow 3d8855ec72 Update README.md 2024-05-14 09:39:45 +08:00
shadowcz007 ac9231d9f3 Update image_mixlab.js 2024-05-13 21:43:20 +08:00
shadowcz007 37c0a56d89 add ComparingTwoFrames 2024-05-13 21:33:44 +08:00
shadowcz007 010915dac4 add help 2024-05-13 11:37:17 +08:00
shadowcz007 83a8b3b970 Update ui_mixlab.js 2024-05-13 10:58:36 +08:00
shadowcz007 f130202aa1 resizeImage 2024-05-12 22:25:08 +08:00
shadowcz007 c8f6800bcd add image-to-text :llava-phi-3-mini-gguf 2024-05-12 17:34:53 +08:00
shadowcz007 078f9f5dd4 Update ui_mixlab.js 2024-05-12 00:03:11 +08:00
shadowcz007 c0de178c7d add re_start 2024-05-11 23:58:05 +08:00
shadowcz007 1c767b538d set n_gpu_layers 2024-05-11 17:53:18 +08:00
shadowcz007 51aab44b5d Update ui_mixlab.js 2024-05-11 14:43:35 +08:00
shadowcz007 b8a0d4a67b Update ui_mixlab.js 2024-05-11 14:17:01 +08:00
shadowcz007 cd0dcfbb8c v0.25.1 2024-05-11 14:12:00 +08:00
shadow fcc9e30eae Update __init__.py 2024-05-11 12:55:14 +08:00
shadowcz007 5f66218a43 修复 sys.stdout.isatty() object has no attribute 'isatty' 2024-05-11 12:28:14 +08:00
shadowcz007 61ef4f9a0f Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-05-11 12:13:34 +08:00
shadowcz007 f0a8734b42 修复 sys.stdout.isatty() object has no attribute 'isatty' 2024-05-11 12:13:31 +08:00
shadow 4fe95ef4ec Update README.md 2024-05-11 08:53:33 +08:00
shadowcz007 2a5148845b yaml 2024-05-10 14:38:32 +08:00
shadowcz007 8fa562caaf Update install.bat 2024-05-09 09:55:34 +08:00
shadowcz007 cab5620cd5 Update index.html 2024-05-08 23:42:48 +08:00
shadowcz007 be38d36677 Update index.html 2024-05-08 23:32:20 +08:00
shadowcz007 69236fca89 Update __init__.py 2024-05-08 22:55:42 +08:00
shadowcz007 4a4f376bfd Update __init__.py 2024-05-08 22:53:39 +08:00
shadowcz007 fd9718fe24 Update __init__.py 2024-05-08 22:50:30 +08:00
shadowcz007 26a6e11212 llama_cpp 2024-05-08 22:42:34 +08:00
shadowcz007 de1a669f6e Update README.md 2024-05-08 10:11:16 +08:00
shadowcz007 e482c9e5c4 0.25.0 text-to-text for prompt 2024-05-07 21:26:47 +08:00
shadowcz007 4e96a77a41 Update README.md 2024-05-07 21:12:27 +08:00
shadowcz007 3346290e5c Update ImageNode.py 2024-05-07 20:59:22 +08:00
shadowcz007 dcac593efe Update install.bat 2024-05-07 12:49:03 +08:00
shadowcz007 7248d0de02 update 2024-05-07 12:14:42 +08:00
shadowcz007 8ad3ce632c Update index.html 2024-05-07 00:08:35 +08:00
shadowcz007 9398b02562 Update ImageNode.py 2024-05-06 22:13:09 +08:00
shadowcz007 164e4da99d Update ImageNode.py 2024-05-05 18:04:00 +08:00
shadowcz007 0025ea6119 output defaultImage 2024-05-05 13:16:14 +08:00
shadowcz007 eef53a5165 Update ImageNode.py 2024-05-05 12:59:21 +08:00
shadowcz007 9ea066d948 composite_images add position 2024-05-05 12:20:57 +08:00
shadowcz007 4bd900c4a1 Update __init__.py 2024-05-04 09:44:01 +08:00
shadowcz007 736cd2bebd 兼容旧版comfyui 2024-05-04 09:37:25 +08:00
shadowcz007 cd658c2a60 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-05-03 22:08:43 +08:00
shadowcz007 9730658f21 Update __init__.py 2024-05-03 22:08:40 +08:00
shadow 3fa107acb1 Update ui_mixlab.js 2024-05-03 18:35:39 +08:00
shadowcz007 7da22179a0 text-to-text 2024-05-03 00:30:16 +08:00
shadowcz007 fc7b71ee78 Update ui_mixlab.js 2024-05-02 19:02:33 +08:00
shadowcz007 b79573bf1f Update __init__.py 2024-05-02 19:02:28 +08:00
shadowcz007 a01db6f7c0 Update README.md 2024-05-02 12:09:47 +08:00
shadowcz007 117d58c58e v0.24.0 2024-05-02 12:06:19 +08:00
shadowcz007 ff6626ed89 add llama.cpp & Local LLM -Phi-3 & llama3
Phi-3
llama3
2024-05-02 12:02:47 +08:00
shadowcz007 10c798a440 Update index.html 2024-05-02 10:54:20 +08:00
shadowcz007 9097d87819 Update index.html 2024-05-02 10:46:09 +08:00
shadowcz007 d901f503d1 fixbug 2024-05-02 10:06:23 +08:00
shadowcz007 65f2b6ce6f update https 2024-05-02 09:07:28 +08:00
shadowcz007 d949fe8bf1 Update __init__.py 2024-05-01 22:26:46 +08:00
shadowcz007 15cfb48550 Update __init__.py 2024-05-01 20:52:50 +08:00
shadowcz007 447dc6d4c3 fixbug 2024-05-01 15:55:43 +08:00
shadowcz007 c92e43b920 Update checkVersion_mixlab.js 2024-04-29 00:16:55 +08:00
shadowcz007 be6c32b0e0 fixbug 2024-04-29 00:15:47 +08:00
shadowcz007 3d2062e810 add TripoSRModel 2024-04-29 00:15:32 +08:00
shadowcz007 dd816e95cd Update README.md 2024-04-27 23:43:40 +08:00
shadowcz007 6d1b51890d Update checkVersion_mixlab.js 2024-04-27 23:38:00 +08:00
shadowcz007 11f03ec99a 支持正片叠底 2024-04-25 21:23:40 +08:00
shadowcz007 bd192f43e7 优化 2024-04-24 23:10:03 +08:00
shadowcz007 36e4b11983 gridout can export mask 2024-04-23 12:14:50 +08:00
shadowcz007 97397ba8c2 Update ImageNode.py 2024-04-22 12:32:48 +08:00
shadowcz007 052eee4111 Update __init__.py 2024-04-22 12:31:40 +08:00
shadowcz007 e319496044 v0.22.0
- 优化ImageColorTransfer
- 支持动态提示
- 添加更多节点支持,AppInfo自动填充id
- LoadImagesToBatch加载
- zhipuai 按需安装
2024-04-22 09:52:45 +08:00
shadowcz007 b83b63c362 添加支持的节点自动填充id 2024-04-21 21:51:07 +08:00
shadowcz007 4d6b1675bb Load Images to Batch加载文件最大宽度1024 2024-04-21 21:50:56 +08:00
shadowcz007 13110fab39 支持动态提示 2024-04-21 20:58:44 +08:00
shadowcz007 74ea509848 fixbug 2024-04-20 23:05:59 +08:00
shadowcz007 192bff9d2c optimize ImageColorTransfer and support batching, 2024-04-20 22:23:34 +08:00
shadowcz007 44ed8812dc output add TransparentImage 2024-04-19 22:27:53 +08:00
shadowcz007 200696ba21 Update ui_mixlab.js 2024-04-19 20:44:27 +08:00
shadowcz007 51cf3b0c04 install Zhipuai as needed 2024-04-19 11:41:56 +08:00
shadowcz007 6a4831c83b add nodes map for appinfo 2024-04-18 18:24:58 +08:00
shadowcz007 c5e7ed95a3 Update index.html 2024-04-18 16:10:02 +08:00
shadowcz007 42a97fa4d9 Update image_mixlab.js 2024-04-18 16:00:43 +08:00
shadowcz007 45240d0012 Update image_mixlab.js 2024-04-18 16:00:18 +08:00
shadowcz007 8ed085febd Update image_mixlab.js 2024-04-18 15:57:34 +08:00
shadowcz007 37803ea61b Update smart_connect.js 2024-04-18 09:06:42 +08:00
shadowcz007 acd416952c Update image_mixlab.js 2024-04-18 07:48:50 +08:00
shadowcz007 6ec46cbc44 Update ui_mixlab.js 2024-04-17 22:55:16 +08:00
shadowcz007 a9d971e476 Update __init__.py 2024-04-17 22:55:12 +08:00
shadowcz007 9fe064675d paste appinfo data / 支持粘贴appinfo导出的数据 2024-04-17 16:28:32 +08:00
shadowcz007 c84fa467d0 Update requirements.txt 2024-04-17 15:36:49 +08:00
shadowcz007 3d7a55f6d3 SaveImageAndMetadata支持格式化文件名@bakkhos8 2024-04-17 15:01:05 +08:00
shadowcz007 d49baa1540 Update checkVersion_mixlab.js 2024-04-15 23:40:01 +08:00
shadowcz007 5c686af842 Update ImageNode.py 2024-04-15 09:02:49 +08:00
shadowcz007 22425b5bc6 Update index.html 2024-04-13 00:12:05 +08:00
shadowcz007 b6d9b338d2 Update index.html 2024-04-12 21:55:56 +08:00
shadowcz007 a191a13751 Update index.html 2024-04-12 21:48:31 +08:00
shadowcz007 3eccdbcc9b Update index.html 2024-04-12 21:43:21 +08:00
shadowcz007 50063903f9 mixlab app add 3D 2024-04-12 16:10:05 +08:00
shadowcz007 95a1b70533 Update app_mixlab.js 2024-04-08 17:14:30 +08:00
shadowcz007 c3679ac90b Update image_mixlab.js 2024-04-07 11:49:26 +08:00
shadowcz007 29e48eb6a2 Update ui_mixlab.js 2024-04-07 10:32:13 +08:00
124 changed files with 142309 additions and 6060 deletions
+21
View File
@@ -0,0 +1,21 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+1
View File
@@ -0,0 +1 @@
mixlabnodes.com
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 shadow
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+163 -61
View File
@@ -1,30 +1,79 @@
> 适配了最新版comfyui的py3.11 ,torch 2.1.2+cu121
![](https://img.shields.io/github/release/shadowcz007/comfyui-mixlab-nodes)
> 适配了最新版 comfyui 的 py3.11 ,torch 2.3.1+cu121
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
####
商务合作请联系 389570357@qq.com
For business cooperation, please contact email 389570357@qq.com
##### `最新`:
- 增加 MiniCPM-V 2.6 int4
This is the int4 quantized version of MiniCPM-V 2.6.
Running with int4 version would use lower GPU memory (about 7GB).
- 移动端适配、修改 app 模式的 Mask 编辑器
- 增加 p5.js 作为输入节点
[workflow](./workflow/p5workflow.json)
[workflow2](./workflow/p5-video-workflow.json)
- App 模式增加 batch prompt,批量提示词,可以把动态提示词批量组成后运行
![alt text](./assets/1722517810720.png)
- 增加 API Key Input 节点,用于管理 LLM 的 Key,同时优化 LLM 相关节点,为后续 agent 模式做准备
- 增加 SiliconflowLLM,可以使用由 Siliconflow 提供的免费 LLM
<!-- - ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/` -->
<!-- - 右键菜单支持 text-to-text,方便对 prompt 词补全 -->
<!--
强烈推荐:
[Phi-3-mini-4k-instruct-function-calling-GGUF](https://huggingface.co/nold/Phi-3-mini-4k-instruct-function-calling-GGUF)
[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main),备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
- 右键菜单支持 image-to-text,使用多模态模型,多模态使用 [llava-phi-3-mini-gguf](https://huggingface.co/xtuner/llava-phi-3-mini-gguf/tree/main),注意需要把llava-phi-3-mini-mmproj-f16.gguf也下载
![](./assets/prompt_ai_setup.png)
![](./assets/prompt-ai.png) -->
#### `相关插件推荐`
[comfyui-liveportrait](https://github.com/shadowcz007/comfyui-liveportrait)
[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS)
[comfyui-sound-lab](https://github.com/shadowcz007/comfyui-sound-lab)
[comfyui-Image-reward](https://github.com/shadowcz007/comfyui-Image-reward)
[comfyui-ultralytics-yolo](https://github.com/shadowcz007/comfyui-ultralytics-yolo)
[comfyui-moondream](https://github.com/shadowcz007/comfyui-moondream)
[comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg)
<!-- [comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg) -->
## 🚀🚗🚚🏃 Workflow-to-APP
## 🚀🚗🚚🏃 Workflow-to-APP
- 新增AppInfo节点,可以通过简单的配置,把workflow转变为一个Web APP。
- 支持多个web app 切换
- 发布为app的workflow,可以在右键里再次编辑了
- web app可以设置分类,在comfyui右键菜单可以编辑更新web app
- 新增 AppInfo 节点,可以通过简单的配置,把 workflow 转变为一个 Web APP。
- 支持多个 web app 切换
- 发布为 app 的 workflow,可以在右键里再次编辑了
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
- 支持动态提示
- 支持把输出显示到 comfyui 背景(TouchDesigner 风格)
- 如果转为 web app 打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是 zip 包下载会多了个-main 的后缀,需要去掉)
![](./assets/微信图片_20240421205440.png)
- Support multiple web app switching.
- Add the AppInfo node, which allows you to transform the workflow into a web app by simple configuration.
- The workflow, which is now released as an app, can also be edited again by right-clicking.
- The web app can be configured with categories, and the web app can be edited and updated in the right-click menu of ComfyUI.
![](./assets/0-m-app.png)
![](./assets/appinfo-readme.png)
@@ -32,59 +81,96 @@
![](./assets/appinfo-2.png)
Example:
- workflow
![APP info](./workflow/appinfo-workflow.svg)
[text-to-image](./workflow/Text-to-Image-app.json)
![APP info](./workflow/appinfo-workflow.svg)
[text-to-image](./workflow/Text-to-Image-app.json)
APP-JSON:
- [text-to-image](./example/Text-to-Image_3.json)
- [image-to-image](./example/Image-to-Image_2.json)
- text-to-text
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS_LoadVideo、CLIPTextEncode、PromptSlide、TextInput_、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS*LoadVideo、CLIPTextEncode、PromptSlide、TextInput*、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT、VHS_VideoCombine、PromptImage
> seed统一输入控件,支持:SamplerCustom、KSampler
> seed 统一输入控件,支持:SamplerCustom、KSampler
> 配套[ps插件](https://github.com/shadowcz007/comfyui-ps-plugin)
> 配套[ps 插件](https://github.com/shadowcz007/comfyui-ps-plugin)
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用https,端口8189这个服务( 感谢 @Damien 反馈问题)
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用 https,端口 8189 这个服务( 感谢 @Damien 反馈问题)
> If you encounter difficulties in uploading images, please check the following: for local network or cloud services, please use HTTPS and the service on port 8189. (Thanks to @Damien for reporting the issue.)
## 🏃🚗🚚🚀 Real-time Design
## 🏃🚗🚚🚀 Real-time Design
> ScreenShareNode & FloatingVideoNode. Now comfyui supports capturing screen pixel streams from any software and can be used for LCM-Lora integration. Let's get started with implementation and design! 💻🌐
![screenshare](./assets/screenshare.png)
https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43e-410a-ab3a-1952b7b4e7da
<!-- [ScreenShareNode](./workflow/2-screeshare.json) -->
[ScreenShareNode & FloatingVideoNode](./workflow/3-FloatVideo-workflow.json)
!! Please use the address with HTTPS (https://127.0.0.1).
### SpeechRecognition & SpeechSynthesis
![f](./assets/audio-workflow.svg)
[Voice + Real-time Face Swap Workflow](./workflow/语音+实时换脸workflow.json)
- Preview Audio
[text-to-audio](./workflow/text-to-audio-base-workflow.json)
### GPT
> Support for calling multiple GPTs.ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
![gpt-workflow.svg](./assets/gpt-workflow.svg)
> Support for calling multiple GPTs.Local LLM 、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
[workflow-5](./workflow/5-gpt-workflow.json)
[LLM_base_workflow](./workflow/LLM_base_workflow.json)
- SiliconflowLLM
- ChatGPTOpenAI
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
Model download,move to :`models/llamafile/`
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main)
备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
> 如果碰到安装失败,可以尝试手动安装
```
../../../python_embeded/python.exe -s -m pip install llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
../../../python_embeded/python.exe -s -m pip install llama-cpp-python[server]
```
> [Mac](https://llama-cpp-python.readthedocs.io/en/latest/install/macos/)
```
pip uninstall llama-cpp-python -y
CMAKE_ARGS="-DLLAMA_METAL=on" pip install -U llama-cpp-python --no-cache-dir
pip install 'llama-cpp-python[server]'
```
```
pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
``` -->
## Prompt
> PromptSlide
![](./assets/prompt_weight.png)
> ![](./assets/prompt_weight.png)
<!-- ![](./workflow/promptslide-appinfo-workflow.svg) -->
@@ -98,95 +184,114 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
> PromptImage & PromptSimplification,Assist in simplifying prompt words, comparing images and prompt word nodes.
> ChinesePrompt && PromptGenerate,中文prompt节点,直接用中文书写你的prompt
> ChinesePrompt && PromptGenerate,中文 prompt 节点,直接用中文书写你的 prompt
![](./assets/ChinesePrompt_workflow.svg)
### Layers
> A new layer class node has been added, allowing you to separate the image into layers. After merging the images, you can input the controlnet for further processing.
> The composite images node overlays a foreground image onto a background image at specified positions and scales, with optional blending modes and masking capabilities. position : 'overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"
![layers](./assets/layers-workflow.svg)
![poster](./assets/poster-workflow.svg)
### 3D
![](./assets/3d-workflow.png)
![](./assets/3d_app.png)
[workflow](./assets/Image-to-3D_1.json)
![](./assets/3dimage.png)
[workflow](./workflow/3D-workflow.json)
### Image
#### LoadImagesToBatch
> Upload multiple images for batch input into the IP adapter.
> Upload multiple images for batch input into the IP adapter.
#### LoadImagesFromLocal
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop.
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop.
![watch](./assets/4-loadfromlocal-watcher-workflow.svg)
[workflow-4](./workflow/4-loadfromlocal-watcher-workflow.json)
#### LoadImagesFromURL
> Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed.
#### TextImage
> [下载字体](https://drxie.github.io/OSFCC/)放到 `custom_nodes/comfyui-mixlab-nodes/assets/fonts`
#### MiniCPM-VQA Simple
This is the int4 quantized version of MiniCPM-V 2.6.
Running with int4 version would use lower GPU memory (about 7GB).
[模型](https://huggingface.co/openbmb/MiniCPM-V-2_6-int4)
![alt text](assets/1724308322276.png)
### Style
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
![](./assets/VisualStylePrompting.png)
> StyleAligned , Modified from [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
> StyleAligned , Modified from [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
### Utils
> The Color node provides a color picker for easy color selection, the Font node offers built-in font selection for use with TextImage to generate text images, and the DynamicDelayByText node allows delayed execution based on the length of the input text.
- [添加了DynamicDelayByText功能,可以根据输入文本的长度进行延迟执行。](./workflow/audio-chatgpt-workflow.json)
- [添加了 DynamicDelayByText 功能,可以根据输入文本的长度进行延迟执行。](./workflow/audio-chatgpt-workflow.json)
- [Added DynamicDelayByText, enabling delayed execution based on input text length.](./workflow/audio-chatgpt-workflow.json)
- [使用CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
- [使用 CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
- [CkptNames compare the effects of different models.](./workflow/ckpts-image-workflow.json)
### Other Nodes
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
![main](./assets/all-workflow.svg)
![main2](./assets/detect-face-all.png)
[workflow-1](./workflow/1-workflow.json)
> TransparentImage
![TransparentImage](./assets/TransparentImage.png)
> FeatheredMask、SmoothMask
Add edges to an image.
![FeatheredMask](./assets/FlVou_Y6kaGWYoEj1Tn0aTd4AjMI.jpg)
> LaMaInpainting(需要手动安装)
> LaMaInpainting
- simple-lama-inpainting 里的 pillow 造成冲突,暂时从依赖里移除,如果有安装 simple-lama-inpainting ,节点会自动添加,没有,则不会自动添加。
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
- [问题汇总](https://github.com/shadowcz007/comfyui-mixlab-nodes/issues/294)
> rembgNode
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
*** briarmbg *** model was developed by BRlA Al and can be used as an open-source model for non-commercial purposes
**_ briarmbg _** model was developed by BRlA Al and can be used as an open-source model for non-commercial purposes
### Improvement
### Improvement
- Add "help" option to the context menu for each node.
- Add "Nodes Map" option to the global context menu.
@@ -197,18 +302,21 @@ An improvement has been made to directly redirect to GitHub to search for missin
![node-not-found](./assets/node-not-found.png)
### Models
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:models/rembg
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : models/lama
- [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in `models/triposr/facebook/dino-vitb16`
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to : models/clip_interrogator/Salesforce/blip-image-captioning-base
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:`models/rembg`
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:prompt_generator/text2image-prompt-generator
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : `models/lama`
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:prompt_generator/opus-mt-zh-en
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to :`models/clip_interrogator/Salesforce/blip-image-captioning-base`
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:`models/prompt_generator/text2image-prompt-generator`
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:`models/prompt_generator/opus-mt-zh-en`
## Installation
@@ -224,40 +332,35 @@ git clone https://github.com/shadowcz007/comfyui-mixlab-nodes.git
Install the requirements:
run directly:
```
cd ComfyUI/custom_nodes/comfyui-mixlab-nodes
install.bat
```
or install the requirements using:
```
../../../python_embeded/python.exe -s -m pip install -r requirements.txt
```
If you are using a venv, make sure you have it activated before installation and use:
```
pip3 install -r requirements.txt
```
#### Chinese community
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab无界社区
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab 无界社区
####
####
File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
#### discussions:
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
<picture>
<source
@@ -277,4 +380,3 @@ File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
src="https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date"
/>
</picture>
+760 -137
View File
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 537 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 340 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 135 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 210 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.2 MiB

File diff suppressed because one or more lines are too long
Binary file not shown.
Binary file not shown.

After

Width:  |  Height:  |  Size: 75 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 63 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 965 KiB

+12876 -554
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+6
View File
@@ -10,6 +10,12 @@ if exist "%python_exec%" (
for /f "delims=" %%i in (%requirements_txt%) do (
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
)
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
) else (
echo Installing with system Python
for /f "delims=" %%i in (%requirements_txt%) do (
+57 -37
View File
@@ -1,6 +1,7 @@
import os
import folder_paths
import torchaudio
class SpeechRecognition:
@classmethod
@@ -55,46 +56,65 @@ class SpeechSynthesis:
return {"ui": {"text": text}, "result": (text,)}
#
class GamePal:
class AudioPlayNode:
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append = ""
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input_text": ("STRING",{"multiline": True,"default": ""}),
},
"optional": {
"input_num": ("INT",{
"default":100,
"min": -1, #Minimum value
"max": 0xffffffffffffffff, #Maximum value
"step": 1, #Slider's step
"display": "slider" # Cosmetic only: display as "number" or "slider"
}),
"python_code": ("STRING",{"multiline": True,"default": "result= 1 if 'Mixlab' in input_text else 0"}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("INT",)
return {"required": {
"audio": ("AUDIO",),
},
}
RETURN_TYPES = ()
FUNCTION = "run"
OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
CATEGORY = "♾️Mixlab/Audio"
def run(self, input_text,input_num,python_code):
exec(python_code)
res=None
try:
# 可能会引发异常的代码
res=result
except:
# 处理异常的代码
print('')
INPUT_IS_LIST = False
OUTPUT_IS_LIST = ()
print(res)
OUTPUT_NODE = True
def run(self,audio):
# print(session_history)
return {"ui": {"text": [input_text],"num":[input_num]}, "result": (res,)}
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if is_tensor and (not 'audio_path' in audio):
filename_prefix=""
# 保存
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
results = list()
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
torchaudio.save(os.path.join(full_output_folder, file), audio['waveform'].squeeze(0), audio["sample_rate"])
results.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
else:
results=[{
"filename": audio['filename'],
"subfolder":audio['subfolder'],
"type": audio['type'],
"audio_path":audio['audio_path']
}]
# print(audio)
return {"ui": {"audio":results}}
+395 -32
View File
@@ -4,8 +4,72 @@ import urllib.error
import re,json,os,string,random
import folder_paths
import hashlib
import codecs
from zhipuai import ZhipuAI
import codecs,sys
import importlib.util
import subprocess
python = sys.executable
# 从文本中提取json
def extract_json_strings(text):
json_strings = []
brace_level = 0
json_str = ''
in_json = False
for char in text:
if char == '{':
brace_level += 1
in_json = True
if in_json:
json_str += char
if char == '}':
brace_level -= 1
if in_json and brace_level == 0:
json_strings.append(json_str)
json_str = ''
in_json = False
return json_strings[0] if len(json_strings)>0 else "{}"
def is_installed(package, package_overwrite=None,auto_install=True):
is_has=False
try:
spec = importlib.util.find_spec(package)
is_has=spec is not None
except ModuleNotFoundError:
pass
package = package_overwrite or package
if spec is None:
if auto_install==True:
print(f"Installing {package}...")
# 清华源 -i https://pypi.tuna.tsinghua.edu.cn/simple
command = f'"{python}" -m pip install {package}'
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ)
is_has=True
if result.returncode != 0:
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
is_has=False
else:
print(package+'## OK')
return is_has
# def is_installed(package):
# try:
# spec = importlib.util.find_spec(package)
# except ModuleNotFoundError:
# return False
# return spec is not None
def get_unique_hash(string):
hash_object = hashlib.sha1(string.encode())
@@ -44,28 +108,123 @@ def azure_client(key,url):
def openai_client(key,url):
client = openai.OpenAI(
api_key=key,
base_url=url
api_key=key,
base_url=url
)
return client
def ZhipuAI_client(key):
try:
if is_installed('zhipuai')==True:
from zhipuai import ZhipuAI
except:
print("#install zhipuai error")
client = ZhipuAI(
api_key=key, # 填写您的 APIKey
)
return client
# 优先使用phi
def phi_sort(lst):
return sorted(lst, key=lambda x: x.lower().count('phi'), reverse=True)
def get_llama_path():
try:
return folder_paths.get_folder_paths('llamafile')[0]
except:
return os.path.join(folder_paths.models_dir, "llamafile")
# def get_llama_models():
# res=[]
# model_path=get_llama_path()
# if os.path.exists(model_path):
# files = os.listdir(model_path)
# for file in files:
# if os.path.isfile(os.path.join(model_path, file)):
# res.append(file)
# res=phi_sort(res)
# return res
# llama_modes_list=get_llama_models()
# llama_modes_list=[]
# def get_llama_model_path(file_name):
# model_path=get_llama_path()
# mp=os.path.join(model_path,file_name)
# return mp
# def llama_cpp_client(file_name):
# try:
# if is_installed('llama_cpp')==False:
# import subprocess
# # 安装
# print('#pip install llama-cpp-python')
# result = subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python',
# '--extra-index-url',
# 'https://abetlen.github.io/llama-cpp-python/whl/cu121'
# ], capture_output=True, text=True)
# #检查命令执行结果
# if result.returncode == 0:
# print("#install success")
# from llama_cpp import Llama
# subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python[server]'
# ], capture_output=True, text=True)
# else:
# print("#install error")
# else:
# from llama_cpp import Llama
# except:
# print("#install llama-cpp-python error")
# if file_name:
# mp=get_llama_model_path(file_name)
# # file_name=get_llama_models()[0]
# # model_path=os.path.join(folder_paths.models_dir, "llamafile")
# # mp=os.path.join(model_path,file_name)
# llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
# return llm
if is_installed('json_repair'):
from json_repair import repair_json
def chat(client, model_name,messages ):
print('#chat',model_name,messages)
try_count = 0
while True:
try_count += 1
try:
response = client.chat.completions.create(
model=model_name,
messages=messages
)
if hasattr(client, "chat"):
response = client.chat.completions.create(
model=model_name,
messages=messages
)
else:
# 是llama的
response = client.create_chat_completion_openai_v1(
messages=messages,
# response_format={
# "type": "json_object",
# },
# temperature=0.7,
)
break
except openai.AuthenticationError as ex:
raise ex
@@ -74,7 +233,8 @@ def chat(client, model_name,messages ):
raise ex
time.sleep(3)
continue
# print(response.keys())
finish_reason = response.choices[0].finish_reason
if finish_reason != "stop":
raise RuntimeError("API finished with unexpected reason: " + finish_reason)
@@ -88,6 +248,36 @@ def chat(client, model_name,messages ):
return content
llm_apis=[
{
"value": "https://api.openai.com/v1",
"label": "openai"
},
{
"value": "https://openai.api2d.net/v1",
"label": "api2d"
},
# {
# "value": "https://docs-test-001.openai.azure.com",
# "label": "https://docs-test-001.openai.azure.com"
# },
{
"value": "https://api.moonshot.cn/v1",
"label": "Kimi"
},
{
"value": "https://api.deepseek.com/v1",
"label": "DeepSeek-V2"
},
{
"value": "https://api.siliconflow.cn/v1",
"label": "SiliconCloud"
}]
llm_apis_dict = {api["label"]: api["value"] for api in llm_apis}
class ChatGPTNode:
def __init__(self):
# self.__client = OpenAI()
@@ -97,33 +287,60 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=[
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-2024-05-13",
"gpt-4",
"gpt-4-0314",
"gpt-4-0613",
"gpt-3.5-turbo-0301",
"gpt-3.5-turbo-0613",
"gpt-3.5-turbo-16k-0613",
"qwen-turbo",
"qwen-plus",
"qwen-long",
"qwen-max",
"qwen-max-longcontext",
"glm-4",
"glm-3-turbo",
"moonshot-v1-8k",
"moonshot-v1-32k",
"moonshot-v1-128k",
"deepseek-chat",
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
]
return {
"required": {
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
# "api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
# "api_key":("STRING", {"forceInput": True,}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ([
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4"],
{"default": "gpt-3.5-turbo"}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
"api_url":(list(llm_apis_dict.keys()),
{"default": list(llm_apis_dict.keys())[0]}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
"optional":{
"api_key":("STRING", {"forceInput": True,}),
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
"custom_api_url":("STRING", {"forceInput": True,}), #适合自定义model
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
@@ -135,12 +352,29 @@ class ChatGPTNode:
def generate_contextual_text(self,
api_key,
api_url,
# api_key,
prompt,
system_content,
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
model,
seed,
context_size,
api_url,
api_key=None,
custom_model_name=None,
custom_api_url=None,
):
if custom_model_name!=None:
model=custom_model_name
api_url=llm_apis_dict[api_url] if api_url in llm_apis_dict else ""
if custom_api_url!=None:
api_url=custom_api_url
if api_key==None:
api_key="lm_studio"
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
@@ -153,7 +387,7 @@ class ChatGPTNode:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
print("api_key,api_url",api_key,api_url)
#
if is_azure_url(api_url):
client=azure_client(api_key,api_url)
@@ -162,9 +396,12 @@ class ChatGPTNode:
if model == "glm-4" :
client = ZhipuAI_client(api_key) # 使用 Zhipuai 的接口
print('using Zhipuai interface')
# elif model in llama_modes_list:
# #
# client=llama_cpp_client(model)
else :
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
print('using ChatGPT interface')
# print('using ChatGPT interface',api_key,api_url)
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
@@ -180,6 +417,7 @@ class ChatGPTNode:
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
@@ -200,6 +438,93 @@ class ChatGPTNode:
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class SiliconflowFreeNode:
def __init__(self):
# self.__client = OpenAI()
self.session_history = [] # 用于存储会话历史的列表
# self.seed=0
self.system_content="You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible."
@classmethod
def INPUT_TYPES(cls):
model_list= [
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
]
return {
"required": {
"api_key":("STRING", {"forceInput": True,}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
},
"optional":{
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
RETURN_NAMES = ("text","messages","session_history",)
FUNCTION = "generate_contextual_text"
CATEGORY = "♾️Mixlab/GPT"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,False,)
def generate_contextual_text(self,
api_key,
prompt,
system_content,
model,
seed,context_size,custom_model_name=None):
if custom_model_name!=None:
model=custom_model_name
api_url="https://api.siliconflow.cn/v1"
# 把系统信息和初始信息添加到会话历史中
if system_content:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
#
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
# print('using ChatGPT interface',api_key,api_url)
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
def crop_list_tail(lst, size):
if size >= len(lst):
return lst
elif size==0:
return []
else:
return lst[-size:]
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class ShowTextForGPT:
@classmethod
@@ -361,3 +686,41 @@ class TextSplitByDelimiter:
arr= arr[start_index:start_index + max_count * (skip_every+1):(skip_every+1)]
return (arr,)
class JsonRepair:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"json_string":("STRING", {"forceInput": True,}),
"key":("STRING", {"multiline": False,"dynamicPrompts": False,"default": ""}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("STRING","STRING",)
RETURN_NAMES = ("json_string","value",)
FUNCTION = "run"
# OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,False,)
CATEGORY = "♾️Mixlab/GPT"
def run(self, json_string,key=""):
json_string=extract_json_strings(json_string)
# print(json_string)
good_json_string = repair_json(json_string)
# 将 JSON 字符串解析为 Python 对象
data = json.loads(good_json_string)
v=""
if key!="" and (key in data):
v=data[key]
# 将 Python 对象转换回 JSON 字符串,确保中文字符不被转义
json_str_with_chinese = json.dumps(data, ensure_ascii=False)
return (json_str_with_chinese,v,)
+9 -3
View File
@@ -70,14 +70,20 @@ def load_caption_model(model_path,config,t='blip-base'):
return (caption_model,caption_processor)
def get_clip_interrogator_path():
try:
return folder_paths.get_folder_paths('clip_interrogator')[0]
except:
return os.path.join(folder_paths.models_dir, "clip_interrogator")
caption_model_path=os.path.join(folder_paths.models_dir, "clip_interrogator/Salesforce/blip-image-captioning-base")
cache_path=get_clip_interrogator_path()
caption_model_path=os.path.join(cache_path, "Salesforce","blip-image-captioning-base")
if not os.path.exists(caption_model_path):
print(f"## clip_interrogator_model not found: {caption_model_path}, pls download from https://huggingface.co/Salesforce/blip-image-captioning-base")
caption_model_path='Salesforce/blip-image-captioning-base'
cache_path=os.path.join(folder_paths.models_dir, "clip_interrogator")
# Tensor to PIL
def tensor2pil(image):
View File
+699 -192
View File
File diff suppressed because it is too large Load Diff
+7 -4
View File
@@ -42,8 +42,13 @@ else:
_available=True
llma_model_path=os.path.join(folder_paths.models_dir, "lama/big-lama.pt")
def get_lama_path():
try:
return folder_paths.get_folder_paths('lama')[0]
except:
return os.path.join(folder_paths.models_dir, "lama")
llma_model_path=os.path.join(get_lama_path(), "big-lama.pt")
if not os.path.exists(llma_model_path):
os.environ['LAMA_MODEL']=''
print(f"## lama torchscript model not found: {llma_model_path},pls download from https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt")
@@ -80,8 +85,6 @@ class LaMaInpainting:
"image": ("IMAGE",),
"mask": ("MASK",),
},
}
RETURN_TYPES = ("IMAGE",)
+32 -5
View File
@@ -2,16 +2,14 @@
import scipy.ndimage
import torch
from nodes import MAX_RESOLUTION
import numpy as np
# from PIL import Image, ImageDraw
from PIL import Image, ImageOps
from comfy.cli_args import args
import cv2
import cv2,os
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
import folder_paths,random
# Tensor to PIL
def tensor2pil(image):
@@ -71,6 +69,35 @@ def combine(destination, source, x, y):
return output
class PreviewMask_(SaveImage):
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append =''.join(random.choice("abcdehijklmnopqrstupvxyzfg") for x in range(5))
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
}
}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
# 运行的函数
def run(self, mask ):
img=tensor2pil(mask)
img=img.convert('RGB')
img=pil2tensor(img)
return self.save_images(img, 'temp_', None, None)
class OutlineMask:
@classmethod
+127
View File
@@ -0,0 +1,127 @@
# Referenced some code:https://github.com/IuvenisSapiens/ComfyUI_MiniCPM-V-2_6-int4
import os
import torch
import folder_paths
from transformers import AutoTokenizer, AutoModel
from torchvision.transforms.v2 import ToPILImage
from decord import VideoReader, cpu # pip install decord
from PIL import Image
def get_model_path(n=""):
try:
return folder_paths.get_folder_paths(n)[0]
except:
return os.path.join(folder_paths.models_dir, n)
class MiniCPM_VQA_Simple:
def __init__(self):
self.model_checkpoint = None
self.tokenizer = None
self.model = None
self.device = (
torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
)
self.bf16_support = (
torch.cuda.is_available()
and torch.cuda.get_device_capability(self.device)[0] >= 8
)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"text": ("STRING", {"default": "", "multiline": True}),
"seed": ("INT", {"default": -1}), # add seed parameter, default is -1
"temperature": (
"FLOAT",
{
"default": 0.7,
},
),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "inference"
CATEGORY = "♾️Mixlab/Image"
def inference(
self,
images,
text,
seed, # add seed parameter, default is -1
temperature,
keep_model_loaded,
):
if seed != -1:
torch.manual_seed(seed)
model_id = "openbmb/MiniCPM-V-2_6-int4"
self.model_checkpoint = os.path.join( get_model_path("prompt_generator"), os.path.basename(model_id))
if not os.path.exists(self.model_checkpoint):
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=model_id,
local_dir=self.model_checkpoint,
local_dir_use_symlinks=False,
endpoint='https://hf-mirror.com'
)
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_checkpoint,
trust_remote_code=True,
low_cpu_mem_usage=True,
)
if self.model is None:
self.model = AutoModel.from_pretrained(
self.model_checkpoint,
trust_remote_code=True,
low_cpu_mem_usage=True,
attn_implementation="sdpa",
torch_dtype=torch.bfloat16 if self.bf16_support else torch.float16,
)
with torch.no_grad():
images = images.permute([0, 3, 1, 2])
images = [ToPILImage()(img).convert("RGB") for img in images]
msgs = [{"role": "user", "content": images + [text]}]
params = {"use_image_id": False, }
# offload model to CPU
# self.model = self.model.to(torch.device("cpu"))
# self.model.eval()
result = self.model.chat(
image=None,
msgs=msgs,
tokenizer=self.tokenizer,
sampling=True,
# top_k=top_k,
# top_p=top_p,
temperature=temperature,
# repetition_penalty=repetition_penalty,
# max_new_tokens=max_new_tokens,
**params,
)
# offload model to GPU
# self.model = self.model.to(torch.device("cpu"))
# self.model.eval()
if not keep_model_loaded:
del self.tokenizer # release tokenizer memory
del self.model # release model memory
self.tokenizer = None # set tokenizer to None
self.model = None # set model to None
torch.cuda.empty_cache() # release GPU memory
torch.cuda.ipc_collect()
return (result,)
+104
View File
@@ -0,0 +1,104 @@
import torch
import numpy as np
from PIL import Image,ImageSequence,ImageOps
import base64
import io
import comfy.utils
import folder_paths
import node_helpers
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# Convert PIL to Tensor
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def load_image_to_tensor( image):
image_path = folder_paths.get_annotated_filepath(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (output_image, output_mask)
class P5Input:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"frames":("IMAGEBASE64",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("frames",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
OUTPUT_NODE = True
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self, frames):
ims=[]
for im in frames['images']:
# print(im)
if 'type' in im and (not f"[{im['type']}]" in im['name']):
im['name']=im['name']+" "+f"[{im['type']}]"
output_image, output_mask = load_image_to_tensor(im['name'])
ims.append(output_image)
if len(ims)==0:
image1 = Image.new('RGB', (512, 512), color='black')
return (pil2tensor(image1),)
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
image1 = torch.cat((image1, image2), dim=0)
# 用于节点提示:p5节点提示有多少帧
return {"ui": {"_info": [len(frames['images'])]}, "result": (image1,)}
+7 -1
View File
@@ -18,7 +18,13 @@ import json
# req = request.Request("http://127.0.0.1:8188/prompt", data=data)
# request.urlopen(req)
embeddings_path=os.path.join(folder_paths.models_dir, "embeddings")
def get_model_path(n=""):
try:
return folder_paths.get_folder_paths(n)[0]
except:
return os.path.join(folder_paths.models_dir, n)
embeddings_path=get_model_path("embeddings")
def get_files_with_extension(directory, extension):
+36 -22
View File
@@ -467,15 +467,37 @@ class BriaRMBG(nn.Module):
def get_U2NET_model_path():
try:
return folder_paths.get_folder_paths('rembg')[0]
except:
return os.path.join(folder_paths.models_dir, "rembg")
U2NET_HOME=os.path.join(folder_paths.models_dir, "rembg")
U2NET_HOME=get_U2NET_model_path()
os.environ["U2NET_HOME"] = U2NET_HOME
global _available
_available=False
def get_rembg_models(path):
"""从目录中获取文件并提取文件名
Args:
path: 目录路径
Returns:
文件名列表
"""
filenames = []
for root, _, files in os.walk(path):
for filename in files:
# 过滤隐藏文件
if not filename.startswith('.'):
name, ext = os.path.splitext(os.path.basename(filename))
filenames.append(name)
return filenames
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
@@ -509,8 +531,8 @@ except:
_available=False
def briarmbg_run(images=[]):
mroot=os.path.join(folder_paths.models_dir, "rembg")
def run_briarmbg(images=[]):
mroot=U2NET_HOME
m=os.path.join(mroot,'briarmbg.pth')
if os.path.exists(m)==False:
# 下载
@@ -573,14 +595,15 @@ def briarmbg_run(images=[]):
return (masks,rgba_images,rgb_images)
def run_bg(model_name= "unet",images=[]):
def run_rembg(model_name= "unet",images=[],callback=None):
# model_name = "unet" # "isnet-general-use"
# print('#run_rembg',model_name)
rembg_session = new_session(model_name)
masks=[]
rgba_images=[]
rgb_images=[]
# 进度条
pbar = comfy.utils.ProgressBar(len(images) )
pbar=callback
for img in images:
# use the post_process_mask argument to post process the mask to get better results.
mask = remove(img, session=rembg_session,only_mask=True,post_process_mask=True)
@@ -620,8 +643,9 @@ def run_bg(model_name= "unet",images=[]):
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
rgb_images.append(rgb_image)
pbar.update(1)
if pbar:
pbar.update(1)
return (masks,rgba_images,rgb_images)
@@ -643,17 +667,7 @@ class RembgNode_:
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"model_name": ([
"briarmbg",
"u2net",
"u2netp",
"u2net_human_seg",
"u2net_cloth_seg",
"silueta",
"isnet-general-use",
"isnet-anime",
],),
"model_name": (get_rembg_models(U2NET_HOME),),
},
}
@@ -681,9 +695,9 @@ class RembgNode_:
images.append(im)
if model_name=='briarmbg':
masks,rgba_images,rgb_images=briarmbg_run(images)
masks,rgba_images,rgb_images=run_briarmbg(images)
else:
masks,rgba_images,rgb_images=run_bg(model_name,images)
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
masks=[pil2tensor(m) for m in masks]
+6 -6
View File
@@ -90,7 +90,7 @@ class ScreenShareNode:
} }
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Screen"
@@ -109,7 +109,7 @@ class FloatingVideo:
@classmethod
def INPUT_TYPES(s):
return { "required":{
"images": ("IMAGE",)
"image": ("IMAGE",)
}, }
# RETURN_TYPES = ('IMAGE','MASK')
@@ -124,16 +124,16 @@ class FloatingVideo:
# OUTPUT_IS_LIST = (False,False,)
# 运行的函数
def run(self,images):
def run(self,image):
results = list()
for image in images:
image=tensor2pil(image)
for im in image:
im=tensor2pil(im)
# image_base64 = base64.b64encode(image.tobytes())
buffered = BytesIO()
image.save(buffered, format="JPEG")
im.save(buffered, format="JPEG")
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
results.append(image_base64)
+19 -11
View File
@@ -18,19 +18,25 @@ from lark import Lark, Transformer, v_args
global _available
_available=True
def get_text_generator_path():
try:
return folder_paths.get_folder_paths('prompt_generator')[0]
except:
return os.path.join(folder_paths.models_dir, "prompt_generator")
text_generator_model_path=os.path.join(folder_paths.models_dir, "prompt_generator/text2image-prompt-generator")
prompt_generator=get_text_generator_path()
text_generator_model_path=os.path.join(prompt_generator, "text2image-prompt-generator")
if not os.path.exists(text_generator_model_path):
print(f"## text_generator_model not found: {text_generator_model_path}, pls download from https://huggingface.co/succinctly/text2image-prompt-generator/tree/main")
text_generator_model_path='succinctly/text2image-prompt-generator'
zh_en_model_path=os.path.join(folder_paths.models_dir, "prompt_generator/opus-mt-zh-en")
zh_en_model_path=os.path.join(prompt_generator, "opus-mt-zh-en")
if not os.path.exists(zh_en_model_path):
print(f"## zh_en_model not found: {zh_en_model_path}, pls download from https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main")
zh_en_model_path='Helsinki-NLP/opus-mt-zh-en'
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
@@ -274,7 +280,7 @@ class ChinesePrompt:
},
"optional":{
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
},
@@ -325,13 +331,15 @@ class ChinesePrompt:
for t in texts:
if t:
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
# print('t',t)
result = parser.parse(t).children
# print('en_result',result)
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
en_texts.append(result[0])
try:
result = parser.parse(t).children
en_texts.append(result[0])
except:
print(f"Error parsing '{t}'")
t = translate(str(t))
en_texts.append(t)
zh_en_model.to('cpu')
print("test en_text",en_texts)
@@ -378,7 +386,7 @@ class PromptGenerate:
"optional":{
"multiple": (["off","on"],),
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
},
}
+180
View File
@@ -0,0 +1,180 @@
import sys
from os import path
sys.path.insert(0, path.dirname(__file__))
from PIL import Image
import numpy as np
import torch
from folder_paths import get_folder_paths, get_full_path, get_save_image_path, get_output_directory,models_dir
from comfy.model_management import get_torch_device
from .tsr.system import TSR
import comfy.utils
def get_triposr_model_path():
try:
return path.join(get_folder_paths('triposr')[0],'model.ckpt')
except:
return path.join(path.join(models_dir, "triposr"),'model.ckpt')
triposr_model_path=get_triposr_model_path()
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# Convert PIL to Tensor
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def fill_background(image):
im = np.array(image).astype(np.float32) / 255.0
im = im[:, :, :3] * im[:, :, 3:4] + (1 - im[:, :, 3:4]) * 0.5
im = Image.fromarray((im * 255.0).astype(np.uint8))
return im
class LoadTripoSRModel:
def __init__(self):
self.initialized_model = None
@classmethod
def INPUT_TYPES(s):
return {
"required": {
# "model": (get_filename_list("checkpoints"),),
"chunk_size": ("INT", {"default": 8192, "min": 0, "max": 10000})
}
}
RETURN_TYPES = ("TRIPOSR_MODEL",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, chunk_size):
device = get_torch_device()
if not torch.cuda.is_available():
device = "cpu"
if not self.initialized_model:
# triposr_model_path
print("#Loading TripoSR model",triposr_model_path)
self.initialized_model = TSR.from_pretrained_custom(
weight_path=triposr_model_path,
config_path=path.join(path.dirname(__file__), "tsr/config.yaml")
)
self.initialized_model.renderer.set_chunk_size(chunk_size)
self.initialized_model.to(device)
return (self.initialized_model,)
class TripoSRSampler:
def __init__(self):
self.initialized_model = None
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("TRIPOSR_MODEL",),
"image": ("IMAGE",),
"resolution": ("INT", {"default": 256, "min": 128, "max": 12288}),
"threshold": ("FLOAT", {"default": 25.0, "min": 0.0, "step": 0.01}),
"device":(["auto","cpu"],),
},
"optional": {
"mask": ("MASK",)
}
}
RETURN_TYPES = ("MESH",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, model, image, resolution, threshold,device='auto', mask=None):
reference_image=image
reference_mask=mask
device = get_torch_device()
if not torch.cuda.is_available():
device = "cpu"
if device=='cpu':
device = "cpu"
print('#TripoSRSampler device',device)
to_images=[]
for i in range(len(reference_image)):
image = reference_image[i]
if reference_mask is not None:
mask = reference_mask[i].unsqueeze(2)
image = torch.cat((image, mask), dim=2).detach().cpu().numpy()
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8))
image = fill_background(image)
else:
image = tensor2pil(image)
image = image.convert('RGB')
to_images.append(image)
# 进度条
pbar = comfy.utils.ProgressBar(len(to_images))
def callback(c):
pbar.update(1)
scene_codes = model(to_images, device)
meshes = model.extract_mesh(scene_codes, resolution=resolution, threshold=threshold,callback=callback)
del model
return (meshes,)
class SaveTripoSRMesh:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESH",),
# "format":(["glb","obj"],),
"filename_prefix":("STRING", {"multiline": False,"default": "TripoSR_"})
}
}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, mesh,filename_prefix):
format='glb'
saved = list()
full_output_folder, filename, counter, subfolder, filename_prefix = get_save_image_path(filename_prefix,
get_output_directory())
for (index, single_mesh) in enumerate(mesh):
filename_with_batch_num = filename.replace("%batch_num%", str(index))
file = f"{filename_with_batch_num}_{counter:05}_.{format}"
single_mesh.apply_transform(np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]]))
single_mesh.export(path.join(full_output_folder, file))
saved.append({
"filename": file,
"type": "output",
"subfolder": subfolder
})
return {"ui": {"mesh": saved}}
+29 -9
View File
@@ -133,7 +133,7 @@ def get_font_files(directory):
return font_files
r_directory = os.path.join(os.path.dirname(__file__), '../assets/')
r_directory = os.path.join(os.path.dirname(__file__), '..','assets','/')
font_files = get_font_files(r_directory)
# print(font_files)
@@ -181,6 +181,28 @@ class ColorInput:
return (h,r,g,b,a,)
class KeyInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key":("KEY",),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("key",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,key):
return (key,)
class FontInput:
@classmethod
@@ -284,7 +306,7 @@ class FloatSlider:
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ('weight(0-1)',)
RETURN_NAMES = ('FLOAT',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
@@ -297,9 +319,7 @@ class FloatSlider:
number = min_value
elif number > max_value:
number = max_value
scaled_number = (number - min_value) / (max_value - min_value)
return (scaled_number,)
return (number,)
class IntNumber:
@classmethod
@@ -568,7 +588,7 @@ class AppInfo:
},
"optional":{
"IMAGE": ("IMAGE",),
"image": ("IMAGE",),
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
"version":("INT", {
"default": 1,
@@ -596,12 +616,12 @@ class AppInfo:
INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (True,)
def run(self,name,input_ids,output_ids,IMAGE,description,version,share_prefix,link,category,auto_save):
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save):
name=name[0]
im=None
if IMAGE:
im=IMAGE[0][0]
if image:
im=image[0][0]
#TODO batch 的方式需要处理
im=create_temp_file(im)
# image [img,] img[batch,w,h,a] 列表里面是batch,
+379 -70
View File
@@ -17,9 +17,128 @@ import folder_paths
from comfy.k_diffusion.utils import FolderOfImages
from comfy.utils import common_upscale
import torchaudio
import base64
import mimetypes
def get_frames(frame_count, frames, revert=False):
if not revert:
if frame_count <= len(frames):
return frames[:frame_count]
else:
return [frames[i % len(frames)] for i in range(frame_count)]
else:
extended_frames = frames + frames[-2:0:-1] # 正向加反向中间部分
if frame_count <= len(extended_frames):
return extended_frames[:frame_count]
else:
return [extended_frames[i % len(extended_frames)] for i in range(frame_count)]
# # 示例用法
# frames = ["frame1", "frame2", "frame3"]
# frame_count = 2
# result = get_frames(frame_count, frames, revert=False)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame1', 'frame2', 'frame3', 'frame1']
# result = get_frames(frame_count, frames, revert=True)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame2', 'frame1', 'frame2', 'frame3']
def get_mime_type(file_path):
# 获取文件的 MIME 类型
mime_type, _ = mimetypes.guess_type(file_path)
# 如果无法猜测类型,返回默认类型
if mime_type is None:
return 'application/octet-stream'
return mime_type
# import subprocess
# from imageio_ffmpeg import get_ffmpeg_exe
def save_audio_base64s_to_file(base64_audios, output_folder, file_name):
# Ensure the output folder exists
if not os.path.exists(output_folder):
os.makedirs(output_folder)
decoded_audios=[]
for a in base64_audios:
# If the base64 string contains a header, remove it
if ',' in a:
a = a.split(',')[1]
# 解码 base64 数据
a=base64.b64decode(a)
decoded_audios.append(a)
# 拼接音频数据
combined_audio = b''.join(decoded_audios)
# Create the full file path
file_path = os.path.join(output_folder, file_name)
# Write the decoded audio to the file
with open(file_path, 'wb') as audio_file:
audio_file.write(combined_audio)
return file_path
# Example usage
# base64_audio = "data:audio/wav;base64,UklGRiQAAABXQVZFZm10IBAAAAABAAEAIlYAAESsAAACABAAZGF0YQAAAAA="
# output_folder = "audio_files"
# file_name = "output.wav"
# file_path = save_audio_base64_to_file(base64_audio, output_folder, file_name)
# print(f"Audio saved to: {file_path}")
# 写一个python文件,用来 判断文件夹内命名为 所有chat_tts开头的文件数量(chat_tts_00001),并输出新的编号
def get_new_counter(full_output_folder, filename_prefix):
# 获取目录中的所有文件
files = os.listdir(full_output_folder)
# 过滤出以 filename_prefix 开头并且后续部分为数字的文件
filtered_files = []
for f in files:
if f.startswith(filename_prefix):
# 去掉文件名中的前缀和后缀,只保留中间的数字部分
base_name = f[len(filename_prefix)+1:]
number_part = base_name.split('.')[0] # 假设文件名中只有一个点,即扩展名
if number_part.isdigit():
filtered_files.append(int(number_part))
if not filtered_files:
return 1
# 获取最大的编号
max_number = max(filtered_files)
# 新的编号
return max_number + 1
def crop_audio(input_file, start_time, duration):
# Load the audio file
audio_tensor, sample_rate = torchaudio.load(input_file)
# Convert start_time and duration from seconds to sample indices
start_sample = int(start_time * sample_rate)
end_sample = start_sample + int(duration * sample_rate)
# Perform the slicing
cropped_audio_tensor = audio_tensor[:, start_sample:end_sample]
# Save the cropped audio to a new file
torchaudio.save(input_file, cropped_audio_tensor, sample_rate)
return input_file
def generate_folder_name(directory,video_path):
# Get the directory and filename from the video path
_, filename = os.path.split(video_path)
@@ -60,6 +179,9 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
# 打印当前片段的起始帧和结束帧
print(f"Segment {i+1}: Start Frame {start_frame}, End Frame {end_frame}")
if end_frame<start_frame:
break
# 保存当前片段为一个视频文件
segment_video_path = f"{output_dir}/segment_{i+1}.avi"
@@ -68,6 +190,7 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
segment_video = cv2.VideoWriter(segment_video_path, fourcc, fps, (int(video_capture.get(cv2.CAP_PROP_FRAME_WIDTH)),
int(video_capture.get(cv2.CAP_PROP_FRAME_HEIGHT))))
for frame_num in range(start_frame, end_frame):
ret, frame = video_capture.read()
if ret:
@@ -87,7 +210,7 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
folder_paths.folder_names_and_paths["video_formats"] = (
[
os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "video_formats"),
os.path.join(os.path.dirname(os.path.abspath(__file__)), ".", "video_formats"),
],
[".json"]
)
@@ -101,6 +224,25 @@ if ffmpeg_path is None:
except:
print("ffmpeg could not be found. Outputs that require it have been disabled")
def combine_audio_video(audio_path, video_path, output_path):
command = [
ffmpeg_path,
'-i', video_path,
'-i', audio_path,
'-c:v', 'copy',
'-c:a', 'aac',
'-shortest',
output_path
]
subprocess.run(command, check=True)
return output_path
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
@@ -262,7 +404,7 @@ class LoadVideoAndSegment:
files.append(f)
return {"required": {
"video": (sorted(files), {"video_upload": True}),
"video_segment_frames": ("INT", {"default": 10, "min": 1, "step": 1}),
"video_segment_frames": ("INT", {"default": 10, "min": -1, "step": 1}),
"transition_frames": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
@@ -332,63 +474,6 @@ class LoadVideoAndSegment:
video_path = folder_paths.get_annotated_filepath(video)
# check if video is a gif - will need to use cv fallback to read frames
# use cv fallback if ffmpeg not installed or gif
# if ffmpeg_path is None:
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# otherwise, continue with ffmpeg
# args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
# try:
# with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
# for line in proc.stderr.readlines():
# match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
# if match is not None:
# size = [int(match.group(1)), int(match.group(2))]
# break
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
# "-pix_fmt", "rgb24"]
# vfilters = []
# if skip_first_frames > 0:
# vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
# if frame_load_cap > 0:
# vfilters.append(f"select=gt({frame_load_cap}\\,n)")
# #manually calculate aspect ratio to ensure reads remain aligned
# if len(vfilters) > 0:
# args_all_frames += ["-vf", ",".join(vfilters)]
# args_all_frames += ["-f", "rawvideo", "-"]
# images = []
# try:
# with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
# #Manually buffer enough bytes for an image
# bpi = size[0]*size[1]*3
# current_bytes = bytearray(bpi)
# current_offset=0
# while True:
# bytes_read = proc.stdout.read(bpi - current_offset)
# if bytes_read is None:#sleep to wait for more data
# time.sleep(.2)
# continue
# if len(bytes_read) == 0:#EOF
# break
# current_bytes[current_offset:len(bytes_read)] = bytes_read
# current_offset+=len(bytes_read)
# if current_offset == bpi:
# images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
# current_offset = 0
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# imgs=split_list(images,video_segment_frames,transition_frames)
# temp path
tp=folder_paths.get_temp_directory()
basename = os.path.basename(video_path) # 获取文件名
@@ -396,15 +481,22 @@ class LoadVideoAndSegment:
folder_path = create_folder(tp,name_without_extension)
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
if video_segment_frames==-1:
# 不切割视频
scenes_video=[video_path]
# 读取视频文件
video_capture = cv2.VideoCapture(video_path)
# 获取视频的总帧数和帧率
total_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
fps = video_capture.get(cv2.CAP_PROP_FPS)
else:
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
# imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
# images = torch.from_numpy(np.stack(images))
return (scenes_video,len(scenes_video), total_frames,fps,)
@@ -422,7 +514,113 @@ class LoadVideoAndSegment:
return "Invalid image file: {}".format(video)
return True
class LoadAndCombinedAudio_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audios": ("AUDIOBASE64",),
"start_time": ("FLOAT" , {"default": 0, "min": 0, "max": 10000000, "step": 0.01}),
"duration": ("FLOAT" , {"default": 10, "min": -1, "max": 10000000, "step": 0.01}),
},
}
CATEGORY = "♾️Mixlab/Audio"
RETURN_TYPES = ("STRING","AUDIO",)
RETURN_NAMES = ("audio_file_path","audio",)
FUNCTION = "run"
def run(self,audios, start_time, duration):
output_dir = folder_paths.get_output_directory()
counter=get_new_counter(output_dir,'audio_')
audio_file_name = f"audio_{counter:05}.wav"
audio_file=save_audio_base64s_to_file(audios['base64'],output_dir,audio_file_name)
# duration == -1 则不裁切
if duration > -1:
crop_audio(audio_file, start_time, duration)
waveform, sample_rate = torchaudio.load(audio_file)
audio = {
"filename": audio_file_name,
"subfolder": "",
"type": "output",
"audio_path":audio_file,
"waveform": waveform.unsqueeze(0),
"sample_rate": sample_rate}
return (audio_file,audio ,)
class CombineAudioVideo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"video": ("SCENE_VIDEO",),
"audio": ("AUDIO", ),
},
}
CATEGORY = "♾️Mixlab/Video"
OUTPUT_NODE = True
FUNCTION = "run"
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("SCENE_VIDEO",)
def run(self,video, audio):
output_dir = folder_paths.get_output_directory()
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if "audio_path" in audio:
is_tensor=False
audio_file_path=audio["audio_path"]
if is_tensor:
filename_prefix="audio_tmp"
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
filename_prefix,
folder_paths.get_temp_directory())
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
audio_file_path=os.path.join(full_output_folder, file)
torchaudio.save(audio_file_path, audio['waveform'].squeeze(0), audio["sample_rate"])
# 获取文件名和扩展名
base, ext = os.path.splitext(video)
counter=get_new_counter(output_dir,'video_final_')
v_file = f"video_final_{counter:05}{ext}"
v_file_path=os.path.join(output_dir, v_file)
combine_audio_video(audio_file_path,video,v_file_path)
previews = [
{
"filename": v_file,
"subfolder": "",
"type": "output",
"format": get_mime_type(v_file),
}
]
return {"ui": {"gifs": previews},"result":(v_file_path,)}
# The code is based on ComfyUI-VideoHelperSuite modification.
class VideoCombine_Adv:
@@ -433,6 +631,7 @@ class VideoCombine_Adv:
ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")]
else:
ffmpeg_formats = []
# ffmpeg_formats =["video/"+x for x in ['webm', 'mp4', 'mkv']]
return {
"required": {
"image_batch": ("IMAGE",),
@@ -453,7 +652,8 @@ class VideoCombine_Adv:
},
}
RETURN_TYPES = ()
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("scenes_video",)
OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Video"
FUNCTION = "run"
@@ -622,7 +822,7 @@ class VideoCombine_Adv:
"format": format,
}
]
return {"ui": {"gifs": previews}}
return {"ui": {"gifs": previews},"result":(file_path,)}
class VAEEncodeForInpaint_Frames:
@@ -689,4 +889,113 @@ class VAEEncodeForInpaint_Frames:
result.append({"samples":t, "noise_mask": (mask_erosion[:,:,:x,:y].round())})
return (result, )
return (result, )
class GenerateFramesByCount:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"frames": ('IMAGE',),
"frame_count": ("INT", {"default": 72, "min": 1, "step": 1}),
"revert" :("BOOLEAN", {"default": True},),
},}
RETURN_TYPES = ('IMAGE',)
RETURN_NAMES = ("frames",)
FUNCTION = "r"
CATEGORY = "♾️Mixlab/Video"
# INPUT_IS_LIST = True
def r(self, frames, frame_count, revert):
image_list = [frames[i:i + 1, ...] for i in range(frames.shape[0])]
image_list=get_frames(frame_count,image_list,revert)
images = torch.cat(image_list, dim=0)
return (images,)
class scenesNode_:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"scenes_video": ('SCENE_VIDEO',),
"index": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
RETURN_TYPES = ('IMAGE','INT',)
RETURN_NAMES = ("video frames (batch)","count",)
# OUTPUT_IS_LIST = (False,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = True
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
# print('#video',video)
try:
video_cap = cv2.VideoCapture(video)
if not video_cap.isOpened():
raise ValueError(f"{video} could not be loaded with cv fallback.")
# set video_cap to look at start_index frame
images = []
total_frame_count = 0
frames_added = 0
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
target_frame_time = base_frame_time
time_offset=0.0
while video_cap.isOpened():
if time_offset < target_frame_time:
is_returned, frame = video_cap.read()
# if didn't return frame, video has ended
if not is_returned:
break
time_offset += base_frame_time
if time_offset < target_frame_time:
continue
time_offset -= target_frame_time
# if not at start_index, skip doing anything with frame
total_frame_count += 1
if total_frame_count <= skip_first_frames:
continue
# TODO: do whatever operations need to happen, like force_size, etc
# opencv loads images in BGR format (yuck), so need to convert to RGB for ComfyUI use
# follow up: can videos ever have an alpha channel?
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# convert frame to comfyui's expected format (taken from comfy's load image code)
image = Image.fromarray(frame)
image = ImageOps.exif_transpose(image)
image = np.array(image, dtype=np.float32) / 255.0
image = torch.from_numpy(image)[None,]
images.append(image)
frames_added += 1
# if cap exists and we've reached it, stop processing frames
if frame_load_cap > 0 and frames_added >= frame_load_cap:
break
finally:
video_cap.release()
images = torch.cat(images, dim=0)
return (images, frames_added,)
def run(self, scenes_video,index):
print('#scenes_video',index,scenes_video)
index=index[0]
if len(scenes_video) > index:
vp=scenes_video[index]
else:
vp=scenes_video[-1]
return self.load_video_cv_fallback(vp,0,0)
+172
View File
@@ -0,0 +1,172 @@
import torch
from PIL import Image, ImageOps, ImageSequence, ImageFile
from PIL.PngImagePlugin import PngInfo
import numpy as np
import os
import folder_paths
import node_helpers
import hashlib
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# tensor 取hash值
def tensor_to_hash(tensor):
# 将 Tensor 转换为 NumPy 数组
np_array = tensor.cpu().numpy()
# 将 NumPy 数组转换为字节数据
byte_data = np_array.tobytes()
# 计算哈希值
hash_value = hashlib.md5(byte_data).hexdigest()
return hash_value
def create_temp_file(image):
output_dir = folder_paths.get_temp_directory()
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path('material', output_dir)
image=tensor2pil(image)
image_file = f"{filename}_{counter:05}.png"
image_path=os.path.join(full_output_folder, image_file)
image.save(image_path,compress_level=4)
return (image_path,[{
"filename": image_file,
"subfolder": subfolder,
"type": "temp"
}])
# image - tensor - 文件路径
# loadImage的方法( 文件路径 - image-mask )
class EditMask:
def __init__(self):
self.image_id = None
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("IMAGE",), # 表示一个张量
},
"optional":{
"image_update": ("IMAGE_FILE",)
},
}
CATEGORY = "♾️Mixlab/Mask"
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("image", "mask")
FUNCTION = "edit"
OUTPUT_NODE = True
def edit(self, image,image_update=None):
# 根据image输入来判断是否是新的图片
if self.image_id==None:
self.image_id=tensor_to_hash(image)
image_update=None
else:
image_id=tensor_to_hash(image)
if image_id!=self.image_id:
image_update=None
self.image_id=image_id
image_path=None
# print('#image_update',self.image_id,image_update)
if image_update==None:
print('--')
else:
if 'images' in image_update:
images=image_update['images']
filename=images[0]['filename']
subfolder=images[0]['subfolder']
type=images[0]['type']
name, base_dir=folder_paths.annotated_filepath(filename)
if type.endswith("output"):
base_dir = folder_paths.get_output_directory()
elif type.endswith("input"):
base_dir = folder_paths.get_input_directory()
elif type.endswith("temp"):
base_dir = folder_paths.get_temp_directory()
#base_dir = folder_paths.get_input_directory()
# print(base_dir,subfolder, name)
image_path = os.path.join(base_dir,subfolder, name)
if image_path==None:
image_path,images=create_temp_file(image)
print('#image_path',os.path.exists(image_path),image_path)
# image_path = folder_paths.get_annotated_filepath(image) #文件名
if not os.path.exists(image_path):
image_path,images=create_temp_file(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
# 尺寸不对,需要按照image来
mask = torch.zeros((h, w), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return {"ui":{"images": images},"result": (output_image, output_mask)}
# return (output_image, output_mask)
+38
View File
@@ -0,0 +1,38 @@
cond_image_size: 512
image_tokenizer_cls: tsr.models.tokenizers.image.DINOSingleImageTokenizer
image_tokenizer:
pretrained_model_name_or_path: "facebook/dino-vitb16"
tokenizer_cls: tsr.models.tokenizers.triplane.Triplane1DTokenizer
tokenizer:
plane_size: 32
num_channels: 1024
backbone_cls: tsr.models.transformer.transformer_1d.Transformer1D
backbone:
in_channels: ${tokenizer.num_channels}
num_attention_heads: 16
attention_head_dim: 64
num_layers: 16
cross_attention_dim: 768
post_processor_cls: tsr.models.network_utils.TriplaneUpsampleNetwork
post_processor:
in_channels: 1024
out_channels: 40
decoder_cls: tsr.models.network_utils.NeRFMLP
decoder:
in_channels: 120 # 3 * 40
n_neurons: 64
n_hidden_layers: 9
activation: silu
renderer_cls: tsr.models.nerf_renderer.TriplaneNeRFRenderer
renderer:
radius: 0.87 # slightly larger than 0.5 * sqrt(3)
feature_reduction: concat
density_activation: exp
density_bias: -1.0
num_samples_per_ray: 128
+51
View File
@@ -0,0 +1,51 @@
from typing import Callable, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from skimage import measure
class IsosurfaceHelper(nn.Module):
points_range: Tuple[float, float] = (0, 1)
@property
def grid_vertices(self) -> torch.FloatTensor:
raise NotImplementedError
class MarchingCubeHelper(IsosurfaceHelper):
def __init__(self, resolution: int) -> None:
super().__init__()
self.resolution = resolution
#self.mc_func: Callable = marching_cubes
self._grid_vertices: Optional[torch.FloatTensor] = None
@property
def grid_vertices(self) -> torch.FloatTensor:
if self._grid_vertices is None:
# keep the vertices on CPU so that we can support very large resolution
x, y, z = (
torch.linspace(*self.points_range, self.resolution),
torch.linspace(*self.points_range, self.resolution),
torch.linspace(*self.points_range, self.resolution),
)
x, y, z = torch.meshgrid(x, y, z, indexing="ij")
verts = torch.cat(
[x.reshape(-1, 1), y.reshape(-1, 1), z.reshape(-1, 1)], dim=-1
).reshape(-1, 3)
self._grid_vertices = verts
return self._grid_vertices
def forward(
self,
level: torch.FloatTensor,
) -> Tuple[torch.FloatTensor, torch.LongTensor]:
level = -level.view(self.resolution, self.resolution, self.resolution)
v_pos, t_pos_idx, _, __ = measure.marching_cubes((level.detach().cpu() if level.is_cuda else level.detach()).numpy(), 0.0) #self.mc_func(level.detach(), 0.0)
v_pos = torch.from_numpy(v_pos.copy()).type(torch.FloatTensor).to(level.device)
t_pos_idx = torch.from_numpy(t_pos_idx.copy()).type(torch.LongTensor).to(level.device)
v_pos = v_pos[..., [0, 1, 2]]
t_pos_idx = t_pos_idx[..., [1, 0, 2]]
v_pos = v_pos / (self.resolution - 1.0)
return v_pos, t_pos_idx
+180
View File
@@ -0,0 +1,180 @@
from dataclasses import dataclass
from typing import Dict
import torch
import torch.nn.functional as F
from einops import rearrange, reduce
from ..utils import (
BaseModule,
chunk_batch,
get_activation,
rays_intersect_bbox,
scale_tensor,
)
class TriplaneNeRFRenderer(BaseModule):
@dataclass
class Config(BaseModule.Config):
radius: float
feature_reduction: str = "concat"
density_activation: str = "trunc_exp"
density_bias: float = -1.0
color_activation: str = "sigmoid"
num_samples_per_ray: int = 128
randomized: bool = False
cfg: Config
def configure(self) -> None:
assert self.cfg.feature_reduction in ["concat", "mean"]
self.chunk_size = 0
def set_chunk_size(self, chunk_size: int):
assert (
chunk_size >= 0
), "chunk_size must be a non-negative integer (0 for no chunking)."
self.chunk_size = chunk_size
def query_triplane(
self,
decoder: torch.nn.Module,
positions: torch.Tensor,
triplane: torch.Tensor,
) -> Dict[str, torch.Tensor]:
input_shape = positions.shape[:-1]
positions = positions.view(-1, 3)
# positions in (-radius, radius)
# normalized to (-1, 1) for grid sample
positions = scale_tensor(
positions, (-self.cfg.radius, self.cfg.radius), (-1, 1)
)
def _query_chunk(x):
indices2D: torch.Tensor = torch.stack(
(x[..., [0, 1]], x[..., [0, 2]], x[..., [1, 2]]),
dim=-3,
)
out: torch.Tensor = F.grid_sample(
rearrange(triplane, "Np Cp Hp Wp -> Np Cp Hp Wp", Np=3),
rearrange(indices2D, "Np N Nd -> Np () N Nd", Np=3),
align_corners=False,
mode="bilinear",
)
if self.cfg.feature_reduction == "concat":
out = rearrange(out, "Np Cp () N -> N (Np Cp)", Np=3)
elif self.cfg.feature_reduction == "mean":
out = reduce(out, "Np Cp () N -> N Cp", Np=3, reduction="mean")
else:
raise NotImplementedError
net_out: Dict[str, torch.Tensor] = decoder(out)
return net_out
if self.chunk_size > 0:
net_out = chunk_batch(_query_chunk, self.chunk_size, positions)
else:
net_out = _query_chunk(positions)
net_out["density_act"] = get_activation(self.cfg.density_activation)(
net_out["density"] + self.cfg.density_bias
)
net_out["color"] = get_activation(self.cfg.color_activation)(
net_out["features"]
)
net_out = {k: v.view(*input_shape, -1) for k, v in net_out.items()}
return net_out
def _forward(
self,
decoder: torch.nn.Module,
triplane: torch.Tensor,
rays_o: torch.Tensor,
rays_d: torch.Tensor,
**kwargs,
):
rays_shape = rays_o.shape[:-1]
rays_o = rays_o.view(-1, 3)
rays_d = rays_d.view(-1, 3)
n_rays = rays_o.shape[0]
t_near, t_far, rays_valid = rays_intersect_bbox(rays_o, rays_d, self.cfg.radius)
t_near, t_far = t_near[rays_valid], t_far[rays_valid]
t_vals = torch.linspace(
0, 1, self.cfg.num_samples_per_ray + 1, device=triplane.device
)
t_mid = (t_vals[:-1] + t_vals[1:]) / 2.0
z_vals = t_near * (1 - t_mid[None]) + t_far * t_mid[None] # (N_rays, N_samples)
xyz = (
rays_o[:, None, :] + z_vals[..., None] * rays_d[..., None, :]
) # (N_rays, N_sample, 3)
mlp_out = self.query_triplane(
decoder=decoder,
positions=xyz,
triplane=triplane,
)
eps = 1e-10
# deltas = z_vals[:, 1:] - z_vals[:, :-1] # (N_rays, N_samples)
deltas = t_vals[1:] - t_vals[:-1] # (N_rays, N_samples)
alpha = 1 - torch.exp(
-deltas * mlp_out["density_act"][..., 0]
) # (N_rays, N_samples)
accum_prod = torch.cat(
[
torch.ones_like(alpha[:, :1]),
torch.cumprod(1 - alpha[:, :-1] + eps, dim=-1),
],
dim=-1,
)
weights = alpha * accum_prod # (N_rays, N_samples)
comp_rgb_ = (weights[..., None] * mlp_out["color"]).sum(dim=-2) # (N_rays, 3)
opacity_ = weights.sum(dim=-1) # (N_rays)
comp_rgb = torch.zeros(
n_rays, 3, dtype=comp_rgb_.dtype, device=comp_rgb_.device
)
opacity = torch.zeros(n_rays, dtype=opacity_.dtype, device=opacity_.device)
comp_rgb[rays_valid] = comp_rgb_
opacity[rays_valid] = opacity_
comp_rgb += 1 - opacity[..., None]
comp_rgb = comp_rgb.view(*rays_shape, 3)
return comp_rgb
def forward(
self,
decoder: torch.nn.Module,
triplane: torch.Tensor,
rays_o: torch.Tensor,
rays_d: torch.Tensor,
) -> Dict[str, torch.Tensor]:
if triplane.ndim == 4:
comp_rgb = self._forward(decoder, triplane, rays_o, rays_d)
else:
comp_rgb = torch.stack(
[
self._forward(decoder, triplane[i], rays_o[i], rays_d[i])
for i in range(triplane.shape[0])
],
dim=0,
)
return comp_rgb
def train(self, mode=True):
self.randomized = mode and self.cfg.randomized
return super().train(mode=mode)
def eval(self):
self.randomized = False
return super().eval()
+124
View File
@@ -0,0 +1,124 @@
from dataclasses import dataclass
from typing import Optional
import torch
import torch.nn as nn
from einops import rearrange
from ..utils import BaseModule
class TriplaneUpsampleNetwork(BaseModule):
@dataclass
class Config(BaseModule.Config):
in_channels: int
out_channels: int
cfg: Config
def configure(self) -> None:
self.upsample = nn.ConvTranspose2d(
self.cfg.in_channels, self.cfg.out_channels, kernel_size=2, stride=2
)
def forward(self, triplanes: torch.Tensor) -> torch.Tensor:
triplanes_up = rearrange(
self.upsample(
rearrange(triplanes, "B Np Ci Hp Wp -> (B Np) Ci Hp Wp", Np=3)
),
"(B Np) Co Hp Wp -> B Np Co Hp Wp",
Np=3,
)
return triplanes_up
class NeRFMLP(BaseModule):
@dataclass
class Config(BaseModule.Config):
in_channels: int
n_neurons: int
n_hidden_layers: int
activation: str = "relu"
bias: bool = True
weight_init: Optional[str] = "kaiming_uniform"
bias_init: Optional[str] = None
cfg: Config
def configure(self) -> None:
layers = [
self.make_linear(
self.cfg.in_channels,
self.cfg.n_neurons,
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
),
self.make_activation(self.cfg.activation),
]
for i in range(self.cfg.n_hidden_layers - 1):
layers += [
self.make_linear(
self.cfg.n_neurons,
self.cfg.n_neurons,
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
),
self.make_activation(self.cfg.activation),
]
layers += [
self.make_linear(
self.cfg.n_neurons,
4, # density 1 + features 3
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
)
]
self.layers = nn.Sequential(*layers)
def make_linear(
self,
dim_in,
dim_out,
bias=True,
weight_init=None,
bias_init=None,
):
layer = nn.Linear(dim_in, dim_out, bias=bias)
if weight_init is None:
pass
elif weight_init == "kaiming_uniform":
torch.nn.init.kaiming_uniform_(layer.weight, nonlinearity="relu")
else:
raise NotImplementedError
if bias:
if bias_init is None:
pass
elif bias_init == "zero":
torch.nn.init.zeros_(layer.bias)
else:
raise NotImplementedError
return layer
def make_activation(self, activation):
if activation == "relu":
return nn.ReLU(inplace=True)
elif activation == "silu":
return nn.SiLU(inplace=True)
else:
raise NotImplementedError
def forward(self, x):
inp_shape = x.shape[:-1]
x = x.reshape(-1, x.shape[-1])
features = self.layers(x)
features = features.reshape(*inp_shape, -1)
out = {"density": features[..., 0:1], "features": features[..., 1:4]}
return out
+72
View File
@@ -0,0 +1,72 @@
from dataclasses import dataclass
import torch
import torch.nn as nn
from einops import rearrange
from huggingface_hub import hf_hub_download
from transformers.models.vit.modeling_vit import ViTModel
from ...utils import BaseModule
import os
import folder_paths
model_path=os.path.join(folder_paths.models_dir,'triposr')
class DINOSingleImageTokenizer(BaseModule):
@dataclass
class Config(BaseModule.Config):
pretrained_model_name_or_path: str = "facebook/dino-vitb16"
enable_gradient_checkpointing: bool = False
cfg: Config
def configure(self) -> None:
print('#Loading ViTModel:',os.path.join(model_path,self.cfg.pretrained_model_name_or_path))
self.model: ViTModel = ViTModel(
ViTModel.config_class.from_pretrained(
hf_hub_download(
repo_id=self.cfg.pretrained_model_name_or_path,
filename="config.json",
local_dir=model_path,
endpoint='https://hf-mirror.com'
)
)
)
if self.cfg.enable_gradient_checkpointing:
self.model.encoder.gradient_checkpointing = True
self.register_buffer(
"image_mean",
torch.as_tensor([0.485, 0.456, 0.406]).reshape(1, 1, 3, 1, 1),
persistent=False,
)
self.register_buffer(
"image_std",
torch.as_tensor([0.229, 0.224, 0.225]).reshape(1, 1, 3, 1, 1),
persistent=False,
)
def forward(self, images: torch.FloatTensor, **kwargs) -> torch.FloatTensor:
packed = False
if images.ndim == 4:
packed = True
images = images.unsqueeze(1)
batch_size, n_input_views = images.shape[:2]
images = (images - self.image_mean) / self.image_std
out = self.model(
rearrange(images, "B N C H W -> (B N) C H W"), interpolate_pos_encoding=True
)
local_features, global_features = out.last_hidden_state, out.pooler_output
local_features = local_features.permute(0, 2, 1)
local_features = rearrange(
local_features, "(B N) Ct Nt -> B N Ct Nt", B=batch_size
)
if packed:
local_features = local_features.squeeze(1)
return local_features
def detokenize(self, *args, **kwargs):
raise NotImplementedError
+45
View File
@@ -0,0 +1,45 @@
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
from einops import rearrange, repeat
from ...utils import BaseModule
class Triplane1DTokenizer(BaseModule):
@dataclass
class Config(BaseModule.Config):
plane_size: int
num_channels: int
cfg: Config
def configure(self) -> None:
self.embeddings = nn.Parameter(
torch.randn(
(3, self.cfg.num_channels, self.cfg.plane_size, self.cfg.plane_size),
dtype=torch.float32,
)
* 1
/ math.sqrt(self.cfg.num_channels)
)
def forward(self, batch_size: int) -> torch.Tensor:
return rearrange(
repeat(self.embeddings, "Np Ct Hp Wp -> B Np Ct Hp Wp", B=batch_size),
"B Np Ct Hp Wp -> B Ct (Np Hp Wp)",
)
def detokenize(self, tokens: torch.Tensor) -> torch.Tensor:
batch_size, Ct, Nt = tokens.shape
assert Nt == self.cfg.plane_size**2 * 3
assert Ct == self.cfg.num_channels
return rearrange(
tokens,
"B Ct (Np Hp Wp) -> B Np Ct Hp Wp",
Np=3,
Hp=self.cfg.plane_size,
Wp=self.cfg.plane_size,
)
+653
View File
@@ -0,0 +1,653 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
class Attention(nn.Module):
r"""
A cross attention layer.
Parameters:
query_dim (`int`):
The number of channels in the query.
cross_attention_dim (`int`, *optional*):
The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`.
heads (`int`, *optional*, defaults to 8):
The number of heads to use for multi-head attention.
dim_head (`int`, *optional*, defaults to 64):
The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0):
The dropout probability to use.
bias (`bool`, *optional*, defaults to False):
Set to `True` for the query, key, and value linear layers to contain a bias parameter.
upcast_attention (`bool`, *optional*, defaults to False):
Set to `True` to upcast the attention computation to `float32`.
upcast_softmax (`bool`, *optional*, defaults to False):
Set to `True` to upcast the softmax computation to `float32`.
cross_attention_norm (`str`, *optional*, defaults to `None`):
The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`.
cross_attention_norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups to use for the group norm in the cross attention.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
norm_num_groups (`int`, *optional*, defaults to `None`):
The number of groups to use for the group norm in the attention.
spatial_norm_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the spatial normalization.
out_bias (`bool`, *optional*, defaults to `True`):
Set to `True` to use a bias in the output linear layer.
scale_qk (`bool`, *optional*, defaults to `True`):
Set to `True` to scale the query and key by `1 / sqrt(dim_head)`.
only_cross_attention (`bool`, *optional*, defaults to `False`):
Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if
`added_kv_proj_dim` is not `None`.
eps (`float`, *optional*, defaults to 1e-5):
An additional value added to the denominator in group normalization that is used for numerical stability.
rescale_output_factor (`float`, *optional*, defaults to 1.0):
A factor to rescale the output by dividing it with this value.
residual_connection (`bool`, *optional*, defaults to `False`):
Set to `True` to add the residual connection to the output.
_from_deprecated_attn_block (`bool`, *optional*, defaults to `False`):
Set to `True` if the attention block is loaded from a deprecated state dict.
processor (`AttnProcessor`, *optional*, defaults to `None`):
The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and
`AttnProcessor` otherwise.
"""
def __init__(
self,
query_dim: int,
cross_attention_dim: Optional[int] = None,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias: bool = False,
upcast_attention: bool = False,
upcast_softmax: bool = False,
cross_attention_norm: Optional[str] = None,
cross_attention_norm_num_groups: int = 32,
added_kv_proj_dim: Optional[int] = None,
norm_num_groups: Optional[int] = None,
out_bias: bool = True,
scale_qk: bool = True,
only_cross_attention: bool = False,
eps: float = 1e-5,
rescale_output_factor: float = 1.0,
residual_connection: bool = False,
_from_deprecated_attn_block: bool = False,
processor: Optional["AttnProcessor"] = None,
out_dim: int = None,
):
super().__init__()
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.query_dim = query_dim
self.cross_attention_dim = (
cross_attention_dim if cross_attention_dim is not None else query_dim
)
self.upcast_attention = upcast_attention
self.upcast_softmax = upcast_softmax
self.rescale_output_factor = rescale_output_factor
self.residual_connection = residual_connection
self.dropout = dropout
self.fused_projections = False
self.out_dim = out_dim if out_dim is not None else query_dim
# we make use of this private variable to know whether this class is loaded
# with an deprecated state dict so that we can convert it on the fly
self._from_deprecated_attn_block = _from_deprecated_attn_block
self.scale_qk = scale_qk
self.scale = dim_head**-0.5 if self.scale_qk else 1.0
self.heads = out_dim // dim_head if out_dim is not None else heads
# for slice_size > 0 the attention score computation
# is split across the batch axis to save memory
# You can set slice_size with `set_attention_slice`
self.sliceable_head_dim = heads
self.added_kv_proj_dim = added_kv_proj_dim
self.only_cross_attention = only_cross_attention
if self.added_kv_proj_dim is None and self.only_cross_attention:
raise ValueError(
"`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`."
)
if norm_num_groups is not None:
self.group_norm = nn.GroupNorm(
num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True
)
else:
self.group_norm = None
self.spatial_norm = None
if cross_attention_norm is None:
self.norm_cross = None
elif cross_attention_norm == "layer_norm":
self.norm_cross = nn.LayerNorm(self.cross_attention_dim)
elif cross_attention_norm == "group_norm":
if self.added_kv_proj_dim is not None:
# The given `encoder_hidden_states` are initially of shape
# (batch_size, seq_len, added_kv_proj_dim) before being projected
# to (batch_size, seq_len, cross_attention_dim). The norm is applied
# before the projection, so we need to use `added_kv_proj_dim` as
# the number of channels for the group norm.
norm_cross_num_channels = added_kv_proj_dim
else:
norm_cross_num_channels = self.cross_attention_dim
self.norm_cross = nn.GroupNorm(
num_channels=norm_cross_num_channels,
num_groups=cross_attention_norm_num_groups,
eps=1e-5,
affine=True,
)
else:
raise ValueError(
f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'"
)
linear_cls = nn.Linear
self.linear_cls = linear_cls
self.to_q = linear_cls(query_dim, self.inner_dim, bias=bias)
if not self.only_cross_attention:
# only relevant for the `AddedKVProcessor` classes
self.to_k = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
self.to_v = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
else:
self.to_k = None
self.to_v = None
if self.added_kv_proj_dim is not None:
self.add_k_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
self.add_v_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
self.to_out = nn.ModuleList([])
self.to_out.append(linear_cls(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
# set attention processor
# We use the AttnProcessor2_0 by default when torch 2.x is used which uses
# torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention
# but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1
if processor is None:
processor = (
AttnProcessor2_0()
if hasattr(F, "scaled_dot_product_attention") and self.scale_qk
else AttnProcessor()
)
self.set_processor(processor)
def set_processor(self, processor: "AttnProcessor") -> None:
self.processor = processor
def forward(
self,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
**cross_attention_kwargs,
) -> torch.Tensor:
r"""
The forward method of the `Attention` class.
Args:
hidden_states (`torch.Tensor`):
The hidden states of the query.
encoder_hidden_states (`torch.Tensor`, *optional*):
The hidden states of the encoder.
attention_mask (`torch.Tensor`, *optional*):
The attention mask to use. If `None`, no mask is applied.
**cross_attention_kwargs:
Additional keyword arguments to pass along to the cross attention.
Returns:
`torch.Tensor`: The output of the attention layer.
"""
# The `Attention` class can call different attention processors / attention functions
# here we simply pass along all tensors to the selected processor class
# For standard processors that are defined here, `**cross_attention_kwargs` is empty
return self.processor(
self,
hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=attention_mask,
**cross_attention_kwargs,
)
def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor:
r"""
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads`
is the number of heads initialized while constructing the `Attention` class.
Args:
tensor (`torch.Tensor`): The tensor to reshape.
Returns:
`torch.Tensor`: The reshaped tensor.
"""
head_size = self.heads
batch_size, seq_len, dim = tensor.shape
tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim)
tensor = tensor.permute(0, 2, 1, 3).reshape(
batch_size // head_size, seq_len, dim * head_size
)
return tensor
def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor:
r"""
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is
the number of heads initialized while constructing the `Attention` class.
Args:
tensor (`torch.Tensor`): The tensor to reshape.
out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is
reshaped to `[batch_size * heads, seq_len, dim // heads]`.
Returns:
`torch.Tensor`: The reshaped tensor.
"""
head_size = self.heads
batch_size, seq_len, dim = tensor.shape
tensor = tensor.reshape(batch_size, seq_len, head_size, dim // head_size)
tensor = tensor.permute(0, 2, 1, 3)
if out_dim == 3:
tensor = tensor.reshape(batch_size * head_size, seq_len, dim // head_size)
return tensor
def get_attention_scores(
self,
query: torch.Tensor,
key: torch.Tensor,
attention_mask: torch.Tensor = None,
) -> torch.Tensor:
r"""
Compute the attention scores.
Args:
query (`torch.Tensor`): The query tensor.
key (`torch.Tensor`): The key tensor.
attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied.
Returns:
`torch.Tensor`: The attention probabilities/scores.
"""
dtype = query.dtype
if self.upcast_attention:
query = query.float()
key = key.float()
if attention_mask is None:
baddbmm_input = torch.empty(
query.shape[0],
query.shape[1],
key.shape[1],
dtype=query.dtype,
device=query.device,
)
beta = 0
else:
baddbmm_input = attention_mask
beta = 1
attention_scores = torch.baddbmm(
baddbmm_input,
query,
key.transpose(-1, -2),
beta=beta,
alpha=self.scale,
)
del baddbmm_input
if self.upcast_softmax:
attention_scores = attention_scores.float()
attention_probs = attention_scores.softmax(dim=-1)
del attention_scores
attention_probs = attention_probs.to(dtype)
return attention_probs
def prepare_attention_mask(
self,
attention_mask: torch.Tensor,
target_length: int,
batch_size: int,
out_dim: int = 3,
) -> torch.Tensor:
r"""
Prepare the attention mask for the attention computation.
Args:
attention_mask (`torch.Tensor`):
The attention mask to prepare.
target_length (`int`):
The target length of the attention mask. This is the length of the attention mask after padding.
batch_size (`int`):
The batch size, which is used to repeat the attention mask.
out_dim (`int`, *optional*, defaults to `3`):
The output dimension of the attention mask. Can be either `3` or `4`.
Returns:
`torch.Tensor`: The prepared attention mask.
"""
head_size = self.heads
if attention_mask is None:
return attention_mask
current_length: int = attention_mask.shape[-1]
if current_length != target_length:
if attention_mask.device.type == "mps":
# HACK: MPS: Does not support padding by greater than dimension of input tensor.
# Instead, we can manually construct the padding tensor.
padding_shape = (
attention_mask.shape[0],
attention_mask.shape[1],
target_length,
)
padding = torch.zeros(
padding_shape,
dtype=attention_mask.dtype,
device=attention_mask.device,
)
attention_mask = torch.cat([attention_mask, padding], dim=2)
else:
# TODO: for pipelines such as stable-diffusion, padding cross-attn mask:
# we want to instead pad by (0, remaining_length), where remaining_length is:
# remaining_length: int = target_length - current_length
# TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
if out_dim == 3:
if attention_mask.shape[0] < batch_size * head_size:
attention_mask = attention_mask.repeat_interleave(head_size, dim=0)
elif out_dim == 4:
attention_mask = attention_mask.unsqueeze(1)
attention_mask = attention_mask.repeat_interleave(head_size, dim=1)
return attention_mask
def norm_encoder_hidden_states(
self, encoder_hidden_states: torch.Tensor
) -> torch.Tensor:
r"""
Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the
`Attention` class.
Args:
encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder.
Returns:
`torch.Tensor`: The normalized encoder hidden states.
"""
assert (
self.norm_cross is not None
), "self.norm_cross must be defined to call self.norm_encoder_hidden_states"
if isinstance(self.norm_cross, nn.LayerNorm):
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
elif isinstance(self.norm_cross, nn.GroupNorm):
# Group norm norms along the channels dimension and expects
# input to be in the shape of (N, C, *). In this case, we want
# to norm along the hidden dimension, so we need to move
# (batch_size, sequence_length, hidden_size) ->
# (batch_size, hidden_size, sequence_length)
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
else:
assert False
return encoder_hidden_states
@torch.no_grad()
def fuse_projections(self, fuse=True):
is_cross_attention = self.cross_attention_dim != self.query_dim
device = self.to_q.weight.data.device
dtype = self.to_q.weight.data.dtype
if not is_cross_attention:
# fetch weight matrices.
concatenated_weights = torch.cat(
[self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]
)
in_features = concatenated_weights.shape[1]
out_features = concatenated_weights.shape[0]
# create a new single projection layer and copy over the weights.
self.to_qkv = self.linear_cls(
in_features, out_features, bias=False, device=device, dtype=dtype
)
self.to_qkv.weight.copy_(concatenated_weights)
else:
concatenated_weights = torch.cat(
[self.to_k.weight.data, self.to_v.weight.data]
)
in_features = concatenated_weights.shape[1]
out_features = concatenated_weights.shape[0]
self.to_kv = self.linear_cls(
in_features, out_features, bias=False, device=device, dtype=dtype
)
self.to_kv.weight.copy_(concatenated_weights)
self.fused_projections = fuse
class AttnProcessor:
r"""
Default processor for performing attention-related computations.
"""
def __call__(
self,
attn: Attention,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
residual = hidden_states
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(
batch_size, channel, height * width
).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape
if encoder_hidden_states is None
else encoder_hidden_states.shape
)
attention_mask = attn.prepare_attention_mask(
attention_mask, sequence_length, batch_size
)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
1, 2
)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(
encoder_hidden_states
)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
query = attn.head_to_batch_dim(query)
key = attn.head_to_batch_dim(key)
value = attn.head_to_batch_dim(value)
attention_probs = attn.get_attention_scores(query, key, attention_mask)
hidden_states = torch.bmm(attention_probs, value)
hidden_states = attn.batch_to_head_dim(hidden_states)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(
batch_size, channel, height, width
)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class AttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0."
)
def __call__(
self,
attn: Attention,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
residual = hidden_states
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(
batch_size, channel, height * width
).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape
if encoder_hidden_states is None
else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(
attention_mask, sequence_length, batch_size
)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(
batch_size, attn.heads, -1, attention_mask.shape[-1]
)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
1, 2
)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(
encoder_hidden_states
)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(
batch_size, -1, attn.heads * head_dim
)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(
batch_size, channel, height, width
)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
@@ -0,0 +1,334 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
from .attention import Attention
class BasicTransformerBlock(nn.Module):
r"""
A basic Transformer block.
Parameters:
dim (`int`): The number of channels in the input and output.
num_attention_heads (`int`): The number of heads to use for multi-head attention.
attention_head_dim (`int`): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
attention_bias (:
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
only_cross_attention (`bool`, *optional*):
Whether to use only cross-attention layers. In this case two cross attention layers are used.
double_self_attention (`bool`, *optional*):
Whether to use two self-attention layers. In this case no cross attention layers are used.
upcast_attention (`bool`, *optional*):
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
"""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
dropout=0.0,
cross_attention_dim: Optional[int] = None,
activation_fn: str = "geglu",
attention_bias: bool = False,
only_cross_attention: bool = False,
double_self_attention: bool = False,
upcast_attention: bool = False,
norm_elementwise_affine: bool = True,
norm_type: str = "layer_norm",
final_dropout: bool = False,
):
super().__init__()
self.only_cross_attention = only_cross_attention
assert norm_type == "layer_norm"
# Define 3 blocks. Each block has its own normalization layer.
# 1. Self-Attn
self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.attn1 = Attention(
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
upcast_attention=upcast_attention,
)
# 2. Cross-Attn
if cross_attention_dim is not None or double_self_attention:
# We currently only use AdaLayerNormZero for self attention where there will only be one attention block.
# I.e. the number of returned modulation chunks from AdaLayerZero would not make sense if returned during
# the second cross attention block.
self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.attn2 = Attention(
query_dim=dim,
cross_attention_dim=(
cross_attention_dim if not double_self_attention else None
),
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
) # is self-attn if encoder_hidden_states is none
else:
self.norm2 = None
self.attn2 = None
# 3. Feed-forward
self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
)
# let chunk size default to None
self._chunk_size = None
self._chunk_dim = 0
def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int):
# Sets chunk feed-forward
self._chunk_size = chunk_size
self._chunk_dim = dim
def forward(
self,
hidden_states: torch.FloatTensor,
attention_mask: Optional[torch.FloatTensor] = None,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
encoder_attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
# Notice that normalization is always applied before the real computation in the following blocks.
# 0. Self-Attention
norm_hidden_states = self.norm1(hidden_states)
attn_output = self.attn1(
norm_hidden_states,
encoder_hidden_states=(
encoder_hidden_states if self.only_cross_attention else None
),
attention_mask=attention_mask,
)
hidden_states = attn_output + hidden_states
# 3. Cross-Attention
if self.attn2 is not None:
norm_hidden_states = self.norm2(hidden_states)
attn_output = self.attn2(
norm_hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=encoder_attention_mask,
)
hidden_states = attn_output + hidden_states
# 4. Feed-forward
norm_hidden_states = self.norm3(hidden_states)
if self._chunk_size is not None:
# "feed_forward_chunk_size" can be used to save memory
if norm_hidden_states.shape[self._chunk_dim] % self._chunk_size != 0:
raise ValueError(
f"`hidden_states` dimension to be chunked: {norm_hidden_states.shape[self._chunk_dim]} has to be divisible by chunk size: {self._chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`."
)
num_chunks = norm_hidden_states.shape[self._chunk_dim] // self._chunk_size
ff_output = torch.cat(
[
self.ff(hid_slice)
for hid_slice in norm_hidden_states.chunk(
num_chunks, dim=self._chunk_dim
)
],
dim=self._chunk_dim,
)
else:
ff_output = self.ff(norm_hidden_states)
hidden_states = ff_output + hidden_states
return hidden_states
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
activation_fn: str = "geglu",
final_dropout: bool = False,
):
super().__init__()
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
linear_cls = nn.Linear
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh")
elif activation_fn == "geglu":
act_fn = GEGLU(dim, inner_dim)
elif activation_fn == "geglu-approximate":
act_fn = ApproximateGELU(dim, inner_dim)
self.net = nn.ModuleList([])
# project in
self.net.append(act_fn)
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(linear_cls(inner_dim, dim_out))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
"""
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none"):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out)
self.approximate = approximate
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate, approximate=self.approximate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(
dtype=gate.dtype
)
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states = self.gelu(hidden_states)
return hidden_states
class GEGLU(nn.Module):
r"""
A variant of the gated linear unit activation function from https://arxiv.org/abs/2002.05202.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
"""
def __init__(self, dim_in: int, dim_out: int):
super().__init__()
linear_cls = nn.Linear
self.proj = linear_cls(dim_in, dim_out * 2)
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
def forward(self, hidden_states, scale: float = 1.0):
args = ()
hidden_states, gate = self.proj(hidden_states, *args).chunk(2, dim=-1)
return hidden_states * self.gelu(gate)
class ApproximateGELU(nn.Module):
r"""
The approximate form of Gaussian Error Linear Unit (GELU). For more details, see section 2:
https://arxiv.org/abs/1606.08415.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
"""
def __init__(self, dim_in: int, dim_out: int):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.proj(x)
return x * torch.sigmoid(1.702 * x)
@@ -0,0 +1,219 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from dataclasses import dataclass
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
from ...utils import BaseModule
from .basic_transformer_block import BasicTransformerBlock
class Transformer1D(BaseModule):
@dataclass
class Config(BaseModule.Config):
num_attention_heads: int = 16
attention_head_dim: int = 88
in_channels: Optional[int] = None
out_channels: Optional[int] = None
num_layers: int = 1
dropout: float = 0.0
norm_num_groups: int = 32
cross_attention_dim: Optional[int] = None
attention_bias: bool = False
activation_fn: str = "geglu"
only_cross_attention: bool = False
double_self_attention: bool = False
upcast_attention: bool = False
norm_type: str = "layer_norm"
norm_elementwise_affine: bool = True
gradient_checkpointing: bool = False
cfg: Config
def configure(self) -> None:
self.num_attention_heads = self.cfg.num_attention_heads
self.attention_head_dim = self.cfg.attention_head_dim
inner_dim = self.num_attention_heads * self.attention_head_dim
linear_cls = nn.Linear
# 2. Define input layers
self.in_channels = self.cfg.in_channels
self.norm = torch.nn.GroupNorm(
num_groups=self.cfg.norm_num_groups,
num_channels=self.cfg.in_channels,
eps=1e-6,
affine=True,
)
self.proj_in = linear_cls(self.cfg.in_channels, inner_dim)
# 3. Define transformers blocks
self.transformer_blocks = nn.ModuleList(
[
BasicTransformerBlock(
inner_dim,
self.num_attention_heads,
self.attention_head_dim,
dropout=self.cfg.dropout,
cross_attention_dim=self.cfg.cross_attention_dim,
activation_fn=self.cfg.activation_fn,
attention_bias=self.cfg.attention_bias,
only_cross_attention=self.cfg.only_cross_attention,
double_self_attention=self.cfg.double_self_attention,
upcast_attention=self.cfg.upcast_attention,
norm_type=self.cfg.norm_type,
norm_elementwise_affine=self.cfg.norm_elementwise_affine,
)
for d in range(self.cfg.num_layers)
]
)
# 4. Define output layers
self.out_channels = (
self.cfg.in_channels
if self.cfg.out_channels is None
else self.cfg.out_channels
)
self.proj_out = linear_cls(inner_dim, self.cfg.in_channels)
self.gradient_checkpointing = self.cfg.gradient_checkpointing
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
):
"""
The [`Transformer1DModel`] forward method.
Args:
hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.FloatTensor` of shape `(batch size, channel, height, width)` if continuous):
Input `hidden_states`.
encoder_hidden_states ( `torch.FloatTensor` of shape `(batch size, sequence len, embed dims)`, *optional*):
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
self-attention.
attention_mask ( `torch.Tensor`, *optional*):
An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask
is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large
negative values to the attention scores corresponding to "discard" tokens.
encoder_attention_mask ( `torch.Tensor`, *optional*):
Cross-attention mask applied to `encoder_hidden_states`. Two formats supported:
* Mask `(batch, sequence_length)` True = keep, False = discard.
* Bias `(batch, 1, sequence_length)` 0 = keep, -10000 = discard.
If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format
above. This bias will be added to the cross-attention scores.
Returns:
torch.FloatTensor
"""
# ensure attention_mask is a bias, and give it a singleton query_tokens dimension.
# we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward.
# we can tell by counting dims; if ndim == 2: it's a mask rather than a bias.
# expects mask of shape:
# [batch, key_tokens]
# adds singleton query_tokens dimension:
# [batch, 1, key_tokens]
# this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes:
# [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn)
# [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn)
if attention_mask is not None and attention_mask.ndim == 2:
# assume that mask is expressed as:
# (1 = keep, 0 = discard)
# convert mask into a bias that can be added to attention scores:
# (keep = +0, discard = -10000.0)
attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0
attention_mask = attention_mask.unsqueeze(1)
# convert encoder_attention_mask to a bias the same way we do for attention_mask
if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2:
encoder_attention_mask = (
1 - encoder_attention_mask.to(hidden_states.dtype)
) * -10000.0
encoder_attention_mask = encoder_attention_mask.unsqueeze(1)
# 1. Input
batch, _, seq_len = hidden_states.shape
residual = hidden_states
hidden_states = self.norm(hidden_states)
inner_dim = hidden_states.shape[1]
hidden_states = hidden_states.permute(0, 2, 1).reshape(
batch, seq_len, inner_dim
)
hidden_states = self.proj_in(hidden_states)
# 2. Blocks
for block in self.transformer_blocks:
if self.training and self.gradient_checkpointing:
hidden_states = torch.utils.checkpoint.checkpoint(
block,
hidden_states,
attention_mask,
encoder_hidden_states,
encoder_attention_mask,
use_reentrant=False,
)
else:
hidden_states = block(
hidden_states,
attention_mask=attention_mask,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
)
# 3. Output
hidden_states = self.proj_out(hidden_states)
hidden_states = (
hidden_states.reshape(batch, seq_len, inner_dim)
.permute(0, 2, 1)
.contiguous()
)
output = hidden_states + residual
return output
+218
View File
@@ -0,0 +1,218 @@
import math
import os
from dataclasses import dataclass, field
from typing import List, Union
import numpy as np
import PIL.Image
import torch
import torch.nn.functional as F
import trimesh
from einops import rearrange
from huggingface_hub import hf_hub_download
from omegaconf import OmegaConf
from PIL import Image
from .models.isosurface import MarchingCubeHelper
from .utils import (
BaseModule,
ImagePreprocessor,
find_class,
get_spherical_cameras,
scale_tensor,
)
class TSR(BaseModule):
@dataclass
class Config(BaseModule.Config):
cond_image_size: int
image_tokenizer_cls: str
image_tokenizer: dict
tokenizer_cls: str
tokenizer: dict
backbone_cls: str
backbone: dict
post_processor_cls: str
post_processor: dict
decoder_cls: str
decoder: dict
renderer_cls: str
renderer: dict
cfg: Config
@classmethod
def from_pretrained(
cls, pretrained_model_name_or_path: str, config_name: str, weight_name: str
):
if os.path.isdir(pretrained_model_name_or_path):
config_path = os.path.join(pretrained_model_name_or_path, config_name)
weight_path = os.path.join(pretrained_model_name_or_path, weight_name)
else:
config_path = hf_hub_download(
repo_id=pretrained_model_name_or_path, filename=config_name
)
weight_path = hf_hub_download(
repo_id=pretrained_model_name_or_path, filename=weight_name
)
cfg = OmegaConf.load(config_path)
OmegaConf.resolve(cfg)
model = cls(cfg)
ckpt = torch.load(weight_path, map_location="cpu")
model.load_state_dict(ckpt)
return model
@classmethod
def from_pretrained_custom(
cls, weight_path: str, config_path: str
):
cfg = OmegaConf.load(config_path)
OmegaConf.resolve(cfg)
model = cls(cfg)
ckpt = torch.load(weight_path, map_location="cpu")
model.load_state_dict(ckpt)
return model
def configure(self):
self.image_tokenizer = find_class(self.cfg.image_tokenizer_cls)(
self.cfg.image_tokenizer
)
self.tokenizer = find_class(self.cfg.tokenizer_cls)(self.cfg.tokenizer)
self.backbone = find_class(self.cfg.backbone_cls)(self.cfg.backbone)
self.post_processor = find_class(self.cfg.post_processor_cls)(
self.cfg.post_processor
)
self.decoder = find_class(self.cfg.decoder_cls)(self.cfg.decoder)
self.renderer = find_class(self.cfg.renderer_cls)(self.cfg.renderer)
self.image_processor = ImagePreprocessor()
self.isosurface_helper = None
def forward(
self,
image: Union[
PIL.Image.Image,
np.ndarray,
torch.FloatTensor,
List[PIL.Image.Image],
List[np.ndarray],
List[torch.FloatTensor],
],
device: str,
) -> torch.FloatTensor:
rgb_cond = self.image_processor(image, self.cfg.cond_image_size)[:, None].to(
device
)
batch_size = rgb_cond.shape[0]
input_image_tokens: torch.Tensor = self.image_tokenizer(
rearrange(rgb_cond, "B Nv H W C -> B Nv C H W", Nv=1),
)
input_image_tokens = rearrange(
input_image_tokens, "B Nv C Nt -> B (Nv Nt) C", Nv=1
)
tokens: torch.Tensor = self.tokenizer(batch_size)
tokens = self.backbone(
tokens,
encoder_hidden_states=input_image_tokens,
)
scene_codes = self.post_processor(self.tokenizer.detokenize(tokens))
return scene_codes
def render(
self,
scene_codes,
n_views: int,
elevation_deg: float = 0.0,
camera_distance: float = 1.9,
fovy_deg: float = 40.0,
height: int = 256,
width: int = 256,
return_type: str = "pil",
):
rays_o, rays_d = get_spherical_cameras(
n_views, elevation_deg, camera_distance, fovy_deg, height, width
)
rays_o, rays_d = rays_o.to(scene_codes.device), rays_d.to(scene_codes.device)
def process_output(image: torch.FloatTensor):
if return_type == "pt":
return image
elif return_type == "np":
return image.detach().cpu().numpy()
elif return_type == "pil":
return Image.fromarray(
(image.detach().cpu().numpy() * 255.0).astype(np.uint8)
)
else:
raise NotImplementedError
images = []
for scene_code in scene_codes:
images_ = []
for i in range(n_views):
with torch.no_grad():
image = self.renderer(
self.decoder, scene_code, rays_o[i], rays_d[i]
)
images_.append(process_output(image))
images.append(images_)
return images
def set_marching_cubes_resolution(self, resolution: int):
if (
self.isosurface_helper is not None
and self.isosurface_helper.resolution == resolution
):
return
self.isosurface_helper = MarchingCubeHelper(resolution)
def extract_mesh(self, scene_codes, resolution: int = 256, threshold: float = 25.0,callback=None):
self.set_marching_cubes_resolution(resolution)
meshes = []
for scene_code in scene_codes:
with torch.no_grad():
density = self.renderer.query_triplane(
self.decoder,
scale_tensor(
self.isosurface_helper.grid_vertices.to(scene_codes.device),
self.isosurface_helper.points_range,
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
),
scene_code,
)["density_act"]
v_pos, t_pos_idx = self.isosurface_helper(-(density - threshold))
v_pos = scale_tensor(
v_pos,
self.isosurface_helper.points_range,
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
)
with torch.no_grad():
color = self.renderer.query_triplane(
self.decoder,
v_pos,
scene_code,
)["color"]
mesh = trimesh.Trimesh(
vertices=v_pos.cpu().numpy(),
faces=t_pos_idx.cpu().numpy(),
vertex_colors=color.cpu().numpy(),
)
meshes.append(mesh)
if callback:
callback(len(meshes))
return meshes
+475
View File
@@ -0,0 +1,475 @@
import importlib
import math
from collections import defaultdict
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import imageio
import numpy as np
import PIL.Image
#import rembg
import torch
import torch.nn as nn
import torch.nn.functional as F
import trimesh
from omegaconf import DictConfig, OmegaConf
#from PIL import Image
def parse_structured(fields: Any, cfg: Optional[Union[dict, DictConfig]] = None) -> Any:
scfg = OmegaConf.merge(OmegaConf.structured(fields), cfg)
return scfg
def find_class(cls_string):
module_string = ".".join(cls_string.split(".")[:-1])
cls_name = cls_string.split(".")[-1]
module = importlib.import_module(module_string, package=None)
cls = getattr(module, cls_name)
return cls
def get_intrinsic_from_fov(fov, H, W, bs=-1):
focal_length = 0.5 * H / np.tan(0.5 * fov)
intrinsic = np.identity(3, dtype=np.float32)
intrinsic[0, 0] = focal_length
intrinsic[1, 1] = focal_length
intrinsic[0, 2] = W / 2.0
intrinsic[1, 2] = H / 2.0
if bs > 0:
intrinsic = intrinsic[None].repeat(bs, axis=0)
return torch.from_numpy(intrinsic)
class BaseModule(nn.Module):
@dataclass
class Config:
pass
cfg: Config # add this to every subclass of BaseModule to enable static type checking
def __init__(
self, cfg: Optional[Union[dict, DictConfig]] = None, *args, **kwargs
) -> None:
super().__init__()
self.cfg = parse_structured(self.Config, cfg)
self.configure(*args, **kwargs)
def configure(self, *args, **kwargs) -> None:
raise NotImplementedError
class ImagePreprocessor:
def convert_and_resize(
self,
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
size: int,
):
if isinstance(image, PIL.Image.Image):
image = torch.from_numpy(np.array(image).astype(np.float32) / 255.0)
elif isinstance(image, np.ndarray):
if image.dtype == np.uint8:
image = torch.from_numpy(image.astype(np.float32) / 255.0)
else:
image = torch.from_numpy(image)
elif isinstance(image, torch.Tensor):
pass
batched = image.ndim == 4
if not batched:
image = image[None, ...]
image = F.interpolate(
image.permute(0, 3, 1, 2),
(size, size),
mode="bilinear",
align_corners=False,
antialias=True,
).permute(0, 2, 3, 1)
if not batched:
image = image[0]
return image
def __call__(
self,
image: Union[
PIL.Image.Image,
np.ndarray,
torch.FloatTensor,
List[PIL.Image.Image],
List[np.ndarray],
List[torch.FloatTensor],
],
size: int,
) -> Any:
if isinstance(image, (np.ndarray, torch.FloatTensor)) and image.ndim == 4:
image = self.convert_and_resize(image, size)
else:
if not isinstance(image, list):
image = [image]
image = [self.convert_and_resize(im, size) for im in image]
image = torch.stack(image, dim=0)
return image
def rays_intersect_bbox(
rays_o: torch.Tensor,
rays_d: torch.Tensor,
radius: float,
near: float = 0.0,
valid_thresh: float = 0.01,
):
input_shape = rays_o.shape[:-1]
rays_o, rays_d = rays_o.view(-1, 3), rays_d.view(-1, 3)
rays_d_valid = torch.where(
rays_d.abs() < 1e-6, torch.full_like(rays_d, 1e-6), rays_d
)
if type(radius) in [int, float]:
radius = torch.FloatTensor(
[[-radius, radius], [-radius, radius], [-radius, radius]]
).to(rays_o.device)
radius = (
1.0 - 1.0e-3
) * radius # tighten the radius to make sure the intersection point lies in the bounding box
interx0 = (radius[..., 1] - rays_o) / rays_d_valid
interx1 = (radius[..., 0] - rays_o) / rays_d_valid
t_near = torch.minimum(interx0, interx1).amax(dim=-1).clamp_min(near)
t_far = torch.maximum(interx0, interx1).amin(dim=-1)
# check wheter a ray intersects the bbox or not
rays_valid = t_far - t_near > valid_thresh
t_near[torch.where(~rays_valid)] = 0.0
t_far[torch.where(~rays_valid)] = 0.0
t_near = t_near.view(*input_shape, 1)
t_far = t_far.view(*input_shape, 1)
rays_valid = rays_valid.view(*input_shape)
return t_near, t_far, rays_valid
def chunk_batch(func: Callable, chunk_size: int, *args, **kwargs) -> Any:
if chunk_size <= 0:
return func(*args, **kwargs)
B = None
for arg in list(args) + list(kwargs.values()):
if isinstance(arg, torch.Tensor):
B = arg.shape[0]
break
assert (
B is not None
), "No tensor found in args or kwargs, cannot determine batch size."
out = defaultdict(list)
out_type = None
# max(1, B) to support B == 0
for i in range(0, max(1, B), chunk_size):
out_chunk = func(
*[
arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
for arg in args
],
**{
k: arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
for k, arg in kwargs.items()
},
)
if out_chunk is None:
continue
out_type = type(out_chunk)
if isinstance(out_chunk, torch.Tensor):
out_chunk = {0: out_chunk}
elif isinstance(out_chunk, tuple) or isinstance(out_chunk, list):
chunk_length = len(out_chunk)
out_chunk = {i: chunk for i, chunk in enumerate(out_chunk)}
elif isinstance(out_chunk, dict):
pass
else:
print(
f"Return value of func must be in type [torch.Tensor, list, tuple, dict], get {type(out_chunk)}."
)
exit(1)
for k, v in out_chunk.items():
v = v if torch.is_grad_enabled() else v.detach()
out[k].append(v)
if out_type is None:
return None
out_merged: Dict[Any, Optional[torch.Tensor]] = {}
for k, v in out.items():
if all([vv is None for vv in v]):
# allow None in return value
out_merged[k] = None
elif all([isinstance(vv, torch.Tensor) for vv in v]):
out_merged[k] = torch.cat(v, dim=0)
else:
raise TypeError(
f"Unsupported types in return value of func: {[type(vv) for vv in v if not isinstance(vv, torch.Tensor)]}"
)
if out_type is torch.Tensor:
return out_merged[0]
elif out_type in [tuple, list]:
return out_type([out_merged[i] for i in range(chunk_length)])
elif out_type is dict:
return out_merged
ValidScale = Union[Tuple[float, float], torch.FloatTensor]
def scale_tensor(dat: torch.FloatTensor, inp_scale: ValidScale, tgt_scale: ValidScale):
if inp_scale is None:
inp_scale = (0, 1)
if tgt_scale is None:
tgt_scale = (0, 1)
if isinstance(tgt_scale, torch.FloatTensor):
assert dat.shape[-1] == tgt_scale.shape[-1]
dat = (dat - inp_scale[0]) / (inp_scale[1] - inp_scale[0])
dat = dat * (tgt_scale[1] - tgt_scale[0]) + tgt_scale[0]
return dat
def get_activation(name) -> Callable:
if name is None:
return lambda x: x
name = name.lower()
if name == "none":
return lambda x: x
elif name == "exp":
return lambda x: torch.exp(x)
elif name == "sigmoid":
return lambda x: torch.sigmoid(x)
elif name == "tanh":
return lambda x: torch.tanh(x)
elif name == "softplus":
return lambda x: F.softplus(x)
else:
try:
return getattr(F, name)
except AttributeError:
raise ValueError(f"Unknown activation function: {name}")
def get_ray_directions(
H: int,
W: int,
focal: Union[float, Tuple[float, float]],
principal: Optional[Tuple[float, float]] = None,
use_pixel_centers: bool = True,
normalize: bool = True,
) -> torch.FloatTensor:
"""
Get ray directions for all pixels in camera coordinate.
Reference: https://www.scratchapixel.com/lessons/3d-basic-rendering/
ray-tracing-generating-camera-rays/standard-coordinate-systems
Inputs:
H, W, focal, principal, use_pixel_centers: image height, width, focal length, principal point and whether use pixel centers
Outputs:
directions: (H, W, 3), the direction of the rays in camera coordinate
"""
pixel_center = 0.5 if use_pixel_centers else 0
if isinstance(focal, float):
fx, fy = focal, focal
cx, cy = W / 2, H / 2
else:
fx, fy = focal
assert principal is not None
cx, cy = principal
i, j = torch.meshgrid(
torch.arange(W, dtype=torch.float32) + pixel_center,
torch.arange(H, dtype=torch.float32) + pixel_center,
indexing="xy",
)
directions = torch.stack([(i - cx) / fx, -(j - cy) / fy, -torch.ones_like(i)], -1)
if normalize:
directions = F.normalize(directions, dim=-1)
return directions
def get_rays(
directions,
c2w,
keepdim=False,
normalize=False,
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
# Rotate ray directions from camera coordinate to the world coordinate
assert directions.shape[-1] == 3
if directions.ndim == 2: # (N_rays, 3)
if c2w.ndim == 2: # (4, 4)
c2w = c2w[None, :, :]
assert c2w.ndim == 3 # (N_rays, 4, 4) or (1, 4, 4)
rays_d = (directions[:, None, :] * c2w[:, :3, :3]).sum(-1) # (N_rays, 3)
rays_o = c2w[:, :3, 3].expand(rays_d.shape)
elif directions.ndim == 3: # (H, W, 3)
assert c2w.ndim in [2, 3]
if c2w.ndim == 2: # (4, 4)
rays_d = (directions[:, :, None, :] * c2w[None, None, :3, :3]).sum(
-1
) # (H, W, 3)
rays_o = c2w[None, None, :3, 3].expand(rays_d.shape)
elif c2w.ndim == 3: # (B, 4, 4)
rays_d = (directions[None, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
-1
) # (B, H, W, 3)
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
elif directions.ndim == 4: # (B, H, W, 3)
assert c2w.ndim == 3 # (B, 4, 4)
rays_d = (directions[:, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
-1
) # (B, H, W, 3)
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
if normalize:
rays_d = F.normalize(rays_d, dim=-1)
if not keepdim:
rays_o, rays_d = rays_o.reshape(-1, 3), rays_d.reshape(-1, 3)
return rays_o, rays_d
def get_spherical_cameras(
n_views: int,
elevation_deg: float,
camera_distance: float,
fovy_deg: float,
height: int,
width: int,
):
azimuth_deg = torch.linspace(0, 360.0, n_views + 1)[:n_views]
elevation_deg = torch.full_like(azimuth_deg, elevation_deg)
camera_distances = torch.full_like(elevation_deg, camera_distance)
elevation = elevation_deg * math.pi / 180
azimuth = azimuth_deg * math.pi / 180
# convert spherical coordinates to cartesian coordinates
# right hand coordinate system, x back, y right, z up
# elevation in (-90, 90), azimuth from +x to +y in (-180, 180)
camera_positions = torch.stack(
[
camera_distances * torch.cos(elevation) * torch.cos(azimuth),
camera_distances * torch.cos(elevation) * torch.sin(azimuth),
camera_distances * torch.sin(elevation),
],
dim=-1,
)
# default scene center at origin
center = torch.zeros_like(camera_positions)
# default camera up direction as +z
up = torch.as_tensor([0, 0, 1], dtype=torch.float32)[None, :].repeat(n_views, 1)
fovy = torch.full_like(elevation_deg, fovy_deg) * math.pi / 180
lookat = F.normalize(center - camera_positions, dim=-1)
right = F.normalize(torch.cross(lookat, up), dim=-1)
up = F.normalize(torch.cross(right, lookat), dim=-1)
c2w3x4 = torch.cat(
[torch.stack([right, up, -lookat], dim=-1), camera_positions[:, :, None]],
dim=-1,
)
c2w = torch.cat([c2w3x4, torch.zeros_like(c2w3x4[:, :1])], dim=1)
c2w[:, 3, 3] = 1.0
# get directions by dividing directions_unit_focal by focal length
focal_length = 0.5 * height / torch.tan(0.5 * fovy)
directions_unit_focal = get_ray_directions(
H=height,
W=width,
focal=1.0,
)
directions = directions_unit_focal[None, :, :, :].repeat(n_views, 1, 1, 1)
directions[:, :, :, :2] = (
directions[:, :, :, :2] / focal_length[:, None, None, None]
)
# must use normalize=True to normalize directions here
rays_o, rays_d = get_rays(directions, c2w, keepdim=True, normalize=True)
return rays_o, rays_d
# def remove_background(
# image: PIL.Image.Image,
# rembg_session: Any = None,
# force: bool = False,
# **rembg_kwargs,
# ) -> PIL.Image.Image:
# do_remove = True
# if image.mode == "RGBA" and image.getextrema()[3][0] < 255:
# do_remove = False
# do_remove = do_remove or force
# if do_remove:
# image = rembg.remove(image, session=rembg_session, **rembg_kwargs)
# return image
def resize_foreground(
image: PIL.Image.Image,
ratio: float,
) -> PIL.Image.Image:
image = np.array(image)
assert image.shape[-1] == 4
alpha = np.where(image[..., 3] > 0)
y1, y2, x1, x2 = (
alpha[0].min(),
alpha[0].max(),
alpha[1].min(),
alpha[1].max(),
)
# crop the foreground
fg = image[y1:y2, x1:x2]
# pad to square
size = max(fg.shape[0], fg.shape[1])
ph0, pw0 = (size - fg.shape[0]) // 2, (size - fg.shape[1]) // 2
ph1, pw1 = size - fg.shape[0] - ph0, size - fg.shape[1] - pw0
new_image = np.pad(
fg,
((ph0, ph1), (pw0, pw1), (0, 0)),
mode="constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
# compute padding according to the ratio
new_size = int(new_image.shape[0] / ratio)
# pad to size, double side
ph0, pw0 = (new_size - size) // 2, (new_size - size) // 2
ph1, pw1 = new_size - size - ph0, new_size - size - pw0
new_image = np.pad(
new_image,
((ph0, ph1), (pw0, pw1), (0, 0)),
mode="constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
new_image = PIL.Image.fromarray(new_image)
return new_image
def save_video(
frames: List[PIL.Image.Image],
output_path: str,
fps: int = 30,
):
# use imageio to save video
frames = [np.array(frame) for frame in frames]
writer = imageio.get_writer(output_path, fps=fps)
for frame in frames:
writer.append_data(frame)
writer.close()
def to_gradio_3d_orientation(mesh):
mesh.apply_transform(trimesh.transformations.rotation_matrix(-np.pi/2, [1, 0, 0]))
mesh.apply_scale([1, 1, -1])
mesh.apply_transform(trimesh.transformations.rotation_matrix(np.pi/2, [0, 1, 0]))
return mesh
+10
View File
@@ -0,0 +1,10 @@
{
"main_pass":
[
"-n", "-c:v", "libsvtav1",
"-pix_fmt", "yuv420p10le",
"-crf", "23"
],
"extension": "webm",
"environment": {"SVT_LOG": "1"}
}
+9
View File
@@ -0,0 +1,9 @@
{
"main_pass":
[
"-n", "-c:v", "libx264",
"-pix_fmt", "yuv420p",
"-crf", "19"
],
"extension": "mp4"
}
+11
View File
@@ -0,0 +1,11 @@
{
"main_pass":
[
"-n", "-c:v", "libx265",
"-pix_fmt", "yuv420p10le",
"-preset", "medium",
"-crf", "22",
"-x265-params", "log-level=quiet"
],
"extension": "mp4"
}
+9
View File
@@ -0,0 +1,9 @@
{
"main_pass":
[
"-n",
"-pix_fmt", "yuv420p",
"-crf", "23"
],
"extension": "webm"
}
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "comfyui-mixlab-nodes"
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
version = "0.39.0"
license = "MIT"
dependencies = ["numpy", "pyOpenSSL", "watchdog", "opencv-python-headless", "matplotlib", "openai", "simple-lama-inpainting", "clip-interrogator==0.6.0", "transformers>=4.36.0", "lark-parser", "imageio-ffmpeg", "rembg[gpu]", "omegaconf==2.3.0", "Pillow>=9.5.0", "einops==0.7.0", "trimesh>=4.0.5", "huggingface-hub", "scikit-image"]
[project.urls]
Repository = "https://github.com/shadowcz007/comfyui-mixlab-nodes"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "shadow"
DisplayName = "comfyui-mixlab-nodes"
Icon = ""
+16 -3
View File
@@ -4,9 +4,22 @@ watchdog
opencv-python-headless
matplotlib
openai
simple-lama-inpainting
# simple-lama-inpainting
clip-interrogator==0.6.0
transformers>=4.36.0
zhipuai
lark-parser
imageio-ffmpeg
imageio-ffmpeg
rembg[gpu]
omegaconf==2.3.0
Pillow>=9.5.0
einops==0.7.0
trimesh>=4.0.5
huggingface-hub
scikit-image
torchaudio
soundfile>=0.12.1
json-repair
decord
bitsandbytes
accelerate
+355 -237
View File
@@ -2,6 +2,8 @@ import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
import { loadExternalScript } from './common.js'
const getLocalData = key => {
let data = {}
try {
@@ -26,7 +28,8 @@ const setLocalDataOfWin = (key, value) => {
localStorage.setItem(key, JSON.stringify(value))
// window[key] = value
}
async function uploadImage (blob, fileType = '.svg', filename) {
async function uploadImage_ (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
const body = new FormData()
body.append(
@@ -41,13 +44,17 @@ async function uploadImage (blob, fileType = '.svg', filename) {
// console.log(resp)
let data = await resp.json()
return data
}
async function uploadImage (blob, fileType = '.svg', filename) {
let data = await uploadImage_(blob, fileType, filename)
let { name, subfolder } = data
let src = api.apiURL(
`/view?filename=${encodeURIComponent(
name
)}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
return src
}
@@ -94,7 +101,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -171,6 +181,42 @@ async function changeMaterial (
targetMaterial.pbrMetallicRoughness.baseColorTexture.setTexture(targetTexture)
}
function inputFileClick (isFileURL = false, isGlb = false) {
return new Promise((res, rej) => {
// 创建一个input元素
var input = document.createElement('input')
input.type = 'file'
input.accept = isGlb ? '.glb' : 'image/*'
// 监听input的change事件
input.addEventListener('change', function () {
// 获取上传的文件
var file = input.files[0]
if (isFileURL) {
res(URL.createObjectURL(file))
return
}
// 创建一个FileReader对象来读取文件
var reader = new FileReader()
// 监听FileReader的load事件
reader.addEventListener('load', async () => {
let base64 = reader.result
input.remove()
res(base64)
})
// 读取文件
reader.readAsDataURL(file)
})
// 触发input的点击事件
input.click()
})
}
app.registerExtension({
name: 'Mixlab.3D.3DImage',
async getCustomWidgets (app) {
@@ -189,7 +235,7 @@ app.registerExtension({
let d = getLocalData('_mixlab_3d_image')
// console.log('serializeValue', node)
if (d && d[node.id]) {
let { url, bg, material } = d[node.id]
let { url, bg, material, images } = d[node.id]
let data = {}
if (url) {
data.image = await parseImage(url)
@@ -205,6 +251,10 @@ app.registerExtension({
data.material = await parseImage(material)
}
if (images) {
data.images = images
}
return JSON.parse(JSON.stringify(data))
} else {
return {}
@@ -221,6 +271,11 @@ app.registerExtension({
if (nodeType.comfyClass == '3DImage') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
await loadExternalScript(
'/mixlab/app/lib/model-viewer.min.js',
'module'
)
orig_nodeCreated?.apply(this, arguments)
const uploadWidget = this.widgets.filter(w => w.name == 'upload')[0]
@@ -243,39 +298,29 @@ app.registerExtension({
const inputDiv = (key, placeholder, preview) => {
let div = document.createElement('div')
const ip = document.createElement('input')
ip.type = 'file'
const ip = document.createElement('button')
ip.className = `${'comfy-multiline-input'} ${placeholder}`
div.style = `display: flex;
align-items: center;
margin: 6px 8px;
margin-top: 0;`
ip.placeholder = placeholder
// ip.value = value
ip.style = `outline: none;
border: none;
padding: 4px;
width: 60%;cursor: pointer;
width: 100px;cursor: pointer;
height: 32px;`
const label = document.createElement('label')
label.style = 'font-size: 10px;min-width:32px'
label.innerText = placeholder
div.appendChild(label)
ip.innerText = placeholder
div.appendChild(ip)
let that = this,
filename = new Date().getTime()
let that = this
ip.addEventListener('change', async event => {
const file = event.target.files[0]
const reader = new FileReader()
filename = new Date().getTime()
// 读取文件内容
reader.onload = async e => {
const fileURL = URL.createObjectURL(file)
// console.log('文件URL: ', fileURL)
let html = `<model-viewer src="${fileURL}"
ip.addEventListener('click', async event => {
let fileURL = await inputFileClick(true, true)
// console.log('文件URL: ', fileURL)
let html = `<model-viewer src="${fileURL}"
oncontextmenu="return false;"
min-field-of-view="0deg" max-field-of-view="180deg"
shadow-intensity="1"
camera-controls
@@ -285,230 +330,303 @@ app.registerExtension({
<div>Variant: <select class="variant"></select></div>
<div>Material: <select class="material"></select></div>
<div>Material: <div class="material_img"> </div></div>
<div><button class="bg">BG</button></div>
<div>
<button class="bg">BG</button>
</div>
<div>
<input class="ddcap_step" type="number" min="1" max="20" step="1" value="1">
<input class="total_images" type="number" min="1" max="180" step="1" value="40">
<input class="ddcap_range" type="range" min="-180" max="180" step="1" value="0">
<input class="ddcap_range_top" type="range" min="-180" max="180" step="1" value="0">
<button class="ddcap">Capture Rotational Screenshots</button></div>
<div><button class="export">Export GLB</button></div>
</div></model-viewer>`
preview.innerHTML = html
if (that.size[1] < 400) {
that.setSize([that.size[0], that.size[1] + 300])
app.canvas.draw(true, true)
}
const modelViewerVariants = preview.querySelector('model-viewer')
const select = preview.querySelector('.variant')
const selectMaterial = preview.querySelector('.material')
const material_img = preview.querySelector('.material_img')
const bg = preview.querySelector('.bg')
const exportGLB = preview.querySelector('.export')
if (modelViewerVariants) {
modelViewerVariants.style.width = `${that.size[0] - 24}px`
modelViewerVariants.style.height = `${that.size[1] - 48}px`
}
modelViewerVariants.addEventListener('load', async () => {
const names = modelViewerVariants.availableVariants
// 变量
for (const name of names) {
const option = document.createElement('option')
option.value = name
option.textContent = name
select.appendChild(option)
}
// Adds a default option.
if (names.length === 0) {
const option = document.createElement('option')
option.value = 'default'
option.textContent = 'Default'
select.appendChild(option)
}
// 材质
extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
})
let timer = null
const delay = 500 // 延迟时间,单位为毫秒
async function checkCameraChange () {
let dd = getLocalData(key)
let base64Data = modelViewerVariants.toDataURL()
const contentType = getContentTypeFromBase64(base64Data)
const blob = await base64ToBlobFromURL(base64Data, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let url = await uploadImage(blob, '.png')
// console.log(url)
let bg_blob = await base64ToBlobFromURL(
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
)
let url_bg = await uploadImage(bg_blob, '.png')
// console.log('url_bg',url_bg)
if (!dd[that.id]) {
dd[that.id] = { url, bg: url_bg }
} else {
dd[that.id] = { ...dd[that.id], url }
}
// 材质贴图
let thumbUrl = material_img.getAttribute('src')
if (thumbUrl) {
let tb = await base64ToBlobFromURL(thumbUrl)
let tUrl = await uploadImage(tb, '.png')
// console.log('材质贴图', tUrl, thumbUrl)
dd[that.id].material = tUrl
}
setLocalDataOfWin(key, dd)
}
function startTimer () {
if (timer) clearTimeout(timer)
timer = setTimeout(checkCameraChange, delay)
}
modelViewerVariants.addEventListener('camera-change', startTimer)
select.addEventListener('input', async event => {
modelViewerVariants.variantName =
event.target.value === 'default' ? null : event.target.value
// 材质
await extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
checkCameraChange()
})
selectMaterial.addEventListener('input', event => {
// console.log(selectMaterial.value)
material_img.setAttribute('src', selectMaterial.value)
if (selectMaterial.getAttribute('data-new-material')) {
let index =
~~selectMaterial.selectedOptions[0].getAttribute(
'data-index'
)
changeMaterial(
modelViewerVariants,
modelViewerVariants.model.materials[index],
selectMaterial.getAttribute('data-new-material')
)
}
checkCameraChange()
})
bg.addEventListener('click', () => {
// 创建一个input元素
var input = document.createElement('input')
input.type = 'file'
// 监听input的change事件
input.addEventListener('change', function () {
// 获取上传的文件
var file = input.files[0]
// 创建一个FileReader对象来读取文件
var reader = new FileReader()
// 监听FileReader的load事件
reader.addEventListener('load', async () => {
let base64 = reader.result
// 将读取的文件内容设置为div的背景
preview.style.backgroundImage = 'url(' + base64 + ')'
const contentType = getContentTypeFromBase64(base64)
const blob = await base64ToBlobFromURL(base64, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let bg_url = await uploadImage(blob, '.png')
let bg_img = await createImage(base64)
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (!dd[that.id]) dd[that.id] = { url: '', bg: bg_url }
dd[that.id] = {
...dd[that.id],
bg: bg_url,
bg_w: bg_img.naturalWidth,
bg_h: bg_img.naturalHeight
}
setLocalDataOfWin(key, dd)
// 更新尺寸
let w = that.size[0] - 24,
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
})
// 读取文件
reader.readAsDataURL(file)
})
// 触发input的点击事件
input.click()
})
exportGLB.addEventListener('click', async () => {
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], 'export.glb')
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
uploadWidget.value = await uploadWidget.serializeValue()
// 更新尺寸
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 24,
h = (w * bg_h) / bg_w
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
}
}
preview.innerHTML = html
if (that.size[1] < 400) {
that.setSize([that.size[0], that.size[1] + 300])
app.canvas.draw(true, true)
}
// 以文本形式读取文件
reader.readAsDataURL(file)
const modelViewerVariants = preview.querySelector('model-viewer')
const select = preview.querySelector('.variant')
const selectMaterial = preview.querySelector('.material')
const material_img = preview.querySelector('.material_img')
const bg = preview.querySelector('.bg')
const exportGLB = preview.querySelector('.export')
const ddcap_step = preview.querySelector('.ddcap_step')
const total_images = preview.querySelector('.total_images')
const ddcap_range = preview.querySelector('.ddcap_range')
const ddcap_range_top = preview.querySelector('.ddcap_range_top')
const ddCap = preview.querySelector('.ddcap')
const sleep = (t = 1000) => {
return new Promise((res, rej) => {
return setTimeout(() => {
res(t)
}, t)
})
}
async function captureImage (isUrl = true) {
let base64Data = modelViewerVariants.toDataURL()
const contentType = getContentTypeFromBase64(base64Data)
const blob = await base64ToBlobFromURL(base64Data, contentType)
if (isUrl) return await uploadImage(blob, '.png')
return await uploadImage_(blob, '.png')
}
async function captureImages (angleIncrement = 1, totalImages = 12) {
// 记录初始旋转角度
const initialCameraOrbit =
modelViewerVariants.cameraOrbit.split(' ')
console.log(
'#captureImages',
initialCameraOrbit,
angleIncrement * totalImages
)
// const totalImages = 12
// const angleIncrement = totalRotation / totalImages // Each increment in degrees
let currentAngle =
Number(initialCameraOrbit[0].replace('deg', '')) -
(angleIncrement * totalImages) / 2 // Start from the leftmost angle
let frames = []
modelViewerVariants.removeAttribute('camera-controls')
for (let i = 0; i < totalImages; i++) {
modelViewerVariants.cameraOrbit = `${currentAngle}deg ${initialCameraOrbit[1]} ${initialCameraOrbit[2]}`
await sleep(1000)
console.log(`Capturing image at angle: ${currentAngle}deg`)
let file = await captureImage(false)
frames.push(file)
currentAngle += angleIncrement
}
await sleep(1000)
// 恢复到初始旋转角度
modelViewerVariants.cameraOrbit = initialCameraOrbit.join(' ')
modelViewerVariants.setAttribute('camera-controls', '')
return frames
}
ddCap.addEventListener('click', async e => {
const angleIncrement = Number(ddcap_step.value),
totalImages = Number(total_images.value)
let images = await captureImages(angleIncrement, totalImages)
// console.log(images)
let dd = getLocalData(key)
dd[that.id].images = images
setLocalDataOfWin(key, dd)
})
ddcap_range.addEventListener('input', async e => {
// console.log(ddcap_range.value)
const initialCameraOrbit =
modelViewerVariants.cameraOrbit.split(' ')
modelViewerVariants.cameraOrbit = `${ddcap_range.value}deg ${initialCameraOrbit[1]} ${initialCameraOrbit[2]}`
modelViewerVariants.setAttribute('camera-controls', '')
})
ddcap_range_top.addEventListener('input', async e => {
// console.log(ddcap_range.value)
const initialCameraOrbit =
modelViewerVariants.cameraOrbit.split(' ')
modelViewerVariants.cameraOrbit = `${initialCameraOrbit[0]} ${ddcap_range_top.value}deg ${initialCameraOrbit[2]}`
modelViewerVariants.setAttribute('camera-controls', '')
})
if (modelViewerVariants) {
modelViewerVariants.style.width = `${that.size[0] - 48}px`
modelViewerVariants.style.height = `${that.size[1] - 48}px`
}
modelViewerVariants.addEventListener('load', async () => {
const names = modelViewerVariants.availableVariants
// 变量
for (const name of names) {
const option = document.createElement('option')
option.value = name
option.textContent = name
select.appendChild(option)
}
// Adds a default option.
if (names.length === 0) {
const option = document.createElement('option')
option.value = 'default'
option.textContent = 'Default'
select.appendChild(option)
}
// 材质
extractMaterial(modelViewerVariants, selectMaterial, material_img)
})
let timer = null
const delay = 500 // 延迟时间,单位为毫秒
async function checkCameraChange () {
let dd = getLocalData(key)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let url = await captureImage()
let bg_blob = await base64ToBlobFromURL(
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
)
let url_bg = await uploadImage(bg_blob, '.png')
// console.log('url_bg',url_bg)
if (!dd[that.id]) {
dd[that.id] = { url, bg: url_bg }
} else {
dd[that.id] = { ...dd[that.id], url }
}
// 材质贴图
let thumbUrl = material_img.getAttribute('src')
if (thumbUrl) {
let tb = await base64ToBlobFromURL(thumbUrl)
let tUrl = await uploadImage(tb, '.png')
// console.log('材质贴图', tUrl, thumbUrl)
dd[that.id].material = tUrl
}
setLocalDataOfWin(key, dd)
}
function startTimer () {
if (timer) clearTimeout(timer)
timer = setTimeout(checkCameraChange, delay)
}
modelViewerVariants.addEventListener('camera-change', startTimer)
select.addEventListener('input', async event => {
modelViewerVariants.variantName =
event.target.value === 'default' ? null : event.target.value
// 材质
await extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
checkCameraChange()
})
selectMaterial.addEventListener('input', event => {
// console.log(selectMaterial.value)
material_img.setAttribute('src', selectMaterial.value)
if (selectMaterial.getAttribute('data-new-material')) {
let index =
~~selectMaterial.selectedOptions[0].getAttribute('data-index')
changeMaterial(
modelViewerVariants,
modelViewerVariants.model.materials[index],
selectMaterial.getAttribute('data-new-material')
)
}
checkCameraChange()
})
//更新bg
const updateBgData = (id, key, url, w, h) => {
let dd = getLocalData(key)
// console.log(dd[that.id],url)
if (!dd[id]) dd[id] = { url: '', bg: url }
dd[id] = {
...dd[id],
bg: url,
bg_w: w,
bg_h: h
}
setLocalDataOfWin(key, dd)
}
bg.addEventListener('click', async () => {
//更新bg
updateBgData(that.id, key, '', 0, 0)
preview.style.backgroundImage = 'none'
let base64 = await inputFileClick(false, false)
// 将读取的文件内容设置为div的背景
preview.style.backgroundImage = 'url(' + base64 + ')'
const contentType = getContentTypeFromBase64(base64)
const blob = await base64ToBlobFromURL(base64, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let bg_url = await uploadImage(blob, '.png')
let bg_img = await createImage(base64)
//更新bg
updateBgData(
that.id,
key,
bg_url,
bg_img.naturalWidth,
bg_img.naturalHeight
)
// 更新尺寸
let w = that.size[0] - 48,
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
})
exportGLB.addEventListener('click', async () => {
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], 'export.glb')
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
uploadWidget.value = await uploadWidget.serializeValue()
// 更新尺寸
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 48,
h = (w * bg_h) / bg_w
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
}
}
})
return div
}
let preview = document.createElement('div')
preview.className = 'preview'
preview.style = `margin-top: 12px;display: flex;
preview.style = `margin-top: 12px;
display: flex;
justify-content: center;
align-items: center;background-repeat: no-repeat;background-size: contain;`
align-items: center;background-repeat: no-repeat;
background-size: contain;`
let upload = inputDiv('_mixlab_3d_image', '3D Model', preview)
@@ -527,7 +645,7 @@ app.registerExtension({
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 24,
let w = that.size[0] - 48,
h = (w * bg_h) / bg_w
if (modelViewerVariants) {
@@ -561,7 +679,7 @@ app.registerExtension({
const r = onExecuted?.apply?.(this, arguments)
let div = this.widgets.filter(d => d.div)[0]?.div
console.log('Test', this.widgets)
// console.log('Test', this.widgets)
let material = message.material[0]
if (material) {
+90 -56
View File
@@ -2,8 +2,12 @@ import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
import { td_bg } from './td_background.js'
// console.log('td_bg', td_bg)
import { getUrl, base64Df, get_position_style, getObjectInfo } from './common.js'
//本机安装的插件节点全集
window._nodesAll = null
const parseImageToBase64 = url => {
return new Promise((res, rej) => {
@@ -24,38 +28,6 @@ const parseImageToBase64 = url => {
})
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 12 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'flex-start'
}
}
async function drawImageToCanvas (imageUrl, sFactor = 320) {
var canvas = document.createElement('canvas')
var ctx = canvas.getContext('2d')
@@ -166,6 +138,25 @@ async function extractInputAndOutputData (
if (node.type == 'Color') {
}
// 语音输入的支持
if (node.type == 'LoadAndCombinedAudio_') {
// if (
// data[id].widgets_values &&
// data[id].widgets_values[0] &&
// data[id].widgets_values[0].base64 &&
// data[id].widgets_values[0].base64.length > 0
// ) {
// options.defaultBase64 = data[id].widgets_values[0].base64
// }
input[inputIds.indexOf(id)] = {
...data[id],
title: node.title,
id,
options
}
}
if (node.type === 'LoadImage') {
// loadImage的mask支持
let output = node.outputs.filter(ot => ot.type == 'MASK')[0]
@@ -174,7 +165,7 @@ async function extractInputAndOutputData (
options.hasMask = true
}
// loadImage的默认图,转为base64
let imgurl = app.graph.getNodeById(id).imgs[0].src
let imgurl = app.graph.getNodeById(id).imgs[0].src + '&channel=rgb'
options.defaultImage = await drawImageToCanvas(imgurl, 512)
console.log('#loadImage的默认图', options)
@@ -189,15 +180,36 @@ async function extractInputAndOutputData (
// input.push()
}
if (outputIds.includes(id)) {
let options = {}
//输出的默认图
if (
node.type === 'SaveImageAndMetadata_' &&
app.graph.getNodeById(id).imgs
) {
// SaveImageAndMetadata_的默认图,转为base64
let imgurl = app.graph.getNodeById(id).imgs[0].src
options.defaultImage = await drawImageToCanvas(imgurl, 512)
console.log('#SaveImageAndMetadata_的默认图', options)
}
// let node = app.graph.getNodeById(id)
// output.push()
output[outputIds.indexOf(id)] = { ...data[id], title: node.title, id }
output[outputIds.indexOf(id)] = {
...data[id],
title: node.title,
id,
options
}
}
if (
node.type === 'KSampler' ||
node.type == 'SamplerCustom' ||
node.type === 'ChinesePrompt_Mix'
node.type === 'ChinesePrompt_Mix' ||
node.type === 'Seed_' ||
node.type === 'SiliconflowLLM' ||
node.type === 'ChatGPTOpenAI'
) {
// seed 的类型收集
try {
@@ -217,13 +229,6 @@ async function extractInputAndOutputData (
return { input, output, seed, seedTitle }
}
function getUrl () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
const getLocalData = key => {
let data = {}
try {
@@ -266,7 +271,10 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
}
async function save (json, download = false, showInfo = true) {
console.log('####SAVE', json[0])
let nodesAll = window._nodesAll || (await getObjectInfo())
console.log('####SAVE', nodesAll, json[0])
const name = json[0],
version = json[5],
share_prefix = json[6], //用于分享的功能扩展
@@ -287,6 +295,13 @@ async function save (json, download = false, showInfo = true) {
try {
let data = await app.graphToPrompt()
//从output数据里把工作流的节点,插件数据统计出来
data.nodesMap = {}
for (const id in data.output) {
data.nodesMap[data.output[id].class_type] =
nodesAll[data.output[id].class_type]
}
let { input, output, seed, seedTitle } = await extractInputAndOutputData(
data,
inputIds,
@@ -349,11 +364,11 @@ async function save (json, download = false, showInfo = true) {
function getInputsAndOutputs () {
const inputs =
`LoadImage ImagesPrompt_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
`LoadImage LoadImagesToBatch ImagesPrompt_ LoadAndCombinedAudio_ LoadVideoAndSegment_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
' '
),
outputs =
`PreviewImage,SaveImage,ShowTextForGPT,VHS_VideoCombine,Image Save,SaveImageAndMetadata_`.split(
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,CombineAudioVideo,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
','
)
@@ -378,6 +393,11 @@ function getInputsAndOutputs () {
app.registerExtension({
name: 'Mixlab.utils.AppInfo',
init () {
if (!window._nodesAll) {
getObjectInfo().then(r => (window._nodesAll = r))
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'AppInfo') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
@@ -392,20 +412,19 @@ app.registerExtension({
const { input, output } = getInputsAndOutputs()
input_ids.value = input.join('\n')
output_ids.value = output.join('\n')
const widget = {
type: 'div',
name: 'AppInfoRun',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(
Object.assign(this.div.style, {
...get_position_style(
ctx,
widget_width,
node.size[1] - widget_height,
node.size[1]
)
)
),
zIndex: 1
})
}
}
@@ -450,6 +469,21 @@ app.registerExtension({
}
})
//td bg
const tdBG = document.createElement('button')
tdBG.innerText = 'Canvas Mode'
tdBG.style = style
tdBG.style.marginLeft = '12px'
tdBG.addEventListener('click', () => {
td_bg.toggle()
if (td_bg.running) {
tdBG.style.background = 'yellow'
} else {
tdBG.style.background = 'transparent'
}
})
// author
let author = document.createElement('div')
// author.style=`display: flex`
@@ -606,6 +640,7 @@ app.registerExtension({
btns.appendChild(btn)
btns.appendChild(download)
btns.appendChild(tdBG)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
@@ -634,9 +669,8 @@ app.registerExtension({
}
const div = this.widgets.filter(w => w.div)[0].div
Array.from(
div.querySelectorAll('button'),
b => (b.style.background = 'yellow')
Array.from(div.querySelectorAll('button'), b =>
b.innerText != 'Canvas Mode' ? (b.style.background = 'yellow') : ''
)
} catch (error) {}
}
+219 -1
View File
@@ -19,7 +19,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -396,3 +399,218 @@ app.registerExtension({
}
}
})
// 上传音频转为base64
async function uploadAndConvertAudio (file) {
if (!file) {
alert('Please select a WAV file.')
return
}
if (file.type !== 'audio/wav') {
alert('Only WAV files are supported.')
return
}
try {
const base64Audio = await readFileAsDataURL(file)
return base64Audio
} catch (error) {
console.error('Error reading file:', error)
alert('Error reading file.')
}
}
function readFileAsDataURL (file) {
return new Promise((resolve, reject) => {
const reader = new FileReader()
reader.onload = function (event) {
resolve(event.target.result)
}
reader.onerror = function (error) {
reject(error)
}
reader.readAsDataURL(file)
})
}
const createInputAudioForBatch = (base64, widget) => {
// Create an audio element
let audio = document.createElement('audio')
audio.src = base64
audio.controls = true
audio.style = 'width: 120px; display: block'
// Create a delete button
let deleteButton = document.createElement('button')
deleteButton.textContent = 'Delete'
deleteButton.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
margin-left: 10px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
// Create a container for the audio and delete button
let container = document.createElement('div')
container.appendChild(audio)
container.appendChild(deleteButton)
container.style = `display: flex;margin-top: 12px;`
// Add event listener for the delete button
deleteButton.addEventListener('click', e => {
let newValue = []
let items = widget.value?.base64 || []
for (const v of items) {
if (v != base64) newValue.push(v)
}
widget.value.base64 = newValue
container.remove()
})
return container
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadAndCombinedAudio_',
async getCustomWidgets (app) {
return {
AUDIOBASE64 (node, inputName, inputData, app) {
// console.log('##node', node)
const widget = {
value: {
base64: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 122] // a method to compute the current size of the widget
}
// serializeValue (nodeId, widgetIndex) {
// return widget.value
// },
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadAndCombinedAudio_') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
let audiosWidget = this.widgets.filter(w => w.name == 'audios')[0]
const widget = {
type: 'div',
name: 'audio_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
)
},
serialize: false
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
let audioPreview = document.createElement('div')
let audiosDiv = document.createElement('div') //显示图片
audiosDiv.className = 'audios_preview'
audiosDiv.style = `width: calc(100% - 14px);
display: flex;
flex-wrap: wrap;
padding: 7px; justify-content: space-between;
align-items: center;`
const btn = document.createElement('button')
btn.innerText = 'Upload Audio'
btn.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
btn.addEventListener('click', e => {
e.preventDefault()
let inputAudio = document.createElement('input')
inputAudio.type = 'file'
inputAudio.accept = "audio/*"
inputAudio.style.display = 'none'
inputAudio.addEventListener('change', async e => {
e.preventDefault()
const file = e.target.files[0]
let base64 = await uploadAndConvertAudio(file)
if (!audiosWidget.value) audiosWidget.value = { base64: [] }
audiosWidget.value.base64.push(base64)
let a = createInputAudioForBatch(base64, audiosWidget)
audiosDiv.appendChild(a)
})
inputAudio.click()
inputAudio.remove()
})
widget.div.appendChild(audioPreview)
audioPreview.appendChild(audiosDiv)
audioPreview.appendChild(btn)
// audioPreview.appendChild(inputAudio)
this.addCustomWidget(widget)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
try {
// document.removeEventListener('wheel', handleMouseWheel)
} catch (error) {
console.log(error)
}
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadAndCombinedAudio_') {
// await sleep(0)
let audiosWidget = node.widgets.filter(w => w.name === 'audios')[0]
let audioPreview = node.widgets.filter(w => w.name == 'audio_base64')[0]
let pre = audioPreview.div.querySelector('.audios_preview')
for (const d of audiosWidget.value?.base64 || []) {
let im = createInputAudioForBatch(d, audiosWidget)
pre.appendChild(im)
}
}
}
})
+173
View File
@@ -0,0 +1,173 @@
import { getUrl } from './common.js'
async function* completion (url, messages, controller) {
let data = {
model: 'gpt-3.5-turbo-16k',
messages,
temperature: 0.05,
stream: true
}
// if (imageNode) {
// data = { ...data, image_data: [imageNode] }
// }
// let controller = new AbortController()
let response = await fetch(url, {
method: 'POST',
body: JSON.stringify(data),
headers: {
Connection: 'keep-alive',
'Content-Type': 'application/json',
Accept: 'text/event-stream'
},
signal: controller.signal
})
const reader = response.body.getReader()
const decoder = new TextDecoder()
let content = ''
let leftover = '' // Buffer for partially read lines
try {
let cont = true
while (cont) {
let result = await reader.read()
if (result.done) {
break
}
// Add any leftover data to the current chunk of data
const text = leftover + decoder.decode(result.value)
// Check if the last character is a line break
const endsWithLineBreak = text.endsWith('\n')
// Split the text into lines
let lines = text.split('\n')
// If the text doesn't end with a line break, then the last line is incomplete
// Store it in leftover to be added to the next chunk of data
if (!endsWithLineBreak) {
leftover = lines.pop()
} else {
leftover = '' // Reset leftover if we have a line break at the end
}
// Parse all sse events and add them to result
const regex = /^(\S+):\s(.*)$/gm
for (const line of lines) {
const match = regex.exec(line)
if (match) {
result[match[1]] = match[2]
// since we know this is llama.cpp, let's just decode the json in data
if (result.data) {
result.data = JSON.parse(result.data)
// console.log('#result.data',result.data)
content += result.data.choices[0].delta?.content || ''
// yield
yield result
// if we got a stop token from server, we will break here
if (result.data.choices[0].finish_reason == 'stop') {
if (result.data.generation_settings) {
// generation_settings = result.data.generation_settings;
}
cont = false
break
}
}
}
}
}
} catch (e) {
console.error('llama error: ', e)
throw e
} finally {
controller.abort()
}
return content
// return (await response.json()).content
}
export async function completion_ (apiKey, url, messages, controller, callback) {
let request = await chatCompletion(apiKey, url, messages, controller)
for await (const chunk of request) {
if (callback) callback(chunk)
}
}
export async function* chatCompletion (apiKey, url, messages, controller) {
url = `${getUrl()}/chat/completions`
const requestBody = {
model: '01-ai/Yi-1.5-9B-Chat-16K',
messages: messages,
stream: true,
key: apiKey
}
let response = await fetch(url, {
method: 'POST',
headers: {
'Content-Type': 'application/json',
Authorization: `Bearer ${apiKey}`
},
body: JSON.stringify(requestBody),
mode: 'cors', // This is to ensure the request is made with CORS
signal: controller.signal
})
const reader = response.body.getReader()
const decoder = new TextDecoder()
let content = ''
let leftover = '' // Buffer for partially read lines
try {
let cont = true
while (cont) {
let result = await reader.read()
if (result.done) {
break
}
// Add any leftover data to the current chunk of data
const text = leftover + decoder.decode(result.value)
// Check if the last character is a line break
const endsWithLineBreak = text.endsWith('\r\n')
// Split the text into lines
let lines = text.split('\r\n')
// If the text doesn't end with a line break, then the last line is incomplete
// Store it in leftover to be added to the next chunk of data
if (!endsWithLineBreak) {
leftover = lines.pop()
} else {
leftover = '' // Reset leftover if we have a line break at the end
}
for (const line of lines) {
if (line) {
content += line
yield line // Yield the trimmed line
} else {
cont = false
break
}
}
}
} catch (e) {
console.error('chat error: ', e)
throw e
} finally {
controller.abort()
}
return content
}
+8 -2
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.20.0'
const version = 'v0.39.0'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
@@ -17,7 +17,13 @@ fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
return
if (latestVersion && latestVersion != version) {
localStorage.setItem('_mixlab_nodes_vesion', latestVersion)
app.ui.dialog.show(`<h4 style="font-size: 18px;">${repoName} <br>
app.ui.dialog.show(`<a style="color: white;
font-size: 18px;
font-weight: 800;
letter-spacing: 2px;
}"
href="https://discord.gg/cXs9vZSqeK">Welcome to Mixlab nodes discord</a>
<h4 style="font-size: 18px;">${repoName} <br>
Latest release version: ${latestVersion}</h4>
<p>Please proceed to the official repository to download the latest version.</p>
<a style="color: #2196F3;
+175
View File
@@ -0,0 +1,175 @@
export const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
export function getUrl () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
// 更新或者获取key
export const updateLLMAPIKey = async key => {
try {
const res = await fetch(`${getUrl()}/mixlab/llm_api_key`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
key: key || null
})
})
const data = await res.json()
if (!res.ok) {
console.error('Error:', data.error)
return
}
if (key) {
console.log('API key saved successfully:', data.message)
return key
} else {
console.log('Retrieved API key:', data.key)
return data.key
}
} catch (error) {
console.error('Request failed:', error)
}
}
//获取当前系统的插件,节点清单
export function getObjectInfo () {
return new Promise(async (resolve, reject) => {
let url = getUrl()
try {
const response = await fetch(`${url}/object_info`)
const data = await response.json()
resolve(data)
} catch (error) {
reject(error)
}
})
}
export function get_position_style (
ctx,
widget_width,
y,
node_height,
left = 44
) {
const MARGIN = 0 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const scaleX = elRect.width / ctx.canvas.width
const scaleY = elRect.height / ctx.canvas.height
const transform = new DOMMatrix()
.scaleSelf(scaleX, scaleY)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `${left}px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'flex-start',
zIndex: 99
}
}
export function loadExternalScript (url, type) {
return new Promise((resolve, reject) => {
const existingScript = document.querySelector(`script[src="${url}"]`)
if (existingScript) {
existingScript.onload = () => {
resolve()
}
existingScript.onerror = reject
return
}
const script = document.createElement('script')
script.src = url
if (type) script.type = type // Add this line to load the script as an ES module
script.onload = () => {
resolve()
}
script.onerror = reject
document.head.appendChild(script)
})
}
export async function getQueue () {
try {
const res = await fetch(`${getUrl()}/queue`)
const data = await res.json()
// console.log(data.queue_running,data.queue_pending)
return {
// Running action uses a different endpoint for cancelling
Running: data.queue_running.length,
Pending: data.queue_pending.length
}
} catch (error) {
console.error(error)
return { Running: 0, Pending: 0 }
}
}
export async function interrupt () {
const resp = await fetch(`${getUrl()}/interrupt`, {
method: 'POST'
})
}
export async function sleep (t = 200) {
return new Promise((res, rej) => {
setTimeout(() => {
res(true)
}, t)
})
}
export function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
export const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
export const saveLocalData = (key, id, val) => {
let data = getLocalData(key)
data[id] = val
localStorage.setItem(key, JSON.stringify(data))
}
+8 -203
View File
@@ -1,205 +1,5 @@
import { app } from '../../../scripts/app.js'
// import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
async function getConfig () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
const res = await fetch(`${url}/mixlab`, {
method: 'POST'
})
return await res.json()
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
app.registerExtension({
name: 'Mixlab.GPT.ChatGPTOpenAI',
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##inputData', inputData)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
return data[node.id] || 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
},
URL (node, inputName, inputData, app) {
// console.log('node', inputName, inputData[0])
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {
// a method to draw the widget (ctx is a CanvasRenderingContext2D)
},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_url')
return data[node.id] || 'https://api.openai.com/v1'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ChatGPTOpenAI') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const api_key = this.widgets.filter(w => w.name == 'api_key')[0]
const api_url = this.widgets.filter(w => w.name == 'api_url')[0]
console.log('ChatGPTOpenAI nodeData', this.widgets)
const widget = {
type: 'div',
name: 'chatgptdiv',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, api_key.y, node.size[1])
)
}
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder) => {
let div = document.createElement('div')
const ip = document.createElement('input')
ip.type = placeholder === 'Key' ? 'password' : 'text'
ip.className = `${'comfy-multiline-input'} ${placeholder}`
div.style = `display: flex;
align-items: center;
margin: 6px 8px;
margin-top: 0;`
ip.placeholder = placeholder
ip.value = placeholder
ip.style = `margin-left: 24px;
outline: none;
border: none;
padding: 4px;width: 100%;`
const label = document.createElement('label')
label.style = 'font-size: 10px;min-width:32px'
label.innerText = placeholder
div.appendChild(label)
div.appendChild(ip)
ip.addEventListener('change', () => {
let data = getLocalData(key)
data[this.id] = ip.value.trim()
localStorage.setItem(key, JSON.stringify(data))
console.log(this.id, key)
})
return div
}
let inputKey = inputDiv('_mixlab_api_key', 'Key')
let inputUrl = inputDiv('_mixlab_api_url', 'URL')
widget.div.appendChild(inputKey)
widget.div.appendChild(inputUrl)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputUrl.remove()
inputKey.remove()
widget.div.remove()
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
// Fires every time a node is constructed
// You can modify widgets/add handlers/etc here
if (node.type === 'ChatGPTOpenAI') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key'),
url = getLocalData('_mixlab_api_url')
let id = node.id
// console.log('ChatGPTOpenAI serialize_widgets', this)
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
widget.div.querySelector('.URL').value =
url[id] || 'https://api.openai.com/v1'
}
}
})
app.registerExtension({
name: 'Mixlab.GPT.ShowTextForGPT',
@@ -209,13 +9,16 @@ app.registerExtension({
text = text.filter(t => t && t?.trim())
if (this.widgets) {
// console.log('#ShowTextForGPT',this.widgets)
// const pos = this.widgets.findIndex(w => w.name === 'text')
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name == 'show_text') this.widgets[i].onRemove?.()
if (this.widgets[i].name == 'show_text')
this.widgets[i].onRemove?.()
}
this.widgets.length = 1
this.widgets.length = 2
}
// console.log('ShowTextForGPT',text)
for (let list of text) {
if (list) {
// console.log('#####', list)
@@ -228,6 +31,8 @@ app.registerExtension({
w.inputEl.readOnly = true
w.inputEl.style.opacity = 0.6
// w.inputEl.style.display='none'
try {
if (typeof list != 'string') {
let data = JSON.parse(list)
+313 -47
View File
@@ -1,7 +1,41 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
// import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import { applyTextReplacements } from '../../../scripts/utils.js'
import { loadExternalScript, get_position_style } from './common.js'
function loadImageToCanvas (base64Image) {
var img = new Image()
var canvas = document.createElement('canvas')
var ctx = canvas.getContext('2d')
return new Promise((res, rej) => {
img.onload = function () {
// 等比例缩放图片
var width = img.width
var height = img.height
var max_width = 1024
if (width > max_width) {
height *= max_width / width
width = max_width
}
// 设置canvas尺寸
canvas.width = width
canvas.height = height
// 在canvas上绘制图片
ctx.drawImage(img, 0, 0, width, height)
// 将canvas转换为base64图片数据
var canvasData = canvas.toDataURL()
res(canvasData) // canvas转换后的base64图片数据
}
img.src = base64Image
})
}
async function uploadImage (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
@@ -56,37 +90,40 @@ function getContentTypeFromBase64 (base64Data) {
// const blob = base64ToBlob(base64Data, contentType);
// console.log(blob);
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
// function get_position_style (ctx, widget_width, y, node_height) {
// const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
// /* Create a transform that deals with all the scrolling and zooming */
// const elRect = ctx.canvas.getBoundingClientRect()
// const transform = new DOMMatrix()
// .scaleSelf(
// elRect.width / ctx.canvas.width,
// elRect.height / ctx.canvas.height
// )
// .multiplySelf(ctx.getTransform())
// .translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
// return {
// transformOrigin: '0 0',
// transform: transform,
// left:
// document.querySelector('.comfy-menu').style.display === 'none'
// ? `60px`
// : `0`,
// top: `0`,
// cursor: 'pointer',
// position: 'absolute',
// maxWidth: `${widget_width - MARGIN * 2}px`,
// // maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
// width: `${widget_width - MARGIN * 2}px`,
// // height: `${node_height * 0.3 - MARGIN * 2}px`,
// // background: '#EEEEEE',
// display: 'flex',
// flexDirection: 'column',
// // alignItems: 'center',
// justifyContent: 'space-around'
// }
// }
const getLocalData = key => {
let data = {}
@@ -312,7 +349,7 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
get_position_style(ctx, widget_width, 44, node.size[1], 36)
)
}
}
@@ -498,7 +535,7 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
get_position_style(ctx, widget_width, y, node.size[1], 36)
)
}
}
@@ -632,17 +669,29 @@ const createInputImageForBatch = (base64, widget) => {
im.addEventListener('click', e => {
let newValue = []
let items=widget.value?.base64||[];
let items = widget.value?.base64 || []
for (const v of items) {
if (v != base64) newValue.push(v)
}
widget.value.base64=newValue;
widget.value.base64 = newValue
im.remove()
})
return im
}
// 添加新图片
const addBase64ToWidgetForLoadImagesToBatch = (
base64,
imagesWidget,
imagesDiv
) => {
if (!imagesWidget.value.base64) imagesWidget.value.base64 = []
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadImagesToBatch',
async getCustomWidgets (app) {
@@ -651,7 +700,7 @@ app.registerExtension({
// console.log('##node', node)
const widget = {
value: {
base64:[]
base64: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
@@ -659,7 +708,7 @@ app.registerExtension({
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
}
// serializeValue (nodeId, widgetIndex) {
// return widget.value
// },
@@ -674,6 +723,7 @@ app.registerExtension({
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadImagesToBatch') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
@@ -685,10 +735,10 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
get_position_style(ctx, widget_width, 44, node.size[1], 44)
)
},
serialize:false
serialize: false
}
widget.div = $el('div', {})
@@ -703,25 +753,53 @@ app.registerExtension({
flex-wrap: wrap;
padding: 7px; justify-content: space-between;
align-items: center;`
let inputImage = document.createElement('input')
inputImage.type = 'file'
inputImage.style.display = 'none'
inputImage.addEventListener('change', e => {
e.preventDefault()
const file = e.target.files[0]
const reader = new FileReader()
reader.onload = async event => {
const base64 = event.target.result
let base64 = event.target.result
//压缩图片,控制1024以内
base64 = await loadImageToCanvas(base64)
// console.log(base64)
if (!imagesWidget.value) imagesWidget.value = {base64:[]}
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64,imagesWidget)
imagesDiv.appendChild(im)
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
addBase64ToWidgetForLoadImagesToBatch(
base64,
imagesWidget,
imagesDiv
)
}
reader.readAsDataURL(file)
})
// 如果是复制的,有数据 , 这个不生效,取不到数据, 需要在nodeCreated里获取
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
const btn = document.createElement('button')
btn.innerText = 'Upload Image'
btn.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
btn.addEventListener('click', e => {
e.preventDefault()
inputImage.click()
})
widget.div.appendChild(imagePreview)
imagePreview.appendChild(imagesDiv)
imagePreview.appendChild(btn)
imagePreview.appendChild(inputImage)
this.addCustomWidget(widget)
@@ -744,18 +822,206 @@ app.registerExtension({
this.serialize_widgets = true //需要保存参数
}
}
if (nodeData.name === 'SaveImageAndMetadata_') {
const onNodeCreated = nodeType.prototype.onNodeCreated
// /web/extensions/core/saveImageExtraOutput.js
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
const widget = this.widgets.find(w => w.name === 'filename_prefix')
widget.serializeValue = () => {
return applyTextReplacements(app, widget.value)
}
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
console.log('##onExecuted', this, message)
//TODO 是否 保存base64
if (message.base64) {
if (Array.isArray(message.base64)) {
}
}
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadImagesToBatch') {
// await sleep(0)
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
let imagesDiv = imagePreview.div.querySelector('.images_preview')
let pre = imagePreview.div.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d,imagesWidget)
pre.appendChild(im)
let im = createInputImageForBatch(d, imagesWidget)
imagesDiv.appendChild(im)
}
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
// console.log('#LoadImagesToBatch', node.type)
if (node.type === 'LoadImagesToBatch') {
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
let imagesDiv = imagePreview?.div?.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
imagesDiv.appendChild(im)
}
}
}, 1000)
}
})
// 如何引入css
app.registerExtension({
name: 'Mixlab.output.ComparingTwoFrames_',
init () {
loadExternalScript('/mixlab/app/lib/juxtapose.min.js')
$el('link', {
rel: 'stylesheet',
href: '/mixlab/app/lib/juxtapose.css',
parent: document.head
})
$el('style', {
textContent: `
.juxtapose-name{
display: none!important;
}
`,
parent: document.body
})
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ComparingTwoFrames_') {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
this.size = [400, this.size[1]]
console.log('##onNodeCreated', this)
const widget = {
type: 'div',
name: 'preview',
draw (ctx, node, widget_width, y, widget_height) {
let s = get_position_style(ctx, widget_width, 44, node.size[1], 36)
delete s.height
Object.assign(this.div.style, s)
},
serialize: false
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
this.serialize_widgets = true //需要保存参数
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
console.log('##onExecuted', this, message)
this.widgets[0].div.id = 'mix_comparingtowframes_' + this.id
let after_image = message.after_images[0]
let before_image = message.before_images[0]
after_image = `${window.location.protocol}//${
window.location.hostname
}:${window.location.port}/view?filename=${encodeURIComponent(
after_image.filename
)}&type=${after_image.type}&subfolder=${encodeURIComponent(
after_image.subfolder
)}&t=${+new Date()}`
before_image = `${window.location.protocol}//${
window.location.hostname
}:${window.location.port}/view?filename=${encodeURIComponent(
before_image.filename
)}&type=${before_image.type}&subfolder=${encodeURIComponent(
before_image.subfolder
)}&t=${+new Date()}`
this.widgets[0].div.innerHTML = ''
let slider = new juxtapose.JXSlider(
'#mix_comparingtowframes_' + this.id,
[
{
src: before_image,
label: 'Before'
},
{
src: after_image,
label: 'After'
}
],
{
animate: true,
showLabels: true,
showCredits: false,
startingPosition: '50%',
makeResponsive: false
}
)
this.widgets_values = [
{
src: before_image,
label: 'Before'
},
{
src: after_image,
label: 'After'
}
]
this.size = [this.size[0], 300]
}
}
},
async loadedGraphNode (node, app) {
// console.log('##loadedGraphNode', node)
if (node.type === 'ComparingTwoFrames_') {
// node.widgets[0].div.id = 'mix_comparingtowframes_' + node.id
// if (node.widgets_values && node.widgets_values[0]) {
// node.widgets[0].div.innerHTML = ''
// let slider = new juxtapose.JXSlider(
// '#mix_comparingtowframes_' + node.id,
// node.widgets_values,
// {
// animate: true,
// showLabels: true,
// showCredits: false,
// startingPosition: '50%',
// makeResponsive: false
// }
// )
// }
}
}
})
+4 -1
View File
@@ -81,7 +81,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
+16 -88
View File
@@ -3,31 +3,17 @@ import { app } from '../../../scripts/app.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
import {
getQueue,
interrupt,
get_position_style,
base64Df,
getUrl,
createImage,
sleep
} from './common.js'
async function getQueue () {
try {
const res = await fetch(`${url}/queue`)
const data = await res.json()
// console.log(data.queue_running,data.queue_pending)
return {
// Running action uses a different endpoint for cancelling
Running: data.queue_running.length,
Pending: data.queue_pending.length
}
} catch (error) {
console.error(error)
return { Running: 0, Pending: 0 }
}
}
async function interrupt () {
const resp = await fetch(`${url}/interrupt`, {
method: 'POST'
})
}
// let url = getUrl()
async function clipboardWriteImage (win, url) {
const canvas = document.createElement('canvas')
@@ -208,22 +194,6 @@ async function shareScreen (
}
}
async function sleep (t = 200) {
return new Promise((res, rej) => {
setTimeout(() => {
res(true)
}, t)
})
}
function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
async function compareImages (threshold, previousImage, currentImage) {
// 将 base64 转换为 Image 对象
var previousImg = await createImage(previousImage)
@@ -458,44 +428,6 @@ async function requestCamera () {
return false
}
/*
A method that returns the required style for the html
*/
function get_position_style (ctx, widget_width, y, node_height, top) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `${top}px`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
app.registerExtension({
name: 'Mixlab.image.ScreenShareNode',
async getCustomWidgets (app) {
@@ -593,17 +525,12 @@ app.registerExtension({
type: 'HTML', // whatever
name: 'sreen_share', // whatever
draw (ctx, node, widget_width, y, widget_height) {
// console.log('ScreenSHare', y, widget_height)
// console.log('ScreenSHare', node)
Object.assign(
this.card.style,
get_position_style(
ctx,
widget_width,
widget_height * 5,
node.size[1],
40
)
get_position_style(ctx, widget_width, y, node.size[1], 40)
)
}
}
@@ -1043,12 +970,13 @@ async function setArea (src) {
div.innerHTML = `
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
height: 100vh;
z-index:999999;
z-index:99999999999999;
width: 100%;'>
<img id='ml_video' style='position: absolute;
height: ${displayHeight}px;user-select: none;
-webkit-user-drag: none;
outline: 2px solid #eaeaea;
left: 0;
box-shadow: 8px 9px 17px #575757;' />
<div id='ml_selection' style='position: absolute;
border: 2px dashed red;
@@ -1267,7 +1195,7 @@ app.registerExtension({
})
widget.PictureInPicture = $el('button', {
innerText: 'PictureInPicture',
innerText: 'Picture In Picture',
style: {
display: 'pictureInPictureEnabled' in document ? 'block' : 'none',
cursor: 'pointer',
+195
View File
@@ -0,0 +1,195 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
import { get_position_style } from './common.js'
function base64ToBlobFromURL (base64URL, contentType) {
return fetch(base64URL).then(response => response.blob())
}
async function uploadImage (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
const body = new FormData()
body.append(
'image',
new File([blob], (filename || new Date().getTime()) + fileType)
)
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body
})
// console.log(resp)
let data = await resp.json()
let { name, subfolder } = data
// let src = api.apiURL(
// `/view?filename=${encodeURIComponent(
// name
// )}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
// )
return data
}
// 上传得到url
async function uploadBase64ToFile (base64) {
let bg_blob = await base64ToBlobFromURL(base64)
let url = await uploadImage(bg_blob, '.png')
return url
}
const p5InputNode = {
name: 'Mixlab.Comfy.P5Input',
async getCustomWidgets (app) {
return {
IMAGEBASE64 (node, inputName, inputData, app) {
const widget = {
value: {
images: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
size: [320, 120], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
}
}
node.addCustomWidget(widget)
return widget
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'P5Input') {
// console.log('P5Input')
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'image_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(
ctx,
widget_width - 24,
44,
node.size[1] * 2.8,
44
)
)
},
serialize: false
}
widget.div = $el('div', {})
widget.div.style = `margin:12px;width:400px;height:480px;background:white`
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
// window.removeEventListener('message', ms)
return onRemoved?.()
}
// 节点的大小控制
this.setSize([480, 560])
app.canvas.draw(true, true)
const onResize = this.onResize
this.onResize = () => {
// 设置最小尺寸
if (
Math.max(this.size[0], 480) != this.size[0] &&
Math.max(this.size[1], 560) != this.size[1]
) {
this.setSize([
Math.max(this.size[0], 480),
Math.max(this.size[1], 560)
])
}
return onResize?.apply(this, arguments)
}
this.serialize_widgets = true //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// console.log('##onExecuted', this, message._info)
// app.graph.getNodeById(8).widgets[1].div.querySelector('iframe').contentWindow.postMessage('Hello from parent', '*');
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'P5Input') {
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
let widget = node.widgets?.filter(w => w.name == 'image_base64')[0]
let framesWidget = node.widgets?.filter(w => w.name == 'frames')[0]
if (node.type === 'P5Input' && widget) {
console.log('#nodeCreated P5Input')
if (framesWidget && !framesWidget.value)
framesWidget.value = { images: [] }
framesWidget.value._seed = Math.random()
let nodeId = node.id
//延迟才能获得this.id
widget.div.innerHTML = `<iframe src="mixlab/app/p5_export/p5.html?id=${nodeId}"
style="border:0;width:100%;height:100%;"
></iframe>`
// 监听来自iframe的消息
const ms = async event => {
const data = event.data
console.log('#P5 Input #', data)
if (
data.from === 'p5.widget' &&
data.status === 'save' &&
data.frames &&
data.frames.length >= 0 &&
data.nodeId == nodeId &&
data.id != framesWidget.value.id
) {
const frames = data.frames
//workflow会存储到local,会卡死
framesWidget.value.images = []
for (const f of frames) {
let file = await uploadBase64ToFile(f)
framesWidget.value.images.push(file)
}
// framesWidget.value.base64 = frames
// framesWidget.value._seed = Math.random()
node.title = 'P5 Input #' + frames.length
framesWidget.value.id = data.id
}
}
window.addEventListener('message', ms)
}
}, 1000)
}
}
app.registerExtension(p5InputNode)
+17 -10
View File
@@ -3,7 +3,7 @@ import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import PhotoSwipeLightbox from '/extensions/comfyui-mixlab-nodes/lib/photoswipe-lightbox.esm.min.js'
import PhotoSwipeLightbox from '/mixlab/app/lib/photoswipe-lightbox.esm.min.js'
function loadCSS (url) {
var link = document.createElement('link')
link.rel = 'stylesheet'
@@ -40,14 +40,14 @@ function loadCSS (url) {
// Append the style element to the document head
document.head.appendChild(style)
}
loadCSS('/extensions/comfyui-mixlab-nodes/lib/photoswipe.min.css')
loadCSS('/mixlab/app/lib/photoswipe.min.css')
function initLightBox () {
const lightbox = new PhotoSwipeLightbox({
gallery: '.prompt_image_output',
children: 'a',
pswpModule: () =>
import('/extensions/comfyui-mixlab-nodes/lib/photoswipe.esm.min.js')
import('/mixlab/app/lib/photoswipe.esm.min.js')
})
lightbox.on('uiRegister', function () {
@@ -100,7 +100,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -178,7 +181,7 @@ app.registerExtension({
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const mutable_prompt = this.widgets.filter(
w => w.name == 'mutable_prompt'
)[0]
@@ -190,7 +193,12 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
get_position_style(
ctx,
widget_width,
y + widget_height + 24,
node.size[1]
)
)
}
}
@@ -207,7 +215,7 @@ app.registerExtension({
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid; height: 30px;min-width: 122px;
border-style: solid;height: 30px;min-width: 122px;
`
// const btn=document.createElement('button');
@@ -266,7 +274,6 @@ app.registerExtension({
},
async loadedGraphNode (node, app) {
if (node.type === 'RandomPrompt') {
}
}
})
@@ -408,7 +415,7 @@ const _createResult = async (node, widget, message) => {
const width = node.size[0] * 0.5 - 12
let height_add = 0
for (let index = 0; index < message._images.length; index++) {
const imgs = message._images[index]
@@ -559,7 +566,7 @@ app.registerExtension({
let cards = widget.div.querySelectorAll('.card')
if (cards.length == 0) node.size = [280, 120]
if(widget.value) _createResult(node, widget, widget.value)
if (widget.value) _createResult(node, widget, widget.value)
}
}
})
+206
View File
@@ -0,0 +1,206 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 14 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
// outline: '1px solid red',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
app.registerExtension({
name: 'Mixlab.3D.SaveTripoSRMesh',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'SaveTripoSRMesh') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 88, node.size[1])
)
}
// value: [],
// async serializeValue (nodeId, widgetIndex) {
// return widget.value
// }
}
widget.div = $el('div', {})
widget.div.style.width = `120px`
document.body.appendChild(widget.div)
// preview.style = `margin-top: 12px;display: flex;
// justify-content: center;
// align-items: center;background-repeat: no-repeat;background-size: contain;`
this.addCustomWidget(widget)
const onResize = this.onResize
this.onResize = () => {
widget.div.style.width = `${this.size[0]}px`
widget.div.style.height = `${this.size[1] - 112}px`
let mvs = widget.div.querySelectorAll('model-viewer')
for (const m of mvs) {
m.style.height = `${Math.round(
(this.size[1] - 112) / mvs.length
)}px`
// console.log(m.style.height)
}
// console.log('resize', this.size)
return onResize?.apply(this, arguments)
}
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
// this.isVirtualNode = true
this.serialize_widgets = false //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const r = onExecuted?.apply?.(this, arguments)
let widget = this.widgets.filter(d => d.name == 'preview')[0]
console.log('Test', widget, message)
let meshes = message.mesh
widget.div.innerHTML = ''
for (const mesh of meshes) {
if (mesh) {
const { filename, subfolder, type } = mesh
const fileURL = api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
let modelViewer = document.createElement('div')
modelViewer.innerHTML = `<model-viewer src="${fileURL}"
min-field-of-view="0deg" max-field-of-view="180deg"
shadow-intensity="1"
camera-controls
touch-action="pan-y"
style="width:100%;margin:4px;min-height:88px"
>
<div class="controls">
<div><button class="export" style="
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);cursor: pointer;">Export GLB</button></div>
</div></model-viewer>`
widget.div.appendChild(modelViewer)
let modelViewerVariants= modelViewer
.querySelector('model-viewer');
modelViewer
.querySelector('.export')
.addEventListener('click', async e => {
e.preventDefault()
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], filename)
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
}
}
// widget.value = [meshes]
this.onResize?.(this.size)
return r
}
}
},
async loadedGraphNode (node, app) {
const sleep = (t = 1000) => {
return new Promise((res, rej) => {
setTimeout(() => res(1), t)
})
}
// if (node.type === 'SaveTripoSRMesh') {
// await sleep(0)
// let widget = node.widgets.filter(w => w.name === 'preview')[0]
// widget.div.innerHTML = ''
// for (const mesh of widget.value) {
// if (mesh) {
// const { filename, subfolder, type } = mesh
// const fileURL = api.apiURL(
// `/view?filename=${encodeURIComponent(
// filename
// )}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
// )
// let modelViewer = document.createElement('div')
// modelViewer.innerHTML = `<model-viewer src="${fileURL}"
// min-field-of-view="0deg" max-field-of-view="180deg"
// shadow-intensity="1"
// camera-controls
// touch-action="pan-y">
// <div class="controls">
// <div><button class="export">Export GLB</button></div>
// </div></model-viewer>`
// widget.div.appendChild(modelViewer)
// }
// }
// }
}
})
+28 -1
View File
@@ -46,6 +46,18 @@ const smart_connect_config_input = [
node_widget_name: 'image',
inputNodeName: 'LoadImage',
inputNode_output_name: 'IMAGE'
},
{
node_type: 'TripoSRSampler_',
node_widget_name: 'image',
inputNodeName: 'LoadImagesToBatch',
inputNode_output_name: 'IMAGE'
},
{
node_type: 'TripoSRSampler_',
node_widget_name: 'mask',
inputNodeName: 'RembgNode_Mix',
inputNode_output_name: 'masks'
}
]
@@ -74,6 +86,18 @@ const smart_connect_config_output = [
outputNodeName: 'SaveImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'AppInfo',
outputNode_input_name: 'IMAGE'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'SaveImageAndMetadata_',
outputNode_input_name: 'images'
},
{
node_type: 'Moondream',
node_output_name: 'STRING',
@@ -181,7 +205,10 @@ export function smart_init () {
]
let node_slotType = config[0]
// 如果input没有,则创建
if (!node.inputs?.filter(inp => inp.name === widget.name)[0]||!node.inputs)
if (
!node.inputs?.filter(inp => inp.name === widget.name)[0] ||
!node.inputs
)
convertToInput(node, widget, config)
input_node.connectByType(inputNode_slot, node, node_slotType)
}
+298
View File
@@ -0,0 +1,298 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import WaveSurfer from 'https://cdn.jsdelivr.net/npm/wavesurfer.js@7/dist/wavesurfer.esm.js'
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: '0',
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
//把文件转为url访问
const parseUrl = data => {
let { filename, subfolder, type, prompt } = data
return {
url: api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
),
prompt
}
}
const createWaveSurfer = (wavesurfer, id,url) => {
// Create an instance of WaveSurfer
if (wavesurfer) {
wavesurfer.destroy()
}
wavesurfer = WaveSurfer.create({
container: '#' + id,
waveColor: 'rgb(200, 0, 200)',
progressColor: 'rgb(100, 0, 100)',
// Set a bar width
barWidth: 10,
// Optionally, specify the spacing between bars
barGap: 2,
// And the bar radius
barRadius: 6,
url
})
wavesurfer._auto = true
// 监听播放结束事件,重新开始播放以实现循环播放
wavesurfer.on('finish', function () {
// console.log(wavesurfer)
if (wavesurfer._auto) wavesurfer.play()
})
wavesurfer.on('interaction', () => {
wavesurfer._auto = false
if (!wavesurfer.isPlaying()) wavesurfer.play()
})
// 获取当前播放时间的峰值
wavesurfer.on('audioprocess', () => {
if (wavesurfer.isPlaying()&&wavesurfer.getDecodedData()) {
const channelData = wavesurfer.getDecodedData().getChannelData(0);
const currentTime = wavesurfer.getCurrentTime()
// console.log(wavesurfer)
const sampleRate = wavesurfer.getDecodedData().sampleRate
// 定义要分析的时间窗口(例如1秒)
const windowSize = 1
const startSample = Math.floor(currentTime * sampleRate)
const endSample = Math.min(
startSample + windowSize * sampleRate,
channelData.length
)
let peak = 0
for (let i = startSample; i < endSample; i++) {
const value = Math.abs(channelData[i])
if (value > peak) {
peak = value
}
}
// console.log('Current Peak:', peak)
}
})
return wavesurfer
}
//更新gui
function updateWaveWidgetValue (widgets, id, url, prompt, wavesurfer) {
let widget = widgets.filter(w => w.name == 'AudioPlay')[0]
// 手动更新widget值
widget.value = [url, prompt]
if (widget.div) {
widget.div.querySelector('.wave').id = `AudioPlay_${id}`
}
wavesurfer = createWaveSurfer(wavesurfer, `AudioPlay_${id}`,url)
wavesurfer.on('ready', duration => {
console.log('Audio duration: ' + duration + ' seconds')
if (widget.div) {
widget.div.setAttribute('data-url', url)
widget.div.querySelector('.link').setAttribute('href', url)
widget.div.querySelector(
'.info'
).innerHTML = `<span style="font-size: 12px;
margin: 8px;">${duration.toFixed(
2
)} seconds</span> <br><span style="font-size: 14px;">${prompt||''}</span> <br>`
}
})
wavesurfer.load(url)
// console.log('updateWaveWidgetValue' ,url,wavesurfer)
return wavesurfer
}
app.registerExtension({
name: 'SoundLab.AudioPlay',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'AudioPlay') {
let that = this
// console.log('that', that)
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'AudioPlay',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
)
}
}
// console.log('AudioPlay nodeData', this)
widget.div = $el('div', {})
document.body.appendChild(widget.div)
// wave
const waveDiv = document.createElement('div')
waveDiv.className = 'wave'
waveDiv.style.minHeight = '172px'
widget.div.appendChild(waveDiv)
//prompt 相关信息展示
const infoDiv = document.createElement('div')
infoDiv.className = 'info'
infoDiv.style.marginBottom = '20px'
widget.div.appendChild(infoDiv)
// 按钮的区域
let btns = document.createElement('div')
btns.className = 'btns'
btns.style = `display: flex;
width: 100%;
justify-content: space-between;`
widget.div.appendChild(btns)
//play button
const playBtn = document.createElement('a')
playBtn.innerText = 'Play/Pause'
playBtn.style = `
display: flex;
padding: 4px 15px;
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);
text-decoration: none;
border-radius: 5px;
transition: background-color 0.3s ease 0s;
`
playBtn.addEventListener('click', e => {
e.preventDefault()
if (that[`wavesurfer_${this.id}`]) {
that[`wavesurfer_${this.id}`]?.playPause()
that[`wavesurfer_${this.id}`]._auto = true
}
})
btns.appendChild(playBtn)
const urlLink = document.createElement('a')
urlLink.className = 'link'
urlLink.innerText = 'URL'
urlLink.setAttribute('target', '_blank')
urlLink.style = `display: flex;
padding: 4px 15px;
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);
text-decoration: none;
border-radius: 5px;
transition: background-color 0.3s ease 0s;`
// urlLink.style.minHeight = '200px'
btns.appendChild(urlLink)
//todo 导出视频 that[`wavesurfer_${this.id}`].renderer.exportImage('image/png',1,'dataURL')
// https://github.com/diffusion-studio/ffmpeg-js
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
this.size = [this.size[0], 280]
this.serialize_widgets = true //需保存widget的值
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
const audio = message.audio
console.log('#onExecuted', `AudioPlay_${this.id}`, message,audio)
try {
let { url, prompt } = parseUrl(audio[0])
that[`wavesurfer_${this.id}`] = updateWaveWidgetValue(
this.widgets,
this.id,
url,
prompt,
that[`wavesurfer_${this.id}`]
)
that[`wavesurfer_${this.id}`]?.playPause()
} catch (error) {
console.log(error)
}
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'AudioPlay') {
let widget = node.widgets.filter(w => w.name == 'AudioPlay')[0]
if (widget.value) {
let [url, prompt] = widget.value
this[`wavesurfer_${node.id}`] = updateWaveWidgetValue(
node.widgets,
node.id,
url,
prompt,
this[`wavesurfer_${node.id}`]
)
}
console.log('#loadedGraphNode', node)
}
}
})
+323
View File
@@ -0,0 +1,323 @@
// touchdesigner的背景效果,把appinfo的输出,选择一张图片作为背景
window._bg_img = null
/**
* draws the back canvas (the one containing the background and the connections)
* @method drawBackCanvas
**/
LGraphCanvas.prototype.drawBackCanvas = function () {
var canvas = this.bgcanvas
if (
canvas.width != this.canvas.width ||
canvas.height != this.canvas.height
) {
canvas.width = this.canvas.width
canvas.height = this.canvas.height
}
if (!this.bgctx) {
this.bgctx = this.bgcanvas.getContext('2d')
}
var ctx = this.bgctx
if (ctx.start) {
ctx.start()
}
var viewport = this.viewport || [0, 0, ctx.canvas.width, ctx.canvas.height]
//clear
if (this.clear_background) {
ctx.clearRect(viewport[0], viewport[1], viewport[2], viewport[3])
}
//show subgraph stack header
if (this._graph_stack && this._graph_stack.length) {
ctx.save()
var parent_graph = this._graph_stack[this._graph_stack.length - 1]
var subgraph_node = this.graph._subgraph_node
ctx.strokeStyle = subgraph_node.bgcolor
ctx.lineWidth = 10
ctx.strokeRect(1, 1, canvas.width - 2, canvas.height - 2)
ctx.lineWidth = 1
ctx.font = '40px Arial'
ctx.textAlign = 'center'
ctx.fillStyle = subgraph_node.bgcolor || '#AAA'
var title = ''
for (var i = 1; i < this._graph_stack.length; ++i) {
title += this._graph_stack[i]._subgraph_node.getTitle() + ' >> '
}
ctx.fillText(title + subgraph_node.getTitle(), canvas.width * 0.5, 40)
ctx.restore()
}
var bg_already_painted = false
if (this.onRenderBackground) {
bg_already_painted = this.onRenderBackground(canvas, ctx)
}
//reset in case of error
if (!this.viewport) {
ctx.restore()
ctx.setTransform(1, 0, 0, 1, 0, 0)
}
this.visible_links.length = 0
if (this.graph) {
//apply transformations
ctx.save()
this.ds.toCanvasContext(ctx)
//render BG
if (
this.ds.scale < 1 &&
!bg_already_painted &&
this.clear_background_color
) {
ctx.fillStyle = this.clear_background_color
ctx.fillRect(
this.visible_area[0],
this.visible_area[1],
this.visible_area[2],
this.visible_area[3]
)
}
// 主要修改
if (this.background_image && this.ds.scale > 0.5 && !bg_already_painted) {
if (this.zoom_modify_alpha) {
//使得 alpha 越接近0时变化越缓慢。
let alpha = (1.0 - 0.5 / this.ds.scale) * this.editor_alpha
ctx.globalAlpha = Math.min(Math.max(0, Math.sqrt(alpha)), 1)
// console.log((1.0 - 0.5 / this.ds.scale) * this.editor_alpha)
} else {
ctx.globalAlpha = this.editor_alpha
}
ctx.imageSmoothingEnabled = ctx.imageSmoothingEnabled = false // ctx.mozImageSmoothingEnabled =
if (!this._bg_img || this._bg_img.name != this.background_image) {
this._bg_img = new Image()
this._bg_img.name = this.background_image
this._bg_img.src = this.background_image
var that = this
this._bg_img.onload = function () {
that.draw(true, true)
}
}
var pattern = null
if (this._pattern == null && this._bg_img.width > 0) {
pattern = ctx.createPattern(this._bg_img, 'repeat')
this._pattern_img = this._bg_img
this._pattern = pattern
} else {
pattern = this._pattern
}
if (pattern) {
ctx.fillStyle = pattern
ctx.fillRect(
this.visible_area[0],
this.visible_area[1],
this.visible_area[2],
this.visible_area[3]
)
ctx.fillStyle = 'transparent'
}
ctx.globalAlpha = 1.0
ctx.imageSmoothingEnabled = ctx.imageSmoothingEnabled = true //= ctx.mozImageSmoothingEnabled
}
//groups
if (this.graph._groups.length && !this.live_mode) {
this.drawGroups(canvas, ctx)
}
if (this.onDrawBackground) {
this.onDrawBackground(ctx, this.visible_area)
}
if (this.onBackgroundRender) {
//LEGACY
console.error(
'WARNING! onBackgroundRender deprecated, now is named onDrawBackground '
)
this.onBackgroundRender = null
}
//DEBUG: show clipping area
//ctx.fillStyle = "red";
//ctx.fillRect( this.visible_area[0] + 10, this.visible_area[1] + 10, this.visible_area[2] - 20, this.visible_area[3] - 20);
//bg
if (this.render_canvas_border) {
ctx.strokeStyle = '#235'
ctx.strokeRect(0, 0, canvas.width, canvas.height)
}
if (this.render_connections_shadows) {
ctx.shadowColor = '#000'
ctx.shadowOffsetX = 0
ctx.shadowOffsetY = 0
ctx.shadowBlur = 6
} else {
ctx.shadowColor = 'rgba(0,0,0,0)'
}
//draw connections
if (!this.live_mode) {
this.drawConnections(ctx)
}
ctx.shadowColor = 'rgba(0,0,0,0)'
//restore state
ctx.restore()
}
if (ctx.finish) {
ctx.finish()
}
this.dirty_bgcanvas = false
this.dirty_canvas = true //to force to repaint the front canvas with the bgcanvas
}
function imgToCanvasBase64 (img) {
const canvas = document.createElement('canvas')
const ctx = canvas.getContext('2d')
canvas.width = img.width
canvas.height = img.height
ctx.drawImage(img, 0, 0)
const base64 = canvas.toDataURL('image/png')
return base64
}
// 使用示例
function convertImageToBase64 (img) {
// const img = new Image()
// img.src = 'path/to/your/image.jpg' // 替换为你的图片路径
// console.log('convertImageToBase64',img)
try {
const base64 = imgToCanvasBase64(img)
return base64
} catch (error) {
console.error(error)
}
}
function getInputsAndOutputs () {
const outputs =
`PreviewImage,SaveImage,TransparentImage,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_`.split(
','
)
let outputsId = []
for (let node of app.graph._nodes) {
if (outputs.includes(node.type)) {
outputsId.push(node.id)
}
}
return outputsId
}
function getRandomElement (arr) {
const randomIndex = Math.floor(Math.random() * arr.length)
return arr[randomIndex]
}
async function getBG () {
var outputs = []
for (let id of app.graph
.getNodeById(50)
.widgets.filter(w => w.name === 'output_ids')[0]
.value.split('\n')) {
if (getInputsAndOutputs().map(Number).includes(Number(id))) {
if (app.graph.getNodeById(id).imgs && app.graph.getNodeById(id).imgs[0]) {
let b = convertImageToBase64(app.graph.getNodeById(id).imgs[0])
// console.log(b)
outputs.push(b)
}
}
}
var BACKGROUND_IMAGE = getRandomElement(outputs),
CLEAR_BACKGROUND_COLOR = 'rgba(0,0,0,0.9)'
if (!window._bg_img) {
window._bg_img = app.canvas._bg_img.src
}
// let img=new Image();
// img.src=BACKGROUND_IMAGE;
//去掉透明度过度
// app.canvas.zoom_modify_alpha=false;
//整体透明度
app.canvas.editor_alpha = 1.1
// app.canvas._pattern=ctx.createPattern(img, "no-repeat");
app.canvas.updateBackground(BACKGROUND_IMAGE, CLEAR_BACKGROUND_COLOR)
app.canvas.draw(true, true)
}
class BgRunner {
constructor () {
this.intervalId = null
this.running = false
}
// 要运行的方法
bg () {
console.log('方法bg正在运行')
getBG()
}
// 启动bg方法每秒运行一次
start () {
if (!this.running) {
this.intervalId = setInterval(() => this.bg(), 1500)
this.running = true
}
}
// 停止bg方法的运行
stop () {
if (this.running) {
clearInterval(this.intervalId)
this.intervalId = null
this.running = false
if (window._bg_img) {
var BACKGROUND_IMAGE = window._bg_img,
CLEAR_BACKGROUND_COLOR = 'rgba(0,0,0,1)'
app.canvas.editor_alpha = 1
app.canvas.updateBackground(BACKGROUND_IMAGE, CLEAR_BACKGROUND_COLOR)
app.canvas.draw(true, true)
}
}
}
// 切换start和stop
toggle () {
if (this.running) {
this.stop()
} else {
this.start()
}
}
// 获取运行状态
isRunning () {
return this.running
}
}
// 示例用法
// const runner = new BgRunner();
// runner.start();
// setTimeout(() => runner.stop(), 5000);
export const td_bg = new BgRunner()
+1025 -203
View File
File diff suppressed because it is too large Load Diff
+163 -67
View File
@@ -1,46 +1,13 @@
import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { $el } from '../../../scripts/ui.js'
import {
loadExternalScript,
updateLLMAPIKey,
get_position_style,
getLocalData
} from './common.js'
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
loadExternalScript('/mixlab/app/lib/pickr.min.js')
function hexToRGBA (hexColor) {
var hex = hexColor.replace('#', '')
@@ -62,7 +29,7 @@ app.registerExtension({
init () {
$el('link', {
rel: 'stylesheet',
href: '/extensions/comfyui-mixlab-nodes/lib/classic.min.css',
href: '/mixlab/app/lib/classic.min.css',
parent: document.head
})
@@ -122,7 +89,7 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
// console.log('Color nodeData', this.widgets)
// console.log('Color nodeData', this.div)
const widget = {
type: 'div',
@@ -273,19 +240,19 @@ app.registerExtension({
})
const min_max = node => {
if(node.widgets){
if (node.widgets) {
const min_value = node.widgets.filter(w => w.name === 'min_value')[0]
const max_value = node.widgets.filter(w => w.name === 'max_value')[0]
const number = node.widgets.filter(w => w.name === 'number')[0]
if (number) {
number.options.min = min_value.value
number.options.max = max_value.value
number.value = Math.min(number.options.max, number.value)
number.value = Math.max(number.options.min, number.value)
}
if (min_value)
min_value.callback = e => {
number.options.min = e
@@ -297,22 +264,18 @@ const min_max = node => {
number.value = e
}
}
}
app.registerExtension({
name: 'Mixlab.utils.FloatSlider',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'FloatSlider') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated;
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
min_max(this)
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'FloatSlider') {
@@ -323,7 +286,6 @@ app.registerExtension({
app.registerExtension({
name: 'Mixlab.utils.IntNumber',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'IntNumber') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
@@ -331,7 +293,6 @@ app.registerExtension({
min_max(this)
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'IntNumber') {
@@ -340,22 +301,157 @@ app.registerExtension({
}
})
app.registerExtension({
name: 'Mixlab.utils.TESTNODE_',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'TESTNODE_') {
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
console.log('##',message)
};
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
console.log('##', message)
}
}
},
}
})
app.registerExtension({
name: 'Mixlab.utils.KeyInput',
init () {},
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##node', node)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 24], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_llm_api_key')
return data[node.id] || 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'KeyInput') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const rowHeight = this.rowHeight
const widget = {
type: 'div',
name: 'input_key',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 24, node.size[1])
)
}
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder) => {
let div = document.createElement('div')
div.style = `
display: flex;
align-items: center;
margin: 6px 8px;
margin-top:0px;
height:44px;
width:220px;
`
const ip = document.createElement('input')
ip.type = 'password'
ip.className = `${'comfy-multiline-input'} ${placeholder}`
ip.placeholder = placeholder
// ip.value = placeholder
ip.style = `margin-left:8px;
outline: none;
border: none;
padding:12px;
width: 100%;
`
div.appendChild(ip)
ip.addEventListener('change', () => {
let data = getLocalData(key)
data[this.id] = ip.value.trim()
localStorage.setItem(key, JSON.stringify(data))
updateLLMAPIKey(data[this.id])
})
return div
}
let inputKey = inputDiv('_mixlab_llm_api_key', 'Key')
widget.div.appendChild(inputKey)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputKey.remove()
widget.div.remove()
return onRemoved?.()
}
// const processMouseWheel=app.canvas.processMouseWheel
// app.canvas.processMouseWheel=()=>{
// console.log(app.canvas.ds.scale)
// return processMouseWheel?.()
// }
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'KeyInput') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_llm_api_key')
let id = node.id
if (widget.div.querySelector('.Key'))
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
if (apiKey[id]) updateLLMAPIKey(apiKey[id])
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
// console.log('#LoadImagesToBatch', node.type)
if (node.type === 'KeyInput') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_llm_api_key')
let id = node.id
if (widget.div.querySelector('.Key'))
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
if (apiKey[id]) updateLLMAPIKey(apiKey[id])
}
}, 1000)
}
})
+49 -50
View File
@@ -6,8 +6,6 @@ import { $el } from '../../../scripts/ui.js'
// The code is based on ComfyUI-VideoHelperSuite modification.
function injectCSS (css) {
// 检查页面中是否已经存在具有相同内容的style标签
const existingStyle = document.querySelector('style')
@@ -45,7 +43,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -240,15 +241,7 @@ app.registerExtension({
}
})
function offsetDOMWidget(
widget,
ctx,
node,
widgetWidth,
widgetY,
height
) {
function offsetDOMWidget (widget, ctx, node, widgetWidth, widgetY, height) {
const margin = 10
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
@@ -270,18 +263,18 @@ function offsetDOMWidget(
position: 'absolute',
background: !node.color ? '' : node.color,
color: !node.color ? '' : 'white',
zIndex: 5, //app.graph._nodes.indexOf(node),
zIndex: 5 //app.graph._nodes.indexOf(node),
})
}
export const hasWidgets = (node) => {
export const hasWidgets = node => {
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
return false
}
return true
}
export const cleanupNode = (node) => {
export const cleanupNode = node => {
if (!hasWidgets(node)) {
return
}
@@ -298,43 +291,43 @@ export const cleanupNode = (node) => {
}
}
const CreatePreviewElement = (name, val, format) => {
const [type] = format.split('/')
const createPreviewElement = (name, val, format) => {
const [type] = format.split('/')
const w = {
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth)
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch)
},
computeSize: function (_) {
const ratio = this.inputRatio || 1
const width = Math.max(220, this.parent.size[0])
return [width, (width / ratio + 10)]
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove()
}
},
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth)
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch)
},
computeSize: function (_) {
const ratio = this.inputRatio || 1
const width = Math.max(220, this.parent.size[0])
return [width, width / ratio + 10]
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove()
}
}
w.inputEl = document.createElement(type === 'video' ? 'video' : 'img')
w.inputEl.src = w.value
if (type === 'video') {
w.inputEl.setAttribute('type', 'video/webm');
w.inputEl.autoplay = true
w.inputEl.loop = true
w.inputEl.controls = false;
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight
}
document.body.appendChild(w.inputEl)
return w
}
w.inputEl = document.createElement(type === 'video' ? 'video' : 'img')
w.inputEl.src = w.value
if (type === 'video' || format.match('.mp4')) {
w.inputEl.setAttribute('type', 'video/webm')
w.inputEl.autoplay = true
w.inputEl.loop = true
w.inputEl.controls = true
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight
}
document.body.appendChild(w.inputEl)
return w
}
app.registerExtension({
name: 'Mixlab.Video.ImageListReplace',
@@ -469,12 +462,17 @@ app.registerExtension({
}
}
if (nodeData?.name == 'VideoCombine_Adv') {
if (
nodeData?.name == 'VideoCombine_Adv' ||
nodeData?.name == 'CombineAudioVideo'
) {
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const prefix = 'vhs_gif_preview_'
const r = onExecuted ? onExecuted.apply(this, message) : undefined
if(!this.widgets) this.widgets=[]
if (this.widgets) {
const pos = this.widgets.findIndex(w => w.name === `${prefix}_0`)
if (pos !== -1) {
@@ -489,12 +487,13 @@ app.registerExtension({
'/view?' + new URLSearchParams(params).toString()
)
const w = this.addCustomWidget(
CreatePreviewElement(
createPreviewElement(
`${prefix}_${i}`,
previewUrl,
params.format || 'image/gif'
)
)
console.log(w)
w.parent = this
})
}
+237
View File
@@ -0,0 +1,237 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
function base64ToBlobFromURL (base64URL, contentType) {
return fetch(base64URL).then(response => response.blob())
}
async function uploadImage (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
const body = new FormData()
body.append(
'image',
new File([blob], (filename || new Date().getTime()) + fileType)
)
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body
})
// console.log(resp)
let data = await resp.json()
return data
}
// 上传得到url
async function uploadBase64ToFile (base64) {
let bg_blob = await base64ToBlobFromURL(base64)
let url = await uploadImage(bg_blob, '.png')
return url
}
class Visualizer {
constructor (node, container, visualSrc) {
this.node = node
this.iframe = document.createElement('iframe')
Object.assign(this.iframe, {
scrolling: 'no',
overflow: 'hidden'
})
this.iframe.src = '/mixlab/app/' + visualSrc + '.html'
console.log('#Visualizer', container, this.iframe)
container.appendChild(this.iframe)
}
updateVisual (params) {
console.log('#updateVisual', params, this.iframe)
// const iframeDocument = this.iframe.contentWindow.document
// const previewScript = iframeDocument.getElementById('visualizer')
// previewScript.setAttribute(
// 'reference_image',
// JSON.stringify(params.reference_image)
// )
// previewScript.setAttribute('depth_map', JSON.stringify(params.depth_map))
// Update the reference image and depth map
this.iframe.contentWindow.postMessage(params, '*')
}
remove () {
this.container.remove()
}
}
function createVisualizer (node, inputName, typeName, inputData, app) {
node.name = inputName
const widget = {
type: typeName,
name: 'preview3d',
callback: () => {},
draw: function (ctx, node, widgetWidth, widgetY, widgetHeight) {
const margin = 10
const top_offset = 5
const visible = app.canvas.ds.scale > 0.5 && this.type === typeName
const w = widgetWidth - margin * 4
const clientRectBound = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
clientRectBound.width / ctx.canvas.width,
clientRectBound.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(margin, margin + widgetY)
Object.assign(this.visualizer.style, {
left: `${transform.a * margin + transform.e}px`,
top: `${transform.d + transform.f + top_offset}px`,
width: `${w * transform.a}px`,
height: `${
w * transform.d - widgetHeight - margin * 15 * transform.d
}px`,
position: 'absolute',
overflow: 'hidden',
zIndex: app.graph._nodes.indexOf(node)
})
Object.assign(this.visualizer.children[0].style, {
transformOrigin: '50% 50%',
width: '100%',
height: '100%',
border: '0 none'
})
this.visualizer.hidden = !visible
}
}
const container = document.createElement('div')
container.id = `Comfy3D_${inputName}`
node.visualizer = new Visualizer(node, container, typeName)
widget.visualizer = container
widget.parent = node
document.body.appendChild(widget.visualizer)
node.addCustomWidget(widget)
node.updateParameters = params => {
// console.log('#updateParameters', params)
params.id = node.id
// node.visualizer = new Visualizer(node, container, typeName)
node.visualizer.updateVisual(params)
}
// Events for drawing backgound
node.onDrawBackground = function (ctx) {
if (!this.flags.collapsed) {
node.visualizer.iframe.hidden = false
} else {
node.visualizer.iframe.hidden = true
}
}
// Make sure visualization iframe is always inside the node when resize the node
node.onResize = function () {
let [w, h] = this.size
if (w <= 600) w = 600
if (h <= 500) h = 500
if (w > 600) {
h = w - 100
}
this.size = [w, h]
}
// Events for remove nodes
node.onRemoved = () => {
for (let w in node.widgets) {
if (node.widgets[w].visualizer) {
node.widgets[w].visualizer.remove()
}
}
}
return {
widget: widget
}
}
function registerVisualizer (nodeType, nodeData, nodeClassName, typeName) {
if (nodeData.name == nodeClassName) {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined
let Preview3DNode = app.graph._nodes.filter(
wi => wi.type == nodeClassName
)
let nodeName = `Preview3DNode_${Preview3DNode.length}`
const result = await createVisualizer.apply(this, [
this,
nodeName,
typeName,
{},
app
])
this.setSize([600, 500])
return r
}
nodeType.prototype.onExecuted = async function (message) {
// Check if reference image and depth map are available
if (message.reference_image && message.depth_map) {
const params = {}
params.reference_image = message.reference_image[0]
params.depth_map = message.depth_map[0]
this.updateParameters(params)
}
}
}
}
app.registerExtension({
name: 'Mixlab.nodes.depthviewer',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
registerVisualizer(nodeType, nodeData, 'DepthViewer', 'threeVisualizer')
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
let widget = node.widgets?.filter(w => w.name == 'preview3d')[0]
let framesWidget = node.widgets?.filter(w => w.name == 'frames')[0]
if (node.type === 'DepthViewer' && widget) {
let nodeId = node.id
//延迟才能获得this.id
widget.visualizer.querySelector('iframe').src += '?id=' + nodeId
// console.log('DepthViewer',widget)
window.addEventListener('message', async event => {
// 检查消息的来源,确保消息来自可信的源
console.log(event)
const { id, imgs } = event.data
if (id == nodeId) {
framesWidget.value = { images: [] }
for (const f of imgs) {
let file = await uploadBase64ToFile(f)
framesWidget.value.images.push(file)
}
// framesWidget.value.base64 = frames
framesWidget.value._seed = Math.random()
node.title = 'Input #' + imgs.length
}
})
}
}, 1000)
}
})
+5 -2
View File
@@ -34,7 +34,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -125,7 +128,7 @@ app.registerExtension({
window._mixlab_file_path_watcher = json.event_type
// widget.card.innerText = window._mixlab_file_path_watcher || ''
//运行
// document.querySelector('#queue-button').click()
if (app) app.queuePrompt()
}
})
}, 1000)
-1077
View File
File diff suppressed because one or more lines are too long
+22
View File
@@ -0,0 +1,22 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Mixlab AR</title>
</head>
<body>
<script type="module">
import { api } from "/mixlab/app/javascript/api.js";
import Command from '/mixlab/app/javascript/command.js'
</script>
</body>
</html>
File diff suppressed because it is too large Load Diff
+482
View File
@@ -0,0 +1,482 @@
class ComfyApi extends EventTarget {
#registered = new Set();
constructor() {
super();
this.api_host = location.host;
this.api_base = location.pathname.split('/').slice(0, -1).join('/');
this.initialClientId = sessionStorage.getItem("clientId");
}
apiURL(route) {
return this.api_base + route;
}
fetchApi(route, options) {
if (!options) {
options = {};
}
if (!options.headers) {
options.headers = {};
}
options.headers["Comfy-User"] = this.user;
return fetch(this.apiURL(route), options);
}
addEventListener(type, callback, options) {
super.addEventListener(type, callback, options);
this.#registered.add(type);
}
/**
* Poll status for colab and other things that don't support websockets.
*/
#pollQueue() {
setInterval(async () => {
try {
const resp = await this.fetchApi("/prompt");
const status = await resp.json();
this.dispatchEvent(new CustomEvent("status", { detail: status }));
} catch (error) {
this.dispatchEvent(new CustomEvent("status", { detail: null }));
}
}, 1000);
}
/**
* Creates and connects a WebSocket for realtime updates
* @param {boolean} isReconnect If the socket is connection is a reconnect attempt
*/
#createSocket(isReconnect) {
if (this.socket) {
return;
}
let opened = false;
let existingSession = window.name;
if (existingSession) {
existingSession = "?clientId=" + existingSession;
}
this.socket = new WebSocket(
`ws${window.location.protocol === "https:" ? "s" : ""}://${this.api_host}${this.api_base}/ws${existingSession}`
);
this.socket.binaryType = "arraybuffer";
this.socket.addEventListener("open", () => {
opened = true;
if (isReconnect) {
this.dispatchEvent(new CustomEvent("reconnected"));
}
});
this.socket.addEventListener("error", () => {
if (this.socket) this.socket.close();
if (!isReconnect && !opened) {
this.#pollQueue();
}
});
this.socket.addEventListener("close", () => {
setTimeout(() => {
this.socket = null;
this.#createSocket(true);
}, 300);
if (opened) {
this.dispatchEvent(new CustomEvent("status", { detail: null }));
this.dispatchEvent(new CustomEvent("reconnecting"));
}
});
this.socket.addEventListener("message", (event) => {
try {
if (event.data instanceof ArrayBuffer) {
const view = new DataView(event.data);
const eventType = view.getUint32(0);
const buffer = event.data.slice(4);
switch (eventType) {
case 1:
const view2 = new DataView(event.data);
const imageType = view2.getUint32(0)
let imageMime
switch (imageType) {
case 1:
default:
imageMime = "image/jpeg";
break;
case 2:
imageMime = "image/png"
}
const imageBlob = new Blob([buffer.slice(4)], { type: imageMime });
this.dispatchEvent(new CustomEvent("b_preview", { detail: imageBlob }));
break;
default:
throw new Error(`Unknown binary websocket message of type ${eventType}`);
}
}
else {
const msg = JSON.parse(event.data);
switch (msg.type) {
case "status":
if (msg.data.sid) {
this.clientId = msg.data.sid;
window.name = this.clientId; // use window name so it isnt reused when duplicating tabs
sessionStorage.setItem("clientId", this.clientId); // store in session storage so duplicate tab can load correct workflow
}
this.dispatchEvent(new CustomEvent("status", { detail: msg.data.status }));
break;
case "progress":
this.dispatchEvent(new CustomEvent("progress", { detail: msg.data }));
break;
case "executing":
this.dispatchEvent(new CustomEvent("executing", { detail: msg.data.node }));
break;
case "executed":
this.dispatchEvent(new CustomEvent("executed", { detail: msg.data }));
break;
case "execution_start":
this.dispatchEvent(new CustomEvent("execution_start", { detail: msg.data }));
break;
case "execution_success":
this.dispatchEvent(new CustomEvent("execution_success", { detail: msg.data }));
break;
case "execution_error":
this.dispatchEvent(new CustomEvent("execution_error", { detail: msg.data }));
break;
case "execution_cached":
this.dispatchEvent(new CustomEvent("execution_cached", { detail: msg.data }));
break;
default:
if (this.#registered.has(msg.type)) {
this.dispatchEvent(new CustomEvent(msg.type, { detail: msg.data }));
} else {
throw new Error(`Unknown message type ${msg.type}`);
}
}
}
} catch (error) {
console.warn("Unhandled message:", event.data, error);
}
});
}
/**
* Initialises sockets and realtime updates
*/
init() {
this.#createSocket();
}
/**
* Gets a list of extension urls
* @returns An array of script urls to import
*/
async getExtensions() {
const resp = await this.fetchApi("/extensions", { cache: "no-store" });
return await resp.json();
}
/**
* Gets a list of embedding names
* @returns An array of script urls to import
*/
async getEmbeddings() {
const resp = await this.fetchApi("/embeddings", { cache: "no-store" });
return await resp.json();
}
/**
* Loads node object definitions for the graph
* @returns The node definitions
*/
async getNodeDefs() {
const resp = await this.fetchApi("/object_info", { cache: "no-store" });
return await resp.json();
}
/**
*
* @param {number} number The index at which to queue the prompt, passing -1 will insert the prompt at the front of the queue
* @param {object} prompt The prompt data to queue
*/
async queuePrompt(number, { output, workflow }) {
const body = {
client_id: this.clientId,
prompt: output,
extra_data: { extra_pnginfo: { workflow } },
};
if (number === -1) {
body.front = true;
} else if (number != 0) {
body.number = number;
}
const res = await this.fetchApi("/prompt", {
method: "POST",
headers: {
"Content-Type": "application/json",
},
body: JSON.stringify(body),
});
if (res.status !== 200) {
throw {
response: await res.json(),
};
}
return await res.json();
}
/**
* Loads a list of items (queue or history)
* @param {string} type The type of items to load, queue or history
* @returns The items of the specified type grouped by their status
*/
async getItems(type) {
if (type === "queue") {
return this.getQueue();
}
return this.getHistory();
}
/**
* Gets the current state of the queue
* @returns The currently running and queued items
*/
async getQueue() {
try {
const res = await this.fetchApi("/queue");
const data = await res.json();
return {
// Running action uses a different endpoint for cancelling
Running: data.queue_running.map((prompt) => ({
prompt,
remove: { name: "Cancel", cb: () => api.interrupt() },
})),
Pending: data.queue_pending.map((prompt) => ({ prompt })),
};
} catch (error) {
console.error(error);
return { Running: [], Pending: [] };
}
}
/**
* Gets the prompt execution history
* @returns Prompt history including node outputs
*/
async getHistory(max_items=200) {
try {
const res = await this.fetchApi(`/history?max_items=${max_items}`);
return { History: Object.values(await res.json()) };
} catch (error) {
console.error(error);
return { History: [] };
}
}
/**
* Gets system & device stats
* @returns System stats such as python version, OS, per device info
*/
async getSystemStats() {
const res = await this.fetchApi("/system_stats");
return await res.json();
}
/**
* Sends a POST request to the API
* @param {*} type The endpoint to post to
* @param {*} body Optional POST data
*/
async #postItem(type, body) {
try {
await this.fetchApi("/" + type, {
method: "POST",
headers: {
"Content-Type": "application/json",
},
body: body ? JSON.stringify(body) : undefined,
});
} catch (error) {
console.error(error);
}
}
/**
* Deletes an item from the specified list
* @param {string} type The type of item to delete, queue or history
* @param {number} id The id of the item to delete
*/
async deleteItem(type, id) {
await this.#postItem(type, { delete: [id] });
}
/**
* Clears the specified list
* @param {string} type The type of list to clear, queue or history
*/
async clearItems(type) {
await this.#postItem(type, { clear: true });
}
/**
* Interrupts the execution of the running prompt
*/
async interrupt() {
await this.#postItem("interrupt", null);
}
/**
* Gets user configuration data and where data should be stored
* @returns { Promise<{ storage: "server" | "browser", users?: Promise<string, unknown>, migrated?: boolean }> }
*/
async getUserConfig() {
return (await this.fetchApi("/users")).json();
}
/**
* Creates a new user
* @param { string } username
* @returns The fetch response
*/
createUser(username) {
return this.fetchApi("/users", {
method: "POST",
headers: {
"Content-Type": "application/json",
},
body: JSON.stringify({ username }),
});
}
/**
* Gets all setting values for the current user
* @returns { Promise<string, unknown> } A dictionary of id -> value
*/
async getSettings() {
return (await this.fetchApi("/settings")).json();
}
/**
* Gets a setting for the current user
* @param { string } id The id of the setting to fetch
* @returns { Promise<unknown> } The setting value
*/
async getSetting(id) {
return (await this.fetchApi(`/settings/${encodeURIComponent(id)}`)).json();
}
/**
* Stores a dictionary of settings for the current user
* @param { Record<string, unknown> } settings Dictionary of setting id -> value to save
* @returns { Promise<void> }
*/
async storeSettings(settings) {
return this.fetchApi(`/settings`, {
method: "POST",
body: JSON.stringify(settings)
});
}
/**
* Stores a setting for the current user
* @param { string } id The id of the setting to update
* @param { unknown } value The value of the setting
* @returns { Promise<void> }
*/
async storeSetting(id, value) {
return this.fetchApi(`/settings/${encodeURIComponent(id)}`, {
method: "POST",
body: JSON.stringify(value)
});
}
/**
* Gets a user data file for the current user
* @param { string } file The name of the userdata file to load
* @param { RequestInit } [options]
* @returns { Promise<Response> } The fetch response object
*/
async getUserData(file, options) {
return this.fetchApi(`/userdata/${encodeURIComponent(file)}`, options);
}
/**
* Stores a user data file for the current user
* @param { string } file The name of the userdata file to save
* @param { unknown } data The data to save to the file
* @param { RequestInit & { overwrite?: boolean, stringify?: boolean, throwOnError?: boolean } } [options]
* @returns { Promise<Response> }
*/
async storeUserData(file, data, options = { overwrite: true, stringify: true, throwOnError: true }) {
const resp = await this.fetchApi(`/userdata/${encodeURIComponent(file)}?overwrite=${options?.overwrite}`, {
method: "POST",
body: options?.stringify ? JSON.stringify(data) : data,
...options,
});
if (resp.status !== 200 && options?.throwOnError !== false) {
throw new Error(`Error storing user data file '${file}': ${resp.status} ${(await resp).statusText}`);
}
return resp;
}
/**
* Deletes a user data file for the current user
* @param { string } file The name of the userdata file to delete
*/
async deleteUserData(file) {
const resp = await this.fetchApi(`/userdata/${encodeURIComponent(file)}`, {
method: "DELETE",
});
if (resp.status !== 204) {
throw new Error(`Error removing user data file '${file}': ${resp.status} ${(resp).statusText}`);
}
}
/**
* Move a user data file for the current user
* @param { string } source The userdata file to move
* @param { string } dest The destination for the file
*/
async moveUserData(source, dest, options = { overwrite: false }) {
const resp = await this.fetchApi(`/userdata/${encodeURIComponent(source)}/move/${encodeURIComponent(dest)}?overwrite=${options?.overwrite}`, {
method: "POST",
});
return resp;
}
/**
* @overload
* Lists user data files for the current user
* @param { string } dir The directory in which to list files
* @param { boolean } [recurse] If the listing should be recursive
* @param { true } [split] If the paths should be split based on the os path separator
* @returns { Promise<string[][]>> } The list of split file paths in the format [fullPath, ...splitPath]
*/
/**
* @overload
* Lists user data files for the current user
* @param { string } dir The directory in which to list files
* @param { boolean } [recurse] If the listing should be recursive
* @param { false | undefined } [split] If the paths should be split based on the os path separator
* @returns { Promise<string[]>> } The list of files
*/
async listUserData(dir, recurse, split) {
const resp = await this.fetchApi(
`/userdata?${new URLSearchParams({
recurse,
dir,
split,
})}`
);
if (resp.status === 404) return [];
if (resp.status !== 200) {
throw new Error(`Error getting user data list '${dir}': ${resp.status} ${resp.statusText}`);
}
return resp.json();
}
}
export const api = new ComfyApi();
+697
View File
@@ -0,0 +1,697 @@
function get_url () {
// 如果有缓存记录
let hostUrl = localStorage.getItem('_hostUrl') || ''
if (hostUrl) {
return hostUrl
}
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
function getFilenameAndCategoryFromUrl (url) {
const queryString = url.split('?')[1]
if (!queryString) {
return {}
}
const params = new URLSearchParams(queryString)
const filename = params.get('filename')
? decodeURIComponent(params.get('filename'))
: null
const category = params.get('category')
? decodeURIComponent(params.get('category') || '')
: ''
return { category, filename }
}
async function get_my_app (category = '', filename = null) {
let url = get_url()
const res = await fetch(`${url}/mixlab/workflow`, {
method: 'POST',
mode: 'cors', // 允许跨域请求
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
task: 'my_app',
filename,
category
})
})
let result = await res.json()
let data = []
try {
for (const res of result.data) {
let { output, app } = res.data
if (app.filename)
data.push({
...app,
data: output,
date: res.date
})
}
} catch (error) {}
return data
}
async function getAppInit () {
const { category, filename } = getFilenameAndCategoryFromUrl(
window.location.href
)
return await get_my_app(category, filename)
}
function success (isSuccess, btn, text) {
isSuccess ? (btn.innerText = 'success') : text
setTimeout(() => {
btn.innerText = text
}, 5000)
}
async function interrupt () {
try {
await fetch(`${get_url()}/interrupt`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: undefined
})
} catch (error) {
console.error(error)
}
return true
}
async function getQueue (clientId) {
try {
const res = await fetch(`${get_url()}/queue`)
const data = await res.json()
return {
// Running action uses a different endpoint for cancelling
Running: Array.from(data.queue_running, prompt => {
if (prompt[3].client_id === clientId) {
let prompt_id = prompt[1]
return {
prompt_id,
remove: () => interrupt()
}
}
}),
Pending: data.queue_pending.map(prompt => ({ prompt }))
}
} catch (error) {
console.error(error)
return { Running: [], Pending: [] }
}
}
// 请求历史数据
async function getPromptResult (category) {
let url = get_url()
try {
const response = await fetch(`${url}/mixlab/prompt_result`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
action: 'all'
})
})
if (response.ok) {
const data = await response.json()
console.log('#getPromptResult:', category, data)
return data.result.filter(r => r.appInfo.category == category)
// 处理返回的数据
} else {
console.log('Error:', response.status)
// 处理错误情况
}
} catch (error) {
console.log('Error:', error)
// 处理异常情况
}
}
// 新的运行工作流的接口
function queuePromptNew (
filename,
category,
seed,
input,
client_id,
apps = null
) {
let url = get_url()
// var filename = "Text-to-Image_1.json", category = "";
// 随机seed
// promptWorkflow = randomSeed(seed, promptWorkflow);
let d = { filename, category, seed, input, client_id }
if (apps) {
d.apps = apps
}
const data = JSON.stringify(d)
return new Promise((res, rej) => {
fetch(`${url}/mixlab/prompt`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: data
})
.then(response => {
if (!response.ok) {
// Handle HTTP error responses
if (response.status === 400) {
return response.json().then(errorData => {
// Process the error data
console.error('Error 400:', errorData)
alert(JSON.stringify(errorData, null, 2))
res(null)
})
}
throw new Error('Network response was not ok')
}
return response.json() // Process the response data
})
.then(data => {
// Handle the response data
console.log('Success:', data)
res(true)
})
.catch(error => {
// Handle fetch errors
console.error('Fetch error:', error)
res(null)
})
})
}
// 保存历史数据
async function savePromptResult (data) {
let url = get_url()
try {
const response = await fetch(`${url}/mixlab/prompt_result`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
action: 'save',
data
})
})
if (response.ok) {
const res = await response.json()
console.log('Response:', res)
return res
// 处理返回的数据
} else {
console.log('Error:', response.status)
// 处理错误情况
}
} catch (error) {
console.log('Error:', error)
// 处理异常情况
}
}
async function uploadImage (blob, fileType = '.png', filename) {
const body = new FormData()
body.append(
'image',
new File([blob], (filename || new Date().getTime()) + fileType)
)
const url = get_url()
const resp = await fetch(`${url}/upload/image`, {
method: 'POST',
body
})
let data = await resp.json()
// console.log(data)
let { name, subfolder } = data
let src = `${url}/view?filename=${encodeURIComponent(
name
)}&type=input&subfolder=${subfolder}&rand=${Math.random()}`
return { url: src, name }
}
async function uploadMask (arrayBuffer, imgurl) {
const body = new FormData()
const filename = 'clipspace-mask-' + performance.now() + '.png'
let original_url = new URL(imgurl)
const original_ref = { filename: original_url.searchParams.get('filename') }
let original_subfolder = original_url.searchParams.get('subfolder')
if (original_subfolder) original_ref.subfolder = original_subfolder
let original_type = original_url.searchParams.get('type')
if (original_type) original_ref.type = original_type
body.append('image', arrayBuffer, filename)
body.append('original_ref', JSON.stringify(original_ref))
body.append('type', 'input')
body.append('subfolder', 'clipspace')
const url = get_url()
const resp = await fetch(`${url}/upload/mask`, {
method: 'POST',
body
})
// console.log(resp)
let data = await resp.json()
let { name, subfolder, type } = data
let src = `${url}/view?filename=${encodeURIComponent(
name
)}&type=${type}&subfolder=${subfolder}&rand=${Math.random()}`
return { url: src, name: 'clipspace/' + name }
}
const parseImageToBase64 = url => {
return new Promise((res, rej) => {
fetch(url)
.then(response => response.blob())
.then(blob => {
const reader = new FileReader()
reader.onloadend = () => {
const base64data = reader.result
res(base64data)
// 在这里可以将base64数据用于进一步处理或显示图片
}
reader.readAsDataURL(blob)
})
.catch(error => {
console.log('发生错误:', error)
})
})
}
function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
function convertImageToBlackBasedOnAlpha (image) {
const canvas = document.createElement('canvas')
const ctx = canvas.getContext('2d')
// Draw the image onto the canvas
canvas.width = image.width
canvas.height = image.height
ctx.drawImage(image, 0, 0)
// Get the image data from the canvas
const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height)
const pixels = imageData.data
// Modify the RGB values based on the alpha channel
for (let i = 0; i < pixels.length; i += 4) {
const alpha = pixels[i + 3]
if (alpha !== 0) {
// Set non-transparent pixels to black
// 蒙版是黑色?
pixels[i] = 0 // Red
pixels[i + 1] = 255 // Green
pixels[i + 2] = 0 // Blue
}
}
// Put the modified image data back onto the canvas
ctx.putImageData(imageData, 0, 0)
// Convert the modified canvas to base64 data URL
const base64ImageData = canvas.toDataURL('image/png') // Replace 'png' with your desired image format
return base64ImageData
}
const blobToBase64 = blob => {
return new Promise((res, rej) => {
const reader = new FileReader()
reader.onloadend = () => {
const base64data = reader.result
res(base64data)
// 在这里可以将base64数据用于进一步处理或显示图片
}
reader.readAsDataURL(blob)
})
}
function base64ToBlob (base64) {
// 去除base64编码中的前缀
const base64WithoutPrefix = base64.replace(/^data:image\/\w+;base64,/, '')
// 将base64编码转换为字节数组
const byteCharacters = atob(base64WithoutPrefix)
// 创建一个存储字节数组的数组
const byteArrays = []
// 将字节数组放入数组中
for (let offset = 0; offset < byteCharacters.length; offset += 1024) {
const slice = byteCharacters.slice(offset, offset + 1024)
const byteNumbers = new Array(slice.length)
for (let i = 0; i < slice.length; i++) {
byteNumbers[i] = slice.charCodeAt(i)
}
const byteArray = new Uint8Array(byteNumbers)
byteArrays.push(byteArray)
}
// 创建blob对象
const blob = new Blob(byteArrays, { type: 'image/png' }) // 根据实际情况设置MIME类型
return blob
}
async function calculateImageHash (blob) {
const buffer = await blob.arrayBuffer()
const hashBuffer = await crypto.subtle.digest('SHA-256', buffer)
const hashArray = Array.from(new Uint8Array(hashBuffer))
const hashHex = hashArray
.map(byte => byte.toString(16).padStart(2, '0'))
.join('')
return hashHex
}
// 获取 rembg 模型
async function get_rembg_models () {
try {
const response = await fetch(`${get_url()}/mixlab/folder_paths`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
type: 'rembg'
})
})
const data = await response.json()
// console.log(data)
return data.names
} catch (error) {
console.error(error)
}
}
//自动抠图
async function run_rembg (model, base64) {
try {
const response = await fetch(`${get_url()}/mixlab/rembg`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
model,
base64
})
})
const data = await response.json()
// console.log(data)
return data.data
} catch (error) {
console.error(error)
}
}
function copyHtmlWithImagesToClipboard (data, cb) {
// 创建一个临时div元素
const tempDiv = document.createElement('div')
// 将HTML字符串赋值给div的innerHTML属性
tempDiv.innerHTML = data
// 获取div中的所有图像元素
const images = tempDiv.getElementsByTagName('img')
// 遍历图像元素,并将图像数据转换为Base64编码
for (let i = 0; i < images.length; i++) {
const image = images[i]
const canvas = document.createElement('canvas')
const context = canvas.getContext('2d')
// 设置canvas尺寸与图像尺寸相同
canvas.width = image.width
canvas.height = image.height
// 在canvas上绘制图像
context.drawImage(image, 0, 0)
// 将canvas转换为Base64编码
const imageData = canvas.toDataURL()
// 将Base64编码替换图像元素的src属性
image.src = imageData
}
let richText = tempDiv.innerHTML
// 创建一个新的Blob对象,并将富文本字符串作为数据传递进去
const blob = new Blob([richText], { type: 'text/html' })
// 创建一个ClipboardItem对象,并将Blob对象添加到其中
const clipboardItem = new ClipboardItem({ 'text/html': blob })
// 使用Clipboard API将内容复制到剪贴板
navigator.clipboard
.write([clipboardItem])
.then(() => {
console.log('富文本已成功复制到剪贴板')
tempDiv.remove()
if (cb) cb(true)
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
tempDiv.remove()
if (cb) cb(false)
})
}
function copyImagesToClipboard (html, cb) {
const tempDiv = document.createElement('div')
tempDiv.innerHTML = html
const images = tempDiv.querySelectorAll('img')
const promises = Array.from(images).map(image => {
return new Promise(resolve => {
const img = new Image()
img.src = image.src
img.onload = () => {
const canvas = document.createElement('canvas')
const context = canvas.getContext('2d')
canvas.width = img.width
canvas.height = img.height
context.drawImage(img, 0, 0)
canvas.toBlob(blob => {
const clipboardItem = new ClipboardItem({ 'image/png': blob })
navigator.clipboard
.write([clipboardItem])
.then(() => {
resolve()
tempDiv.remove()
if (cb) cb(true)
})
.catch(error => {
reject(error)
tempDiv.remove()
if (cb) cb(false)
})
})
}
})
})
Promise.all([...promises])
.then(() => {
console.log('所有图片已成功复制到剪贴板')
if (cb) cb(true)
tempDiv.remove()
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
if (cb) cb(false)
tempDiv.remove()
})
}
function copyTextToClipboard (html, cb) {
const tempDiv = document.createElement('div')
tempDiv.innerHTML = html
const text = tempDiv.innerText
const textData = new ClipboardItem({
'text/plain': new Blob([text], { type: 'text/plain' })
})
navigator.clipboard
.write([textData])
.then(() => {
console.log('所有文本已成功复制到剪贴板', text)
if (cb) cb(true)
tempDiv.remove()
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
if (cb) cb(false)
tempDiv.remove()
})
}
// ComfyUI\web\extensions\core\dynamicPrompts.js
// 官方实现修改
// Allows for simple dynamic prompt replacement
// Inputs in the format {a|b} will have a random value of a or b chosen when the prompt is queued.
/*
* Strips C-style line and block comments from a string
*/
function dynamicPrompts (prompt) {
prompt = prompt.replace(/\/\*[\s\S]*?\*\/|\/\/.*/g, '')
while (
prompt.replace('\\{', '').includes('{') &&
prompt.replace('\\}', '').includes('}')
) {
const startIndex = prompt.replace('\\{', '00').indexOf('{')
const endIndex = prompt.replace('\\}', '00').indexOf('}')
const optionsString = prompt.substring(startIndex + 1, endIndex)
const options = optionsString.split('|')
const randomIndex = Math.floor(Math.random() * options.length)
const randomOption = options[randomIndex]
prompt =
prompt.substring(0, startIndex) +
randomOption +
prompt.substring(endIndex + 1)
}
return prompt
}
// 遍历所有组合,语法同 动态提示
function generateAllCombinations (prompt) {
prompt = prompt.replace(/\/\*[\s\S]*?\*\/|\/\/.*/g, '')
// Helper function to get all combinations
function getAllCombinations (parts) {
if (parts.length === 0) return ['']
const [firstPart, ...restParts] = parts
const restCombinations = getAllCombinations(restParts)
const allCombinations = []
firstPart.forEach(option => {
restCombinations.forEach(combination => {
allCombinations.push(option + combination)
})
})
return allCombinations
}
// Split prompt into static parts and dynamic parts
let parts = []
let startIndex = 0
while (
prompt.replace('\\{', '').includes('{') &&
prompt.replace('\\}', '').includes('}')
) {
startIndex = prompt.replace('\\{', '00').indexOf('{')
const endIndex = prompt.replace('\\}', '00').indexOf('}')
const staticPart = prompt.substring(0, startIndex)
const optionsString = prompt.substring(startIndex + 1, endIndex)
const options = optionsString.split('|')
parts.push([staticPart])
parts.push(options)
prompt = prompt.substring(endIndex + 1)
}
// Add the remaining static part
parts.push([prompt])
// Get all combinations
const combinations = getAllCombinations(parts)
return combinations
}
const _textNodes = [
'TextInput_',
'CLIPTextEncode',
'PromptSimplification',
'ChinesePrompt_Mix'
],
_loraNodes = ['CheckpointLoaderSimple', 'LoraLoader'],
_numberNodes = ['FloatSlider', 'IntNumber'],
_slideNodes = ['PromptSlide'],
_imageNodes = [
'LoadImage',
'VHS_LoadVideo',
'ImagesPrompt_',
'LoadImagesToBatch'
],
_colorNodes = ['Color'],
_audioNodes = ['LoadAndCombinedAudio_']
export default {
get_url,
get_my_app,
getAppInit,
getFilenameAndCategoryFromUrl,
success,
interrupt,
getQueue,
queuePromptNew,
savePromptResult,
uploadImage,
uploadMask,
run_rembg,
get_rembg_models,
parseImageToBase64,
createImage,
convertImageToBlackBasedOnAlpha,
blobToBase64,
base64ToBlob,
calculateImageHash,
copyHtmlWithImagesToClipboard,
copyImagesToClipboard,
copyTextToClipboard,
dynamicPrompts,
generateAllCombinations,
_textNodes,
_loraNodes,
_numberNodes,
_slideNodes,
_imageNodes,
_colorNodes,
_audioNodes
}
+347
View File
@@ -0,0 +1,347 @@
/* juxtapose - v1.2.2 - 2020-09-03
* Copyright (c) 2020 Alex Duner and Northwestern University Knight Lab
*/
div.juxtapose {
width: 100%;
font-family: Helvetica, Arial, sans-serif;
}
div.jx-slider {
width: 100%;
height: 100%;
position: relative;
overflow: hidden;
cursor: pointer;
color: #f3f3f3;
}
div.jx-handle {
position: absolute;
height: 100%;
width: 40px;
cursor: col-resize;
z-index: 15;
margin-left: -20px;
}
.vertical div.jx-handle {
height: 40px;
width: 100%;
cursor: row-resize;
margin-top: -20px;
margin-left: 0;
}
div.jx-control {
height: 100%;
margin-right: auto;
margin-left: auto;
width: 3px;
background-color: currentColor;
}
.vertical div.jx-control {
height: 3px;
width: 100%;
background-color: currentColor;
position: relative;
top: 50%;
transform: translateY(-50%);
}
div.jx-controller {
position: absolute;
margin: auto;
top: 0;
bottom: 0;
height: 60px;
width: 9px;
margin-left: -3px;
background-color: currentColor;
}
.vertical div.jx-controller {
height: 9px;
width: 100px;
margin-left: auto;
margin-right: auto;
top: -3px;
position: relative;
}
div.jx-arrow {
position: absolute;
margin: auto;
top: 0;
bottom: 0;
width: 0;
height: 0;
transition: all .2s ease;
}
.vertical div.jx-arrow {
position: absolute;
margin: 0 auto;
left: 0;
right: 0;
width: 0;
height: 0;
transition: all .2s ease;
}
div.jx-arrow.jx-left {
left: 2px;
border-style: solid;
border-width: 8px 8px 8px 0;
border-color: transparent currentColor transparent transparent;
}
div.jx-arrow.jx-right {
right: 2px;
border-style: solid;
border-width: 8px 0 8px 8px;
border-color: transparent transparent transparent currentColor;
}
.vertical div.jx-arrow.jx-left {
left: 0px;
top: 2px;
border-style: solid;
border-width: 0px 8px 8px 8px;
border-color: transparent transparent currentColor transparent;
}
.vertical div.jx-arrow.jx-right {
right: 0px;
top: auto;
bottom: 2px;
border-style: solid;
border-width: 8px 8px 0 8px;
border-color: currentColor transparent transparent transparent;
}
div.jx-handle:hover div.jx-arrow.jx-left,
div.jx-handle:active div.jx-arrow.jx-left {
left: -1px;
}
div.jx-handle:hover div.jx-arrow.jx-right,
div.jx-handle:active div.jx-arrow.jx-right {
right: -1px;
}
.vertical div.jx-handle:hover div.jx-arrow.jx-left,
.vertical div.jx-handle:active div.jx-arrow.jx-left {
left: 0px;
top: 0px;
}
.vertical div.jx-handle:hover div.jx-arrow.jx-right,
.vertical div.jx-handle:active div.jx-arrow.jx-right {
right: 0px;
bottom: 0px;
}
div.jx-image {
position: absolute;
height: 100%;
display: inline-block;
top: 0;
overflow: hidden;
-webkit-backface-visibility: hidden;
}
.vertical div.jx-image {
width: 100%;
left: 0;
top: auto;
}
div.jx-image img {
height: 100%;
width: auto;
z-index: 5;
position: absolute;
margin-bottom: 0;
max-height: none;
max-width: none;
max-height: initial;
max-width: initial;
}
.vertical div.jx-image img {
height: auto;
width: 100%;
}
div.jx-image.jx-left {
left: 0;
background-position: left;
}
div.jx-image.jx-left img {
left: 0;
}
div.jx-image.jx-right {
right: 0;
background-position: right;
}
div.jx-image.jx-right img {
right: 0;
bottom: 0;
}
.veritcal div.jx-image.jx-left {
top: 0;
background-position: top;
}
.veritcal div.jx-image.jx-left img {
top: 0;
}
.vertical div.jx-image.jx-right {
bottom: 0;
background-position: bottom;
}
.veritcal div.jx-image.jx-right img {
bottom: 0;
}
div.jx-image div.jx-label {
font-size: 1em;
padding: .25em .75em;
position: relative;
display: inline-block;
top: 0;
background-color: #000; /* IE 8 */
background-color: rgba(0,0,0,.7);
color: white;
z-index: 10;
white-space: nowrap;
line-height: 18px;
vertical-align: middle;
}
div.jx-image.jx-left div.jx-label {
float: left;
left: 0;
}
div.jx-image.jx-right div.jx-label {
float: right;
right: 0;
}
.vertical div.jx-image div.jx-label {
display: table;
position: absolute;
}
.vertical div.jx-image.jx-right div.jx-label {
left: 0;
bottom: 0;
top: auto;
}
div.jx-credit {
line-height: 1.1;
font-size: 0.75em;
}
div.jx-credit em {
font-weight: bold;
font-style: normal;
}
/* Animation */
div.jx-image.transition {
transition: width .5s ease;
}
div.jx-handle.transition {
transition: left .5s ease;
}
.vertical div.jx-image.transition {
transition: height .5s ease;
}
.vertical div.jx-handle.transition {
transition: top .5s ease;
}
/* Knight Lab Credit */
a.jx-knightlab {
background-color: #000; /* IE 8 */
background-color: rgba(0,0,0,.25);
bottom: 0;
display: table;
height: 14px;
line-height: 14px;
padding: 1px 4px 1px 5px;
position: absolute;
right: 0;
text-decoration: none;
z-index: 10;
}
a.jx-knightlab div.knightlab-logo {
display: inline-block;
vertical-align: middle;
height: 8px;
width: 8px;
background-color: #c34528;
transform: rotate(45deg);
-ms-transform: rotate(45deg);
-webkit-transform: rotate(45deg);
top: -1.25px;
position: relative;
cursor: pointer;
}
a.jx-knightlab:hover {
background-color: #000; /* IE 8 */
background-color: rgba(0,0,0,.35);
}
a.jx-knightlab:hover div.knightlab-logo {
background-color: #ce4d28;
}
a.jx-knightlab span.juxtapose-name {
display: table-cell;
margin: 0;
padding: 0;
font-family: Helvetica, Arial, sans-serif;
font-weight: 300;
color: white;
font-size: 10px;
padding-left: 0.375em;
vertical-align: middle;
line-height: normal;
text-shadow: none;
}
/* keyboard accessibility */
div.jx-controller:focus,
div.jx-image.jx-left div.jx-label:focus,
div.jx-image.jx-right div.jx-label:focus,
a.jx-knightlab:focus {
background: #eae34a;
color: #000;
}
a.jx-knightlab:focus span.juxtapose-name{
color: #000;
border: none;
}
+8
View File
File diff suppressed because one or more lines are too long
+1094
View File
File diff suppressed because one or more lines are too long
View File
File diff suppressed because it is too large Load Diff
+144
View File
@@ -0,0 +1,144 @@
/**
* https://github.com/google/model-viewer/blob/master/packages/model-viewer/src/three-components/EnvironmentScene.ts
*/
import {
BackSide,
BoxGeometry,
Mesh,
MeshBasicMaterial,
MeshStandardMaterial,
PointLight,
Scene,
} from './three.module.js';
class RoomEnvironment extends Scene {
constructor( renderer = null ) {
super();
const geometry = new BoxGeometry();
geometry.deleteAttribute( 'uv' );
const roomMaterial = new MeshStandardMaterial( { side: BackSide } );
const boxMaterial = new MeshStandardMaterial();
const mainLight = new PointLight( 0xffffff, 900, 28, 2 );
mainLight.position.set( 0.418, 16.199, 0.300 );
this.add( mainLight );
const room = new Mesh( geometry, roomMaterial );
room.position.set( - 0.757, 13.219, 0.717 );
room.scale.set( 31.713, 28.305, 28.591 );
this.add( room );
const box1 = new Mesh( geometry, boxMaterial );
box1.position.set( - 10.906, 2.009, 1.846 );
box1.rotation.set( 0, - 0.195, 0 );
box1.scale.set( 2.328, 7.905, 4.651 );
this.add( box1 );
const box2 = new Mesh( geometry, boxMaterial );
box2.position.set( - 5.607, - 0.754, - 0.758 );
box2.rotation.set( 0, 0.994, 0 );
box2.scale.set( 1.970, 1.534, 3.955 );
this.add( box2 );
const box3 = new Mesh( geometry, boxMaterial );
box3.position.set( 6.167, 0.857, 7.803 );
box3.rotation.set( 0, 0.561, 0 );
box3.scale.set( 3.927, 6.285, 3.687 );
this.add( box3 );
const box4 = new Mesh( geometry, boxMaterial );
box4.position.set( - 2.017, 0.018, 6.124 );
box4.rotation.set( 0, 0.333, 0 );
box4.scale.set( 2.002, 4.566, 2.064 );
this.add( box4 );
const box5 = new Mesh( geometry, boxMaterial );
box5.position.set( 2.291, - 0.756, - 2.621 );
box5.rotation.set( 0, - 0.286, 0 );
box5.scale.set( 1.546, 1.552, 1.496 );
this.add( box5 );
const box6 = new Mesh( geometry, boxMaterial );
box6.position.set( - 2.193, - 0.369, - 5.547 );
box6.rotation.set( 0, 0.516, 0 );
box6.scale.set( 3.875, 3.487, 2.986 );
this.add( box6 );
// -x right
const light1 = new Mesh( geometry, createAreaLightMaterial( 50 ) );
light1.position.set( - 16.116, 14.37, 8.208 );
light1.scale.set( 0.1, 2.428, 2.739 );
this.add( light1 );
// -x left
const light2 = new Mesh( geometry, createAreaLightMaterial( 50 ) );
light2.position.set( - 16.109, 18.021, - 8.207 );
light2.scale.set( 0.1, 2.425, 2.751 );
this.add( light2 );
// +x
const light3 = new Mesh( geometry, createAreaLightMaterial( 17 ) );
light3.position.set( 14.904, 12.198, - 1.832 );
light3.scale.set( 0.15, 4.265, 6.331 );
this.add( light3 );
// +z
const light4 = new Mesh( geometry, createAreaLightMaterial( 43 ) );
light4.position.set( - 0.462, 8.89, 14.520 );
light4.scale.set( 4.38, 5.441, 0.088 );
this.add( light4 );
// -z
const light5 = new Mesh( geometry, createAreaLightMaterial( 20 ) );
light5.position.set( 3.235, 11.486, - 12.541 );
light5.scale.set( 2.5, 2.0, 0.1 );
this.add( light5 );
// +y
const light6 = new Mesh( geometry, createAreaLightMaterial( 100 ) );
light6.position.set( 0.0, 20.0, 0.0 );
light6.scale.set( 1.0, 0.1, 1.0 );
this.add( light6 );
}
dispose() {
const resources = new Set();
this.traverse( ( object ) => {
if ( object.isMesh ) {
resources.add( object.geometry );
resources.add( object.material );
}
} );
for ( const resource of resources ) {
resource.dispose();
}
}
}
function createAreaLightMaterial( intensity ) {
const material = new MeshBasicMaterial();
material.color.setScalar( intensity );
return material;
}
export { RoomEnvironment };
File diff suppressed because one or more lines are too long
+316
View File
@@ -0,0 +1,316 @@
import * as THREE from './three/three.module.js'
import { api } from '../../../scripts/api.js'
import { OrbitControls } from './three/OrbitControls.js'
import { RoomEnvironment } from './three/RoomEnvironment.js'
const visualizer = document.getElementById('visualizer')
const container = document.getElementById('container')
const progressDialog = document.getElementById('progress-dialog')
const progressIndicator = document.getElementById('progress-indicator')
const renderer = new THREE.WebGLRenderer({
antialias: true,
extensions: {
derivatives: true
}
})
renderer.setPixelRatio(window.devicePixelRatio)
renderer.setSize(window.innerWidth, window.innerHeight)
if (container) container.appendChild(renderer.domElement)
const pmremGenerator = new THREE.PMREMGenerator(renderer)
// scene
const scene = new THREE.Scene()
scene.background = new THREE.Color(0x000000)
scene.environment = pmremGenerator.fromScene(
new RoomEnvironment(renderer),
0.04
).texture
const ambientLight = new THREE.AmbientLight(0xffffff)
const camera = new THREE.PerspectiveCamera(
40,
window.innerWidth / window.innerHeight,
0.1,
1000
)
camera.position.set(0, 0, 10)
const pointLight = new THREE.PointLight(0xffffff, 15)
camera.add(pointLight)
const controls = new OrbitControls(camera, renderer.domElement)
controls.target.set(0, 0, 0)
controls.update()
controls.enablePan = true
controls.enableDamping = true
// Handle window resize event
window.onresize = function () {
camera.aspect = window.innerWidth / window.innerHeight
camera.updateProjectionMatrix()
renderer.setSize(window.innerWidth, window.innerHeight)
}
var lastReferenceImage = ''
var lastDepthMap = ''
var needUpdate = false
function frameUpdate () {
var referenceImage = visualizer?.getAttribute('reference_image')
var depthMap = visualizer?.getAttribute('depth_map')
if (referenceImage == lastReferenceImage && depthMap == lastDepthMap) {
if (needUpdate) {
controls.update()
renderer.render(scene, camera)
}
requestAnimationFrame(frameUpdate)
} else {
needUpdate = false
scene.clear()
if (progressDialog) progressDialog.open = true
lastReferenceImage = referenceImage
lastDepthMap = depthMap
if (lastReferenceImage && lastReferenceImage != 'undefined') {
// console.log('lastReferenceImage',typeof(lastReferenceImage),lastDepthMap)
main(JSON.parse(lastReferenceImage), JSON.parse(lastDepthMap))
}
}
}
const onProgress = function (xhr) {
if (xhr.lengthComputable) {
progressIndicator.value = (xhr.loaded / xhr.total) * 100
}
}
const onError = function (e) {
console.error(e)
}
async function main (referenceImageParams, depthMapParams) {
let referenceTexture, depthTexture
let imageWidth = 10 // Default width
let imageHeight = 10 // Default height, will be updated based on the image's aspect ratio
// console.log('#referenceImageParams', referenceImageParams)
if (referenceImageParams?.filename) {
const referenceImageUrl = api
.apiURL('/view?' + new URLSearchParams(referenceImageParams))
.replace(/extensions.*\//, '')
const referenceImageExt = referenceImageParams.filename.slice(
referenceImageParams.filename.lastIndexOf('.') + 1
)
if (
referenceImageExt === 'png' ||
referenceImageExt === 'jpg' ||
referenceImageExt === 'jpeg'
) {
const referenceImageLoader = new THREE.TextureLoader()
referenceTexture = await new Promise((resolve, reject) => {
referenceImageLoader.load(
referenceImageUrl,
texture => {
// Once the image is loaded, update the width and height based on the image's aspect ratio
imageWidth = 10 // Keep the width as 10
imageHeight = texture.image.height / (texture.image.width / 10)
resolve(texture)
},
undefined,
reject
)
})
}
}
if (depthMapParams?.filename) {
const depthMapUrl = api
.apiURL('/view?' + new URLSearchParams(depthMapParams))
.replace(/extensions.*\//, '')
const depthMapExt = depthMapParams.filename.slice(
depthMapParams.filename.lastIndexOf('.') + 1
)
if (
depthMapExt === 'png' ||
depthMapExt === 'jpg' ||
depthMapExt === 'jpeg'
) {
const depthMapLoader = new THREE.TextureLoader()
depthTexture = await depthMapLoader.loadAsync(depthMapUrl)
}
}
if (referenceTexture && depthTexture) {
const depthMaterial = new THREE.ShaderMaterial({
uniforms: {
referenceTexture: { value: referenceTexture },
depthTexture: { value: depthTexture },
depthScale: { value: 5.0 },
ambientLightColor: { value: new THREE.Color(0.2, 0.2, 0.2) },
lightPosition: { value: new THREE.Vector3(2, 2, 2) },
lightColor: { value: new THREE.Color(1, 1, 1) },
lightIntensity: { value: 1.0 },
shininess: { value: 30 }
},
vertexShader: `
uniform sampler2D depthTexture;
uniform float depthScale;
varying vec2 vUv;
varying float vDepth;
varying vec3 vNormal;
varying vec3 vViewPosition;
void main() {
vUv = uv;
float depth = texture2D(depthTexture, uv).r;
vec3 displacement = normal * depth * depthScale;
vec3 displacedPosition = position + displacement;
vec4 worldPosition = modelMatrix * vec4(displacedPosition, 1.0);
vNormal = normalize(normalMatrix * normal);
vViewPosition = (viewMatrix * worldPosition).xyz;
gl_Position = projectionMatrix * viewMatrix * worldPosition;
vDepth = depth;
}
`,
fragmentShader: `
uniform sampler2D referenceTexture;
varying vec2 vUv;
varying float vDepth;
void main() {
vec4 referenceColor = texture2D(referenceTexture, vUv);
// Directly use reference color without fog
gl_FragColor = referenceColor;
}
`
})
const planeGeometry = new THREE.PlaneGeometry(
imageWidth,
imageHeight,
200,
200
)
const depthMesh = new THREE.Mesh(planeGeometry, depthMaterial)
scene.add(depthMesh)
}
needUpdate = true
scene.add(ambientLight)
scene.add(camera)
progressDialog?.close()
frameUpdate()
}
document
.getElementById('screenshotButton')
?.addEventListener('click', takeScreenshot)
const sleep = (t = 1000) => {
return new Promise((res, rej) => {
setTimeout(() => res(1), t)
})
}
// 方法:旋转摄像机并拍摄图片 // 每次旋转的角度增量,转换为弧度
async function captureImages (
totalFrames = 40,
angleIncrement = THREE.MathUtils.degToRad(0.5)
) {
// 计算场景中所有物体的中心点
const box = new THREE.Box3().setFromObject(scene)
const center = new THREE.Vector3()
box.getCenter(center)
// 计算当前相机距离中心点的半径
const radius = camera.position.distanceTo(center)
// 存储图片的数组
let images = []
// 记录初始相机位置和朝向
const initialPosition = camera.position.clone()
const initialTarget = center.clone()
// 计算当前相机的初始角度
const initialAngle = Math.atan2(
camera.position.z - center.z,
camera.position.x - center.x
)
// 起始角度为从当前角度往左旋转 20 度的位置
const startAngle = initialAngle - (angleIncrement * totalFrames) / 2
for (let i = 0; i < totalFrames; i++) {
const angle = startAngle + i * angleIncrement
// 计算相机的位置
camera.position.x = center.x + radius * Math.cos(angle)
camera.position.z = center.z + radius * Math.sin(angle)
camera.position.y = initialPosition.y // 保持相机高度不变
camera.lookAt(center) // 相机看向中心点
// 渲染当前帧
renderer.render(scene, camera)
// 将当前帧保存为图片
const imgData = renderer.domElement.toDataURL('image/png')
images.push(imgData)
// 等待一段时间
await new Promise(resolve => setTimeout(resolve, 500))
}
// 恢复相机到初始位置和朝向
camera.position.copy(initialPosition)
camera.lookAt(initialTarget)
return images
}
async function takeScreenshot () {
// 更新相机的矩阵,以确保其世界矩阵是最新的
camera.updateMatrixWorld()
const imgs = await captureImages()
// 获取当前网页的 URL
const currentUrl = window.location.href
// 创建一个 URL 对象
const url = new URL(currentUrl)
// 使用 URLSearchParams 获取参数
const params = new URLSearchParams(url.search)
// 获取参数 'id' 的值
const id = params.get('id')
window.parent.postMessage({ imgs, id }, '*')
}
main()
window.addEventListener('message', event => {
// 这里可以添加对来源的验证,以确保安全
// console.log('Message received from parent page:', event.data)
let { reference_image, depth_map } = event.data
if (reference_image && depth_map) {
visualizer?.setAttribute('reference_image', JSON.stringify(reference_image))
visualizer?.setAttribute('depth_map', JSON.stringify(depth_map))
frameUpdate()
}
})
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,6 @@
<!DOCTYPE html>
<meta charset="utf-8">
<!-- <link href="https://fonts.googleapis.com/css?family=Montserrat" rel="stylesheet"> -->
<title>p5.js-widget</title>
<div id="app-holder"></div>
<script src="./main.bundle.js"></script>

Some files were not shown because too many files have changed in this diff Show More